writer.py 5.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144
  1. """Atomic writer coordination for planning rows and the legacy query contract."""
  2. from __future__ import annotations
  3. from dataclasses import dataclass, field
  4. from typing import Any
  5. from uuid import UUID, uuid4
  6. from query_planning.domain import GenerationResult, KnowledgeNeedDraft
  7. from query_planning.ports import LegacyQueryBatchSink, QueryPlanningStore
  8. @dataclass(frozen=True)
  9. class QueryBatchWriteSpec:
  10. name: str
  11. source_type: str
  12. generation_method: str
  13. target_platforms: tuple[str, ...]
  14. metadata: dict[str, Any] = field(default_factory=dict)
  15. def __post_init__(self) -> None:
  16. if not self.name.strip():
  17. raise ValueError("query batch name must not be empty")
  18. if not self.target_platforms:
  19. raise ValueError("query batch requires at least one target platform")
  20. @dataclass(frozen=True)
  21. class QueryBatchWriteResult:
  22. plan_id: UUID
  23. batch: Any
  24. queries: tuple[Any, ...]
  25. class NullQueryPlanningStore:
  26. """Compatibility adapter for legacy fakes that expose no DB connection."""
  27. def create_plan(self, **kwargs: Any) -> UUID:
  28. return uuid4()
  29. def add_knowledge_need(self, *, plan_id: UUID, need: KnowledgeNeedDraft) -> UUID:
  30. return uuid4()
  31. def link_plan_batch(self, *, plan_id: UUID, batch_id: UUID) -> None:
  32. return None
  33. def link_query_need(self, *, query_id: UUID, knowledge_need_id: UUID) -> None:
  34. return None
  35. class QueryBatchWriter:
  36. def __init__(self, *, legacy_sink: LegacyQueryBatchSink, planning_store: QueryPlanningStore) -> None:
  37. self.legacy_sink = legacy_sink
  38. self.planning_store = planning_store
  39. def write(self, result: GenerationResult, spec: QueryBatchWriteSpec) -> QueryBatchWriteResult:
  40. if not result.selected_queries:
  41. raise ValueError("cannot persist an empty query selection")
  42. declared_need_keys = [need.need_key for need in result.knowledge_needs]
  43. if len(set(declared_need_keys)) != len(declared_need_keys):
  44. raise ValueError("knowledge need keys must be unique inside one plan")
  45. referenced_need_keys = {
  46. need_key
  47. for candidate in result.selected_queries
  48. for need_key in candidate.knowledge_need_keys
  49. }
  50. unknown_need_keys = referenced_need_keys - set(declared_need_keys)
  51. if unknown_need_keys:
  52. raise ValueError(
  53. "query references unknown knowledge need key(s): "
  54. + ", ".join(sorted(unknown_need_keys))
  55. )
  56. stats = result.stats
  57. plan_id = self.planning_store.create_plan(
  58. generator_kind=result.plan.generator_kind.value,
  59. status="ready",
  60. input_snapshot=result.plan.input_snapshot,
  61. generator_config=result.plan.generator_config,
  62. max_queries=result.request.max_queries,
  63. generated_count=stats.generated_count,
  64. normalized_count=stats.normalized_count,
  65. unique_count=stats.unique_count,
  66. selected_count=stats.selected_count,
  67. dropped_count=stats.dropped_count,
  68. metadata=result.plan.metadata,
  69. )
  70. need_ids = {
  71. need.need_key: self.planning_store.add_knowledge_need(plan_id=plan_id, need=need)
  72. for need in result.knowledge_needs
  73. }
  74. batch_metadata = dict(spec.metadata)
  75. batch_metadata["query_planning"] = {
  76. "plan_id": str(plan_id),
  77. "generator_kind": result.plan.generator_kind.value,
  78. "generated_count": stats.generated_count,
  79. "unique_count": stats.unique_count,
  80. "selected_count": stats.selected_count,
  81. "dropped_count": stats.dropped_count,
  82. }
  83. batch = self.legacy_sink.create_query_batch(
  84. name=spec.name,
  85. source_type=spec.source_type,
  86. generation_method=spec.generation_method,
  87. target_platforms=list(spec.target_platforms),
  88. status="ready",
  89. metadata=batch_metadata,
  90. )
  91. if getattr(batch, "id", None) is None:
  92. raise RuntimeError("legacy query batch sink returned a batch without id")
  93. written: list[Any] = []
  94. for sort_order, candidate in enumerate(result.selected_queries):
  95. query = self.legacy_sink.add_query(
  96. batch_id=batch.id,
  97. query_text=candidate.query_text,
  98. axes=candidate.axes,
  99. keep=True,
  100. filter_reason=candidate.filter_reason,
  101. status="ready",
  102. sort_order=sort_order,
  103. metadata=candidate.metadata,
  104. )
  105. if getattr(query, "id", None) is None:
  106. raise RuntimeError("legacy query batch sink returned a query without id")
  107. written.append(query)
  108. self.planning_store.link_plan_batch(plan_id=plan_id, batch_id=batch.id)
  109. for query, candidate in zip(written, result.selected_queries, strict=True):
  110. for need_key in candidate.knowledge_need_keys:
  111. need_id = need_ids.get(need_key)
  112. if need_id is not None:
  113. self.planning_store.link_query_need(
  114. query_id=query.id,
  115. knowledge_need_id=need_id,
  116. )
  117. return QueryBatchWriteResult(plan_id=plan_id, batch=batch, queries=tuple(written))
  118. def planning_store_for_repository(repo: Any) -> QueryPlanningStore:
  119. conn = getattr(repo, "conn", None)
  120. if conn is None:
  121. return NullQueryPlanningStore()
  122. from query_planning.repositories.postgres import PostgresQueryPlanningStore
  123. return PostgresQueryPlanningStore(conn)