diff --git a/src/backend/tests/unit/test_unified_models.py b/src/backend/tests/unit/test_unified_models.py index addc4c22e9..07b0ef3bc8 100644 --- a/src/backend/tests/unit/test_unified_models.py +++ b/src/backend/tests/unit/test_unified_models.py @@ -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( diff --git a/src/lfx/src/lfx/base/embeddings/embeddings_class.py b/src/lfx/src/lfx/base/embeddings/embeddings_class.py index 4a71a5b639..582598c057 100644 --- a/src/lfx/src/lfx/base/embeddings/embeddings_class.py +++ b/src/lfx/src/lfx/base/embeddings/embeddings_class.py @@ -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})" + ) diff --git a/src/lfx/src/lfx/base/models/unified_models/instantiation.py b/src/lfx/src/lfx/base/models/unified_models/instantiation.py index 756fc24a24..a00c9bdf53 100644 --- a/src/lfx/src/lfx/base/models/unified_models/instantiation.py +++ b/src/lfx/src/lfx/base/models/unified_models/instantiation.py @@ -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. 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(