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="agent_plan"): service.generate( _request(kind=GeneratorKind.AGENT_PLAN, 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