test_repositories.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289
  1. from __future__ import annotations
  2. import asyncio
  3. from datetime import UTC, datetime
  4. import pytest
  5. from agent.orchestration import ArtifactRef, EvidenceQuery
  6. from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker
  7. from script_build_host.agents.validation import ScriptBuildArtifactEvidenceReader
  8. from script_build_host.domain.artifacts import (
  9. ArtifactState,
  10. DirectionGoal,
  11. EvidenceRecordV1,
  12. ScriptDirectionArtifactV1,
  13. )
  14. from script_build_host.domain.errors import (
  15. ArtifactAlreadyFrozen,
  16. ArtifactDigestMismatch,
  17. ArtifactNotFound,
  18. ArtifactOwnershipMismatch,
  19. InputHashConflict,
  20. )
  21. from script_build_host.domain.input_snapshot import ScriptBuildInput
  22. from script_build_host.domain.records import BuildStatus, PublicationState, PublicationType
  23. from script_build_host.repositories.legacy_state import SqlAlchemyLegacyBuildStateRepository
  24. from script_build_host.repositories.sqlalchemy import (
  25. SqlAlchemyInputSnapshotRepository,
  26. SqlAlchemyMissionBindingRepository,
  27. SqlAlchemyPublicationRepository,
  28. SqlAlchemyScriptBusinessArtifactRepository,
  29. )
  30. def _input(build_id: int = 1, *, topic_result: str = "result") -> ScriptBuildInput:
  31. return ScriptBuildInput(
  32. script_build_id=build_id,
  33. execution_id=10,
  34. topic_build_id=20,
  35. topic_id=30,
  36. topic={"topic": {"id": 30, "result": topic_result}, "points": [], "sources": []},
  37. account={"account_name": "account", "source": "personal_config"},
  38. prompt_manifest=({"biz_type": "planner", "content_sha256": "sha256:" + "0" * 64},),
  39. datasource_manifest={"pattern": {"version": "fixture"}},
  40. model_manifest={"planner": {"model": "fake"}},
  41. )
  42. def _evidence() -> EvidenceRecordV1:
  43. return EvidenceRecordV1(
  44. evidence_id="evidence-1",
  45. source_type="decode",
  46. tool_name="search_script_decode_case",
  47. query={"return_field": "summary", "top_k": 3},
  48. source_refs=("decode://case/1",),
  49. raw_artifact_ref=None,
  50. summary="source-backed summary",
  51. supports=("criterion-1",),
  52. confidence="medium",
  53. limitations=("fixture",),
  54. content_sha256="",
  55. created_at=datetime(2026, 7, 19, tzinfo=UTC),
  56. )
  57. @pytest.mark.asyncio
  58. async def test_snapshot_freeze_is_idempotent_and_conflicts_by_version(
  59. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  60. ) -> None:
  61. _, sessions = database
  62. repository = SqlAlchemyInputSnapshotRepository(sessions)
  63. first = await repository.freeze(_input())
  64. second = await repository.freeze(_input())
  65. assert first.snapshot_id == second.snapshot_id
  66. assert first.canonical_sha256 == second.canonical_sha256
  67. same_hash_new_version = await repository.freeze(_input(), version=2)
  68. assert same_hash_new_version.snapshot_id == first.snapshot_id
  69. with pytest.raises(InputHashConflict):
  70. await repository.freeze(_input(topic_result="changed"))
  71. with pytest.raises(Exception) as cross_build:
  72. await repository.get(first.snapshot_id, script_build_id=999)
  73. assert getattr(cross_build.value, "code", None) == "BUILD_NOT_FOUND"
  74. @pytest.mark.asyncio
  75. async def test_concurrent_same_snapshot_freeze_returns_one_version(
  76. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  77. ) -> None:
  78. _, sessions = database
  79. repository = SqlAlchemyInputSnapshotRepository(sessions)
  80. first, second = await asyncio.gather(repository.freeze(_input()), repository.freeze(_input()))
  81. assert first.snapshot_id == second.snapshot_id
  82. @pytest.mark.asyncio
  83. async def test_binding_create_is_idempotent_and_active_pointer_is_build_scoped(
  84. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  85. ) -> None:
  86. _, sessions = database
  87. repository = SqlAlchemyMissionBindingRepository(sessions)
  88. values = {
  89. "script_build_id": 1,
  90. "root_trace_id": "root-1",
  91. "input_snapshot_id": 2,
  92. "engine_version": "0.4.0",
  93. "schema_version": "phase-one/v1",
  94. }
  95. first, second = await asyncio.gather(repository.create(**values), repository.create(**values))
  96. assert first.binding_id == second.binding_id
  97. updated = await repository.set_active_direction(script_build_id=1, artifact_version_id=99)
  98. assert updated.active_direction_artifact_version_id == 99
  99. @pytest.mark.asyncio
  100. async def test_artifact_freeze_ref_ownership_digest_and_spec_replay(
  101. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  102. ) -> None:
  103. _, sessions = database
  104. repository = SqlAlchemyScriptBusinessArtifactRepository(sessions)
  105. version, reference = await repository.freeze(
  106. script_build_id=1,
  107. task_id="task-1",
  108. attempt_id="attempt-1",
  109. spec_version=1,
  110. artifact=_evidence(),
  111. )
  112. assert version.state is ArtifactState.FROZEN
  113. assert reference.uri == f"script-build://artifact-versions/{version.artifact_version_id}"
  114. assert reference.kind == "evidence"
  115. assert reference.version == str(version.artifact_version_id)
  116. assert reference.digest and reference.digest.startswith("sha256:")
  117. loaded = await repository.read_by_ref(
  118. reference, script_build_id=1, task_id="task-1", attempt_id="attempt-1"
  119. )
  120. assert loaded == version
  121. assert await repository.get_by_id(version.artifact_version_id, script_build_id=1) == version
  122. with pytest.raises(ArtifactNotFound):
  123. await repository.get_by_id(version.artifact_version_id, script_build_id=2)
  124. with pytest.raises(ArtifactOwnershipMismatch):
  125. await repository.read_by_ref(reference, script_build_id=2)
  126. tampered = ArtifactRef(
  127. uri=reference.uri,
  128. kind=reference.kind,
  129. version=reference.version,
  130. digest="sha256:" + "f" * 64,
  131. )
  132. with pytest.raises(ArtifactDigestMismatch):
  133. await repository.read_by_ref(tampered, script_build_id=1)
  134. assert not await repository.verify_digest(tampered, script_build_id=1)
  135. with pytest.raises(ArtifactAlreadyFrozen):
  136. await repository.freeze(
  137. script_build_id=1,
  138. task_id="task-1",
  139. attempt_id="attempt-1",
  140. spec_version=2,
  141. artifact=_evidence(),
  142. )
  143. class Bindings:
  144. async def get_by_root(self, _root: str):
  145. return type("Binding", (), {"script_build_id": 1})()
  146. reader = ScriptBuildArtifactEvidenceReader(repository, Bindings()) # type: ignore[arg-type]
  147. with pytest.raises(ArtifactOwnershipMismatch):
  148. await reader.read(
  149. reference,
  150. EvidenceQuery(
  151. root_trace_id="root-1",
  152. task_id="other-task",
  153. attempt_id="other-attempt",
  154. snapshot_id="snapshot-1",
  155. validation_id="validation-1",
  156. query="evidence",
  157. limit=1,
  158. ),
  159. )
  160. @pytest.mark.asyncio
  161. async def test_artifact_same_attempt_and_digest_is_concurrently_idempotent(
  162. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  163. ) -> None:
  164. _, sessions = database
  165. repository = SqlAlchemyScriptBusinessArtifactRepository(sessions)
  166. async def freeze() -> tuple[object, ArtifactRef]:
  167. return await repository.freeze(
  168. script_build_id=1,
  169. task_id="task-1",
  170. attempt_id="attempt-1",
  171. spec_version=1,
  172. artifact=_evidence(),
  173. )
  174. first, second = await asyncio.gather(freeze(), freeze())
  175. assert first[1] == second[1]
  176. @pytest.mark.asyncio
  177. async def test_direction_and_publication_are_idempotent_and_error_is_redacted(
  178. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  179. ) -> None:
  180. _, sessions = database
  181. artifacts = SqlAlchemyScriptBusinessArtifactRepository(sessions)
  182. direction, _ = await artifacts.freeze(
  183. script_build_id=1,
  184. task_id="direction-task",
  185. attempt_id="direction-attempt",
  186. spec_version=1,
  187. artifact=ScriptDirectionArtifactV1(
  188. goals=(DirectionGoal("goal-1", "make the topic clear", "source-backed"),),
  189. evidence_refs=("script-build://artifact-versions/1",),
  190. legacy_markdown="# Direction",
  191. ),
  192. )
  193. publications = SqlAlchemyPublicationRepository(sessions)
  194. first = await publications.prepare(
  195. script_build_id=1,
  196. publication_type=PublicationType.DIRECTION,
  197. accept_decision_id="decision-1",
  198. artifact_version_id=direction.artifact_version_id,
  199. expected_sha256=direction.canonical_sha256,
  200. )
  201. second = await publications.prepare(
  202. script_build_id=1,
  203. publication_type=PublicationType.DIRECTION,
  204. accept_decision_id="decision-1",
  205. artifact_version_id=direction.artifact_version_id,
  206. expected_sha256=direction.canonical_sha256,
  207. )
  208. assert first.publication_id == second.publication_id
  209. failed = await publications.mark_failed(
  210. first.publication_id,
  211. error_code="UPSTREAM",
  212. error_summary="mysql://reader:password@db.example/app Bearer raw-token",
  213. )
  214. assert failed.state is PublicationState.FAILED
  215. assert "password" not in (failed.last_error_summary or "")
  216. assert "raw-token" not in (failed.last_error_summary or "")
  217. published = await publications.mark_published(first.publication_id)
  218. assert published.state is PublicationState.PUBLISHED
  219. assert published.publication_revision == 1
  220. replayed = await publications.mark_published(first.publication_id)
  221. assert replayed.publication_revision == 1
  222. still_published = await publications.mark_failed(
  223. first.publication_id,
  224. error_code="LATE_FAILURE",
  225. error_summary="must not downgrade a committed publication",
  226. )
  227. assert still_published.state is PublicationState.PUBLISHED
  228. assert still_published.publication_revision == 1
  229. with pytest.raises(Exception, match="only an accepted direction"):
  230. await publications.prepare(
  231. script_build_id=1,
  232. publication_type=PublicationType.FINAL,
  233. accept_decision_id="decision-1",
  234. artifact_version_id=direction.artifact_version_id,
  235. expected_sha256=direction.canonical_sha256,
  236. )
  237. assert (
  238. await artifacts.get_by_id(direction.artifact_version_id, script_build_id=1)
  239. ).state is ArtifactState.PUBLISHED
  240. observed = await publications.get_by_build(1, publication_type=PublicationType.DIRECTION)
  241. assert observed == published
  242. @pytest.mark.asyncio
  243. async def test_legacy_status_read_write_and_direction_projection(
  244. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  245. ) -> None:
  246. _, sessions = database
  247. repository = SqlAlchemyLegacyBuildStateRepository(sessions)
  248. build_id = await repository.create(
  249. execution_id=10,
  250. topic_build_id=20,
  251. topic_id=30,
  252. agent_type="AigcAgent",
  253. agent_config={"model": "fake"},
  254. data_source_url=None,
  255. strategies_config={"always_on": [7], "on_demand": []},
  256. )
  257. assert await repository.get_status(build_id) is BuildStatus.RUNNING
  258. await repository.set_status(build_id, BuildStatus.STOPPING)
  259. assert await repository.get_status(build_id) is BuildStatus.STOPPING
  260. await repository.project_direction(build_id, "# accepted direction")
  261. await repository.set_status(build_id, BuildStatus.PARTIAL)
  262. assert await repository.get_status(build_id) is BuildStatus.PARTIAL
  263. with pytest.raises(Exception, match="cannot project script build success"):
  264. await repository.set_status(build_id, BuildStatus.SUCCESS)