From 09acff1b170ddf962667f4ecff9703781f4983b5 Mon Sep 17 00:00:00 2001 From: himavarshagoutham Date: Wed, 24 Jun 2026 17:24:31 -0400 Subject: [PATCH] feat(embeddings): populate available_models from configured providers Restore EmbeddingsWithModels wrapping in get_embeddings so multi-model consumers (e.g. OpenSearch multimodal) get dedicated instances for every embedding model on the configured provider catalog, including live models. --- src/backend/tests/unit/test_unified_models.py | 46 +++++++- .../models/unified_models/instantiation.py | 106 +++++++++++++++++- 2 files changed, 147 insertions(+), 5 deletions(-) diff --git a/src/backend/tests/unit/test_unified_models.py b/src/backend/tests/unit/test_unified_models.py index 17f8f08e90..79affa4270 100644 --- a/src/backend/tests/unit/test_unified_models.py +++ b/src/backend/tests/unit/test_unified_models.py @@ -1,6 +1,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from lfx.base.embeddings.embeddings_class import EmbeddingsWithModels from lfx.base.models import models_dev_catalog from lfx.base.models.unified_models import ( _get_all_provider_mapped_fields, @@ -414,6 +415,15 @@ def test_get_all_provider_mapped_fields_is_cached(): # --------------------------------------------------------------------------- +@pytest.fixture(autouse=True) +def _skip_available_models_catalog(monkeypatch): + """Keep existing get_embeddings tests focused on the primary instance.""" + monkeypatch.setattr( + "lfx.base.models.unified_models.instantiation._get_provider_embedding_model_names", + lambda provider, user_id: [], + ) + + def test_get_embeddings_passthrough_embeddings_object(): """An already-instantiated Embeddings object should be returned as-is.""" from langchain_core.embeddings import Embeddings as BaseEmbeddings @@ -515,7 +525,9 @@ def test_get_embeddings_falls_back_when_metadata_stripped(mock_get_class, mock_g kwargs = fake_class.call_args.kwargs assert kwargs["model"] == "text-embedding-3-small" assert kwargs["api_key"] == "sk-test" - assert result == "embeddings-instance" + assert isinstance(result, EmbeddingsWithModels) + assert result.embeddings == "embeddings-instance" + assert result.available_models == {"text-embedding-3-small": "embeddings-instance"} @patch("lfx.base.models.unified_models.get_api_key_for_provider") @@ -529,13 +541,43 @@ def test_get_embeddings_openai_basic(mock_get_class, mock_get_api_key): result = get_embeddings([_make_openai_embedding_model()], api_key="sk-test") - assert result is mock_instance + assert isinstance(result, EmbeddingsWithModels) + assert result.embeddings is mock_instance + assert result.available_models == {"text-embedding-3-small": mock_instance} mock_get_class.assert_called_once_with("OpenAIEmbeddings") kwargs = mock_embedding_class.call_args.kwargs assert kwargs["model"] == "text-embedding-3-small" assert kwargs["api_key"] == "sk-test" # pragma: allowlist secret +@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 +): + 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"], + ) + mock_get_api_key.return_value = "sk-test" + primary = MagicMock(name="primary") + secondary = MagicMock(name="secondary") + mock_embedding_class = MagicMock(side_effect=[primary, secondary]) + 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 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" + + @pytest.mark.parametrize( ("env_values", "expected_base_url"), [ 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 1333b4157a..a62b5d2eef 100644 --- a/src/lfx/src/lfx/base/models/unified_models/instantiation.py +++ b/src/lfx/src/lfx/base/models/unified_models/instantiation.py @@ -2,13 +2,19 @@ from __future__ import annotations +import contextlib import os from typing import TYPE_CHECKING, Any -from lfx.base.models.model_utils import _to_str +from lfx.base.embeddings.embeddings_class import EmbeddingsWithModels +from lfx.base.models.model_utils import _to_str, replace_with_live_models +from lfx.log.logger import logger 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 .model_catalog import get_unified_models_detailed from .provider_queries import model_provider_metadata if TYPE_CHECKING: @@ -275,6 +281,80 @@ def get_llm( raise +def _get_provider_embedding_model_names( + 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``. + """ + provider_models = get_unified_models_detailed( + providers=[provider], + model_type="embeddings", + include_deprecated=False, + include_unsupported=False, + ) + + if user_id: + with contextlib.suppress(Exception): + enabled_providers = run_until_complete(_fetch_enabled_providers_for_user(user_id)) + if provider in enabled_providers: + replace_with_live_models( + provider_models, + user_id, + {provider}, + "embeddings", + model_provider_metadata, + ) + + model_names: list[str] = [] + 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) + return model_names + + +def _build_available_embedding_models( + embedding_class: type, + kwargs: dict[str, Any], + param_mapping: dict[str, str], + 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} + + 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: + 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, + provider, + exc_info=True, + ) + + return available_models + + def get_embeddings( model, user_id: UUID | str | None = None, @@ -293,7 +373,12 @@ def get_embeddings( watsonx_input_text=None, ollama_base_url=None, ) -> Any: - """Instantiate an embeddings model from a model selection dict.""" + """Instantiate an embeddings model from a model selection dict. + + 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. + """ # Resolve helpers via package namespace so tests patching # lfx.base.models.unified_models. keep working. from lfx.base.models import unified_models as unified_models_module @@ -462,7 +547,7 @@ def get_embeddings( kwargs[param_mapping[param_name]] = param_value try: - return embedding_class(**kwargs) + primary_instance = embedding_class(**kwargs) except Exception as e: if provider == "IBM WatsonX" and ("url" in str(e).lower() or "project" in str(e).lower()): msg = ( @@ -471,3 +556,18 @@ def get_embeddings( ) raise ValueError(msg) from e raise + + available_models = _build_available_embedding_models( + embedding_class=embedding_class, + kwargs=kwargs, + param_mapping=param_mapping, + provider=provider, + user_id=user_id, + primary_model_name=model_name, + primary_instance=primary_instance, + ) + + return EmbeddingsWithModels( + embeddings=primary_instance, + available_models=available_models, + )