test_classification_coarse.py 2.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990
  1. from __future__ import annotations
  2. import acquisition.classification.coarse as coarse
  3. import acquisition.classify as legacy_classify
  4. from core.config import PgConfig, Settings
  5. def _settings() -> Settings:
  6. return Settings(
  7. pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
  8. aiddit_crawler_base_url="http://crawler.test",
  9. crawler_timeout=30,
  10. openrouter_timeout_seconds=90,
  11. openrouter_model="m",
  12. openrouter_base_url="http://openrouter.test",
  13. openrouter_api_key="k",
  14. llm_model="m",
  15. max_cards=12,
  16. frames_dir="f",
  17. douyin_ratio="540p",
  18. data_dir="data",
  19. )
  20. def test_formal_classify_imgtext_passes_http_image_urls(monkeypatch):
  21. captured = {}
  22. def fake_judge(messages, settings, timeout):
  23. captured["messages"] = messages
  24. return 1, "ok", "先定受众", ""
  25. monkeypatch.setattr(coarse, "_judge", fake_judge)
  26. result = coarse.classify_imgtext(
  27. {
  28. "platform": "xiaohongshu",
  29. "title": "脚本创作",
  30. "body_text": "讲怎么设计开头",
  31. "images": ["https://cdn.test/a.jpg"],
  32. },
  33. _settings(),
  34. )
  35. assert result[0] == 1
  36. user_content = captured["messages"][1]["content"]
  37. assert {"type": "image_url", "image_url": {"url": "https://cdn.test/a.jpg"}} in user_content
  38. def test_legacy_classify_imgtext_also_passes_http_image_urls(monkeypatch):
  39. captured = {}
  40. def fake_judge(messages, settings, timeout):
  41. captured["messages"] = messages
  42. return 1, "ok", "先定受众", ""
  43. monkeypatch.setattr(legacy_classify, "_judge", fake_judge)
  44. result = legacy_classify.classify_imgtext(
  45. {
  46. "platform": "weixin",
  47. "title": "公众号选题",
  48. "body_text": "讲怎么做标题",
  49. "images": ["https://cdn.test/w.jpg"],
  50. },
  51. _settings(),
  52. )
  53. assert result[0] == 1
  54. user_content = captured["messages"][1]["content"]
  55. assert {"type": "image_url", "image_url": {"url": "https://cdn.test/w.jpg"}} in user_content
  56. def test_coarse_classify_item_maps_judge_result(monkeypatch):
  57. monkeypatch.setattr(
  58. coarse,
  59. "classify_imgtext",
  60. lambda payload, settings: (0, "不是创作知识", "", ""),
  61. )
  62. result = coarse.coarse_classify_item(
  63. platform="weixin",
  64. title="普通知识",
  65. body_text="只是作品介绍",
  66. image_urls=["https://cdn.test/a.jpg"],
  67. settings=_settings(),
  68. )
  69. assert result.is_creation_knowledge is False
  70. assert result.label == "not_creation"
  71. assert result.status == "classified"