377 lines
13 KiB
Python
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
|