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"