From 22d327d21fd8f70283760a98d0edac33700bdf94 Mon Sep 17 00:00:00 2001 From: alex s <46074070+bearsyankees@users.noreply.github.com> Date: Fri, 10 Jul 2026 18:36:33 -0400 Subject: [PATCH] feat(settings): add force_required_tool_choice to LlmSettings (#730) feat(inputs): implement logic for required tool choice based on model test(inputs): add tests for force_required_tool_choice behavior test(runner): update tests to include force_required_tool_choice in settings --- strix/config/settings.py | 4 ++++ strix/core/inputs.py | 14 +++++++++++++- strix/core/runner.py | 1 + tests/test_config_loader.py | 1 + tests/test_inputs.py | 25 ++++++++++++++++++++++++- tests/test_runner_rate_limit.py | 6 +++++- 6 files changed, 48 insertions(+), 3 deletions(-) diff --git a/strix/config/settings.py b/strix/config/settings.py index 91fbdef1..9d16e083 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -36,6 +36,10 @@ class LlmSettings(BaseSettings): ), ) reasoning_effort: ReasoningEffort = Field(default="high", alias="STRIX_REASONING_EFFORT") + force_required_tool_choice: bool = Field( + default=False, + alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE", + ) timeout: int = Field(default=300, alias="LLM_TIMEOUT") diff --git a/strix/core/inputs.py b/strix/core/inputs.py index 2bcc077d..f723c52b 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -8,7 +8,11 @@ from typing import TYPE_CHECKING, Any from agents.model_settings import ModelSettings from openai.types.shared import Reasoning -from strix.config.models import DEFAULT_MODEL_RETRY, model_supports_reasoning +from strix.config.models import ( + DEFAULT_MODEL_RETRY, + is_known_openai_bare_model, + model_supports_reasoning, +) if TYPE_CHECKING: @@ -18,6 +22,11 @@ if TYPE_CHECKING: DEFAULT_MAX_TURNS = 500 +def _accepts_required_tool_choice(model_name: str | None) -> bool: + name = (model_name or "").strip().lower() + return name.startswith("openai/") or is_known_openai_bare_model(name) + + def build_root_task(scan_config: dict[str, Any]) -> str: targets = scan_config.get("targets", []) or [] diff_scope = scan_config.get("diff_scope") or {} @@ -111,6 +120,7 @@ def make_model_settings( reasoning_effort: ReasoningEffort | None, *, model_name: str, + force_required_tool_choice: bool = False, ) -> ModelSettings: model_settings = ModelSettings( parallel_tool_calls=False, @@ -125,6 +135,8 @@ def make_model_settings( model_settings = model_settings.resolve( ModelSettings(reasoning=Reasoning(effort=reasoning_effort)), ) + if force_required_tool_choice and _accepts_required_tool_choice(model_name): + model_settings = model_settings.resolve(ModelSettings(tool_choice="required")) return model_settings diff --git a/strix/core/runner.py b/strix/core/runner.py index 90becf87..13052b6c 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -158,6 +158,7 @@ async def run_strix_scan( model_settings = make_model_settings( settings.llm.reasoning_effort, model_name=resolved_model, + force_required_tool_choice=settings.llm.force_required_tool_choice, ) run_config = RunConfig( model=resolved_model, diff --git a/tests/test_config_loader.py b/tests/test_config_loader.py index 56e644f3..be0bfb8f 100644 --- a/tests/test_config_loader.py +++ b/tests/test_config_loader.py @@ -26,6 +26,7 @@ _LLM_ENV_KEYS = [ "LITELLM_BASE_URL", "OLLAMA_API_BASE", "STRIX_REASONING_EFFORT", + "STRIX_FORCE_REQUIRED_TOOL_CHOICE", "LLM_TIMEOUT", "PERPLEXITY_API_KEY", # RuntimeSettings diff --git a/tests/test_inputs.py b/tests/test_inputs.py index c8fcf601..ca1ba0d9 100644 --- a/tests/test_inputs.py +++ b/tests/test_inputs.py @@ -7,7 +7,7 @@ from typing import Any import pytest -from strix.core.inputs import build_root_task, child_initial_input +from strix.core.inputs import build_root_task, child_initial_input, make_model_settings def _child_kwargs(parent_history: list[Any]) -> dict[str, Any]: @@ -112,3 +112,26 @@ def test_build_root_task_diff_scope() -> None: assert "Scope Constraints:" in task assert "3 changed file(s)" in task assert "2 deleted file(s)" in task + + +@pytest.mark.parametrize("model_name", ["openai/o3", "gpt-4o"]) +def test_make_model_settings_forces_required_tool_choice_for_openai_models( + model_name: str, +) -> None: + settings = make_model_settings( + "none", + model_name=model_name, + force_required_tool_choice=True, + ) + + assert settings.tool_choice == "required" + + +def test_make_model_settings_skips_required_tool_choice_for_non_openai_models() -> None: + settings = make_model_settings( + "none", + model_name="anthropic/claude-3-7-sonnet-latest", + force_required_tool_choice=True, + ) + + assert settings.tool_choice is None diff --git a/tests/test_runner_rate_limit.py b/tests/test_runner_rate_limit.py index ed28d250..7e708eb5 100644 --- a/tests/test_runner_rate_limit.py +++ b/tests/test_runner_rate_limit.py @@ -33,7 +33,11 @@ async def test_persistent_rate_limit_stops_gracefully( monkeypatch.setattr(runner, "set_scan_id", lambda _scan_id: None) settings = types.SimpleNamespace( - llm=types.SimpleNamespace(model="openai/gpt-4o", reasoning_effort="high") + llm=types.SimpleNamespace( + model="openai/gpt-4o", + reasoning_effort="high", + force_required_tool_choice=False, + ) ) monkeypatch.setattr(runner, "load_settings", lambda: settings) monkeypatch.setattr(runner, "configure_sdk_model_defaults", lambda _settings: None)