from __future__ import annotations import asyncio from hashlib import sha256 from types import SimpleNamespace from typing import Any, Mapping import pytest from agent import ( AgentPreset, AgentRole, AgentRunner, FileSystemArtifactStore, FileSystemTaskStore, FileSystemTraceStore, RoleSystemPromptOverride, RoleSystemPromptResolver, wire_orchestration, ) from agent.core.presets import register_preset from agent.orchestration.executor import LocalAgentExecutor from agent.orchestration.models import FailureCode def _override(content: str) -> RoleSystemPromptOverride: return RoleSystemPromptOverride( content=content, prompt_identity=f"sha256:{sha256(content.encode('utf-8')).hexdigest()}", ) def _worker_context(**updates: Any) -> dict[str, Any]: context = { "worker_preset": "script_worker", "worker_trace_id": "worker-trace", "root_trace_id": "root", "task_id": "task", "spec_version": 1, "attempt_id": "attempt", "task_spec": {"objective": "write"}, "continue_trace_id": None, } context.update(updates) return context def _validator_context(**updates: Any) -> dict[str, Any]: context = { "validator_preset": "script_validator", "validator_trace_id": "validator-trace", "root_trace_id": "root", "task_id": "task", "spec_version": 1, "attempt_id": "attempt", "snapshot_id": "snapshot", "validation_id": "validation", "task_spec": {"objective": "validate"}, "artifact_snapshot": {"artifact_refs": []}, } context.update(updates) return context class PromptResolver: def __init__(self, override: RoleSystemPromptOverride | None) -> None: self.override = override self.calls: list[tuple[AgentRole, str, Mapping[str, Any]]] = [] async def resolve(self, *, role, preset, context): with pytest.raises(TypeError): context["unsafe"] = True self.calls.append((role, preset, context)) return self.override class CapturingRunner: trace_store = None task_coordinator = None def __init__(self) -> None: self.configs = [] async def run_result(self, *, messages, config): self.configs.append(config) return {"status": "completed", "summary": "done"} def test_system_prompt_types_are_public_and_validate_content_identity(): assert RoleSystemPromptResolver.__name__ == "RoleSystemPromptResolver" value = _override("frozen prompt") assert value.content == "frozen prompt" with pytest.raises(ValueError, match="does not match"): RoleSystemPromptOverride( content="frozen prompt", prompt_identity=f"sha256:{'0' * 64}", ) with pytest.raises(ValueError, match="lowercase"): RoleSystemPromptOverride(content="x", prompt_identity=f"sha256:{'A' * 64}") @pytest.mark.asyncio async def test_new_role_trace_receives_only_prompt_and_protected_identity(): runner = CapturingRunner() resolver = PromptResolver(_override("business policy")) result = await LocalAgentExecutor( runner, role_system_prompt_resolver=resolver ).run_worker(_worker_context(untrusted="not protected")) assert result.status == "completed" assert resolver.calls[0][0:2] == (AgentRole.WORKER, "script_worker") config = runner.configs[0] assert config.system_prompt == "business policy" assert ( config.context["role_prompt_identity"] == _override("business policy").prompt_identity ) assert "untrusted" not in config.context assert config.tools is None and config.tool_groups is None @pytest.mark.asyncio async def test_new_validator_trace_receives_frozen_prompt_identity(): runner = CapturingRunner() expected = _override("validator business policy") resolver = PromptResolver(expected) result = await LocalAgentExecutor( runner, role_system_prompt_resolver=resolver ).run_validator(_validator_context(untrusted="not protected")) assert result.status == "completed" assert resolver.calls[0][0:2] == (AgentRole.VALIDATOR, "script_validator") config = runner.configs[0] assert config.system_prompt == expected.content assert config.context["role_prompt_identity"] == expected.prompt_identity assert config.context["validation_id"] == "validation" assert "untrusted" not in config.context @pytest.mark.asyncio async def test_missing_or_empty_prompt_resolver_keeps_existing_defaults(): no_resolver_runner = CapturingRunner() empty_resolver_runner = CapturingRunner() no_resolver = await LocalAgentExecutor(no_resolver_runner).run_worker( _worker_context() ) empty_resolver = await LocalAgentExecutor( empty_resolver_runner, role_system_prompt_resolver=PromptResolver(None), ).run_worker(_worker_context()) assert no_resolver.status == empty_resolver.status == "completed" assert no_resolver_runner.configs[0].system_prompt is None assert empty_resolver_runner.configs[0].system_prompt is None assert "role_prompt_identity" not in no_resolver_runner.configs[0].context assert "role_prompt_identity" not in empty_resolver_runner.configs[0].context @pytest.mark.asyncio async def test_invalid_prompt_resolver_result_is_an_executor_failure(): class InvalidResolver: async def resolve(self, **kwargs): return {"content": "unsafe", "prompt_identity": "forged"} runner = CapturingRunner() result = await LocalAgentExecutor( runner, role_system_prompt_resolver=InvalidResolver() ).run_worker(_worker_context()) assert result.status == "failed" assert "must return RoleSystemPromptOverride or None" in result.error assert runner.configs == [] @pytest.mark.asyncio async def test_prompt_resolver_exception_is_normalized_as_executor_failure(): class BrokenResolver: async def resolve(self, **kwargs): raise RuntimeError("frozen prompt manifest unavailable") runner = CapturingRunner() result = await LocalAgentExecutor( runner, role_system_prompt_resolver=BrokenResolver() ).run_validator(_validator_context()) assert result.status == "failed" assert result.error == "frozen prompt manifest unavailable" assert result.execution_stats.failure_code == FailureCode.EXECUTOR_ERROR assert runner.configs == [] @pytest.mark.asyncio async def test_concurrent_roles_keep_distinct_frozen_prompts_and_identities(): class PerTaskResolver: async def resolve(self, *, role, preset, context): del preset await asyncio.sleep(0) return _override(f"{role.value} policy for {context['task_id']}") runner = CapturingRunner() executor = LocalAgentExecutor(runner, role_system_prompt_resolver=PerTaskResolver()) worker, validator = await asyncio.gather( executor.run_worker( _worker_context(task_id="worker-task", worker_trace_id="worker-trace") ), executor.run_validator( _validator_context( task_id="validator-task", validator_trace_id="validator-trace" ) ), ) assert worker.status == validator.status == "completed" configs = {config.name: config for config in runner.configs} worker_config = configs["Worker worker-task"] validator_config = configs["Validator validator-task"] assert worker_config.system_prompt == "worker policy for worker-task" assert validator_config.system_prompt == "validator policy for validator-task" assert ( worker_config.context["role_prompt_identity"] == _override("worker policy for worker-task").prompt_identity ) assert ( validator_config.context["role_prompt_identity"] == _override("validator policy for validator-task").prompt_identity ) @pytest.mark.asyncio async def test_real_runner_persists_worker_and_validator_prompt_identity(tmp_path): worker_preset = "test_prompt_identity_worker" validator_preset = "test_prompt_identity_validator" register_preset( worker_preset, AgentPreset(role=AgentRole.WORKER, allowed_tools=[], max_iterations=1), ) register_preset( validator_preset, AgentPreset(role=AgentRole.VALIDATOR, allowed_tools=[], max_iterations=1), ) async def llm_call(**_kwargs): return {"content": "no terminal tool", "tool_calls": None} store = FileSystemTraceStore(str(tmp_path)) runner = AgentRunner( trace_store=store, llm_call=llm_call, task_coordinator=object(), ) class PerRoleResolver: async def resolve(self, *, role, preset, context): del preset, context return _override(f"frozen {role.value} prompt") executor = LocalAgentExecutor(runner, role_system_prompt_resolver=PerRoleResolver()) worker = await executor.run_worker(_worker_context(worker_preset=worker_preset)) validator = await executor.run_validator( _validator_context(validator_preset=validator_preset) ) assert worker.status == validator.status == "failed" worker_trace = await store.get_trace("worker-trace") validator_trace = await store.get_trace("validator-trace") assert ( worker_trace.context["role_prompt_identity"] == _override("frozen worker prompt").prompt_identity ) assert ( validator_trace.context["role_prompt_identity"] == _override("frozen validator prompt").prompt_identity ) worker_path = await store.get_main_path_messages( worker_trace.trace_id, worker_trace.head_sequence ) validator_path = await store.get_main_path_messages( validator_trace.trace_id, validator_trace.head_sequence ) assert "frozen worker prompt" in str(worker_path[0].content) assert "frozen validator prompt" in str(validator_path[0].content) @pytest.mark.asyncio async def test_repair_validates_identity_without_resetting_system_prompt(): expected = _override("original frozen prompt") trace = SimpleNamespace( context={"role_prompt_identity": expected.prompt_identity}, total_tokens=0, total_cost=0.0, model="worker-model", ) class Coordinator: async def validate_continue_from(self, *args): return "worker-trace" class Store: async def get_trace(self, trace_id): return trace async def update_trace(self, trace_id, **updates): trace.context = updates["context"] class Runner(CapturingRunner): def __init__(self): super().__init__() self.trace_store = Store() self.task_coordinator = Coordinator() runner = Runner() result = await LocalAgentExecutor( runner, role_system_prompt_resolver=PromptResolver(expected) ).run_worker( _worker_context( attempt_id="attempt-2", prior_attempt_id="attempt-1", continue_trace_id="worker-trace", ) ) assert result.status == "completed" assert runner.configs[0].system_prompt is None assert trace.context["role_prompt_identity"] == expected.prompt_identity assert trace.context["attempt_id"] == "attempt-2" @pytest.mark.asyncio async def test_repair_prompt_mismatch_fails_before_trace_context_mutation(): original = _override("original") trace = SimpleNamespace( context={"role_prompt_identity": original.prompt_identity, "attempt_id": "old"}, total_tokens=0, total_cost=0.0, model="worker-model", ) class Coordinator: async def validate_continue_from(self, *args): return "worker-trace" class Store: updated = False async def get_trace(self, trace_id): return trace async def update_trace(self, trace_id, **updates): self.updated = True runner = CapturingRunner() runner.trace_store = Store() runner.task_coordinator = Coordinator() result = await LocalAgentExecutor( runner, role_system_prompt_resolver=PromptResolver(_override("changed")) ).run_worker( _worker_context( attempt_id="new", prior_attempt_id="old", continue_trace_id="worker-trace", ) ) assert result.status == "failed" assert "ROLE_SYSTEM_PROMPT_MISMATCH" in result.error assert runner.trace_store.updated is False assert runner.configs == [] @pytest.mark.asyncio async def test_wire_injects_optional_system_prompt_resolver(tmp_path): runner = CapturingRunner() resolver = PromptResolver(_override("wired prompt")) coordinator = wire_orchestration( runner, FileSystemTaskStore(str(tmp_path)), FileSystemArtifactStore(str(tmp_path)), role_system_prompt_resolver=resolver, ) result = await coordinator.executor.run_worker(_worker_context()) assert result.status == "completed" assert runner.configs[0].system_prompt == "wired prompt"