Files
intl_news/tests/test_embedding.py
T
simon 9aa44e610d feat: LLM/Embedding 重试次数纳入 system.yaml 统一配置
- system.yaml: llm.max_retries(死配置) 改为 llm.max_attempts;embedding 段新增 max_attempts
- llm/client.py: LLMConfig 新增 max_attempts(默认 3 兜底),load_llm_config 读取配置
- llm/extractor.py: translate_and_extract(_async) max_attempts=None 时取 config.max_attempts
- embedding/client.py: load_embedding_config 读取 embedding.max_attempts
- scheduler/reporter.py: _call_llm_simple max_retries=None 时取 max_attempts-1
- 新增 2 个配置读取测试;已同步 pi5 验证
2026-08-05 08:48:43 +08:00

309 lines
9.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""M5 向量生成模块单元测试。"""
from unittest.mock import MagicMock
import pytest
from openai import OpenAI
from embedding.client import (
EmbeddingConfig,
_chunked,
embed_batch,
load_embedding_config,
)
from embedding.embedder import compose_text, embed_article
from embedding.models import EmbeddingError, EmbeddingResult
from llm.models import EnTranslatedArticle, EventExtraction, Sentiment
# --------------------------------------------------------------------------- #
# 辅助工厂
# --------------------------------------------------------------------------- #
def _make_article(
*,
url_hash: str = "abc123",
source_id: str = "reuters",
source_name: str = "Reuters",
url: str = "https://example.com/1",
title: str = "Fed Holds Rates Steady",
title_zh: str = "美联储维持利率不变",
content_en: str = "The Fed held rates steady.",
content_zh: str = "美联储维持利率不变,理由是经济增长稳健。市场应声上涨。",
events: list[EventExtraction] | None = None,
) -> EnTranslatedArticle:
if events is None:
events = [
EventExtraction(
event_type="央行决议",
stock_codes=[],
sentiment=Sentiment.POSITIVE,
importance=5,
summary_zh="美联储维持利率不变,市场反弹",
)
]
return EnTranslatedArticle(
source_id=source_id,
source_name=source_name,
url=url,
url_hash=url_hash,
title=title,
title_zh=title_zh,
content_en=content_en,
content_zh=content_zh,
events=events,
provider="deepseek",
model="deepseek-v4-flash",
)
# --------------------------------------------------------------------------- #
# EmbeddingConfig
# --------------------------------------------------------------------------- #
class TestEmbeddingConfig:
"""EmbeddingConfig 测试。"""
def test_valid_config(self):
cfg = EmbeddingConfig(
model="text-embedding-v3",
api_key="sk-test",
)
assert cfg.provider == "dashscope"
assert cfg.dimension == 1024
def test_missing_api_key_raises(self):
with pytest.raises(EmbeddingError, match="API_KEY"):
EmbeddingConfig(api_key="")
class TestLoadEmbeddingConfig:
"""load_embedding_config 测试。"""
def test_from_env(self, monkeypatch):
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-dashscope-test")
cfg = load_embedding_config()
assert cfg.provider == "dashscope"
assert cfg.api_key == "sk-dashscope-test"
def test_max_attempts_from_system_yaml(self, monkeypatch):
"""重试次数取自 system.yaml embedding.max_attempts(当前 3)。"""
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-dashscope-test")
cfg = load_embedding_config()
assert cfg.max_attempts == 3
def test_fallback_to_qwen_key(self, monkeypatch):
monkeypatch.delenv("DASHSCOPE_API_KEY", raising=False)
monkeypatch.setenv("QWEN_API_KEY", "sk-qwen-key")
cfg = load_embedding_config()
assert cfg.api_key == "sk-qwen-key"
def test_missing_key_raises(self, monkeypatch):
monkeypatch.delenv("DASHSCOPE_API_KEY", raising=False)
monkeypatch.delenv("QWEN_API_KEY", raising=False)
with pytest.raises(EmbeddingError, match="API_KEY"):
load_embedding_config()
# --------------------------------------------------------------------------- #
# _chunked
# --------------------------------------------------------------------------- #
class TestChunked:
"""_chunked 分块工具测试。"""
def test_empty(self):
assert _chunked([], 10) == []
def test_single_chunk(self):
assert _chunked(["a", "b", "c"], 10) == [["a", "b", "c"]]
def test_multiple_chunks(self):
assert _chunked(["a", "b", "c", "d", "e"], 2) == [
["a", "b"], ["c", "d"], ["e"]
]
def test_exact_fit(self):
assert _chunked(["a", "b", "c", "d"], 2) == [
["a", "b"], ["c", "d"]
]
# --------------------------------------------------------------------------- #
# embed_batchmock
# --------------------------------------------------------------------------- #
class TestEmbedBatch:
"""embed_batch 测试(mock DashScope API)。"""
def test_empty_list(self):
cfg = EmbeddingConfig(api_key="sk-test")
client = MagicMock(spec=OpenAI)
result = embed_batch(client, cfg, [])
assert result == []
def test_single_text(self):
cfg = EmbeddingConfig(api_key="sk-test", batch_size=5)
client = MagicMock(spec=OpenAI)
# Mock 返回
mock_resp = MagicMock()
mock_resp.data = [MagicMock(embedding=[0.1] * 1024)]
client.embeddings.create.return_value = mock_resp
result = embed_batch(client, cfg, ["测试文本"])
assert len(result) == 1
assert len(result[0]) == 1024
def test_batch_split(self):
"""超过 batch_size 自动分多批。"""
cfg = EmbeddingConfig(api_key="sk-test", batch_size=3)
client = MagicMock(spec=OpenAI)
def _make_mock(**kwargs):
input_texts = kwargs.get("input", [])
resp = MagicMock()
resp.data = [MagicMock(embedding=[0.5] * 1024) for _ in input_texts]
return resp
client.embeddings.create.side_effect = _make_mock
texts = ["a", "b", "c", "d", "e"] # 需要分 3+2 两批
result = embed_batch(client, cfg, texts)
assert len(result) == 5
assert client.embeddings.create.call_count == 2
# --------------------------------------------------------------------------- #
# compose_text
# --------------------------------------------------------------------------- #
class TestComposeText:
"""compose_text 测试。"""
def test_basic_composition(self):
article = _make_article()
text = compose_text(article)
assert "标题:" in text
assert "美联储维持利率不变" in text
assert "事件:" in text
assert "央行决议" in text
assert "正文:" in text
def test_no_events(self):
article = _make_article(events=[])
text = compose_text(article)
assert "事件:" not in text
assert "标题:" in text
assert "正文:" in text
def test_with_stock_codes(self):
article = _make_article(events=[
EventExtraction(
event_type="财报披露",
stock_codes=["AAPL", "TSLA"],
sentiment=Sentiment.POSITIVE,
importance=4,
summary_zh="苹果财报超预期",
)
])
text = compose_text(article)
assert "AAPL" in text
assert "TSLA" in text
assert "positive" in text
def test_truncation(self):
"""超长文本应被截断。"""
long_content = "这是测试正文。" * 500 # ~2500 chars
article = _make_article(content_zh=long_content)
text = compose_text(article, max_chars=500)
assert len(text) <= 500
def test_empty_article(self):
article = _make_article(title_zh="", content_zh="", events=[])
text = compose_text(article)
# 不会崩溃即可
assert isinstance(text, str)
# --------------------------------------------------------------------------- #
# embed_articlemock
# --------------------------------------------------------------------------- #
class TestEmbedArticle:
"""embed_article 测试(mock)。"""
def test_successful_embed(self):
cfg = EmbeddingConfig(api_key="sk-test")
client = MagicMock(spec=OpenAI)
mock_resp = MagicMock()
mock_resp.data = [MagicMock(embedding=[0.1] * 1024)]
client.embeddings.create.return_value = mock_resp
article = _make_article()
result = embed_article(client, cfg, article)
assert isinstance(result, EmbeddingResult)
assert result.url_hash == article.url_hash
assert result.dimension == 1024
assert len(result.vector) == 1024
assert result.provider == "dashscope"
assert len(result.embedded_text) > 0
# --------------------------------------------------------------------------- #
# EmbeddingResult
# --------------------------------------------------------------------------- #
class TestEmbeddingResult:
"""EmbeddingResult 模型测试。"""
def test_minimal(self):
r = EmbeddingResult(
url_hash="abc",
source_id="reuters",
vector=[0.5] * 1024,
model="text-embedding-v3",
)
assert r.dimension == 1024
assert len(r.vector) == 1024
def test_serialization(self):
r = EmbeddingResult(
url_hash="abc",
source_id="reuters",
vector=[0.1, 0.2, 0.3],
dimension=3,
embedded_text="测试",
model="text-embedding-v3",
)
data = r.model_dump_json()
assert "url_hash" in data
assert "vector" in data
# --------------------------------------------------------------------------- #
# EmbeddingError
# --------------------------------------------------------------------------- #
class TestEmbeddingError:
"""EmbeddingError 异常测试。"""
def test_basic(self):
err = EmbeddingError("测试错误", attempts=3)
assert err.reason == "测试错误"
assert err.attempts == 3
assert str(err) == "测试错误"
def test_default_attempts(self):
err = EmbeddingError("错误")
assert err.attempts == 0