diff --git a/src/backend/base/langflow/initial_setup/starter_projects/Knowledge Retrieval.json b/src/backend/base/langflow/initial_setup/starter_projects/Knowledge Retrieval.json index da41c0a688..987ed57d59 100644 --- a/src/backend/base/langflow/initial_setup/starter_projects/Knowledge Retrieval.json +++ b/src/backend/base/langflow/initial_setup/starter_projects/Knowledge Retrieval.json @@ -531,7 +531,7 @@ "last_updated": "2025-08-26T16:19:16.681Z", "legacy": false, "metadata": { - "code_hash": "f6b6a4eaf1c4", + "code_hash": "8b5ca1f38f6e", "dependencies": { "dependencies": [ { @@ -641,7 +641,7 @@ "show": true, "title_case": false, "type": "code", - "value": "import json\nfrom pathlib import Path\nfrom typing import Any\n\nimport chromadb\nimport chromadb.api.client\nfrom cryptography.fernet import InvalidToken\nfrom langchain_chroma import Chroma\nfrom langflow.services.auth.utils import decrypt_api_key\nfrom langflow.services.database.models.user.crud import get_user_by_id\nfrom pydantic import SecretStr\n\nfrom lfx.base.knowledge_bases.knowledge_base_utils import get_knowledge_bases\nfrom lfx.base.models.unified_models import (\n get_model_provider_variable_mapping,\n get_provider_all_variables,\n)\nfrom lfx.custom import Component\nfrom lfx.io import BoolInput, DropdownInput, IntInput, MessageTextInput, Output, SecretStrInput\nfrom lfx.log.logger import logger\nfrom lfx.schema.data import Data\nfrom lfx.schema.dataframe import DataFrame\nfrom lfx.services.deps import get_settings_service, get_variable_service, session_scope\nfrom lfx.utils.validate_cloud import raise_error_if_astra_cloud_disable_component\n\n_KNOWLEDGE_BASES_ROOT_PATH: Path | None = None\n\n# Error message to raise if we're in Astra cloud environment and the component is not supported.\nastra_error_msg = \"Knowledge retrieval is not supported in Astra cloud environment.\"\n\n\ndef _get_knowledge_bases_root_path() -> Path:\n \"\"\"Lazy load the knowledge bases root path from settings.\"\"\"\n global _KNOWLEDGE_BASES_ROOT_PATH # noqa: PLW0603\n if _KNOWLEDGE_BASES_ROOT_PATH is None:\n settings = get_settings_service().settings\n knowledge_directory = settings.knowledge_bases_dir\n if not knowledge_directory:\n msg = \"Knowledge bases directory is not set in the settings.\"\n raise ValueError(msg)\n _KNOWLEDGE_BASES_ROOT_PATH = Path(knowledge_directory).expanduser()\n return _KNOWLEDGE_BASES_ROOT_PATH\n\n\nclass KnowledgeBaseComponent(Component):\n display_name = \"Knowledge Base\"\n description = \"Search and retrieve data from knowledge.\"\n icon = \"download\"\n name = \"KnowledgeBase\"\n\n inputs = [\n DropdownInput(\n name=\"knowledge_base\",\n display_name=\"Knowledge\",\n info=\"Select the knowledge to load data from.\",\n required=True,\n options=[],\n refresh_button=True,\n real_time_refresh=True,\n ),\n SecretStrInput(\n name=\"api_key\",\n display_name=\"Embedding Provider API Key\",\n info=\"API key for the embedding provider to generate embeddings.\",\n advanced=True,\n required=False,\n ),\n MessageTextInput(\n name=\"search_query\",\n display_name=\"Search Query\",\n info=\"Optional search query to filter knowledge base data.\",\n tool_mode=True,\n ),\n IntInput(\n name=\"top_k\",\n display_name=\"Top K Results\",\n info=\"Number of top results to return from the knowledge base.\",\n value=5,\n advanced=True,\n required=False,\n ),\n BoolInput(\n name=\"include_metadata\",\n display_name=\"Include Metadata\",\n info=\"Whether to include all metadata in the output. If false, only content is returned.\",\n value=True,\n advanced=False,\n ),\n BoolInput(\n name=\"include_embeddings\",\n display_name=\"Include Embeddings\",\n info=\"Whether to include embeddings in the output. Only applicable if 'Include Metadata' is enabled.\",\n value=False,\n advanced=True,\n ),\n ]\n\n outputs = [\n Output(\n name=\"retrieve_data\",\n display_name=\"Results\",\n method=\"retrieve_data\",\n info=\"Returns the data from the selected knowledge base.\",\n ),\n ]\n\n async def update_build_config(self, build_config, field_value, field_name=None): # noqa: ARG002\n # Check if we're in Astra cloud environment and raise an error if we are.\n raise_error_if_astra_cloud_disable_component(astra_error_msg)\n if field_name == \"knowledge_base\":\n # Update the knowledge base options dynamically\n build_config[\"knowledge_base\"][\"options\"] = await get_knowledge_bases(\n _get_knowledge_bases_root_path(),\n user_id=self.user_id, # Use the user_id from the component context\n )\n\n # If the selected knowledge base is not available, reset it\n if build_config[\"knowledge_base\"][\"value\"] not in build_config[\"knowledge_base\"][\"options\"]:\n build_config[\"knowledge_base\"][\"value\"] = None\n\n return build_config\n\n def _get_kb_metadata(self, kb_path: Path) -> dict:\n \"\"\"Load and process knowledge base metadata.\"\"\"\n # Check if we're in Astra cloud environment and raise an error if we are.\n raise_error_if_astra_cloud_disable_component(astra_error_msg)\n metadata: dict[str, Any] = {}\n metadata_file = kb_path / \"embedding_metadata.json\"\n if not metadata_file.exists():\n logger.warning(f\"Embedding metadata file not found at {metadata_file}\")\n return metadata\n\n try:\n with metadata_file.open(\"r\", encoding=\"utf-8\") as f:\n metadata = json.load(f)\n except json.JSONDecodeError:\n logger.error(f\"Error decoding JSON from {metadata_file}\")\n return {}\n\n # Decrypt API key if it exists\n if \"api_key\" in metadata and metadata.get(\"api_key\"):\n settings_service = get_settings_service()\n try:\n decrypted_key = decrypt_api_key(metadata[\"api_key\"], settings_service)\n metadata[\"api_key\"] = decrypted_key\n except (InvalidToken, TypeError, ValueError) as e:\n logger.error(f\"Could not decrypt API key. Please provide it manually. Error: {e}\")\n metadata[\"api_key\"] = None\n return metadata\n\n async def _resolve_provider_variables(self, provider: str) -> dict[str, str]:\n \"\"\"Resolve all global variables for a provider using the async session.\n\n This avoids the run_until_complete thread dance by doing the lookup\n directly in the already-running async context.\n \"\"\"\n import os\n\n result: dict[str, str] = {}\n provider_vars = get_provider_all_variables(provider)\n if not provider_vars or not self.user_id:\n return result\n\n async with session_scope() as session:\n variable_service = get_variable_service()\n if variable_service is None:\n return result\n\n user_id = self.user_id if not isinstance(self.user_id, str) else __import__(\"uuid\").UUID(self.user_id)\n for var_info in provider_vars:\n var_key = var_info.get(\"variable_key\")\n if not var_key:\n continue\n try:\n value = await variable_service.get_variable(\n user_id=user_id,\n name=var_key,\n field=\"\",\n session=session,\n )\n if value and str(value).strip():\n result[var_key] = str(value)\n except (ValueError, Exception): # noqa: BLE001\n env_value = os.environ.get(var_key)\n if env_value and env_value.strip():\n result[var_key] = env_value\n return result\n\n async def _resolve_api_key(self, provider: str) -> str | None:\n \"\"\"Resolve the API key for the given provider.\n\n Priority: user override > metadata (decrypted) > global variable.\n \"\"\"\n provider_variable_map = get_model_provider_variable_mapping()\n variable_name = provider_variable_map.get(provider)\n if not variable_name or not self.user_id:\n return None\n\n async with session_scope() as session:\n variable_service = get_variable_service()\n if variable_service is None:\n return None\n try:\n user_id = self.user_id if not isinstance(self.user_id, str) else __import__(\"uuid\").UUID(self.user_id)\n return await variable_service.get_variable(\n user_id=user_id,\n name=variable_name,\n field=\"\",\n session=session,\n )\n except (ValueError, Exception): # noqa: BLE001\n return None\n\n def _build_embeddings(self, metadata: dict, *, api_key: str | None = None, provider_vars: dict | None = None):\n \"\"\"Build embedding model from metadata.\n\n Args:\n metadata: The knowledge base embedding metadata.\n api_key: Pre-resolved API key (user override > metadata > global).\n provider_vars: Pre-resolved provider variables (for Ollama/WatsonX).\n \"\"\"\n provider = metadata.get(\"embedding_provider\")\n model = metadata.get(\"embedding_model\")\n chunk_size = metadata.get(\"chunk_size\")\n\n # Handle various providers\n if provider == \"OpenAI\":\n from langchain_openai import OpenAIEmbeddings\n\n if not api_key:\n msg = (\n \"OpenAI API key is required. Provide it in the component's advanced settings\"\n \" or configure it globally.\"\n )\n raise ValueError(msg)\n openai_kwargs: dict = {\"model\": model, \"api_key\": api_key}\n if chunk_size is not None:\n openai_kwargs[\"chunk_size\"] = chunk_size\n return OpenAIEmbeddings(**openai_kwargs)\n if provider == \"HuggingFace\":\n from langchain_huggingface import HuggingFaceEmbeddings\n\n return HuggingFaceEmbeddings(\n model=model,\n )\n if provider == \"Cohere\":\n from langchain_cohere import CohereEmbeddings\n\n if not api_key:\n msg = \"Cohere API key is required when using Cohere provider\"\n raise ValueError(msg)\n return CohereEmbeddings(\n model=model,\n cohere_api_key=api_key,\n )\n if provider == \"Google Generative AI\":\n from langchain_google_genai import GoogleGenerativeAIEmbeddings\n\n if not api_key:\n msg = (\n \"Google API key is required. Provide it in the component's advanced settings\"\n \" or configure it globally.\"\n )\n raise ValueError(msg)\n return GoogleGenerativeAIEmbeddings(\n model=model,\n google_api_key=api_key,\n )\n if provider == \"Ollama\":\n from langchain_ollama import OllamaEmbeddings\n\n all_vars = provider_vars or {}\n base_url = all_vars.get(\"OLLAMA_BASE_URL\")\n kwargs: dict = {\"model\": model}\n if base_url:\n kwargs[\"base_url\"] = base_url\n return OllamaEmbeddings(**kwargs)\n if provider == \"IBM WatsonX\":\n from langchain_ibm import WatsonxEmbeddings\n\n all_vars = provider_vars or {}\n watsonx_apikey = api_key or all_vars.get(\"WATSONX_APIKEY\")\n watsonx_project_id = all_vars.get(\"WATSONX_PROJECT_ID\")\n watsonx_url = all_vars.get(\"WATSONX_URL\")\n if not watsonx_apikey:\n msg = (\n \"IBM WatsonX API key is required. Provide it in the component's advanced settings\"\n \" or configure it globally.\"\n )\n raise ValueError(msg)\n kwargs = {\"model_id\": model, \"apikey\": watsonx_apikey}\n if watsonx_project_id:\n kwargs[\"project_id\"] = watsonx_project_id\n if watsonx_url:\n kwargs[\"url\"] = watsonx_url\n return WatsonxEmbeddings(**kwargs)\n if provider == \"Custom\":\n # For custom embedding models, we would need additional configuration\n msg = \"Custom embedding models not yet supported\"\n raise NotImplementedError(msg)\n msg = f\"Embedding provider '{provider}' is not supported for retrieval.\"\n raise NotImplementedError(msg)\n\n async def retrieve_data(self) -> DataFrame:\n \"\"\"Retrieve data from the selected knowledge base by reading the Chroma collection.\n\n Returns:\n A DataFrame containing the data rows from the knowledge base.\n \"\"\"\n # Check if we're in Astra cloud environment and raise an error if we are.\n raise_error_if_astra_cloud_disable_component(astra_error_msg)\n # Get the current user\n async with session_scope() as db:\n if not self.user_id:\n msg = \"User ID is required for fetching Knowledge Base data.\"\n raise ValueError(msg)\n current_user = await get_user_by_id(db, self.user_id)\n if not current_user:\n msg = f\"User with ID {self.user_id} not found.\"\n raise ValueError(msg)\n kb_user = current_user.username\n kb_path = _get_knowledge_bases_root_path() / kb_user / self.knowledge_base\n\n metadata = self._get_kb_metadata(kb_path)\n if not metadata:\n msg = f\"Metadata not found for knowledge base: {self.knowledge_base}. Ensure it has been indexed.\"\n raise ValueError(msg)\n\n # Resolve API key: user override > metadata (decrypted) > global variable\n provider = metadata.get(\"embedding_provider\")\n runtime_api_key = self.api_key.get_secret_value() if isinstance(self.api_key, SecretStr) else self.api_key\n api_key = runtime_api_key or metadata.get(\"api_key\")\n if not api_key and provider:\n api_key = await self._resolve_api_key(provider)\n\n # Resolve provider-specific variables (e.g. base_url for Ollama, project_id for WatsonX)\n provider_vars: dict[str, str] = {}\n if provider in {\"Ollama\", \"IBM WatsonX\"}:\n provider_vars = await self._resolve_provider_variables(provider)\n\n # Build the embedder for the knowledge base\n embedding_function = self._build_embeddings(metadata, api_key=api_key, provider_vars=provider_vars)\n\n # Clear Chroma's singleton client cache to avoid \"different settings\"\n # conflicts when ingestion and retrieval run in the same process.\n chromadb.api.client.SharedSystemClient.clear_system_cache()\n chroma = Chroma(\n persist_directory=str(kb_path),\n embedding_function=embedding_function,\n collection_name=self.knowledge_base,\n )\n\n # If a search query is provided, perform a similarity search\n if self.search_query:\n # Use the search query to perform a similarity search\n logger.info(f\"Performing similarity search with query: {self.search_query}\")\n results = chroma.similarity_search_with_score(\n query=self.search_query or \"\",\n k=self.top_k,\n )\n else:\n results = chroma.similarity_search(\n query=self.search_query or \"\",\n k=self.top_k,\n )\n\n # For each result, make it a tuple to match the expected output format\n results = [(doc, 0) for doc in results] # Assign a dummy score of 0\n\n # If include_embeddings is enabled, get embeddings for the results\n id_to_embedding = {}\n if self.include_embeddings and results:\n doc_ids = [doc[0].metadata.get(\"_id\") for doc in results if doc[0].metadata.get(\"_id\")]\n\n # Only proceed if we have valid document IDs\n if doc_ids:\n # Access underlying collection to get embeddings\n collection = chroma._collection # noqa: SLF001\n embeddings_result = collection.get(where={\"_id\": {\"$in\": doc_ids}}, include=[\"metadatas\", \"embeddings\"])\n\n # Create a mapping from document ID to embedding\n for i, metadata in enumerate(embeddings_result.get(\"metadatas\", [])):\n if metadata and \"_id\" in metadata:\n id_to_embedding[metadata[\"_id\"]] = embeddings_result[\"embeddings\"][i]\n\n # Build output data based on include_metadata setting\n data_list = []\n for doc in results:\n kwargs = {\n \"content\": doc[0].page_content,\n }\n if self.search_query:\n kwargs[\"_score\"] = -1 * doc[1]\n if self.include_metadata:\n # Include all metadata, embeddings, and content\n kwargs.update(doc[0].metadata)\n if self.include_embeddings:\n kwargs[\"_embeddings\"] = id_to_embedding.get(doc[0].metadata.get(\"_id\"))\n\n data_list.append(Data(**kwargs))\n\n # Return the DataFrame containing the data\n return DataFrame(data=data_list)\n" + "value": "import json\nimport os\nimport uuid\nfrom pathlib import Path\nfrom typing import Any\n\nimport chromadb\nimport chromadb.api.client\nfrom cryptography.fernet import InvalidToken\nfrom langchain_chroma import Chroma\nfrom langflow.services.auth.utils import decrypt_api_key\nfrom langflow.services.database.models.user.crud import get_user_by_id\nfrom pydantic import SecretStr\n\nfrom lfx.base.knowledge_bases.knowledge_base_utils import get_knowledge_bases\nfrom lfx.base.models.unified_models import (\n get_model_provider_variable_mapping,\n get_provider_all_variables,\n)\nfrom lfx.custom import Component\nfrom lfx.io import BoolInput, DropdownInput, IntInput, MessageTextInput, Output, SecretStrInput\nfrom lfx.log.logger import logger\nfrom lfx.schema.data import Data\nfrom lfx.schema.dataframe import DataFrame\nfrom lfx.services.deps import get_settings_service, get_variable_service, session_scope\nfrom lfx.utils.validate_cloud import raise_error_if_astra_cloud_disable_component\n\n_KNOWLEDGE_BASES_ROOT_PATH: Path | None = None\n\n# Error message to raise if we're in Astra cloud environment and the component is not supported.\nastra_error_msg = \"Knowledge retrieval is not supported in Astra cloud environment.\"\n\n\ndef _get_knowledge_bases_root_path() -> Path:\n \"\"\"Lazy load the knowledge bases root path from settings.\"\"\"\n global _KNOWLEDGE_BASES_ROOT_PATH # noqa: PLW0603\n if _KNOWLEDGE_BASES_ROOT_PATH is None:\n settings = get_settings_service().settings\n knowledge_directory = settings.knowledge_bases_dir\n if not knowledge_directory:\n msg = \"Knowledge bases directory is not set in the settings.\"\n raise ValueError(msg)\n _KNOWLEDGE_BASES_ROOT_PATH = Path(knowledge_directory).expanduser()\n return _KNOWLEDGE_BASES_ROOT_PATH\n\n\nclass KnowledgeBaseComponent(Component):\n display_name = \"Knowledge Base\"\n description = \"Search and retrieve data from knowledge.\"\n icon = \"download\"\n name = \"KnowledgeBase\"\n\n inputs = [\n DropdownInput(\n name=\"knowledge_base\",\n display_name=\"Knowledge\",\n info=\"Select the knowledge to load data from.\",\n required=True,\n options=[],\n refresh_button=True,\n real_time_refresh=True,\n ),\n SecretStrInput(\n name=\"api_key\",\n display_name=\"Embedding Provider API Key\",\n info=\"API key for the embedding provider to generate embeddings.\",\n advanced=True,\n required=False,\n ),\n MessageTextInput(\n name=\"search_query\",\n display_name=\"Search Query\",\n info=\"Optional search query to filter knowledge base data.\",\n tool_mode=True,\n ),\n IntInput(\n name=\"top_k\",\n display_name=\"Top K Results\",\n info=\"Number of top results to return from the knowledge base.\",\n value=5,\n advanced=True,\n required=False,\n ),\n BoolInput(\n name=\"include_metadata\",\n display_name=\"Include Metadata\",\n info=\"Whether to include all metadata in the output. If false, only content is returned.\",\n value=True,\n advanced=False,\n ),\n BoolInput(\n name=\"include_embeddings\",\n display_name=\"Include Embeddings\",\n info=\"Whether to include embeddings in the output. Only applicable if 'Include Metadata' is enabled.\",\n value=False,\n advanced=True,\n ),\n ]\n\n outputs = [\n Output(\n name=\"retrieve_data\",\n display_name=\"Results\",\n method=\"retrieve_data\",\n info=\"Returns the data from the selected knowledge base.\",\n ),\n ]\n\n async def update_build_config(self, build_config, field_value, field_name=None): # noqa: ARG002\n # Check if we're in Astra cloud environment and raise an error if we are.\n raise_error_if_astra_cloud_disable_component(astra_error_msg)\n if field_name == \"knowledge_base\":\n # Update the knowledge base options dynamically\n build_config[\"knowledge_base\"][\"options\"] = await get_knowledge_bases(\n _get_knowledge_bases_root_path(),\n user_id=self.user_id, # Use the user_id from the component context\n )\n\n # If the selected knowledge base is not available, reset it\n if build_config[\"knowledge_base\"][\"value\"] not in build_config[\"knowledge_base\"][\"options\"]:\n build_config[\"knowledge_base\"][\"value\"] = None\n\n return build_config\n\n @property\n def _user_uuid(self) -> uuid.UUID | None:\n \"\"\"Return self.user_id as a UUID, converting from str if necessary.\"\"\"\n if not self.user_id:\n return None\n return self.user_id if isinstance(self.user_id, uuid.UUID) else uuid.UUID(self.user_id)\n\n def _get_kb_metadata(self, kb_path: Path) -> dict:\n \"\"\"Load and process knowledge base metadata.\"\"\"\n # Check if we're in Astra cloud environment and raise an error if we are.\n raise_error_if_astra_cloud_disable_component(astra_error_msg)\n metadata: dict[str, Any] = {}\n metadata_file = kb_path / \"embedding_metadata.json\"\n if not metadata_file.exists():\n logger.warning(f\"Embedding metadata file not found at {metadata_file}\")\n return metadata\n\n try:\n with metadata_file.open(\"r\", encoding=\"utf-8\") as f:\n metadata = json.load(f)\n except json.JSONDecodeError:\n logger.error(f\"Error decoding JSON from {metadata_file}\")\n return {}\n\n # Decrypt API key if it exists\n if \"api_key\" in metadata and metadata.get(\"api_key\"):\n settings_service = get_settings_service()\n try:\n decrypted_key = decrypt_api_key(metadata[\"api_key\"], settings_service)\n metadata[\"api_key\"] = decrypted_key\n except (InvalidToken, TypeError, ValueError) as e:\n logger.error(f\"Could not decrypt API key. Please provide it manually. Error: {e}\")\n metadata[\"api_key\"] = None\n return metadata\n\n async def _resolve_provider_variables(self, provider: str) -> dict[str, str]:\n \"\"\"Resolve all global variables for a provider using the async session.\n\n This avoids the run_until_complete thread dance by doing the lookup\n directly in the already-running async context.\n \"\"\"\n result: dict[str, str] = {}\n provider_vars = get_provider_all_variables(provider)\n user_id = self._user_uuid\n if not provider_vars or not user_id:\n return result\n\n async with session_scope() as session:\n variable_service = get_variable_service()\n if variable_service is None:\n return result\n\n for var_info in provider_vars:\n var_key = var_info.get(\"variable_key\")\n if not var_key:\n continue\n try:\n value = await variable_service.get_variable(\n user_id=user_id,\n name=var_key,\n field=\"\",\n session=session,\n )\n if value and str(value).strip():\n result[var_key] = str(value)\n except (ValueError, KeyError, AttributeError) as e:\n logger.debug(f\"Variable service lookup failed for '{var_key}', falling back to environment: {e}\")\n env_value = os.environ.get(var_key)\n if env_value and env_value.strip():\n result[var_key] = env_value\n return result\n\n async def _resolve_api_key(self, provider: str) -> str | None:\n \"\"\"Resolve the API key for the given provider.\n\n Priority: user override > metadata (decrypted) > global variable.\n \"\"\"\n provider_variable_map = get_model_provider_variable_mapping()\n variable_name = provider_variable_map.get(provider)\n user_id = self._user_uuid\n if not variable_name or not user_id:\n return None\n\n async with session_scope() as session:\n variable_service = get_variable_service()\n if variable_service is None:\n return None\n try:\n return await variable_service.get_variable(\n user_id=user_id,\n name=variable_name,\n field=\"\",\n session=session,\n )\n except (ValueError, KeyError, AttributeError):\n return None\n\n def _build_embeddings(self, metadata: dict, *, api_key: str | None = None, provider_vars: dict | None = None):\n \"\"\"Build embedding model from metadata.\n\n Args:\n metadata: The knowledge base embedding metadata.\n api_key: Pre-resolved API key (user override > metadata > global).\n provider_vars: Pre-resolved provider variables (for Ollama/WatsonX).\n \"\"\"\n provider = metadata.get(\"embedding_provider\")\n model = metadata.get(\"embedding_model\")\n chunk_size = metadata.get(\"chunk_size\")\n\n # Handle various providers\n if provider == \"OpenAI\":\n from langchain_openai import OpenAIEmbeddings\n\n if not api_key:\n msg = (\n \"OpenAI API key is required. Provide it in the component's advanced settings\"\n \" or configure it globally.\"\n )\n raise ValueError(msg)\n openai_kwargs: dict = {\"model\": model, \"api_key\": api_key}\n if chunk_size is not None:\n openai_kwargs[\"chunk_size\"] = chunk_size\n return OpenAIEmbeddings(**openai_kwargs)\n if provider == \"HuggingFace\":\n from langchain_huggingface import HuggingFaceEmbeddings\n\n return HuggingFaceEmbeddings(\n model=model,\n )\n if provider == \"Cohere\":\n from langchain_cohere import CohereEmbeddings\n\n if not api_key:\n msg = \"Cohere API key is required when using Cohere provider\"\n raise ValueError(msg)\n return CohereEmbeddings(\n model=model,\n cohere_api_key=api_key,\n )\n if provider == \"Google Generative AI\":\n from langchain_google_genai import GoogleGenerativeAIEmbeddings\n\n if not api_key:\n msg = (\n \"Google API key is required. Provide it in the component's advanced settings\"\n \" or configure it globally.\"\n )\n raise ValueError(msg)\n return GoogleGenerativeAIEmbeddings(\n model=model,\n google_api_key=api_key,\n )\n if provider == \"Ollama\":\n from langchain_ollama import OllamaEmbeddings\n\n all_vars = provider_vars or {}\n base_url = all_vars.get(\"OLLAMA_BASE_URL\")\n kwargs: dict = {\"model\": model}\n if base_url:\n kwargs[\"base_url\"] = base_url\n return OllamaEmbeddings(**kwargs)\n if provider == \"IBM WatsonX\":\n from langchain_ibm import WatsonxEmbeddings\n\n all_vars = provider_vars or {}\n watsonx_apikey = api_key or all_vars.get(\"WATSONX_APIKEY\")\n watsonx_project_id = all_vars.get(\"WATSONX_PROJECT_ID\")\n watsonx_url = all_vars.get(\"WATSONX_URL\")\n if not watsonx_apikey:\n msg = (\n \"IBM WatsonX API key is required. Provide it in the component's advanced settings\"\n \" or configure it globally.\"\n )\n raise ValueError(msg)\n kwargs = {\"model_id\": model, \"apikey\": watsonx_apikey}\n if watsonx_project_id:\n kwargs[\"project_id\"] = watsonx_project_id\n if watsonx_url:\n kwargs[\"url\"] = watsonx_url\n return WatsonxEmbeddings(**kwargs)\n if provider == \"Custom\":\n # For custom embedding models, we would need additional configuration\n msg = \"Custom embedding models not yet supported\"\n raise NotImplementedError(msg)\n msg = f\"Embedding provider '{provider}' is not supported for retrieval.\"\n raise NotImplementedError(msg)\n\n async def retrieve_data(self) -> DataFrame:\n \"\"\"Retrieve data from the selected knowledge base by reading the Chroma collection.\n\n Returns:\n A DataFrame containing the data rows from the knowledge base.\n \"\"\"\n # Check if we're in Astra cloud environment and raise an error if we are.\n raise_error_if_astra_cloud_disable_component(astra_error_msg)\n # Get the current user\n async with session_scope() as db:\n if not self.user_id:\n msg = \"User ID is required for fetching Knowledge Base data.\"\n raise ValueError(msg)\n current_user = await get_user_by_id(db, self.user_id)\n if not current_user:\n msg = f\"User with ID {self.user_id} not found.\"\n raise ValueError(msg)\n kb_user = current_user.username\n kb_path = _get_knowledge_bases_root_path() / kb_user / self.knowledge_base\n\n metadata = self._get_kb_metadata(kb_path)\n if not metadata:\n msg = f\"Metadata not found for knowledge base: {self.knowledge_base}. Ensure it has been indexed.\"\n raise ValueError(msg)\n\n # Resolve API key: user override > metadata (decrypted) > global variable\n provider = metadata.get(\"embedding_provider\")\n runtime_api_key = self.api_key.get_secret_value() if isinstance(self.api_key, SecretStr) else self.api_key\n api_key = runtime_api_key or metadata.get(\"api_key\")\n if not api_key and provider:\n api_key = await self._resolve_api_key(provider)\n\n # Resolve provider-specific variables (e.g. base_url for Ollama, project_id for WatsonX)\n provider_vars: dict[str, str] = {}\n if provider in {\"Ollama\", \"IBM WatsonX\"}:\n provider_vars = await self._resolve_provider_variables(provider)\n\n # Build the embedder for the knowledge base\n embedding_function = self._build_embeddings(metadata, api_key=api_key, provider_vars=provider_vars)\n\n # Clear Chroma's singleton client cache to avoid \"different settings\"\n # conflicts when ingestion and retrieval run in the same process.\n chromadb.api.client.SharedSystemClient.clear_system_cache()\n chroma = Chroma(\n persist_directory=str(kb_path),\n embedding_function=embedding_function,\n collection_name=self.knowledge_base,\n )\n\n # If a search query is provided, perform a similarity search\n if self.search_query:\n # Use the search query to perform a similarity search\n logger.info(\"Performing similarity search\")\n results = chroma.similarity_search_with_score(\n query=self.search_query or \"\",\n k=self.top_k,\n )\n else:\n results = chroma.similarity_search(\n query=self.search_query or \"\",\n k=self.top_k,\n )\n\n # For each result, make it a tuple to match the expected output format\n results = [(doc, 0) for doc in results] # Assign a dummy score of 0\n\n # If include_embeddings is enabled, get embeddings for the results\n id_to_embedding = {}\n if self.include_embeddings and results:\n doc_ids = [doc[0].metadata.get(\"_id\") for doc in results if doc[0].metadata.get(\"_id\")]\n\n # Only proceed if we have valid document IDs\n if doc_ids:\n # Access underlying collection to get embeddings\n collection = chroma._collection # noqa: SLF001\n embeddings_result = collection.get(where={\"_id\": {\"$in\": doc_ids}}, include=[\"metadatas\", \"embeddings\"])\n\n # Create a mapping from document ID to embedding\n for i, metadata in enumerate(embeddings_result.get(\"metadatas\", [])):\n if metadata and \"_id\" in metadata:\n id_to_embedding[metadata[\"_id\"]] = embeddings_result[\"embeddings\"][i]\n\n # Build output data based on include_metadata setting\n data_list = []\n for doc in results:\n kwargs = {\n \"content\": doc[0].page_content,\n }\n if self.search_query:\n kwargs[\"_score\"] = -1 * doc[1]\n if self.include_metadata:\n # Include all metadata, embeddings, and content\n kwargs.update(doc[0].metadata)\n if self.include_embeddings:\n kwargs[\"_embeddings\"] = id_to_embedding.get(doc[0].metadata.get(\"_id\"))\n\n data_list.append(Data(**kwargs))\n\n # Return the DataFrame containing the data\n return DataFrame(data=data_list)\n" }, "include_embeddings": { "_input_type": "BoolInput",