test_query_builder.py 7.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216
  1. from __future__ import annotations
  2. import json
  3. from uuid import uuid4
  4. import pytest
  5. from acquisition.domain import Query, QueryBatch
  6. from acquisition.queries.builder import (
  7. QueryBuildOptions,
  8. TREES,
  9. build_creation_query_batch,
  10. persist_query_batch,
  11. )
  12. from core.config import PgConfig, Settings
  13. def _settings() -> Settings:
  14. return Settings(
  15. pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
  16. aiddit_crawler_base_url="http://crawler.test",
  17. crawler_timeout=30,
  18. openrouter_timeout_seconds=90,
  19. openrouter_model="m",
  20. openrouter_base_url="http://openrouter.test",
  21. openrouter_api_key="k",
  22. llm_model="m",
  23. max_cards=12,
  24. frames_dir="f",
  25. douyin_ratio="540p",
  26. data_dir="",
  27. )
  28. class FakeRepo:
  29. def __init__(self) -> None:
  30. self.batch_kwargs = None
  31. self.queries = []
  32. def create_query_batch(self, **kwargs):
  33. self.batch_kwargs = kwargs
  34. return QueryBatch(id=uuid4(), **kwargs)
  35. def add_query(self, **kwargs):
  36. self.queries.append(kwargs)
  37. return Query(id=uuid4(), **kwargs)
  38. def _expected_l3_l4_names(source_type: str) -> set[str]:
  39. rows = json.loads(TREES.read_text("utf-8"))
  40. names: set[str] = set()
  41. for row in rows:
  42. if row.get("source_type") != source_type:
  43. continue
  44. depth = len([part for part in (row.get("path") or "").split("/") if part])
  45. if depth in (3, 4) and row.get("name"):
  46. names.add(row["name"])
  47. return names
  48. def test_build_creation_query_batch_defaults_to_first_two_families():
  49. generated = build_creation_query_batch(
  50. _settings(),
  51. options=QueryBuildOptions(per=2, batch_n=4),
  52. )
  53. assert [family["key"] for family in generated["families"]] == ["f1", "f2"]
  54. assert generated["metadata"]["active_family_keys"] == ["f1", "f2"]
  55. assert sum(len(family["items"]) for family in generated["families"]) == 4
  56. def test_build_creation_query_batch_uses_all_l3_l4_substance_and_form_nodes():
  57. generated = build_creation_query_batch(
  58. _settings(),
  59. options=QueryBuildOptions(per=2, batch_n=999),
  60. )
  61. assert set(generated["axis_values"]["实质"]) == _expected_l3_l4_names("实质")
  62. assert set(generated["axis_values"]["形式"]) == _expected_l3_l4_names("形式")
  63. assert {"原声录音", "动物声音", "场景原声", "制作音效", "环境音", "符号", "文字符号"} <= set(
  64. generated["axis_values"]["实质"]
  65. )
  66. assert {"配乐", "语音"} <= set(generated["axis_values"]["形式"])
  67. def test_build_creation_query_batch_exposes_substance_and_form_axis_trees():
  68. generated = build_creation_query_batch(
  69. _settings(),
  70. options=QueryBuildOptions(per=2, batch_n=999),
  71. )
  72. for axis in ("实质", "形式"):
  73. tree = generated["axis_trees"][axis]
  74. assert tree
  75. assert all(node["level"] == 3 for node in tree)
  76. assert all(child["level"] == 4 for node in tree for child in node["children"])
  77. tree_names = {node["name"] for node in tree} | {
  78. child["name"] for node in tree for child in node["children"]
  79. }
  80. assert tree_names <= set(generated["axis_values"][axis])
  81. def test_build_creation_query_batch_expands_active_families_as_cartesian_products():
  82. generated = build_creation_query_batch(
  83. _settings(),
  84. options=QueryBuildOptions(per=0, batch_n=0),
  85. )
  86. by_key = {family["key"]: family for family in generated["families"]}
  87. f1_items = by_key["f1"]["items"]
  88. f2_items = by_key["f2"]["items"]
  89. assert len(f1_items) == len(generated["axis_values"]["实质"]) * 2 * 3 * 3
  90. assert len(f2_items) == len(generated["axis_values"]["形式"]) * 2 * 3 * 3
  91. assert {
  92. item["parts"]["知识类型"]
  93. for item in f1_items
  94. if item["parts"]["实质"] == "公共安全"
  95. and item["parts"]["模态"] == "视频"
  96. and item["parts"]["业务阶段"] == "灵感"
  97. } == {"怎么做", "有哪些", "为什么"}
  98. def test_build_creation_query_batch_can_explicitly_enable_reserved_families():
  99. all_keys = (
  100. "f1",
  101. "f2",
  102. "f4",
  103. "f3",
  104. "f5",
  105. "a_shi",
  106. "a_xing",
  107. "a_both",
  108. "a_purpose",
  109. "a_tail",
  110. "b_shi",
  111. "b_xing",
  112. "b_both",
  113. "b_purpose",
  114. "b_tail",
  115. )
  116. generated = build_creation_query_batch(
  117. _settings(),
  118. options=QueryBuildOptions(
  119. per=1,
  120. batch_n=4,
  121. active_family_keys=all_keys,
  122. ),
  123. )
  124. assert [family["key"] for family in generated["families"]] == list(all_keys)
  125. assert generated["metadata"]["active_family_keys"] == list(all_keys)
  126. def test_build_creation_query_batch_rejects_unknown_family_key():
  127. with pytest.raises(ValueError, match="unknown query family"):
  128. build_creation_query_batch(
  129. _settings(),
  130. options=QueryBuildOptions(
  131. active_family_keys=("f1", "no_such_family"),
  132. ),
  133. )
  134. def test_build_creation_query_batch_keeps_all_queries():
  135. generated = build_creation_query_batch(
  136. _settings(),
  137. options=QueryBuildOptions(per=3, batch_n=4),
  138. )
  139. assert "query_filter_enabled" not in generated["metadata"]
  140. assert "query_filter_prompt_version" not in generated["metadata"]
  141. assert all(item["keep"] is True for family in generated["families"] for item in family["items"])
  142. assert all(item["reason"] == "" for family in generated["families"] for item in family["items"])
  143. def test_persist_query_batch_writes_formal_batch_and_query_contract():
  144. repo = FakeRepo()
  145. generated = {
  146. "metadata": {
  147. "active_family_keys": ["f1", "f2"],
  148. },
  149. "families": [
  150. {
  151. "key": "f1",
  152. "name": "实质 × 模态 × 业务阶段 × 知识类型",
  153. "axes": ["实质", "模态"],
  154. "items": [
  155. {
  156. "query": "反转 视频 脚本 怎么做",
  157. "parts": {"实质": "反转", "模态": "视频"},
  158. "keep": True,
  159. "reason": "可搜创作方法",
  160. }
  161. ],
  162. }
  163. ],
  164. }
  165. batch, count = persist_query_batch(repo, generated, name="formal-demo")
  166. assert batch.name == "formal-demo"
  167. assert count == 1
  168. assert repo.batch_kwargs["status"] == "ready"
  169. assert repo.batch_kwargs["target_platforms"] == ["xiaohongshu", "weixin", "douyin"]
  170. assert repo.batch_kwargs["metadata"]["active_family_keys"] == ["f1", "f2"]
  171. row = repo.queries[0]
  172. assert row["batch_id"] == batch.id
  173. assert row["query_text"] == "反转 视频 脚本 怎么做"
  174. assert row["axes"] == {"实质": "反转", "模态": "视频"}
  175. assert row["keep"] is True
  176. assert row["filter_reason"] == "可搜创作方法"
  177. assert row["metadata"]["family_key"] == "f1"
  178. assert "valid" not in row["metadata"]
  179. assert "relevant" not in row["metadata"]
  180. assert "query_filter_prompt_version" not in row["metadata"]