Initial commit

This commit is contained in:
2026-07-18 15:51:01 +08:00
commit f2c80c5a9c
799 changed files with 133475 additions and 0 deletions
+169
View File
@@ -0,0 +1,169 @@
"""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 抽取结果 + 调用元信息。"""
# ---- 来源标识 ----
source_id: str
url: str
url_hash: str
title: str
publish_time: datetime | None = None
# ---- 抽取结果 ----
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
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