Skip to content

Commit 2eefb24

Browse files
committed
feat(models): fail closed on denied provider runtime
1 parent f7f2493 commit 2eefb24

2 files changed

Lines changed: 120 additions & 39 deletions

File tree

src/lfx/src/lfx/base/models/unified_models/instantiation.py

Lines changed: 44 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -90,33 +90,32 @@ def get_llm(
9090
overrides: dict[str, Any] | None = None,
9191
provider_policy: ModelProviderPolicySnapshot | None = None,
9292
) -> Any:
93-
# Resolve helpers via package namespace so tests patching
94-
# lfx.base.models.unified_models.<name> keep working.
95-
from lfx.base.models import unified_models as unified_models_module
96-
9793
# Coerce provider-specific string params (Message/Data may leak through StrInput)
9894
ollama_base_url = _to_str(ollama_base_url)
9995
watsonx_url = _to_str(watsonx_url)
10096
watsonx_project_id = _to_str(watsonx_project_id)
10197

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

112-
# Safely extract model configuration
113-
if not model or not isinstance(model, list) or len(model) == 0:
112+
if isinstance(model, BaseLanguageModel):
113+
return model
114+
except ImportError:
115+
pass
114116
msg = "A model selection is required"
115117
raise ValueError(msg)
116118

117-
# Extract the first model (only one expected)
118-
model = model[0]
119-
120119
# Extract model configuration from metadata
121120
model_name = model.get("name")
122121
provider = model.get("provider")
@@ -129,17 +128,21 @@ def get_llm(
129128
)
130129
raise ValueError(msg)
131130
if provider_policy is None:
131+
from lfx.base.models.provider_registry import get_registry_snapshot
132132
from lfx.services.model_provider_policy import ModelProviderPolicyPurpose, resolve_model_provider_policy
133133

134-
from .provider_queries import get_model_providers
135-
136134
provider_policy = resolve_model_provider_policy(
137135
user_id=user_id,
138-
providers=(*get_model_providers(), provider),
136+
providers=(*get_registry_snapshot().provider_ids, provider),
139137
purpose=ModelProviderPolicyPurpose.USE,
140138
)
141139
provider_policy.require(provider)
142140

