| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144 |
- """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)
|