You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为asyncio Socket服务器实现KeyboardInterrupt下的优雅关闭?

实现Asyncio服务器的优雅关闭(处理KeyboardInterrupt)

问题背景

当服务器无客户端连接时,外层的try-except可以正常捕获KeyboardInterrupt,但存在活跃客户端时,conn_handler协程会抛出未处理的KeyboardInterrupt异常——要么需要按两次Ctrl+C才能关闭,要么主进程捕获异常后仍会输出"Task exception was never retrieved"的报错信息,需要实现一次Ctrl+C即可优雅关闭服务器,同时妥善处理现有连接。

现有代码

服务器代码

import asyncio
import json

async def conn_handler(reader, writer):
            addr = writer.get_extra_info('peername')
            print(f"{addr} connected")

            while True:
                    data_len = await reader.read(2)
                    data_len = int.from_bytes(data_len, byteorder="big")

                    if data_len == 0:
                        break

                    if data := await reader.readexactly(data_len):
                        msg = json.loads(data)
                        print(f"received: {msg}")
                        
                        msg = json.dumps(msg)                        
                        msg_len = len(msg).to_bytes(2, byteorder="big")

                        writer.write(msg_len + msg.encode())
                        await writer.drain()
                        print(f"sent: {msg}")

            print(f"{addr} closed")


async def start_server():
    server = await asyncio.start_server(conn_handler, "0.0.0.0", 10999)

    print(f'Serving on {server.sockets[0].getsockname()}')

    async with server:
        await server.serve_forever()

try:
    asyncio.run(start_server())
except KeyboardInterrupt:
    print("keyboard interrupt occured")

客户端代码

import asyncio
import json
import uuid

outstanding_msgs = []
msgs_count = 0

async def send_msgs_loop(writer):
    with open("msgs.json") as f:
        msgs = json.load(f)

        global msgs_count
        msgs_count = len(msgs)

        for msg in msgs:
            msg["id"] = str(uuid.uuid4())
            outstanding_msgs.append(msg["id"])

            msg = json.dumps(msg)
            msg_len = len(msg).to_bytes(2, byteorder="big")

            writer.write(msg_len + msg.encode())
            await writer.drain()

            print(f"sent: {msg}")

async def read_msgs_loop(reader):
    received_count = 0
    while True:
        msg_len = await reader.read(2)
        msg_len = int.from_bytes(msg_len, byteorder="big")

        if msg_len == 0:
            break

        if msg := await reader.readexactly(msg_len):
            msg = json.loads(msg)
            if msg["id"] in outstanding_msgs:
                received_count += 1
                outstanding_msgs.remove(msg["id"])
                print(f"received: {msg}")

        if not outstanding_msgs and received_count == msgs_count:
            print("all responses received")
            break

async def start_client():
    reader, writer = await asyncio.open_connection("localhost", 10999)

    await asyncio.gather(send_msgs_loop(writer), read_msgs_loop(reader))

    writer.close()
    await writer.wait_closed()

if __name__ == '__main__':    
    asyncio.run(start_client(), debug=False)

msgs.json示例

[
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"},
        {"amount": 30, "card": "4444444442", "terminal": "EFG"},
        {"amount": 10, "card": "1213212312", "terminal": "ABC"},
        {"amount": 25, "card": "5555555552", "terminal": "CDE"}
]

解决方案

核心思路:通过信号触发服务器关闭,而非让异常在协程中传播,同时跟踪所有活跃连接任务,确保它们能优雅退出。具体步骤:

  1. 服务器启动后注册信号处理器,捕获SIGINT(对应Ctrl+C)
  2. 维护集合跟踪所有conn_handler任务,避免任务异常未被检索
  3. 在conn_handler中处理连接相关异常(如ConnectionResetError、asyncio.CancelledError),确保循环正常退出
  4. 收到关闭信号时,先停止接受新连接,再等待现有连接任务完成(或超时强制关闭)

修改后的服务器代码

import asyncio
import json
import signal
from typing import Set

# 跟踪所有活跃的连接处理任务
active_tasks: Set[asyncio.Task] = set()

async def conn_handler(reader, writer):
    addr = writer.get_extra_info('peername')
    print(f"{addr} connected")
    task = asyncio.current_task()
    active_tasks.add(task)
    
    try:
        while True:
            try:
                # 读取数据长度,超时或连接断开时抛出异常
                data_len = await asyncio.wait_for(reader.read(2), timeout=5.0)
                if not data_len:
                    break
                data_len = int.from_bytes(data_len, byteorder="big")
                
                if data_len == 0:
                    break
                
                data = await asyncio.wait_for(reader.readexactly(data_len), timeout=5.0)
                msg = json.loads(data)
                print(f"received: {msg}")
                
                # 构造响应
                msg = json.dumps(msg)                        
                msg_len = len(msg).to_bytes(2, byteorder="big")
                writer.write(msg_len + msg.encode())
                await writer.drain()
                print(f"sent: {msg}")
            except (asyncio.TimeoutError, ConnectionResetError, asyncio.CancelledError):
                # 处理连接超时、重置或被取消的情况
                break
            except Exception as e:
                print(f"Error handling connection {addr}: {e}")
                break
    finally:
        # 关闭连接并从任务集合中移除
        writer.close()
        await writer.wait_closed()
        active_tasks.discard(task)
        print(f"{addr} closed")

async def shutdown(server: asyncio.Server):
    print("\nStarting graceful shutdown...")
    # 停止接受新连接
    server.close()
    await server.wait_closed()
    
    # 等待所有活跃连接任务完成,最多等待10秒
    if active_tasks:
        print(f"Waiting for {len(active_tasks)} active connections to finish...")
        done, pending = await asyncio.wait(active_tasks, timeout=10.0)
        # 取消未完成的任务
        for task in pending:
            task.cancel()
            try:
                await task
            except asyncio.CancelledError:
                pass
    print("Server shutdown complete")

async def start_server():
    server = await asyncio.start_server(conn_handler, "0.0.0.0", 10999)
    print(f'Serving on {server.sockets[0].getsockname()}')
    
    # 注册信号处理器,捕获SIGINT(Ctrl+C)和SIGTERM
    loop = asyncio.get_running_loop()
    for sig in (signal.SIGINT, signal.SIGTERM):
        loop.add_signal_handler(sig, lambda: asyncio.create_task(shutdown(server)))
    
    async with server:
        await server.serve_forever()

if __name__ == "__main__":
    try:
        asyncio.run(start_server())
    except KeyboardInterrupt:
        # 兼容Windows下信号处理的延迟情况
        print("Keyboard interrupt received")

关键修改说明

  1. 任务跟踪:新增active_tasks集合,每个conn_handler启动时将自身任务加入集合,退出时移除,避免任务异常未被检索。
  2. 信号处理:注册SIGINT和SIGTERM信号处理器,触发时调用shutdown协程,优雅关闭服务器。
  3. 异常处理:在conn_handler中捕获连接相关异常,确保循环能正常退出,同时关闭客户端连接。
  4. 超时控制:给reader.read和reader.readexactly添加超时,避免协程无限阻塞。
  5. 优雅等待:关闭服务器后,等待现有连接任务完成,超时则强制取消未完成任务。

验证方式

  1. 启动修改后的服务器
  2. 运行客户端发送消息
  3. 在客户端运行过程中按一次Ctrl+C,服务器会输出关闭日志,无未处理任务异常,客户端会收到断开信号后退出。

内容的提问来源于stack exchange,提问作者coolio

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 01:19:49