feat: Add Hybrid Search functionality to AstraDB + AstraPy / LangChain Updates (#7358)
* feat: Add Hybrid Search functionality and AstraPy 2.0 and associated deps (#7357) * astrapy 2.0 tentative full pass * Update the create collection function --------- Co-authored-by: Stefano Lottini <stefano.lottini@datastax.com> * Update deps * Update uv.lock * Fix linting errors in astradb * Update package lock * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * Add basic UI scaffolding for hybrid search * [autofix.ci] apply automated fixes * Continue to clean up component * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * Fix the keyspace compatibility * [autofix.ci] apply automated fixes * feat: add nodeId, nodeClass, and handleNodeClass props to dropdown an… (#7406) feat: add nodeId, nodeClass, and handleNodeClass props to dropdown and string render components Co-authored-by: deon-sanchez <deon.sanchez@datastax.com> * Update uv.lock * Update uv.lock * Add hybrid search support in collection creation * [autofix.ci] apply automated fixes * Updates from review comments * [autofix.ci] apply automated fixes * Add in lexical search support * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * Detect collection hybrid params * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * Pass lexical terms at search time * [autofix.ci] apply automated fixes * Update test_astra_component.py * Add Query Input and Mixin on backend * Adds Query on supported types * Adds types for query modal and component * Adds size for new query modal * Adds query modal * Adds query component * Adds query component on parameter render * [autofix.ci] apply automated fixes * Feedback from review * [autofix.ci] apply automated fixes * ✨ (switch-case-size.ts): Update height value to 'h-fit' for 'small-query' case to improve responsiveness ✨ (queryInputComponent.spec.ts): Add unit test for user interaction with query input component, including updating code and testing functionality * Switch to multiline for lexical terms * [autofix.ci] apply automated fixes * Create Hybrid Search RAG.json * Update Hybrid Search RAG.json * Added queryInput in vectorstore model * Added queryInput in lexical terms * Update model.py * Update Hybrid Search RAG.json * Add query support in field validation * fix: bump Astra Assistants version to support AstraPy 2.0 (#7535) 2.2.12 Co-authored-by: phact <estevezsebastian@gmail.com> * Update uv.lock * Fixed QueryInput not receiving text from handle * Set search type to similarity search when hybrid * Always set to similarity when we have the reranker * [autofix.ci] apply automated fixes * Add logging for hybrid search support * Update starter projects * Update Hybrid Search RAG.json * Added dropdown toggle on backend * Added toggle on dropdown on frontend * Added showing only value if there is just one option in the dropdown * Added toggle to Dropdown Input on Astra Db * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes (attempt 2/3) * init toggle value as true or false * Change it to send null value if toggle is disabled * Added resizer on search query * Added Search Hybrid, Lexical and Vector icons * Added icons and new Lexical Search on Dropdown Input of Astra DB * Updated starter projects * Changed descriptions on astradb component * Changed starter projects * Lexical search option for dropdown * Update astradb.py * Update starter projects * One small lexical update * Update astradb.py * Update projects * [autofix.ci] apply automated fixes * Fixed dropdown changing when toggle is off * Update astradb.py * [autofix.ci] apply automated fixes * Don't show lexical terms on new collection creation * ✨ (actionsMainPage-shard-0.spec.ts): add functionality to add flow to test on empty langflow button click ✨ (filterEdge-shard-1.spec.ts): add functionality to add flow to test on empty langflow button click ♻️ (await-bootstrap-test.ts): refactor code to reuse addFlowToTestOnEmptyLangflow function for adding flow to test on empty langflow button click * [autofix.ci] apply automated fixes * 🐛 (filterEdge-shard-1.spec.ts): fix incorrect reference to memoriesAstra DB Chat Memory, update to memoriesMem0 Chat Memory for accurate testing data. --------- Co-authored-by: Stefano Lottini <stefano.lottini@datastax.com> Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: deon-sanchez <deon.sanchez@datastax.com> Co-authored-by: Lucas Oliveira <lucas.edu.oli@hotmail.com> Co-authored-by: cristhianzl <cristhian.lousa@gmail.com> Co-authored-by: phact <estevezsebastian@gmail.com>
This commit is contained in:
parent
fb79b80f91
commit
907a594428
30 changed files with 4901 additions and 1804 deletions
|
|
@ -6,7 +6,7 @@ from langflow.custom import Component
|
|||
from langflow.field_typing import Text, VectorStore
|
||||
from langflow.helpers.data import docs_to_data
|
||||
from langflow.inputs.inputs import BoolInput
|
||||
from langflow.io import HandleInput, MultilineInput, Output
|
||||
from langflow.io import HandleInput, Output, QueryInput
|
||||
from langflow.schema import Data, DataFrame
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -62,9 +62,11 @@ class LCVectorStoreComponent(Component):
|
|||
input_types=["Data", "DataFrame"],
|
||||
is_list=True,
|
||||
),
|
||||
MultilineInput(
|
||||
QueryInput(
|
||||
name="search_query",
|
||||
display_name="Search Query",
|
||||
info="Enter a query to run a combined similarity and lexical terms search.",
|
||||
placeholder="Enter a query...",
|
||||
tool_mode=True,
|
||||
),
|
||||
BoolInput(
|
||||
|
|
|
|||
|
|
@ -112,7 +112,7 @@ class AstraVectorizeComponent(Component):
|
|||
if api_key_name:
|
||||
authentication["providerKey"] = api_key_name
|
||||
return {
|
||||
# must match astrapy.info.CollectionVectorServiceOptions
|
||||
# must match astrapy.info.VectorServiceOptions
|
||||
"collection_vector_service_options": {
|
||||
"provider": provider_value,
|
||||
"modelName": self.model_name,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from datetime import datetime, timezone
|
|||
from typing import Any
|
||||
|
||||
from astrapy import Collection, DataAPIClient, Database
|
||||
from astrapy.admin import parse_api_endpoint
|
||||
from langchain.pydantic_v1 import BaseModel, Field, create_model
|
||||
from langchain_core.tools import StructuredTool, Tool
|
||||
|
||||
|
|
@ -195,7 +196,8 @@ class AstraDBToolComponent(LCToolComponent):
|
|||
return self._cached_collection
|
||||
|
||||
try:
|
||||
cached_client = DataAPIClient(self.token)
|
||||
environment = parse_api_endpoint(self.api_endpoint).environment
|
||||
cached_client = DataAPIClient(self.token, environment=environment)
|
||||
cached_db = cached_client.get_database(self.api_endpoint, keyspace=self.keyspace)
|
||||
self._cached_collection = cached_db.get_collection(self.collection_name)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -2,9 +2,11 @@ import re
|
|||
from collections import defaultdict
|
||||
from dataclasses import asdict, dataclass, field
|
||||
|
||||
from astrapy import AstraDBAdmin, DataAPIClient, Database
|
||||
from astrapy.info import CollectionDescriptor
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
from astrapy import DataAPIClient, Database
|
||||
from astrapy.data.info.reranking import RerankServiceOptions
|
||||
from astrapy.info import CollectionDescriptor, CollectionLexicalOptions, CollectionRerankOptions
|
||||
from langchain_astradb import AstraDBVectorStore, VectorServiceOptions
|
||||
from langchain_astradb.utils.astradb import HybridSearchMode, _AstraDBCollectionEnvironment
|
||||
|
||||
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
|
||||
from langflow.base.vectorstores.vector_store_connection_decorator import vector_store_connection
|
||||
|
|
@ -15,6 +17,7 @@ from langflow.io import (
|
|||
DropdownInput,
|
||||
HandleInput,
|
||||
IntInput,
|
||||
QueryInput,
|
||||
SecretStrInput,
|
||||
StrInput,
|
||||
)
|
||||
|
|
@ -136,12 +139,15 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
real_time_refresh=True,
|
||||
input_types=[],
|
||||
),
|
||||
StrInput(
|
||||
DropdownInput(
|
||||
name="environment",
|
||||
display_name="Environment",
|
||||
info="The environment for the Astra DB API Endpoint.",
|
||||
options=["prod", "test", "dev"],
|
||||
value="prod",
|
||||
advanced=True,
|
||||
real_time_refresh=True,
|
||||
combobox=True,
|
||||
),
|
||||
DropdownInput(
|
||||
name="database_name",
|
||||
|
|
@ -157,7 +163,15 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
name="api_endpoint",
|
||||
display_name="Astra DB API Endpoint",
|
||||
info="The API Endpoint for the Astra DB instance. Supercedes database selection.",
|
||||
show=False,
|
||||
),
|
||||
DropdownInput(
|
||||
name="keyspace",
|
||||
display_name="Keyspace",
|
||||
info="Optional keyspace within Astra DB to use for the collection.",
|
||||
advanced=True,
|
||||
options=[],
|
||||
real_time_refresh=True,
|
||||
),
|
||||
DropdownInput(
|
||||
name="collection_name",
|
||||
|
|
@ -168,22 +182,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
real_time_refresh=True,
|
||||
dialog_inputs=asdict(NewCollectionInput()),
|
||||
combobox=True,
|
||||
advanced=True,
|
||||
),
|
||||
StrInput(
|
||||
name="keyspace",
|
||||
display_name="Keyspace",
|
||||
info="Optional keyspace within Astra DB to use for the collection.",
|
||||
advanced=True,
|
||||
),
|
||||
DropdownInput(
|
||||
name="embedding_choice",
|
||||
display_name="Embedding Model or Astra Vectorize",
|
||||
info="Choose an embedding model or use Astra Vectorize.",
|
||||
options=["Embedding Model", "Astra Vectorize"],
|
||||
value="Embedding Model",
|
||||
advanced=True,
|
||||
real_time_refresh=True,
|
||||
show=False,
|
||||
),
|
||||
HandleInput(
|
||||
name="embedding_model",
|
||||
|
|
@ -191,8 +190,40 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
input_types=["Embeddings"],
|
||||
info="Specify the Embedding Model. Not required for Astra Vectorize collections.",
|
||||
required=False,
|
||||
show=False,
|
||||
),
|
||||
*LCVectorStoreComponent.inputs,
|
||||
DropdownInput(
|
||||
name="search_method",
|
||||
display_name="Search Method",
|
||||
info=(
|
||||
"Determine how your content is matched: Vector finds semantic similarity, "
|
||||
"and Hybrid Search (suggested) combines both approaches "
|
||||
"with a reranker."
|
||||
),
|
||||
options=["Hybrid Search", "Vector Search"], # TODO: Restore Lexical Search?
|
||||
options_metadata=[{"icon": "SearchHybrid"}, {"icon": "SearchVector"}],
|
||||
value="Vector Search",
|
||||
advanced=True,
|
||||
real_time_refresh=True,
|
||||
),
|
||||
DropdownInput(
|
||||
name="reranker",
|
||||
display_name="Reranker",
|
||||
info="Post-retrieval model that re-scores results for optimal relevance ranking.",
|
||||
show=False,
|
||||
toggle=True,
|
||||
),
|
||||
QueryInput(
|
||||
name="lexical_terms",
|
||||
display_name="Lexical Terms",
|
||||
info="Add additional terms/keywords to augment search precision.",
|
||||
placeholder="Enter terms to search...",
|
||||
separator=" ",
|
||||
show=False,
|
||||
value="",
|
||||
advanced=True,
|
||||
),
|
||||
IntInput(
|
||||
name="number_of_results",
|
||||
display_name="Number of Search Results",
|
||||
|
|
@ -262,12 +293,15 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
# TODO: Programmatically fetch the regions for each cloud provider
|
||||
return {
|
||||
"dev": {
|
||||
"Amazon Web Services": {
|
||||
"id": "aws",
|
||||
"regions": ["us-west-2"],
|
||||
},
|
||||
"Google Cloud Platform": {
|
||||
"id": "gcp",
|
||||
"regions": ["us-central1"],
|
||||
"regions": ["us-central1", "europe-west4"],
|
||||
},
|
||||
},
|
||||
# TODO: Check test regions
|
||||
"test": {
|
||||
"Google Cloud Platform": {
|
||||
"id": "gcp",
|
||||
|
|
@ -294,18 +328,19 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
def get_vectorize_providers(cls, token: str, environment: str | None = None, api_endpoint: str | None = None):
|
||||
try:
|
||||
# Get the admin object
|
||||
admin = AstraDBAdmin(token=token, environment=environment)
|
||||
db_admin = admin.get_database_admin(api_endpoint=api_endpoint)
|
||||
client = DataAPIClient(environment=environment)
|
||||
admin_client = client.get_admin()
|
||||
db_admin = admin_client.get_database_admin(api_endpoint, token=token)
|
||||
|
||||
# Get the list of embedding providers
|
||||
embedding_providers = db_admin.find_embedding_providers().as_dict()
|
||||
embedding_providers = db_admin.find_embedding_providers()
|
||||
|
||||
vectorize_providers_mapping = {}
|
||||
# Map the provider display name to the provider key and models
|
||||
for provider_key, provider_data in embedding_providers["embeddingProviders"].items():
|
||||
for provider_key, provider_data in embedding_providers.embedding_providers.items():
|
||||
# Get the provider display name and models
|
||||
display_name = provider_data["displayName"]
|
||||
models = [model["name"] for model in provider_data["models"]]
|
||||
display_name = provider_data.display_name
|
||||
models = [model.name for model in provider_data.models]
|
||||
|
||||
# Build our mapping
|
||||
vectorize_providers_mapping[display_name] = [provider_key, models]
|
||||
|
|
@ -325,7 +360,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
environment: str | None = None,
|
||||
keyspace: str | None = None,
|
||||
):
|
||||
client = DataAPIClient(token=token, environment=environment)
|
||||
client = DataAPIClient(environment=environment)
|
||||
|
||||
# Get the admin object
|
||||
admin_client = client.get_admin(token=token)
|
||||
|
|
@ -358,20 +393,14 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
dimension: int | None = None,
|
||||
embedding_generation_provider: str | None = None,
|
||||
embedding_generation_model: str | None = None,
|
||||
reranker: str | None = None,
|
||||
):
|
||||
# Create the data API client
|
||||
client = DataAPIClient(token=token, environment=environment)
|
||||
|
||||
# Get the database object
|
||||
database = client.get_async_database(api_endpoint=api_endpoint, token=token)
|
||||
|
||||
# Build vectorize options, if needed
|
||||
vectorize_options = None
|
||||
if not dimension:
|
||||
vectorize_options = CollectionVectorServiceOptions(
|
||||
provider=cls.get_vectorize_providers(
|
||||
token=token, environment=environment, api_endpoint=api_endpoint
|
||||
).get(embedding_generation_provider, [None, []])[0],
|
||||
providers = cls.get_vectorize_providers(token=token, environment=environment, api_endpoint=api_endpoint)
|
||||
vectorize_options = VectorServiceOptions(
|
||||
provider=providers.get(embedding_generation_provider, [None, []])[0],
|
||||
model_name=embedding_generation_model,
|
||||
)
|
||||
|
||||
|
|
@ -380,44 +409,53 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
msg = "Collection name is required to create a new collection."
|
||||
raise ValueError(msg)
|
||||
|
||||
# Create the collection
|
||||
return await database.create_collection(
|
||||
name=new_collection_name,
|
||||
keyspace=keyspace,
|
||||
dimension=dimension,
|
||||
service=vectorize_options,
|
||||
)
|
||||
# Define the base arguments being passed to the create collection function
|
||||
base_args = {
|
||||
"collection_name": new_collection_name,
|
||||
"token": token,
|
||||
"api_endpoint": api_endpoint,
|
||||
"keyspace": keyspace,
|
||||
"environment": environment,
|
||||
"embedding_dimension": dimension,
|
||||
"collection_vector_service_options": vectorize_options,
|
||||
}
|
||||
|
||||
# Add optional arguments only if environment is "dev"
|
||||
if environment == "dev" and reranker: # TODO: Remove conditional check soon
|
||||
# Split the reranker field into a provider a model name
|
||||
provider, _ = reranker.split("/")
|
||||
base_args["collection_rerank"] = CollectionRerankOptions(
|
||||
service=RerankServiceOptions(provider=provider, model_name=reranker),
|
||||
)
|
||||
base_args["collection_lexical"] = CollectionLexicalOptions(analyzer="STANDARD")
|
||||
|
||||
_AstraDBCollectionEnvironment(**base_args)
|
||||
|
||||
@classmethod
|
||||
def get_database_list_static(cls, token: str, environment: str | None = None):
|
||||
client = DataAPIClient(token=token, environment=environment)
|
||||
client = DataAPIClient(environment=environment)
|
||||
|
||||
# Get the admin object
|
||||
admin_client = client.get_admin(token=token)
|
||||
|
||||
# Get the list of databases
|
||||
db_list = list(admin_client.list_databases())
|
||||
|
||||
# Set the environment properly
|
||||
env_string = ""
|
||||
if environment and environment != "prod":
|
||||
env_string = f"-{environment}"
|
||||
db_list = admin_client.list_databases()
|
||||
|
||||
# Generate the api endpoint for each database
|
||||
db_info_dict = {}
|
||||
for db in db_list:
|
||||
try:
|
||||
# Get the API endpoint for the database
|
||||
api_endpoint = f"https://{db.info.id}-{db.info.region}.apps.astra{env_string}.datastax.com"
|
||||
api_endpoint = db.regions[0].api_endpoint
|
||||
|
||||
# Get the number of collections
|
||||
try:
|
||||
# Get the number of collections in the database
|
||||
num_collections = len(
|
||||
list(
|
||||
client.get_database(
|
||||
api_endpoint=api_endpoint, token=token, keyspace=db.info.keyspace
|
||||
).list_collection_names(keyspace=db.info.keyspace)
|
||||
)
|
||||
client.get_database(
|
||||
api_endpoint,
|
||||
token=token,
|
||||
).list_collection_names()
|
||||
)
|
||||
except Exception: # noqa: BLE001
|
||||
if db.status != "PENDING":
|
||||
|
|
@ -425,8 +463,9 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
num_collections = 0
|
||||
|
||||
# Add the database to the dictionary
|
||||
db_info_dict[db.info.name] = {
|
||||
db_info_dict[db.name] = {
|
||||
"api_endpoint": api_endpoint,
|
||||
"keyspaces": db.keyspaces,
|
||||
"collections": num_collections,
|
||||
"status": db.status if db.status != "ACTIVE" else None,
|
||||
"org_id": db.org_id if db.org_id else None,
|
||||
|
|
@ -437,7 +476,10 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
return db_info_dict
|
||||
|
||||
def get_database_list(self):
|
||||
return self.get_database_list_static(token=self.token, environment=self.environment)
|
||||
return self.get_database_list_static(
|
||||
token=self.token,
|
||||
environment=self.environment,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def get_api_endpoint_static(
|
||||
|
|
@ -492,14 +534,14 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
if keyspace:
|
||||
return keyspace.strip()
|
||||
|
||||
return None
|
||||
return "default_keyspace"
|
||||
|
||||
def get_database_object(self, api_endpoint: str | None = None):
|
||||
try:
|
||||
client = DataAPIClient(token=self.token, environment=self.environment)
|
||||
client = DataAPIClient(environment=self.environment)
|
||||
|
||||
return client.get_database(
|
||||
api_endpoint=api_endpoint or self.get_api_endpoint(),
|
||||
api_endpoint or self.get_api_endpoint(),
|
||||
token=self.token,
|
||||
keyspace=self.get_keyspace(),
|
||||
)
|
||||
|
|
@ -510,15 +552,15 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
def collection_data(self, collection_name: str, database: Database | None = None):
|
||||
try:
|
||||
if not database:
|
||||
client = DataAPIClient(token=self.token, environment=self.environment)
|
||||
client = DataAPIClient(environment=self.environment)
|
||||
|
||||
database = client.get_database(
|
||||
api_endpoint=self.get_api_endpoint(),
|
||||
self.get_api_endpoint(),
|
||||
token=self.token,
|
||||
keyspace=self.get_keyspace(),
|
||||
)
|
||||
|
||||
collection = database.get_collection(collection_name, keyspace=self.get_keyspace())
|
||||
collection = database.get_collection(collection_name)
|
||||
|
||||
return collection.estimated_document_count()
|
||||
except Exception as e: # noqa: BLE001
|
||||
|
|
@ -534,6 +576,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"status": info["status"],
|
||||
"collections": info["collections"],
|
||||
"api_endpoint": info["api_endpoint"],
|
||||
"keyspaces": info["keyspaces"],
|
||||
"org_id": info["org_id"],
|
||||
}
|
||||
for name, info in self.get_database_list().items()
|
||||
|
|
@ -546,13 +589,18 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
def get_provider_icon(cls, collection: CollectionDescriptor | None = None, provider_name: str | None = None) -> str:
|
||||
# Get the provider name from the collection
|
||||
provider_name = provider_name or (
|
||||
collection.options.vector.service.provider
|
||||
if collection and collection.options and collection.options.vector and collection.options.vector.service
|
||||
collection.definition.vector.service.provider
|
||||
if (
|
||||
collection
|
||||
and collection.definition
|
||||
and collection.definition.vector
|
||||
and collection.definition.vector.service
|
||||
)
|
||||
else None
|
||||
)
|
||||
|
||||
# If there is no provider, use the vector store icon
|
||||
if not provider_name or provider_name == "Bring your own":
|
||||
if not provider_name or provider_name.lower() == "bring your own":
|
||||
return "vectorstores"
|
||||
|
||||
# Map provider casings
|
||||
|
|
@ -581,7 +629,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
database = self.get_database_object(api_endpoint=api_endpoint)
|
||||
|
||||
# Get the list of collections
|
||||
collection_list = list(database.list_collections(keyspace=self.get_keyspace()))
|
||||
collection_list = database.list_collections(keyspace=self.get_keyspace())
|
||||
|
||||
# Return the list of collections and metadata associated
|
||||
return [
|
||||
|
|
@ -589,11 +637,15 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"name": col.name,
|
||||
"records": self.collection_data(collection_name=col.name, database=database),
|
||||
"provider": (
|
||||
col.options.vector.service.provider if col.options.vector and col.options.vector.service else None
|
||||
col.definition.vector.service.provider
|
||||
if col.definition.vector and col.definition.vector.service
|
||||
else None
|
||||
),
|
||||
"icon": self.get_provider_icon(collection=col),
|
||||
"model": (
|
||||
col.options.vector.service.model_name if col.options.vector and col.options.vector.service else None
|
||||
col.definition.vector.service.model_name
|
||||
if col.definition.vector and col.definition.vector.service
|
||||
else None
|
||||
),
|
||||
}
|
||||
for col in collection_list
|
||||
|
|
@ -679,7 +731,6 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"""Reset collection list options based on provided configuration."""
|
||||
# Get collection options
|
||||
collection_options = self._initialize_collection_options(api_endpoint=build_config["api_endpoint"]["value"])
|
||||
|
||||
# Update collection configuration
|
||||
collection_config = build_config["collection_name"]
|
||||
collection_config.update(
|
||||
|
|
@ -694,7 +745,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
collection_config["value"] = ""
|
||||
|
||||
# Set advanced status based on database selection
|
||||
collection_config["advanced"] = not build_config["database_name"]["value"]
|
||||
collection_config["show"] = bool(build_config["database_name"]["value"])
|
||||
|
||||
return build_config
|
||||
|
||||
|
|
@ -704,7 +755,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
database_options = self._initialize_database_options()
|
||||
|
||||
# Update cloud provider options
|
||||
env = self.environment or "prod"
|
||||
env = self.environment
|
||||
template = build_config["database_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"]
|
||||
template["02_cloud_provider"]["options"] = list(self.map_cloud_providers()[env].keys())
|
||||
|
||||
|
|
@ -721,10 +772,10 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
if database_config["value"] not in database_config["options"]:
|
||||
database_config["value"] = ""
|
||||
build_config["api_endpoint"]["value"] = ""
|
||||
build_config["collection_name"]["advanced"] = True
|
||||
build_config["collection_name"]["show"] = False
|
||||
|
||||
# Set advanced status based on token presence
|
||||
database_config["advanced"] = not build_config["token"]["value"]
|
||||
database_config["show"] = bool(build_config["token"]["value"])
|
||||
|
||||
return build_config
|
||||
|
||||
|
|
@ -732,12 +783,53 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"""Reset all build configuration options to default empty state."""
|
||||
# Reset database configuration
|
||||
database_config = build_config["database_name"]
|
||||
database_config.update({"options": [], "options_metadata": [], "value": "", "advanced": True})
|
||||
database_config.update({"options": [], "options_metadata": [], "value": "", "show": False})
|
||||
build_config["api_endpoint"]["value"] = ""
|
||||
|
||||
# Reset collection configuration
|
||||
collection_config = build_config["collection_name"]
|
||||
collection_config.update({"options": [], "options_metadata": [], "value": "", "advanced": True})
|
||||
collection_config.update({"options": [], "options_metadata": [], "value": "", "show": False})
|
||||
|
||||
return build_config
|
||||
|
||||
def _handle_hybrid_search_options(self, build_config: dict) -> dict:
|
||||
"""Set hybrid search options in the build configuration."""
|
||||
# Detect what hybrid options are available
|
||||
# Get the admin object
|
||||
client = DataAPIClient(environment=self.environment)
|
||||
admin_client = client.get_admin()
|
||||
db_admin = admin_client.get_database_admin(self.get_api_endpoint(), token=self.token)
|
||||
|
||||
# We will try to get the reranking providers to see if its hybrid emabled
|
||||
try:
|
||||
providers = db_admin.find_reranking_providers()
|
||||
build_config["reranker"]["options"] = [
|
||||
model.name for provider_data in providers.reranking_providers.values() for model in provider_data.models
|
||||
]
|
||||
build_config["reranker"]["options_metadata"] = [
|
||||
{"icon": self.get_provider_icon(provider_name=model.name.split("/")[0])}
|
||||
for provider in providers.reranking_providers.values()
|
||||
for model in provider.models
|
||||
]
|
||||
build_config["reranker"]["value"] = build_config["reranker"]["options"][0]
|
||||
|
||||
# Set the default search field to hybrid search
|
||||
build_config["search_method"]["show"] = True
|
||||
build_config["search_method"]["options"] = ["Hybrid Search", "Vector Search"]
|
||||
build_config["search_method"]["value"] = "Hybrid Search"
|
||||
except Exception as _: # noqa: BLE001
|
||||
build_config["reranker"]["options"] = []
|
||||
build_config["reranker"]["options_metadata"] = []
|
||||
|
||||
# Set the default search field to vector search
|
||||
build_config["search_method"]["show"] = False
|
||||
build_config["search_method"]["options"] = ["Vector Search"]
|
||||
build_config["search_method"]["value"] = "Vector Search"
|
||||
|
||||
# Set reranker and lexical terms options based on search method
|
||||
build_config["reranker"]["show"] = build_config["search_method"]["value"] == "Hybrid Search"
|
||||
if build_config["reranker"]["show"]:
|
||||
build_config["search_type"]["value"] = "Similarity"
|
||||
|
||||
return build_config
|
||||
|
||||
|
|
@ -778,10 +870,31 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
if field_name == "database_name" and not isinstance(field_value, dict):
|
||||
return self._handle_database_selection(build_config, field_value)
|
||||
|
||||
# Keyspace selection change
|
||||
if field_name == "keyspace":
|
||||
return self.reset_collection_list(build_config)
|
||||
|
||||
# Collection selection change
|
||||
if field_name == "collection_name" and not isinstance(field_value, dict):
|
||||
return self._handle_collection_selection(build_config, field_value)
|
||||
|
||||
# Search method selection change
|
||||
if field_name == "search_method":
|
||||
is_vector_search = field_value == "Vector Search"
|
||||
is_autodetect = build_config["autodetect_collection"]["value"]
|
||||
|
||||
# Configure lexical terms (same for both cases)
|
||||
build_config["lexical_terms"]["show"] = not is_vector_search
|
||||
build_config["lexical_terms"]["value"] = "" if is_vector_search else build_config["lexical_terms"]["value"]
|
||||
|
||||
# Toggle search type and score threshold based on search method
|
||||
build_config["search_type"]["show"] = is_vector_search
|
||||
build_config["search_score_threshold"]["show"] = is_vector_search
|
||||
|
||||
# Make sure the search_type is set to "Similarity"
|
||||
if not is_vector_search or is_autodetect:
|
||||
build_config["search_type"]["value"] = "Similarity"
|
||||
|
||||
return build_config
|
||||
|
||||
async def _create_new_database(self, build_config: dict, field_value: dict) -> None:
|
||||
|
|
@ -805,13 +918,14 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"status": "PENDING",
|
||||
"collections": 0,
|
||||
"api_endpoint": None,
|
||||
"keyspaces": [self.get_keyspace()],
|
||||
"org_id": None,
|
||||
}
|
||||
)
|
||||
|
||||
def _update_cloud_regions(self, build_config: dict, field_value: dict) -> dict:
|
||||
"""Update cloud provider regions in build config."""
|
||||
env = self.environment or "prod"
|
||||
env = self.environment
|
||||
cloud_provider = field_value["02_cloud_provider"]
|
||||
|
||||
# Update the region options based on the selected cloud provider
|
||||
|
|
@ -837,6 +951,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
dimension=field_value.get("04_dimension") if embedding_provider == "Bring your own" else None,
|
||||
embedding_generation_provider=embedding_provider,
|
||||
embedding_generation_model=field_value.get("03_embedding_generation_model"),
|
||||
reranker=self.reranker,
|
||||
)
|
||||
except Exception as e:
|
||||
msg = f"Error creating collection: {e}"
|
||||
|
|
@ -849,17 +964,21 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"options": build_config["collection_name"]["options"] + [field_value["01_new_collection_name"]],
|
||||
}
|
||||
)
|
||||
build_config["embedding_choice"]["value"] = "Astra Vectorize" if provider else "Embedding Model"
|
||||
build_config["embedding_model"]["advanced"] = bool(provider)
|
||||
build_config["embedding_model"]["show"] = not bool(provider)
|
||||
build_config["embedding_model"]["required"] = not bool(provider)
|
||||
build_config["collection_name"]["options_metadata"].append(
|
||||
{
|
||||
"records": 0,
|
||||
"provider": provider,
|
||||
"icon": self.get_provider_icon(provider_name=embedding_provider),
|
||||
"icon": self.get_provider_icon(provider_name=provider),
|
||||
"model": field_value.get("03_embedding_generation_model"),
|
||||
}
|
||||
)
|
||||
|
||||
# Make sure we always show the reranker options if the collection is hybrid enabled
|
||||
# And right now they always are
|
||||
build_config["lexical_terms"]["show"] = True
|
||||
|
||||
def _handle_database_selection(self, build_config: dict, field_value: str) -> dict:
|
||||
"""Handle database selection and update related configurations."""
|
||||
build_config = self.reset_database_list(build_config)
|
||||
|
|
@ -878,9 +997,17 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
if not org_id:
|
||||
return build_config
|
||||
|
||||
# Update the list of keyspaces based on the db info
|
||||
build_config["keyspace"]["options"] = build_config["database_name"]["options_metadata"][index]["keyspaces"]
|
||||
build_config["keyspace"]["value"] = (
|
||||
build_config["keyspace"]["options"] and build_config["keyspace"]["options"][0]
|
||||
if build_config["keyspace"]["value"] not in build_config["keyspace"]["options"]
|
||||
else build_config["keyspace"]["value"]
|
||||
)
|
||||
|
||||
# Get the database id for the selected database
|
||||
db_id = self.get_database_id_static(api_endpoint=build_config["api_endpoint"]["value"])
|
||||
keyspace = self.get_keyspace() or "default_keyspace"
|
||||
keyspace = self.get_keyspace()
|
||||
|
||||
# Update the helper text for the embedding provider field
|
||||
template = build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"]
|
||||
|
|
@ -894,6 +1021,9 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
# Reset provider options
|
||||
build_config = self.reset_provider_options(build_config)
|
||||
|
||||
# Handle hybrid search options
|
||||
build_config = self._handle_hybrid_search_options(build_config)
|
||||
|
||||
return self.reset_collection_list(build_config)
|
||||
|
||||
def _handle_collection_selection(self, build_config: dict, field_value: str) -> dict:
|
||||
|
|
@ -901,6 +1031,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
build_config["autodetect_collection"]["value"] = True
|
||||
build_config = self.reset_collection_list(build_config)
|
||||
|
||||
# Reset embedding model if collection selection changes
|
||||
if field_value and field_value not in build_config["collection_name"]["options"]:
|
||||
build_config["collection_name"]["options"].append(field_value)
|
||||
build_config["collection_name"]["options_metadata"].append(
|
||||
|
|
@ -916,10 +1047,30 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
if not field_value:
|
||||
return build_config
|
||||
|
||||
# Get the selected collection index
|
||||
index = build_config["collection_name"]["options"].index(field_value)
|
||||
|
||||
# Set the provider of the selected collection
|
||||
provider = build_config["collection_name"]["options_metadata"][index]["provider"]
|
||||
build_config["embedding_model"]["advanced"] = bool(provider)
|
||||
build_config["embedding_choice"]["value"] = "Astra Vectorize" if provider else "Embedding Model"
|
||||
build_config["embedding_model"]["show"] = not bool(provider)
|
||||
build_config["embedding_model"]["required"] = not bool(provider)
|
||||
|
||||
# Grab the collection object
|
||||
database = self.get_database_object(api_endpoint=build_config["api_endpoint"]["value"])
|
||||
collection = database.get_collection(
|
||||
name=field_value,
|
||||
keyspace=build_config["keyspace"]["value"],
|
||||
)
|
||||
|
||||
# Check if hybrid and lexical are enabled
|
||||
col_options = collection.options()
|
||||
hyb_enabled = col_options.rerank and col_options.rerank.enabled
|
||||
lex_enabled = col_options.lexical and col_options.lexical.enabled
|
||||
user_hyb_enabled = build_config["search_method"]["value"] == "Hybrid Search"
|
||||
|
||||
# Show lexical terms if the collection is hybrid enabled
|
||||
build_config["lexical_terms"]["show"] = hyb_enabled and lex_enabled and user_hyb_enabled
|
||||
|
||||
return build_config
|
||||
|
||||
@check_cached_vector_store
|
||||
|
|
@ -934,11 +1085,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
raise ImportError(msg) from e
|
||||
|
||||
# Get the embedding model and additional params
|
||||
embedding_params = (
|
||||
{"embedding": self.embedding_model}
|
||||
if self.embedding_model and self.embedding_choice == "Embedding Model"
|
||||
else {}
|
||||
)
|
||||
embedding_params = {"embedding": self.embedding_model} if self.embedding_model else {}
|
||||
|
||||
# Get the additional parameters
|
||||
additional_params = self.astradb_vectorstore_kwargs or {}
|
||||
|
|
@ -969,6 +1116,9 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
"ignore_invalid_documents": self.ignore_invalid_documents,
|
||||
}
|
||||
|
||||
# Choose HybridSearchMode based on the selected param
|
||||
hybrid_search_mode = HybridSearchMode.DEFAULT if self.search_method == "Hybrid Search" else HybridSearchMode.OFF
|
||||
|
||||
# Attempt to build the Vector Store object
|
||||
try:
|
||||
vector_store = AstraDBVectorStore(
|
||||
|
|
@ -978,6 +1128,8 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
namespace=database.keyspace,
|
||||
collection_name=self.collection_name,
|
||||
environment=self.environment,
|
||||
# Hybrid Search Parameters
|
||||
hybrid_search=hybrid_search_mode,
|
||||
# Astra DB Usage Tracking Parameters
|
||||
ext_callers=[(f"{langflow_prefix}langflow", __version__)],
|
||||
# Astra DB Vector Store Parameters
|
||||
|
|
@ -1036,14 +1188,18 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
return search_type_mapping.get(self.search_type, "similarity")
|
||||
|
||||
def _build_search_args(self):
|
||||
# Clean up the search query
|
||||
query = self.search_query if isinstance(self.search_query, str) and self.search_query.strip() else None
|
||||
lexical_terms = self.lexical_terms or None
|
||||
|
||||
# Check if we have a search query, and if so set the args
|
||||
if query:
|
||||
args = {
|
||||
"query": query,
|
||||
"search_type": self._map_search_type(),
|
||||
"k": self.number_of_results,
|
||||
"score_threshold": self.search_score_threshold,
|
||||
"lexical_query": lexical_terms,
|
||||
}
|
||||
elif self.advanced_search_filter:
|
||||
args = {
|
||||
|
|
@ -1064,6 +1220,9 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
|
|||
self.log(f"Search input: {self.search_query}")
|
||||
self.log(f"Search type: {self.search_type}")
|
||||
self.log(f"Number of results: {self.number_of_results}")
|
||||
self.log(f"store.hybrid_search: {vector_store.hybrid_search}")
|
||||
self.log(f"Lexical terms: {self.lexical_terms}")
|
||||
self.log(f"Reranker: {self.reranker}")
|
||||
|
||||
try:
|
||||
search_args = self._build_search_args()
|
||||
|
|
|
|||
|
|
@ -194,16 +194,14 @@ class HCDVectorStoreComponent(LCVectorStoreComponent):
|
|||
if not isinstance(self.embedding, dict):
|
||||
embedding_dict = {"embedding": self.embedding}
|
||||
else:
|
||||
from astrapy.info import CollectionVectorServiceOptions
|
||||
from astrapy.info import VectorServiceOptions
|
||||
|
||||
dict_options = self.embedding.get("collection_vector_service_options", {})
|
||||
dict_options["authentication"] = {
|
||||
k: v for k, v in dict_options.get("authentication", {}).items() if k and v
|
||||
}
|
||||
dict_options["parameters"] = {k: v for k, v in dict_options.get("parameters", {}).items() if k and v}
|
||||
embedding_dict = {
|
||||
"collection_vector_service_options": CollectionVectorServiceOptions.from_dict(dict_options)
|
||||
}
|
||||
embedding_dict = {"collection_vector_service_options": VectorServiceOptions.from_dict(dict_options)}
|
||||
collection_embedding_api_key = self.embedding.get("collection_embedding_api_key")
|
||||
if collection_embedding_api_key:
|
||||
embedding_dict["collection_embedding_api_key"] = collection_embedding_api_key
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
|
|
@ -202,6 +202,18 @@ class DropDownMixin(BaseModel):
|
|||
"""Variable that defines if the user can insert custom values in the dropdown."""
|
||||
dialog_inputs: dict[str, Any] | None = None
|
||||
"""Dictionary of dialog inputs for the field. Default is an empty object."""
|
||||
toggle: bool = False
|
||||
"""Variable that defines if a toggle button is shown."""
|
||||
toggle_value: bool | None = None
|
||||
"""Variable that defines the value of the toggle button. Defaults to None."""
|
||||
|
||||
@field_validator("toggle_value")
|
||||
@classmethod
|
||||
def validate_toggle_value(cls, v):
|
||||
if v is not None and not isinstance(v, bool):
|
||||
msg = "toggle_value must be a boolean or None"
|
||||
raise ValueError(msg)
|
||||
return v
|
||||
|
||||
|
||||
class SortableListMixin(BaseModel):
|
||||
|
|
|
|||
|
|
@ -459,6 +459,8 @@ class DropdownInput(BaseInputMixin, DropDownMixin, MetadataTraceMixin, ToolModeM
|
|||
options_metadata (Optional[list[dict[str, str]]): List of dictionaries with metadata for each option.
|
||||
Default is None.
|
||||
combobox (CoalesceBool): Variable that defines if the user can insert custom values in the dropdown.
|
||||
toggle (CoalesceBool): Variable that defines if a toggle button is shown.
|
||||
toggle_value (CoalesceBool | None): Variable that defines the value of the toggle button. Defaults to None.
|
||||
"""
|
||||
|
||||
field_type: SerializableFieldTypes = FieldTypes.TEXT
|
||||
|
|
@ -466,6 +468,8 @@ class DropdownInput(BaseInputMixin, DropDownMixin, MetadataTraceMixin, ToolModeM
|
|||
options_metadata: list[dict[str, Any]] = Field(default_factory=list)
|
||||
combobox: CoalesceBool = False
|
||||
dialog_inputs: dict[str, Any] = Field(default_factory=dict)
|
||||
toggle: bool = False
|
||||
toggle_value: bool | None = None
|
||||
|
||||
|
||||
class ConnectionInput(BaseInputMixin, ConnectionMixin, MetadataTraceMixin, ToolModeMixin):
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ _convert_field_type_to_type: dict[FieldTypes, type] = {
|
|||
FieldTypes.CODE: str,
|
||||
FieldTypes.OTHER: str,
|
||||
FieldTypes.TAB: str,
|
||||
FieldTypes.QUERY: str,
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ DIRECT_TYPES = [
|
|||
"sortableList",
|
||||
"auth",
|
||||
"connect",
|
||||
"query",
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import os
|
|||
|
||||
import pytest
|
||||
from astrapy import DataAPIClient
|
||||
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
|
||||
from langchain_astradb import AstraDBVectorStore, VectorServiceOptions
|
||||
from langchain_core.documents import Document
|
||||
from langflow.components.embeddings import OpenAIEmbeddingsComponent
|
||||
from langflow.components.vectorstores import AstraDBVectorStoreComponent
|
||||
|
|
@ -30,8 +30,8 @@ ALL_COLLECTIONS = [
|
|||
|
||||
@pytest.fixture
|
||||
def astradb_client():
|
||||
api_client = DataAPIClient(token=get_astradb_application_token())
|
||||
client = api_client.get_database(get_astradb_api_endpoint())
|
||||
api_client = DataAPIClient()
|
||||
client = api_client.get_database(get_astradb_api_endpoint(), token=get_astradb_application_token())
|
||||
|
||||
yield client # Provide the client to the test functions
|
||||
|
||||
|
|
@ -106,7 +106,7 @@ def test_astra_vectorize():
|
|||
collection_name=VECTORIZE_COLLECTION,
|
||||
api_endpoint=api_endpoint,
|
||||
token=application_token,
|
||||
collection_vector_service_options=CollectionVectorServiceOptions.from_dict(options),
|
||||
collection_vector_service_options=VectorServiceOptions._from_dict(options),
|
||||
)
|
||||
|
||||
documents = [Document(page_content="test1"), Document(page_content="test2")]
|
||||
|
|
@ -150,7 +150,7 @@ def test_astra_vectorize_with_provider_api_key():
|
|||
collection_name=VECTORIZE_COLLECTION_OPENAI,
|
||||
api_endpoint=api_endpoint,
|
||||
token=application_token,
|
||||
collection_vector_service_options=CollectionVectorServiceOptions.from_dict(options),
|
||||
collection_vector_service_options=VectorServiceOptions._from_dict(options),
|
||||
collection_embedding_api_key=os.getenv("OPENAI_API_KEY"),
|
||||
)
|
||||
documents = [Document(page_content="test1"), Document(page_content="test2")]
|
||||
|
|
@ -195,7 +195,7 @@ def test_astra_vectorize_passes_authentication():
|
|||
collection_name=VECTORIZE_COLLECTION_OPENAI_WITH_AUTH,
|
||||
api_endpoint=api_endpoint,
|
||||
token=application_token,
|
||||
collection_vector_service_options=CollectionVectorServiceOptions.from_dict(options),
|
||||
collection_vector_service_options=VectorServiceOptions._from_dict(options),
|
||||
)
|
||||
|
||||
documents = [Document(page_content="test1"), Document(page_content="test2")]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue