| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262 |
- from __future__ import annotations
- import json
- from uuid import uuid4
- import pytest
- from acquisition.domain import Query, QueryBatch
- from acquisition.queries import builder as query_builder
- from acquisition.queries import filter as query_filter
- from acquisition.queries.builder import (
- QueryBuildOptions,
- TREES,
- build_creation_query_batch,
- persist_query_batch,
- )
- from core.config import PgConfig, Settings
- from core.text_limits import QUERY_FILTER_REASON_MAX_CHARS
- def _settings() -> Settings:
- return Settings(
- pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
- aiddit_crawler_base_url="http://crawler.test",
- crawler_timeout=30,
- openrouter_timeout_seconds=90,
- openrouter_model="m",
- openrouter_base_url="http://openrouter.test",
- openrouter_api_key="k",
- llm_model="m",
- max_cards=12,
- frames_dir="f",
- douyin_ratio="540p",
- data_dir="",
- )
- def test_query_filter_preserves_wide_reason(monkeypatch):
- reason = "理" * (QUERY_FILTER_REASON_MAX_CHARS + 3)
- monkeypatch.setattr(
- query_filter,
- "_chat_content",
- lambda settings, messages: json.dumps(
- [{"idx": 0, "valid": 8, "relevant": True, "reason": reason}],
- ensure_ascii=False,
- ),
- )
- rows = query_filter.filter_queries(["符号 视频 灵感 怎么做"], _settings())
- assert rows[0]["reason"] == "理" * QUERY_FILTER_REASON_MAX_CHARS
- class FakeRepo:
- def __init__(self) -> None:
- self.batch_kwargs = None
- self.queries = []
- def create_query_batch(self, **kwargs):
- self.batch_kwargs = kwargs
- return QueryBatch(id=uuid4(), **kwargs)
- def add_query(self, **kwargs):
- self.queries.append(kwargs)
- return Query(id=uuid4(), **kwargs)
- def _expected_l3_l4_names(source_type: str) -> set[str]:
- rows = json.loads(TREES.read_text("utf-8"))
- names: set[str] = set()
- for row in rows:
- if row.get("source_type") != source_type:
- continue
- depth = len([part for part in (row.get("path") or "").split("/") if part])
- if depth in (3, 4) and row.get("name"):
- names.add(row["name"])
- return names
- def test_build_creation_query_batch_defaults_to_first_two_families():
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(per=2, batch_n=4),
- )
- assert [family["key"] for family in generated["families"]] == ["f1", "f2"]
- assert generated["metadata"]["active_family_keys"] == ["f1", "f2"]
- assert sum(len(family["items"]) for family in generated["families"]) == 4
- def test_build_creation_query_batch_uses_all_l3_l4_substance_and_form_nodes():
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(per=2, batch_n=999),
- )
- assert set(generated["axis_values"]["实质"]) == _expected_l3_l4_names("实质")
- assert set(generated["axis_values"]["形式"]) == _expected_l3_l4_names("形式")
- assert {"原声录音", "动物声音", "场景原声", "制作音效", "环境音", "符号", "文字符号"} <= set(
- generated["axis_values"]["实质"]
- )
- assert {"配乐", "语音"} <= set(generated["axis_values"]["形式"])
- def test_build_creation_query_batch_exposes_substance_and_form_axis_trees():
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(per=2, batch_n=999),
- )
- for axis in ("实质", "形式"):
- tree = generated["axis_trees"][axis]
- assert tree
- assert all(node["level"] == 3 for node in tree)
- assert all(child["level"] == 4 for node in tree for child in node["children"])
- tree_names = {node["name"] for node in tree} | {
- child["name"] for node in tree for child in node["children"]
- }
- assert tree_names <= set(generated["axis_values"][axis])
- def test_build_creation_query_batch_expands_active_families_as_cartesian_products():
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(per=0, batch_n=0),
- )
- by_key = {family["key"]: family for family in generated["families"]}
- f1_items = by_key["f1"]["items"]
- f2_items = by_key["f2"]["items"]
- assert len(f1_items) == len(generated["axis_values"]["实质"]) * 2 * 3 * 3
- assert len(f2_items) == len(generated["axis_values"]["形式"]) * 2 * 3 * 3
- assert {
- item["parts"]["知识类型"]
- for item in f1_items
- if item["parts"]["实质"] == "公共安全"
- and item["parts"]["模态"] == "视频"
- and item["parts"]["业务阶段"] == "灵感"
- } == {"怎么做", "有哪些", "为什么"}
- def test_build_creation_query_batch_can_explicitly_enable_reserved_families():
- all_keys = (
- "f1",
- "f2",
- "f4",
- "f3",
- "f5",
- "a_shi",
- "a_xing",
- "a_both",
- "a_purpose",
- "a_tail",
- "b_shi",
- "b_xing",
- "b_both",
- "b_purpose",
- "b_tail",
- )
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(
- per=1,
- batch_n=4,
- active_family_keys=all_keys,
- ),
- )
- assert [family["key"] for family in generated["families"]] == list(all_keys)
- assert generated["metadata"]["active_family_keys"] == list(all_keys)
- def test_build_creation_query_batch_rejects_unknown_family_key():
- with pytest.raises(ValueError, match="unknown query family"):
- build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(
- active_family_keys=("f1", "no_such_family"),
- ),
- )
- def test_build_creation_query_batch_keeps_all_queries_without_default_llm_filter(monkeypatch):
- def fail_filter(*args, **kwargs):
- raise AssertionError("query filter should be disabled by default")
- monkeypatch.setattr(query_builder, "filter_queries", fail_filter)
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(per=3, batch_n=4),
- )
- assert generated["metadata"]["query_filter_enabled"] is False
- assert all(item["keep"] is True for family in generated["families"] for item in family["items"])
- assert all(item["reason"] == "" for family in generated["families"] for item in family["items"])
- def test_build_creation_query_batch_can_enable_llm_filter(monkeypatch):
- def fake_filter(queries, settings):
- return [
- {"keep": i != 1, "valid": 9 if i != 1 else 5, "relevant": i != 1, "reason": f"r{i}"}
- for i, _ in enumerate(queries)
- ]
- monkeypatch.setattr(query_builder, "filter_queries", fake_filter)
- generated = build_creation_query_batch(
- _settings(),
- options=QueryBuildOptions(per=3, batch_n=4, enable_query_filter=True),
- )
- items = generated["families"][0]["items"]
- assert generated["metadata"]["query_filter_enabled"] is True
- assert [item["keep"] for item in items] == [True, False, True]
- assert [item["reason"] for item in items] == ["r0", "r1", "r2"]
- def test_persist_query_batch_writes_formal_batch_and_query_contract():
- repo = FakeRepo()
- generated = {
- "metadata": {
- "query_filter_prompt_version": "abc123",
- "active_family_keys": ["f1", "f2"],
- },
- "families": [
- {
- "key": "f1",
- "name": "实质 × 模态 × 业务阶段 × 知识类型",
- "axes": ["实质", "模态"],
- "items": [
- {
- "query": "反转 视频 脚本 怎么做",
- "parts": {"实质": "反转", "模态": "视频"},
- "keep": True,
- "valid": 8,
- "relevant": True,
- "reason": "可搜创作方法",
- }
- ],
- }
- ],
- }
- batch, count = persist_query_batch(repo, generated, name="formal-demo")
- assert batch.name == "formal-demo"
- assert count == 1
- assert repo.batch_kwargs["status"] == "ready"
- assert repo.batch_kwargs["target_platforms"] == ["xiaohongshu", "weixin", "douyin"]
- assert repo.batch_kwargs["metadata"]["active_family_keys"] == ["f1", "f2"]
- row = repo.queries[0]
- assert row["batch_id"] == batch.id
- assert row["query_text"] == "反转 视频 脚本 怎么做"
- assert row["axes"] == {"实质": "反转", "模态": "视频"}
- assert row["keep"] is True
- assert row["filter_reason"] == "可搜创作方法"
- assert row["metadata"]["family_key"] == "f1"
- assert row["metadata"]["valid"] == 8
- assert row["metadata"]["query_filter_prompt_version"] == "abc123"
|