Files
news/tests/test_llm.py
T
simon ff911cf6f7 feat: Token Plan 迁移与 .env 热加载,并修复日报 AI 摘要为空
Token Plan 迁移 / 配置热加载:
- configs/llm_models.yaml: 各场景切到 Token Plan(deepseek-v4.1-flash / qwen3.6-flash)
- 新增 configs/runtime_env.py: .env 按 (mtime_ns, size) 热加载并同步 os.environ,
  统一 env_get 取值;llm / embedding / vectorstore / mcp / pipeline 改用 env_get
- configs/loader.py / scripts/run_scheduler.py 等配套调整
- 新增 tests/test_hot_reload.py

日报 AI 摘要为空修复(2026-09-25):
- 根因: 推理模型的 reasoning token 与正文共用 max_tokens, 预算 1500 被"思考"
  占满 -> text_tokens=0 / finish_reason=length, 摘要静默为空且不重试
- daily_report 场景新增 max_tokens(默认 4000, YAML 保存即热生效);
  LLMConfig 支持可选 max_tokens; 分块预算 800 -> 2000
- _llm_call 拆出 _call_once, 正文为空时自动加倍预算重试(上限 16000),
  用尽才降级返回空串; 网络异常重试语义不变
- docs/user-guide.md 新增 FAQ; continuation.md 记录本次排查
- 已重跑 2026-09-25 日报(report_id=357)补回 466 字摘要

测试: 相关用例 56 passed(test_hot_reload 12 passed);
      ruff 无新增问题; 3 个 crawler 既有失败与本改动无关
2026-09-25 11:13:37 +08:00

558 lines
21 KiB
Python

