"""Atomic writer coordination for planning rows and the legacy query contract.""" from __future__ import annotations from dataclasses import dataclass, field from typing import Any from uuid import UUID, uuid4 from query_planning.domain import GenerationResult, KnowledgeNeedDraft from query_planning.ports import LegacyQueryBatchSink, QueryPlanningStore @dataclass(frozen=True) class QueryBatchWriteSpec: name: str source_type: str generation_method: str target_platforms: tuple[str, ...] metadata: dict[str, Any] = field(default_factory=dict) def __post_init__(self) -> None: if not self.name.strip(): raise ValueError("query batch name must not be empty") if not self.target_platforms: raise ValueError("query batch requires at least one target platform") @dataclass(frozen=True) class QueryBatchWriteResult: plan_id: UUID batch: Any queries: tuple[Any, ...] class NullQueryPlanningStore: """Compatibility adapter for legacy fakes that expose no DB connection.""" def create_plan(self, **kwargs: Any) -> UUID: return uuid4() def add_knowledge_need(self, *, plan_id: UUID, need: KnowledgeNeedDraft) -> UUID: return uuid4() def link_plan_batch(self, *, plan_id: UUID, batch_id: UUID) -> None: return None def link_query_need(self, *, query_id: UUID, knowledge_need_id: UUID) -> None: return None class QueryBatchWriter: def __init__(self, *, legacy_sink: LegacyQueryBatchSink, planning_store: QueryPlanningStore) -> None: self.legacy_sink = legacy_sink self.planning_store = planning_store def write(self, result: GenerationResult, spec: QueryBatchWriteSpec) -> QueryBatchWriteResult: if not result.selected_queries: raise ValueError("cannot persist an empty query selection") declared_need_keys = [need.need_key for need in result.knowledge_needs] if len(set(declared_need_keys)) != len(declared_need_keys): raise ValueError("knowledge need keys must be unique inside one plan") referenced_need_keys = { need_key for candidate in result.selected_queries for need_key in candidate.knowledge_need_keys } unknown_need_keys = referenced_need_keys - set(declared_need_keys) if unknown_need_keys: raise ValueError( "query references unknown knowledge need key(s): " + ", ".join(sorted(unknown_need_keys)) ) stats = result.stats plan_id = self.planning_store.create_plan( generator_kind=result.plan.generator_kind.value, status="ready", input_snapshot=result.plan.input_snapshot, generator_config=result.plan.generator_config, max_queries=result.request.max_queries, generated_count=stats.generated_count, normalized_count=stats.normalized_count, unique_count=stats.unique_count, selected_count=stats.selected_count, dropped_count=stats.dropped_count, metadata=result.plan.metadata, ) need_ids = { need.need_key: self.planning_store.add_knowledge_need(plan_id=plan_id, need=need) for need in result.knowledge_needs } batch_metadata = dict(spec.metadata) batch_metadata["query_planning"] = { "plan_id": str(plan_id), "generator_kind": result.plan.generator_kind.value, "generated_count": stats.generated_count, "unique_count": stats.unique_count, "selected_count": stats.selected_count, "dropped_count": stats.dropped_count, } batch = self.legacy_sink.create_query_batch( name=spec.name, source_type=spec.source_type, generation_method=spec.generation_method, target_platforms=list(spec.target_platforms), status="ready", metadata=batch_metadata, ) if getattr(batch, "id", None) is None: raise RuntimeError("legacy query batch sink returned a batch without id") written: list[Any] = [] for sort_order, candidate in enumerate(result.selected_queries): query = self.legacy_sink.add_query( batch_id=batch.id, query_text=candidate.query_text, axes=candidate.axes, keep=True, filter_reason=candidate.filter_reason, status="ready", sort_order=sort_order, metadata=candidate.metadata, ) if getattr(query, "id", None) is None: raise RuntimeError("legacy query batch sink returned a query without id") written.append(query) self.planning_store.link_plan_batch(plan_id=plan_id, batch_id=batch.id) for query, candidate in zip(written, result.selected_queries, strict=True): for need_key in candidate.knowledge_need_keys: need_id = need_ids.get(need_key) if need_id is not None: self.planning_store.link_query_need( query_id=query.id, knowledge_need_id=need_id, ) return QueryBatchWriteResult(plan_id=plan_id, batch=batch, queries=tuple(written)) def planning_store_for_repository(repo: Any) -> QueryPlanningStore: conn = getattr(repo, "conn", None) if conn is None: return NullQueryPlanningStore() from query_planning.repositories.postgres import PostgresQueryPlanningStore return PostgresQueryPlanningStore(conn)