| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307 |
- """Formal creation-query batch builder."""
- from __future__ import annotations
- import json
- import random
- from dataclasses import dataclass
- from itertools import product
- from pathlib import Path
- from typing import Any
- from acquisition.queries.axes import ACTIONS, STAGES, _nonleaf_d4
- from acquisition.repositories.base import AcquisitionRepository
- from core.config import Settings
- ROOT = Path(__file__).resolve().parents[2]
- TREES = ROOT / "scope_trees" / "trees_index.json"
- KTYPE = ["怎么做", "有哪些", "为什么"]
- MODALITY = ["视频", "图文"]
- INTENT = ["灵感", "选题", "脚本"]
- DEFAULT_ACTIVE_FAMILY_KEYS = ("f1", "f2")
- @dataclass(frozen=True)
- class QueryBuildOptions:
- per: int = 0
- batch_n: int = 0
- seed: int = 7
- active_family_keys: tuple[str, ...] = DEFAULT_ACTIVE_FAMILY_KEYS
- def _segs(path: str | None) -> list[str]:
- return [x for x in (path or "").split("/") if x]
- def _leaves(idx: list[dict[str, Any]], source_type: str, under: str | None = None) -> list[str]:
- paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
- if under:
- paths = [(s, nm) for s, nm in paths if under in s]
- all_paths = {"/".join(s) for s, _ in paths}
- out: list[str] = []
- seen: set[str] = set()
- for segs, name in paths:
- if len(segs) < 2:
- continue
- full = "/".join(segs)
- is_leaf = not any(other != full and other.startswith(full + "/") for other in all_paths)
- value = name or segs[-1]
- if is_leaf and value and value not in seen:
- seen.add(value)
- out.append(value)
- return out
- def _nodes_at_depth(
- idx: list[dict[str, Any]],
- source_type: str,
- depths: tuple[int, ...] = (3, 4),
- under: str | None = None,
- ) -> list[str]:
- paths = [(_segs(n.get("path")), n.get("name")) for n in idx if n.get("source_type") == source_type]
- if under:
- paths = [(s, nm) for s, nm in paths if under in s]
- out: list[str] = []
- seen: set[str] = set()
- for segs, name in paths:
- if len(segs) not in depths:
- continue
- value = name or (segs[-1] if segs else "")
- if value and value not in seen:
- seen.add(value)
- out.append(value)
- return out
- def _axis_tree(idx: list[dict[str, Any]], source_type: str) -> list[dict[str, Any]]:
- nodes: dict[str, dict[str, Any]] = {}
- for row in idx:
- if row.get("source_type") != source_type:
- continue
- segs = _segs(row.get("path"))
- if len(segs) not in (3, 4):
- continue
- path = "/" + "/".join(segs)
- nodes[path] = {
- "name": row.get("name") or segs[-1],
- "path": path,
- "level": len(segs),
- "children": [],
- }
- roots: list[dict[str, Any]] = []
- for path, node in nodes.items():
- if node["level"] == 3:
- roots.append(node)
- continue
- parent_path = "/" + "/".join(_segs(path)[:-1])
- parent = nodes.get(parent_path)
- if parent:
- parent["children"].append(node)
- roots.sort(key=lambda node: node["path"])
- for node in roots:
- node["children"].sort(key=lambda child: child["path"])
- return roots
- def _sample_axis(values: list[str], limit: int, rng: random.Random) -> list[str]:
- if limit <= 0 or limit >= len(values):
- return values
- return rng.sample(values, limit)
- def build_creation_query_batch(
- settings: Settings,
- *,
- tree_path: Path = TREES,
- options: QueryBuildOptions | None = None,
- ) -> dict[str, Any]:
- """Build a creation-query batch without writing files or database rows."""
- opts = options or QueryBuildOptions()
- rng = random.Random(opts.seed)
- idx = json.loads(tree_path.read_text("utf-8"))
- shi = _nodes_at_depth(idx, "实质", depths=(3, 4))
- xing = _nodes_at_depth(idx, "形式", depths=(3, 4))
- purpose_pool = _leaves(idx, "作用") + _leaves(idx, "感受") + _leaves(idx, "意图")
- shi_batch = _sample_axis(shi, opts.batch_n, rng)
- xing_batch = _sample_axis(xing, opts.batch_n, rng)
- purpose_batch = _sample_axis(purpose_pool, opts.batch_n, rng)
- f6_zy = _nonleaf_d4("作用", 10, tree_path=tree_path)
- f6_stage_act = [(s, a) for s in STAGES for a in ACTIONS] + [("", "")]
- if not shi_batch or not xing_batch or not purpose_batch or not f6_zy:
- raise RuntimeError("scope tree does not contain enough creation query axes")
- product_pools = {
- "实质": shi_batch,
- "形式": xing_batch,
- "目的": purpose_batch,
- "模态": MODALITY,
- "业务阶段": INTENT,
- "知识类型": KTYPE,
- }
- def gen_product(keys: list[str]):
- for values in product(*(product_pools[key] for key in keys)):
- yield {"parts": dict(zip(keys, values, strict=True))}
- legacy_per = opts.per if opts.per > 0 else 30
- master: list[dict[str, str]] = []
- for i in range(legacy_per):
- stage, action = f6_stage_act[i % len(f6_stage_act)]
- master.append(
- {
- "实质": shi_batch[i % len(shi_batch)],
- "形式": xing_batch[i % len(xing_batch)],
- "目的": purpose_batch[i % len(purpose_batch)],
- "模态": MODALITY[i % len(MODALITY)],
- "业务阶段": INTENT[i % len(INTENT)],
- "知识类型": KTYPE[(i // 3) % len(KTYPE)],
- "阶段": stage or "/",
- "动作": action or "/",
- "_st": stage,
- "_ac": action,
- "作用": f6_zy[i % len(f6_zy)],
- }
- )
- def gen_old(
- i: int,
- *,
- shi_axis: bool = False,
- xing_axis: bool = False,
- purpose_axis: bool = False,
- effect_axis: bool = True,
- ) -> dict[str, Any]:
- row = master[i]
- segment = (row["_st"] + row["_ac"]) if row["_ac"] else ""
- head = (
- ([row["实质"]] if shi_axis else [])
- + ([row["形式"]] if xing_axis else [])
- + ([row["目的"]] if purpose_axis else [])
- )
- tail = ([row["作用"]] if effect_axis else []) + [row["知识类型"]]
- query = " ".join(head + ([segment] if segment else []) + tail)
- parts: dict[str, str] = {}
- if shi_axis:
- parts["实质"] = row["实质"]
- if xing_axis:
- parts["形式"] = row["形式"]
- if purpose_axis:
- parts["目的"] = row["目的"]
- parts["阶段"], parts["动作"] = row["阶段"], row["动作"]
- if effect_axis:
- parts["作用"] = row["作用"]
- parts["知识类型"] = row["知识类型"]
- return {"parts": parts, "query": query}
- families = [
- {"key": "f1", "axes": ["实质", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["实质", "模态", "业务阶段", "知识类型"])},
- {"key": "f2", "axes": ["形式", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["形式", "模态", "业务阶段", "知识类型"])},
- {"key": "f4", "axes": ["作用/感受/意图", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["目的", "模态", "业务阶段", "知识类型"])},
- {"key": "f3", "axes": ["实质", "形式", "模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["实质", "形式", "模态", "业务阶段", "知识类型"])},
- {"key": "f5", "axes": ["模态", "业务阶段", "知识类型"], "items": lambda: gen_product(["模态", "业务阶段", "知识类型"])},
- {"key": "a_shi", "axes": ["实质", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True)},
- {"key": "a_xing", "axes": ["形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True)},
- {"key": "a_both", "axes": ["实质", "形式", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True)},
- {"key": "a_purpose", "axes": ["作用/感受/意图", "阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i, purpose_axis=True)},
- {"key": "a_tail", "axes": ["阶段", "动作", "作用", "知识类型"], "gen": lambda i: gen_old(i)},
- {"key": "b_shi", "axes": ["实质", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, effect_axis=False)},
- {"key": "b_xing", "axes": ["形式", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, xing_axis=True, effect_axis=False)},
- {"key": "b_both", "axes": ["实质", "形式", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, shi_axis=True, xing_axis=True, effect_axis=False)},
- {"key": "b_purpose", "axes": ["作用/感受/意图", "阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, purpose_axis=True, effect_axis=False)},
- {"key": "b_tail", "axes": ["阶段", "动作", "知识类型"], "gen": lambda i: gen_old(i, effect_axis=False)},
- ]
- order = ["实质", "形式", "目的", "模态", "业务阶段", "知识类型"]
- out: dict[str, Any] = {
- "axis_values": {
- "实质": shi,
- "形式": xing,
- "目的池": purpose_pool,
- "业务阶段": INTENT,
- "模态": MODALITY,
- "知识类型": KTYPE,
- "阶段": STAGES,
- "动作": ACTIONS,
- "作用": f6_zy,
- },
- "axis_trees": {
- "实质": _axis_tree(idx, "实质"),
- "形式": _axis_tree(idx, "形式"),
- },
- "metadata": {
- "seed": opts.seed,
- "per": opts.per,
- "batch_n": opts.batch_n,
- "active_family_keys": list(opts.active_family_keys),
- },
- "families": [],
- }
- family_by_key = {family["key"]: family for family in families}
- unknown = [key for key in opts.active_family_keys if key not in family_by_key]
- if unknown:
- raise ValueError(f"unknown query family key(s): {', '.join(unknown)}")
- for family_key in opts.active_family_keys:
- family = family_by_key[family_key]
- name = " × ".join(family["axes"])
- seen: set[str] = set()
- items: list[dict[str, Any]] = []
- generated_items = family["items"]() if "items" in family else (family["gen"](i) for i in range(legacy_per))
- for generated in generated_items:
- parts = generated["parts"]
- query = generated.get("query") or " ".join(parts[k] for k in order if k in parts)
- if query in seen:
- continue
- seen.add(query)
- items.append({"query": query, "parts": parts})
- if opts.per > 0 and len(items) >= opts.per:
- break
- for item in items:
- item.update({"keep": True, "reason": ""})
- out["families"].append(
- {"key": family["key"], "name": name, "axes": family["axes"], "items": items}
- )
- return out
- def persist_query_batch(
- repo: AcquisitionRepository,
- generated: dict[str, Any],
- *,
- name: str,
- source_type: str = "generated",
- generation_method: str = "creation_demo_v1",
- target_platforms: list[str] | None = None,
- ) -> tuple[Any, int]:
- """Persist generated families into formal query batch/query rows."""
- batch = repo.create_query_batch(
- name=name,
- source_type=source_type,
- generation_method=generation_method,
- target_platforms=target_platforms or ["xiaohongshu", "weixin", "douyin"],
- status="ready",
- metadata=generated.get("metadata") or {},
- )
- count = 0
- sort_order = 0
- for family in generated.get("families") or []:
- for item in family.get("items") or []:
- sort_order += 1
- repo.add_query(
- batch_id=batch.id,
- query_text=item["query"],
- axes=item.get("parts") or {},
- keep=bool(item.get("keep", True)),
- filter_reason=item.get("reason") or "",
- status="ready",
- sort_order=sort_order,
- metadata={
- "family_key": family.get("key"),
- "family_name": family.get("name"),
- "family_axes": family.get("axes") or [],
- },
- )
- count += 1
- return batch, count
|