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.
This commit is contained in:
himavarshagoutham
2026-06-24 17:24:31 -04:00
parent 2b7b113049
commit 09acff1b17
2 changed files with 147 additions and 5 deletions

View File

@ -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"),
[

View File

@ -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.<name> 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,
)