From 36b47d84c0f96467fdc3adacc92d0230e61272b7 Mon Sep 17 00:00:00 2001 From: Ousama Ben Younes Date: Mon, 27 Jul 2026 19:12:05 +0000 Subject: [PATCH] feat(config): add per-request token budget --- docs/advanced/configuration.mdx | 5 +++++ strix/config/settings.py | 1 + strix/core/inputs.py | 2 ++ strix/core/runner.py | 1 + tests/test_config_loader.py | 7 ++++++- tests/test_inputs.py | 10 ++++++++++ tests/test_runner_rate_limit.py | 1 + tests/test_runner_root_prompt.py | 25 ++++++++++++++++++++++++- 8 files changed, 50 insertions(+), 2 deletions(-) diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index af98b8b87..d20bc4938 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -31,6 +31,11 @@ Configure Strix using environment variables or a config file. Request timeout in seconds for LLM calls. + + Optional maximum output tokens for each agent LLM request. Leave unset to use + the provider/model default. + + Maximum number of retries for LLM API calls on transient failures. diff --git a/strix/config/settings.py b/strix/config/settings.py index e53d125ce..cca9486e6 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -52,6 +52,7 @@ class LlmSettings(BaseSettings): default=False, alias="LLM_DISABLE_STREAMING", ) + max_tokens: int | None = Field(default=None, gt=0, alias="STRIX_LLM_MAX_TOKENS") timeout: int = Field(default=300, alias="LLM_TIMEOUT") diff --git a/strix/core/inputs.py b/strix/core/inputs.py index 34a2d4b30..9d2df18c1 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -133,11 +133,13 @@ def make_model_settings( request_timeout: float | None = None, prompt_cache: bool = True, extra_headers: dict[str, str] | None = None, + max_tokens: int | None = None, ) -> ModelSettings: model_settings = ModelSettings( parallel_tool_calls=False, retry=DEFAULT_MODEL_RETRY, include_usage=True, + max_tokens=max_tokens, extra_args=request_timeout_extra_args(request_timeout), extra_headers=dict(extra_headers) if extra_headers else None, ) diff --git a/strix/core/runner.py b/strix/core/runner.py index 01725cabb..e2d23df15 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -251,6 +251,7 @@ async def _spill_to_workspace(output_id: str, text: str) -> str | None: request_timeout=settings.llm.timeout, prompt_cache=settings.llm.prompt_cache, extra_headers=settings.llm.extra_headers, + max_tokens=settings.llm.max_tokens, ) run_config = RunConfig( model=resolved_model, diff --git a/tests/test_config_loader.py b/tests/test_config_loader.py index 7e14662ae..4fc6eb37b 100644 --- a/tests/test_config_loader.py +++ b/tests/test_config_loader.py @@ -10,7 +10,7 @@ from pydantic.fields import FieldInfo from strix.config import loader -from strix.config.settings import ContextSettings +from strix.config.settings import ContextSettings, LlmSettings if TYPE_CHECKING: @@ -28,6 +28,7 @@ "OLLAMA_API_BASE", "STRIX_REASONING_EFFORT", "STRIX_FORCE_REQUIRED_TOOL_CHOICE", + "STRIX_LLM_MAX_TOKENS", "LLM_TIMEOUT", "PERPLEXITY_API_KEY", # RuntimeSettings @@ -130,6 +131,10 @@ def test_tool_output_max_bytes_accepts_floor() -> None: assert ContextSettings(STRIX_TOOL_OUTPUT_MAX_BYTES=1024).tool_output_max_bytes == 1024 +def test_llm_max_tokens_env_alias() -> None: + assert LlmSettings(STRIX_LLM_MAX_TOKENS=12_000).max_tokens == 12_000 + + # --------------------------------------------------------------------------- # # _aliases_for # --------------------------------------------------------------------------- # diff --git a/tests/test_inputs.py b/tests/test_inputs.py index 871ea1490..04a6d62fb 100644 --- a/tests/test_inputs.py +++ b/tests/test_inputs.py @@ -266,6 +266,16 @@ def test_make_model_settings_sets_request_timeout() -> None: assert settings.extra_args["timeout"] == 300.0 +def test_make_model_settings_sets_configured_token_budget() -> None: + settings = make_model_settings( + "none", + model_name="gpt-4o", + max_tokens=12_000, + ) + + assert settings.max_tokens == 12_000 + + def test_make_model_settings_omits_timeout_when_unset() -> None: settings = make_model_settings("none", model_name="gpt-4o") diff --git a/tests/test_runner_rate_limit.py b/tests/test_runner_rate_limit.py index 061ad3c55..46c56e6a6 100644 --- a/tests/test_runner_rate_limit.py +++ b/tests/test_runner_rate_limit.py @@ -41,6 +41,7 @@ async def test_persistent_rate_limit_stops_gracefully( timeout=300, prompt_cache=True, extra_headers=None, + max_tokens=None, ), runtime=types.SimpleNamespace(max_context_images=3), ) diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index cd4d4ac84..f17fcac6c 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -49,6 +49,7 @@ def _patch_engine_scaffold( timeout=300, prompt_cache=True, extra_headers=None, + max_tokens=12_000, ), runtime=types.SimpleNamespace(max_context_images=3), ) @@ -74,10 +75,15 @@ async def _cleanup(*_args: Any, **_kwargs: Any) -> None: monkeypatch.setattr(runner, "build_root_task", lambda _scan_config: "task") monkeypatch.setattr(runner, "build_scope_context", lambda _scan_config: scope_context) - monkeypatch.setattr(runner, "make_model_settings", lambda *_args, **_kwargs: object()) captured: dict[str, Any] = {} + def _make_model_settings(*_args: Any, **kwargs: Any) -> object: + captured["model_settings_kwargs"] = kwargs + return object() + + monkeypatch.setattr(runner, "make_model_settings", _make_model_settings) + def _build_strix_agent(**kwargs: Any) -> object: if kwargs.get("is_root") and "kwargs" not in captured: captured["kwargs"] = kwargs @@ -176,3 +182,20 @@ async def test_root_prompt_options_default_to_none( kwargs = captured["kwargs"] assert kwargs["instructions_override"] is None assert kwargs["system_prompt_context"] == {"scope": "built-in"} + + +@pytest.mark.asyncio +async def test_llm_max_tokens_flows_into_model_settings( + monkeypatch: pytest.MonkeyPatch, + tmp_path: Any, +) -> None: + captured = _patch_engine_scaffold(monkeypatch, tmp_path, {"scope": "built-in"}) + + await runner.run_strix_scan( + scan_config={"targets": [], "scan_mode": "deep"}, + scan_id="scan-token-budget", + image="img", + coordinator=AgentCoordinator(), + ) + + assert captured["model_settings_kwargs"]["max_tokens"] == 12_000