mirror of
https://github.com/langflow-ai/langflow.git
synced 2026-07-23 16:10:27 +08:00
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:
@ -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"),
|
||||
[
|
||||
|
||||
@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user