| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376 |
- from __future__ import annotations
- from pathlib import Path
- from uuid import uuid4
- import pytest
- from acquisition.domain import Query, QueryBatch
- from acquisition.queries.builder import persist_query_batch
- from core import db_session
- from query_planning import (
- GenerationRequest,
- GeneratorKind,
- NoQueryCandidatesError,
- QueryBatchWriteSpec,
- QueryBatchWriter,
- UnifiedQueryGenerationService,
- UnsupportedGeneratorError,
- )
- def _request(kind=GeneratorKind.MANUAL, *, candidates, max_queries=None):
- return GenerationRequest(
- generator_kind=kind,
- name="test-plan",
- payload={"candidates": candidates, "input_snapshot": {"case": "test"}},
- max_queries=max_queries,
- )
- def test_normalizes_and_exactly_dedupes_while_preserving_near_queries():
- result = UnifiedQueryGenerationService().generate(
- _request(
- candidates=[
- {
- "query_text": " A B ",
- "axes": {"first": True},
- "metadata": {"family_key": "first"},
- "source_refs": [{"family_key": "first"}],
- },
- {
- "query_text": "a b",
- "priority": 8,
- "axes": {"second": True},
- "metadata": {"family_key": "second"},
- "source_refs": [{"family_key": "second"}],
- },
- {
- "query_text": "a b 方法",
- "source_refs": [{"family_key": "near"}],
- },
- ]
- )
- )
- assert [query.query_text for query in result.selected_queries] == ["A B", "a b 方法"]
- merged = result.selected_queries[0]
- assert merged.axes == {"first": True}
- assert merged.metadata["family_key"] == "first"
- assert merged.priority == 8
- assert merged.metadata["origins"] == [
- {"family_key": "first"},
- {"family_key": "second"},
- ]
- assert result.stats.generated_count == 3
- assert result.stats.unique_count == 2
- assert result.stats.dropped_count == 1
- def test_priority_budget_is_stable_and_default_is_unlimited():
- candidates = [
- {"query_text": "q0", "priority": 0},
- {"query_text": "q1", "priority": 5},
- {"query_text": "q2", "priority": 5},
- {"query_text": "q3", "priority": 1},
- ]
- service = UnifiedQueryGenerationService()
- unlimited = service.generate(_request(candidates=candidates))
- budgeted = service.generate(_request(candidates=candidates, max_queries=2))
- assert [query.query_text for query in unlimited.selected_queries] == ["q1", "q2", "q3", "q0"]
- assert [query.query_text for query in budgeted.selected_queries] == ["q1", "q2"]
- assert budgeted.stats.selected_count == 2
- assert budgeted.stats.dropped_count == 2
- def test_empty_selection_and_unregistered_reserved_generator_fail_explicitly():
- service = UnifiedQueryGenerationService()
- with pytest.raises(NoQueryCandidatesError):
- service.generate(_request(candidates=[" "]))
- with pytest.raises(NoQueryCandidatesError):
- service.generate(_request(candidates=["valid"], max_queries=0))
- with pytest.raises(UnsupportedGeneratorError, match="topic_table"):
- service.generate(
- _request(kind=GeneratorKind.TOPIC_TABLE, candidates=["not allowed yet"])
- )
- class FakeLegacySink:
- def __init__(self):
- self.batch_kwargs = None
- self.query_kwargs = []
- def create_query_batch(self, **kwargs):
- self.batch_kwargs = kwargs
- return QueryBatch(id=uuid4(), **kwargs)
- def add_query(self, **kwargs):
- self.query_kwargs.append(kwargs)
- return Query(id=uuid4(), **kwargs)
- class FakePlanningStore:
- def __init__(self):
- self.plan_id = uuid4()
- self.plans = []
- self.needs = []
- self.plan_batch_links = []
- self.query_need_links = []
- def create_plan(self, **kwargs):
- self.plans.append(kwargs)
- return self.plan_id
- def add_knowledge_need(self, *, plan_id, need):
- need_id = uuid4()
- self.needs.append((plan_id, need, need_id))
- return need_id
- def link_plan_batch(self, **kwargs):
- self.plan_batch_links.append(kwargs)
- def link_query_need(self, **kwargs):
- self.query_need_links.append(kwargs)
- def test_writer_keeps_legacy_search_contract_and_links_plan_need_rows():
- service = UnifiedQueryGenerationService()
- result = service.generate(
- GenerationRequest(
- generator_kind=GeneratorKind.MANUAL,
- name="manual",
- payload={
- "knowledge_needs": [
- {
- "need_key": "need-1",
- "decision_context": "选择画面主线",
- "unknown_information": "透明伞怎样形成柔光",
- "source_ref": {"topic_id": 428},
- }
- ],
- "candidates": [
- {
- "query_text": "透明伞 人像 柔光",
- "axes": {"道具": "透明伞"},
- "metadata": {"family_key": "manual"},
- "source_refs": [{"generator_kind": "manual"}],
- "knowledge_need_keys": ["need-1"],
- }
- ],
- },
- )
- )
- sink = FakeLegacySink()
- planning = FakePlanningStore()
- written = QueryBatchWriter(legacy_sink=sink, planning_store=planning).write(
- result,
- QueryBatchWriteSpec(
- name="manual",
- source_type="manual",
- generation_method="manual_query_api_v1",
- target_platforms=("xiaohongshu",),
- metadata={"source": "test"},
- ),
- )
- assert written.plan_id == planning.plan_id
- assert sink.batch_kwargs["status"] == "ready"
- assert sink.batch_kwargs["target_platforms"] == ["xiaohongshu"]
- assert sink.batch_kwargs["metadata"]["query_planning"]["plan_id"] == str(planning.plan_id)
- query = sink.query_kwargs[0]
- assert query["query_text"] == "透明伞 人像 柔光"
- assert query["keep"] is True
- assert query["status"] == "ready"
- assert query["sort_order"] == 0
- assert query["filter_reason"] is None
- assert query["metadata"]["family_key"] == "manual"
- assert planning.plan_batch_links == [
- {"plan_id": planning.plan_id, "batch_id": written.batch.id}
- ]
- assert len(planning.query_need_links) == 1
- def test_writer_rejects_unknown_need_before_any_write():
- result = UnifiedQueryGenerationService().generate(
- _request(
- candidates=[
- {
- "query_text": "query",
- "knowledge_need_keys": ["missing"],
- }
- ]
- )
- )
- sink = FakeLegacySink()
- planning = FakePlanningStore()
- with pytest.raises(ValueError, match="missing"):
- QueryBatchWriter(legacy_sink=sink, planning_store=planning).write(
- result,
- QueryBatchWriteSpec(
- name="manual",
- source_type="manual",
- generation_method="manual_query_api_v1",
- target_platforms=("xiaohongshu",),
- ),
- )
- assert planning.plans == []
- assert sink.batch_kwargs is None
- def test_outer_transaction_rolls_back_partial_writer_failure(monkeypatch):
- class FakeConnection:
- def __init__(self):
- self.pending = []
- self.committed = []
- self.rollback_called = False
- self.closed = False
- def commit(self):
- self.committed.extend(self.pending)
- self.pending.clear()
- def rollback(self):
- self.rollback_called = True
- self.pending.clear()
- def close(self):
- self.closed = True
- class TransactionalPlanningStore(FakePlanningStore):
- def __init__(self, conn):
- super().__init__()
- self.conn = conn
- def create_plan(self, **kwargs):
- self.conn.pending.append(("plan", kwargs))
- return self.plan_id
- def link_plan_batch(self, **kwargs):
- self.conn.pending.append(("plan_batch", kwargs))
- class FailingLegacySink(FakeLegacySink):
- def __init__(self, conn):
- super().__init__()
- self.conn = conn
- def create_query_batch(self, **kwargs):
- batch = super().create_query_batch(**kwargs)
- self.conn.pending.append(("batch", batch.id))
- return batch
- def add_query(self, **kwargs):
- if len(self.query_kwargs) == 1:
- raise RuntimeError("second query failed")
- query = super().add_query(**kwargs)
- self.conn.pending.append(("query", query.id))
- return query
- result = UnifiedQueryGenerationService().generate(
- _request(candidates=["first query", "second query"])
- )
- conn = FakeConnection()
- monkeypatch.setattr(db_session, "connect", lambda _config: conn)
- with pytest.raises(RuntimeError, match="second query failed"):
- with db_session.transaction(object()) as transaction_conn:
- QueryBatchWriter(
- legacy_sink=FailingLegacySink(transaction_conn),
- planning_store=TransactionalPlanningStore(transaction_conn),
- ).write(
- result,
- QueryBatchWriteSpec(
- name="rollback",
- source_type="manual",
- generation_method="manual_query_api_v1",
- target_platforms=("xiaohongshu",),
- ),
- )
- assert conn.rollback_called is True
- assert conn.pending == []
- assert conn.committed == []
- assert conn.closed is True
- def test_cartesian_legacy_facade_exactly_dedupes_across_families():
- sink = FakeLegacySink()
- generated = {
- "metadata": {"active_family_keys": ["f1", "f2"]},
- "families": [
- {
- "key": "f1",
- "name": "实质 × 模态",
- "axes": ["实质", "模态"],
- "items": [
- {"query": "A B", "parts": {"实质": "A"}, "keep": True}
- ],
- },
- {
- "key": "f2",
- "name": "形式 × 模态",
- "axes": ["形式", "模态"],
- "items": [
- {"query": "a b", "parts": {"形式": "A"}, "keep": True}
- ],
- },
- ],
- }
- batch, count = persist_query_batch(sink, generated, name="cartesian")
- assert batch.id is not None
- assert count == 1
- assert sink.query_kwargs[0]["query_text"] == "A B"
- assert sink.query_kwargs[0]["sort_order"] == 0
- assert sink.query_kwargs[0]["metadata"]["family_key"] == "f1"
- assert [origin["family_key"] for origin in sink.query_kwargs[0]["metadata"]["origins"]] == [
- "f1",
- "f2",
- ]
- def test_migration_is_additive_replayable_and_does_not_force_what_how_why():
- sql = Path("db/migrations/005_query_planning_schema.sql").read_text(encoding="utf-8")
- for table in (
- "query_plans",
- "knowledge_needs",
- "query_plan_batches",
- "query_knowledge_need_links",
- ):
- assert f"CREATE TABLE IF NOT EXISTS creation_knowledge.{table}" in sql
- assert "ALTER TABLE creation_knowledge.query_batches" not in sql
- assert "ALTER TABLE creation_knowledge.queries" not in sql
- assert "005_query_planning_schema" in sql
- assert "trg_query_plans_touch_updated_at" in sql
- assert "TO ck_app" in sql
- assert "particle_type" not in sql
- def test_frozen_search_modules_do_not_depend_on_query_planning():
- frozen_roots = [
- Path("acquisition/runner.py"),
- Path("acquisition/repositories"),
- Path("acquisition/platforms"),
- Path("pipeline"),
- Path("decode_content"),
- ]
- for root in frozen_roots:
- files = [root] if root.is_file() else list(root.rglob("*.py"))
- for path in files:
- assert "query_planning" not in path.read_text(encoding="utf-8"), path
- forbidden = (
- "acquisition.runner",
- "acquisition.platforms",
- "acquisition.search",
- "pipeline",
- "decode_content",
- )
- for path in Path("query_planning").rglob("*.py"):
- source = path.read_text(encoding="utf-8")
- assert not any(f"from {module}" in source or f"import {module}" in source for module in forbidden), path
|