test_query_builder.py 2.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960
  1. from __future__ import annotations
  2. from uuid import uuid4
  3. from acquisition.domain import Query, QueryBatch
  4. from acquisition.queries.builder import persist_query_batch
  5. class FakeRepo:
  6. def __init__(self) -> None:
  7. self.batch_kwargs = None
  8. self.queries = []
  9. def create_query_batch(self, **kwargs):
  10. self.batch_kwargs = kwargs
  11. return QueryBatch(id=uuid4(), **kwargs)
  12. def add_query(self, **kwargs):
  13. self.queries.append(kwargs)
  14. return Query(id=uuid4(), **kwargs)
  15. def test_persist_query_batch_writes_formal_batch_and_query_contract():
  16. repo = FakeRepo()
  17. generated = {
  18. "metadata": {"query_filter_prompt_version": "abc123"},
  19. "families": [
  20. {
  21. "key": "f1",
  22. "name": "实质 × 模态 × 业务阶段 × 知识类型",
  23. "axes": ["实质", "模态"],
  24. "items": [
  25. {
  26. "query": "反转 视频 脚本 怎么做",
  27. "parts": {"实质": "反转", "模态": "视频"},
  28. "keep": True,
  29. "valid": 8,
  30. "relevant": True,
  31. "reason": "可搜创作方法",
  32. }
  33. ],
  34. }
  35. ],
  36. }
  37. batch, count = persist_query_batch(repo, generated, name="formal-demo")
  38. assert batch.name == "formal-demo"
  39. assert count == 1
  40. assert repo.batch_kwargs["status"] == "ready"
  41. assert repo.batch_kwargs["target_platforms"] == ["xiaohongshu", "weixin", "douyin"]
  42. row = repo.queries[0]
  43. assert row["batch_id"] == batch.id
  44. assert row["query_text"] == "反转 视频 脚本 怎么做"
  45. assert row["axes"] == {"实质": "反转", "模态": "视频"}
  46. assert row["keep"] is True
  47. assert row["filter_reason"] == "可搜创作方法"
  48. assert row["metadata"]["family_key"] == "f1"
  49. assert row["metadata"]["valid"] == 8
  50. assert row["metadata"]["query_filter_prompt_version"] == "abc123"