test_query_builder.py 4.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142
  1. from __future__ import annotations
  2. from uuid import uuid4
  3. import pytest
  4. from acquisition.domain import Query, QueryBatch
  5. from acquisition.queries.builder import (
  6. QueryBuildOptions,
  7. build_creation_query_batch,
  8. persist_query_batch,
  9. )
  10. from core.config import PgConfig, Settings
  11. def _settings() -> Settings:
  12. return Settings(
  13. pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
  14. aiddit_crawler_base_url="http://crawler.test",
  15. crawler_timeout=30,
  16. openrouter_timeout_seconds=90,
  17. openrouter_model="m",
  18. openrouter_base_url="http://openrouter.test",
  19. openrouter_api_key="k",
  20. llm_model="m",
  21. max_cards=12,
  22. frames_dir="f",
  23. douyin_ratio="540p",
  24. data_dir="",
  25. )
  26. class FakeRepo:
  27. def __init__(self) -> None:
  28. self.batch_kwargs = None
  29. self.queries = []
  30. def create_query_batch(self, **kwargs):
  31. self.batch_kwargs = kwargs
  32. return QueryBatch(id=uuid4(), **kwargs)
  33. def add_query(self, **kwargs):
  34. self.queries.append(kwargs)
  35. return Query(id=uuid4(), **kwargs)
  36. def test_build_creation_query_batch_defaults_to_first_two_families():
  37. generated = build_creation_query_batch(
  38. _settings(),
  39. options=QueryBuildOptions(per=2, batch_n=4, dry=True),
  40. )
  41. assert [family["key"] for family in generated["families"]] == ["f1", "f2"]
  42. assert generated["metadata"]["active_family_keys"] == ["f1", "f2"]
  43. assert sum(len(family["items"]) for family in generated["families"]) == 4
  44. def test_build_creation_query_batch_can_explicitly_enable_reserved_families():
  45. all_keys = (
  46. "f1",
  47. "f2",
  48. "f4",
  49. "f3",
  50. "f5",
  51. "a_shi",
  52. "a_xing",
  53. "a_both",
  54. "a_purpose",
  55. "a_tail",
  56. "b_shi",
  57. "b_xing",
  58. "b_both",
  59. "b_purpose",
  60. "b_tail",
  61. )
  62. generated = build_creation_query_batch(
  63. _settings(),
  64. options=QueryBuildOptions(
  65. per=1,
  66. batch_n=4,
  67. dry=True,
  68. active_family_keys=all_keys,
  69. ),
  70. )
  71. assert [family["key"] for family in generated["families"]] == list(all_keys)
  72. assert generated["metadata"]["active_family_keys"] == list(all_keys)
  73. def test_build_creation_query_batch_rejects_unknown_family_key():
  74. with pytest.raises(ValueError, match="unknown query family"):
  75. build_creation_query_batch(
  76. _settings(),
  77. options=QueryBuildOptions(
  78. dry=True,
  79. active_family_keys=("f1", "no_such_family"),
  80. ),
  81. )
  82. def test_persist_query_batch_writes_formal_batch_and_query_contract():
  83. repo = FakeRepo()
  84. generated = {
  85. "metadata": {
  86. "query_filter_prompt_version": "abc123",
  87. "active_family_keys": ["f1", "f2"],
  88. },
  89. "families": [
  90. {
  91. "key": "f1",
  92. "name": "实质 × 模态 × 业务阶段 × 知识类型",
  93. "axes": ["实质", "模态"],
  94. "items": [
  95. {
  96. "query": "反转 视频 脚本 怎么做",
  97. "parts": {"实质": "反转", "模态": "视频"},
  98. "keep": True,
  99. "valid": 8,
  100. "relevant": True,
  101. "reason": "可搜创作方法",
  102. }
  103. ],
  104. }
  105. ],
  106. }
  107. batch, count = persist_query_batch(repo, generated, name="formal-demo")
  108. assert batch.name == "formal-demo"
  109. assert count == 1
  110. assert repo.batch_kwargs["status"] == "ready"
  111. assert repo.batch_kwargs["target_platforms"] == ["xiaohongshu", "weixin", "douyin"]
  112. assert repo.batch_kwargs["metadata"]["active_family_keys"] == ["f1", "f2"]
  113. row = repo.queries[0]
  114. assert row["batch_id"] == batch.id
  115. assert row["query_text"] == "反转 视频 脚本 怎么做"
  116. assert row["axes"] == {"实质": "反转", "模态": "视频"}
  117. assert row["keep"] is True
  118. assert row["filter_reason"] == "可搜创作方法"
  119. assert row["metadata"]["family_key"] == "f1"
  120. assert row["metadata"]["valid"] == 8
  121. assert row["metadata"]["query_filter_prompt_version"] == "abc123"