test_query_planning_core.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376
  1. from __future__ import annotations
  2. from pathlib import Path
  3. from uuid import uuid4
  4. import pytest
  5. from acquisition.domain import Query, QueryBatch
  6. from acquisition.queries.builder import persist_query_batch
  7. from core import db_session
  8. from query_planning import (
  9. GenerationRequest,
  10. GeneratorKind,
  11. NoQueryCandidatesError,
  12. QueryBatchWriteSpec,
  13. QueryBatchWriter,
  14. UnifiedQueryGenerationService,
  15. UnsupportedGeneratorError,
  16. )
  17. def _request(kind=GeneratorKind.MANUAL, *, candidates, max_queries=None):
  18. return GenerationRequest(
  19. generator_kind=kind,
  20. name="test-plan",
  21. payload={"candidates": candidates, "input_snapshot": {"case": "test"}},
  22. max_queries=max_queries,
  23. )
  24. def test_normalizes_and_exactly_dedupes_while_preserving_near_queries():
  25. result = UnifiedQueryGenerationService().generate(
  26. _request(
  27. candidates=[
  28. {
  29. "query_text": " A B ",
  30. "axes": {"first": True},
  31. "metadata": {"family_key": "first"},
  32. "source_refs": [{"family_key": "first"}],
  33. },
  34. {
  35. "query_text": "a b",
  36. "priority": 8,
  37. "axes": {"second": True},
  38. "metadata": {"family_key": "second"},
  39. "source_refs": [{"family_key": "second"}],
  40. },
  41. {
  42. "query_text": "a b 方法",
  43. "source_refs": [{"family_key": "near"}],
  44. },
  45. ]
  46. )
  47. )
  48. assert [query.query_text for query in result.selected_queries] == ["A B", "a b 方法"]
  49. merged = result.selected_queries[0]
  50. assert merged.axes == {"first": True}
  51. assert merged.metadata["family_key"] == "first"
  52. assert merged.priority == 8
  53. assert merged.metadata["origins"] == [
  54. {"family_key": "first"},
  55. {"family_key": "second"},
  56. ]
  57. assert result.stats.generated_count == 3
  58. assert result.stats.unique_count == 2
  59. assert result.stats.dropped_count == 1
  60. def test_priority_budget_is_stable_and_default_is_unlimited():
  61. candidates = [
  62. {"query_text": "q0", "priority": 0},
  63. {"query_text": "q1", "priority": 5},
  64. {"query_text": "q2", "priority": 5},
  65. {"query_text": "q3", "priority": 1},
  66. ]
  67. service = UnifiedQueryGenerationService()
  68. unlimited = service.generate(_request(candidates=candidates))
  69. budgeted = service.generate(_request(candidates=candidates, max_queries=2))
  70. assert [query.query_text for query in unlimited.selected_queries] == ["q1", "q2", "q3", "q0"]
  71. assert [query.query_text for query in budgeted.selected_queries] == ["q1", "q2"]
  72. assert budgeted.stats.selected_count == 2
  73. assert budgeted.stats.dropped_count == 2
  74. def test_empty_selection_and_unregistered_reserved_generator_fail_explicitly():
  75. service = UnifiedQueryGenerationService()
  76. with pytest.raises(NoQueryCandidatesError):
  77. service.generate(_request(candidates=[" "]))
  78. with pytest.raises(NoQueryCandidatesError):
  79. service.generate(_request(candidates=["valid"], max_queries=0))
  80. with pytest.raises(UnsupportedGeneratorError, match="topic_table"):
  81. service.generate(
  82. _request(kind=GeneratorKind.TOPIC_TABLE, candidates=["not allowed yet"])
  83. )
  84. class FakeLegacySink:
  85. def __init__(self):
  86. self.batch_kwargs = None
  87. self.query_kwargs = []
  88. def create_query_batch(self, **kwargs):
  89. self.batch_kwargs = kwargs
  90. return QueryBatch(id=uuid4(), **kwargs)
  91. def add_query(self, **kwargs):
  92. self.query_kwargs.append(kwargs)
  93. return Query(id=uuid4(), **kwargs)
  94. class FakePlanningStore:
  95. def __init__(self):
  96. self.plan_id = uuid4()
  97. self.plans = []
  98. self.needs = []
  99. self.plan_batch_links = []
  100. self.query_need_links = []
  101. def create_plan(self, **kwargs):
  102. self.plans.append(kwargs)
  103. return self.plan_id
  104. def add_knowledge_need(self, *, plan_id, need):
  105. need_id = uuid4()
  106. self.needs.append((plan_id, need, need_id))
  107. return need_id
  108. def link_plan_batch(self, **kwargs):
  109. self.plan_batch_links.append(kwargs)
  110. def link_query_need(self, **kwargs):
  111. self.query_need_links.append(kwargs)
  112. def test_writer_keeps_legacy_search_contract_and_links_plan_need_rows():
  113. service = UnifiedQueryGenerationService()
  114. result = service.generate(
  115. GenerationRequest(
  116. generator_kind=GeneratorKind.MANUAL,
  117. name="manual",
  118. payload={
  119. "knowledge_needs": [
  120. {
  121. "need_key": "need-1",
  122. "decision_context": "选择画面主线",
  123. "unknown_information": "透明伞怎样形成柔光",
  124. "source_ref": {"topic_id": 428},
  125. }
  126. ],
  127. "candidates": [
  128. {
  129. "query_text": "透明伞 人像 柔光",
  130. "axes": {"道具": "透明伞"},
  131. "metadata": {"family_key": "manual"},
  132. "source_refs": [{"generator_kind": "manual"}],
  133. "knowledge_need_keys": ["need-1"],
  134. }
  135. ],
  136. },
  137. )
  138. )
  139. sink = FakeLegacySink()
  140. planning = FakePlanningStore()
  141. written = QueryBatchWriter(legacy_sink=sink, planning_store=planning).write(
  142. result,
  143. QueryBatchWriteSpec(
  144. name="manual",
  145. source_type="manual",
  146. generation_method="manual_query_api_v1",
  147. target_platforms=("xiaohongshu",),
  148. metadata={"source": "test"},
  149. ),
  150. )
  151. assert written.plan_id == planning.plan_id
  152. assert sink.batch_kwargs["status"] == "ready"
  153. assert sink.batch_kwargs["target_platforms"] == ["xiaohongshu"]
  154. assert sink.batch_kwargs["metadata"]["query_planning"]["plan_id"] == str(planning.plan_id)
  155. query = sink.query_kwargs[0]
  156. assert query["query_text"] == "透明伞 人像 柔光"
  157. assert query["keep"] is True
  158. assert query["status"] == "ready"
  159. assert query["sort_order"] == 0
  160. assert query["filter_reason"] is None
  161. assert query["metadata"]["family_key"] == "manual"
  162. assert planning.plan_batch_links == [
  163. {"plan_id": planning.plan_id, "batch_id": written.batch.id}
  164. ]
  165. assert len(planning.query_need_links) == 1
  166. def test_writer_rejects_unknown_need_before_any_write():
  167. result = UnifiedQueryGenerationService().generate(
  168. _request(
  169. candidates=[
  170. {
  171. "query_text": "query",
  172. "knowledge_need_keys": ["missing"],
  173. }
  174. ]
  175. )
  176. )
  177. sink = FakeLegacySink()
  178. planning = FakePlanningStore()
  179. with pytest.raises(ValueError, match="missing"):
  180. QueryBatchWriter(legacy_sink=sink, planning_store=planning).write(
  181. result,
  182. QueryBatchWriteSpec(
  183. name="manual",
  184. source_type="manual",
  185. generation_method="manual_query_api_v1",
  186. target_platforms=("xiaohongshu",),
  187. ),
  188. )
  189. assert planning.plans == []
  190. assert sink.batch_kwargs is None
  191. def test_outer_transaction_rolls_back_partial_writer_failure(monkeypatch):
  192. class FakeConnection:
  193. def __init__(self):
  194. self.pending = []
  195. self.committed = []
  196. self.rollback_called = False
  197. self.closed = False
  198. def commit(self):
  199. self.committed.extend(self.pending)
  200. self.pending.clear()
  201. def rollback(self):
  202. self.rollback_called = True
  203. self.pending.clear()
  204. def close(self):
  205. self.closed = True
  206. class TransactionalPlanningStore(FakePlanningStore):
  207. def __init__(self, conn):
  208. super().__init__()
  209. self.conn = conn
  210. def create_plan(self, **kwargs):
  211. self.conn.pending.append(("plan", kwargs))
  212. return self.plan_id
  213. def link_plan_batch(self, **kwargs):
  214. self.conn.pending.append(("plan_batch", kwargs))
  215. class FailingLegacySink(FakeLegacySink):
  216. def __init__(self, conn):
  217. super().__init__()
  218. self.conn = conn
  219. def create_query_batch(self, **kwargs):
  220. batch = super().create_query_batch(**kwargs)
  221. self.conn.pending.append(("batch", batch.id))
  222. return batch
  223. def add_query(self, **kwargs):
  224. if len(self.query_kwargs) == 1:
  225. raise RuntimeError("second query failed")
  226. query = super().add_query(**kwargs)
  227. self.conn.pending.append(("query", query.id))
  228. return query
  229. result = UnifiedQueryGenerationService().generate(
  230. _request(candidates=["first query", "second query"])
  231. )
  232. conn = FakeConnection()
  233. monkeypatch.setattr(db_session, "connect", lambda _config: conn)
  234. with pytest.raises(RuntimeError, match="second query failed"):
  235. with db_session.transaction(object()) as transaction_conn:
  236. QueryBatchWriter(
  237. legacy_sink=FailingLegacySink(transaction_conn),
  238. planning_store=TransactionalPlanningStore(transaction_conn),
  239. ).write(
  240. result,
  241. QueryBatchWriteSpec(
  242. name="rollback",
  243. source_type="manual",
  244. generation_method="manual_query_api_v1",
  245. target_platforms=("xiaohongshu",),
  246. ),
  247. )
  248. assert conn.rollback_called is True
  249. assert conn.pending == []
  250. assert conn.committed == []
  251. assert conn.closed is True
  252. def test_cartesian_legacy_facade_exactly_dedupes_across_families():
  253. sink = FakeLegacySink()
  254. generated = {
  255. "metadata": {"active_family_keys": ["f1", "f2"]},
  256. "families": [
  257. {
  258. "key": "f1",
  259. "name": "实质 × 模态",
  260. "axes": ["实质", "模态"],
  261. "items": [
  262. {"query": "A B", "parts": {"实质": "A"}, "keep": True}
  263. ],
  264. },
  265. {
  266. "key": "f2",
  267. "name": "形式 × 模态",
  268. "axes": ["形式", "模态"],
  269. "items": [
  270. {"query": "a b", "parts": {"形式": "A"}, "keep": True}
  271. ],
  272. },
  273. ],
  274. }
  275. batch, count = persist_query_batch(sink, generated, name="cartesian")
  276. assert batch.id is not None
  277. assert count == 1
  278. assert sink.query_kwargs[0]["query_text"] == "A B"
  279. assert sink.query_kwargs[0]["sort_order"] == 0
  280. assert sink.query_kwargs[0]["metadata"]["family_key"] == "f1"
  281. assert [origin["family_key"] for origin in sink.query_kwargs[0]["metadata"]["origins"]] == [
  282. "f1",
  283. "f2",
  284. ]
  285. def test_migration_is_additive_replayable_and_does_not_force_what_how_why():
  286. sql = Path("db/migrations/005_query_planning_schema.sql").read_text(encoding="utf-8")
  287. for table in (
  288. "query_plans",
  289. "knowledge_needs",
  290. "query_plan_batches",
  291. "query_knowledge_need_links",
  292. ):
  293. assert f"CREATE TABLE IF NOT EXISTS creation_knowledge.{table}" in sql
  294. assert "ALTER TABLE creation_knowledge.query_batches" not in sql
  295. assert "ALTER TABLE creation_knowledge.queries" not in sql
  296. assert "005_query_planning_schema" in sql
  297. assert "trg_query_plans_touch_updated_at" in sql
  298. assert "TO ck_app" in sql
  299. assert "particle_type" not in sql
  300. def test_frozen_search_modules_do_not_depend_on_query_planning():
  301. frozen_roots = [
  302. Path("acquisition/runner.py"),
  303. Path("acquisition/repositories"),
  304. Path("acquisition/platforms"),
  305. Path("pipeline"),
  306. Path("decode_content"),
  307. ]
  308. for root in frozen_roots:
  309. files = [root] if root.is_file() else list(root.rglob("*.py"))
  310. for path in files:
  311. assert "query_planning" not in path.read_text(encoding="utf-8"), path
  312. forbidden = (
  313. "acquisition.runner",
  314. "acquisition.platforms",
  315. "acquisition.search",
  316. "pipeline",
  317. "decode_content",
  318. )
  319. for path in Path("query_planning").rglob("*.py"):
  320. source = path.read_text(encoding="utf-8")
  321. assert not any(f"from {module}" in source or f"import {module}" in source for module in forbidden), path