Files
simon 9aa44e610d feat: LLM/Embedding 重试次数纳入 system.yaml 统一配置
- system.yaml: llm.max_retries(死配置) 改为 llm.max_attempts;embedding 段新增 max_attempts
- llm/client.py: LLMConfig 新增 max_attempts(默认 3 兜底),load_llm_config 读取配置
- llm/extractor.py: translate_and_extract(_async) max_attempts=None 时取 config.max_attempts
- embedding/client.py: load_embedding_config 读取 embedding.max_attempts
- scheduler/reporter.py: _call_llm_simple max_retries=None 时取 max_attempts-1
- 新增 2 个配置读取测试;已同步 pi5 验证
2026-08-05 08:48:43 +08:00

376 lines
13 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""LLM 翻译 + 投资事件抽取主流程。
输入: ProcessedArticleM2/M3 输出)
输出: EnTranslatedArticle(含中英文双语内容 + 抽取事件 + 元信息)
设计:
1. 加载 prompts/translation_and_extraction.md,分离 system/user 模板
2. 单次 LLM 调用完成翻译 + 事件抽取(节省 token)
3. 解析 JSON → Pydantic LLMTranslationOutput 强校验 + 重试
4. 限制正文长度避免触顶 context window
"""
import asyncio
import json
import logging
import re
import time
from pathlib import Path
from openai import AsyncOpenAI, OpenAI
from extractor.models import ProcessedArticle
from llm.client import LLMConfig
from llm.models import (
EnTranslatedArticle,
LLMCallError,
LLMTranslationOutput,
)
logger = logging.getLogger(__name__)
# Prompt 模板路径
DEFAULT_PROMPT_PATH = Path("prompts/translation_and_extraction.md")
# 文章正文截断长度(保守取 8K 字符,避免超出上下文窗口)
MAX_CONTENT_CHARS = 8000
# 重试设置
DEFAULT_MAX_ATTEMPTS = 3
RETRY_BASE_WAIT_SEC = 1.0
RETRY_MAX_WAIT_SEC = 8.0
# Prompt 中 system / user 分隔标记
_SYSTEM_SECTION_START = "## System Prompt"
_USER_SECTION_START = "## User Input"
class PromptTemplate:
"""Prompt 模板加载器。
从 prompts/translation_and_extraction.md 读取模板,
分离 System Prompt 和 User Input 两部分。
User Input 支持 {title} / {source_name} / {publish_time} / {content} 占位符。
"""
def __init__(self, template_path: str | Path = DEFAULT_PROMPT_PATH) -> None:
self._path = Path(template_path)
raw = self._path.read_text(encoding="utf-8")
# 分离 system 和 user 两部分
self._system_prompt, self._user_template = self._parse_template(raw)
@staticmethod
def _parse_template(raw: str) -> tuple[str, str]:
"""解析模板文件,返回 (system_prompt, user_template)。"""
# 找到 ## System Prompt 之后的内容直到 ## User Input
sys_match = re.search(
r"## System Prompt\s*\n(.*?)(?=---\s*\n## User Input)",
raw, re.DOTALL
)
user_match = re.search(
r"## User Input\s*\n(.*)",
raw, re.DOTALL
)
system = sys_match.group(1).strip() if sys_match else ""
user = user_match.group(1).strip() if user_match else raw
return system, user
def render(self, article: ProcessedArticle) -> tuple[str, str]:
"""渲染 Prompt,返回 (system_prompt, user_prompt)。
对正文做截断处理。
"""
content = article.content
if len(content) > MAX_CONTENT_CHARS:
logger.debug(
"文章 %s 超长截断: %d%d",
article.url_hash, len(content), MAX_CONTENT_CHARS,
)
content = content[:MAX_CONTENT_CHARS] + "\n\n[正文过长已截断]"
user_prompt = (
self._user_template
.replace("{title}", article.title)
.replace("{source_name}", article.source_name)
.replace("{publish_time}", article.publish_time or "未知")
.replace("{content}", content)
)
return self._system_prompt, user_prompt
# --------------------------------------------------------------------------- #
# JSON 提取(LLM 偶尔会包 ```json 围栏)
# --------------------------------------------------------------------------- #
def _extract_json_object(text: str) -> str:
"""从 LLM 输出中提取首个 JSON 对象字符串(去围栏 / 取首个 {...})。"""
s = text.strip()
if s.startswith("```"):
s = s.strip("`")
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_translation_json(raw: str) -> LLMTranslationOutput:
"""把 LLM 输出文本解析为 LLMTranslationOutput(可能抛 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 LLMTranslationOutput.model_validate(obj)
except Exception as e:
raise LLMCallError(f"翻译输出 schema 校验失败: {e}") from e
# --------------------------------------------------------------------------- #
# 词数计算
# --------------------------------------------------------------------------- #
def _count_zh_chars(text: str) -> int:
"""统计中文文本字数(汉字 + 数字 + 字母混合)。"""
# 简单统计:去除空白后长度
return len(text.replace(" ", "").replace("\n", "").replace("\r", ""))
# --------------------------------------------------------------------------- #
# 同步 / 异步 LLM 调用
# --------------------------------------------------------------------------- #
def _call_llm_sync(
client: OpenAI,
config: LLMConfig,
system_prompt: str,
user_prompt: str,
) -> tuple[str, dict[str, int | None]]:
"""同步单次 LLM 调用,返回 (raw_text, usage)。"""
resp = client.chat.completions.create(
model=config.model,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
temperature=config.temperature,
max_tokens=config.max_tokens,
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,
system_prompt: str,
user_prompt: str,
) -> tuple[str, dict[str, int | None]]:
"""异步单次 LLM 调用。"""
resp = await client.chat.completions.create(
model=config.model,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
temperature=config.temperature,
max_tokens=config.max_tokens,
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 translate_and_extract(
client: OpenAI,
config: LLMConfig,
article: ProcessedArticle,
*,
template: PromptTemplate | None = None,
max_attempts: int | None = None,
) -> EnTranslatedArticle:
"""同步翻译 + 事件抽取(单篇文章,带重试)。
Args:
client: OpenAI 同步客户端
config: LLM 配置
article: 待处理的英文新闻
template: Prompt 模板,默认加载 prompts/translation_and_extraction.md
max_attempts: 最大尝试次数(含首次);None 时取 config.max_attempts
(来自 system.yaml llm.max_attempts
Returns:
EnTranslatedArticle 含双语内容 + 事件
Raises:
LLMCallError: 所有重试均失败
"""
if max_attempts is None:
max_attempts = config.max_attempts
tpl = template or PromptTemplate()
system_prompt, user_prompt = tpl.render(article)
last_err: Exception | None = None
for attempt in range(1, max_attempts + 1):
try:
raw, usage = _call_llm_sync(client, config, system_prompt, user_prompt)
output = parse_translation_json(raw)
# 检查 content_zh 非空
if not output.content_zh.strip():
raise LLMCallError("LLM 返回的 content_zh 为空")
return EnTranslatedArticle(
source_id=article.source_id,
source_name=article.source_name,
url=article.url,
url_hash=article.url_hash,
title=article.title,
title_zh=output.title_zh,
content_en=article.content,
content_zh=output.content_zh,
publish_time=article.publish_time,
word_count=article.word_count,
word_count_zh=_count_zh_chars(output.content_zh),
events=output.events,
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=%s 尝试 %d/%d: %s",
article.url, attempt, max_attempts, e.reason,
)
except Exception as e:
last_err = e
logger.warning(
"LLM 调用异常 url=%s 尝试 %d/%d: %s: %s",
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)
time.sleep(wait)
raise LLMCallError(
f"LLM 翻译抽取放弃,共 {max_attempts} 次尝试: {last_err}",
attempts=max_attempts,
)
async def translate_and_extract_async(
client: AsyncOpenAI,
config: LLMConfig,
article: ProcessedArticle,
*,
template: PromptTemplate | None = None,
max_attempts: int | None = None,
semaphore: asyncio.Semaphore | None = None,
) -> EnTranslatedArticle:
"""异步翻译 + 事件抽取(批处理用),与同步版逻辑等价。"""
if max_attempts is None:
max_attempts = config.max_attempts
tpl = template or PromptTemplate()
system_prompt, user_prompt = tpl.render(article)
async def _run() -> EnTranslatedArticle:
last_err: Exception | None = None
for attempt in range(1, max_attempts + 1):
try:
raw, usage = await _call_llm_async(client, config, system_prompt, user_prompt)
output = parse_translation_json(raw)
if not output.content_zh.strip():
raise LLMCallError("LLM 返回的 content_zh 为空")
return EnTranslatedArticle(
source_id=article.source_id,
source_name=article.source_name,
url=article.url,
url_hash=article.url_hash,
title=article.title,
title_zh=output.title_zh,
content_en=article.content,
content_zh=output.content_zh,
publish_time=article.publish_time,
word_count=article.word_count,
word_count_zh=_count_zh_chars(output.content_zh),
events=output.events,
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=%s 尝试 %d/%d: %s",
article.url, attempt, max_attempts, e.reason,
)
except Exception as e:
last_err = e
logger.warning(
"LLM 调用异常 url=%s 尝试 %d/%d: %s: %s",
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()