跳转至

优雅终止

在生产环境中,服务需要在收到终止信号时优雅地完成正在处理的请求,而不是直接中断。

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('新消息比老消息更先写入数据库')
    # 不影响正确性,后排部队转作先头部队

整体流程图

用户发消息A → 写入消息A + lock=uuid_a → commit → 开始生成回复A
用户发消息B → 写入消息B + lock=uuid_b → commit → 开始生成回复B
回复A生成中 → check_should_stop() → 发现 lock≠uuid_a → 中断回复A
回复B生成中 → 读取消息A+B作为完整上下文 → 生成基于A+B的回复