Files
news/tests/test_embedding.py
T
2026-07-18 15:51:01 +08:00

377 lines
13 KiB
Python

"""M5 嵌入模块单元测试。
不依赖真实 LLM/HuggingFace,所有 provider 调用通过 mock 注入。
"""
from __future__ import annotations
import asyncio
import json
from datetime import datetime
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock
import pytest
from embedding import (
DASHSCOPE_BATCH_LIMIT,
AsyncEmbeddingProvider,
EmbeddingError,
EmbeddingProvider,
EmbeddingProviderType,
EmbeddingResult,
compose_text,
make_async_provider,
make_sync_provider,
resolve_provider_type,
)
from embedding.base import _from_event_dict
from embedding.remote import (
DashScopeAsyncEmbeddingProvider,
DashScopeEmbeddingProvider,
_chunked,
)
from extractor import Article
# --------------------------------------------------------------------------- #
# fixtures
# --------------------------------------------------------------------------- #
def _article(
*,
url: str = "https://www.cls.cn/detail/1",
url_hash: str = "abc1234567890000",
title: str = "宁德时代签订 100GWh 长期供货协议",
content: str = "宁德时代与某车企签 5 年 100GWh 协议,涉及金额超 1500 亿。" * 3,
publish_time: datetime | None = datetime(2026, 6, 16, 10, 0),
) -> Article:
return Article(
source_id="cls",
url=url,
url_hash=url_hash,
title=title,
content=content,
publish_time=publish_time,
word_count=len(content),
)
def _embedding_response(vectors: list[list[float]]) -> MagicMock:
"""构造与 OpenAI SDK 一致的 embeddings.create 返回。"""
resp = MagicMock()
resp.data = [MagicMock(embedding=v) for v in vectors]
return resp
# --------------------------------------------------------------------------- #
# compose_text
# --------------------------------------------------------------------------- #
def test_compose_text_basic() -> None:
art = _article()
text = compose_text(art)
assert text.startswith("标题:")
assert "正文:" in text
assert art.title in text
assert art.content[:30] in text
def test_compose_text_with_head_and_summary() -> None:
art = _article()
text = compose_text(art, head="[sentiment=positive]", summary="一句话摘要")
assert "[sentiment=positive]" in text
assert "摘要:一句话摘要" in text
def test_compose_text_truncates_overlong() -> None:
art = _article(content="字" * 10000)
text = compose_text(art, max_chars=500)
assert len(text) <= 500
def test_from_event_dict_extracts_head_and_article() -> None:
event_obj = {
"source_id": "cls",
"url": "https://x/1",
"url_hash": "h1",
"title": "宁德合作",
"publish_time": "2026-06-16T10:00:00",
"event": {
"stock_codes": ["300750.SZ"],
"company_names": ["宁德时代"],
"industries": ["动力电池"],
"sentiment": "positive",
"importance": 5,
"event_type": "重大合同",
"summary": "签订 100GWh 协议",
},
}
article, head, summary = _from_event_dict(event_obj)
assert article.url_hash == "h1"
assert article.publish_time == datetime(2026, 6, 16, 10, 0, 0)
assert "sentiment=positive" in head
assert "importance=5" in head
assert "300750.SZ" in head
assert "宁德时代" in head
assert "动力电池" in head
assert summary == "签订 100GWh 协议"
def test_from_event_dict_handles_missing_publish_time() -> None:
article, _, _ = _from_event_dict({"source_id": "x", "url": "u", "url_hash": "h",
"title": "t", "event": {
"sentiment": "neutral",
"importance": 1,
"event_type": "其他",
}})
assert article.publish_time is None
# --------------------------------------------------------------------------- #
# remote 工具
# --------------------------------------------------------------------------- #
def test_chunked_splits_evenly() -> None:
assert _chunked(list(range(7)), 3) == [[0, 1, 2], [3, 4, 5], [6]]
assert _chunked([], 3) == []
assert _chunked([1, 2, 3], 10) == [[1, 2, 3]]
def test_dashscope_batch_limit_is_10() -> None:
assert DASHSCOPE_BATCH_LIMIT == 10
# --------------------------------------------------------------------------- #
# Provider type 解析
# --------------------------------------------------------------------------- #
def test_resolve_provider_type_dashscope_aliases(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("EMBEDDING_PROVIDER", raising=False)
assert resolve_provider_type("dashscope") == EmbeddingProviderType.DASHSCOPE
assert resolve_provider_type("qwen") == EmbeddingProviderType.DASHSCOPE
assert resolve_provider_type("remote") == EmbeddingProviderType.DASHSCOPE
def test_resolve_provider_type_local_aliases(monkeypatch: pytest.MonkeyPatch) -> None:
for name in ("local", "local-bge", "bge", "bge-m3"):
assert resolve_provider_type(name) == EmbeddingProviderType.LOCAL_BGE
def test_resolve_provider_type_default_is_dashscope(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("EMBEDDING_PROVIDER", raising=False)
assert resolve_provider_type() == EmbeddingProviderType.DASHSCOPE
def test_resolve_provider_type_env_override(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("EMBEDDING_PROVIDER", "local-bge")
assert resolve_provider_type() == EmbeddingProviderType.LOCAL_BGE
def test_resolve_provider_type_unknown_raises() -> None:
with pytest.raises(EmbeddingError):
resolve_provider_type("anthropic-emb")
# --------------------------------------------------------------------------- #
# DashScope 同步/异步(mock 网络)
# --------------------------------------------------------------------------- #
def test_dashscope_sync_embed_batch(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeEmbeddingProvider(model="text-embedding-v3")
fake = MagicMock()
fake.embeddings.create = MagicMock(
side_effect=lambda model, input: _embedding_response([[0.1] * 1024] * len(input))
)
provider._client = fake
out = provider.embed_batch(["a", "b", "c"])
assert len(out) == 3
assert all(len(v) == 1024 for v in out)
def test_dashscope_sync_chunks_when_over_batch_limit(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeEmbeddingProvider()
fake = MagicMock()
fake.embeddings.create = MagicMock(
side_effect=lambda model, input: _embedding_response([[0.0] * 1024] * len(input))
)
provider._client = fake
texts = [f"t{i}" for i in range(25)] # > 10 -> 应分 3 批 (10+10+5)
provider.embed_batch(texts)
assert fake.embeddings.create.call_count == 3
def test_dashscope_sync_retries_then_succeeds(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeEmbeddingProvider(max_attempts=3)
fake = MagicMock()
fake.embeddings.create = MagicMock(
side_effect=[
RuntimeError("rate-limit"),
_embedding_response([[0.1] * 1024]),
]
)
provider._client = fake
out = provider.embed_batch(["x"])
assert len(out) == 1
assert fake.embeddings.create.call_count == 2
def test_dashscope_sync_gives_up(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeEmbeddingProvider(max_attempts=2)
fake = MagicMock()
fake.embeddings.create = MagicMock(side_effect=RuntimeError("net"))
provider._client = fake
with pytest.raises(EmbeddingError) as exc:
provider.embed_batch(["x"])
assert exc.value.attempts == 2
def test_dashscope_missing_api_key_raises(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("DASHSCOPE_API_KEY", raising=False)
with pytest.raises(EmbeddingError):
DashScopeEmbeddingProvider()
@pytest.mark.asyncio
async def test_dashscope_async_embed_batch(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeAsyncEmbeddingProvider()
fake = MagicMock()
fake.embeddings.create = AsyncMock(
side_effect=lambda model, input: _embedding_response([[0.1] * 1024] * len(input))
)
provider._client = fake
out = await provider.embed_batch(["a", "b"])
assert len(out) == 2
@pytest.mark.asyncio
async def test_dashscope_async_chunks(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeAsyncEmbeddingProvider()
fake = MagicMock()
fake.embeddings.create = AsyncMock(
side_effect=lambda model, input: _embedding_response([[0.0] * 1024] * len(input))
)
provider._client = fake
texts = [f"t{i}" for i in range(15)] # 2 批 (10+5)
await provider.embed_batch(texts)
assert fake.embeddings.create.await_count == 2
@pytest.mark.asyncio
async def test_dashscope_async_retries(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
provider = DashScopeAsyncEmbeddingProvider(max_attempts=2)
fake = MagicMock()
fake.embeddings.create = AsyncMock(
side_effect=[RuntimeError("transient"), _embedding_response([[0.0] * 1024])]
)
provider._client = fake
out = await provider.embed_batch(["a"])
assert len(out) == 1
# --------------------------------------------------------------------------- #
# factory
# --------------------------------------------------------------------------- #
def test_make_sync_provider_dashscope(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
p = make_sync_provider("dashscope")
assert isinstance(p, EmbeddingProvider)
assert p.name == "dashscope"
assert p.dim == 1024
def test_make_async_provider_dashscope(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test")
p = make_async_provider("qwen")
assert isinstance(p, AsyncEmbeddingProvider)
assert p.name == "dashscope"
def test_make_sync_provider_local_without_st_raises(monkeypatch: pytest.MonkeyPatch) -> None:
"""无 sentence-transformers 时,本地 provider 应给出友好错误。"""
import sys
# 模拟 sentence_transformers 缺失
monkeypatch.setitem(sys.modules, "sentence_transformers", None)
with pytest.raises(EmbeddingError) as exc:
make_sync_provider("local")
assert "sentence-transformers" in str(exc.value)
# --------------------------------------------------------------------------- #
# EmbeddingResult 模型 + 序列化
# --------------------------------------------------------------------------- #
def test_embedding_result_serializes(tmp_path: Path) -> None:
r = EmbeddingResult(
url_hash="abc",
source_id="cls",
title="t",
text="text",
vector=[0.1, 0.2, 0.3],
dim=3,
provider="dashscope",
model="text-embedding-v3",
)
p = tmp_path / "r.json"
p.write_text(r.model_dump_json(), encoding="utf-8")
obj = json.loads(p.read_text(encoding="utf-8"))
assert obj["dim"] == 3
assert obj["vector"] == [0.1, 0.2, 0.3]
def test_embedding_result_short_summary() -> None:
r = EmbeddingResult(
url_hash="abc",
source_id="cls",
title="宁德时代签约",
text="x",
vector=[0.0] * 4,
dim=4,
provider="dashscope",
model="text-embedding-v3",
)
s = r.short_summary()
assert "cls" in s and "dim=4" in s and "dashscope" in s
# --------------------------------------------------------------------------- #
# 集成式: compose_text + 假异步 provider
# --------------------------------------------------------------------------- #
class _FakeAsyncProvider(AsyncEmbeddingProvider):
name = "fake"
model = "fake-1"
dim = 8
async def embed_batch(self, texts: list[str]) -> list[list[float]]:
return [[float(len(t))] * self.dim for t in texts]
@pytest.mark.asyncio
async def test_fake_async_provider_round_trip() -> None:
art = _article()
text = compose_text(art)
async with _FakeAsyncProvider() as p:
v = (await p.embed_batch([text]))[0]
assert len(v) == 8
assert v[0] == float(len(text))
@pytest.mark.asyncio
async def test_fake_async_provider_concurrent_batches() -> None:
p = _FakeAsyncProvider()
res = await asyncio.gather(
p.embed_batch(["a", "bb"]),
p.embed_batch(["ccc"]),
)
assert res[0][0][0] == 1.0
assert res[0][1][0] == 2.0
assert res[1][0][0] == 3.0