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_graph import AstraGraphVectorStoreComponent
from .astradb import AstraDBVectorStoreComponent
from .astradb_graph import AstraDBGraphVectorStoreComponent
from .cassandra import CassandraVectorStoreComponent
from .cassandra_graph import CassandraGraphVectorStoreComponent
from .chroma import ChromaVectorStoreComponent
@ -23,8 +23,8 @@ from .vectara_self_query import VectaraSelfQueryRetriverComponent
from .weaviate import WeaviateVectorStoreComponent
__all__ = [
"AstraGraphVectorStoreComponent",
"AstraVectorStoreComponent",
"AstraDBGraphVectorStoreComponent",
"AstraDBVectorStoreComponent",
"CassandraGraphVectorStoreComponent",
"CassandraVectorStoreComponent",
"ChromaVectorStoreComponent",

View file

@ -4,6 +4,7 @@ from collections import defaultdict
import orjson
from astrapy import DataAPIClient
from astrapy.admin import parse_api_endpoint
from astrapy.exceptions import CollectionNotFoundException
from langchain_astradb import AstraDBVectorStore
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
@ -22,7 +23,7 @@ from langflow.io import (
from langflow.schema import Data
class AstraVectorStoreComponent(LCVectorStoreComponent):
class AstraDBVectorStoreComponent(LCVectorStoreComponent):
display_name: str = "Astra DB"
description: str = "Implementation of Vector Store using Astra DB with search capabilities"
documentation: str = "https://docs.langflow.org/starter-projects-vector-store-rag"
@ -31,6 +32,24 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
_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(
list,
{
@ -61,7 +80,7 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
],
],
"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"]],
"Upstage": ["upstageAI", ["solar-embedding-1-large"]],
"Voyage AI": [
@ -87,20 +106,14 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
value="ASTRA_DB_API_ENDPOINT",
required=True,
),
StrInput(
DropdownInput(
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.",
required=True,
),
MultilineInput(
name="search_input",
display_name="Search Input",
),
DataInput(
name="ingest_data",
display_name="Ingest Data",
is_list=True,
real_time_refresh=True,
refresh_button=True,
options=[],
),
StrInput(
name="keyspace",
@ -108,6 +121,50 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
info="Optional keyspace within Astra DB to use for the collection.",
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(
name="embedding_choice",
display_name="Embedding Model or Astra Vectorize",
@ -172,14 +229,14 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
name="metadata_indexing_include",
display_name="Metadata Indexing Include",
info="Optional list of metadata fields to include in the indexing.",
is_list=True,
list=True,
advanced=True,
),
StrInput(
name="metadata_indexing_exclude",
display_name="Metadata Indexing Exclude",
info="Optional list of metadata fields to exclude from the indexing.",
is_list=True,
list=True,
advanced=True,
),
StrInput(
@ -189,42 +246,6 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
"See https://docs.datastax.com/en/astra-db-serverless/api-reference/collections.html#the-indexing-option",
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):
@ -289,8 +310,110 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
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):
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":
self.del_fields(build_config, ["embedding_model"])
@ -363,7 +486,7 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
new_parameter_1 = DictInput(
name="z_01_model_parameters",
display_name="Model Parameters",
is_list=True,
list=True,
).to_dict()
new_parameter_2 = MessageTextInput(
@ -386,7 +509,7 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
new_parameter_4 = DictInput(
name="z_04_authentication",
display_name="Authentication Parameters",
is_list=True,
list=True,
).to_dict()
self.insert_in_dict(
@ -415,11 +538,10 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
setattr(self, attribute, 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(
"embedding_provider"
)
provider_mapping = self.update_providers_mapping()
provider_value = provider_mapping.get(self.embedding_provider, [None])[0] or kwargs.get("embedding_provider")
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", {})
# 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")
if 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
if not authentication:
@ -466,13 +591,64 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
msg = f"Invalid setup mode: {self.setup_mode}"
raise ValueError(msg) from e
metric_value = self.metric or None
autodetect = False
if self.embedding_choice == "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:
from astrapy.info import CollectionVectorServiceOptions
# Fetch values from kwargs if any self.* attributes are None
dict_options = vectorize_options or self.build_vectorize_options()
# Grab the collection options if available
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
embedding_dict = {
@ -484,12 +660,17 @@ class AstraVectorStoreComponent(LCVectorStoreComponent):
try:
vector_store = AstraDBVectorStore(
collection_name=self.collection_name,
token=self.token,
api_endpoint=self.api_endpoint,
namespace=self.keyspace or None,
environment=parse_api_endpoint(self.api_endpoint).environment if self.api_endpoint else None,
metric=self.metric or None,
collection_name=getattr(self, "collection_name_new", None) or self.collection_name,
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,
bulk_insert_batch_concurrency=self.bulk_insert_batch_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
class AstraGraphVectorStoreComponent(LCVectorStoreComponent):
class AstraDBGraphVectorStoreComponent(LCVectorStoreComponent):
display_name: str = "Astra DB Graph"
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"

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

View file

@ -4,7 +4,7 @@ import pytest
from astrapy.db import AstraDB
from langchain_core.documents import Document
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 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()
results = await run_single_component(
AstraVectorStoreComponent,
AstraDBVectorStoreComponent,
inputs={
"token": application_token,
"api_endpoint": api_endpoint,
@ -69,7 +69,7 @@ async def test_astra_embeds_and_search():
api_endpoint = get_astradb_api_endpoint()
results = await run_single_component(
AstraVectorStoreComponent,
AstraDBVectorStoreComponent,
inputs={
"token": application_token,
"api_endpoint": api_endpoint,
@ -111,7 +111,7 @@ def test_astra_vectorize():
documents = [Document(page_content="test1"), Document(page_content="test2")]
records = [Data.from_document(d) for d in documents]
component = AstraVectorStoreComponent()
component = AstraDBVectorStoreComponent()
vectorize_options = component.build_vectorize_options(**options_comp)
component.build(
@ -167,7 +167,7 @@ def test_astra_vectorize_with_provider_api_key():
documents = [Document(page_content="test1"), Document(page_content="test2")]
records = [Data.from_document(d) for d in documents]
component = AstraVectorStoreComponent()
component = AstraDBVectorStoreComponent()
vectorize_options = component.build_vectorize_options(**options_comp)
component.build(
@ -222,7 +222,7 @@ def test_astra_vectorize_passes_authentication():
documents = [Document(page_content="test1"), Document(page_content="test2")]
records = [Data.from_document(d) for d in documents]
component = AstraVectorStoreComponent()
component = AstraDBVectorStoreComponent()
vectorize_options = component.build_vectorize_options(**options_comp)
component.build(

View file

@ -11,7 +11,7 @@ from langflow.components.outputs import ChatOutput
from langflow.components.processing import ParseDataComponent
from langflow.components.processing.split_text import SplitTextComponent
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.graph.constants import Finish
from langflow.schema import Data
@ -29,7 +29,7 @@ def ingestion_graph():
openai_embeddings.set(
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(
embedding_model=openai_embeddings.build_embeddings,
ingest_data=text_splitter.split_text,
@ -48,7 +48,7 @@ def rag_graph():
openai_embeddings = OpenAIEmbeddingsComponent(_id="openai-embeddings-124")
chat_input = ChatInput(_id="chatinput-123")
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(
search_input=chat_input.message_response,
api_endpoint="https://astra.example.com",