from __future__ import annotations import asyncio from types import SimpleNamespace import pytest from sqlalchemy import event, insert, select from script_build_host.application.input_snapshot_service import ScriptInputSnapshotService from script_build_host.application.mission_factory import ScriptMissionFactory from script_build_host.application.mission_service import ( BuildTransitionGate, ScriptMissionService, StartScriptBuildCommand, ) from script_build_host.domain.ports import PersonaInput from script_build_host.domain.records import BuildStatus, Principal from script_build_host.infrastructure.legacy_tables import ( script_build_record, topic_build_record, topic_build_topic, topic_pattern_execution, ) from script_build_host.infrastructure.tables import ( input_snapshot_table, mission_binding_table, ) from script_build_host.repositories import ( LegacySqlAlchemyInputReader, SqlAlchemyInputSnapshotRepository, SqlAlchemyLegacyBuildStateRepository, SqlAlchemyMissionBindingRepository, ) class _Persona: async def load(self, account_name: str) -> PersonaInput: return PersonaInput(account_name, account_name, (), (), {}) class _Strategies: async def load(self, **_selectors: object) -> tuple[dict[str, object], ...]: return () class _Prompts: async def load(self, _requests: object) -> tuple[dict[str, object], ...]: return () class _Authorizer: def __init__(self) -> None: self.source_checked = False async def require_source_access(self, _principal: Principal, **source: int) -> None: assert source == {"execution_id": 10, "topic_build_id": 20, "topic_id": 30} self.source_checked = True async def require_access(self, _principal: Principal, _build: int) -> None: return None class _StartOnlyMissionService(ScriptMissionService): def __init__(self, **kwargs: object) -> None: super().__init__(**kwargs) # type: ignore[arg-type] self.started_builds: list[int] = [] async def run(self, script_build_id: int) -> None: self.started_builds.append(script_build_id) @pytest.mark.asyncio async def test_start_uses_real_input_snapshot_binding_and_phase_one_write_allowlist( database, ) -> None: engine, sessions = database async with sessions() as session, session.begin(): await session.execute(insert(topic_pattern_execution).values(id=10, status="success")) await session.execute( insert(topic_build_record).values( id=20, execution_id=10, demand="topic demand", status="success", is_deleted=False, personal_config={"account_name": "acct"}, origin="generated", ) ) await session.execute( insert(topic_build_topic).values( id=30, build_id=20, execution_id=10, sort_order=0, result="topic", status="mature", ) ) statements: list[str] = [] def record_sql( _connection: object, _cursor: object, statement: str, _parameters: object, _context: object, _executemany: object, ) -> None: if statement.lstrip().upper().startswith(("INSERT", "UPDATE", "DELETE")): statements.append(statement.lower()) event.listen(engine.sync_engine, "before_cursor_execute", record_sql) try: snapshots = SqlAlchemyInputSnapshotRepository(sessions) bindings = SqlAlchemyMissionBindingRepository(sessions) legacy = SqlAlchemyLegacyBuildStateRepository(sessions) input_service = ScriptInputSnapshotService( legacy_input=LegacySqlAlchemyInputReader(sessions), persona_source=_Persona(), strategy_source=_Strategies(), prompt_source=_Prompts(), snapshots=snapshots, ) authorizer = _Authorizer() service = _StartOnlyMissionService( runner=SimpleNamespace(), coordinator=SimpleNamespace(), factory=ScriptMissionFactory(), input_snapshots=input_service, bindings=bindings, legacy_state=legacy, authorizer=authorizer, direction_reconciler=SimpleNamespace(), transition_gate=BuildTransitionGate(), ) result = await service.start( StartScriptBuildCommand( execution_id=10, topic_build_id=20, topic_id=30, principal=Principal("owner"), runtime_prompt_manifest=( { "preset": "script_planner", "source": "fixture", "content_sha256": "sha256:" + "1" * 64, }, ), model_manifest={"presets": {"script_planner": {"model": "fake-model"}}}, ) ) await asyncio.sleep(0) finally: event.remove(engine.sync_engine, "before_cursor_execute", record_sql) assert authorizer.source_checked assert result.status is BuildStatus.RUNNING assert service.started_builds == [result.script_build_id] binding = await bindings.get_by_build(result.script_build_id) snapshot = await snapshots.get( str(binding.input_snapshot_id), script_build_id=result.script_build_id, ) assert snapshot.topic["execution"]["id"] == 10 assert snapshot.account["account_name"] == "acct" async with sessions() as session: assert await session.scalar(select(script_build_record.c.status)) == "running" assert await session.scalar(select(input_snapshot_table.c.id)) is not None assert await session.scalar(select(mission_binding_table.c.id)) is not None allowed = { "script_build_record", "script_build_input_snapshot", "script_build_mission_binding", } assert statements assert all(any(table in statement for table in allowed) for statement in statements) forbidden = ("paragraph", "element", "link", "round", "branch", "plan_step", "external_log") assert not any(name in statement for statement in statements for name in forbidden)