test_start_integration.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181
  1. from __future__ import annotations
  2. import asyncio
  3. from types import SimpleNamespace
  4. import pytest
  5. from sqlalchemy import event, insert, select
  6. from script_build_host.application.input_snapshot_service import ScriptInputSnapshotService
  7. from script_build_host.application.mission_factory import ScriptMissionFactory
  8. from script_build_host.application.mission_service import (
  9. BuildTransitionGate,
  10. ScriptMissionService,
  11. StartScriptBuildCommand,
  12. )
  13. from script_build_host.domain.ports import PersonaInput
  14. from script_build_host.domain.records import BuildStatus, Principal
  15. from script_build_host.infrastructure.legacy_tables import (
  16. script_build_record,
  17. topic_build_record,
  18. topic_build_topic,
  19. topic_pattern_execution,
  20. )
  21. from script_build_host.infrastructure.tables import (
  22. input_snapshot_table,
  23. mission_binding_table,
  24. )
  25. from script_build_host.repositories import (
  26. LegacySqlAlchemyInputReader,
  27. SqlAlchemyInputSnapshotRepository,
  28. SqlAlchemyLegacyBuildStateRepository,
  29. SqlAlchemyMissionBindingRepository,
  30. )
  31. class _Persona:
  32. async def load(self, account_name: str) -> PersonaInput:
  33. return PersonaInput(account_name, account_name, (), (), {})
  34. class _Strategies:
  35. async def load(self, **_selectors: object) -> tuple[dict[str, object], ...]:
  36. return ()
  37. class _Prompts:
  38. async def load(self, _requests: object) -> tuple[dict[str, object], ...]:
  39. return ()
  40. class _Authorizer:
  41. def __init__(self) -> None:
  42. self.source_checked = False
  43. async def require_source_access(self, _principal: Principal, **source: int) -> None:
  44. assert source == {"execution_id": 10, "topic_build_id": 20, "topic_id": 30}
  45. self.source_checked = True
  46. async def require_access(self, _principal: Principal, _build: int) -> None:
  47. return None
  48. class _StartOnlyMissionService(ScriptMissionService):
  49. def __init__(self, **kwargs: object) -> None:
  50. super().__init__(**kwargs) # type: ignore[arg-type]
  51. self.started_builds: list[int] = []
  52. async def run(self, script_build_id: int) -> None:
  53. self.started_builds.append(script_build_id)
  54. @pytest.mark.asyncio
  55. async def test_start_uses_real_input_snapshot_binding_and_phase_one_write_allowlist(
  56. database,
  57. ) -> None:
  58. engine, sessions = database
  59. async with sessions() as session, session.begin():
  60. await session.execute(insert(topic_pattern_execution).values(id=10, status="success"))
  61. await session.execute(
  62. insert(topic_build_record).values(
  63. id=20,
  64. execution_id=10,
  65. demand="topic demand",
  66. status="success",
  67. is_deleted=False,
  68. personal_config={"account_name": "acct"},
  69. origin="generated",
  70. )
  71. )
  72. await session.execute(
  73. insert(topic_build_topic).values(
  74. id=30,
  75. build_id=20,
  76. execution_id=10,
  77. sort_order=0,
  78. result="topic",
  79. status="mature",
  80. )
  81. )
  82. statements: list[str] = []
  83. def record_sql(
  84. _connection: object,
  85. _cursor: object,
  86. statement: str,
  87. _parameters: object,
  88. _context: object,
  89. _executemany: object,
  90. ) -> None:
  91. if statement.lstrip().upper().startswith(("INSERT", "UPDATE", "DELETE")):
  92. statements.append(statement.lower())
  93. event.listen(engine.sync_engine, "before_cursor_execute", record_sql)
  94. try:
  95. snapshots = SqlAlchemyInputSnapshotRepository(sessions)
  96. bindings = SqlAlchemyMissionBindingRepository(sessions)
  97. legacy = SqlAlchemyLegacyBuildStateRepository(sessions)
  98. input_service = ScriptInputSnapshotService(
  99. legacy_input=LegacySqlAlchemyInputReader(sessions),
  100. persona_source=_Persona(),
  101. strategy_source=_Strategies(),
  102. prompt_source=_Prompts(),
  103. snapshots=snapshots,
  104. )
  105. authorizer = _Authorizer()
  106. service = _StartOnlyMissionService(
  107. runner=SimpleNamespace(),
  108. coordinator=SimpleNamespace(),
  109. factory=ScriptMissionFactory(),
  110. input_snapshots=input_service,
  111. bindings=bindings,
  112. legacy_state=legacy,
  113. authorizer=authorizer,
  114. direction_reconciler=SimpleNamespace(),
  115. transition_gate=BuildTransitionGate(),
  116. )
  117. result = await service.start(
  118. StartScriptBuildCommand(
  119. execution_id=10,
  120. topic_build_id=20,
  121. topic_id=30,
  122. principal=Principal("owner"),
  123. runtime_prompt_manifest=(
  124. {
  125. "preset": "script_planner",
  126. "source": "fixture",
  127. "content_sha256": "sha256:" + "1" * 64,
  128. },
  129. ),
  130. model_manifest={"presets": {"script_planner": {"model": "fake-model"}}},
  131. )
  132. )
  133. await asyncio.sleep(0)
  134. finally:
  135. event.remove(engine.sync_engine, "before_cursor_execute", record_sql)
  136. assert authorizer.source_checked
  137. assert result.status is BuildStatus.RUNNING
  138. assert service.started_builds == [result.script_build_id]
  139. binding = await bindings.get_by_build(result.script_build_id)
  140. snapshot = await snapshots.get(
  141. str(binding.input_snapshot_id),
  142. script_build_id=result.script_build_id,
  143. )
  144. assert snapshot.topic["execution"]["id"] == 10
  145. assert snapshot.account["account_name"] == "acct"
  146. async with sessions() as session:
  147. assert await session.scalar(select(script_build_record.c.status)) == "running"
  148. assert await session.scalar(select(input_snapshot_table.c.id)) is not None
  149. assert await session.scalar(select(mission_binding_table.c.id)) is not None
  150. allowed = {
  151. "script_build_record",
  152. "script_build_input_snapshot",
  153. "script_build_mission_binding",
  154. }
  155. assert statements
  156. assert all(any(table in statement for table in allowed) for statement in statements)
  157. forbidden = ("paragraph", "element", "link", "round", "branch", "plan_step", "external_log")
  158. assert not any(name in statement for statement in statements for name in forbidden)