test_ai_generated_materials.py 2.8 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768
  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. def test_parse_bool_flag_accepts_feishu_yes_no(self):
  21. self.assertTrue(parse_bool_flag("是"))
  22. self.assertFalse(parse_bool_flag("否"))
  23. self.assertFalse(parse_bool_flag(""))
  24. def test_build_generation_prompts_uses_cached_video_topic(self):
  25. load_dotenv(_HERE / ".env")
  26. video_id = int(os.getenv("AI_MATERIAL_TEST_VIDEO_ID", "69131739"))
  27. try:
  28. features = read_cached_video_element_features([video_id]).get(video_id) or []
  29. except Exception as e:
  30. self.skipTest(f"数据库不可用,跳过真实视频特征测试:{e}")
  31. topic = next((f for f in features if f.element_dimension == "解构选题"), None)
  32. if topic is None:
  33. self.skipTest(f"video_id={video_id} 在本地缓存表中没有解构选题")
  34. prompts = build_generation_prompts(
  35. video_id=video_id,
  36. title="不应作为兜底标题",
  37. category="不应作为兜底品类",
  38. features=features,
  39. )
  40. self.assertEqual([p.prompt_type for p in prompts], ["topic"])
  41. self.assertIn(topic.standard_element, prompts[0].prompt_text)
  42. self.assertEqual(prompts[0].feature_hits[0]["original_standard_element"], topic.standard_element)
  43. self.assertEqual(prompts[0].feature_hits[0]["video_description"], topic.standard_element)
  44. self.assertIn("腾讯广告信息流 / 公众号投放", prompts[0].prompt_text)
  45. self.assertNotIn("不应作为兜底标题", prompts[0].prompt_text)
  46. self.assertNotIn("不应作为兜底品类", prompts[0].prompt_text)
  47. self.assertIn("16:9", prompts[0].prompt_text)
  48. def test_public_oss_url_keeps_directory_slashes(self):
  49. self.assertEqual(
  50. _public_oss_url(
  51. "https://rescdn.yishihui.com/",
  52. "auto_put_tencent/image/a b/test.jpg",
  53. ),
  54. "https://rescdn.yishihui.com/auto_put_tencent/image/a%20b/test.jpg",
  55. )
  56. if __name__ == "__main__":
  57. unittest.main()