test_classification_coarse.py 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186
  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. from core.text_limits import CLASSIFY_BODY_MAX_CHARS
  6. def _settings() -> Settings:
  7. return Settings(
  8. pg=PgConfig(host="h", port=5432, user="u", password="p", database="d"),
  9. aiddit_crawler_base_url="http://crawler.test",
  10. crawler_timeout=30,
  11. openrouter_timeout_seconds=90,
  12. openrouter_model="m",
  13. openrouter_base_url="http://openrouter.test",
  14. openrouter_api_key="k",
  15. llm_model="m",
  16. max_cards=12,
  17. frames_dir="f",
  18. douyin_ratio="540p",
  19. data_dir="data",
  20. )
  21. def test_formal_classify_imgtext_passes_http_image_urls(monkeypatch):
  22. captured = {}
  23. def fake_judge(messages, settings, timeout):
  24. captured["messages"] = messages
  25. return 1, "ok", "先定受众", ""
  26. monkeypatch.setattr(coarse, "_judge", fake_judge)
  27. result = coarse.classify_imgtext(
  28. {
  29. "platform": "xiaohongshu",
  30. "title": "脚本创作",
  31. "body_text": "讲怎么设计开头",
  32. "images": ["https://cdn.test/a.jpg"],
  33. },
  34. _settings(),
  35. )
  36. assert result[0] == 1
  37. user_content = captured["messages"][1]["content"]
  38. assert {"type": "image_url", "image_url": {"url": "https://cdn.test/a.jpg"}} in user_content
  39. def test_formal_classify_imgtext_uses_wide_body_limit(monkeypatch):
  40. captured = {}
  41. def fake_judge(messages, settings, timeout):
  42. captured["messages"] = messages
  43. return 1, "ok", "长正文方法", ""
  44. monkeypatch.setattr(coarse, "_judge", fake_judge)
  45. long_body = "甲" * (CLASSIFY_BODY_MAX_CHARS + 7)
  46. coarse.classify_imgtext(
  47. {
  48. "platform": "xiaohongshu",
  49. "title": "长文",
  50. "body_text": long_body,
  51. "images": [],
  52. },
  53. _settings(),
  54. )
  55. user_text = captured["messages"][1]["content"][0]["text"]
  56. assert user_text.count("甲") == CLASSIFY_BODY_MAX_CHARS
  57. def test_legacy_classify_imgtext_also_passes_http_image_urls(monkeypatch):
  58. captured = {}
  59. def fake_judge(messages, settings, timeout):
  60. captured["messages"] = messages
  61. return 1, "ok", "先定受众", ""
  62. monkeypatch.setattr(legacy_classify, "_judge", fake_judge)
  63. result = legacy_classify.classify_imgtext(
  64. {
  65. "platform": "weixin",
  66. "title": "公众号选题",
  67. "body_text": "讲怎么做标题",
  68. "images": ["https://cdn.test/w.jpg"],
  69. },
  70. _settings(),
  71. )
  72. assert result[0] == 1
  73. user_content = captured["messages"][1]["content"]
  74. assert {"type": "image_url", "image_url": {"url": "https://cdn.test/w.jpg"}} in user_content
  75. def test_legacy_classify_video_includes_title_and_body_text(monkeypatch):
  76. captured = {}
  77. def fake_judge(messages, settings, timeout):
  78. captured["messages"] = messages
  79. return 1, "ok", "视频方法", ""
  80. monkeypatch.setattr(legacy_classify, "_judge", fake_judge)
  81. legacy_classify.classify_video(
  82. {
  83. "platform": "douyin",
  84. "title": "视频标题",
  85. "body_text": "视频文案",
  86. "video": "https://cdn.test/video.mp4",
  87. },
  88. _settings(),
  89. )
  90. user_text = captured["messages"][1]["content"][0]["text"]
  91. assert "平台:douyin" in user_text
  92. assert "标题:视频标题" in user_text
  93. assert "正文/文案:视频文案" in user_text
  94. def test_coarse_classify_item_maps_judge_result(monkeypatch):
  95. monkeypatch.setattr(
  96. coarse,
  97. "classify_imgtext",
  98. lambda payload, settings: (0, "不是创作知识", "", ""),
  99. )
  100. result = coarse.coarse_classify_item(
  101. platform="weixin",
  102. title="普通知识",
  103. body_text="只是作品介绍",
  104. image_urls=["https://cdn.test/a.jpg"],
  105. settings=_settings(),
  106. )
  107. assert result.is_creation_knowledge is False
  108. assert result.label == "not_creation"
  109. assert result.status == "classified"
  110. def test_coarse_classify_item_routes_douyin_image_post_to_imgtext(monkeypatch):
  111. called = {}
  112. def fake_imgtext(payload, settings):
  113. called["imgtext"] = payload
  114. return 1, "是创作知识", "图片方法", ""
  115. def fake_video(payload, settings):
  116. called["video"] = payload
  117. return 1, "不应调用", "", ""
  118. monkeypatch.setattr(coarse, "classify_imgtext", fake_imgtext)
  119. monkeypatch.setattr(coarse, "classify_video", fake_video)
  120. result = coarse.coarse_classify_item(
  121. platform="douyin",
  122. content_mode="image_post",
  123. title="图片帖",
  124. body_text="图片正文",
  125. image_urls=["https://cdn.test/dy.jpg"],
  126. settings=_settings(),
  127. )
  128. assert result.is_creation_knowledge is True
  129. assert "imgtext" in called
  130. assert "video" not in called
  131. def test_coarse_classify_item_skips_unsupported_and_missing_video():
  132. unsupported = coarse.coarse_classify_item(
  133. platform="douyin",
  134. content_mode="unsupported",
  135. settings=_settings(),
  136. )
  137. missing_video = coarse.coarse_classify_item(
  138. platform="douyin",
  139. content_mode="video_post",
  140. settings=_settings(),
  141. )
  142. assert unsupported.status == "skipped"
  143. assert unsupported.label == "unsupported_content_mode"
  144. assert missing_video.status == "skipped"
  145. assert missing_video.label == "video_missing"