builder.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263
  1. """Formal creation-query batch builder."""
  2. from __future__ import annotations
  3. import json
  4. import random
  5. from dataclasses import dataclass
  6. from pathlib import Path
  7. from typing import Any
  8. from acquisition.query import ACTIONS, STAGES as OLD_STAGES, _nonleaf_d4
  9. from acquisition.queries.filter import filter_queries, prompt_version
  10. from acquisition.repositories.base import AcquisitionRepository
  11. from core.config import Settings
  12. ROOT = Path(__file__).resolve().parents[2]
  13. TREES = ROOT / "scope_trees" / "trees_index.json"
  14. KTYPE = ["怎么做", "有哪些", "为什么"]
  15. MODALITY = ["视频", "图文"]
  16. INTENT = ["灵感", "选题", "脚本"]
  17. @dataclass(frozen=True)
  18. class QueryBuildOptions:
  19. per: int = 30
  20. batch_n: int = 30
  21. seed: int = 7
  22. dry: bool = False
  23. def _segs(path: str | None) -> list[str]:
  24. return [x for x in (path or "").split("/") if x]
  25. def _leaves(idx: list[dict[str, Any]], source_type: str, under: str | None = None) -> list[str]:
  26. paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
  27. if under:
  28. paths = [(s, nm) for s, nm in paths if under in s]
  29. all_paths = {"/".join(s) for s, _ in paths}
  30. out: list[str] = []
  31. seen: set[str] = set()
  32. for segs, name in paths:
  33. if len(segs) < 2:
  34. continue
  35. full = "/".join(segs)
  36. is_leaf = not any(other != full and other.startswith(full + "/") for other in all_paths)
  37. value = name or segs[-1]
  38. if is_leaf and value and value not in seen:
  39. seen.add(value)
  40. out.append(value)
  41. return out
  42. def _nonleaf(
  43. idx: list[dict[str, Any]],
  44. source_type: str,
  45. depths: tuple[int, ...] = (3, 4),
  46. under: str | None = None,
  47. ) -> list[str]:
  48. paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
  49. if under:
  50. paths = [(s, nm) for s, nm in paths if under in s]
  51. all_paths = {"/".join(s) for s, _ in paths}
  52. out: list[str] = []
  53. seen: set[str] = set()
  54. for segs, name in paths:
  55. if len(segs) not in depths:
  56. continue
  57. full = "/".join(segs)
  58. is_nonleaf = any(other != full and other.startswith(full + "/") for other in all_paths)
  59. value = name or (segs[-1] if segs else "")
  60. if is_nonleaf and value and value not in seen:
  61. seen.add(value)
  62. out.append(value)
  63. return out
  64. def build_creation_query_batch(
  65. settings: Settings,
  66. *,
  67. tree_path: Path = TREES,
  68. options: QueryBuildOptions | None = None,
  69. ) -> dict[str, Any]:
  70. """Build a creation-query batch without writing files or database rows."""
  71. opts = options or QueryBuildOptions()
  72. rng = random.Random(opts.seed)
  73. idx = json.loads(tree_path.read_text("utf-8"))
  74. shi = _nonleaf(idx, "实质", depths=(3, 4), under="理念")
  75. xing = _nonleaf(idx, "形式", depths=(3, 4), under="架构")
  76. purpose_pool = _leaves(idx, "作用") + _leaves(idx, "感受") + _leaves(idx, "意图")
  77. shi_batch = rng.sample(shi, min(opts.batch_n, len(shi)))
  78. xing_batch = rng.sample(xing, min(opts.batch_n, len(xing)))
  79. purpose_batch = rng.sample(purpose_pool, min(opts.batch_n, len(purpose_pool)))
  80. f6_zy = _nonleaf_d4("作用", 10)
  81. f6_stage_act = [(s, a) for s in OLD_STAGES for a in ACTIONS] + [("", "")]
  82. if not shi_batch or not xing_batch or not purpose_batch or not f6_zy:
  83. raise RuntimeError("scope tree does not contain enough creation query axes")
  84. master: list[dict[str, str]] = []
  85. for i in range(opts.per):
  86. stage, action = f6_stage_act[i % len(f6_stage_act)]
  87. master.append(
  88. {
  89. "实质": shi_batch[i % len(shi_batch)],
  90. "形式": xing_batch[i % len(xing_batch)],
  91. "目的": purpose_batch[i % len(purpose_batch)],
  92. "模态": MODALITY[i % len(MODALITY)],
  93. "业务阶段": INTENT[i % len(INTENT)],
  94. "知识类型": KTYPE[(i // 3) % len(KTYPE)],
  95. "阶段": stage or "/",
  96. "动作": action or "/",
  97. "_st": stage,
  98. "_ac": action,
  99. "作用": f6_zy[i % len(f6_zy)],
  100. }
  101. )
  102. def project(i: int, keys: list[str]) -> dict[str, Any]:
  103. row = master[i]
  104. return {"parts": {key: row[key] for key in keys}}
  105. def gen_old(
  106. i: int,
  107. *,
  108. shi_axis: bool = False,
  109. xing_axis: bool = False,
  110. purpose_axis: bool = False,
  111. effect_axis: bool = True,
  112. ) -> dict[str, Any]:
  113. row = master[i]
  114. segment = (row["_st"] + row["_ac"]) if row["_ac"] else ""
  115. head = (
  116. ([row["实质"]] if shi_axis else [])
  117. + ([row["形式"]] if xing_axis else [])
  118. + ([row["目的"]] if purpose_axis else [])
  119. )
  120. tail = ([row["作用"]] if effect_axis else []) + [row["知识类型"]]
  121. query = " ".join(head + ([segment] if segment else []) + tail)
  122. parts: dict[str, str] = {}
  123. if shi_axis:
  124. parts["实质"] = row["实质"]
  125. if xing_axis:
  126. parts["形式"] = row["形式"]
  127. if purpose_axis:
  128. parts["目的"] = row["目的"]
  129. parts["阶段"], parts["动作"] = row["阶段"], row["动作"]
  130. if effect_axis:
  131. parts["作用"] = row["作用"]
  132. parts["知识类型"] = row["知识类型"]
  133. return {"parts": parts, "query": query}
  134. families = [
  135. {"key": "f1", "axes": ["实质", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["实质", "模态", "业务阶段", "知识类型"])},
  136. {"key": "f2", "axes": ["形式", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["形式", "模态", "业务阶段", "知识类型"])},
  137. {"key": "f4", "axes": ["作用/感受/意图", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["目的", "模态", "业务阶段", "知识类型"])},
  138. {"key": "f3", "axes": ["实质", "形式", "模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["实质", "形式", "模态", "业务阶段", "知识类型"])},
  139. {"key": "f5", "axes": ["模态", "业务阶段", "知识类型"], "gen": lambda i: project(i, ["模态", "业务阶段", "知识类型"])},
  140. {"key": "a_shi", "axes": ["实质", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True)},
  141. {"key": "a_xing", "axes": ["形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True)},
  142. {"key": "a_both", "axes": ["实质", "形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True)},
  143. {"key": "a_purpose", "axes": ["作用/感受/意图", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, purpose_axis=True)},
  144. {"key": "a_tail", "axes": ["阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i)},
  145. {"key": "b_shi", "axes": ["实质", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, effect_axis=False)},
  146. {"key": "b_xing", "axes": ["形式", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True, effect_axis=False)},
  147. {"key": "b_both", "axes": ["实质", "形式", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True, effect_axis=False)},
  148. {"key": "b_purpose", "axes": ["作用/感受/意图", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, purpose_axis=True, effect_axis=False)},
  149. {"key": "b_tail", "axes": ["阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, effect_axis=False)},
  150. ]
  151. order = ["实质", "形式", "目的", "模态", "业务阶段", "知识类型"]
  152. out: dict[str, Any] = {
  153. "axis_values": {
  154. "实质": shi,
  155. "形式": xing,
  156. "目的池": purpose_pool,
  157. "业务阶段": INTENT,
  158. "模态": MODALITY,
  159. "知识类型": KTYPE,
  160. "阶段": OLD_STAGES,
  161. "动作": ACTIONS,
  162. "作用": f6_zy,
  163. },
  164. "metadata": {
  165. "seed": opts.seed,
  166. "per": opts.per,
  167. "batch_n": opts.batch_n,
  168. "dry": opts.dry,
  169. "query_filter_prompt_version": prompt_version(),
  170. },
  171. "families": [],
  172. }
  173. for family in families:
  174. name = " × ".join(family["axes"])
  175. seen: set[str] = set()
  176. items: list[dict[str, Any]] = []
  177. for i in range(opts.per):
  178. generated = family["gen"](i)
  179. parts = generated["parts"]
  180. query = generated.get("query") or " ".join(parts[k] for k in order if k in parts)
  181. if query in seen:
  182. continue
  183. seen.add(query)
  184. items.append({"query": query, "parts": parts})
  185. verdicts = (
  186. [
  187. {"keep": True, "valid": None, "relevant": True, "reason": "(dry:未过筛)"}
  188. for _ in items
  189. ]
  190. if opts.dry
  191. else filter_queries([it["query"] for it in items], settings)
  192. )
  193. for item, verdict in zip(items, verdicts):
  194. item.update(verdict)
  195. out["families"].append(
  196. {"key": family["key"], "name": name, "axes": family["axes"], "items": items}
  197. )
  198. return out
  199. def persist_query_batch(
  200. repo: AcquisitionRepository,
  201. generated: dict[str, Any],
  202. *,
  203. name: str,
  204. source_type: str = "generated",
  205. generation_method: str = "creation_demo_v1",
  206. target_platforms: list[str] | None = None,
  207. ) -> tuple[Any, int]:
  208. """Persist generated families into formal query batch/query rows."""
  209. batch = repo.create_query_batch(
  210. name=name,
  211. source_type=source_type,
  212. generation_method=generation_method,
  213. target_platforms=target_platforms or ["xiaohongshu", "weixin", "douyin"],
  214. status="ready",
  215. metadata=generated.get("metadata") or {},
  216. )
  217. count = 0
  218. sort_order = 0
  219. for family in generated.get("families") or []:
  220. for item in family.get("items") or []:
  221. sort_order += 1
  222. repo.add_query(
  223. batch_id=batch.id,
  224. query_text=item["query"],
  225. axes=item.get("parts") or {},
  226. keep=bool(item.get("keep", True)),
  227. filter_reason=item.get("reason") or "",
  228. status="ready",
  229. sort_order=sort_order,
  230. metadata={
  231. "family_key": family.get("key"),
  232. "family_name": family.get("name"),
  233. "family_axes": family.get("axes") or [],
  234. "valid": item.get("valid"),
  235. "relevant": item.get("relevant"),
  236. "query_filter_prompt_version": (generated.get("metadata") or {}).get(
  237. "query_filter_prompt_version"
  238. ),
  239. },
  240. )
  241. count += 1
  242. return batch, count