Initial commit
This commit is contained in:
+169
@@ -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
|
||||
Reference in New Issue
Block a user