|
|
@@ -1,269 +1,23 @@
|
|
|
-"""Formal creation-query batch builder."""
|
|
|
+"""Backward-compatible façade for the query planning bounded context."""
|
|
|
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
|
|
|
+from query_planning import (
|
|
|
+ GenerationRequest,
|
|
|
+ GeneratorKind,
|
|
|
+ QueryBatchWriteSpec,
|
|
|
+ QueryBatchWriter,
|
|
|
+ UnifiedQueryGenerationService,
|
|
|
+ planning_store_for_repository,
|
|
|
+)
|
|
|
+from query_planning.cartesian import (
|
|
|
+ DEFAULT_ACTIVE_FAMILY_KEYS,
|
|
|
+ TREES,
|
|
|
+ QueryBuildOptions,
|
|
|
+ build_creation_query_batch,
|
|
|
+)
|
|
|
|
|
|
|
|
|
def persist_query_batch(
|
|
|
@@ -275,33 +29,73 @@ def persist_query_batch(
|
|
|
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
|
|
|
+ """Persist legacy family JSON through the unified planning and writer flow."""
|
|
|
+ candidates: list[dict[str, Any]] = []
|
|
|
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 [],
|
|
|
- },
|
|
|
+ if not item.get("keep", True):
|
|
|
+ continue
|
|
|
+ candidates.append(
|
|
|
+ {
|
|
|
+ "query_text": item.get("query") or "",
|
|
|
+ "axes": item.get("parts") or {},
|
|
|
+ "filter_reason": item.get("reason") or "",
|
|
|
+ "priority": item.get("priority") or 0,
|
|
|
+ "source_refs": [
|
|
|
+ {
|
|
|
+ "generator_kind": GeneratorKind.CARTESIAN.value,
|
|
|
+ "family_key": family.get("key"),
|
|
|
+ "family_name": family.get("name"),
|
|
|
+ "axes": item.get("parts") or {},
|
|
|
+ }
|
|
|
+ ],
|
|
|
+ "metadata": {
|
|
|
+ "family_key": family.get("key"),
|
|
|
+ "family_name": family.get("name"),
|
|
|
+ "family_axes": family.get("axes") or [],
|
|
|
+ },
|
|
|
+ }
|
|
|
)
|
|
|
- count += 1
|
|
|
- return batch, count
|
|
|
+ platforms = tuple(target_platforms or ["xiaohongshu", "weixin", "douyin"])
|
|
|
+ result = UnifiedQueryGenerationService().generate(
|
|
|
+ GenerationRequest(
|
|
|
+ generator_kind=GeneratorKind.CARTESIAN,
|
|
|
+ name=name,
|
|
|
+ target_platforms=platforms,
|
|
|
+ payload={
|
|
|
+ "candidates": candidates,
|
|
|
+ "generator_config": generated.get("metadata") or {},
|
|
|
+ "input_snapshot": {
|
|
|
+ "active_family_keys": (generated.get("metadata") or {}).get(
|
|
|
+ "active_family_keys", []
|
|
|
+ ),
|
|
|
+ "family_count": len(generated.get("families") or []),
|
|
|
+ "candidate_count": len(candidates),
|
|
|
+ },
|
|
|
+ },
|
|
|
+ metadata={"generation_method": generation_method},
|
|
|
+ )
|
|
|
+ )
|
|
|
+ write_result = QueryBatchWriter(
|
|
|
+ legacy_sink=repo,
|
|
|
+ planning_store=planning_store_for_repository(repo),
|
|
|
+ ).write(
|
|
|
+ result,
|
|
|
+ QueryBatchWriteSpec(
|
|
|
+ name=name,
|
|
|
+ source_type=source_type,
|
|
|
+ generation_method=generation_method,
|
|
|
+ target_platforms=platforms,
|
|
|
+ metadata=generated.get("metadata") or {},
|
|
|
+ ),
|
|
|
+ )
|
|
|
+ return write_result.batch, len(write_result.queries)
|
|
|
+
|
|
|
+
|
|
|
+__all__ = [
|
|
|
+ "DEFAULT_ACTIVE_FAMILY_KEYS",
|
|
|
+ "QueryBuildOptions",
|
|
|
+ "TREES",
|
|
|
+ "build_creation_query_batch",
|
|
|
+ "persist_query_batch",
|
|
|
+]
|