diff --git a/strix/config/models.py b/strix/config/models.py index a8598be1..88349ad6 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -18,7 +18,7 @@ from agents import ( ) from agents.model_settings import ModelSettings from agents.models.fake_id import FAKE_RESPONSES_ID -from agents.models.interface import Model +from agents.models.interface import Model, ModelProvider from agents.models.multi_provider import MultiProvider from agents.models.openai_responses import OpenAIResponsesModel from agents.retry import ( @@ -48,7 +48,7 @@ if TYPE_CHECKING: from agents.agent_output import AgentOutputSchemaBase from agents.handoffs import Handoff from agents.items import ModelResponse, TResponseInputItem, TResponseStreamEvent - from agents.models.interface import ModelProvider, ModelTracing + from agents.models.interface import ModelTracing from agents.retry import ModelRetryAdvice, ModelRetryAdviceRequest from agents.tool import Tool from agents.usage import Usage @@ -445,12 +445,61 @@ def _response_usage(usage: Usage | None) -> ResponseUsage | None: ) +class _CredentialedLitellmProvider(ModelProvider): + """LiteLLM route bound to one endpoint's credentials. + + ``LitellmProvider`` reads them from the process-wide LiteLLM globals, which + belong to the main model; a secondary endpoint needs its own. + """ + + def __init__(self, api_key: str | None, base_url: str | None) -> None: + self._api_key = api_key + self._base_url = base_url + + def get_model(self, model_name: str | None) -> Model: + from agents.extensions.models.litellm_model import LitellmModel + from agents.models.default_models import get_default_model + + return LitellmModel( + model=model_name or get_default_model(), + api_key=self._api_key, + base_url=self._base_url, + ) + + class StrixProvider(MultiProvider): """Route any non-OpenAI prefix through LiteLLM with the prefix preserved, so users type ``deepseek/deepseek-chat`` rather than ``litellm/deepseek/deepseek-chat``. + + ``api_key``/``base_url`` bind every route this provider resolves to one + endpoint, for a secondary model (the dedupe judge) whose endpoint differs + from the main model's process-wide defaults. """ + def __init__( + self, + *, + api_key: str | None = None, + base_url: str | None = None, + **kwargs: Any, + ) -> None: + super().__init__( + openai_api_key=api_key, + openai_base_url=base_url, + # A custom endpoint is OpenAI-compatible, i.e. chat completions; the + # global default is the main model's and may say otherwise. + openai_use_responses=False if base_url else None, + **kwargs, + ) + self._override_api_key = api_key + self._override_base_url = base_url + + def _create_fallback_provider(self, prefix: str) -> ModelProvider: + if prefix == "litellm" and (self._override_api_key or self._override_base_url): + return _CredentialedLitellmProvider(self._override_api_key, self._override_base_url) + return super()._create_fallback_provider(prefix) + def _resolve_prefixed_model( self, *, diff --git a/strix/core/inputs.py b/strix/core/inputs.py index 1f50bce8..3dd0d701 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -268,7 +268,7 @@ def make_model_settings( and model_supports_reasoning(model_name) ): model_settings = model_settings.resolve( - _reasoning_settings(reasoning_effort, model_settings.extra_args), + _reasoning_settings(reasoning_effort), ) if force_required_tool_choice and _accepts_required_tool_choice(model_name): model_settings = model_settings.resolve(ModelSettings(tool_choice="required")) @@ -294,20 +294,19 @@ def _request_headers( return headers or None -def _reasoning_settings( - effort: ReasoningEffort, - extra_args: dict[str, Any] | None, -) -> ModelSettings: +def _reasoning_settings(effort: ReasoningEffort) -> ModelSettings: """``max`` is not in the OpenAI SDK's ``Reasoning.effort`` enum, so send it as a raw body field instead — also keeping it clear of LiteLLM's DeepSeek mapping, which collapses every ``reasoning_effort`` level to plain thinking-enabled. Providers that don't support ``max`` reject the request. + + It goes in ``extra_body``, the field every model implementation forwards as the + request's ``extra_body``; the same value under ``extra_args`` collides with that + keyword and raises before a request is ever sent. """ if effort != "max": return ModelSettings(reasoning=Reasoning(effort=effort)) - return ModelSettings( - extra_args={**(extra_args or {}), "extra_body": {"reasoning_effort": "max"}}, - ) + return ModelSettings(extra_body={"reasoning_effort": "max"}) def _prompt_cache_extra_args(model_name: str) -> dict[str, Any] | None: diff --git a/strix/interface/main.py b/strix/interface/main.py index 96459978..d1c29ea0 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -127,12 +127,10 @@ def _subscription_error_hint(exc: BaseException) -> str | None: async def warm_up_llm(show_model_warning: bool = True) -> None: - from agents.model_settings import ModelSettings from agents.models.interface import ModelTracing from strix.config.models import ( RECOMMENDED_MODEL_NAMES, - StrixProvider, configure_sdk_model_defaults, is_known_openai_bare_model, is_recommended_or_frontier_model, @@ -209,12 +207,11 @@ async def warm_up_llm(show_model_warning: bool = True) -> None: logger.info("LLM warm-up succeeded for model %s", (llm.model or "").strip()) if settings.dedupe.model: - from strix.report.dedupe import _dedupe_extra_args + from strix.report.dedupe import resolve_dedupe_model dedupe_model = settings.dedupe.model.strip() raw_model = dedupe_model - deduper = StrixProvider().get_model(dedupe_model) - deduper_extra = _dedupe_extra_args(settings.dedupe) + deduper = resolve_dedupe_model(settings.dedupe, dedupe_model) # A dedicated dedupe model may route to another provider, which must # never receive the main endpoint's headers; it has its own # DEDUPE_LLM_EXTRA_HEADERS. @@ -226,9 +223,6 @@ async def warm_up_llm(show_model_warning: bool = True) -> None: extra_headers=settings.dedupe.extra_headers, has_tools=False, ) - if deduper_extra: - merged = {**(deduper_settings.extra_args or {}), **deduper_extra} - deduper_settings = deduper_settings.resolve(ModelSettings(extra_args=merged)) await asyncio.wait_for( deduper.get_response( system_instructions="You are a helpful assistant.", diff --git a/strix/report/dedupe.py b/strix/report/dedupe.py index 1cc0a66a..93f9ad8a 100644 --- a/strix/report/dedupe.py +++ b/strix/report/dedupe.py @@ -7,7 +7,6 @@ import logging import re from typing import TYPE_CHECKING, Any -from agents.model_settings import ModelSettings from agents.models.interface import ModelTracing from openai.types.responses import ResponseOutputMessage @@ -22,6 +21,8 @@ from strix.report.state import get_global_report_state if TYPE_CHECKING: from agents.items import ModelResponse + from agents.model_settings import ModelSettings + from agents.models.interface import Model from strix.config.settings import DedupeSettings @@ -29,30 +30,11 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -def _dedupe_extra_args(dedupe: DedupeSettings) -> dict[str, str]: - """Per-call credential + endpoint for the dedupe model. - - Provider env vars and the global base URL are process-wide, so a - shared-provider dedupe key or a distinct dedupe endpoint can't be installed - globally without clobbering (or being clobbered by) the main model's - config. Passing them per call keeps the two apart. Only applies when a - dedicated dedupe model is configured. - """ - if not dedupe.model: - return {} - extra: dict[str, str] = {} - if dedupe.api_key and dedupe.api_key.strip(): - extra["api_key"] = dedupe.api_key.strip() - if dedupe.api_base and dedupe.api_base.strip(): - extra["api_base"] = dedupe.api_base.strip() - return extra - - def _dedupe_model_settings( dedupe: DedupeSettings, model_name: str, request_timeout: float | None ) -> ModelSettings: llm = load_settings().llm - settings = make_model_settings( + return make_model_settings( dedupe.reasoning_effort, model_name=model_name, force_required_tool_choice=False, @@ -64,10 +46,21 @@ def _dedupe_model_settings( extra_headers=dedupe.extra_headers if dedupe.model else llm.extra_headers, has_tools=False, ) - extra = _dedupe_extra_args(dedupe) - if extra: - settings = settings.resolve(ModelSettings(extra_args=extra)) - return settings + + +def resolve_dedupe_model(dedupe: DedupeSettings, model_name: str) -> Model: + """Resolve the dedupe model, bound to its own endpoint when it has one. + + Credentials can't ride on the request: every model implementation already + passes its own ``api_key``/``base_url``, so the same keys in ``extra_args`` + collide with them and raise before anything is sent. A provider bound to the + dedupe endpoint keeps it apart from the main model's process-wide defaults. + """ + api_key = (dedupe.api_key or "").strip() if dedupe.model else "" + api_base = (dedupe.api_base or "").strip() if dedupe.model else "" + if not (api_key or api_base): + return StrixProvider().get_model(model_name) + return StrixProvider(api_key=api_key or None, base_url=api_base or None).get_model(model_name) DEDUPE_SYSTEM_PROMPT = """You are an expert vulnerability report deduplication judge. @@ -371,7 +364,7 @@ async def check_duplicate( configure_sdk_model_defaults(settings) resolved_model = model_name.strip() - model = StrixProvider().get_model(resolved_model) + model = resolve_dedupe_model(dedupe, resolved_model) response = await model.get_response( system_instructions=DEDUPE_SYSTEM_PROMPT, input=user_msg, diff --git a/tests/test_dedupe_model.py b/tests/test_dedupe_model.py index b17946e9..547f434d 100644 --- a/tests/test_dedupe_model.py +++ b/tests/test_dedupe_model.py @@ -7,7 +7,7 @@ from typing import TYPE_CHECKING from strix.config import loader from strix.config.settings import DedupeSettings -from strix.report.dedupe import _dedupe_model_settings +from strix.report.dedupe import _dedupe_model_settings, resolve_dedupe_model if TYPE_CHECKING: @@ -16,32 +16,49 @@ if TYPE_CHECKING: import pytest -def test_dedupe_key_sent_per_call_not_via_global_env() -> None: +def _unwrap(model: object) -> object: + while hasattr(model, "_inner"): + model = model._inner + return model + + +def test_dedupe_key_bound_to_model_client_not_global_env() -> None: dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap", DEDUPE_LLM_API_KEY="dedupe-key") - settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300) - # The key rides on the request, so a shared-provider main key can't clobber - # it (and vice versa) through the global provider env var. - assert (settings.extra_args or {})["api_key"] == "dedupe-key" + model = _unwrap(resolve_dedupe_model(dedupe, "deepseek/cheap")) + # The key is bound to the dedupe model's own client, so a shared-provider + # main key can't clobber it (and vice versa) through the process globals — + # and it never rides on the request, where every model implementation's own + # api_key kwarg would collide with it. + assert model.api_key == "dedupe-key" # type: ignore[attr-defined] -def test_dedupe_settings_omit_api_key_when_unset() -> None: - dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap") +def test_dedupe_settings_carry_no_request_credentials() -> None: + dedupe = DedupeSettings( + STRIX_DEDUPE_MODEL="deepseek/cheap", + DEDUPE_LLM_API_KEY="dedupe-key", + DEDUPE_LLM_API_BASE="https://dedupe.example/v1", + ) settings = _dedupe_model_settings(dedupe, "deepseek/cheap", 300) assert "api_key" not in (settings.extra_args or {}) assert "api_base" not in (settings.extra_args or {}) -def test_dedupe_endpoint_sent_per_call() -> None: +def test_dedupe_endpoint_bound_to_model_client() -> None: dedupe = DedupeSettings( STRIX_DEDUPE_MODEL="openai/cheap", DEDUPE_LLM_API_KEY="dedupe-key", DEDUPE_LLM_API_BASE="https://dedupe.example/v1", ) - settings = _dedupe_model_settings(dedupe, "openai/cheap", 300) - # A distinct dedupe endpoint rides on the request instead of the - # process-wide base URL, so it can't clobber the main model's endpoint. - assert (settings.extra_args or {})["api_base"] == "https://dedupe.example/v1" - assert (settings.extra_args or {})["api_key"] == "dedupe-key" + model = _unwrap(resolve_dedupe_model(dedupe, "openai/cheap")) + client = model._client # type: ignore[attr-defined] + assert client.api_key == "dedupe-key" + assert str(client.base_url).startswith("https://dedupe.example/v1") + + +def test_dedupe_without_credentials_uses_default_provider() -> None: + dedupe = DedupeSettings(STRIX_DEDUPE_MODEL="deepseek/cheap") + model = _unwrap(resolve_dedupe_model(dedupe, "deepseek/cheap")) + assert model.api_key is None # type: ignore[attr-defined] def test_dedicated_dedupe_model_uses_own_headers_not_main() -> None: diff --git a/tests/test_inputs.py b/tests/test_inputs.py index 76cc6bea..5a483edf 100644 --- a/tests/test_inputs.py +++ b/tests/test_inputs.py @@ -153,7 +153,8 @@ def test_max_reasoning_effort_sent_as_raw_body_field() -> None: "max", model_name="deepseek/deepseek-v4-flash", request_timeout=30 ) assert settings.reasoning is None - assert settings.extra_args == {"timeout": 30, "extra_body": {"reasoning_effort": "max"}} + assert settings.extra_args == {"timeout": 30} + assert settings.extra_body == {"reasoning_effort": "max"} def test_conversation_tail_breakpoint_moves_with_appended_transcript() -> None: