test_repositories.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376
  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 import select, update
  7. from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker
  8. from script_build_host.agents.validation import ScriptBuildArtifactEvidenceReader
  9. from script_build_host.domain.artifacts import (
  10. ArtifactState,
  11. DirectionGoal,
  12. EvidenceRecordV1,
  13. ScriptDirectionArtifactV1,
  14. )
  15. from script_build_host.domain.errors import (
  16. ArtifactAlreadyFrozen,
  17. ArtifactDigestMismatch,
  18. ArtifactNotFound,
  19. ArtifactOwnershipMismatch,
  20. InputDigestMismatch,
  21. InputHashConflict,
  22. )
  23. from script_build_host.domain.input_snapshot import ScriptBuildInput
  24. from script_build_host.domain.records import BuildStatus, PublicationState, PublicationType
  25. from script_build_host.infrastructure.legacy_tables import script_build_record
  26. from script_build_host.infrastructure.tables import input_snapshot_table
  27. from script_build_host.repositories.legacy_state import SqlAlchemyLegacyBuildStateRepository
  28. from script_build_host.repositories.sqlalchemy import (
  29. SqlAlchemyInputSnapshotRepository,
  30. SqlAlchemyMissionBindingRepository,
  31. SqlAlchemyPublicationRepository,
  32. SqlAlchemyScriptBusinessArtifactRepository,
  33. )
  34. def _input(build_id: int = 1, *, topic_result: str = "result") -> ScriptBuildInput:
  35. return ScriptBuildInput(
  36. script_build_id=build_id,
  37. execution_id=10,
  38. topic_build_id=20,
  39. topic_id=30,
  40. topic={"topic": {"id": 30, "result": topic_result}, "points": [], "sources": []},
  41. account={"account_name": "account", "source": "personal_config"},
  42. prompt_manifest=({"biz_type": "planner", "content_sha256": "sha256:" + "0" * 64},),
  43. datasource_manifest={"pattern": {"version": "fixture"}},
  44. model_manifest={"planner": {"model": "fake"}},
  45. )
  46. def _evidence() -> EvidenceRecordV1:
  47. return EvidenceRecordV1(
  48. evidence_id="evidence-1",
  49. source_type="decode",
  50. tool_name="search_script_decode_case",
  51. query={"return_field": "summary", "top_k": 3},
  52. source_refs=("decode://case/1",),
  53. raw_artifact_ref=None,
  54. summary="source-backed summary",
  55. supports=("criterion-1",),
  56. confidence="medium",
  57. limitations=("fixture",),
  58. content_sha256="",
  59. created_at=datetime(2026, 7, 19, tzinfo=UTC),
  60. )
  61. @pytest.mark.asyncio
  62. async def test_snapshot_freeze_is_idempotent_and_conflicts_by_version(
  63. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  64. ) -> None:
  65. _, sessions = database
  66. repository = SqlAlchemyInputSnapshotRepository(sessions)
  67. first = await repository.freeze(_input())
  68. second = await repository.freeze(_input())
  69. assert first.snapshot_id == second.snapshot_id
  70. assert first.canonical_sha256 == second.canonical_sha256
  71. same_hash_new_version = await repository.freeze(_input(), version=2)
  72. assert same_hash_new_version.snapshot_id == first.snapshot_id
  73. with pytest.raises(InputHashConflict):
  74. await repository.freeze(_input(topic_result="changed"))
  75. with pytest.raises(Exception) as cross_build:
  76. await repository.get(first.snapshot_id, script_build_id=999)
  77. assert getattr(cross_build.value, "code", None) == "BUILD_NOT_FOUND"
  78. @pytest.mark.asyncio
  79. async def test_concurrent_same_snapshot_freeze_returns_one_version(
  80. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  81. ) -> None:
  82. _, sessions = database
  83. repository = SqlAlchemyInputSnapshotRepository(sessions)
  84. first, second = await asyncio.gather(repository.freeze(_input()), repository.freeze(_input()))
  85. assert first.snapshot_id == second.snapshot_id
  86. @pytest.mark.asyncio
  87. async def test_snapshot_get_recomputes_digest_and_rejects_database_tampering(database) -> None:
  88. _, sessions = database
  89. repository = SqlAlchemyInputSnapshotRepository(sessions)
  90. frozen = await repository.freeze(_input())
  91. async with sessions() as session, session.begin():
  92. await session.execute(
  93. update(input_snapshot_table)
  94. .where(input_snapshot_table.c.id == int(frozen.snapshot_id))
  95. .values(canonical_json={**frozen.to_input().canonical_payload(), "topic_id": 999})
  96. )
  97. with pytest.raises(InputDigestMismatch):
  98. await repository.get(frozen.snapshot_id, script_build_id=1)
  99. @pytest.mark.asyncio
  100. async def test_binding_create_is_idempotent_and_active_pointer_is_build_scoped(
  101. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  102. ) -> None:
  103. _, sessions = database
  104. repository = SqlAlchemyMissionBindingRepository(sessions)
  105. values = {
  106. "script_build_id": 1,
  107. "root_trace_id": "root-1",
  108. "input_snapshot_id": 2,
  109. "engine_version": "0.4.0",
  110. "schema_version": "phase-one/v1",
  111. }
  112. first, second = await asyncio.gather(repository.create(**values), repository.create(**values))
  113. assert first.binding_id == second.binding_id
  114. updated = await repository.set_active_direction(script_build_id=1, artifact_version_id=99)
  115. assert updated.active_direction_artifact_version_id == 99
  116. advanced = await repository.compare_and_set_input_snapshot(
  117. script_build_id=1,
  118. expected_snapshot_id=2,
  119. new_snapshot_id=3,
  120. )
  121. assert advanced.input_snapshot_id == 3
  122. replayed = await repository.compare_and_set_input_snapshot(
  123. script_build_id=1,
  124. expected_snapshot_id=2,
  125. new_snapshot_id=3,
  126. )
  127. assert replayed.input_snapshot_id == 3
  128. with pytest.raises(Exception, match="changed concurrently"):
  129. await repository.compare_and_set_input_snapshot(
  130. script_build_id=1,
  131. expected_snapshot_id=2,
  132. new_snapshot_id=4,
  133. )
  134. @pytest.mark.asyncio
  135. async def test_artifact_freeze_ref_ownership_digest_and_spec_replay(
  136. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  137. ) -> None:
  138. _, sessions = database
  139. repository = SqlAlchemyScriptBusinessArtifactRepository(sessions)
  140. version, reference = await repository.freeze(
  141. script_build_id=1,
  142. task_id="task-1",
  143. attempt_id="attempt-1",
  144. spec_version=1,
  145. artifact=_evidence(),
  146. )
  147. assert version.state is ArtifactState.FROZEN
  148. assert reference.uri == f"script-build://artifact-versions/{version.artifact_version_id}"
  149. assert reference.kind == "evidence"
  150. assert reference.version == str(version.artifact_version_id)
  151. assert reference.digest and reference.digest.startswith("sha256:")
  152. loaded = await repository.read_by_ref(
  153. reference, script_build_id=1, task_id="task-1", attempt_id="attempt-1"
  154. )
  155. assert loaded == version
  156. assert await repository.get_by_id(version.artifact_version_id, script_build_id=1) == version
  157. with pytest.raises(ArtifactNotFound):
  158. await repository.get_by_id(version.artifact_version_id, script_build_id=2)
  159. with pytest.raises(ArtifactOwnershipMismatch):
  160. await repository.read_by_ref(reference, script_build_id=2)
  161. tampered = ArtifactRef(
  162. uri=reference.uri,
  163. kind=reference.kind,
  164. version=reference.version,
  165. digest="sha256:" + "f" * 64,
  166. )
  167. with pytest.raises(ArtifactDigestMismatch):
  168. await repository.read_by_ref(tampered, script_build_id=1)
  169. assert not await repository.verify_digest(tampered, script_build_id=1)
  170. with pytest.raises(ArtifactAlreadyFrozen):
  171. await repository.freeze(
  172. script_build_id=1,
  173. task_id="task-1",
  174. attempt_id="attempt-1",
  175. spec_version=2,
  176. artifact=_evidence(),
  177. )
  178. class Bindings:
  179. async def get_by_root(self, _root: str):
  180. return type("Binding", (), {"script_build_id": 1})()
  181. reader = ScriptBuildArtifactEvidenceReader(repository, Bindings()) # type: ignore[arg-type]
  182. with pytest.raises(ArtifactOwnershipMismatch):
  183. await reader.read(
  184. reference,
  185. EvidenceQuery(
  186. root_trace_id="root-1",
  187. task_id="other-task",
  188. attempt_id="other-attempt",
  189. snapshot_id="snapshot-1",
  190. validation_id="validation-1",
  191. query="evidence",
  192. limit=1,
  193. ),
  194. )
  195. @pytest.mark.asyncio
  196. async def test_artifact_same_attempt_and_digest_is_concurrently_idempotent(
  197. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  198. ) -> None:
  199. _, sessions = database
  200. repository = SqlAlchemyScriptBusinessArtifactRepository(sessions)
  201. async def freeze() -> tuple[object, ArtifactRef]:
  202. return await repository.freeze(
  203. script_build_id=1,
  204. task_id="task-1",
  205. attempt_id="attempt-1",
  206. spec_version=1,
  207. artifact=_evidence(),
  208. )
  209. first, second = await asyncio.gather(freeze(), freeze())
  210. assert first[1] == second[1]
  211. @pytest.mark.asyncio
  212. async def test_direction_and_publication_are_idempotent_and_error_is_redacted(
  213. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  214. ) -> None:
  215. _, sessions = database
  216. artifacts = SqlAlchemyScriptBusinessArtifactRepository(sessions)
  217. direction, _ = await artifacts.freeze(
  218. script_build_id=1,
  219. task_id="direction-task",
  220. attempt_id="direction-attempt",
  221. spec_version=1,
  222. artifact=ScriptDirectionArtifactV1(
  223. goals=(DirectionGoal("goal-1", "make the topic clear", "source-backed"),),
  224. evidence_refs=("script-build://artifact-versions/1",),
  225. legacy_markdown="# Direction",
  226. ),
  227. )
  228. publications = SqlAlchemyPublicationRepository(sessions)
  229. first = await publications.prepare(
  230. script_build_id=1,
  231. publication_type=PublicationType.DIRECTION,
  232. accept_decision_id="decision-1",
  233. artifact_version_id=direction.artifact_version_id,
  234. expected_sha256=direction.canonical_sha256,
  235. )
  236. second = await publications.prepare(
  237. script_build_id=1,
  238. publication_type=PublicationType.DIRECTION,
  239. accept_decision_id="decision-1",
  240. artifact_version_id=direction.artifact_version_id,
  241. expected_sha256=direction.canonical_sha256,
  242. )
  243. assert first.publication_id == second.publication_id
  244. failed = await publications.mark_failed(
  245. first.publication_id,
  246. error_code="UPSTREAM",
  247. error_summary="mysql://reader:password@db.example/app Bearer raw-token",
  248. )
  249. assert failed.state is PublicationState.FAILED
  250. assert "password" not in (failed.last_error_summary or "")
  251. assert "raw-token" not in (failed.last_error_summary or "")
  252. published = await publications.mark_published(first.publication_id)
  253. assert published.state is PublicationState.PUBLISHED
  254. assert published.publication_revision == 1
  255. replayed = await publications.mark_published(first.publication_id)
  256. assert replayed.publication_revision == 1
  257. still_published = await publications.mark_failed(
  258. first.publication_id,
  259. error_code="LATE_FAILURE",
  260. error_summary="must not downgrade a committed publication",
  261. )
  262. assert still_published.state is PublicationState.PUBLISHED
  263. assert still_published.publication_revision == 1
  264. with pytest.raises(Exception, match="already bound to another publication"):
  265. await publications.prepare(
  266. script_build_id=1,
  267. publication_type=PublicationType.FINAL,
  268. accept_decision_id="decision-1",
  269. artifact_version_id=direction.artifact_version_id,
  270. expected_sha256=direction.canonical_sha256,
  271. )
  272. assert (
  273. await artifacts.get_by_id(direction.artifact_version_id, script_build_id=1)
  274. ).state is ArtifactState.PUBLISHED
  275. observed = await publications.get_by_build(1, publication_type=PublicationType.DIRECTION)
  276. assert observed == published
  277. @pytest.mark.asyncio
  278. async def test_legacy_status_read_write_and_direction_projection(
  279. database: tuple[AsyncEngine, async_sessionmaker[AsyncSession]],
  280. ) -> None:
  281. _, sessions = database
  282. repository = SqlAlchemyLegacyBuildStateRepository(sessions)
  283. build_id = await repository.create(
  284. execution_id=10,
  285. topic_build_id=20,
  286. topic_id=30,
  287. agent_type="AigcAgent",
  288. agent_config={"model": "fake"},
  289. data_source_url=None,
  290. strategies_config={"always_on": [7], "on_demand": []},
  291. )
  292. assert await repository.get_status(build_id) is BuildStatus.RUNNING
  293. await repository.set_status(build_id, BuildStatus.STOPPING)
  294. assert await repository.get_status(build_id) is BuildStatus.STOPPING
  295. await repository.project_direction(build_id, "# accepted direction")
  296. await repository.set_status(build_id, BuildStatus.PARTIAL)
  297. assert await repository.get_status(build_id) is BuildStatus.PARTIAL
  298. await repository.set_checkpoint(
  299. build_id,
  300. checkpoint_code="PHASE_ONE_CAPABILITY_BOUNDARY",
  301. summary="direction ready",
  302. root_trace_id="root-1",
  303. )
  304. async with sessions() as session:
  305. row = (
  306. (
  307. await session.execute(
  308. select(script_build_record).where(script_build_record.c.id == build_id)
  309. )
  310. )
  311. .mappings()
  312. .one()
  313. )
  314. assert row["status"] == "partial"
  315. assert row["error_message"] == "PHASE_ONE_CAPABILITY_BOUNDARY"
  316. assert row["summary"] == "direction ready"
  317. assert row["reson_trace_id"] == "root-1"
  318. assert row["end_time"] is not None
  319. await repository.set_status(build_id, BuildStatus.STOPPING)
  320. async with sessions() as session:
  321. stopping_row = (
  322. (
  323. await session.execute(
  324. select(script_build_record).where(script_build_record.c.id == build_id)
  325. )
  326. )
  327. .mappings()
  328. .one()
  329. )
  330. assert stopping_row["status"] == "stopping"
  331. assert stopping_row["error_message"] is None
  332. assert stopping_row["summary"] is None
  333. assert stopping_row["end_time"] is None
  334. await repository.begin_phase(build_id)
  335. async with sessions() as session:
  336. row = (
  337. (
  338. await session.execute(
  339. select(script_build_record).where(script_build_record.c.id == build_id)
  340. )
  341. )
  342. .mappings()
  343. .one()
  344. )
  345. assert row["status"] == "running"
  346. assert row["end_time"] is None
  347. assert row["summary"] is None
  348. with pytest.raises(Exception, match="cannot project script build success"):
  349. await repository.set_status(build_id, BuildStatus.SUCCESS)