test_ai_generated_materials.py 2.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869
  1. import unittest
  2. import os
  3. import sys
  4. from pathlib import Path
  5. from dotenv import load_dotenv
  6. _HERE = Path(__file__).parent
  7. sys.path.insert(0, str(_HERE))
  8. sys.path.insert(0, str(_HERE.parent.parent))
  9. from tools.ai_generated_material import (
  10. _public_oss_url,
  11. build_generation_prompts,
  12. )
  13. from tools.account_material_strategy import normalize_material_source, parse_bool_flag
  14. from tools.video_feature_query import read_cached_video_element_features
  15. class AiGeneratedMaterialsTest(unittest.TestCase):
  16. def test_normalize_material_source_defaults_to_history(self):
  17. self.assertEqual(normalize_material_source(""), "history")
  18. self.assertEqual(normalize_material_source("历史素材"), "history")
  19. self.assertEqual(normalize_material_source("AI生成素材"), "ai_generated")
  20. self.assertEqual(normalize_material_source("外部素材"), "external_recall")
  21. def test_parse_bool_flag_accepts_feishu_yes_no(self):
  22. self.assertTrue(parse_bool_flag("是"))
  23. self.assertFalse(parse_bool_flag("否"))
  24. self.assertFalse(parse_bool_flag(""))
  25. def test_build_generation_prompts_uses_cached_video_topic(self):
  26. load_dotenv(_HERE / ".env")
  27. video_id = int(os.getenv("AI_MATERIAL_TEST_VIDEO_ID", "69131739"))
  28. try:
  29. features = read_cached_video_element_features([video_id]).get(video_id) or []
  30. except Exception as e:
  31. self.skipTest(f"数据库不可用,跳过真实视频特征测试:{e}")
  32. topic = next((f for f in features if f.element_dimension == "解构选题"), None)
  33. if topic is None:
  34. self.skipTest(f"video_id={video_id} 在本地缓存表中没有解构选题")
  35. prompts = build_generation_prompts(
  36. video_id=video_id,
  37. title="不应作为兜底标题",
  38. category="不应作为兜底品类",
  39. features=features,
  40. )
  41. self.assertEqual([p.prompt_type for p in prompts], ["topic"])
  42. self.assertIn(topic.standard_element, prompts[0].prompt_text)
  43. self.assertEqual(prompts[0].feature_hits[0]["original_standard_element"], topic.standard_element)
  44. self.assertEqual(prompts[0].feature_hits[0]["video_description"], topic.standard_element)
  45. self.assertIn("腾讯广告信息流 / 公众号投放", prompts[0].prompt_text)
  46. self.assertNotIn("不应作为兜底标题", prompts[0].prompt_text)
  47. self.assertNotIn("不应作为兜底品类", prompts[0].prompt_text)
  48. self.assertIn("16:9", prompts[0].prompt_text)
  49. def test_public_oss_url_keeps_directory_slashes(self):
  50. self.assertEqual(
  51. _public_oss_url(
  52. "https://rescdn.yishihui.com/",
  53. "auto_put_tencent/image/a b/test.jpg",
  54. ),
  55. "https://rescdn.yishihui.com/auto_put_tencent/image/a%20b/test.jpg",
  56. )
  57. if __name__ == "__main__":
  58. unittest.main()