| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388 |
- 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"
|