feat: Support for Autodetect in AstraDBVectorStore settings (#4869)

* feat: first pass at autodetect updates

* [autofix.ci] apply automated fixes

* Fully support autodetect

---------

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Eric Hare 2024-12-02 08:37:14 -08:00 • committed by GitHub
commit 19d2974904
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 280 additions and 93 deletions

View file

@ -1,5 +1,5 @@
from .astradb import AstraVectorStoreComponent from .astradb import AstraDBVectorStoreComponent
from .astradb_graph import AstraGraphVectorStoreComponent from .astradb_graph import AstraDBGraphVectorStoreComponent
from .cassandra import CassandraVectorStoreComponent from .cassandra import CassandraVectorStoreComponent
from .cassandra_graph import CassandraGraphVectorStoreComponent from .cassandra_graph import CassandraGraphVectorStoreComponent
from .chroma import ChromaVectorStoreComponent from .chroma import ChromaVectorStoreComponent
@ -23,8 +23,8 @@ from .vectara_self_query import VectaraSelfQueryRetriverComponent
from .weaviate import WeaviateVectorStoreComponent from .weaviate import WeaviateVectorStoreComponent
__all__ = [ __all__ = [
"AstraGraphVectorStoreComponent", "AstraDBGraphVectorStoreComponent",
"AstraVectorStoreComponent", "AstraDBVectorStoreComponent",
"CassandraGraphVectorStoreComponent", "CassandraGraphVectorStoreComponent",
"CassandraVectorStoreComponent", "CassandraVectorStoreComponent",
"ChromaVectorStoreComponent", "ChromaVectorStoreComponent",

View file

@ -4,6 +4,7 @@ from collections import defaultdict
import orjson import orjson
from astrapy import DataAPIClient from astrapy import DataAPIClient
from astrapy.admin import parse_api_endpoint from astrapy.admin import parse_api_endpoint
from astrapy.exceptions import CollectionNotFoundException
from langchain_astradb import AstraDBVectorStore from langchain_astradb import AstraDBVectorStore
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
@ -22,7 +23,7 @@ from langflow.io import (
from langflow.schema import Data from langflow.schema import Data
class AstraVectorStoreComponent(LCVectorStoreComponent): class AstraDBVectorStoreComponent(LCVectorStoreComponent):
display_name: str = "Astra DB" display_name: str = "Astra DB"
description: str = "Implementation of Vector Store using Astra DB with search capabilities" description: str = "Implementation of Vector Store using Astra DB with search capabilities"
documentation: str = "https://docs.langflow.org/starter-projects-vector-store-rag" documentation: str = "https://docs.langflow.org/starter-projects-vector-store-rag"
@ -31,6 +32,24 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
_cached_vector_store: AstraDBVectorStore | None = None _cached_vector_store: AstraDBVectorStore | None = None
def list_collections(self):
client = DataAPIClient(token=self.token)
database = client.get_database(
self.api_endpoint,
token=self.token,
)
return database.list_collections()
def _initialize_collection_options(self):
try:
collections = [collection.name for collection in self.list_collections()]
except (CollectionNotFoundException, ConnectionError, ValueError) as _:
collections = []
return [*collections, "+ Create new collection"]
VECTORIZE_PROVIDERS_MAPPING = defaultdict( VECTORIZE_PROVIDERS_MAPPING = defaultdict(
list, list,
{ {
@ -61,7 +80,7 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
], ],
], ],
"Mistral AI": ["mistral", ["mistral-embed"]], "Mistral AI": ["mistral", ["mistral-embed"]],
"NVIDIA": ["nvidia", ["NV-Embed-QA"]], "Nvidia": ["nvidia", ["NV-Embed-QA"]],
"OpenAI": ["openai", ["text-embedding-3-small", "text-embedding-3-large", "text-embedding-ada-002"]], "OpenAI": ["openai", ["text-embedding-3-small", "text-embedding-3-large", "text-embedding-ada-002"]],
"Upstage": ["upstageAI", ["solar-embedding-1-large"]], "Upstage": ["upstageAI", ["solar-embedding-1-large"]],
"Voyage AI": [ "Voyage AI": [
@ -87,20 +106,14 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
value="ASTRA_DB_API_ENDPOINT", value="ASTRA_DB_API_ENDPOINT",
required=True, required=True,
), ),
StrInput( DropdownInput(
name="collection_name", name="collection_name",
display_name="Collection Name", display_name="Collection",
info="The name of the collection within Astra DB where the vectors will be stored.", info="The name of the collection within Astra DB where the vectors will be stored.",
required=True, required=True,
), real_time_refresh=True,
MultilineInput( refresh_button=True,
name="search_input", options=[],
display_name="Search Input",
),
DataInput(
name="ingest_data",
display_name="Ingest Data",
is_list=True,
), ),
StrInput( StrInput(
name="keyspace", name="keyspace",
@ -108,6 +121,50 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
info="Optional keyspace within Astra DB to use for the collection.", info="Optional keyspace within Astra DB to use for the collection.",
advanced=True, advanced=True,
), ),
MultilineInput(
name="search_input",
display_name="Search Input",
),
IntInput(
name="number_of_results",
display_name="Number of Results",
info="Number of results to return.",
advanced=True,
value=4,
),
DropdownInput(
name="search_type",
display_name="Search Type",
info="Search type to use",
options=["Similarity", "Similarity with score threshold", "MMR (Max Marginal Relevance)"],
value="Similarity",
advanced=True,
),
FloatInput(
name="search_score_threshold",
display_name="Search Score Threshold",
info="Minimum similarity score threshold for search results. "
"(when using 'Similarity with score threshold')",
value=0,
advanced=True,
),
NestedDictInput(
name="advanced_search_filter",
display_name="Search Metadata Filter",
info="Optional dictionary of filters to apply to the search query.",
advanced=True,
),
DictInput(
name="search_filter",
display_name="[DEPRECATED] Search Metadata Filter",
info="Deprecated: use advanced_search_filter. Optional dictionary of filters to apply to the search query.",
advanced=True,
list=True,
),
DataInput(
name="ingest_data",
display_name="Ingest Data",
),
DropdownInput( DropdownInput(
name="embedding_choice", name="embedding_choice",
display_name="Embedding Model or Astra Vectorize", display_name="Embedding Model or Astra Vectorize",
@ -172,14 +229,14 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
name="metadata_indexing_include", name="metadata_indexing_include",
display_name="Metadata Indexing Include", display_name="Metadata Indexing Include",
info="Optional list of metadata fields to include in the indexing.", info="Optional list of metadata fields to include in the indexing.",
is_list=True, list=True,
advanced=True, advanced=True,
), ),
StrInput( StrInput(
name="metadata_indexing_exclude", name="metadata_indexing_exclude",
display_name="Metadata Indexing Exclude", display_name="Metadata Indexing Exclude",
info="Optional list of metadata fields to exclude from the indexing.", info="Optional list of metadata fields to exclude from the indexing.",
is_list=True, list=True,
advanced=True, advanced=True,
), ),
StrInput( StrInput(
@ -189,42 +246,6 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
"See https://docs.datastax.com/en/astra-db-serverless/api-reference/collections.html#the-indexing-option", "See https://docs.datastax.com/en/astra-db-serverless/api-reference/collections.html#the-indexing-option",
advanced=True, advanced=True,
), ),
IntInput(
name="number_of_results",
display_name="Number of Results",
info="Number of results to return.",
advanced=True,
value=4,
),
DropdownInput(
name="search_type",
display_name="Search Type",
info="Search type to use",
options=["Similarity", "Similarity with score threshold", "MMR (Max Marginal Relevance)"],
value="Similarity",
advanced=True,
),
FloatInput(
name="search_score_threshold",
display_name="Search Score Threshold",
info="Minimum similarity score threshold for search results. "
"(when using 'Similarity with score threshold')",
value=0,
advanced=True,
),
NestedDictInput(
name="advanced_search_filter",
display_name="Search Metadata Filter",
info="Optional dictionary of filters to apply to the search query.",
advanced=True,
),
DictInput(
name="search_filter",
display_name="[DEPRECATED] Search Metadata Filter",
info="Deprecated: use advanced_search_filter. Optional dictionary of filters to apply to the search query.",
advanced=True,
is_list=True,
),
] ]
def del_fields(self, build_config, field_list): def del_fields(self, build_config, field_list):
@ -289,8 +310,110 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
return self.VECTORIZE_PROVIDERS_MAPPING return self.VECTORIZE_PROVIDERS_MAPPING
def get_collection_options(self):
client = DataAPIClient(token=self.token)
database = client.get_database(
self.api_endpoint,
token=self.token,
)
collection = database.get_collection(self.collection_name)
# Only get the options if the collection exists
try:
collection_options = collection.options()
except CollectionNotFoundException as e:
self.log(f"Collection not found: {e}")
return None
return collection_options.vector
def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None): def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None):
if field_name == "embedding_choice": # Refresh the collection name options
build_config["collection_name"]["options"] = self._initialize_collection_options()
# If the collection name is set to "+ Create new collection", show the advanced options
if field_name == "collection_name" and field_value == "+ Create new collection":
build_config["embedding_choice"]["advanced"] = False
build_config["embedding_choice"]["value"] = "Embedding Model"
new_parameter = StrInput(
name="collection_name_new",
display_name="Collection Name",
required=True,
).to_dict()
self.insert_in_dict(build_config, "embedding_choice", {"collection_name_new": new_parameter})
new_parameter = HandleInput(
name="embedding_model",
display_name="Embedding Model",
input_types=["Embeddings"],
info="Allows an embedding model configuration.",
).to_dict()
self.insert_in_dict(build_config, "collection_name_new", {"embedding_model": new_parameter})
elif field_name == "collection_name" and field_value != "+ Create new collection":
self.del_fields(build_config, ["collection_name_new"])
# Get the collection options
collection_options = self.get_collection_options()
# If the collection options are available, show the advanced options
if collection_options:
build_config["embedding_choice"]["advanced"] = True
if collection_options.service:
for input_field in [
"embedding_provider",
"z_01_model_parameters",
"z_02_api_key_name",
"z_03_provider_api_key",
"z_04_authentication",
]:
build_config[input_field]["advanced"] = False
build_config["embedding_model"]["advanced"] = True
build_config["embedding_provider"]["advanced"] = True
build_config["embedding_choice"]["value"] = "Astra Vectorize"
build_config["embedding_provider"]["value"] = collection_options.service.provider
build_config["model"]["value"] = collection_options.service.model_name
build_config["z_01_model_parameters"]["value"] = collection_options.service.parameters
if collection_options.service.authentication:
build_config["z_02_api_key_name"]["value"] = collection_options.service.authentication.get(
"providerKey"
)
build_config["z_03_provider_api_key"]["value"] = collection_options.service.authentication.get(
"apiKey"
)
build_config["z_04_authentication"]["value"] = collection_options.service.authentication
else:
for input_field in [
"z_01_model_parameters",
"z_02_api_key_name",
"z_03_provider_api_key",
"z_04_authentication",
]:
build_config[input_field]["advanced"] = True
build_config["embedding_model"]["advanced"] = False
build_config["embedding_provider"]["advanced"] = False
build_config["embedding_choice"]["value"] = "Embedding Model"
new_parameter = HandleInput(
name="embedding_model",
display_name="Embedding Model",
input_types=["Embeddings"],
info="Allows an embedding model configuration.",
).to_dict()
self.insert_in_dict(build_config, "embedding_choice", {"embedding_model": new_parameter})
elif field_name == "embedding_choice":
if field_value == "Astra Vectorize": if field_value == "Astra Vectorize":
self.del_fields(build_config, ["embedding_model"]) self.del_fields(build_config, ["embedding_model"])
@ -363,7 +486,7 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
new_parameter_1 = DictInput( new_parameter_1 = DictInput(
name="z_01_model_parameters", name="z_01_model_parameters",
display_name="Model Parameters", display_name="Model Parameters",
is_list=True, list=True,
).to_dict() ).to_dict()
new_parameter_2 = MessageTextInput( new_parameter_2 = MessageTextInput(
@ -386,7 +509,7 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
new_parameter_4 = DictInput( new_parameter_4 = DictInput(
name="z_04_authentication", name="z_04_authentication",
display_name="Authentication Parameters", display_name="Authentication Parameters",
is_list=True, list=True,
).to_dict() ).to_dict()
self.insert_in_dict( self.insert_in_dict(
@ -415,11 +538,10 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
setattr(self, attribute, None) setattr(self, attribute, None)
# Fetch values from kwargs if any self.* attributes are None # Fetch values from kwargs if any self.* attributes are None
provider_value = self.VECTORIZE_PROVIDERS_MAPPING.get(self.embedding_provider, [None])[0] or kwargs.get( provider_mapping = self.update_providers_mapping()
"embedding_provider" provider_value = provider_mapping.get(self.embedding_provider, [None])[0] or kwargs.get("embedding_provider")
)
model_name = self.model or kwargs.get("model") model_name = self.model or kwargs.get("model")
authentication = {**(self.z_04_authentication or kwargs.get("z_04_authentication", {}))} authentication = {**(self.z_04_authentication or {}), **kwargs.get("z_04_authentication", {})}
parameters = self.z_01_model_parameters or kwargs.get("z_01_model_parameters", {}) parameters = self.z_01_model_parameters or kwargs.get("z_01_model_parameters", {})
# Set the API key name if provided # Set the API key name if provided
@ -427,6 +549,9 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
provider_key = self.z_03_provider_api_key or kwargs.get("z_03_provider_api_key") provider_key = self.z_03_provider_api_key or kwargs.get("z_03_provider_api_key")
if api_key_name: if api_key_name:
authentication["providerKey"] = api_key_name authentication["providerKey"] = api_key_name
if authentication:
provider_key = None
authentication["providerKey"] = authentication["providerKey"].split(".")[0]
# Set authentication and parameters to None if no values are provided # Set authentication and parameters to None if no values are provided
if not authentication: if not authentication:
@ -466,13 +591,64 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
msg = f"Invalid setup mode: {self.setup_mode}" msg = f"Invalid setup mode: {self.setup_mode}"
raise ValueError(msg) from e raise ValueError(msg) from e
metric_value = self.metric or None
autodetect = False
if self.embedding_choice == "Embedding Model": if self.embedding_choice == "Embedding Model":
embedding_dict = {"embedding": self.embedding_model} embedding_dict = {"embedding": self.embedding_model}
# Use autodetect if the collection name is NOT set to "+ Create new collection"
elif self.collection_name != "+ Create new collection":
autodetect = True
metric_value = None
setup_mode_value = None
embedding_dict = {}
else: else:
from astrapy.info import CollectionVectorServiceOptions from astrapy.info import CollectionVectorServiceOptions
# Fetch values from kwargs if any self.* attributes are None # Grab the collection options if available
dict_options = vectorize_options or self.build_vectorize_options() collection_options = self.get_collection_options()
# Ensure collection_options and its nested attributes are handled safely
authentication = getattr(self, "z_04_authentication", {}) or (
collection_options.service.authentication if collection_options and collection_options.service else {}
)
# Build the vectorize options dictionary
dict_options = vectorize_options or self.build_vectorize_options(
embedding_provider=(
getattr(self, "embedding_provider", None)
or (
collection_options.service.provider
if collection_options and collection_options.service
else None
)
),
model=(
getattr(self, "model", None)
or (
collection_options.service.model_name
if collection_options and collection_options.service
else None
)
),
z_01_model_parameters=(
getattr(self, "z_01_model_parameters", None)
or (
collection_options.service.parameters
if collection_options and collection_options.service
else None
)
),
z_02_api_key_name=(
getattr(self, "z_02_api_key_name", None)
or (authentication.get("apiKey") if authentication else None)
),
z_03_provider_api_key=(
getattr(self, "z_03_provider_api_key", None)
or (authentication.get("providerKey") if authentication else None)
),
z_04_authentication=authentication,
)
# Set the embedding dictionary # Set the embedding dictionary
embedding_dict = { embedding_dict = {
@ -484,12 +660,17 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
try: try:
vector_store = AstraDBVectorStore( vector_store = AstraDBVectorStore(
collection_name=self.collection_name,
token=self.token, token=self.token,
api_endpoint=self.api_endpoint, api_endpoint=self.api_endpoint,
namespace=self.keyspace or None, namespace=self.keyspace or None,
environment=parse_api_endpoint(self.api_endpoint).environment if self.api_endpoint else None, collection_name=getattr(self, "collection_name_new", None) or self.collection_name,
metric=self.metric or None, autodetect_collection=autodetect,
environment=(
parse_api_endpoint(getattr(self, "api_endpoint", None)).environment
if getattr(self, "api_endpoint", None)
else None
),
metric=metric_value,
batch_size=self.batch_size or None, batch_size=self.batch_size or None,
bulk_insert_batch_concurrency=self.bulk_insert_batch_concurrency or None, bulk_insert_batch_concurrency=self.bulk_insert_batch_concurrency or None,
bulk_insert_overwrite_concurrency=self.bulk_insert_overwrite_concurrency or None, bulk_insert_overwrite_concurrency=self.bulk_insert_overwrite_concurrency or None,

View file

@ -20,7 +20,7 @@ from langflow.io import (
from langflow.schema import Data from langflow.schema import Data
class AstraGraphVectorStoreComponent(LCVectorStoreComponent): class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
display_name: str = "Astra DB Graph" display_name: str = "Astra DB Graph"
description: str = "Implementation of Graph Vector Store using Astra DB" description: str = "Implementation of Graph Vector Store using Astra DB"
documentation: str = "https://python.langchain.com/api_reference/astradb/graph_vectorstores/langchain_astradb.graph_vectorstores.AstraDBGraphVectorStore.html" documentation: str = "https://python.langchain.com/api_reference/astradb/graph_vectorstores/langchain_astradb.graph_vectorstores.AstraDBGraphVectorStore.html"

File diff suppressed because one or more lines are too long

View file

@ -8,7 +8,7 @@ from langflow.components.outputs import ChatOutput
from langflow.components.processing import ParseDataComponent from langflow.components.processing import ParseDataComponent
from langflow.components.processing.split_text import SplitTextComponent from langflow.components.processing.split_text import SplitTextComponent
from langflow.components.prompts import PromptComponent from langflow.components.prompts import PromptComponent
from langflow.components.vectorstores import AstraVectorStoreComponent from langflow.components.vectorstores import AstraDBVectorStoreComponent
from langflow.graph import Graph from langflow.graph import Graph
@ -18,7 +18,7 @@ def ingestion_graph():
text_splitter = SplitTextComponent() text_splitter = SplitTextComponent()
text_splitter.set(data_inputs=file_component.load_files) text_splitter.set(data_inputs=file_component.load_files)
openai_embeddings = OpenAIEmbeddingsComponent() openai_embeddings = OpenAIEmbeddingsComponent()
vector_store = AstraVectorStoreComponent() vector_store = AstraDBVectorStoreComponent()
vector_store.set( vector_store.set(
embedding_model=openai_embeddings.build_embeddings, embedding_model=openai_embeddings.build_embeddings,
ingest_data=text_splitter.split_text, ingest_data=text_splitter.split_text,
@ -31,7 +31,7 @@ def rag_graph():
# RAG Graph # RAG Graph
openai_embeddings = OpenAIEmbeddingsComponent() openai_embeddings = OpenAIEmbeddingsComponent()
chat_input = ChatInput() chat_input = ChatInput()
rag_vector_store = AstraVectorStoreComponent() rag_vector_store = AstraDBVectorStoreComponent()
rag_vector_store.set( rag_vector_store.set(
search_input=chat_input.message_response, search_input=chat_input.message_response,
embedding_model=openai_embeddings.build_embeddings, embedding_model=openai_embeddings.build_embeddings,

View file

@ -4,7 +4,7 @@ import pytest
from astrapy.db import AstraDB from astrapy.db import AstraDB
from langchain_core.documents import Document from langchain_core.documents import Document
from langflow.components.embeddings import OpenAIEmbeddingsComponent from langflow.components.embeddings import OpenAIEmbeddingsComponent
from langflow.components.vectorstores import AstraVectorStoreComponent from langflow.components.vectorstores import AstraDBVectorStoreComponent
from langflow.schema.data import Data from langflow.schema.data import Data
from tests.api_keys import get_astradb_api_endpoint, get_astradb_application_token, get_openai_api_key from tests.api_keys import get_astradb_api_endpoint, get_astradb_application_token, get_openai_api_key
@ -43,7 +43,7 @@ async def test_base(astradb_client: AstraDB):
api_endpoint = get_astradb_api_endpoint() api_endpoint = get_astradb_api_endpoint()
results = await run_single_component( results = await run_single_component(
AstraVectorStoreComponent, AstraDBVectorStoreComponent,
inputs={ inputs={
"token": application_token, "token": application_token,
"api_endpoint": api_endpoint, "api_endpoint": api_endpoint,
@ -69,7 +69,7 @@ async def test_astra_embeds_and_search():
api_endpoint = get_astradb_api_endpoint() api_endpoint = get_astradb_api_endpoint()
results = await run_single_component( results = await run_single_component(
AstraVectorStoreComponent, AstraDBVectorStoreComponent,
inputs={ inputs={
"token": application_token, "token": application_token,
"api_endpoint": api_endpoint, "api_endpoint": api_endpoint,
@ -111,7 +111,7 @@ def test_astra_vectorize():
documents = [Document(page_content="test1"), Document(page_content="test2")] documents = [Document(page_content="test1"), Document(page_content="test2")]
records = [Data.from_document(d) for d in documents] records = [Data.from_document(d) for d in documents]
component = AstraVectorStoreComponent() component = AstraDBVectorStoreComponent()
vectorize_options = component.build_vectorize_options(**options_comp) vectorize_options = component.build_vectorize_options(**options_comp)
component.build( component.build(
@ -167,7 +167,7 @@ def test_astra_vectorize_with_provider_api_key():
documents = [Document(page_content="test1"), Document(page_content="test2")] documents = [Document(page_content="test1"), Document(page_content="test2")]
records = [Data.from_document(d) for d in documents] records = [Data.from_document(d) for d in documents]
component = AstraVectorStoreComponent() component = AstraDBVectorStoreComponent()
vectorize_options = component.build_vectorize_options(**options_comp) vectorize_options = component.build_vectorize_options(**options_comp)
component.build( component.build(
@ -222,7 +222,7 @@ def test_astra_vectorize_passes_authentication():
documents = [Document(page_content="test1"), Document(page_content="test2")] documents = [Document(page_content="test1"), Document(page_content="test2")]
records = [Data.from_document(d) for d in documents] records = [Data.from_document(d) for d in documents]
component = AstraVectorStoreComponent() component = AstraDBVectorStoreComponent()
vectorize_options = component.build_vectorize_options(**options_comp) vectorize_options = component.build_vectorize_options(**options_comp)
component.build( component.build(

View file

@ -11,7 +11,7 @@ from langflow.components.outputs import ChatOutput
from langflow.components.processing import ParseDataComponent from langflow.components.processing import ParseDataComponent
from langflow.components.processing.split_text import SplitTextComponent from langflow.components.processing.split_text import SplitTextComponent
from langflow.components.prompts import PromptComponent from langflow.components.prompts import PromptComponent
from langflow.components.vectorstores import AstraVectorStoreComponent from langflow.components.vectorstores import AstraDBVectorStoreComponent
from langflow.graph import Graph from langflow.graph import Graph
from langflow.graph.graph.constants import Finish from langflow.graph.graph.constants import Finish
from langflow.schema import Data from langflow.schema import Data
@ -29,7 +29,7 @@ def ingestion_graph():
openai_embeddings.set( openai_embeddings.set(
openai_api_key="sk-123", openai_api_base="https://api.openai.com/v1", openai_api_type="openai" openai_api_key="sk-123", openai_api_base="https://api.openai.com/v1", openai_api_type="openai"
) )
vector_store = AstraVectorStoreComponent(_id="vector-store-123") vector_store = AstraDBVectorStoreComponent(_id="vector-store-123")
vector_store.set( vector_store.set(
embedding_model=openai_embeddings.build_embeddings, embedding_model=openai_embeddings.build_embeddings,
ingest_data=text_splitter.split_text, ingest_data=text_splitter.split_text,
@ -48,7 +48,7 @@ def rag_graph():
openai_embeddings = OpenAIEmbeddingsComponent(_id="openai-embeddings-124") openai_embeddings = OpenAIEmbeddingsComponent(_id="openai-embeddings-124")
chat_input = ChatInput(_id="chatinput-123") chat_input = ChatInput(_id="chatinput-123")
chat_input.get_output("message").value = "What is the meaning of life?" chat_input.get_output("message").value = "What is the meaning of life?"
rag_vector_store = AstraVectorStoreComponent(_id="rag-vector-store-123") rag_vector_store = AstraDBVectorStoreComponent(_id="rag-vector-store-123")
rag_vector_store.set( rag_vector_store.set(
search_input=chat_input.message_response, search_input=chat_input.message_response,
api_endpoint="https://astra.example.com", api_endpoint="https://astra.example.com",