| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181 |
- 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)
|