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:
Eric Hare 2025-04-11 11:03:34 -07:00 • committed by GitHub
commit 907a594428
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
30 changed files with 4901 additions and 1804 deletions

View file

@ -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(

View file

@ -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,

View file

@ -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:

View file

@ -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()

View file

@ -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

View file

@ -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):

View file

@ -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):

View file

@ -18,6 +18,7 @@ _convert_field_type_to_type: dict[FieldTypes, type] = {
FieldTypes.CODE: str,
FieldTypes.OTHER: str,
FieldTypes.TAB: str,
FieldTypes.QUERY: str,
}

View file

@ -68,6 +68,7 @@ DIRECT_TYPES = [
"sortableList",
"auth",
"connect",
"query",
]

View file

@ -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")]