diff --git a/strix/agents/prompt.py b/strix/agents/prompt.py index ec0651956..5cfb8711b 100644 --- a/strix/agents/prompt.py +++ b/strix/agents/prompt.py @@ -16,6 +16,10 @@ logger = logging.getLogger(__name__) _PROMPT_DIRNAME = "prompts" +# Marks where the system prompt is split so the part before it can be cached. +# Removed before the prompt is sent. +CACHE_POINT = "" + def _resolve_skills( *, diff --git a/strix/agents/prompts/scope.jinja b/strix/agents/prompts/scope.jinja index 3c81a9343..80e06e78d 100644 --- a/strix/agents/prompts/scope.jinja +++ b/strix/agents/prompts/scope.jinja @@ -1,3 +1,4 @@ + {% if system_prompt_context and system_prompt_context.authorized_targets %} SYSTEM-VERIFIED SCOPE: - The following scope metadata is injected by the platform into the system prompt and is authoritative diff --git a/strix/config/models.py b/strix/config/models.py index ca92b77ad..6eedf3788 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -36,6 +36,7 @@ from openai.types.responses import ( from openai.types.responses.response_usage import ResponseUsage from openai.types.shared import Reasoning +from strix.agents.prompt import CACHE_POINT from strix.config import codex from strix.config.loader import load_settings from strix.config.tool_call_ids import TurnCallIdRewriter, dedupe_input @@ -306,6 +307,9 @@ class _TurnGuardModel(Model): ) -> ModelResponse: sanitized = dedupe_input(input) rewriter = TurnCallIdRewriter(sanitized) + system_instructions, sanitized = _split_cached_prefix( + system_instructions, sanitized, model_settings + ) response = await self._inner.get_response( system_instructions, cast("str | list[TResponseInputItem]", sanitized), @@ -339,6 +343,9 @@ class _TurnGuardModel(Model): ) -> AsyncIterator[TResponseStreamEvent]: sanitized = dedupe_input(input) rewriter = TurnCallIdRewriter(sanitized) + system_instructions, sanitized = _split_cached_prefix( + system_instructions, sanitized, model_settings + ) limiter = self._limiter() stream = self._inner.stream_response( system_instructions, @@ -359,6 +366,27 @@ class _TurnGuardModel(Model): self._log_dropped(limiter) +def _split_cached_prefix( + system_instructions: str | None, + model_input: str | list[Any], + model_settings: ModelSettings, +) -> tuple[str | None, str | list[Any]]: + """Split the system prompt at each ``CACHE_POINT`` on cache-point routes. + + LiteLLM puts a cache point at the end of each system message, so each part + gets its own. Other routes get the prompt with the markers removed. + """ + if not system_instructions or CACHE_POINT not in system_instructions: + return system_instructions, model_input + extra_args = model_settings.extra_args or {} + if "cache_control_injection_points" not in extra_args: + return system_instructions.replace(CACHE_POINT, ""), model_input + if isinstance(model_input, str): + model_input = [{"role": "user", "content": model_input}] + parts = [part for part in system_instructions.split(CACHE_POINT) if part.strip()] + return None, [*({"role": "system", "content": part} for part in parts), *model_input] + + async def _aclose(stream: AsyncIterator[TResponseStreamEvent]) -> None: if isinstance(stream, AsyncGenerator): with contextlib.suppress(Exception): diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index 17a831195..393da669f 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -18,8 +18,10 @@ import strix.tools.mcp as mcp_pkg import strix.tools.notes.tools as notes_tools import strix.tools.todo.tools as todo_tools from strix.agents.prompt import render_system_prompt +from strix.config.models import _split_cached_prefix from strix.core import runner from strix.core.agents import AgentCoordinator +from strix.core.inputs import make_model_settings from strix.runtime import session_manager from strix.tools.mcp import BearerAuth, McpConnectionConfig, McpConnectionRequest @@ -275,3 +277,28 @@ def test_scope_is_rendered_once_at_the_end_of_the_prompt() -> None: assert prompt.count("SYSTEM-VERIFIED SCOPE") == 1 assert prompt.index("") < prompt.index("SYSTEM-VERIFIED SCOPE") + + +def test_scope_is_sent_as_its_own_system_message_on_cache_point_routes() -> None: + settings = make_model_settings(None, model_name="anthropic/claude-sonnet-5-5") + prompt = render_system_prompt( + system_prompt_context={ + "authorized_targets": [{"type": "web_application", "value": "https://target.invalid"}], + }, + ) + + system, model_input = _split_cached_prefix(prompt, "go", settings) + + assert system is None + assert isinstance(model_input, list) + assert [item["role"] for item in model_input] == ["system", "system", "user"] + assert "https://target.invalid" not in model_input[0]["content"] + assert "https://target.invalid" in model_input[1]["content"] + assert "" not in model_input[0]["content"] + model_input[1]["content"] + + +def test_cache_point_marker_is_removed_without_cache_points() -> None: + settings = make_model_settings(None, model_name="openai/gpt-5") + prompt = "shared\n\ntargets" + + assert _split_cached_prefix(prompt, "go", settings) == ("shared\n\ntargets", "go")