from __future__ import annotations import json from math import nan from typing import Any, Mapping import pytest from agent import ( AgentRole, FileSystemArtifactStore, FileSystemTaskStore, RunConfig, RoleRunConfigOverrides, RoleRunConfigResolver, wire_orchestration, ) from agent.core.prompts.orchestration import PLANNER_ROLE_CONTRACT from agent.orchestration.executor import LocalAgentExecutor from agent.orchestration.models import CompletionPolicy, FailureCode class CapturingRunner: trace_store = None def __init__(self) -> None: self.configs = [] self.messages = [] async def run_result(self, *, messages, config): self.messages.append(messages) self.configs.append(config) return {"status": "completed", "summary": "done"} class RecordingResolver: def __init__(self, overrides: RoleRunConfigOverrides) -> None: self.overrides = overrides self.calls = [] async def resolve( self, *, role: AgentRole, preset: str, context: Mapping[str, Any], ) -> RoleRunConfigOverrides: with pytest.raises(TypeError): context["injected"] = True # type: ignore[index] self.calls.append((role, preset, dict(context))) return self.overrides def worker_context(**updates): context = { "worker_preset": "custom_worker", "worker_trace_id": "worker-trace", "root_trace_id": "root", "task_id": "task", "spec_version": 2, "attempt_id": "attempt", "task_spec": {"objective": "test"}, "continue_trace_id": None, } context.update(updates) return context def validator_context(**updates): context = { "validator_preset": "custom_validator", "validator_trace_id": "validator-trace", "root_trace_id": "root", "task_id": "task", "spec_version": 2, "attempt_id": "attempt", "snapshot_id": "snapshot", "validation_id": "validation", "task_spec": {"objective": "test"}, "artifact_snapshot": {"artifact_refs": []}, } context.update(updates) return context def test_role_run_config_types_are_publicly_exported(): assert RoleRunConfigResolver.__name__ == "RoleRunConfigResolver" assert RoleRunConfigOverrides().__dict__ == { "model": None, "temperature": None, "max_iterations": None, } @pytest.mark.asyncio async def test_local_executor_without_resolver_preserves_run_config_defaults(): runner = CapturingRunner() result = await LocalAgentExecutor(runner).run_worker(worker_context()) assert result.status == "completed" config = runner.configs[0] assert config.model == "gpt-4o" assert config.temperature == 0.3 assert config.max_iterations == 200 @pytest.mark.asyncio async def test_worker_resolver_overrides_only_model_run_fields(): runner = CapturingRunner() resolver = RecordingResolver( RoleRunConfigOverrides( model="worker-model", temperature=0.15, max_iterations=17, ) ) result = await LocalAgentExecutor(runner, resolver).run_worker( worker_context(untrusted="resolver-can-see-but-runner-must-not-receive") ) assert result.status == "completed" assert resolver.calls[0][0:2] == (AgentRole.WORKER, "custom_worker") assert resolver.calls[0][2]["untrusted"] == ( "resolver-can-see-but-runner-must-not-receive" ) config = runner.configs[0] assert (config.model, config.temperature, config.max_iterations) == ( "worker-model", 0.15, 17, ) assert config.agent_type == "custom_worker" assert config.completion_policy == CompletionPolicy.EXPLICIT_VALIDATION assert config.new_trace_id == "worker-trace" assert config.parent_trace_id == "root" assert config.tools is None and config.tool_groups is None assert config.parallel_tool_execution is False assert config.enable_memory is False assert config.enable_research_flow is False assert "untrusted" not in config.context @pytest.mark.asyncio async def test_resolver_context_is_recursively_read_only_and_cannot_mutate_prompt(): class NestedMutationResolver: async def resolve(self, *, role, preset, context): del role, preset with pytest.raises(TypeError): context["task_spec"]["objective"] = "mutated" with pytest.raises(AttributeError): context["task_spec"]["refs"].append("mutated") return RoleRunConfigOverrides() runner = CapturingRunner() result = await LocalAgentExecutor(runner, NestedMutationResolver()).run_worker( worker_context(task_spec={"objective": "original", "refs": ["safe"]}) ) assert result.status == "completed" prompt = json.loads(runner.messages[0][0]["content"]) assert prompt["task_spec"] == {"objective": "original", "refs": ["safe"]} assert runner.configs[0].context["task_id"] == "task" def test_trusted_resolver_values_override_preset_model_policy(): from agent import AgentPreset, ToolRegistry from agent.orchestration.policy import DefaultToolPolicy config = RunConfig( temperature=0.7, max_iterations=99, completion_policy=CompletionPolicy.EXPLICIT_VALIDATION, _role_run_config_override_fields=frozenset({"temperature", "max_iterations"}), ) preset = AgentPreset( role=AgentRole.VALIDATOR, allowed_tools=[], temperature=0.0, max_iterations=30, ) resolved = DefaultToolPolicy().resolve(config, preset, ToolRegistry()) assert resolved.temperature == 0.7 assert resolved.max_iterations == 99 @pytest.mark.asyncio async def test_validator_resolver_receives_validator_role_and_preset(): runner = CapturingRunner() resolver = RecordingResolver( RoleRunConfigOverrides(model="validator-model", temperature=0.0) ) result = await LocalAgentExecutor(runner, resolver).run_validator( validator_context() ) assert result.status == "completed" assert resolver.calls[0][0:2] == (AgentRole.VALIDATOR, "custom_validator") config = runner.configs[0] assert config.model == "validator-model" assert config.temperature == 0.0 assert config.max_iterations == 200 assert config.context["validation_id"] == "validation" @pytest.mark.parametrize( "kwargs", [ {"model": " "}, {"model": 1}, {"temperature": nan}, {"temperature": -0.1}, {"temperature": True}, {"max_iterations": 0}, {"max_iterations": True}, ], ) def test_role_run_config_overrides_reject_invalid_values(kwargs): with pytest.raises(ValueError): RoleRunConfigOverrides(**kwargs) @pytest.mark.asyncio async def test_invalid_resolver_result_is_normalized_as_executor_failure(): class InvalidResolver: async def resolve(self, **kwargs): return {"model": "unsafe"} runner = CapturingRunner() result = await LocalAgentExecutor(runner, InvalidResolver()).run_worker( worker_context() ) assert result.status == "failed" assert "must return RoleRunConfigOverrides" in result.error assert result.execution_stats.failure_code == FailureCode.EXECUTOR_ERROR assert runner.configs == [] @pytest.mark.asyncio async def test_resolver_exception_is_normalized_as_executor_failure(): class BrokenResolver: async def resolve(self, **kwargs): raise RuntimeError("model manifest unavailable") runner = CapturingRunner() result = await LocalAgentExecutor(runner, BrokenResolver()).run_validator( validator_context() ) assert result.status == "failed" assert result.error == "model manifest unavailable" assert result.execution_stats.failure_code == FailureCode.EXECUTOR_ERROR assert runner.configs == [] @pytest.mark.asyncio async def test_wire_orchestration_injects_role_run_config_resolver(tmp_path): runner = CapturingRunner() runner.task_coordinator = None resolver = RecordingResolver(RoleRunConfigOverrides(model="wired-model")) coordinator = wire_orchestration( runner, FileSystemTaskStore(str(tmp_path)), FileSystemArtifactStore(str(tmp_path)), role_run_config_resolver=resolver, ) result = await coordinator.executor.run_worker(worker_context()) assert runner.task_coordinator is coordinator assert result.status == "completed" assert runner.configs[0].model == "wired-model" def test_planner_contract_allows_host_specific_dispatch_tool_name(): assert "dispatch_tasks" not in PLANNER_ROLE_CONTRACT assert "task-dispatch tool exposed by your preset" in PLANNER_ROLE_CONTRACT