diff --git a/docs/advanced/configuration.mdx b/docs/advanced/configuration.mdx index a444b497..61128d6b 100644 --- a/docs/advanced/configuration.mdx +++ b/docs/advanced/configuration.mdx @@ -37,6 +37,11 @@ Configure Strix using environment variables or a config file. Request timeout in seconds for LLM calls. + + Seconds the startup connection check waits for the model to answer before + failing with `LLM CONNECTION FAILED`. Scan requests use `LLM_TIMEOUT`. + + Maximum number of retries for LLM API calls on transient failures. diff --git a/strix/config/settings.py b/strix/config/settings.py index d0474612..54243b74 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -77,6 +77,7 @@ class LlmSettings(BaseSettings): alias="LLM_DISABLE_STREAMING", ) timeout: int = Field(default=300, alias="LLM_TIMEOUT") + preflight_timeout: int = Field(default=30, ge=1, alias="LLM_PREFLIGHT_TIMEOUT") stream_idle_timeout: int = Field(default=300, ge=0, alias="LLM_STREAM_IDLE_TIMEOUT") max_tool_calls_per_turn: int = Field( default=32, diff --git a/strix/interface/main.py b/strix/interface/main.py index 1ebccffb..18f9113d 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -36,6 +36,7 @@ from strix.interface.interactive import ( from strix.interface.scan_setup import ( ModelConnectionError, preflight_model_connection, + preflight_request, prepare_run, telemetry_start, ) @@ -106,13 +107,10 @@ def _subscription_error_hint(exc: BaseException) -> str | None: async def warm_up_llm() -> None: - from agents.models.interface import ModelTracing - from strix.config.models import ( configure_sdk_model_defaults, is_known_openai_bare_model, ) - from strix.core.inputs import make_model_settings console = Console() logger.info("Warming up LLM connection") @@ -166,28 +164,11 @@ async def warm_up_llm() -> None: # 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. - deduper_settings = make_model_settings( - None, + await preflight_request( + deduper, model_name=dedupe_model, - request_timeout=llm.timeout, - prompt_cache=False, extra_headers=settings.dedupe.extra_headers, - has_tools=False, - ) - await asyncio.wait_for( - deduper.get_response( - system_instructions="You are a helpful assistant.", - input="Reply with just 'OK'.", - model_settings=deduper_settings, - tools=[], - output_schema=None, - handoffs=[], - tracing=ModelTracing.DISABLED, - previous_response_id=None, - conversation_id=None, - prompt=None, - ), - timeout=llm.timeout, + timeout=llm.preflight_timeout, ) logger.info("LLM warm-up succeeded for dedupe model %s", dedupe_model) diff --git a/strix/interface/scan_setup.py b/strix/interface/scan_setup.py index 2b070bd0..d7d0aab6 100644 --- a/strix/interface/scan_setup.py +++ b/strix/interface/scan_setup.py @@ -46,6 +46,8 @@ from strix.utils.api_spec import ( if TYPE_CHECKING: import argparse + from agents.models.interface import Model + logger = logging.getLogger(__name__) HOST_GATEWAY_HOSTNAME = "host.docker.internal" @@ -104,38 +106,61 @@ async def preflight_model_connection( settings: Settings | None = None, ) -> None: """Verify the configured model route before starting a scan.""" - from agents.models.interface import ModelTracing - from strix.config.models import StrixProvider, configure_sdk_model_defaults - from strix.core.inputs import make_model_settings resolved_settings = load_settings() if settings is None else settings check_header_safe_credentials(resolved_settings) configure_sdk_model_defaults(resolved_settings) model = StrixProvider().get_model(model_name) + await preflight_request( + model, + model_name=model_name, + extra_headers=resolved_settings.llm.extra_headers, + timeout=resolved_settings.llm.preflight_timeout, + ) + + +async def preflight_request( + model: Model, + *, + model_name: str, + extra_headers: dict[str, str] | None, + timeout: int, +) -> None: + """Send one tiny request to ``model`` and fail if it does not answer in ``timeout`` seconds.""" + from agents.models.interface import ModelTracing + + from strix.core.inputs import make_model_settings + request_settings = make_model_settings( None, model_name=model_name, - request_timeout=resolved_settings.llm.timeout, + request_timeout=timeout, prompt_cache=False, - extra_headers=resolved_settings.llm.extra_headers, + extra_headers=extra_headers, has_tools=False, ) - await asyncio.wait_for( - model.get_response( - system_instructions="You are a helpful assistant.", - input="Reply with just 'OK'.", - model_settings=request_settings, - tools=[], - output_schema=None, - handoffs=[], - tracing=ModelTracing.DISABLED, - previous_response_id=None, - conversation_id=None, - prompt=None, - ), - timeout=resolved_settings.llm.timeout, - ) + try: + await asyncio.wait_for( + model.get_response( + system_instructions="You are a helpful assistant.", + input="Reply with just 'OK'.", + model_settings=request_settings, + tools=[], + output_schema=None, + handoffs=[], + tracing=ModelTracing.DISABLED, + previous_response_id=None, + conversation_id=None, + prompt=None, + ), + timeout=timeout, + ) + except TimeoutError: + raise TimeoutError( + f"{model_name} did not answer within {timeout}s " + "(LLM_PREFLIGHT_TIMEOUT). Check LLM_API_BASE and that the endpoint is reachable." + ) from None def build_targets_info(args: argparse.Namespace) -> None: diff --git a/tests/test_preflight_timeout.py b/tests/test_preflight_timeout.py new file mode 100644 index 00000000..28f40b3f --- /dev/null +++ b/tests/test_preflight_timeout.py @@ -0,0 +1,84 @@ +"""The startup model check has its own short timeout; scan requests keep LLM_TIMEOUT.""" + +from __future__ import annotations + +import asyncio +import time +from typing import TYPE_CHECKING, Any + +import pytest + +from strix.config import load_settings, loader +from strix.interface import scan_setup +from strix.interface.scan_setup import preflight_model_connection, preflight_request + + +if TYPE_CHECKING: + from pathlib import Path + + +def _fresh_settings(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + monkeypatch.setattr(loader, "_cached", None) + monkeypatch.setattr(loader, "_override", tmp_path / "no-cli-config.json") + + +class _NeverAnswers: + def __init__(self) -> None: + self.request_timeouts: list[float | None] = [] + + async def get_response(self, *, model_settings: Any, **_: Any) -> None: + self.request_timeouts.append((model_settings.extra_args or {}).get("timeout")) + await asyncio.sleep(3600) + + +def test_preflight_timeout_defaults_to_30s_and_is_separate_from_llm_timeout( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + monkeypatch.delenv("LLM_PREFLIGHT_TIMEOUT", raising=False) + monkeypatch.setenv("LLM_TIMEOUT", "600") + _fresh_settings(monkeypatch, tmp_path) + llm = load_settings().llm + assert llm.timeout == 600 + assert llm.preflight_timeout == 30 + + monkeypatch.setenv("LLM_PREFLIGHT_TIMEOUT", "7") + _fresh_settings(monkeypatch, tmp_path) + assert load_settings().llm.preflight_timeout == 7 + + +def test_preflight_request_fails_after_its_own_timeout_with_a_clear_message() -> None: + model = _NeverAnswers() + started = time.monotonic() + with pytest.raises(TimeoutError) as excinfo: + asyncio.run( + preflight_request( + model, # type: ignore[arg-type] + model_name="openai/gpt-4o", + extra_headers=None, + timeout=1, + ) + ) + assert time.monotonic() - started < 5 + message = str(excinfo.value) + assert "openai/gpt-4o did not answer within 1s" in message + assert "LLM_PREFLIGHT_TIMEOUT" in message + assert model.request_timeouts == [1] + + +def test_preflight_model_connection_uses_the_preflight_timeout( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + monkeypatch.setenv("STRIX_LLM", "openai/gpt-4o") + monkeypatch.setenv("LLM_API_KEY", "sk-test") + monkeypatch.setenv("LLM_TIMEOUT", "600") + monkeypatch.setenv("LLM_PREFLIGHT_TIMEOUT", "12") + _fresh_settings(monkeypatch, tmp_path) + seen: dict[str, Any] = {} + + async def fake_request(_model: Any, **kwargs: Any) -> None: + seen.update(kwargs) + + monkeypatch.setattr(scan_setup, "preflight_request", fake_request) + asyncio.run(preflight_model_connection("openai/gpt-4o", settings=load_settings())) + assert seen["timeout"] == 12 + assert seen["model_name"] == "openai/gpt-4o"