test_input_snapshot_service.py 5.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160
  1. from __future__ import annotations
  2. from types import SimpleNamespace
  3. from typing import Any
  4. import pytest
  5. from script_build_host.application.input_snapshot_service import (
  6. AssembleInputRequest,
  7. ScriptInputSnapshotService,
  8. )
  9. from script_build_host.domain.input_snapshot import ScriptBuildInput, ScriptBuildInputSnapshotV1
  10. from script_build_host.domain.ports import PersonaInput, PromptRequest
  11. from script_build_host.domain.records import Principal
  12. from script_build_host.infrastructure.canonical_json import canonical_sha256
  13. from script_build_host.repositories.sqlalchemy import SqlAlchemyInputSnapshotRepository
  14. class FakeLegacyInput:
  15. def __init__(self) -> None:
  16. self.calls: list[tuple[int, int, int]] = []
  17. async def read_topic_graph(
  18. self, *, execution_id: int, topic_build_id: int, topic_id: int
  19. ) -> dict[str, Any]:
  20. self.calls.append((execution_id, topic_build_id, topic_id))
  21. return {
  22. "build_record": {
  23. "personal_config": {"account_name": "account"},
  24. "strategies_config": {"always_on": [7], "on_demand": [8]},
  25. "agent_config": {"api_key": "hidden", "model": "fake"},
  26. },
  27. "topic": {"id": topic_id},
  28. "points": [],
  29. "composition_items": [],
  30. "relations": [],
  31. "sources": [],
  32. }
  33. class FakePersona:
  34. async def load(self, account_name: str) -> PersonaInput:
  35. return PersonaInput(account_name, account_name, ({"point": "one"},), (), {"p": "sha256:x"})
  36. class FakeStrategies:
  37. async def load(
  38. self, *, always_on: tuple[Any, ...], on_demand: tuple[Any, ...]
  39. ) -> tuple[dict[str, Any], ...]:
  40. return ({"always_on": list(always_on), "on_demand": list(on_demand)},)
  41. class FakePrompts:
  42. async def load(self, requests: tuple[Any, ...]) -> tuple[dict[str, Any], ...]:
  43. return ({"source": "fixture", "count": len(requests)},)
  44. class FakeSnapshots:
  45. def __init__(self) -> None:
  46. self.frozen: ScriptBuildInput | None = None
  47. async def freeze(
  48. self, assembled: ScriptBuildInput, *, version: int = 1
  49. ) -> ScriptBuildInputSnapshotV1:
  50. self.frozen = assembled
  51. raise NotImplementedError
  52. async def get(self, snapshot_id: str, *, script_build_id: int) -> ScriptBuildInputSnapshotV1:
  53. raise NotImplementedError
  54. @pytest.mark.asyncio
  55. async def test_validate_source_is_read_only_and_assemble_redacts_secrets() -> None:
  56. legacy = FakeLegacyInput()
  57. service = ScriptInputSnapshotService(
  58. legacy_input=legacy,
  59. persona_source=FakePersona(),
  60. strategy_source=FakeStrategies(),
  61. prompt_source=FakePrompts(),
  62. snapshots=FakeSnapshots(),
  63. )
  64. await service.validate_source(execution_id=10, topic_build_id=20, topic_id=30)
  65. assembled = await service.assemble(
  66. AssembleInputRequest(
  67. script_build_id=1,
  68. execution_id=10,
  69. topic_build_id=20,
  70. topic_id=30,
  71. principal=Principal("user-1"),
  72. strategies_always_on=(7,),
  73. strategies_on_demand=(8,),
  74. datasource_manifest={"endpoint": "https://api.example?q=ok&token=hidden"},
  75. model_manifest={"planner": {"model": "fake", "api_key": "hidden"}},
  76. )
  77. )
  78. assert legacy.calls == [(10, 20, 30), (10, 20, 30)]
  79. assert "hidden" not in str(assembled.canonical_payload())
  80. assert assembled.strategies[0] == {"always_on": [7], "on_demand": [8]}
  81. @pytest.mark.asyncio
  82. async def test_extend_prompt_lineage_is_append_only_and_preserves_business_digest(
  83. database,
  84. ) -> None:
  85. _, sessions = database
  86. repository = SqlAlchemyInputSnapshotRepository(sessions)
  87. parent = await repository.freeze(
  88. ScriptBuildInput(
  89. script_build_id=1,
  90. execution_id=10,
  91. topic_build_id=20,
  92. topic_id=30,
  93. topic={"topic": {"id": 30}},
  94. account={"account_name": "acct"},
  95. prompt_manifest=(
  96. {
  97. "preset": "script_planner",
  98. "role": "planner",
  99. "content": "phase one",
  100. "content_sha256": canonical_sha256("phase one").wire,
  101. },
  102. ),
  103. )
  104. )
  105. class Prompts:
  106. async def load(self, _requests):
  107. return (
  108. {
  109. "preset": "script_compose_worker",
  110. "role": "worker",
  111. "content": "phase two",
  112. "content_sha256": canonical_sha256("phase two").wire,
  113. },
  114. )
  115. service = ScriptInputSnapshotService(
  116. legacy_input=SimpleNamespace(),
  117. persona_source=SimpleNamespace(),
  118. strategy_source=SimpleNamespace(),
  119. prompt_source=Prompts(),
  120. snapshots=repository,
  121. )
  122. child = await service.extend_prompt_lineage(
  123. parent,
  124. prompt_requests=(
  125. PromptRequest("compose", "compose.md", "script_compose_worker", "worker"),
  126. ),
  127. model_manifest={"presets": {"script_compose_worker": {"model": "fake"}}},
  128. )
  129. assert child.snapshot_id != parent.snapshot_id
  130. assert child.parent_snapshot_id == parent.snapshot_id
  131. assert child.parent_snapshot_sha256 == parent.canonical_sha256
  132. assert (
  133. child.business_input_sha256 == canonical_sha256(parent.to_input().business_payload()).wire
  134. )
  135. assert {item["preset"] for item in child.prompt_manifest} == {
  136. "script_planner",
  137. "script_compose_worker",
  138. }