Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 44 additions & 39 deletions src/lfx/src/lfx/base/models/unified_models/instantiation.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,33 +90,32 @@ def get_llm(
overrides: dict[str, Any] | None = None,
provider_policy: ModelProviderPolicySnapshot | None = None,
) -> Any:
# 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

# Coerce provider-specific string params (Message/Data may leak through StrInput)
ollama_base_url = _to_str(ollama_base_url)
watsonx_url = _to_str(watsonx_url)
watsonx_project_id = _to_str(watsonx_project_id)

# Check if model is already a BaseLanguageModel instance (from a connection)
try:
from langchain_core.language_models import BaseLanguageModel

if isinstance(model, BaseLanguageModel):
# Model is already instantiated, return it directly
return model
except ImportError:
pass
# List-shaped selections carry the provider identity needed for policy
# enforcement. Gate that common runtime path before importing even the
# LangChain base class, provider SDKs, or credential-resolution helpers.
if isinstance(model, list):
if not model:
msg = "A model selection is required"
raise ValueError(msg)
model = model[0]
else:
# Preserve the existing connection-object passthrough. Only this
# identity-free path needs the LangChain base-class import up front.
try:
from langchain_core.language_models import BaseLanguageModel

# Safely extract model configuration
if not model or not isinstance(model, list) or len(model) == 0:
if isinstance(model, BaseLanguageModel):
return model
except ImportError:
pass
msg = "A model selection is required"
raise ValueError(msg)

# Extract the first model (only one expected)
model = model[0]

# Extract model configuration from metadata
model_name = model.get("name")
provider = model.get("provider")
Expand All @@ -129,17 +128,21 @@ def get_llm(
)
raise ValueError(msg)
if provider_policy is None:
from lfx.base.models.provider_registry import get_registry_snapshot
from lfx.services.model_provider_policy import ModelProviderPolicyPurpose, resolve_model_provider_policy

from .provider_queries import get_model_providers

provider_policy = resolve_model_provider_policy(
user_id=user_id,
providers=(*get_model_providers(), provider),
providers=(*get_registry_snapshot().provider_ids, provider),
purpose=ModelProviderPolicyPurpose.USE,
)
provider_policy.require(provider)

# Resolve helpers through the package namespace only after policy passes so
# tests can patch lfx.base.models.unified_models.<name> and denied requests
# cannot trigger runtime integration imports.
from lfx.base.models import unified_models as unified_models_module

