feat: Generic Callback Dialog Input for Custom Component (#6236)

* force dialog

* Reimplement backend dialog

* Update astradb.py

* Clean up dropdown options

* Remove unused import

* [autofix.ci] apply automated fixes

* Update astradb.py

* Ruff fixes

* Update Vector Store RAG.json

* [autofix.ci] apply automated fixes

* fix: Conditionally render custom option dialog in dropdown

* ✨ (NodeDialogComponent/index.tsx): Add support for passing 'name' prop to NodeDialog component to improve customization and flexibility
📝 (NodeDialogComponent/index.tsx): Update comments and remove unused import to improve code readability and maintainability
🔧 (dropdownComponent/index.tsx): Pass 'name' prop to Dropdown component to enhance customization and flexibility

* ✨ Refactor NodeDialog component to improve state management and payload handling

* Update astradb.py

* [autofix.ci] apply automated fixes

* ✨ Enhance NodeDialog and Dropdown components with improved payload handling and type safety

* Add DB creation functionality

* First version of create

* Update astradb.py

* Fix ruff errors

* Update Vector Store RAG.json

* [autofix.ci] apply automated fixes

* Update astradb.py

* [autofix.ci] apply automated fixes

* Update astradb.py

* [autofix.ci] apply automated fixes

* Update astradb.py

* Update astradb.py

* Update astradb.py

* Update Vector Store RAG.json

* [autofix.ci] apply automated fixes

* Update astradb.py

* [autofix.ci] apply automated fixes

* feat: Enhance dropdown and node dialog with loading states and improved UX

* refactor: Improve error handling in NodeDialog component

* refactor: Update default excluded keys in dropdown metadata filter

* [autofix.ci] apply automated fixes

* refactor: Update Vector Store RAG starter project JSON with formatting and connection ID corrections

* Hide fields that aren't relevant yet

* [autofix.ci] apply automated fixes

* Update Vector Store RAG.json

* [autofix.ci] apply automated fixes

* Update astradb.py

* feat: Improve dropdown component with loading states and enhanced UX

* Update astradb.py

* [autofix.ci] apply automated fixes

* Update astradb.py

* Simon feedback

* [autofix.ci] apply automated fixes

* feat: Enhance dropdown and UI components with status indicators and loading states

* refactor: Update dropdown metadata filtering to exclude 'icon' key

* fix: Conditionally render dropdown icon when available

* fix: Improve dropdown icon rendering with null checks

* chore: Remove debug console log in dropdown component

* Add support for icons in the dropdowns

* Update astradb.py

* Update Vector Store RAG.json

* [autofix.ci] apply automated fixes

* feat: Enhance dropdown status display and color handling

* feat: Add auto-close functionality to node dialog and expand status color handling

* feat: Add real-time template refresh for node dialog fields

* refactor: Improve node dialog component state management and naming

* Async for create collection

* [autofix.ci] apply automated fixes

* Dynamic provider list generation

* Update astradb.py

* [autofix.ci] apply automated fixes

* [autofix.ci] apply automated fixes (attempt 2/3)

* Update astradb.py

* [autofix.ci] apply automated fixes

---------

Co-authored-by: Eric Hare <ericrhare@gmail.com>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: cristhianzl <cristhian.lousa@gmail.com>
This commit is contained in:
Deon Sanchez 2025-02-14 19:10:12 -07:00 • committed by GitHub
commit c902fb9e11
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1314 additions and 538 deletions

View file

@ -1,8 +1,8 @@
import os
from collections import defaultdict
from dataclasses import dataclass, field
from dataclasses import asdict, dataclass, field
from astrapy import AstraDBAdmin, DataAPIClient, Database
from astrapy.info import CollectionDescriptor
from langchain_astradb import AstraDBVectorStore, CollectionVectorServiceOptions
from langflow.base.vectorstores.model import LCVectorStoreComponent, check_cached_vector_store
@ -36,22 +36,24 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
default_factory=lambda: {
"data": {
"node": {
"description": "Create a new database in Astra DB.",
"display_name": "Create New Database",
"name": "create_database",
"description": "",
"display_name": "Create new database",
"field_order": ["new_database_name", "cloud_provider", "region"],
"template": {
"new_database_name": StrInput(
name="new_database_name",
display_name="New Database Name",
display_name="Name",
info="Name of the new database to create in Astra DB.",
required=True,
),
"cloud_provider": DropdownInput(
name="cloud_provider",
display_name="Cloud Provider",
display_name="Cloud provider",
info="Cloud provider for the new database.",
options=["Amazon Web Services", "Google Cloud Platform", "Microsoft Azure"],
required=True,
real_time_refresh=True,
),
"region": DropdownInput(
name="region",
@ -73,8 +75,9 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
default_factory=lambda: {
"data": {
"node": {
"description": "Create a new collection in Astra DB.",
"display_name": "Create New Collection",
"name": "create_collection",
"description": "",
"display_name": "Create new collection",
"field_order": [
"new_collection_name",
"embedding_generation_provider",
@ -83,23 +86,31 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
"template": {
"new_collection_name": StrInput(
name="new_collection_name",
display_name="New Collection Name",
display_name="Name",
info="Name of the new collection to create in Astra DB.",
required=True,
),
"embedding_generation_provider": DropdownInput(
name="embedding_generation_provider",
display_name="Embedding Generation Provider",
display_name="Embedding generation method",
info="Provider to use for generating embeddings.",
options=[],
real_time_refresh=True,
required=True,
options=["Bring your own", "Nvidia"],
),
"embedding_generation_model": DropdownInput(
name="embedding_generation_model",
display_name="Embedding Generation Model",
display_name="Embedding model",
info="Model to use for generating embeddings.",
options=[],
required=True,
options=[],
),
"dimension": IntInput(
name="dimension",
display_name="Dimensions (Required only for `Bring your own`)",
info="Dimensions of the embeddings to generate.",
required=False,
value=1024,
),
},
},
@ -125,17 +136,18 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
real_time_refresh=True,
),
DropdownInput(
name="api_endpoint",
name="database_name",
display_name="Database",
info="The Database / API Endpoint for the Astra DB instance.",
info="The Database name for the Astra DB instance.",
required=True,
refresh_button=True,
real_time_refresh=True,
dialog_inputs=asdict(NewDatabaseInput()),
combobox=True,
),
StrInput(
name="d_api_endpoint",
display_name="Database API Endpoint",
name="api_endpoint",
display_name="Astra DB API Endpoint",
info="The API Endpoint for the Astra DB instance. Supercedes database selection.",
advanced=True,
),
@ -146,8 +158,9 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
required=True,
refresh_button=True,
real_time_refresh=True,
# dialog_inputs=asdict(NewCollectionInput()),
dialog_inputs=asdict(NewCollectionInput()),
combobox=True,
advanced=True,
),
StrInput(
name="keyspace",
@ -238,6 +251,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
@classmethod
def map_cloud_providers(cls):
# TODO: Programmatically fetch the regions for each cloud provider
return {
"Amazon Web Services": {
"id": "aws",
@ -254,54 +268,87 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
}
@classmethod
def create_database_api(
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)
# Get the list of embedding providers
embedding_providers = db_admin.find_embedding_providers().as_dict()
vectorize_providers_mapping = {}
# Map the provider display name to the provider key and models
for provider_key, provider_data in embedding_providers["embeddingProviders"].items():
# Get the provider display name and models
display_name = provider_data["displayName"]
models = [model["name"] for model in provider_data["models"]]
# Build our mapping
vectorize_providers_mapping[display_name] = [provider_key, models]
# Sort the resulting dictionary
return defaultdict(list, dict(sorted(vectorize_providers_mapping.items())))
except Exception as e:
msg = f"Error fetching vectorize providers: {e}"
raise ValueError(msg) from e
@classmethod
async def create_database_api(
cls,
token: str,
new_database_name: str,
cloud_provider: str,
region: str,
token: str,
environment: str | None = None,
keyspace: str | None = None,
):
client = DataAPIClient(token=token)
client = DataAPIClient(token=token, environment=environment)
# Get the admin object
admin_client = client.get_admin(token=token)
# Call the create database function
return admin_client.create_database(
return await admin_client.async_create_database(
name=new_database_name,
cloud_provider=cloud_provider,
cloud_provider=cls.map_cloud_providers()[cloud_provider]["id"],
region=region,
keyspace=keyspace,
wait_until_active=False,
)
@classmethod
def create_collection_api(
async def create_collection_api(
cls,
token: str,
database_name: str,
new_collection_name: str,
token: str,
api_endpoint: str,
environment: str | None = None,
keyspace: str | None = None,
dimension: int | None = None,
embedding_generation_provider: str | None = None,
embedding_generation_model: str | None = None,
):
# Create the data API client
client = DataAPIClient(token=token)
api_endpoint = cls.get_api_endpoint_static(token=token, database_name=database_name)
# Get the database object
database = client.get_database(api_endpoint=api_endpoint, token=token)
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=embedding_generation_provider,
provider=cls.get_vectorize_providers(
token=token, environment=environment, api_endpoint=api_endpoint
).get(embedding_generation_provider, [None, []])[0],
model_name=embedding_generation_model,
authentication=None,
parameters=None,
)
# Create the collection
return database.create_collection(
return await database.create_collection(
name=new_collection_name,
keyspace=keyspace,
dimension=dimension,
service=vectorize_options,
)
@ -325,16 +372,28 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
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"
db_info_dict[db.info.name] = {
"api_endpoint": api_endpoint,
"collections": len(
# Get the number of collections
try:
num_collections = len(
list(
client.get_database(
api_endpoint=api_endpoint, token=token, keyspace=db.info.keyspace
).list_collection_names(keyspace=db.info.keyspace)
)
),
)
except Exception: # noqa: BLE001
num_collections = 0
if db.status != "PENDING":
continue
# Add the database to the dictionary
db_info_dict[db.info.name] = {
"api_endpoint": api_endpoint,
"collections": num_collections,
"status": db.status if db.status != "ACTIVE" else None,
}
except Exception: # noqa: BLE001, S110
pass
@ -364,15 +423,20 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
if not database_name:
return None
# Otherwise, get the URL from the database list
return cls.get_database_list_static(token=token, environment=environment).get(database_name).get("api_endpoint")
# Grab the database object
db = cls.get_database_list_static(token=token, environment=environment).get(database_name)
if not db:
return None
def get_api_endpoint(self, *, api_endpoint: str | None = None):
# Otherwise, get the URL from the database list
return db.get("api_endpoint")
def get_api_endpoint(self):
return self.get_api_endpoint_static(
token=self.token,
environment=self.environment,
api_endpoint=api_endpoint or self.d_api_endpoint,
database_name=self.api_endpoint,
api_endpoint=self.api_endpoint,
database_name=self.database_name,
)
def get_keyspace(self):
@ -388,7 +452,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
client = DataAPIClient(token=self.token, environment=self.environment)
return client.get_database(
api_endpoint=self.get_api_endpoint(api_endpoint=api_endpoint),
api_endpoint=api_endpoint or self.get_api_endpoint(),
token=self.token,
keyspace=self.get_keyspace(),
)
@ -415,40 +479,15 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
return None
def get_vectorize_providers(self):
try:
self.log("Dynamically updating list of Vectorize providers.")
# Get the admin object
admin = AstraDBAdmin(token=self.token)
db_admin = admin.get_database_admin(api_endpoint=self.get_api_endpoint())
# Get the list of embedding providers
embedding_providers = db_admin.find_embedding_providers().as_dict()
vectorize_providers_mapping = {}
# Map the provider display name to the provider key and models
for provider_key, provider_data in embedding_providers["embeddingProviders"].items():
display_name = provider_data["displayName"]
models = [model["name"] for model in provider_data["models"]]
# TODO: https://astra.datastax.com/api/v2/graphql
vectorize_providers_mapping[display_name] = [provider_key, models]
# Sort the resulting dictionary
return defaultdict(list, dict(sorted(vectorize_providers_mapping.items())))
except Exception as e: # noqa: BLE001
self.log(f"Error fetching Vectorize providers: {e}")
return {}
def _initialize_database_options(self):
try:
return [
{
"name": name,
"status": info["status"],
"collections": info["collections"],
"api_endpoint": info["api_endpoint"],
"icon": "data",
}
for name, info in self.get_database_list().items()
]
@ -456,7 +495,35 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
msg = f"Error fetching database options: {e}"
raise ValueError(msg) from e
@classmethod
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
else None
)
# If there is no provider, use the vector store icon
if not provider_name or provider_name == "bring your own":
return "vectorstores"
# Special case for certain models
# TODO: Add more icons
if provider_name == "nvidia":
return "NVIDIA"
if provider_name == "openai":
return "OpenAI"
# Title case on the provider for the icon if no special case
return provider_name.title()
def _initialize_collection_options(self, api_endpoint: str | None = None):
# Nothing to generate if we don't have an API endpoint yet
api_endpoint = api_endpoint or self.get_api_endpoint()
if not api_endpoint:
return []
# Retrieve the database object
database = self.get_database_object(api_endpoint=api_endpoint)
@ -471,7 +538,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
"provider": (
col.options.vector.service.provider if col.options.vector and col.options.vector.service else None
),
"icon": "",
"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
),
@ -479,9 +546,53 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
for col in collection_list
]
def reset_provider_options(self, build_config: dict):
# Get the list of vectorize providers
vectorize_providers = self.get_vectorize_providers(
token=self.token,
environment=self.environment,
api_endpoint=build_config["api_endpoint"]["value"],
)
# If the collection is set, allow user to see embedding options
build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["options"] = ["Bring your own", "Nvidia", *[key for key in vectorize_providers if key != "Nvidia"]]
# For all not Bring your own or Nvidia providers, add metadata saying configure in Astra DB Portal
provider_options = build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["options"]
# Go over each possible provider and add metadata to configure in Astra DB Portal
for provider in provider_options:
# Skip Bring your own and Nvidia, automatically configured
if provider in ["Bring your own", "Nvidia"]:
build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["options_metadata"].append({"icon": self.get_provider_icon(provider_name=provider.lower())})
continue
# Add metadata to configure in Astra DB Portal
build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["options_metadata"].append({" ": "Configure in Astra DB Portal"})
# And allow the user to see the models based on a selected provider
embedding_provider = build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["value"]
# Set the options for the embedding model based on the provider
build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_model"
]["options"] = vectorize_providers.get(embedding_provider, [[], []])[1]
return build_config
def reset_collection_list(self, build_config: dict):
# Get the list of options we have based on the token provided
collection_options = self._initialize_collection_options()
collection_options = self._initialize_collection_options(api_endpoint=build_config["api_endpoint"]["value"])
# If we retrieved options based on the token, show the dropdown
build_config["collection_name"]["options"] = [col["name"] for col in collection_options]
@ -490,7 +601,11 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
]
# Reset the selected collection
build_config["collection_name"]["value"] = ""
if build_config["collection_name"]["value"] not in build_config["collection_name"]["options"]:
build_config["collection_name"]["value"] = ""
# If we have a database, collection name should not be advanced
build_config["collection_name"]["advanced"] = not build_config["database_name"]["value"]
return build_config
@ -499,84 +614,171 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
database_options = self._initialize_database_options()
# If we retrieved options based on the token, show the dropdown
build_config["api_endpoint"]["options"] = [db["name"] for db in database_options]
build_config["api_endpoint"]["options_metadata"] = [
build_config["database_name"]["options"] = [db["name"] for db in database_options]
build_config["database_name"]["options_metadata"] = [
{k: v for k, v in db.items() if k not in ["name"]} for db in database_options
]
# Reset the selected database
build_config["api_endpoint"]["value"] = ""
if build_config["database_name"]["value"] not in build_config["database_name"]["options"]:
build_config["database_name"]["value"] = ""
build_config["api_endpoint"]["value"] = ""
build_config["collection_name"]["advanced"] = True
# If we have a token, database name should not be advanced
build_config["database_name"]["advanced"] = not build_config["token"]["value"]
return build_config
def reset_build_config(self, build_config: dict):
# Reset the list of databases we have based on the token provided
build_config["api_endpoint"]["options"] = []
build_config["api_endpoint"]["options_metadata"] = []
build_config["database_name"]["options"] = []
build_config["database_name"]["options_metadata"] = []
build_config["database_name"]["value"] = ""
build_config["database_name"]["advanced"] = True
build_config["api_endpoint"]["value"] = ""
build_config["api_endpoint"]["name"] = "Database"
# Reset the list of collections and metadata associated
build_config["collection_name"]["options"] = []
build_config["collection_name"]["options_metadata"] = []
build_config["collection_name"]["value"] = ""
build_config["collection_name"]["advanced"] = True
return build_config
def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None):
# When the component first executes, this is the update refresh call
first_run = field_name == "collection_name" and not field_value and not build_config["api_endpoint"]["options"]
async def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None):
# Callback for database creation
if field_name == "database_name" and isinstance(field_value, dict) and "new_database_name" in field_value:
try:
await self.create_database_api(
new_database_name=field_value["new_database_name"],
token=self.token,
keyspace=self.get_keyspace(),
environment=self.environment,
cloud_provider=field_value["cloud_provider"],
region=field_value["region"],
)
except Exception as e:
msg = f"Error creating database: {e}"
raise ValueError(msg) from e
# If the token has not been provided, simply return
# Add the new database to the list of options
build_config["database_name"]["options"] = build_config["database_name"]["options"] + [
field_value["new_database_name"]
]
build_config["database_name"]["options_metadata"] = build_config["database_name"]["options_metadata"] + [
{"status": "PENDING"}
]
return self.reset_collection_list(build_config)
# This is the callback required to update the list of regions for a cloud provider
if field_name == "database_name" and isinstance(field_value, dict) and "new_database_name" not in field_value:
cloud_provider = field_value["cloud_provider"]
build_config["database_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"]["region"][
"options"
] = self.map_cloud_providers()[cloud_provider]["regions"]
return build_config
# Callback for the creation of collections
if field_name == "collection_name" and isinstance(field_value, dict) and "new_collection_name" in field_value:
try:
# Get the dimension if its a BYO provider
dimension = (
field_value["dimension"]
if field_value["embedding_generation_provider"] == "Bring your own"
else None
)
# Create the collection
await self.create_collection_api(
new_collection_name=field_value["new_collection_name"],
token=self.token,
api_endpoint=build_config["api_endpoint"]["value"],
environment=self.environment,
keyspace=self.get_keyspace(),
dimension=dimension,
embedding_generation_provider=field_value["embedding_generation_provider"],
embedding_generation_model=field_value["embedding_generation_model"],
)
except Exception as e:
msg = f"Error creating collection: {e}"
raise ValueError(msg) from e
# Add the new collection to the list of options
build_config["collection_name"]["value"] = field_value["new_collection_name"]
build_config["collection_name"]["options"].append(field_value["new_collection_name"])
# Get the provider and model for the new collection
generation_provider = field_value["embedding_generation_provider"]
provider = generation_provider if generation_provider != "Bring your own" else None
generation_model = field_value["embedding_generation_model"]
model = generation_model if generation_model else None
# Add the new collection to the list of options
icon = "NVIDIA" if provider == "Nvidia" else "vectorstores"
build_config["collection_name"]["options_metadata"] = build_config["collection_name"][
"options_metadata"
] + [{"records": 0, "provider": provider, "icon": icon, "model": model}]
return build_config
# Callback to update the model list based on the embedding provider
if (
field_name == "collection_name"
and isinstance(field_value, dict)
and "new_collection_name" not in field_value
):
return self.reset_provider_options(build_config)
# When the component first executes, this is the update refresh call
first_run = field_name == "collection_name" and not field_value and not build_config["database_name"]["options"]
# If the token has not been provided, simply return the empty build config
if not self.token:
return self.reset_build_config(build_config)
# If this is the first execution of the component, reset and build database list
if first_run or field_name in ["token", "environment"]:
# Reset the build config to ensure we are starting fresh
build_config = self.reset_build_config(build_config)
build_config = self.reset_database_list(build_config)
# Get list of regions for a given cloud provider
"""
cloud_provider = (
build_config["api_endpoint"]["dialog_inputs"]["fields"]["data"]["node"]["template"]["cloud_provider"][
"value"
]
or "Amazon Web Services"
)
build_config["api_endpoint"]["dialog_inputs"]["fields"]["data"]["node"]["template"]["region"][
"options"
] = self.map_cloud_providers()[cloud_provider]["regions"]
"""
return build_config
return self.reset_database_list(build_config)
# Refresh the collection name options
if field_name == "api_endpoint":
if field_name == "database_name" and not isinstance(field_value, dict):
# If missing, refresh the database options
if not build_config["api_endpoint"]["options"] or not field_value:
return self.update_build_config(build_config, field_value=self.token, field_name="token")
if field_value not in build_config["database_name"]["options"]:
build_config = await self.update_build_config(build_config, field_value=self.token, field_name="token")
build_config["database_name"]["value"] = ""
else:
# Find the position of the selected database to align with metadata
index_of_name = build_config["database_name"]["options"].index(field_value)
# Set the underlying api endpoint value of the database
if field_value in build_config["api_endpoint"]["options"]:
index_of_name = build_config["api_endpoint"]["options"].index(field_value)
build_config["d_api_endpoint"]["value"] = build_config["api_endpoint"]["options_metadata"][
# Initializing database condition
pending = build_config["database_name"]["options_metadata"][index_of_name]["status"] == "PENDING"
if pending:
return self.update_build_config(build_config, field_value=self.token, field_name="token")
# Set the API endpoint based on the selected database
build_config["api_endpoint"]["value"] = build_config["database_name"]["options_metadata"][
index_of_name
]["api_endpoint"]
else:
build_config["d_api_endpoint"]["value"] = ""
# Reset the provider options
build_config = self.reset_provider_options(build_config)
# Reset the list of collections we have based on the token provided
return self.reset_collection_list(build_config)
# Hide embedding model option if opriona_metadata provider is not null
if field_name == "collection_name" and field_value:
if field_name == "collection_name" and not isinstance(field_value, dict):
# Assume we will be autodetecting the collection:
build_config["autodetect_collection"]["value"] = True
# Reload the collection list
build_config = self.reset_collection_list(build_config)
# Set the options for collection name to be the field value if its a new collection
if field_value not in build_config["collection_name"]["options"]:
if field_value and field_value not in build_config["collection_name"]["options"]:
# Add the new collection to the list of options
build_config["collection_name"]["options"].append(field_value)
build_config["collection_name"]["options_metadata"].append(
@ -598,36 +800,8 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
build_config["embedding_model"]["advanced"] = False
build_config["embedding_choice"]["value"] = "Embedding Model"
# For the final step, get the list of vectorize providers
"""
vectorize_providers = self.get_vectorize_providers()
if not vectorize_providers:
return build_config
# Allow the user to see the embedding provider options
provider_options = build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["options"]
if not provider_options:
# If the collection is set, allow user to see embedding options
build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["options"] = ["Bring your own", "Nvidia", *[key for key in vectorize_providers if key != "Nvidia"]]
# And allow the user to see the models based on a selected provider
model_options = build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_model"
]["options"]
if not model_options:
embedding_provider = build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_provider"
]["value"]
build_config["collection_name"]["dialog_inputs"]["fields"]["data"]["node"]["template"][
"embedding_generation_model"
]["options"] = vectorize_providers.get(embedding_provider, [[], []])[1]
"""
return build_config
@check_cached_vector_store
@ -654,11 +828,11 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
# Get Langflow version and platform information
__version__ = get_version_info()["version"]
langflow_prefix = ""
if os.getenv("AWS_EXECUTION_ENV") == "AWS_ECS_FARGATE": # TODO: More precise way of detecting
langflow_prefix = "ds-"
# if os.getenv("AWS_EXECUTION_ENV") == "AWS_ECS_FARGATE": # TODO: More precise way of detecting
# langflow_prefix = "ds-"
# Get the database object
database = self.get_database_object(api_endpoint=self.d_api_endpoint)
database = self.get_database_object()
autodetect = self.collection_name in database.list_collection_names() and self.autodetect_collection
# Bundle up the auto-detect parameters
@ -714,7 +888,7 @@ class AstraDBVectorStoreComponent(LCVectorStoreComponent):
if documents and self.deletion_field:
self.log(f"Deleting documents where {self.deletion_field}")
try:
database = self.get_database_object(api_endpoint=self.d_api_endpoint)
database = self.get_database_object()
collection = database.get_collection(self.collection_name, keyspace=database.keyspace)
delete_values = list({doc.metadata[self.deletion_field] for doc in documents})
self.log(f"Deleting documents where {self.deletion_field} matches {delete_values}.")

File diff suppressed because one or more lines are too long