优雅终止
在生产环境中,服务需要在收到终止信号时优雅地完成正在处理的请求,而不是直接中断。
GracefulRunner 模式
核心思路:捕获 SIGINT/SIGTERM → 设置标志位 → 转发信号给 uvicorn → 业务代码轮询标志位来中断长任务。
from loguru import logger
import signal
import os
import time
class GracefulRunner:
def __init__(self):
self._keep_running = True
self.original_sigint_handler = None
self.original_sigterm_handler = None
def __bool__(self):
"""支持 `if not runner:` 这种优雅写法"""
return self._keep_running
def setup_signal_handlers(self):
self.original_sigint_handler = signal.signal(signal.SIGINT, self._handle_signal)
self.original_sigterm_handler = signal.signal(signal.SIGTERM, self._handle_signal)
def _handle_signal(self, signum, frame):
if not self._keep_running:
return
logger.info(f"收到信号 {signum},标记停止运行...")
self._keep_running = False
import random
time.sleep(random.random()) # 避免多worker同时抢占文件描述符
# 转发信号给 uvicorn,让它正常走关闭流程
if signum == signal.SIGINT and self.original_sigint_handler:
if callable(self.original_sigint_handler):
self.original_sigint_handler(signum, frame)
else:
signal.signal(signal.SIGINT, signal.SIG_DFL)
os.kill(os.getpid(), signal.SIGINT)
elif signum == signal.SIGTERM and self.original_sigterm_handler:
if callable(self.original_sigterm_handler):
self.original_sigterm_handler(signum, frame)
else:
signal.signal(signal.SIGTERM, signal.SIG_DFL)
os.kill(os.getpid(), signal.SIGTERM)
runner = GracefulRunner()
在 lifespan 中注册
@asynccontextmanager
async def lifespan(app: FastAPI):
runner.setup_signal_handlers()
yield
logger.success('服务已优雅关闭')
在流式响应/长任务中检查
关键点:任何长时间运行的循环都应该检查 runner 状态。
from .utils.runner import runner
# 简单用法:在循环中检查
async for chunk in stream:
if not runner:
logger.warning("收到终止信号,正在优雅退出...")
break
yield chunk
StreamController:流式响应的中断控制
对于流式 API,通常需要检查多种中断条件(不仅是后端终止,还有用户发新消息等)。可以封装成控制器:
from dataclasses import dataclass, field
from enum import Enum, auto
class StopReason(Enum):
BACKEND_SHUTDOWN = auto() # 后端终止信号
NEW_MESSAGE = auto() # 用户发送新消息
@dataclass
class StopSignal:
reason: StopReason
data: str # 要返回给客户端的数据
@dataclass
class StreamController:
session: AsyncSession
session_id: str
lock: Any
check_interval: float = 5.0
_last_check_time: float = field(default=0.0, init=False)
async def check_should_stop(self) -> StopSignal | None:
# 1. 后端终止 — 立即检查
if not runner:
return StopSignal(reason=StopReason.BACKEND_SHUTDOWN, data=...)
# 2. 其他条件 — 按间隔检查(避免频繁查库)
now = time.time()
if now - self._last_check_time < self.check_interval:
return None
self._last_check_time = now
# 例:检查是否有新消息覆盖了当前回复
await self.session.refresh(self._ticket)
if self._ticket.lock != self.lock:
return StopSignal(reason=StopReason.NEW_MESSAGE, data=...)
return None
使用方式:
controller = StreamController(session, session_id, lock)
async for chunk in stream:
if stop := await controller.check_should_stop():
yield stop.data
await stream.aclose()
return
yield process(chunk)
乐观锁:允许用户连发消息
在对话场景中,用户可以连续发送多条消息。新消息到达时,正在生成的旧回复应该被中断,让 AI 基于完整上下文重新回复。通过乐观锁实现这个机制。
核心设计
在 Session 表上放一个 lock 字段(UUID),每次收到新消息时更新它。正在运行的流式任务定期检查这个值是否变了,变了就说明有新消息进来,应该中断当前生成。
class SessionTable(SQLModel, table=True):
session_id: str = Field(primary_key=True)
lock: UUID | None = Field(default=None) # 乐观锁:每条新消息刷新
# ...
新消息到达时:写入消息 + 刷新锁
关键:消息写入和锁更新必须在同一次 commit 中,保证原子性。
@router.post("/send")
async def send_message(request: ReceivedMessage, session: ...):
lock = uuid4() # 生成新锁
# 获取 session,更新锁
ticket = await session.get(SessionTable, session_id)
ticket.lock = lock
session.add(ticket)
# 同时写入新消息
db_message = MessageTable.model_validate(request)
session.add(db_message)
# 一次 commit,保证锁和消息同时生效
await session.commit()
# 后续用这个 lock 创建 StreamController
controller = StreamController(session, session_id, lock)
流式生成中:检查锁是否被覆盖
StreamController 定期从数据库 refresh session 记录,对比 lock 值:
# 在 StreamController.check_should_stop() 中:
await self.session.refresh(self._ticket)
if self._ticket.lock != self.lock:
# lock 变了 → 有新消息进来了 → 中断当前生成
return StopSignal(reason=StopReason.NEW_MESSAGE, data=...)
连发消息时的消息合并
用户连发多条消息后,最后一个请求拿到锁,它需要把之前所有未回复的消息一起作为输入:
# 查询该 session 下所有有效消息
stmt = (
select(MessageTable)
.where(
MessageTable.session_id == session_id,
MessageTable.send_status.in_([
MessageStatus.SUCCESS,
MessageStatus.PROCESSING,
MessageStatus.INITIALIZING
]),
)
.order_by(MessageTable.created_at)
)
db_messages = (await session.exec(stmt)).all()
# 从最后一个 checkpoint 开始,收集所有新消息作为输入
new_messages = get_new_messages(db_messages)
时序保护
极端情况下,新消息可能比老消息先写入数据库。检测到这种情况时不报错,而是容忍——因为锁机制保证了只有最新的请求会继续执行:
if db_messages[-1].message_id != db_message.message_id:
logger.warning('新消息比老消息更先写入数据库')
# 不影响正确性,后排部队转作先头部队