"""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