"""M4 LLM 投资事件抽取测试。
不依赖真实 LLM API,所有调用通过 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 extractor import Article
from llm import (
EVENT_TYPES,
EventExtraction,
ExtractedEvent,
LLMCallError,
PromptTemplate,
Sentiment,
extract_event,
extract_event_async,
load_llm_config,
parse_event_json,
)
from llm.client import LLMConfig
from llm.extractor import _extract_json_object
# --------------------------------------------------------------------------- #
# fixtures
# --------------------------------------------------------------------------- #
def _article(
*,
title: str = "宁德时代签订 100GWh 长期供货协议",
content: str = "宁德时代(300750)与某车企签 5 年 100GWh 协议,涉及金额超 1500 亿。" * 3,
publish_time: datetime | None = datetime(2026, 6, 16, 10, 0),
) -> Article:
return Article(
source_id="cls",
url="https://www.cls.cn/detail/1",
url_hash="abc1234567890000",
title=title,
content=content,
publish_time=publish_time,
word_count=len(content),
)
@pytest.fixture
def fake_config() -> LLMConfig:
return LLMConfig(
provider="deepseek",
model="deepseek-chat",
api_key="sk-fake",
base_url="https://api.deepseek.com",
)
def _mock_completion(content: str, prompt_tokens: int = 100, completion_tokens: int = 50) -> MagicMock:
"""构造与 openai SDK 返回兼容的 mock 对象。"""
msg = MagicMock()
msg.content = content
choice = MagicMock()
choice.message = msg
usage = MagicMock()
usage.prompt_tokens = prompt_tokens
usage.completion_tokens = completion_tokens
resp = MagicMock()
resp.choices = [choice]
resp.usage = usage
return resp
# --------------------------------------------------------------------------- #
# EventExtraction 模型校验
# --------------------------------------------------------------------------- #
def test_event_extraction_minimal_valid() -> None:
e = EventExtraction(sentiment="positive", importance=4, event_type="重大合同")
assert e.sentiment == Sentiment.POSITIVE
assert e.importance == 4
def test_event_extraction_normalizes_stock_codes() -> None:
e = EventExtraction(
stock_codes=["300750.sz", " 300750.SZ ", "abc", "12345", "600519"],
sentiment="positive", importance=3, event_type="其他",
)
# 大小写 / 空白被规范;非法被过滤;去重
assert e.stock_codes == ["300750.SZ", "600519"]
def test_event_extraction_filters_empty_lists() -> None:
e = EventExtraction(
stock_codes=[], company_names=["", " ", "宁德时代", "宁德时代"],
industries=[],
sentiment="neutral", importance=1, event_type="其他",
)
assert e.company_names == ["宁德时代"]
assert e.stock_codes == []
def test_event_extraction_rejects_importance_out_of_range() -> None:
with pytest.raises(Exception): # noqa: B017
EventExtraction(sentiment="positive", importance=0, event_type="其他")
with pytest.raises(Exception): # noqa: B017
EventExtraction(sentiment="positive", importance=6, event_type="其他")
def test_event_extraction_normalizes_blank_event_type() -> None:
e = EventExtraction(sentiment="neutral", importance=1, event_type=" ")
assert e.event_type == "其他"
def test_event_types_constant_includes_common() -> None:
for must in ["业绩预告", "合作签约", "监管处罚", "其他"]:
assert must in EVENT_TYPES
# --------------------------------------------------------------------------- #
# ExtractedEvent.sources 多源字段
# --------------------------------------------------------------------------- #
def test_extracted_event_sources_defaults_to_main_source() -> None:
"""未提供 sources 时兜底为 [source_id](兼容旧产物)。"""
ev = ExtractedEvent(
source_id="cls", url="https://x/1", url_hash="h1", title="t",
event=EventExtraction(sentiment="positive", importance=3, event_type="重大合同"),
provider="deepseek", model="m",
)
assert ev.sources == ["cls"]
def test_extracted_event_sources_keeps_main_first_and_dedup() -> None:
"""sources 保主源居首、去重保序。"""
ev = ExtractedEvent(
source_id="cls", url="https://x/1", url_hash="h1", title="t",
sources=["sina", "cls", "eastmoney", "sina"],
event=EventExtraction(sentiment="neutral", importance=2, event_type="其他"),
provider="deepseek", model="m",
)
assert ev.sources == ["cls", "sina", "eastmoney"]
# --------------------------------------------------------------------------- #
# JSON 提取与解析
# --------------------------------------------------------------------------- #
def test_extract_json_object_strips_fence() -> None:
s = '```json\n{"a": 1}\n```'
assert _extract_json_object(s) == '{"a": 1}'
def test_extract_json_object_picks_first_object() -> None:
s = '前置说明\n{"a": 1}\n更多文字'
assert _extract_json_object(s) == '{"a": 1}'
def test_extract_json_object_handles_nested() -> None:
s = '{"a": {"b": 2}}'
assert _extract_json_object(s) == '{"a": {"b": 2}}'
def test_parse_event_json_ok() -> None:
raw = json.dumps({
"stock_codes": ["300750.SZ"],
"company_names": ["宁德时代"],
"industries": ["动力电池"],
"sentiment": "positive",
"importance": 5,
"event_type": "重大合同",
"summary": "签订长期供货协议",
})
e = parse_event_json(raw)
assert e.sentiment == Sentiment.POSITIVE
assert e.stock_codes == ["300750.SZ"]
def test_parse_event_json_invalid_json_raises() -> None:
with pytest.raises(LLMCallError):
parse_event_json("not a json")
def test_parse_event_json_non_object_raises() -> None:
with pytest.raises(LLMCallError):
parse_event_json('["array"]')
def test_parse_event_json_schema_invalid_raises() -> None:
with pytest.raises(LLMCallError):
parse_event_json('{"importance": 99}') # 缺 sentiment + event_type 且 importance 越界
# --------------------------------------------------------------------------- #
# PromptTemplate
# --------------------------------------------------------------------------- #
def test_prompt_template_renders_placeholders(tmp_path: Path) -> None:
tpl_file = tmp_path / "tpl.md"
tpl_file.write_text(
"标题:{title}\n时间:{publish_time}\n源:{source_name}\n内容:\n{content}\nEND",
encoding="utf-8",
)
tpl = PromptTemplate(tpl_file)
art = _article()
rendered = tpl.render(art)
assert "标题:" + art.title in rendered
assert "2026-06-16" in rendered
assert "源:财联社" in rendered or "源:cls" in rendered # source_name 默认空,落到 source_id
assert art.content[:30] in rendered
def test_prompt_template_truncates_long_content(tmp_path: Path) -> None:
tpl_file = tmp_path / "tpl.md"
tpl_file.write_text("{content}", encoding="utf-8")
tpl = PromptTemplate(tpl_file)
art = _article(content="字" * 20000)
rendered = tpl.render(art)
assert "[正文过长已截断]" in rendered
assert len(rendered) < 20000
def test_prompt_template_default_path_loads() -> None:
"""项目内置 prompts/event_extraction.md 必须可加载,作为回归保护。"""
real = Path("prompts/event_extraction.md")
if not real.is_file():
pytest.skip("prompts/event_extraction.md 未找到")
tpl = PromptTemplate(real)
out = tpl.render(_article())
assert "{title}" not in out
assert "{content}" not in out
# --------------------------------------------------------------------------- #
# extract_event(同步,带重试)
# --------------------------------------------------------------------------- #
def test_extract_event_succeeds_first_try(fake_config: LLMConfig) -> None:
client = MagicMock()
raw = json.dumps({
"stock_codes": ["300750.SZ"],
"company_names": ["宁德时代"],
"industries": ["动力电池"],
"sentiment": "positive",
"importance": 5,
"event_type": "重大合同",
"summary": "100GWh 合作",
})
client.chat.completions.create = MagicMock(return_value=_mock_completion(raw))
result = extract_event(client, fake_config, _article())
assert isinstance(result, ExtractedEvent)
assert result.attempts == 1
assert result.event.sentiment == Sentiment.POSITIVE
assert result.provider == "deepseek"
assert result.prompt_tokens == 100
def test_extract_event_retries_on_invalid_json(fake_config: LLMConfig) -> None:
"""第 1/2 次返回非法 JSON,第 3 次成功。"""
valid = json.dumps({
"sentiment": "neutral", "importance": 1, "event_type": "其他",
})
client = MagicMock()
client.chat.completions.create = MagicMock(side_effect=[
_mock_completion("not a json"),
_mock_completion('{"sentiment":"???"}'), # schema 校验失败
_mock_completion(valid),
])
result = extract_event(client, fake_config, _article(), max_attempts=3)
assert result.attempts == 3
assert client.chat.completions.create.call_count == 3
def test_extract_event_gives_up_after_max(fake_config: LLMConfig) -> None:
client = MagicMock()
client.chat.completions.create = MagicMock(
return_value=_mock_completion("not a json")
)
with pytest.raises(LLMCallError) as exc:
extract_event(client, fake_config, _article(), max_attempts=2)
assert exc.value.attempts == 2
assert client.chat.completions.create.call_count == 2
def test_extract_event_handles_network_exception(fake_config: LLMConfig) -> None:
client = MagicMock()
valid = json.dumps({
"sentiment": "negative", "importance": 3, "event_type": "监管处罚",
})
client.chat.completions.create = MagicMock(side_effect=[
TimeoutError("net hang"),
_mock_completion(valid),
])
result = extract_event(client, fake_config, _article(), max_attempts=2)
assert result.attempts == 2
assert result.event.sentiment == Sentiment.NEGATIVE
# --------------------------------------------------------------------------- #
# extract_event_async
# --------------------------------------------------------------------------- #
@pytest.mark.asyncio
async def test_extract_event_async_succeeds(fake_config: LLMConfig) -> None:
client = MagicMock()
valid = json.dumps({
"stock_codes": ["600519"],
"company_names": ["贵州茅台"],
"sentiment": "neutral",
"importance": 2,
"event_type": "财报披露",
"summary": "披露半年报",
})
client.chat.completions.create = AsyncMock(return_value=_mock_completion(valid))
sem = asyncio.Semaphore(2)
result = await extract_event_async(
client, fake_config, _article(), semaphore=sem,
)
assert result.event.stock_codes == ["600519"]
assert result.event.event_type == "财报披露"
@pytest.mark.asyncio
async def test_extract_event_async_retries(fake_config: LLMConfig) -> None:
valid = json.dumps({"sentiment": "positive", "importance": 4, "event_type": "其他"})
client = MagicMock()
client.chat.completions.create = AsyncMock(side_effect=[
ValueError("transient"),
_mock_completion(valid),
])
result = await extract_event_async(
client, fake_config, _article(), max_attempts=2,
)
assert result.attempts == 2
# --------------------------------------------------------------------------- #
# load_llm_config
# --------------------------------------------------------------------------- #
def test_load_llm_config_deepseek_from_env(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LLM_PROVIDER", "deepseek")
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-test-deepseek")
monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-v4-flash")
monkeypatch.delenv("LLM_MODEL", raising=False)
cfg = load_llm_config()
assert cfg.provider == "deepseek"
assert cfg.api_key == "sk-test-deepseek"
assert cfg.model == "deepseek-v4-flash"
assert "deepseek" in cfg.base_url
def test_load_llm_config_qwen_from_env(
monkeypatch: pytest.MonkeyPatch, tmp_path
) -> None:
"""qwen 的 key 兜底链:QWEN_API_KEY 缺失时回退 DASHSCOPE_API_KEY。
真实 .env 里带有 QWEN_API_KEY,会盖过 DASHSCOPE_API_KEY,所以本用例先把
热加载指向一个空 .env,测完再切回真实 .env。
"""
from configs import runtime_env
empty = tmp_path / ".env"
empty.write_text("", encoding="utf-8")
monkeypatch.setenv(runtime_env.ENV_FILE_OVERRIDE, str(empty))
runtime_env.reset_cache()
monkeypatch.delenv("QWEN_API_KEY", raising=False)
monkeypatch.delenv("QWEN_BASE_URL", raising=False)
monkeypatch.setenv("LLM_PROVIDER", "qwen")
monkeypatch.setenv("DASHSCOPE_API_KEY", "sk-test-qwen")
monkeypatch.setenv("QWEN_MODEL", "qwen-plus")
monkeypatch.delenv("LLM_MODEL", raising=False)
cfg = load_llm_config()
assert cfg.provider == "qwen"
assert cfg.api_key == "sk-test-qwen"
assert cfg.model == "qwen-plus"
assert "dashscope" in cfg.base_url or "aliyuncs" in cfg.base_url
# 切回真实 .env,避免影响后续用例(DB / Qdrant 等直接读 os.environ 的测试)
monkeypatch.delenv(runtime_env.ENV_FILE_OVERRIDE, raising=False)
runtime_env.reset_cache()
runtime_env.ensure_env_loaded(force=True)
def test_load_llm_config_unknown_provider_raises(monkeypatch: pytest.MonkeyPatch) -> None:
with pytest.raises(ValueError):
load_llm_config(provider="anthropic")
def test_load_llm_config_missing_key_raises(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.delenv("DEEPSEEK_API_KEY", raising=False)
monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-v4-flash")
with pytest.raises(ValueError, match="API key"):
load_llm_config(provider="deepseek")
def test_load_llm_config_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -> None:
"""去掉内置默认模型后:未显式配置模型必须报错(不再回退 deepseek-chat)。"""
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-test")
monkeypatch.delenv("DEEPSEEK_MODEL", raising=False)
monkeypatch.delenv("LLM_MODEL", raising=False)
with pytest.raises(ValueError, match="模型"):
load_llm_config(provider="deepseek")
# --------------------------------------------------------------------------- #
# load_llm_config —— configs/llm_models.yaml 场景配置
# --------------------------------------------------------------------------- #
def _patch_scene(monkeypatch: pytest.MonkeyPatch, cfg: dict) -> None:
"""替换场景加载,模拟 configs/llm_models.yaml 中的某场景配置。"""
monkeypatch.setattr(
"llm.client.load_scene_config",
lambda scene: cfg if scene == "daily_report" else {},
)
def test_load_llm_config_scene_overrides_env(monkeypatch: pytest.MonkeyPatch) -> None:
"""YAML 场景配置优先于 .env:provider / model / temperature / timeout / max_attempts。"""
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env")
monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-env-model")
monkeypatch.setenv("QWEN_API_KEY", "sk-qwen")
_patch_scene(monkeypatch, {
"provider": "qwen",
"model": "qwen-max",
"temperature": 0.5,
"timeout_sec": 99,
"max_attempts": 5,
})
cfg = load_llm_config(scene="daily_report")
assert cfg.provider == "qwen"
assert cfg.model == "qwen-max"
assert cfg.api_key == "sk-qwen"
assert cfg.temperature == 0.5
assert cfg.timeout_sec == 99
assert cfg.max_attempts == 5
def test_load_llm_config_scene_blank_fields_fall_back_to_env(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""YAML 场景未配置的字段(如 model 留空)回退 .env,保持向后兼容。"""
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env")
monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-env-model")
_patch_scene(monkeypatch, {"provider": "deepseek", "model": "", "temperature": 0.7})
cfg = load_llm_config(scene="daily_report")
assert cfg.provider == "deepseek"
assert cfg.model == "deepseek-env-model"
assert cfg.api_key == "sk-env"
assert cfg.temperature == 0.7
def test_load_llm_config_scene_api_key_env_name(monkeypatch: pytest.MonkeyPatch) -> None:
"""api_key_env 指向自定义环境变量时,优先使用该变量。"""
monkeypatch.setenv("MY_CUSTOM_KEY", "sk-custom")
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-default")
monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-m")
_patch_scene(monkeypatch, {
"provider": "deepseek",
"model": "deepseek-scene-m",
"api_key_env": "MY_CUSTOM_KEY",
})
cfg = load_llm_config(scene="daily_report")
assert cfg.api_key == "sk-custom"
assert cfg.model == "deepseek-scene-m"
def test_load_llm_config_scene_explicit_args_win(monkeypatch: pytest.MonkeyPatch) -> None:
"""CLI/显式参数优先级最高,覆盖 YAML 场景。"""
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env")
_patch_scene(monkeypatch, {"provider": "qwen", "model": "qwen-max"})
monkeypatch.setenv("QWEN_API_KEY", "sk-qwen")
cfg = load_llm_config(provider="deepseek", model="deepseek-chat", scene="daily_report")
assert cfg.provider == "deepseek"
assert cfg.model == "deepseek-chat"
def test_load_llm_config_scene_missing_model_raises(monkeypatch: pytest.MonkeyPatch) -> None:
"""场景与 .env 都未配置模型时必须报错(无内置兜底)。"""
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-test")
monkeypatch.delenv("DEEPSEEK_MODEL", raising=False)
monkeypatch.delenv("LLM_MODEL", raising=False)
_patch_scene(monkeypatch, {"provider": "deepseek", "model": ""})
with pytest.raises(ValueError, match="模型"):
load_llm_config(scene="daily_report")
def test_load_llm_config_real_yaml_parseable() -> None:
"""真实 configs/llm_models.yaml 必须可解析且包含全部场景(回归保护)。"""
from configs.loader import load_defaults, load_scene_config
for scene in ("event_extraction", "daily_report", "stock_report", "embedding"):
assert isinstance(load_scene_config(scene), dict)
assert isinstance(load_defaults(), dict)
def test_load_llm_config_temperature_zero_is_respected(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""temperature=0 是合法配置,不应被 or 链回退默认。"""
monkeypatch.setenv("DEEPSEEK_API_KEY", "sk-env")
monkeypatch.setenv("DEEPSEEK_MODEL", "deepseek-m")
_patch_scene(monkeypatch, {"provider": "deepseek", "model": "deepseek-m", "temperature": 0})
cfg = load_llm_config(scene="daily_report")
assert cfg.temperature == 0.0
def test_load_llm_config_scene_max_attempts_zero() -> None:
"""max_attempts=0 由 _pick_int 显式处理。"""
from llm.client import _pick_int
assert _pick_int({"max_attempts": 0}, "max_attempts", 3) == 0
assert _pick_int({"max_attempts": ""}, "max_attempts", 3) == 3
assert _pick_int({}, "max_attempts", 3) == 3
# --------------------------------------------------------------------------- #
# max_tokens 场景配置(推理模型 reasoning 与正文共用预算,回归 9-25 无 AI 摘要)
# --------------------------------------------------------------------------- #
def test_load_llm_config_scene_max_tokens(monkeypatch: pytest.MonkeyPatch) -> None:
"""scenes.<scene>.max_tokens 应写入 LLMConfig;未配置则为 None(调用方走默认)。"""
monkeypatch.setenv("QWEN_API_KEY", "sk-qwen")
_patch_scene(monkeypatch, {"provider": "qwen", "model": "deepseek-v4.1-flash",
"max_tokens": 4000})
cfg = load_llm_config(scene="daily_report")
assert cfg.max_tokens == 4000
_patch_scene(monkeypatch, {"provider": "qwen", "model": "deepseek-v4.1-flash"})
assert load_llm_config(scene="daily_report").max_tokens is None
def test_pick_optional_int() -> None:
"""_pick_optional_int:未配置/空值/非法 → None;数字字符串可解析。"""
from llm.client import _pick_optional_int
assert _pick_optional_int({"max_tokens": 4000}, "max_tokens") == 4000
assert _pick_optional_int({"max_tokens": "4000"}, "max_tokens") == 4000
assert _pick_optional_int({"max_tokens": ""}, "max_tokens") is None
assert _pick_optional_int({"max_tokens": "abc"}, "max_tokens") is None
assert _pick_optional_int({}, "max_tokens") is None
def test_real_yaml_daily_report_max_tokens_has_reasoning_headroom() -> None:
"""真实配置的日报预算必须大于旧的 1500(为 reasoning token 留余量)。"""
from configs.loader import load_scene_config
scene = load_scene_config("daily_report")
max_tokens = int(scene.get("max_tokens") or 0)
assert max_tokens >= 4000, f"daily_report.max_tokens 过小: {max_tokens}"