diff --git a/agent/auxiliary_client.py b/agent/auxiliary_client.py index e3c87068f1..a85fb25731 100644 --- a/agent/auxiliary_client.py +++ b/agent/auxiliary_client.py @@ -931,6 +931,18 @@ def _fast_model_from_catalog(provider_id: str) -> str: return "" +def _nous_policy_blocks(model_id: str) -> bool: + """True when the org's model policy does not admit *model_id*.""" + try: + from hermes_cli.models import nous_policy_allowed_ids, restrict_to_nous_policy + + allowed = nous_policy_allowed_ids() + return bool(allowed) and not restrict_to_nous_policy([model_id], allowed) + except Exception: + logger.debug("Nous policy check unavailable", exc_info=True) + return False + + # Default auxiliary models for direct API-key providers (cheap/fast for side tasks) def _get_aux_model_for_provider(provider_id: str, *, prefer_fast: bool = False) -> str: """Return the cheap auxiliary model for a provider. @@ -958,21 +970,26 @@ def _get_aux_model_for_provider(provider_id: str, *, prefer_fast: bool = False) except Exception: pass + picked = "" if prefer_fast: - catalog_pick = _fast_model_from_catalog(provider_id) - if catalog_pick: - return catalog_pick - if profile is not None: + picked = _fast_model_from_catalog(provider_id) + if not picked and profile is not None: try: - live = profile.resolve_aux_model() - if live: - return live + picked = profile.resolve_aux_model() or "" except Exception: logger.debug("resolve_aux_model failed for %s", provider_id, exc_info=True) - if profile is not None and profile.default_aux_model: - return profile.default_aux_model - return _API_KEY_PROVIDER_AUX_MODELS_FALLBACK.get(provider_id, "") + if not picked and profile is not None and profile.default_aux_model: + picked = profile.default_aux_model + if not picked: + picked = _API_KEY_PROVIDER_AUX_MODELS_FALLBACK.get(provider_id, "") + + # Steps 2-4 are policy-blind: resolve_aux_model queries a public + # recommendation and the rest are hardcoded. A blocked pick is refused at + # request time, so drop it and let the caller keep the main model. + if picked and provider_id.strip().lower() == "nous" and _nous_policy_blocks(picked): + return "" + return picked diff --git a/tests/hermes_cli/test_nous_policy_surfaces.py b/tests/hermes_cli/test_nous_policy_surfaces.py index 1fd43725d3..8e30f0efc8 100644 --- a/tests/hermes_cli/test_nous_policy_surfaces.py +++ b/tests/hermes_cli/test_nous_policy_surfaces.py @@ -229,3 +229,58 @@ class TestPolicyNoticeIsShown: monkeypatch.setattr(account_mod, "nous_policy_present", lambda: False) TestLoginNous()._run(monkeypatch, tmp_path) assert "restricts which models" not in capsys.readouterr().out + + +class TestAuxFallbackRespectsPolicy: + """Steps 2-4 of the aux ladder are policy-blind: `resolve_aux_model` queries + a public recommendation and the rest are hardcoded.""" + + def _patch(self, monkeypatch, *, allowed, recommended): + import agent.auxiliary_client as aux + import providers + + monkeypatch.setattr(models_mod, "nous_policy_allowed_ids", lambda **_k: allowed) + monkeypatch.setattr( + models_mod, "_resolve_nous_pricing_credentials", + lambda: ("sk", "https://inference.example.com"), + ) + # No fast-family match, so the catalog step yields nothing. + monkeypatch.setattr( + models_mod, "fetch_models_with_pricing", + lambda **_k: {"vendor/allowed-large": {}}, + ) + + class _Profile: + default_aux_model = "" + + def resolve_aux_model(self, **_k): + return recommended + + monkeypatch.setattr(providers, "get_provider_profile", lambda _p: _Profile()) + return aux + + def test_blocked_recommendation_is_not_used(self, monkeypatch): + aux = self._patch( + monkeypatch, allowed={"vendor/allowed-large"}, + recommended="vendor/blocked-haiku", + ) + assert aux._get_aux_model_for_provider("nous", prefer_fast=True) == "" + + def test_allowed_recommendation_still_used(self, monkeypatch): + aux = self._patch( + monkeypatch, allowed={"vendor/allowed-large", "vendor/ok-haiku"}, + recommended="vendor/ok-haiku", + ) + assert ( + aux._get_aux_model_for_provider("nous", prefer_fast=True) + == "vendor/ok-haiku" + ) + + def test_ungoverned_org_is_unaffected(self, monkeypatch): + aux = self._patch( + monkeypatch, allowed=None, recommended="vendor/anything" + ) + assert ( + aux._get_aux_model_for_provider("nous", prefer_fast=True) + == "vendor/anything" + )