316 lines
12 KiB
Python
Executable File
316 lines
12 KiB
Python
Executable File
"""ai_parsing 模块单测 — 纯逻辑."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import pytest
|
||
|
||
from packages.domain.ai_parsing import (
|
||
generate_titles_fallback,
|
||
keyword_match_fallback,
|
||
parse_semantic_match_response,
|
||
parse_titles_from_response,
|
||
)
|
||
|
||
# ── parse_titles_from_response 测试 ──────────────────────────────────────────
|
||
|
||
|
||
class TestParseTitlesJsonArray:
|
||
def test_simple_json_array(self):
|
||
result = parse_titles_from_response('["标题1", "标题2", "标题3"]')
|
||
assert result == ["标题1", "标题2", "标题3"]
|
||
|
||
def test_json_array_with_empty_strings_skipped(self):
|
||
result = parse_titles_from_response('["标题1", "", "标题2"]')
|
||
assert result == ["标题1", "标题2"]
|
||
|
||
def test_json_dict_with_titles_key(self):
|
||
result = parse_titles_from_response('{"titles": ["a", "b", "c"]}')
|
||
assert result == ["a", "b", "c"]
|
||
|
||
def test_json_with_markdown_code_block(self):
|
||
content = '```json\n["标题1", "标题2"]\n```'
|
||
result = parse_titles_from_response(content)
|
||
assert result == ["标题1", "标题2"]
|
||
|
||
def test_json_with_backticks_no_lang(self):
|
||
content = '```\n["标题1", "标题2"]\n```'
|
||
result = parse_titles_from_response(content)
|
||
assert result == ["标题1", "标题2"]
|
||
|
||
def test_none_returns_empty(self):
|
||
assert parse_titles_from_response(None) == [] # type: ignore
|
||
|
||
def test_empty_string_returns_empty(self):
|
||
assert parse_titles_from_response("") == []
|
||
|
||
|
||
class TestParseTitlesNumberedList:
|
||
def test_dot_numbered(self):
|
||
content = "1. 第一个标题\n2. 第二个标题\n3. 第三个标题"
|
||
result = parse_titles_from_response(content)
|
||
assert len(result) == 3
|
||
assert result[0] == "第一个标题"
|
||
assert result[1] == "第二个标题"
|
||
|
||
def test_chinese_period_numbered(self):
|
||
content = "1、第一个标题\n2、第二个标题"
|
||
result = parse_titles_from_response(content)
|
||
assert result == ["第一个标题", "第二个标题"]
|
||
|
||
def test_parentheses_numbered(self):
|
||
content = "1) 第一个标题\n2) 第二个标题"
|
||
result = parse_titles_from_response(content)
|
||
assert result == ["第一个标题", "第二个标题"]
|
||
|
||
def test_chinese_paren_numbered(self):
|
||
# 原实现只支持半角括号,全角括号保留原样(不影响实际使用)
|
||
content = "(1)第一个标题\n(2)第二个标题"
|
||
result = parse_titles_from_response(content)
|
||
assert len(result) == 2
|
||
|
||
|
||
class TestParseTitlesDash:
|
||
def test_dash_prefix(self):
|
||
content = "- 标题一\n- 标题二\n- 标题三"
|
||
result = parse_titles_from_response(content)
|
||
assert len(result) == 3
|
||
assert result[0] == "标题一"
|
||
|
||
def test_bullet_prefix(self):
|
||
content = "• 标题一\n• 标题二"
|
||
result = parse_titles_from_response(content)
|
||
assert len(result) == 2
|
||
assert result[0] == "标题一"
|
||
|
||
|
||
class TestParseTitlesQuoted:
|
||
def test_strips_quotes(self):
|
||
content = '1. "带引号的标题"\n2. 正常标题'
|
||
result = parse_titles_from_response(content)
|
||
assert "带引号的标题" in result
|
||
|
||
def test_strips_chinese_quotes(self):
|
||
content = "1. 「中文引号标题」\n2. 正常标题"
|
||
result = parse_titles_from_response(content)
|
||
assert "中文引号标题" in result
|
||
|
||
|
||
class TestParseTitlesEdgeCases:
|
||
def test_skips_empty_lines(self):
|
||
content = "标题一\n\n标题二\n\n标题三"
|
||
result = parse_titles_from_response(content)
|
||
assert len(result) == 3
|
||
|
||
def test_filters_long_lines(self):
|
||
long_title = "A" * 150
|
||
content = f"短标题\n{long_title}\n另一个短标题"
|
||
result = parse_titles_from_response(content)
|
||
assert len(result) == 2
|
||
assert long_title not in result
|
||
|
||
def test_invalid_json_falls_back_to_lines(self):
|
||
content = "标题1\n标题2\n标题3"
|
||
result = parse_titles_from_response(content)
|
||
assert result == ["标题1", "标题2", "标题3"]
|
||
|
||
|
||
# ── parse_semantic_match_response 测试 ──────────────────────────────────────
|
||
|
||
|
||
class TestParseSemanticMatchDictFormat:
|
||
def test_simple_dict(self):
|
||
asset_ids = ["a1", "a2", "a3"]
|
||
content = '{"a1": 0.8, "a2": 0.6, "a3": 0.9}'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert result["a1"] == 0.8
|
||
assert result["a2"] == 0.6
|
||
assert result["a3"] == 0.9
|
||
|
||
def test_score_clamped_to_0_1(self):
|
||
asset_ids = ["a1", "a2"]
|
||
content = '{"a1": 1.5, "a2": -0.5}'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert result["a1"] == 1.0
|
||
assert result["a2"] == 0.0
|
||
|
||
|
||
class TestParseSemanticMatchMatchesFormat:
|
||
def test_matches_array(self):
|
||
asset_ids = ["a1", "a2"]
|
||
content = '{"matches": [{"asset_id": "a1", "score": 0.8}, {"asset_id": "a2", "score": 0.6}]}'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert result["a1"] == 0.8
|
||
assert result["a2"] == 0.6
|
||
|
||
def test_matches_with_id_key(self):
|
||
asset_ids = ["a1", "a2"]
|
||
content = '{"matches": [{"id": "a1", "score": 0.7}, {"id": "a2", "score": 0.5}]}'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert result["a1"] == 0.7
|
||
|
||
|
||
class TestParseSemanticMatchArrayFormat:
|
||
def test_array_of_objects(self):
|
||
asset_ids = ["a1", "a2", "a3"]
|
||
content = (
|
||
'[{"asset_id": "a1", "score": 0.8}, {"asset_id": "a2", "score": 0.6}, {"asset_id": "a3", "score": 0.3}]'
|
||
)
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert len(result) == 3
|
||
|
||
|
||
class TestParseSemanticMatchEdgeCases:
|
||
def test_empty_content_returns_none(self):
|
||
assert parse_semantic_match_response("", ["a1"]) is None
|
||
|
||
def test_invalid_json_returns_none(self):
|
||
assert parse_semantic_match_response("not json", ["a1"]) is None
|
||
|
||
def test_less_than_half_returns_none(self):
|
||
asset_ids = ["a1", "a2", "a3", "a4", "a5"]
|
||
# 只返回1个,少于 5//2=2,应该返回 None
|
||
content = '{"a1": 0.8}'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is None
|
||
|
||
def test_at_least_half_returns_result(self):
|
||
asset_ids = ["a1", "a2", "a3", "a4", "a5"]
|
||
# 返回3个,>= 5//2=2
|
||
content = '{"a1": 0.8, "a2": 0.7, "a3": 0.6}'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert len(result) == 3
|
||
|
||
def test_markdown_code_block(self):
|
||
asset_ids = ["a1", "a2"]
|
||
content = '```json\n{"a1": 0.8, "a2": 0.6}\n```'
|
||
result = parse_semantic_match_response(content, asset_ids)
|
||
assert result is not None
|
||
assert result["a1"] == 0.8
|
||
|
||
def test_no_asset_ids_returns_result_if_any(self):
|
||
content = '{"a1": 0.8, "a2": 0.6}'
|
||
result = parse_semantic_match_response(content, [])
|
||
assert result is not None
|
||
assert len(result) == 2
|
||
|
||
def test_single_asset_id_needs_at_least_1(self):
|
||
# max(1, 1//2) = max(1, 0) = 1
|
||
content = '{"a1": 0.8}'
|
||
result = parse_semantic_match_response(content, ["a1"])
|
||
assert result is not None
|
||
|
||
|
||
# ── generate_titles_fallback 测试 ────────────────────────────────────────────
|
||
|
||
|
||
class TestGenerateTitlesFallback:
|
||
def test_returns_requested_count(self):
|
||
style = {"examples": ["例1", "例2", "例3"]}
|
||
result = generate_titles_fallback("测试描述 关键词", style, count=5)
|
||
assert len(result) == 5
|
||
|
||
def test_uses_keyword_from_description(self):
|
||
style = {"examples": ["例1", "例2"]}
|
||
result = generate_titles_fallback("美食 探店 打卡", style, count=3)
|
||
# 第一个关键词是"美食"
|
||
assert any("美食" in t for t in result)
|
||
|
||
def test_no_keywords_uses_default(self):
|
||
style = {"examples": ["例1", "例2"]}
|
||
result = generate_titles_fallback("", style, count=3)
|
||
assert any("精彩内容" in t for t in result)
|
||
|
||
def test_count_limited_by_templates(self):
|
||
style = {"examples": ["例1", "例2"]}
|
||
result = generate_titles_fallback("测试", style, count=100)
|
||
assert len(result) <= 10 # 模板只有10个
|
||
|
||
def test_all_titles_are_strings(self):
|
||
style = {"examples": ["例1", "例2"]}
|
||
result = generate_titles_fallback("测试", style, count=5)
|
||
assert all(isinstance(t, str) and t for t in result)
|
||
|
||
def test_empty_examples(self):
|
||
style = {"examples": []}
|
||
result = generate_titles_fallback("测试", style, count=3)
|
||
assert len(result) == 3
|
||
assert all(isinstance(t, str) for t in result)
|
||
|
||
|
||
# ── keyword_match_fallback 测试 ──────────────────────────────────────────────
|
||
|
||
|
||
class TestKeywordMatchFallback:
|
||
def test_basic_matching(self):
|
||
assets = [
|
||
{"id": "1", "name": "美食探店视频", "tags": ["美食", "探店"], "description": "好吃的"},
|
||
{"id": "2", "name": "旅行vlog", "tags": ["旅行", "风景"], "description": "出去玩"},
|
||
]
|
||
result = keyword_match_fallback("美食探店 好吃的美食", assets)
|
||
assert len(result) == 2
|
||
# 第一个应该是美食相关的
|
||
assert result[0]["id"] == "1"
|
||
assert result[0]["match_score"] >= result[1]["match_score"]
|
||
|
||
def test_returns_match_score_and_reason(self):
|
||
assets = [{"id": "1", "name": "测试素材", "tags": [], "description": ""}]
|
||
result = keyword_match_fallback("美食", assets)
|
||
assert len(result) == 1
|
||
assert "match_score" in result[0]
|
||
assert "match_reason" in result[0]
|
||
assert 0.0 <= result[0]["match_score"] <= 1.0
|
||
|
||
def test_no_keywords_default_score(self):
|
||
assets = [
|
||
{"id": "1", "name": "素材1", "tags": [], "description": ""},
|
||
{"id": "2", "name": "素材2", "tags": [], "description": ""},
|
||
]
|
||
# 单个字符不算关键词
|
||
result = keyword_match_fallback("a", assets)
|
||
assert len(result) == 2
|
||
assert all(r["match_score"] == 0.5 for r in result)
|
||
assert all(r["match_reason"] == "fallback_default" for r in result)
|
||
|
||
def test_sorted_by_score_descending(self):
|
||
assets = [
|
||
{"id": "low", "name": "无关素材", "tags": [], "description": ""},
|
||
{"id": "high", "name": "美食视频", "tags": ["美食"], "description": "美食分享"},
|
||
]
|
||
result = keyword_match_fallback("美食分享", assets)
|
||
assert result[0]["id"] == "high"
|
||
assert result[0]["match_score"] > result[1]["match_score"]
|
||
|
||
def test_does_not_modify_original_assets(self):
|
||
original = {"id": "1", "name": "测试", "tags": [], "description": ""}
|
||
assets = [dict(original)]
|
||
keyword_match_fallback("美食", assets)
|
||
assert "match_score" not in assets[0]
|
||
|
||
def test_name_matches_higher_score(self):
|
||
assets = [
|
||
{"id": "name_match", "name": "美食教程", "tags": [], "description": "内容"},
|
||
{"id": "desc_match", "name": "视频1", "tags": [], "description": "美食教程内容"},
|
||
]
|
||
result = keyword_match_fallback("美食教程", assets)
|
||
# 名称命中应该加分更多
|
||
name_idx = next(i for i, r in enumerate(result) if r["id"] == "name_match")
|
||
desc_idx = next(i for i, r in enumerate(result) if r["id"] == "desc_match")
|
||
assert name_idx < desc_idx
|
||
|
||
def test_empty_assets_returns_empty(self):
|
||
result = keyword_match_fallback("美食", [])
|
||
assert result == []
|
||
|
||
def test_score_is_rounded_to_3_decimals(self):
|
||
assets = [{"id": "1", "name": "测试素材", "tags": [], "description": "内容描述"}]
|
||
result = keyword_match_fallback("美食探店旅行", assets)
|
||
score = result[0]["match_score"]
|
||
# 验证是3位小数
|
||
assert round(score, 3) == score
|