builder.py 11 KB

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