mirror of
https://github.com/langflow-ai/langflow.git
synced 2026-07-23 16:10:27 +08:00
feat(embeddings): build available_models from all configured providers
Populate EmbeddingsWithModels.available_models with dedicated Embeddings instances for every enabled embedding model across all configured providers, not only the provider selected in the component.
This commit is contained in:
@ -552,28 +552,44 @@ def test_get_embeddings_openai_basic(mock_get_class, mock_get_api_key):
|
||||
|
||||
@patch("lfx.base.models.unified_models.get_api_key_for_provider")
|
||||
@patch("lfx.base.models.unified_models.get_embedding_class")
|
||||
def test_get_embeddings_populates_available_models_from_provider_catalog(mock_get_class, mock_get_api_key, monkeypatch):
|
||||
def test_get_embeddings_populates_available_models_from_all_configured_providers(
|
||||
mock_get_class, mock_get_api_key, monkeypatch
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
"lfx.base.models.unified_models.instantiation._get_configured_embedding_providers",
|
||||
lambda _user_id, _selected_provider: ["OpenAI", "Google Generative AI"],
|
||||
)
|
||||
|
||||
def _embedding_names_for_provider(provider, _user_id):
|
||||
if provider == "OpenAI":
|
||||
return ["text-embedding-3-small", "text-embedding-3-large"]
|
||||
if provider == "Google Generative AI":
|
||||
return ["models/text-embedding-004"]
|
||||
return []
|
||||
|
||||
monkeypatch.setattr(
|
||||
"lfx.base.models.unified_models.instantiation._get_provider_embedding_model_names",
|
||||
lambda _provider, _user_id: ["text-embedding-3-small", "text-embedding-3-large"],
|
||||
_embedding_names_for_provider,
|
||||
)
|
||||
mock_get_api_key.return_value = "sk-test"
|
||||
primary = MagicMock(name="primary")
|
||||
secondary = MagicMock(name="secondary")
|
||||
mock_embedding_class = MagicMock(side_effect=[primary, secondary])
|
||||
openai_secondary = MagicMock(name="openai-secondary")
|
||||
google_embedding = MagicMock(name="google-embedding")
|
||||
mock_embedding_class = MagicMock(side_effect=[primary, openai_secondary, google_embedding])
|
||||
mock_get_class.return_value = mock_embedding_class
|
||||
|
||||
result = get_embeddings([_make_openai_embedding_model()], api_key="sk-test")
|
||||
|
||||
assert isinstance(result, EmbeddingsWithModels)
|
||||
assert result.embeddings is primary
|
||||
assert set(result.available_models.keys()) == {"text-embedding-3-small", "text-embedding-3-large"}
|
||||
assert set(result.available_models.keys()) == {
|
||||
"text-embedding-3-small",
|
||||
"text-embedding-3-large",
|
||||
"models/text-embedding-004",
|
||||
}
|
||||
assert result.available_models["text-embedding-3-small"] is primary
|
||||
assert result.available_models["text-embedding-3-large"] is secondary
|
||||
|
||||
calls = mock_embedding_class.call_args_list
|
||||
assert calls[0].kwargs["model"] == "text-embedding-3-small"
|
||||
assert calls[1].kwargs["model"] == "text-embedding-3-large"
|
||||
assert result.available_models["text-embedding-3-large"] is openai_secondary
|
||||
assert result.available_models["models/text-embedding-004"] is google_embedding
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
||||
@ -12,7 +12,8 @@ class EmbeddingsWithModels(Embeddings):
|
||||
|
||||
Attributes:
|
||||
embeddings: The primary LangChain Embeddings instance (used as fallback).
|
||||
available_models: Dict mapping model names to their dedicated Embeddings instances.
|
||||
available_models: Dict mapping embedding model names to dedicated Embeddings instances
|
||||
across every configured provider in Model Providers settings.
|
||||
Each model has its own pre-configured instance with specific parameters.
|
||||
"""
|
||||
|
||||
@ -25,9 +26,9 @@ class EmbeddingsWithModels(Embeddings):
|
||||
|
||||
Args:
|
||||
embeddings: The primary LangChain Embeddings instance (used as default/fallback).
|
||||
available_models: Dict mapping model names to dedicated Embeddings instances.
|
||||
Each value should be a fully configured Embeddings object ready to use.
|
||||
Defaults to empty dict if not provided.
|
||||
available_models: Dict mapping embedding model names to dedicated Embeddings instances
|
||||
from all configured providers. Each value should be a fully configured
|
||||
Embeddings object ready to use. Defaults to empty dict if not provided.
|
||||
"""
|
||||
super().__init__()
|
||||
self.embeddings = embeddings
|
||||
@ -110,4 +111,6 @@ class EmbeddingsWithModels(Embeddings):
|
||||
|
||||
def __repr__(self) -> str:
|
||||
"""Return string representation of the wrapper."""
|
||||
return f"EmbeddingsWithModels(embeddings={self.embeddings!r}, available_models={self.available_models!r})"
|
||||
return (
|
||||
f"EmbeddingsWithModels(embeddings={self.embeddings!r}, available_models={self.available_models!r})"
|
||||
)
|
||||
|
||||
@ -13,13 +13,15 @@ from lfx.services.variable.request_scope import is_env_fallback_disabled
|
||||
from lfx.utils.async_helpers import run_until_complete
|
||||
|
||||
from .class_registry import EMBEDDING_PARAM_MAPPINGS, EMBEDDING_PROVIDER_CLASS_MAPPING
|
||||
from .credentials import _fetch_enabled_providers_for_user
|
||||
from .credentials import _fetch_enabled_providers_for_user, _get_model_status
|
||||
from .model_catalog import get_unified_models_detailed
|
||||
from .provider_queries import model_provider_metadata
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from uuid import UUID
|
||||
|
||||
from langchain_core.embeddings import Embeddings
|
||||
|
||||
|
||||
def _env_if_allowed(key: str) -> str | None:
|
||||
"""Return ``os.environ.get(key)`` unless the active request disables env fallback.
|
||||
@ -306,19 +308,13 @@ def get_llm(
|
||||
raise
|
||||
|
||||
|
||||
def _get_provider_embedding_model_names(
|
||||
def _get_provider_catalog_models(
|
||||
provider: str,
|
||||
user_id: UUID | str | None,
|
||||
) -> list[str]:
|
||||
"""Return all embedding model names for a provider from the configured catalog.
|
||||
|
||||
Unlike ``get_embedding_model_options``, this does not filter by user
|
||||
default/disabled/explicitly-enabled preferences — callers use it to build
|
||||
the full ``available_models`` map on ``EmbeddingsWithModels``.
|
||||
"""
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return catalog model entries for a provider (LLM + embedding, not type-filtered)."""
|
||||
provider_models = get_unified_models_detailed(
|
||||
providers=[provider],
|
||||
model_type="embeddings",
|
||||
include_deprecated=False,
|
||||
include_unsupported=False,
|
||||
)
|
||||
@ -331,51 +327,293 @@ def _get_provider_embedding_model_names(
|
||||
provider_models,
|
||||
user_id,
|
||||
{provider},
|
||||
"embeddings",
|
||||
None,
|
||||
model_provider_metadata,
|
||||
)
|
||||
|
||||
model_names: list[str] = []
|
||||
catalog_models: list[dict[str, Any]] = []
|
||||
for provider_data in provider_models:
|
||||
if provider_data.get("provider") != provider:
|
||||
continue
|
||||
for model_data in provider_data.get("models", []):
|
||||
name = model_data.get("model_name")
|
||||
if name:
|
||||
model_names.append(name)
|
||||
catalog_models.extend(provider_data.get("models", []))
|
||||
return catalog_models
|
||||
|
||||
|
||||
def _get_provider_enabled_model_names(
|
||||
provider: str,
|
||||
user_id: UUID | str | None,
|
||||
) -> list[str]:
|
||||
"""Return model names enabled for a provider in Model Providers settings.
|
||||
|
||||
Includes both LLM and embedding models (e.g. gpt-5 and text-embedding-3-small).
|
||||
When *user_id* is absent, returns all non-deprecated catalog models for the provider.
|
||||
"""
|
||||
catalog_models = _get_provider_catalog_models(provider, user_id)
|
||||
|
||||
disabled_models: set[str] = set()
|
||||
explicitly_enabled_models: set[str] = set()
|
||||
enabled_providers: set[str] = set()
|
||||
if user_id:
|
||||
with contextlib.suppress(Exception):
|
||||
disabled_models, explicitly_enabled_models = run_until_complete(_get_model_status(user_id))
|
||||
with contextlib.suppress(Exception):
|
||||
enabled_providers = run_until_complete(_fetch_enabled_providers_for_user(user_id))
|
||||
|
||||
apply_user_prefs = bool(user_id and enabled_providers and provider in enabled_providers)
|
||||
|
||||
model_names: list[str] = []
|
||||
for model_data in catalog_models:
|
||||
model_name = model_data.get("model_name")
|
||||
if not model_name:
|
||||
continue
|
||||
|
||||
if apply_user_prefs:
|
||||
metadata = model_data.get("metadata", {})
|
||||
is_default = metadata.get("default", False)
|
||||
if not is_default and model_name not in explicitly_enabled_models:
|
||||
continue
|
||||
if model_name in disabled_models:
|
||||
continue
|
||||
|
||||
model_names.append(model_name)
|
||||
return model_names
|
||||
|
||||
|
||||
def _build_available_embedding_models(
|
||||
embedding_class: type,
|
||||
kwargs: dict[str, Any],
|
||||
param_mapping: dict[str, str],
|
||||
def _is_embedding_catalog_model(model_data: dict[str, Any]) -> bool:
|
||||
"""Return True when a catalog entry is an embedding model."""
|
||||
model_type = model_data.get("metadata", {}).get("model_type", "llm")
|
||||
return model_type == "embeddings"
|
||||
|
||||
|
||||
def _get_configured_embedding_providers(
|
||||
user_id: UUID | str | None,
|
||||
selected_provider: str,
|
||||
) -> list[str]:
|
||||
"""Return embedding-capable providers configured in Model Providers."""
|
||||
if not user_id:
|
||||
return [selected_provider]
|
||||
|
||||
with contextlib.suppress(Exception):
|
||||
enabled_providers = run_until_complete(_fetch_enabled_providers_for_user(user_id))
|
||||
providers = sorted(p for p in enabled_providers if p in EMBEDDING_PROVIDER_CLASS_MAPPING)
|
||||
if selected_provider not in providers:
|
||||
providers.insert(0, selected_provider)
|
||||
return providers
|
||||
|
||||
return [selected_provider]
|
||||
|
||||
|
||||
def _get_provider_embedding_model_names(
|
||||
provider: str,
|
||||
user_id: UUID | str | None,
|
||||
primary_model_name: str,
|
||||
primary_instance: Any,
|
||||
) -> dict[str, Any]:
|
||||
"""Build dedicated embedding instances for every model on the configured provider."""
|
||||
available_models: dict[str, Any] = {primary_model_name: primary_instance}
|
||||
) -> list[str]:
|
||||
"""Return enabled embedding model names for a single provider."""
|
||||
catalog_by_name = {
|
||||
model_data.get("model_name"): model_data
|
||||
for model_data in _get_provider_catalog_models(provider, user_id)
|
||||
if model_data.get("model_name")
|
||||
}
|
||||
|
||||
model_param_key = param_mapping.get("model") or param_mapping.get("model_id")
|
||||
if not model_param_key:
|
||||
return available_models
|
||||
|
||||
for model_name in _get_provider_embedding_model_names(provider, user_id):
|
||||
if model_name in available_models:
|
||||
model_names: list[str] = []
|
||||
for model_name in _get_provider_enabled_model_names(provider, user_id):
|
||||
model_data = catalog_by_name.get(model_name)
|
||||
if model_data is None or not _is_embedding_catalog_model(model_data):
|
||||
continue
|
||||
model_kwargs = dict(kwargs)
|
||||
model_kwargs[model_param_key] = model_name
|
||||
try:
|
||||
available_models[model_name] = embedding_class(**model_kwargs)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug(
|
||||
"Failed to instantiate embedding model %s for provider %s; skipping",
|
||||
model_name,
|
||||
model_names.append(model_name)
|
||||
return model_names
|
||||
|
||||
|
||||
def _compose_embedding_kwargs(
|
||||
provider: str,
|
||||
model_name: str,
|
||||
user_id: UUID | str | None,
|
||||
unified_models_module: Any,
|
||||
*,
|
||||
selected_provider: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
component_api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
dimensions: int | None = None,
|
||||
chunk_size: int | None = None,
|
||||
request_timeout: float | None = None,
|
||||
max_retries: int | None = None,
|
||||
show_progress_bar: bool | None = None,
|
||||
model_kwargs: dict[str, Any] | None = None,
|
||||
watsonx_url: str | None = None,
|
||||
watsonx_project_id: str | None = None,
|
||||
watsonx_truncate_input_tokens: int | None = None,
|
||||
watsonx_input_text: bool | None = None,
|
||||
ollama_base_url: str | None = None,
|
||||
) -> tuple[type, dict[str, Any]] | None:
|
||||
"""Build kwargs for a provider/model pair. Returns None when credentials are missing."""
|
||||
metadata = metadata or {}
|
||||
api_key_override = component_api_key if provider == selected_provider else None
|
||||
api_key = unified_models_module.get_api_key_for_provider(user_id, provider, api_key_override)
|
||||
if not api_key and provider != "Ollama":
|
||||
return None
|
||||
|
||||
embedding_class_name = metadata.get("embedding_class") or EMBEDDING_PROVIDER_CLASS_MAPPING.get(provider)
|
||||
if not embedding_class_name:
|
||||
return None
|
||||
|
||||
param_mapping: dict[str, str] = metadata.get("param_mapping") or EMBEDDING_PARAM_MAPPINGS.get(provider, {})
|
||||
if not param_mapping:
|
||||
return None
|
||||
|
||||
embedding_class = unified_models_module.get_embedding_class(embedding_class_name)
|
||||
|
||||
api_base_value = _to_str(api_base) if provider == selected_provider else None
|
||||
if provider == "OpenAI" and not api_base_value:
|
||||
api_base_value = _to_str(os.environ.get("OPENAI_EMBEDDINGS_API_BASE")) or _to_str(
|
||||
os.environ.get("OPENAI_API_BASE")
|
||||
)
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
if "model" in param_mapping:
|
||||
kwargs[param_mapping["model"]] = model_name
|
||||
elif "model_id" in param_mapping:
|
||||
kwargs[param_mapping["model_id"]] = model_name
|
||||
|
||||
if "api_key" in param_mapping and api_key:
|
||||
kwargs[param_mapping["api_key"]] = api_key
|
||||
|
||||
use_component_overrides = provider == selected_provider
|
||||
optional_params: dict[str, Any] = {
|
||||
"api_base": api_base_value if use_component_overrides else None,
|
||||
"dimensions": dimensions if use_component_overrides else None,
|
||||
"chunk_size": chunk_size if use_component_overrides else None,
|
||||
"request_timeout": request_timeout if use_component_overrides else None,
|
||||
"max_retries": max_retries if use_component_overrides else None,
|
||||
"show_progress_bar": show_progress_bar if use_component_overrides else None,
|
||||
"model_kwargs": model_kwargs if use_component_overrides else None,
|
||||
}
|
||||
|
||||
if provider in {"IBM WatsonX", "IBM watsonx.ai"}:
|
||||
watsonx_provider_vars = unified_models_module.get_all_variables_for_provider(user_id, provider)
|
||||
url_value = (
|
||||
(watsonx_url if use_component_overrides else None)
|
||||
or watsonx_provider_vars.get("WATSONX_URL")
|
||||
or _env_if_allowed("WATSONX_URL")
|
||||
)
|
||||
pid_value = (
|
||||
(watsonx_project_id if use_component_overrides else None)
|
||||
or watsonx_provider_vars.get("WATSONX_PROJECT_ID")
|
||||
or _env_if_allowed("WATSONX_PROJECT_ID")
|
||||
)
|
||||
if url_value and pid_value:
|
||||
if "url" in param_mapping:
|
||||
kwargs[param_mapping["url"]] = url_value
|
||||
if "project_id" in param_mapping:
|
||||
kwargs[param_mapping["project_id"]] = pid_value
|
||||
|
||||
if use_component_overrides:
|
||||
watsonx_params = {}
|
||||
if watsonx_truncate_input_tokens is not None:
|
||||
try:
|
||||
from ibm_watsonx_ai.metanames import EmbedTextParamsMetaNames
|
||||
|
||||
watsonx_params[EmbedTextParamsMetaNames.TRUNCATE_INPUT_TOKENS] = int(watsonx_truncate_input_tokens)
|
||||
except ImportError:
|
||||
watsonx_params["truncate_input_tokens"] = int(watsonx_truncate_input_tokens)
|
||||
if watsonx_input_text is not None:
|
||||
try:
|
||||
from ibm_watsonx_ai.metanames import EmbedTextParamsMetaNames
|
||||
|
||||
watsonx_params[EmbedTextParamsMetaNames.RETURN_OPTIONS] = {"input_text": bool(watsonx_input_text)}
|
||||
except ImportError:
|
||||
watsonx_params["return_options"] = {"input_text": bool(watsonx_input_text)}
|
||||
if watsonx_params:
|
||||
kwargs["params"] = watsonx_params
|
||||
|
||||
if provider == "Ollama" and "base_url" in param_mapping:
|
||||
provider_vars = unified_models_module.get_all_variables_for_provider(user_id, provider)
|
||||
base_url_value = (
|
||||
(ollama_base_url if use_component_overrides else None)
|
||||
or provider_vars.get("OLLAMA_BASE_URL")
|
||||
or _env_if_allowed("OLLAMA_BASE_URL")
|
||||
or "http://localhost:11434"
|
||||
)
|
||||
kwargs[param_mapping["base_url"]] = base_url_value
|
||||
|
||||
for param_name, param_value in optional_params.items():
|
||||
if param_value is not None and param_name in param_mapping:
|
||||
if (
|
||||
param_name == "request_timeout"
|
||||
and provider == "Google Generative AI"
|
||||
and isinstance(param_value, (int, float))
|
||||
):
|
||||
kwargs[param_mapping[param_name]] = {"timeout": param_value}
|
||||
else:
|
||||
kwargs[param_mapping[param_name]] = param_value
|
||||
|
||||
return embedding_class, kwargs
|
||||
|
||||
|
||||
def _build_available_embedding_models(
|
||||
*,
|
||||
selected_provider: str,
|
||||
primary_model_name: str,
|
||||
primary_instance: Embeddings,
|
||||
user_id: UUID | str | None,
|
||||
unified_models_module: Any,
|
||||
metadata: dict[str, Any],
|
||||
component_api_key: str | None,
|
||||
api_base: str | None,
|
||||
dimensions: int | None,
|
||||
chunk_size: int | None,
|
||||
request_timeout: float | None,
|
||||
max_retries: int | None,
|
||||
show_progress_bar: bool | None,
|
||||
model_kwargs: dict[str, Any] | None,
|
||||
watsonx_url: str | None,
|
||||
watsonx_project_id: str | None,
|
||||
watsonx_truncate_input_tokens: int | None,
|
||||
watsonx_input_text: bool | None,
|
||||
ollama_base_url: str | None,
|
||||
) -> dict[str, Embeddings]:
|
||||
"""Build embedding instances for every enabled embedding model on configured providers."""
|
||||
available_models: dict[str, Embeddings] = {primary_model_name: primary_instance}
|
||||
|
||||
for provider in _get_configured_embedding_providers(user_id, selected_provider):
|
||||
provider_metadata = metadata if provider == selected_provider else {}
|
||||
for model_name in _get_provider_embedding_model_names(provider, user_id):
|
||||
if model_name in available_models:
|
||||
continue
|
||||
|
||||
composed = _compose_embedding_kwargs(
|
||||
provider,
|
||||
exc_info=True,
|
||||
model_name,
|
||||
user_id,
|
||||
unified_models_module,
|
||||
selected_provider=selected_provider,
|
||||
metadata=provider_metadata,
|
||||
component_api_key=component_api_key,
|
||||
api_base=api_base,
|
||||
dimensions=dimensions,
|
||||
chunk_size=chunk_size,
|
||||
request_timeout=request_timeout,
|
||||
max_retries=max_retries,
|
||||
show_progress_bar=show_progress_bar,
|
||||
model_kwargs=model_kwargs,
|
||||
watsonx_url=watsonx_url,
|
||||
watsonx_project_id=watsonx_project_id,
|
||||
watsonx_truncate_input_tokens=watsonx_truncate_input_tokens,
|
||||
watsonx_input_text=watsonx_input_text,
|
||||
ollama_base_url=ollama_base_url,
|
||||
)
|
||||
if composed is None:
|
||||
continue
|
||||
|
||||
embedding_class, model_kwargs_dict = composed
|
||||
try:
|
||||
available_models[model_name] = embedding_class(**model_kwargs_dict)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug(
|
||||
"Failed to instantiate embedding model %s for provider %s; skipping",
|
||||
model_name,
|
||||
provider,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
return available_models
|
||||
|
||||
@ -402,7 +640,7 @@ def get_embeddings(
|
||||
|
||||
Returns an :class:`~lfx.base.embeddings.embeddings_class.EmbeddingsWithModels`
|
||||
wrapper containing the primary instance for the selected model and an
|
||||
``available_models`` map of all embedding models for the configured provider.
|
||||
``available_models`` map of enabled embedding models from every configured provider.
|
||||
"""
|
||||
# Resolve helpers via package namespace so tests patching
|
||||
# lfx.base.models.unified_models.<name> keep working.
|
||||
@ -431,13 +669,8 @@ def get_embeddings(
|
||||
model_name = model_dict.get("name")
|
||||
provider = model_dict.get("provider")
|
||||
metadata = model_dict.get("metadata", {})
|
||||
api_base_value = _to_str(api_base)
|
||||
if provider == "OpenAI" and not api_base_value:
|
||||
api_base_value = _to_str(os.environ.get("OPENAI_EMBEDDINGS_API_BASE")) or _to_str(
|
||||
os.environ.get("OPENAI_API_BASE")
|
||||
)
|
||||
|
||||
# --- resolve API key -----------------------------------------------------
|
||||
# --- resolve API key for the selected provider ---------------------------
|
||||
api_key = unified_models_module.get_api_key_for_provider(user_id, provider, api_key)
|
||||
if not api_key and provider != "Ollama":
|
||||
provider_variable_map = unified_models_module.get_model_provider_variable_mapping()
|
||||
@ -452,124 +685,47 @@ def get_embeddings(
|
||||
msg = "Embedding model name is required"
|
||||
raise ValueError(msg)
|
||||
|
||||
# Get embedding class from metadata. Selections persisted via the
|
||||
# generic ``/models`` catalog (e.g. saved into ``KnowledgeBase.model_selection``
|
||||
# by the KB upload flow) lack the enriched embedding metadata, so we
|
||||
# fall back to deriving it from the provider name. Both lookups
|
||||
# share the same source of truth in ``class_registry``.
|
||||
embedding_class_name = metadata.get("embedding_class") or EMBEDDING_PROVIDER_CLASS_MAPPING.get(provider)
|
||||
if not embedding_class_name:
|
||||
msg = (
|
||||
f"No embedding class defined in metadata for {model_name} (provider: {provider}). "
|
||||
"Add the provider to EMBEDDING_PROVIDER_CLASS_MAPPING or re-select the model."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
embedding_class = unified_models_module.get_embedding_class(embedding_class_name)
|
||||
|
||||
# --- build kwargs from param_mapping -------------------------------------
|
||||
param_mapping: dict[str, str] = metadata.get("param_mapping") or EMBEDDING_PARAM_MAPPINGS.get(provider, {})
|
||||
if not param_mapping:
|
||||
msg = (
|
||||
f"Parameter mapping not found in metadata for model '{model_name}' (provider: {provider}). "
|
||||
"This usually means the model was saved with an older format that is no longer recognized. "
|
||||
"Please re-select the embedding model in the component configuration."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
|
||||
kwargs: dict[str, Any] = {}
|
||||
|
||||
# Model name
|
||||
if "model" in param_mapping:
|
||||
kwargs[param_mapping["model"]] = model_name
|
||||
elif "model_id" in param_mapping:
|
||||
kwargs[param_mapping["model_id"]] = model_name
|
||||
|
||||
# API key
|
||||
if "api_key" in param_mapping and api_key:
|
||||
kwargs[param_mapping["api_key"]] = api_key
|
||||
|
||||
# Optional parameters - only add when both a value is supplied *and* the
|
||||
# provider's param_mapping declares the corresponding key.
|
||||
optional_params: dict[str, Any] = {
|
||||
"api_base": api_base_value or None,
|
||||
"dimensions": int(dimensions) if dimensions else None,
|
||||
"chunk_size": int(chunk_size) if chunk_size else None,
|
||||
"request_timeout": float(request_timeout) if request_timeout else None,
|
||||
"max_retries": int(max_retries) if max_retries else None,
|
||||
"show_progress_bar": show_progress_bar,
|
||||
"model_kwargs": model_kwargs if model_kwargs else None,
|
||||
}
|
||||
|
||||
# Watson-specific parameters
|
||||
if provider in {"IBM WatsonX", "IBM watsonx.ai"}:
|
||||
watsonx_provider_vars = unified_models_module.get_all_variables_for_provider(user_id, provider)
|
||||
url_value = watsonx_url or watsonx_provider_vars.get("WATSONX_URL") or _env_if_allowed("WATSONX_URL")
|
||||
pid_value = (
|
||||
watsonx_project_id
|
||||
or watsonx_provider_vars.get("WATSONX_PROJECT_ID")
|
||||
or _env_if_allowed("WATSONX_PROJECT_ID")
|
||||
)
|
||||
|
||||
has_url = bool(url_value)
|
||||
has_project_id = bool(pid_value)
|
||||
|
||||
if has_url and has_project_id:
|
||||
if "url" in param_mapping:
|
||||
kwargs[param_mapping["url"]] = url_value
|
||||
if "project_id" in param_mapping:
|
||||
kwargs[param_mapping["project_id"]] = pid_value
|
||||
elif has_url or has_project_id:
|
||||
missing = "project ID (WATSONX_PROJECT_ID)" if has_url else "URL (WATSONX_URL)"
|
||||
provided = "URL" if has_url else "project ID"
|
||||
composed = _compose_embedding_kwargs(
|
||||
provider,
|
||||
model_name,
|
||||
user_id,
|
||||
unified_models_module,
|
||||
selected_provider=provider,
|
||||
metadata=metadata,
|
||||
component_api_key=api_key,
|
||||
api_base=_to_str(api_base),
|
||||
dimensions=int(dimensions) if dimensions else None,
|
||||
chunk_size=int(chunk_size) if chunk_size else None,
|
||||
request_timeout=float(request_timeout) if request_timeout else None,
|
||||
max_retries=int(max_retries) if max_retries else None,
|
||||
show_progress_bar=show_progress_bar,
|
||||
model_kwargs=model_kwargs if model_kwargs else None,
|
||||
watsonx_url=watsonx_url,
|
||||
watsonx_project_id=watsonx_project_id,
|
||||
watsonx_truncate_input_tokens=watsonx_truncate_input_tokens,
|
||||
watsonx_input_text=watsonx_input_text,
|
||||
ollama_base_url=ollama_base_url,
|
||||
)
|
||||
if composed is None:
|
||||
embedding_class_name = metadata.get("embedding_class") or EMBEDDING_PROVIDER_CLASS_MAPPING.get(provider)
|
||||
if not embedding_class_name:
|
||||
msg = (
|
||||
f"IBM WatsonX requires both a URL and project ID. "
|
||||
f"You provided a watsonx {provided} but no {missing}. "
|
||||
f"Please configure the missing value in the component or set the environment variable."
|
||||
f"No embedding class defined in metadata for {model_name} (provider: {provider}). "
|
||||
"Add the provider to EMBEDDING_PROVIDER_CLASS_MAPPING or re-select the model."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
param_mapping = metadata.get("param_mapping") or EMBEDDING_PARAM_MAPPINGS.get(provider, {})
|
||||
if not param_mapping:
|
||||
msg = (
|
||||
f"Parameter mapping not found in metadata for model '{model_name}' (provider: {provider}). "
|
||||
"This usually means the model was saved with an older format that is no longer recognized. "
|
||||
"Please re-select the embedding model in the component configuration."
|
||||
)
|
||||
raise ValueError(msg)
|
||||
msg = f"{provider} API key is required."
|
||||
raise ValueError(msg)
|
||||
|
||||
# Build WatsonX embed params (truncate_input_tokens, return_options)
|
||||
watsonx_params = {}
|
||||
if watsonx_truncate_input_tokens is not None:
|
||||
try:
|
||||
from ibm_watsonx_ai.metanames import EmbedTextParamsMetaNames
|
||||
|
||||
watsonx_params[EmbedTextParamsMetaNames.TRUNCATE_INPUT_TOKENS] = int(watsonx_truncate_input_tokens)
|
||||
except ImportError:
|
||||
watsonx_params["truncate_input_tokens"] = int(watsonx_truncate_input_tokens)
|
||||
if watsonx_input_text is not None:
|
||||
try:
|
||||
from ibm_watsonx_ai.metanames import EmbedTextParamsMetaNames
|
||||
|
||||
watsonx_params[EmbedTextParamsMetaNames.RETURN_OPTIONS] = {"input_text": bool(watsonx_input_text)}
|
||||
except ImportError:
|
||||
watsonx_params["return_options"] = {"input_text": bool(watsonx_input_text)}
|
||||
if watsonx_params:
|
||||
kwargs["params"] = watsonx_params
|
||||
|
||||
# Ollama-specific parameters
|
||||
if provider == "Ollama" and "base_url" in param_mapping:
|
||||
provider_vars = unified_models_module.get_all_variables_for_provider(user_id, provider)
|
||||
base_url_value = (
|
||||
ollama_base_url
|
||||
or provider_vars.get("OLLAMA_BASE_URL")
|
||||
or _env_if_allowed("OLLAMA_BASE_URL")
|
||||
or "http://localhost:11434"
|
||||
)
|
||||
kwargs[param_mapping["base_url"]] = base_url_value
|
||||
|
||||
# Add optional parameters if they have values and are mapped
|
||||
for param_name, param_value in optional_params.items():
|
||||
if param_value is not None and param_name in param_mapping:
|
||||
# Google wraps timeout in a dict
|
||||
if (
|
||||
param_name == "request_timeout"
|
||||
and provider == "Google Generative AI"
|
||||
and isinstance(param_value, (int, float))
|
||||
):
|
||||
kwargs[param_mapping[param_name]] = {"timeout": param_value}
|
||||
else:
|
||||
kwargs[param_mapping[param_name]] = param_value
|
||||
embedding_class, kwargs = composed
|
||||
|
||||
try:
|
||||
primary_instance = embedding_class(**kwargs)
|
||||
@ -583,13 +739,25 @@ def get_embeddings(
|
||||
raise
|
||||
|
||||
available_models = _build_available_embedding_models(
|
||||
embedding_class=embedding_class,
|
||||
kwargs=kwargs,
|
||||
param_mapping=param_mapping,
|
||||
provider=provider,
|
||||
user_id=user_id,
|
||||
selected_provider=provider,
|
||||
primary_model_name=model_name,
|
||||
primary_instance=primary_instance,
|
||||
user_id=user_id,
|
||||
unified_models_module=unified_models_module,
|
||||
metadata=metadata,
|
||||
component_api_key=api_key,
|
||||
api_base=_to_str(api_base),
|
||||
dimensions=int(dimensions) if dimensions else None,
|
||||
chunk_size=int(chunk_size) if chunk_size else None,
|
||||
request_timeout=float(request_timeout) if request_timeout else None,
|
||||
max_retries=int(max_retries) if max_retries else None,
|
||||
show_progress_bar=show_progress_bar,
|
||||
model_kwargs=model_kwargs if model_kwargs else None,
|
||||
watsonx_url=watsonx_url,
|
||||
watsonx_project_id=watsonx_project_id,
|
||||
watsonx_truncate_input_tokens=watsonx_truncate_input_tokens,
|
||||
watsonx_input_text=watsonx_input_text,
|
||||
ollama_base_url=ollama_base_url,
|
||||
)
|
||||
|
||||
return EmbeddingsWithModels(
|
||||
|
||||
Reference in New Issue
Block a user