141+
# Resolve helpers through the package namespace only after policy passes so
142+
# tests can patch lfx.base.models.unified_models.<name> and denied requests
143+
# cannot trigger runtime integration imports.
144+
from lfx.base.models import unified_models as unified_models_module
145+
143146
# Policy is evaluated against the submitted identity, then all runtime
144147
# wiring for a known provider is resolved from its canonical registry
145148
# entry. Stored flow metadata is not authoritative for class or parameter
@@ -864,30 +867,30 @@ def get_embeddings(
864867
wrapper containing the primary instance for the selected model and an
865868
``available_models`` map of enabled embedding models from every configured provider.
866869
"""
867-
# Resolve helpers via package namespace so tests patching
868-
# lfx.base.models.unified_models.<name> keep working.
869-
from lfx.base.models import unified_models as unified_models_module
870-
871870
# Coerce provider-specific string params
872871
ollama_base_url = _to_str(ollama_base_url)
873872
watsonx_url = _to_str(watsonx_url)
874873
watsonx_project_id = _to_str(watsonx_project_id)
875874

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

885-
# Validate input
886-
if not model or not isinstance(model, list) or len(model) == 0:
887+
if isinstance(model, BaseEmbeddings):
888+
return model
889+
except ImportError:
890+
pass
887891
msg = "An embedding model selection is required"
888892
raise ValueError(msg)
889893

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

901-
from .provider_queries import get_model_providers
902-
903905
provider_policy = resolve_model_provider_policy(
904906
user_id=user_id,
905-
providers=(*get_model_providers(), provider),
907+
providers=(*get_registry_snapshot().provider_ids, provider),
906908
purpose=ModelProviderPolicyPurpose.USE,
907909
)
908910
provider_policy.require(provider)
909911

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

912917
provider_id = resolve_provider_id(provider)

src/lfx/tests/unit/services/model_provider_policy/test_policy.py

Lines changed: 76 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from __future__ import annotations
22

3+
import builtins
34
from dataclasses import FrozenInstanceError
45
from unittest.mock import AsyncMock, patch
56

@@ -16,6 +17,13 @@
1617
set_current_model_provider_policy_context,
1718
)
1819

20+
_DENIED_RUNTIME_CATALOG_CALLS: list[None] = []
21+
22+
23+
def _denied_runtime_catalog_loader():
24+
_DENIED_RUNTIME_CATALOG_CALLS.append(None)
25+
return [{"name": "blocked-model", "model_type": "llm"}]
26+
1927

2028
def _restricted_snapshot(*allowed: str) -> ModelProviderPolicySnapshot:
2129
return ModelProviderPolicySnapshot(
@@ -514,13 +522,27 @@ def test_context_attributes_are_deeply_immutable():
514522

515523
def test_runtime_denies_provider_before_credential_resolution(monkeypatch):
516524
credential_lookup_called = False
525+
class_lookup_called = False
526+
runtime_imports = []
527+
real_import = builtins.__import__
517528

518529
def _credential_lookup(*_args, **_kwargs):
519530
nonlocal credential_lookup_called
520531
credential_lookup_called = True
521532
return "secret"
522533

534+
def _track_import(name, *args, **kwargs):
535+
runtime_imports.append(name)
536+
return real_import(name, *args, **kwargs)
537+
538+
def _class_lookup(*_args, **_kwargs):
539+
nonlocal class_lookup_called
540+
class_lookup_called = True
541+
return object
542+
523543
monkeypatch.setattr("lfx.base.models.unified_models.get_api_key_for_provider", _credential_lookup)
544+
monkeypatch.setattr("lfx.base.models.unified_models.get_model_class", _class_lookup)
545+
monkeypatch.setattr(builtins, "__import__", _track_import)
524546

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

532554
assert exc_info.value.code == "policy_blocked"
533555
assert credential_lookup_called is False
556+
assert class_lookup_called is False
557+
assert "langchain_core.language_models" not in runtime_imports
558+
assert not any(name.startswith("langchain_anthropic") for name in runtime_imports)
534559

535560

536561
def test_embedding_runtime_denies_provider_before_credential_resolution(monkeypatch):
537562
credential_lookup_called = False
563+
class_lookup_called = False
564+
runtime_imports = []
565+
real_import = builtins.__import__
538566

539567
def _credential_lookup(*_args, **_kwargs):
540568
nonlocal credential_lookup_called
541569
credential_lookup_called = True
542570
return "secret"
543571

572+
def _track_import(name, *args, **kwargs):
573+
runtime_imports.append(name)
574+
return real_import(name, *args, **kwargs)
575+
576+
def _class_lookup(*_args, **_kwargs):
577+
nonlocal class_lookup_called
578+
class_lookup_called = True
579+
return object
580+
544581
monkeypatch.setattr("lfx.base.models.unified_models.get_api_key_for_provider", _credential_lookup)
582+
monkeypatch.setattr("lfx.base.models.unified_models.get_embedding_class", _class_lookup)
583+
monkeypatch.setattr(builtins, "__import__", _track_import)
545584

546585
with pytest.raises(ModelProviderPolicyError):
547586
get_embeddings(
@@ -551,6 +590,43 @@ def _credential_lookup(*_args, **_kwargs):
551590
)
552591

553592
assert credential_lookup_called is False
593+
assert class_lookup_called is False
594+
assert "langchain_core.embeddings" not in runtime_imports
595+
assert not any(name.startswith("langchain_anthropic") for name in runtime_imports)
596+
597+
598+
@pytest.mark.parametrize("instantiate", [get_llm, get_embeddings])
599+
def test_denied_extension_runtime_does_not_execute_catalog_loader(monkeypatch, instantiate):
600+
from lfx.base.models.provider_registry import ProviderSpec, register_provider, unregister_provider
601+
602+
provider = "Denied Runtime Extension"
603+
register_provider(
604+
ProviderSpec(
605+
name=provider,
606+
provider_id="denied-runtime-extension",
607+
metadata={
608+
"icon": "Bot",
609+
"variables": [],
610+
"mapping": {"model_class": "ChatOpenAI", "model_param": "model"},
611+
},
612+
catalog_loader=f"{__name__}:_denied_runtime_catalog_loader",
613+
)
614+
)
615+
service = ModelProviderPolicyService()
616+
service.set_approved_provider_ids({"openai"})
617+
monkeypatch.setattr("lfx.services.deps.get_model_provider_policy_service", lambda: service)
618+
_DENIED_RUNTIME_CATALOG_CALLS.clear()
619+
620+
try:
621+
with pytest.raises(ModelProviderPolicyError):
622+
instantiate(
623+
[{"name": "blocked-model", "provider": provider, "metadata": {}}],
624+
user_id="user-1",
625+
)
626+
finally:
627+
unregister_provider(provider)
628+
629+
assert _DENIED_RUNTIME_CATALOG_CALLS == []
554630

555631

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

0 commit comments

Comments
 (0)