# Policy is evaluated against the submitted identity, then all runtime
# wiring for a known provider is resolved from its canonical registry
# entry. Stored flow metadata is not authoritative for class or parameter
Expand Down Expand Up @@ -864,30 +867,30 @@ def get_embeddings(
wrapper containing the primary instance for the selected model and an
``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.
from lfx.base.models import unified_models as unified_models_module

# Coerce provider-specific string params
ollama_base_url = _to_str(ollama_base_url)
watsonx_url = _to_str(watsonx_url)
watsonx_project_id = _to_str(watsonx_project_id)

# Passthrough: already-instantiated Embeddings object from a connection
try:
from langchain_core.embeddings import Embeddings as BaseEmbeddings

if isinstance(model, BaseEmbeddings):
return model
except ImportError:
pass
# Gate list-shaped selections before importing a LangChain base class,
# provider SDK, or credential-resolution helper.
if isinstance(model, list):
if not model:
msg = "An embedding model selection is required"
raise ValueError(msg)
model_dict = model[0]
else:
# Preserve passthrough for already-instantiated connection objects.
try:
from langchain_core.embeddings import Embeddings as BaseEmbeddings

# Validate input
if not model or not isinstance(model, list) or len(model) == 0:
if isinstance(model, BaseEmbeddings):
return model
except ImportError:
pass
msg = "An embedding model selection is required"
raise ValueError(msg)

model_dict = model[0]
model_name = model_dict.get("name")
provider = model_dict.get("provider")
metadata = model_dict.get("metadata", {})
Expand All @@ -896,17 +899,19 @@ def get_embeddings(
msg = "The selected embedding model is missing a provider"
raise ValueError(msg)
if provider_policy is None:
from lfx.base.models.provider_registry import get_registry_snapshot
from lfx.services.model_provider_policy import ModelProviderPolicyPurpose, resolve_model_provider_policy

from .provider_queries import get_model_providers

provider_policy = resolve_model_provider_policy(
user_id=user_id,
providers=(*get_model_providers(), provider),
providers=(*get_registry_snapshot().provider_ids, provider),
purpose=ModelProviderPolicyPurpose.USE,
)
provider_policy.require(provider)

# Resolve helpers through the patchable package namespace only after the
# provider has been authorized for runtime use.
from lfx.base.models import unified_models as unified_models_module
from lfx.base.models.provider_registry import provider_name_for_id, resolve_provider_id

provider_id = resolve_provider_id(provider)
Expand Down
76 changes: 76 additions & 0 deletions src/lfx/tests/unit/services/model_provider_policy/test_policy.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import builtins
from dataclasses import FrozenInstanceError
from unittest.mock import AsyncMock, patch

Expand All @@ -16,6 +17,13 @@
set_current_model_provider_policy_context,
)

_DENIED_RUNTIME_CATALOG_CALLS: list[None] = []


def _denied_runtime_catalog_loader():
_DENIED_RUNTIME_CATALOG_CALLS.append(None)
return [{"name": "blocked-model", "model_type": "llm"}]


def _restricted_snapshot(*allowed: str) -> ModelProviderPolicySnapshot:
return ModelProviderPolicySnapshot(
Expand Down Expand Up @@ -514,13 +522,27 @@ def test_context_attributes_are_deeply_immutable():

def test_runtime_denies_provider_before_credential_resolution(monkeypatch):
credential_lookup_called = False
class_lookup_called = False
runtime_imports = []
real_import = builtins.__import__

def _credential_lookup(*_args, **_kwargs):
nonlocal credential_lookup_called
credential_lookup_called = True
return "secret"

def _track_import(name, *args, **kwargs):
runtime_imports.append(name)
return real_import(name, *args, **kwargs)

def _class_lookup(*_args, **_kwargs):
nonlocal class_lookup_called
class_lookup_called = True
return object

monkeypatch.setattr("lfx.base.models.unified_models.get_api_key_for_provider", _credential_lookup)
monkeypatch.setattr("lfx.base.models.unified_models.get_model_class", _class_lookup)
monkeypatch.setattr(builtins, "__import__", _track_import)

with pytest.raises(ModelProviderPolicyError) as exc_info:
get_llm(
Expand All @@ -531,17 +553,34 @@ def _credential_lookup(*_args, **_kwargs):

assert exc_info.value.code == "policy_blocked"
assert credential_lookup_called is False
assert class_lookup_called is False
assert "langchain_core.language_models" not in runtime_imports
assert not any(name.startswith("langchain_anthropic") for name in runtime_imports)


def test_embedding_runtime_denies_provider_before_credential_resolution(monkeypatch):
credential_lookup_called = False
class_lookup_called = False
runtime_imports = []
real_import = builtins.__import__

def _credential_lookup(*_args, **_kwargs):
nonlocal credential_lookup_called
credential_lookup_called = True
return "secret"

def _track_import(name, *args, **kwargs):
runtime_imports.append(name)
return real_import(name, *args, **kwargs)

def _class_lookup(*_args, **_kwargs):
nonlocal class_lookup_called
class_lookup_called = True
return object

monkeypatch.setattr("lfx.base.models.unified_models.get_api_key_for_provider", _credential_lookup)
monkeypatch.setattr("lfx.base.models.unified_models.get_embedding_class", _class_lookup)
monkeypatch.setattr(builtins, "__import__", _track_import)

with pytest.raises(ModelProviderPolicyError):
get_embeddings(
Expand All @@ -551,6 +590,43 @@ def _credential_lookup(*_args, **_kwargs):
)

assert credential_lookup_called is False
assert class_lookup_called is False
assert "langchain_core.embeddings" not in runtime_imports
assert not any(name.startswith("langchain_anthropic") for name in runtime_imports)


@pytest.mark.parametrize("instantiate", [get_llm, get_embeddings])
def test_denied_extension_runtime_does_not_execute_catalog_loader(monkeypatch, instantiate):
from lfx.base.models.provider_registry import ProviderSpec, register_provider, unregister_provider

provider = "Denied Runtime Extension"
register_provider(
ProviderSpec(
name=provider,
provider_id="denied-runtime-extension",
metadata={
"icon": "Bot",
"variables": [],
"mapping": {"model_class": "ChatOpenAI", "model_param": "model"},
},
catalog_loader=f"{__name__}:_denied_runtime_catalog_loader",
)
)
service = ModelProviderPolicyService()
service.set_approved_provider_ids({"openai"})
monkeypatch.setattr("lfx.services.deps.get_model_provider_policy_service", lambda: service)
_DENIED_RUNTIME_CATALOG_CALLS.clear()

try:
with pytest.raises(ModelProviderPolicyError):
instantiate(
[{"name": "blocked-model", "provider": provider, "metadata": {}}],
user_id="user-1",
)
finally:
unregister_provider(provider)

assert _DENIED_RUNTIME_CATALOG_CALLS == []


@pytest.mark.parametrize("runtime", [get_llm, get_embeddings], ids=["llm", "embeddings"])
Expand Down
Loading