builder.py 13 KB

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