"""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.queries.filter import filter_queries, prompt_version 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 enable_query_filter: bool = False 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, "query_filter_enabled": opts.enable_query_filter, "active_family_keys": list(opts.active_family_keys), "query_filter_prompt_version": prompt_version(), }, "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 verdicts = ( filter_queries([it["query"] for it in items], settings) if opts.enable_query_filter else [{"keep": True, "valid": None, "relevant": True, "reason": ""} for _ in items] ) for item, verdict in zip(items, verdicts): item.update(verdict) 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 [], "valid": item.get("valid"), "relevant": item.get("relevant"), "query_filter_prompt_version": (generated.get("metadata") or {}).get( "query_filter_prompt_version" ), }, ) count += 1 return batch, count