| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288 |
- from __future__ import annotations
- import json
- from math import nan
- from typing import Any, Mapping
- import pytest
- from agent import (
- AgentRole,
- FileSystemArtifactStore,
- FileSystemTaskStore,
- RoleRunConfigOverrides,
- RoleRunConfigResolver,
- RunConfig,
- 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():
- for framework_tool_name in ("task_plan", "task_decide", "dispatch_tasks"):
- assert framework_tool_name not in PLANNER_ROLE_CONTRACT
- assert "tools exposed by your preset" in PLANNER_ROLE_CONTRACT
- assert "preset-specific" in PLANNER_ROLE_CONTRACT
- assert "plan-inspection" in PLANNER_ROLE_CONTRACT
- assert "explicit host capability boundary may require Root BLOCK" in (
- PLANNER_ROLE_CONTRACT
- )
|