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:
@@ -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
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user