docs: 文档清理与重构 — 统一为 3 个核心文档

- 删除 5 个过时/残留文档(project_plan/agent_prompt/optimization_plan/report_db_design/deploy/README)
- 新建 docs/architecture.md(项目架构:11 包职责+数据模型+配置+产物)
- 重写 docs/user-guide.md(CLI 全量+增量/断点续跑+MCP+FAQ)
- 重写 README.md(精简入口+文档索引)
- 更新 continuation.md(追加本次记录)
- 更新 .gitignore(排除 data/* 运行产物)
This commit is contained in:
2026-08-22 17:10:39 +08:00
commit 65ead54b4f
101 changed files with 21093 additions and 0 deletions
+64
View File
@@ -0,0 +1,64 @@
"""LLM 投资事件抽取模块 (M4)。
公共 API:
- load_llm_config / make_sync_client / make_async_client
- extract_event / extract_event_async
- PromptTemplate / parse_event_json
- EventExtraction / ExtractedEvent / Sentiment / EVENT_TYPES / LLMCallError
"""
from .client import (
DEFAULT_MAX_ATTEMPTS,
DEFAULT_TEMPERATURE,
DEFAULT_TIMEOUT_SEC,
SCENE_DAILY_REPORT,
SCENE_EVENT_EXTRACTION,
SCENE_STOCK_REPORT,
LLMConfig,
load_llm_config,
make_async_client,
make_sync_client,
)
from .extractor import (
DEFAULT_PROMPT_PATH,
MAX_CONTENT_CHARS,
PromptTemplate,
extract_event,
extract_event_async,
parse_event_json,
)
from .models import (
EVENT_TYPES,
MAX_IMPORTANCE,
MIN_IMPORTANCE,
EventExtraction,
ExtractedEvent,
LLMCallError,
Sentiment,
)
__all__ = [
"DEFAULT_MAX_ATTEMPTS",
"DEFAULT_PROMPT_PATH",
"DEFAULT_TEMPERATURE",
"DEFAULT_TIMEOUT_SEC",
"EVENT_TYPES",
"MAX_CONTENT_CHARS",
"MAX_IMPORTANCE",
"MIN_IMPORTANCE",
"SCENE_DAILY_REPORT",
"SCENE_EVENT_EXTRACTION",
"SCENE_STOCK_REPORT",
"EventExtraction",
"ExtractedEvent",
"LLMCallError",
"LLMConfig",
"PromptTemplate",
"Sentiment",
"extract_event",
"extract_event_async",
"load_llm_config",
"make_async_client",
"make_sync_client",
"parse_event_json",
]
+208
View File
@@ -0,0 +1,208 @@
"""LLM 客户端抽象与工厂。
支持 DeepSeek 和 Qwen(百炼),两者均为 OpenAI 兼容接口,共用 openai SDK。
配置来源(优先级从高到低):
1. 代码 / CLI 显式参数(provider / model)
2. configs/llm_models.yaml 场景配置(scene 参数,见 configs/loader.py)
3. 环境变量 / .env(LLM_PROVIDER、DEEPSEEK_MODEL 等,向后兼容)
4. 代码内置默认值
环境变量(兜底):
LLM_PROVIDER = deepseek | qwen (默认 deepseek)
DeepSeek: DEEPSEEK_API_KEY / DEEPSEEK_BASE_URL / DEEPSEEK_MODEL
Qwen: QWEN_API_KEY / QWEN_BASE_URL / QWEN_MODEL
(QWEN_API_KEY -> DASHSCOPE_API_KEY 兜底)
模型必须显式配置(provider 对应的 *_MODEL 或 LLM_MODEL),不再提供内置默认模型。
LLM_TEMPERATURE / LLM_TIMEOUT_SEC
"""
from __future__ import annotations
import os
from dataclasses import dataclass
from loguru import logger
from openai import AsyncOpenAI, OpenAI
from configs.loader import load_defaults, load_scene_config
# 默认基址
_DEEPSEEK_DEFAULT_BASE = "https://api.deepseek.com"
_QWEN_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
# 抽取任务默认参数
DEFAULT_TIMEOUT_SEC = 60.0
DEFAULT_TEMPERATURE = 0.1
DEFAULT_MAX_ATTEMPTS = 3
# 场景名 -> configs/llm_models.yaml 中 scenes 的 key
SCENE_EVENT_EXTRACTION = "event_extraction"
SCENE_DAILY_REPORT = "daily_report"
SCENE_STOCK_REPORT = "stock_report"
@dataclass
class LLMConfig:
"""LLM 调用配置(provider / model / api_key / base_url / 参数)。"""
provider: str # "deepseek" / "qwen"
model: str
api_key: str
base_url: str
timeout_sec: float = DEFAULT_TIMEOUT_SEC
temperature: float = DEFAULT_TEMPERATURE
max_attempts: int = DEFAULT_MAX_ATTEMPTS # 单次任务失败重试次数
def __post_init__(self) -> None:
if not self.api_key:
raise ValueError(f"LLM provider={self.provider} 的 API key 为空")
def _read_env(key: str, default: str | None = None) -> str | None:
val = os.environ.get(key)
if val is None or val.strip() == "":
return default
return val.strip()
def _first_env(keys: list[str | None]) -> str | None:
"""按顺序返回第一个非空的环境变量值。"""
for k in keys:
if not k:
continue
v = _read_env(k)
if v:
return v
return None
def _num(value: object) -> float | None:
"""把 YAML 数字/字符串安全转 float;非法或为空返回 None。"""
if value is None or value == "":
return None
try:
return float(value)
except (TypeError, ValueError):
return None
def load_llm_config(
provider: str | None = None,
*,
model: str | None = None,
scene: str | None = None,
) -> LLMConfig:
"""按优先级构造 LLMConfig:显式参数 > YAML 场景 > 环境变量 > 内置默认。
scene 对应 configs/llm_models.yaml 中 scenes 的 key
(event_extraction / daily_report / stock_report),该场景未配置的字段
回退到环境变量,保持向后兼容。
"""
sc = load_scene_config(scene or "")
dflt = load_defaults()
p = (
provider
or sc.get("provider")
or _read_env("LLM_PROVIDER", "deepseek")
or "deepseek"
).lower()
# 各 provider 的 api_key / base_url / model 环境变量链
provider_envs: dict[str, tuple[list[str | None], list[str | None], list[str | None]]] = {
"deepseek": (
[sc.get("api_key_env"), "DEEPSEEK_API_KEY"],
[sc.get("base_url_env"), "DEEPSEEK_BASE_URL"],
["DEEPSEEK_MODEL", "LLM_MODEL"],
),
"qwen": (
[sc.get("api_key_env"), "QWEN_API_KEY", "DASHSCOPE_API_KEY"],
[sc.get("base_url_env"), "QWEN_BASE_URL"],
["QWEN_MODEL", "LLM_MODEL"],
),
}
if p == "deepseek":
key_envs, base_envs, model_envs = provider_envs["deepseek"]
default_base = _DEEPSEEK_DEFAULT_BASE
elif p in ("qwen", "dashscope"):
key_envs, base_envs, model_envs = provider_envs["qwen"]
default_base = _QWEN_DEFAULT_BASE
p = "qwen" # 内部统一用 qwen
else:
raise ValueError(f"未知 LLM provider: {p!r},仅支持 deepseek / qwen")
api_key = _first_env(key_envs) or ""
base_url = _first_env(base_envs) or default_base
# 模型优先级:显式参数 > YAML 场景 > 环境变量;模型必须显式配置,无内置兜底
m = model or sc.get("model") or _first_env(model_envs)
if not m:
env_hint = "/".join(v for v in model_envs if v)
raise ValueError(
f"未配置 LLM 模型(场景 {scene or 'default'}): "
f"请在 configs/llm_models.yaml 的 model 或 .env 设置 {env_hint}"
)
timeout = _pick_float(sc, dflt, "timeout_sec", "LLM_TIMEOUT_SEC", DEFAULT_TIMEOUT_SEC)
temperature = _pick_float(sc, dflt, "temperature", "LLM_TEMPERATURE", DEFAULT_TEMPERATURE)
max_attempts = _pick_int(sc, "max_attempts", DEFAULT_MAX_ATTEMPTS)
return LLMConfig(
provider=p,
model=m,
api_key=api_key,
base_url=base_url,
timeout_sec=timeout,
temperature=temperature,
max_attempts=max_attempts,
)
def _pick_float(
sc: dict,
dflt: dict,
sc_key: str,
env_key: str,
default: float,
) -> float:
"""数值参数选择:YAML 场景 > 环境变量 > YAML defaults > 内置默认(零值合法)。"""
v = _num(sc.get(sc_key))
if v is not None:
return v
v = _num(_read_env(env_key))
if v is not None:
return v
v = _num(dflt.get(sc_key))
return v if v is not None else default
def _pick_int(sc: dict, sc_key: str, default: int) -> int:
v = _num(sc.get(sc_key))
return int(v) if v is not None else default
def make_sync_client(config: LLMConfig) -> OpenAI:
"""构造同步 OpenAI 客户端(指向 DeepSeek/Qwen 兼容端点)。"""
logger.debug(
"初始化同步 LLM 客户端: provider={} model={} base_url={}",
config.provider, config.model, config.base_url,
)
return OpenAI(
api_key=config.api_key,
base_url=config.base_url,
timeout=config.timeout_sec,
)
def make_async_client(config: LLMConfig) -> AsyncOpenAI:
"""构造异步 OpenAI 客户端(用于批处理高并发)。"""
logger.debug(
"初始化异步 LLM 客户端: provider={} model={} base_url={}",
config.provider, config.model, config.base_url,
)
return AsyncOpenAI(
api_key=config.api_key,
base_url=config.base_url,
timeout=config.timeout_sec,
)
+307
View File
@@ -0,0 +1,307 @@
"""LLM 投资事件抽取主流程。
输入:Article(M2/M3 输出)
输出:ExtractedEvent(含 Pydantic 校验过的 EventExtraction)
设计:
1. 加载 prompts/event_extraction.md,字符串替换填入文章字段;
2. 调用 LLM JSON mode (response_format={"type":"json_object"});
3. 解析 JSON -> Pydantic EventExtraction(强校验)+ 重试;
4. 限制正文长度避免触顶 context window。
"""
from __future__ import annotations
import asyncio
import json
from pathlib import Path
from typing import Any, Protocol
from loguru import logger
from openai import AsyncOpenAI, OpenAI
from extractor import Article
from .client import LLMConfig
from .models import EventExtraction, ExtractedEvent, LLMCallError
# Prompt 模板默认路径
DEFAULT_PROMPT_PATH = Path("prompts/event_extraction.md")
# 文章正文截断长度(防止超出上下文窗口,DeepSeek/Qwen 都支持 32K+,这里保守取 8K 字符)
MAX_CONTENT_CHARS = 8000
# 重试设置
DEFAULT_MAX_ATTEMPTS = 3
RETRY_BASE_WAIT_SEC = 1.0
RETRY_MAX_WAIT_SEC = 8.0
# --------------------------------------------------------------------------- #
# Prompt 渲染
# --------------------------------------------------------------------------- #
class PromptTemplate:
"""Prompt 模板加载器,支持 {placeholder} 字符串替换。"""
def __init__(self, template_path: str | Path = DEFAULT_PROMPT_PATH) -> None:
self._path = Path(template_path)
self._template = self._path.read_text(encoding="utf-8")
def render(self, article: Article) -> str:
content = article.content
if len(content) > MAX_CONTENT_CHARS:
logger.debug(
"文章 {} 超长截断: {} -> {}",
article.url_hash, len(content), MAX_CONTENT_CHARS,
)
content = content[:MAX_CONTENT_CHARS] + "\n\n[正文过长已截断]"
publish_time_str = (
article.publish_time.strftime("%Y-%m-%d %H:%M")
if article.publish_time
else "未知"
)
return (
self._template
.replace("{title}", article.title)
.replace("{publish_time}", publish_time_str)
.replace("{source_name}", article.source_name or article.source_id)
.replace("{content}", content)
)
# --------------------------------------------------------------------------- #
# JSON 提取(LLM 偶尔会包 ```json 围栏)
# --------------------------------------------------------------------------- #
def _extract_json_object(text: str) -> str:
"""从 LLM 输出中提取首个 JSON 对象字符串(去围栏 / 取首个 {...})。"""
s = text.strip()
if s.startswith("```"):
# 去除 ```json ... ``` 围栏
s = s.strip("`")
# 可能形如 "json\n{...}"
if s.lower().startswith("json"):
s = s[4:].lstrip("\n").lstrip()
# 末尾可能还有 ```
if s.endswith("```"):
s = s[:-3]
# 取首个 { 到对应 }
start = s.find("{")
if start < 0:
return s
depth = 0
for i in range(start, len(s)):
if s[i] == "{":
depth += 1
elif s[i] == "}":
depth -= 1
if depth == 0:
return s[start : i + 1]
return s[start:]
def parse_event_json(raw: str) -> EventExtraction:
"""把 LLM 输出文本解析为 EventExtraction(可能抛 LLMCallError)。"""
payload = _extract_json_object(raw)
try:
obj = json.loads(payload)
except json.JSONDecodeError as e:
raise LLMCallError(f"JSON 解析失败: {e}") from e
if not isinstance(obj, dict):
raise LLMCallError(f"JSON 顶层非对象: {type(obj).__name__}")
try:
return EventExtraction.model_validate(obj)
except Exception as e: # noqa: BLE001 - pydantic ValidationError 等多类型
raise LLMCallError(f"事件 schema 校验失败: {e}") from e
# --------------------------------------------------------------------------- #
# Article -> ExtractedEvent
# --------------------------------------------------------------------------- #
class _SyncChat(Protocol):
def chat(self, *args: Any, **kwargs: Any) -> Any: ...
def _call_llm_sync(
client: OpenAI,
config: LLMConfig,
prompt: str,
) -> tuple[str, dict[str, int | None]]:
"""同步单次 LLM 调用,返回 (raw_text, usage)。usage 含 prompt_tokens / completion_tokens。"""
resp = client.chat.completions.create(
model=config.model,
messages=[
{
"role": "system",
"content": "你是 A 股投资研究助手,严格按用户指定的 JSON 格式输出。",
},
{"role": "user", "content": prompt},
],
temperature=config.temperature,
response_format={"type": "json_object"},
)
text = resp.choices[0].message.content or ""
usage = {
"prompt_tokens": getattr(resp.usage, "prompt_tokens", None) if resp.usage else None,
"completion_tokens": (
getattr(resp.usage, "completion_tokens", None) if resp.usage else None
),
}
return text, usage
async def _call_llm_async(
client: AsyncOpenAI,
config: LLMConfig,
prompt: str,
) -> tuple[str, dict[str, int | None]]:
"""异步单次 LLM 调用。"""
resp = await client.chat.completions.create(
model=config.model,
messages=[
{
"role": "system",
"content": "你是 A 股投资研究助手,严格按用户指定的 JSON 格式输出。",
},
{"role": "user", "content": prompt},
],
temperature=config.temperature,
response_format={"type": "json_object"},
)
text = resp.choices[0].message.content or ""
usage = {
"prompt_tokens": getattr(resp.usage, "prompt_tokens", None) if resp.usage else None,
"completion_tokens": (
getattr(resp.usage, "completion_tokens", None) if resp.usage else None
),
}
return text, usage
def extract_event(
client: OpenAI,
config: LLMConfig,
article: Article,
*,
template: PromptTemplate | None = None,
max_attempts: int | None = None,
sources: list[str] | None = None,
) -> ExtractedEvent:
"""同步抽取单篇文章的事件(带重试)。
max_attempts 为 None 时使用 config.max_attempts(来自 YAML/环境变量配置)。
sources 为该新闻全部来源(来自去重层多源记录);None 时兜底 [article.source_id]。
"""
tpl = template or PromptTemplate()
prompt = tpl.render(article)
max_attempts = max_attempts or config.max_attempts
last_err: Exception | None = None
for attempt in range(1, max_attempts + 1):
try:
raw, usage = _call_llm_sync(client, config, prompt)
event = parse_event_json(raw)
return ExtractedEvent(
source_id=article.source_id,
url=article.url,
url_hash=article.url_hash,
title=article.title,
publish_time=article.publish_time,
sources=sources or [article.source_id],
event=event,
provider=config.provider,
model=config.model,
attempts=attempt,
prompt_tokens=usage.get("prompt_tokens"),
completion_tokens=usage.get("completion_tokens"),
)
except LLMCallError as e:
last_err = e
logger.warning(
"LLM 抽取失败 url={} 尝试 {}/{}: {}",
article.url, attempt, max_attempts, e.reason,
)
except Exception as e: # noqa: BLE001 - 网络/限流等
last_err = e
logger.warning(
"LLM 调用异常 url={} 尝试 {}/{}: {}: {}",
article.url, attempt, max_attempts, type(e).__name__, e,
)
if attempt < max_attempts:
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
import time
time.sleep(wait)
raise LLMCallError(
f"LLM 抽取放弃,共 {max_attempts} 次尝试: {last_err}",
attempts=max_attempts,
)
async def extract_event_async(
client: AsyncOpenAI,
config: LLMConfig,
article: Article,
*,
template: PromptTemplate | None = None,
max_attempts: int | None = None,
sources: list[str] | None = None,
semaphore: asyncio.Semaphore | None = None,
) -> ExtractedEvent:
"""异步抽取(批处理用),与同步版逻辑等价。
max_attempts 为 None 时使用 config.max_attempts。
sources 为该新闻全部来源;None 时兜底 [article.source_id]。
"""
tpl = template or PromptTemplate()
prompt = tpl.render(article)
max_attempts = max_attempts or config.max_attempts
async def _run() -> ExtractedEvent:
last_err: Exception | None = None
for attempt in range(1, max_attempts + 1):
try:
raw, usage = await _call_llm_async(client, config, prompt)
event = parse_event_json(raw)
return ExtractedEvent(
source_id=article.source_id,
url=article.url,
url_hash=article.url_hash,
title=article.title,
publish_time=article.publish_time,
sources=sources or [article.source_id],
event=event,
provider=config.provider,
model=config.model,
attempts=attempt,
prompt_tokens=usage.get("prompt_tokens"),
completion_tokens=usage.get("completion_tokens"),
)
except LLMCallError as e:
last_err = e
logger.warning(
"LLM 抽取失败 url={} 尝试 {}/{}: {}",
article.url, attempt, max_attempts, e.reason,
)
except Exception as e: # noqa: BLE001
last_err = e
logger.warning(
"LLM 调用异常 url={} 尝试 {}/{}: {}: {}",
article.url, attempt, max_attempts, type(e).__name__, e,
)
if attempt < max_attempts:
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
await asyncio.sleep(wait)
raise LLMCallError(
f"LLM 抽取放弃,共 {max_attempts} 次尝试: {last_err}",
attempts=max_attempts,
)
if semaphore is None:
return await _run()
async with semaphore:
return await _run()
+188
View File
@@ -0,0 +1,188 @@
"""LLM 投资事件抽取的数据模型 (M4)。
EventExtraction 是 LLM 严格输出 schema(JSON mode 解析后用 Pydantic 校验)。
ExtractedEvent 把 EventExtraction 与原文章元数据合并,作为 M4 最终落盘格式。
"""
from __future__ import annotations
import re
from datetime import datetime
from enum import StrEnum
from typing import Self
from pydantic import BaseModel, Field, field_validator, model_validator
class Sentiment(StrEnum):
"""事件情绪倾向。"""
POSITIVE = "positive" # 利好
NEUTRAL = "neutral" # 中性
NEGATIVE = "negative" # 利空
# 事件类型枚举(Prompt 中也会展示给 LLM)
EVENT_TYPES: tuple[str, ...] = (
"业绩预告",
"业绩快报",
"财报披露",
"合作签约",
"投资并购",
"重大合同",
"产品发布",
"技术突破",
"监管处罚",
"诉讼仲裁",
"股东减持",
"股东增持",
"回购",
"分红",
"高管变动",
"资产重组",
"停牌复牌",
"ST警示",
"退市风险",
"宏观政策",
"行业政策",
"国际局势",
"其他",
)
# A 股股票代码:6 位数字(000xxx/300xxx/600xxx 等),也可带 .SH/.SZ/.BJ 后缀
_STOCK_CODE_RE = re.compile(r"^\d{6}(\.(SH|SZ|BJ))?$")
# 重要程度合理区间(LLM 偶尔会给 0/6/10,这里夹紧)
MIN_IMPORTANCE = 1
MAX_IMPORTANCE = 5
class EventExtraction(BaseModel):
"""LLM 输出的 JSON 直接映射到此模型。"""
stock_codes: list[str] = Field(
default_factory=list,
description="A 股 6 位代码,允许带 .SH/.SZ/.BJ 后缀;无相关股票时为空",
)
company_names: list[str] = Field(
default_factory=list, description="涉及公司中文简称,无关时为空"
)
industries: list[str] = Field(
default_factory=list, description="所属行业(申万二级粒度优先);无关时为空"
)
sentiment: Sentiment = Field(..., description="positive/neutral/negative")
importance: int = Field(
..., ge=MIN_IMPORTANCE, le=MAX_IMPORTANCE, description="1-5 重要程度"
)
event_type: str = Field(..., description="事件类型,见 EVENT_TYPES")
summary: str = Field(
default="",
max_length=200,
description="一句话事件摘要(≤ 100 字),便于人工浏览",
)
@field_validator("stock_codes")
@classmethod
def _strip_and_validate_stock_codes(cls, v: list[str]) -> list[str]:
"""剔除空字符串、统一大写、过滤明显非法格式。"""
cleaned: list[str] = []
for code in v:
s = (code or "").strip().upper().replace(" ", "")
if not s:
continue
if _STOCK_CODE_RE.match(s):
cleaned.append(s)
# 去重保持顺序
seen: set[str] = set()
out: list[str] = []
for c in cleaned:
if c not in seen:
seen.add(c)
out.append(c)
return out
@field_validator("company_names", "industries")
@classmethod
def _strip_text_lists(cls, v: list[str]) -> list[str]:
cleaned = [(s or "").strip() for s in v]
cleaned = [s for s in cleaned if s]
seen: set[str] = set()
out: list[str] = []
for c in cleaned:
if c not in seen:
seen.add(c)
out.append(c)
return out
@field_validator("event_type")
@classmethod
def _normalize_event_type(cls, v: str) -> str:
s = (v or "").strip()
if not s:
return "其他"
return s
@model_validator(mode="after")
def _post_check(self) -> Self:
"""中性情绪时 importance 应较低(1-3),纠正常见误判。"""
# 不强制纠正,留作后续校准。占位,便于以后扩展。
return self
class ExtractedEvent(BaseModel):
"""落盘格式:文章元数据 + LLM 抽取结果 + 调用元信息。
sources: 该唯一新闻的全部来源(主源 source_id 居首)。
来自去重层多源记录(M3 sources.json / uniques JSON 的 sources 字段),
旧产物无此字段时兜底为 [source_id]。
"""
# ---- 来源标识 ----
source_id: str
url: str
url_hash: str
title: str
publish_time: datetime | None = None
sources: list[str] = Field(
default_factory=list,
description="全部来源(主源居首,去重保序);旧产物无字段时兜底为 [source_id]",
)
# ---- 抽取结果 ----
event: EventExtraction
# ---- 调用元信息 ----
provider: str = Field(..., description="deepseek / qwen 等")
model: str
extracted_at: datetime = Field(default_factory=datetime.now)
attempts: int = Field(default=1, ge=1, description="LLM 实际调用次数(含重试)")
prompt_tokens: int | None = None
completion_tokens: int | None = None
@model_validator(mode="after")
def _ensure_sources(self) -> Self:
"""保证 sources 非空、去重且以主源 source_id 开头。"""
seen: list[str] = []
for s in [self.source_id, *self.sources]:
if s and s not in seen:
seen.append(s)
self.sources = seen
return self
def short_summary(self) -> str:
ev = self.event
codes = ",".join(ev.stock_codes) or "-"
return (
f"[{self.source_id}] {self.title[:30]} "
f"-> {ev.sentiment.value}/{ev.importance}/{ev.event_type} "
f"({codes})"
)
class LLMCallError(Exception):
"""LLM 调用失败(网络 / 解析 / 校验)。"""
def __init__(self, reason: str, *, attempts: int = 0) -> None:
super().__init__(reason)
self.reason = reason
self.attempts = attempts