diff --git a/strix/config/models.py b/strix/config/models.py index 768a8f02..213bc5ff 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -59,6 +59,42 @@ DEFAULT_MODEL_RETRY = ModelRetrySettings( ), ) +RECOMMENDED_MODEL_NAMES = ( + "openai/gpt-5.6", + "openai/gpt-5.6-sol", + "openai/gpt-5.6-terra", + "openai/gpt-5.5", + "openai/gpt-5.5-pro", + "openai/gpt-5.4", + "openai/gpt-5.3-codex", + "anthropic/claude-fable-5", + "anthropic/claude-opus-4-8", + "anthropic/claude-opus-4-7", + "anthropic/claude-sonnet-5", + "anthropic/claude-sonnet-4-6", + "vertex_ai/gemini-3.1-pro-preview", + "gemini/gemini-3.1-pro-preview", + "deepseek/deepseek-v4-pro", + "deepseek/deepseek-v4-flash", + "dashscope/qwen3.7-max-2026-06-08", + "moonshot/kimi-k2.7-code", + "moonshot/kimi-k2.6", +) + +_RECOMMENDED_MODEL_NAME_SET = frozenset(name.lower() for name in RECOMMENDED_MODEL_NAMES) + +FRONTIER_MODEL_FAMILIES = ( + (("azure", "azure_ai", "bedrock_mantle", "openai"), ("gpt-5",)), + ( + ("anthropic", "azure_ai", "bedrock", "claude", "databricks", "snowflake", "vertex_ai"), + ("claude-fable-5", "claude-opus-4", "claude-sonnet-5", "claude-sonnet-4"), + ), + (("google", "gemini", "vertex_ai"), ("gemini-3",)), + (("deepseek",), ("deepseek-v4", "deepseek-r1", "deepseek-reasoner")), + (("alibaba", "dashscope", "qwen"), ("qwen3.7", "qwen3.5", "qwen3-max")), + (("moonshot", "moonshotai", "kimi"), ("kimi-k2.7", "kimi-k2.6", "kimi-k2.5")), +) + def configure_sdk_model_defaults(settings: Settings) -> None: """Apply Strix config to SDK-native defaults.""" @@ -180,6 +216,78 @@ def model_supports_reasoning(model_name: str) -> bool: return bool(entry and entry.get("supports_reasoning")) +def is_recommended_or_frontier_model(model_name: str) -> bool: + """Return whether a model is recommended or in a frontier model family.""" + name = _normalized_model_name(model_name) + if not name: + return False + if name in _RECOMMENDED_MODEL_NAME_SET: + return True + provider_name, bare_model_name = _split_model_provider(name) + return any( + _matches_frontier_family(provider_name, bare_model_name, provider_markers, prefixes) + for provider_markers, prefixes in FRONTIER_MODEL_FAMILIES + ) + + +def _normalized_model_name(model_name: str) -> str: + name = model_name.strip().lower() + for prefix in ("litellm/", "any-llm/"): + if name.startswith(prefix): + name = name[len(prefix) :] + break + return name + + +def _split_model_provider(model_name: str) -> tuple[str | None, str]: + if "/" not in model_name: + return None, model_name + provider_name, bare_model_name = model_name.rsplit("/", 1) + return provider_name, bare_model_name + + +def _matches_frontier_family( + provider_name: str | None, + model_name: str, + provider_markers: tuple[str, ...], + model_prefixes: tuple[str, ...], +) -> bool: + if not _matches_model_prefix(model_name, model_prefixes): + return False + if provider_name is None: + return True + return _contains_provider_marker( + provider_name, provider_markers, split_compound_names=True + ) or _contains_provider_marker(model_name, provider_markers) + + +def _matches_model_prefix(model_name: str, model_prefixes: tuple[str, ...]) -> bool: + return any( + candidate.startswith(prefix) + for candidate in _model_name_candidates(model_name) + for prefix in model_prefixes + ) + + +def _model_name_candidates(model_name: str) -> tuple[str, ...]: + if "." not in model_name: + return (model_name,) + suffixes = tuple( + model_name.split(".", index)[-1] for index in range(1, model_name.count(".") + 1) + ) + return (model_name, *suffixes) + + +def _contains_provider_marker( + value: str, provider_markers: tuple[str, ...], *, split_compound_names: bool = False +) -> bool: + parts = set(value.replace(".", "/").split("/")) + if split_compound_names: + for separator in ("_", "-"): + parts.update(piece for part in tuple(parts) for piece in part.split(separator)) + return any(marker in parts for marker in provider_markers) + + def is_known_openai_bare_model(model_name: str) -> bool: import litellm diff --git a/strix/core/hooks.py b/strix/core/hooks.py index 6b0d5924..64562d74 100644 --- a/strix/core/hooks.py +++ b/strix/core/hooks.py @@ -28,7 +28,10 @@ class ReportUsageHooks(RunHooks[dict[str, Any]]): def __init__(self, *, model: str, max_budget_usd: float | None = None) -> None: import math - if max_budget_usd is not None and (not math.isfinite(max_budget_usd) or max_budget_usd <= 0): + + if max_budget_usd is not None and ( + not math.isfinite(max_budget_usd) or max_budget_usd <= 0 + ): raise ValueError("max_budget_usd must be a finite number greater than 0") self._model = model self._max_budget_usd = max_budget_usd diff --git a/strix/interface/main.py b/strix/interface/main.py index 54041818..2403200d 100644 --- a/strix/interface/main.py +++ b/strix/interface/main.py @@ -23,9 +23,11 @@ from strix.config import ( persist_current, ) from strix.config.models import ( + RECOMMENDED_MODEL_NAMES, StrixProvider, configure_sdk_model_defaults, is_known_openai_bare_model, + is_recommended_or_frontier_model, ) from strix.core.paths import run_dir_for, runtime_state_dir from strix.interface.cli import run_cli @@ -264,7 +266,7 @@ def _provider_import_hint(exc: BaseException, model: str) -> str | None: return None -async def warm_up_llm() -> None: +async def warm_up_llm(show_model_warning: bool = True) -> None: console = Console() logger.info("Warming up LLM connection") @@ -306,6 +308,32 @@ async def warm_up_llm() -> None: ) sys.exit(1) + if show_model_warning and raw_model and not is_recommended_or_frontier_model(raw_model): + warn_text = Text() + warn_text.append("MODEL QUALITY WARNING", style="bold yellow") + warn_text.append("\n\n", style="white") + warn_text.append(f"'{raw_model}'", style="bold cyan") + warn_text.append( + " is not a recommended frontier model for Strix.\nSecurity scans work best with:\n", + style="white", + ) + for recommended_model in RECOMMENDED_MODEL_NAMES: + warn_text.append(f"• {recommended_model}\n", style="bold cyan") + warn_text.append( + "\nYou can continue, but weaker models may miss vulnerabilities " + "or produce lower-quality findings.", + style="white", + ) + console.print( + Panel( + warn_text, + title="[bold white]STRIX", + title_align="left", + border_style="yellow", + padding=(1, 2), + ), + ) + model = StrixProvider().get_model(raw_model) await asyncio.wait_for( model.get_response( @@ -827,7 +855,7 @@ def main() -> None: pull_docker_image() validate_environment() - asyncio.run(warm_up_llm()) + asyncio.run(warm_up_llm(show_model_warning=args.non_interactive)) persist_current() diff --git a/strix/interface/tui/app.py b/strix/interface/tui/app.py index 496e99c7..73a7ec0d 100644 --- a/strix/interface/tui/app.py +++ b/strix/interface/tui/app.py @@ -31,6 +31,7 @@ from textual.widgets import Button, Label, Static, TextArea, Tree from textual.widgets.tree import TreeNode from strix.config import load_settings +from strix.config.models import is_recommended_or_frontier_model from strix.core.hooks import BudgetExceededError from strix.core.runner import run_strix_scan from strix.interface.tui.live_view import TuiLiveView @@ -116,9 +117,16 @@ class SplashScreen(Static): # type: ignore[misc] self._animation_timer: Timer | None = None self._panel_static: Static | None = None self._version = "dev" + self._non_frontier_model: str | None = None def compose(self) -> ComposeResult: self._version = get_package_version() + try: + model = (load_settings().llm.model or "").strip() + except Exception: + model = "" + if model and not is_recommended_or_frontier_model(model): + self._non_frontier_model = model self._animation_step = 0 start_line = self._build_start_line_text(self._animation_step) panel = self._build_panel(start_line) @@ -145,7 +153,7 @@ class SplashScreen(Static): # type: ignore[misc] self._panel_static.update(panel) def _build_panel(self, start_line: Text) -> Panel: - content = Group( + rows = [ Align.center(Text(self.BANNER.strip("\n"), style=self.PRIMARY_GREEN, justify="center")), Align.center(Text(" ")), Align.center(self._build_welcome_text()), @@ -155,9 +163,26 @@ class SplashScreen(Static): # type: ignore[misc] Align.center(start_line.copy()), Align.center(Text(" ")), Align.center(self._build_url_text()), - ) + ] + if self._non_frontier_model: + rows.extend( + ( + Align.center(Text(" ")), + Align.center(self._build_model_warning_text(self._non_frontier_model)), + ) + ) - return Panel.fit(content, border_style=self.PRIMARY_GREEN, padding=(1, 6)) + return Panel.fit(Group(*rows), border_style=self.PRIMARY_GREEN, padding=(1, 6)) + + @staticmethod + def _build_model_warning_text(model: str) -> Text: + text = Text("⚠ ", style=Style(color="yellow", bold=True)) + text.append(model, style=Style(color="cyan", bold=True)) + text.append( + " is not a recommended frontier model - pentest quality could be degraded", + style=Style(color="yellow"), + ) + return text def _build_url_text(self) -> Text: return Text("strix.ai", style=Style(color=self.PRIMARY_GREEN, bold=True)) diff --git a/strix/runtime/docker_client.py b/strix/runtime/docker_client.py index 13ac05bb..2321fa50 100644 --- a/strix/runtime/docker_client.py +++ b/strix/runtime/docker_client.py @@ -49,7 +49,7 @@ logger = logging.getLogger(__name__) class StrixDockerSandboxClient(DockerSandboxClient): # Host directories to bind-mount into the container, set by the docker # backend before ``create()``. Each item is ``{source, target, read_only}``. - strix_bind_mounts: list[dict[str, Any]] = [] # overridden per-instance in backends.py + strix_bind_mounts: list[dict[str, Any]] | None = None async def _create_container( self, diff --git a/tests/test_models.py b/tests/test_models.py new file mode 100644 index 00000000..dc86e98f --- /dev/null +++ b/tests/test_models.py @@ -0,0 +1,66 @@ +"""Tests for LLM model recommendation helpers.""" + +from __future__ import annotations + +import pytest + +from strix.config.models import RECOMMENDED_MODEL_NAMES, is_recommended_or_frontier_model + + +@pytest.mark.parametrize("model_name", RECOMMENDED_MODEL_NAMES) +def test_recommended_models_are_accepted(model_name: str) -> None: + assert is_recommended_or_frontier_model(model_name) + + +def test_recommended_models_are_matched_case_insensitively() -> None: + assert is_recommended_or_frontier_model("Vertex_AI/Gemini-3-Pro-Preview") + + +@pytest.mark.parametrize( + "model_name", + [ + "gpt-5.5", + "litellm/openai/gpt-5.4-pro", + "azure_ai/gpt-5.5-pro", + "bedrock_mantle/openai.gpt-5.5", + "anthropic/claude-opus-4-8", + "anthropic.claude-opus-4-8", + "anthropic/claude-opus-4-7", + "anthropic/claude-fable-5", + "anthropic/claude-sonnet-5", + "vertex_ai/claude-sonnet-5@default", + "vertex_ai/claude-sonnet-4-6@default", + "any-llm/anthropic/claude-sonnet-4-6", + "vertex_ai/gemini-3.1-pro-preview", + "openrouter/google/gemini-3.1-pro-preview", + "deepseek/deepseek-v4-pro", + "deepseek/deepseek-r1-0528", + "deepseek/deepseek-reasoner", + "dashscope/qwen3-max-2026-01-23", + "qwen3.7-max", + "moonshot/kimi-k2.6", + "kimi-k2.7-code", + ], +) +def test_frontier_model_families_are_accepted(model_name: str) -> None: + assert is_recommended_or_frontier_model(model_name) + + +@pytest.mark.parametrize( + "model_name", + [ + "", + "openai/gpt-4.1", + "anthropic/claude-3-5-sonnet-latest", + "ollama/llama3.1", + "deepseek/deepseek-chat", + "custom-ollama/gpt-5-mini-local", + "custom-provider/claude-opus-4-local", + "xai/grok-4.5", + "openrouter/x-ai/grok-4", + "mistral/mistral-medium-3-5", + "mistral/magistral-medium-latest", + ], +) +def test_non_frontier_models_are_rejected(model_name: str) -> None: + assert not is_recommended_or_frontier_model(model_name)