Files
daily-paper/app/services/claude_backend.py
T

88 lines
2.7 KiB
Python
Raw 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.
"""Claude CLI 后端 — 调用 claude CLI 子进程生成总结。
和 pi_client.py 对称的接口,复用 prompt 构建、PDF 文本提取、JSON 提取逻辑。
"""
from __future__ import annotations
import asyncio
import logging
import uuid
from app.config import settings
logger = logging.getLogger(__name__)
class ClaudeTimeoutError(Exception):
pass
class ClaudeProcessError(Exception):
def __init__(self, returncode: int, stderr: str):
self.returncode = returncode
self.stderr = stderr
super().__init__(f"claude exited with code {returncode}: {stderr[:500]}")
async def call_claude(
prompt: str,
session_id: str | None = None,
fix_errors: list[str] | None = None,
) -> tuple[str, str]:
"""调用 claude CLI print 模式,返回 (stdout 文本, session_id)。
和 call_pi() 对称的接口,但 claude CLI 不需要文件路径和 pdf_mode——
所有内容已在 prompt 中准备好。
Args:
prompt: 完整的 prompt 文本
session_id: session ID(首次为 None 时自动生成)
fix_errors: 上一轮验证错误列表(用于重试)
"""
if session_id is None:
# claude CLI 的 --session-id 要求合法 UUID;非 UUID 会被拒绝(Invalid session ID
session_id = str(uuid.uuid4())
cmd = [settings.CLAUDE_BIN, "-p", "--output-format", "text"]
if fix_errors and session_id:
# 重试:延续 session
cmd += ["--session-id", session_id, "--continue"]
else:
cmd += ["--session-id", session_id]
logger.info(
"Calling claude (session=%s, fix=%s)",
session_id,
bool(fix_errors),
)
# prompt 走 stdin,避免长文本(论文 PDF 动辄数十 KB)超出系统 ARG_MAX
proc = await asyncio.create_subprocess_exec(
*cmd,
stdin=asyncio.subprocess.PIPE,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
)
try:
stdout, stderr = await asyncio.wait_for(
proc.communicate(input=prompt.encode("utf-8")),
timeout=settings.SUMMARY_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
proc.kill()
await proc.wait()
raise ClaudeTimeoutError(
f"claude timed out after {settings.SUMMARY_TIMEOUT_SECONDS}s"
)
if proc.returncode != 0:
# claude CLI 把 API 错误输出到 stdout(stderr 常为空),优先用有内容的
detail = stderr.decode("utf-8", errors="replace")
if not detail.strip():
detail = stdout.decode("utf-8", errors="replace")
raise ClaudeProcessError(proc.returncode, detail)
return stdout.decode("utf-8", errors="replace"), session_id