test_orchestration_run_config_resolver.py 9.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288
  1. from __future__ import annotations
  2. import json
  3. from math import nan
  4. from typing import Any, Mapping
  5. import pytest
  6. from agent import (
  7. AgentRole,
  8. FileSystemArtifactStore,
  9. FileSystemTaskStore,
  10. RoleRunConfigOverrides,
  11. RoleRunConfigResolver,
  12. RunConfig,
  13. wire_orchestration,
  14. )
  15. from agent.core.prompts.orchestration import PLANNER_ROLE_CONTRACT
  16. from agent.orchestration.executor import LocalAgentExecutor
  17. from agent.orchestration.models import CompletionPolicy, FailureCode
  18. class CapturingRunner:
  19. trace_store = None
  20. def __init__(self) -> None:
  21. self.configs = []
  22. self.messages = []
  23. async def run_result(self, *, messages, config):
  24. self.messages.append(messages)
  25. self.configs.append(config)
  26. return {"status": "completed", "summary": "done"}
  27. class RecordingResolver:
  28. def __init__(self, overrides: RoleRunConfigOverrides) -> None:
  29. self.overrides = overrides
  30. self.calls = []
  31. async def resolve(
  32. self,
  33. *,
  34. role: AgentRole,
  35. preset: str,
  36. context: Mapping[str, Any],
  37. ) -> RoleRunConfigOverrides:
  38. with pytest.raises(TypeError):
  39. context["injected"] = True # type: ignore[index]
  40. self.calls.append((role, preset, dict(context)))
  41. return self.overrides
  42. def worker_context(**updates):
  43. context = {
  44. "worker_preset": "custom_worker",
  45. "worker_trace_id": "worker-trace",
  46. "root_trace_id": "root",
  47. "task_id": "task",
  48. "spec_version": 2,
  49. "attempt_id": "attempt",
  50. "task_spec": {"objective": "test"},
  51. "continue_trace_id": None,
  52. }
  53. context.update(updates)
  54. return context
  55. def validator_context(**updates):
  56. context = {
  57. "validator_preset": "custom_validator",
  58. "validator_trace_id": "validator-trace",
  59. "root_trace_id": "root",
  60. "task_id": "task",
  61. "spec_version": 2,
  62. "attempt_id": "attempt",
  63. "snapshot_id": "snapshot",
  64. "validation_id": "validation",
  65. "task_spec": {"objective": "test"},
  66. "artifact_snapshot": {"artifact_refs": []},
  67. }
  68. context.update(updates)
  69. return context
  70. def test_role_run_config_types_are_publicly_exported():
  71. assert RoleRunConfigResolver.__name__ == "RoleRunConfigResolver"
  72. assert RoleRunConfigOverrides().__dict__ == {
  73. "model": None,
  74. "temperature": None,
  75. "max_iterations": None,
  76. }
  77. @pytest.mark.asyncio
  78. async def test_local_executor_without_resolver_preserves_run_config_defaults():
  79. runner = CapturingRunner()
  80. result = await LocalAgentExecutor(runner).run_worker(worker_context())
  81. assert result.status == "completed"
  82. config = runner.configs[0]
  83. assert config.model == "gpt-4o"
  84. assert config.temperature == 0.3
  85. assert config.max_iterations == 200
  86. @pytest.mark.asyncio
  87. async def test_worker_resolver_overrides_only_model_run_fields():
  88. runner = CapturingRunner()
  89. resolver = RecordingResolver(
  90. RoleRunConfigOverrides(
  91. model="worker-model",
  92. temperature=0.15,
  93. max_iterations=17,
  94. )
  95. )
  96. result = await LocalAgentExecutor(runner, resolver).run_worker(
  97. worker_context(untrusted="resolver-can-see-but-runner-must-not-receive")
  98. )
  99. assert result.status == "completed"
  100. assert resolver.calls[0][0:2] == (AgentRole.WORKER, "custom_worker")
  101. assert resolver.calls[0][2]["untrusted"] == (
  102. "resolver-can-see-but-runner-must-not-receive"
  103. )
  104. config = runner.configs[0]
  105. assert (config.model, config.temperature, config.max_iterations) == (
  106. "worker-model",
  107. 0.15,
  108. 17,
  109. )
  110. assert config.agent_type == "custom_worker"
  111. assert config.completion_policy == CompletionPolicy.EXPLICIT_VALIDATION
  112. assert config.new_trace_id == "worker-trace"
  113. assert config.parent_trace_id == "root"
  114. assert config.tools is None and config.tool_groups is None
  115. assert config.parallel_tool_execution is False
  116. assert config.enable_memory is False
  117. assert config.enable_research_flow is False
  118. assert "untrusted" not in config.context
  119. @pytest.mark.asyncio
  120. async def test_resolver_context_is_recursively_read_only_and_cannot_mutate_prompt():
  121. class NestedMutationResolver:
  122. async def resolve(self, *, role, preset, context):
  123. del role, preset
  124. with pytest.raises(TypeError):
  125. context["task_spec"]["objective"] = "mutated"
  126. with pytest.raises(AttributeError):
  127. context["task_spec"]["refs"].append("mutated")
  128. return RoleRunConfigOverrides()
  129. runner = CapturingRunner()
  130. result = await LocalAgentExecutor(runner, NestedMutationResolver()).run_worker(
  131. worker_context(task_spec={"objective": "original", "refs": ["safe"]})
  132. )
  133. assert result.status == "completed"
  134. prompt = json.loads(runner.messages[0][0]["content"])
  135. assert prompt["task_spec"] == {"objective": "original", "refs": ["safe"]}
  136. assert runner.configs[0].context["task_id"] == "task"
  137. def test_trusted_resolver_values_override_preset_model_policy():
  138. from agent import AgentPreset, ToolRegistry
  139. from agent.orchestration.policy import DefaultToolPolicy
  140. config = RunConfig(
  141. temperature=0.7,
  142. max_iterations=99,
  143. completion_policy=CompletionPolicy.EXPLICIT_VALIDATION,
  144. _role_run_config_override_fields=frozenset({"temperature", "max_iterations"}),
  145. )
  146. preset = AgentPreset(
  147. role=AgentRole.VALIDATOR,
  148. allowed_tools=[],
  149. temperature=0.0,
  150. max_iterations=30,
  151. )
  152. resolved = DefaultToolPolicy().resolve(config, preset, ToolRegistry())
  153. assert resolved.temperature == 0.7
  154. assert resolved.max_iterations == 99
  155. @pytest.mark.asyncio
  156. async def test_validator_resolver_receives_validator_role_and_preset():
  157. runner = CapturingRunner()
  158. resolver = RecordingResolver(
  159. RoleRunConfigOverrides(model="validator-model", temperature=0.0)
  160. )
  161. result = await LocalAgentExecutor(runner, resolver).run_validator(
  162. validator_context()
  163. )
  164. assert result.status == "completed"
  165. assert resolver.calls[0][0:2] == (AgentRole.VALIDATOR, "custom_validator")
  166. config = runner.configs[0]
  167. assert config.model == "validator-model"
  168. assert config.temperature == 0.0
  169. assert config.max_iterations == 200
  170. assert config.context["validation_id"] == "validation"
  171. @pytest.mark.parametrize(
  172. "kwargs",
  173. [
  174. {"model": " "},
  175. {"model": 1},
  176. {"temperature": nan},
  177. {"temperature": -0.1},
  178. {"temperature": True},
  179. {"max_iterations": 0},
  180. {"max_iterations": True},
  181. ],
  182. )
  183. def test_role_run_config_overrides_reject_invalid_values(kwargs):
  184. with pytest.raises(ValueError):
  185. RoleRunConfigOverrides(**kwargs)
  186. @pytest.mark.asyncio
  187. async def test_invalid_resolver_result_is_normalized_as_executor_failure():
  188. class InvalidResolver:
  189. async def resolve(self, **kwargs):
  190. return {"model": "unsafe"}
  191. runner = CapturingRunner()
  192. result = await LocalAgentExecutor(runner, InvalidResolver()).run_worker(
  193. worker_context()
  194. )
  195. assert result.status == "failed"
  196. assert "must return RoleRunConfigOverrides" in result.error
  197. assert result.execution_stats.failure_code == FailureCode.EXECUTOR_ERROR
  198. assert runner.configs == []
  199. @pytest.mark.asyncio
  200. async def test_resolver_exception_is_normalized_as_executor_failure():
  201. class BrokenResolver:
  202. async def resolve(self, **kwargs):
  203. raise RuntimeError("model manifest unavailable")
  204. runner = CapturingRunner()
  205. result = await LocalAgentExecutor(runner, BrokenResolver()).run_validator(
  206. validator_context()
  207. )
  208. assert result.status == "failed"
  209. assert result.error == "model manifest unavailable"
  210. assert result.execution_stats.failure_code == FailureCode.EXECUTOR_ERROR
  211. assert runner.configs == []
  212. @pytest.mark.asyncio
  213. async def test_wire_orchestration_injects_role_run_config_resolver(tmp_path):
  214. runner = CapturingRunner()
  215. runner.task_coordinator = None
  216. resolver = RecordingResolver(RoleRunConfigOverrides(model="wired-model"))
  217. coordinator = wire_orchestration(
  218. runner,
  219. FileSystemTaskStore(str(tmp_path)),
  220. FileSystemArtifactStore(str(tmp_path)),
  221. role_run_config_resolver=resolver,
  222. )
  223. result = await coordinator.executor.run_worker(worker_context())
  224. assert runner.task_coordinator is coordinator
  225. assert result.status == "completed"
  226. assert runner.configs[0].model == "wired-model"
  227. def test_planner_contract_allows_host_specific_dispatch_tool_name():
  228. for framework_tool_name in ("task_plan", "task_decide", "dispatch_tasks"):
  229. assert framework_tool_name not in PLANNER_ROLE_CONTRACT
  230. assert "tools exposed by your preset" in PLANNER_ROLE_CONTRACT
  231. assert "preset-specific" in PLANNER_ROLE_CONTRACT
  232. assert "plan-inspection" in PLANNER_ROLE_CONTRACT
  233. assert "explicit host capability boundary may require Root BLOCK" in (
  234. PLANNER_ROLE_CONTRACT
  235. )