test_agent_surface.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119
  1. from __future__ import annotations
  2. from datetime import UTC, datetime
  3. from types import SimpleNamespace
  4. import pytest
  5. from agent import AgentRole, ToolRegistry, get_preset
  6. from script_build_host.agents.model_resolver import (
  7. ScriptRoleRunConfigResolver,
  8. SnapshotModelManifestSource,
  9. )
  10. from script_build_host.agents.presets import register_script_presets
  11. from script_build_host.domain.input_snapshot import ScriptBuildInputSnapshotV1
  12. from script_build_host.tools.registry import register_script_tools
  13. def test_script_presets_are_exact_and_never_expose_legacy_control_tools() -> None:
  14. register_script_presets()
  15. expected = {
  16. "script_planner",
  17. "script_direction_worker",
  18. "script_pattern_retrieval_worker",
  19. "script_decode_retrieval_worker",
  20. "script_external_retrieval_worker",
  21. "script_knowledge_retrieval_worker",
  22. "script_retrieval_validator",
  23. "script_candidate_validator",
  24. }
  25. forbidden = {
  26. "agent",
  27. "evaluate",
  28. "bash_command",
  29. "write_file",
  30. "dispatch_tasks",
  31. "implement_paths",
  32. "begin_round",
  33. "multipath",
  34. }
  35. for name in expected:
  36. preset = get_preset(name)
  37. assert preset.role in {AgentRole.PLANNER, AgentRole.WORKER, AgentRole.VALIDATOR}
  38. assert forbidden.isdisjoint(set(preset.allowed_tools or []))
  39. assert set(preset.allowed_tools or []).isdisjoint(set(preset.denied_tools or []))
  40. class _Gateway:
  41. coordinator = SimpleNamespace(task_store=None)
  42. def test_legacy_retrieval_tool_schemas_keep_old_argument_shapes() -> None:
  43. registry = ToolRegistry()
  44. register_script_tools(registry, _Gateway()) # type: ignore[arg-type]
  45. schemas = {
  46. item["function"]["name"]: item["function"]["parameters"] for item in registry.get_schemas()
  47. }
  48. decode = schemas["search_script_decode_case"]
  49. assert decode["required"] == ["return_field"]
  50. assert decode["properties"]["match_fields"]["type"] == "object"
  51. assert decode["properties"]["top_k"]["default"] == 3
  52. external = schemas["external_search_case"]
  53. assert external["required"] == ["keyword"]
  54. assert external["properties"]["platform_channel"]["default"] == "xhs"
  55. knowledge = schemas["search_knowledge"]
  56. assert knowledge["required"] == ["keyword"]
  57. assert knowledge["properties"]["max_count"]["default"] == 3
  58. assert schemas["query_pattern_qa"]["required"] == ["message"]
  59. assert schemas["load_images"]["required"] == ["image_urls"]
  60. class _Bindings:
  61. async def get_by_root(self, root_trace_id: str) -> SimpleNamespace:
  62. assert root_trace_id == "root-a"
  63. return SimpleNamespace(script_build_id=9, input_snapshot_id=3)
  64. class _Snapshots:
  65. async def get(self, snapshot_id: str, *, script_build_id: int) -> ScriptBuildInputSnapshotV1:
  66. assert (snapshot_id, script_build_id) == ("3", 9)
  67. return ScriptBuildInputSnapshotV1(
  68. snapshot_id="3",
  69. script_build_id=9,
  70. execution_id=1,
  71. topic_build_id=2,
  72. topic_id=3,
  73. topic={},
  74. account={},
  75. persona_points=(),
  76. section_patterns=(),
  77. strategies=(),
  78. prompt_manifest=(),
  79. datasource_manifest={},
  80. model_manifest={
  81. "presets": {
  82. "script_direction_worker": {
  83. "model": "frozen-direction-model",
  84. "temperature": 0.2,
  85. "max_iterations": 18,
  86. }
  87. }
  88. },
  89. canonical_sha256="sha256:" + "a" * 64,
  90. created_at=datetime.now(UTC),
  91. )
  92. @pytest.mark.asyncio
  93. async def test_role_model_resolver_reads_each_roots_frozen_manifest() -> None:
  94. resolver = ScriptRoleRunConfigResolver(
  95. {}, manifest_source=SnapshotModelManifestSource(_Bindings(), _Snapshots())
  96. )
  97. value = await resolver.resolve(
  98. role=AgentRole.WORKER,
  99. preset="script_direction_worker",
  100. context={"root_trace_id": "root-a"},
  101. )
  102. assert value.model == "frozen-direction-model"
  103. assert value.temperature == 0.2
  104. assert value.max_iterations == 18