diff --git a/endpoint/tcp_server.py b/endpoint/tcp_server.py index 62c1828..acf96c4 100644 --- a/endpoint/tcp_server.py +++ b/endpoint/tcp_server.py @@ -1,161 +1,84 @@ -import socket +import asyncio import signal -import threading -from concurrent.futures import ThreadPoolExecutor - +import platform +from config.logger_config import logger from config.setting import settings from model.index import UNKNOWN_MESSAGE from processor.tcp_processor import TcpMessageProcessor -from config.logger_config import logger class TCPServer: def __init__(self): self.host = settings.host self.port = settings.port - - self.max_workers = settings.max_workers - self.timeout = settings.timeout - - self.server_socket = None + self.server = None + self.clients = set() self.running = False - self.thread_pool = ThreadPoolExecutor(max_workers=self.max_workers) self.message_processor = TcpMessageProcessor() - # 已经连接的客户端 - self.clients = {} - self.client_lock = threading.Lock() - - self.server_thread = None - - signal.signal(signal.SIGTERM, self._handle_signal) - signal.signal(signal.SIGINT, self._handle_signal) - - def _handle_signal(self, signum, frame): - """处理终止信号,触发优雅关闭""" - logger.info(f"收到信号 {signum},准备关闭服务器...") - self.running = False - - def _start_loop(self): - """启动服务器""" - try: - self.server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - self.server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) - self.server_socket.bind((self.host, self.port)) - self.server_socket.listen(5) - self.server_socket.settimeout(1.0) - self.running = True - - logger.info(f"TCP服务器启动,监听 {self.host}:{self.port} " - f"(最大线程: {self.max_workers}, 超时: {self.timeout}s)") - - while self.running: - try: - client_socket, client_address = self.server_socket.accept() - client_ip = client_address[0] - client_socket.settimeout(self.timeout) - logger.info(f"新连接: {client_address}") - - # 存储客户端 - with self.client_lock: - self.clients[client_ip] = client_socket - - # 提交到线程池处理 - self.thread_pool.submit(self.handle_client, client_socket, client_ip) - except socket.timeout: - continue - except Exception as e: - if self.running: - logger.error(f"接受连接失败: {str(e)}") - except Exception as e: - logger.error(f"TCP服务器启动失败: {str(e)}") - finally: - self.stop() - - def start(self): - self.server_thread = threading.Thread(target=self._start_loop, daemon=True) - self.server_thread.start() - - def stop(self): - if not self.running: - return - - self.running = False - logger.info("开始关闭服务器...") - - # 移除客户端 - with self.client_lock: - for client_socket in self.clients.values(): - try: - client_socket.close() - except Exception as e: - logger.warning(f"关闭客户端连接失败: {str(e)}") - self.clients.clear() - - # 关闭线程池 - self.thread_pool.shutdown(wait=True) - logger.info("所有客户端处理线程已结束") - - # 关闭连接 - if self.server_socket: - self.server_socket.close() - logger.info(f"服务器已关闭({self.host}:{self.port})") - - def handle_client(self, client_socket, client_ip): + async def handle_client(self, reader, writer): """处理客户端连接""" + addr = writer.get_extra_info('peername') + logger.info(f"新客户端连接: {addr}") + self.clients.add(writer) + try: - while True: - data = client_socket.recv(1024) + while self.running: + data = await reader.read(1024) if not data: - logger.info(f"客户端 {client_ip} 主动断开连接") - break + break # 连接断开 message = data.decode('utf-8').strip() - logger.info(f"收到 {client_ip} 的消息: {message}") + logger.info(f"收到 {addr} 的消息: {message}") # 放到消息处理器里面处理 response = self.message_processor.process(message) # 回复消息 if response != UNKNOWN_MESSAGE: - client_socket.sendall(response.encode('utf-8')) - logger.info(f"回复 {client_ip}: {response}") - except socket.timeout: - logger.warning(f"客户端 {client_ip} 超时未活动") + writer.write(response.encode()) + await writer.drain() + logger.info(f"回复 {addr}: {response}") except Exception as e: - logger.error(f"处理 {client_ip} 出错: {str(e)}") + logger.error(f"客户端 {addr} 处理错误: {e}") finally: - # 异常情况下关闭连接 - with self.client_lock: - if client_ip in self.clients: - del self.clients[client_ip] + if writer in self.clients: + self.clients.remove(writer) + writer.close() + await writer.wait_closed() + logger.warning(f"客户端 {addr} 断开") - try: - client_socket.close() - logger.info(f"客户端 {client_ip} 连接已关闭") - except Exception as e: - logger.warning(f"关闭 {client_ip} 连接失败: {str(e)}") + async def start(self): + """启动服务器""" + self.running = True + self.server = await asyncio.start_server(self.handle_client, self.host, self.port) - def send_to_client(self, client_ip, message): - # 先获取客户端连接(加锁保护) - with self.client_lock: - client_socket = self.clients.get(client_ip) - if not client_socket: - logger.warning(f"客户端 {client_ip} 不存在或已断开连接") - return False + # 仅在非Windows平台设置信号处理 + if platform.system() != 'Windows': + loop = asyncio.get_running_loop() + for sig in (signal.SIGINT, signal.SIGTERM): + loop.add_signal_handler(sig, self.stop) - # 发送消息 - try: - client_socket.sendall(message.encode('utf-8')) - logger.info(f"主动发送消息给 {client_ip}: {message}") - return True - except Exception as e: - logger.error(f"向 {client_ip} 发送消息失败: {str(e)}") - # 发送失败时移除无效连接 - with self.client_lock: - if client_ip in self.clients: - del self.clients[client_ip] - return False + logger.info(f"TCP服务器已启动 {self.host}:{self.port}") + async with self.server: + await self.server.serve_forever() + + async def _safe_stop(self): + """安全的停止服务器""" + self.running = False + if self.server: + self.server.close() + await self.server.wait_closed() + + # 关闭所有客户端连接 + for writer in list(self.clients): + writer.close() + await writer.wait_closed() + + def stop(self): + """停止服务器""" + logger.info("正在关闭TCP服务器...") + asyncio.create_task(self._safe_stop()) tcp_server = TCPServer() diff --git a/main.py b/main.py index bf7a5c1..4286fb7 100644 --- a/main.py +++ b/main.py @@ -1,19 +1,16 @@ -from endpoint.tcp_server import tcp_server +import asyncio +from endpoint.tcp_server import TCPServer from config.logger_config import logger -import time if __name__ == "__main__": + server = TCPServer() try: - tcp_server.start() - - try: - while True: - time.sleep(1) - except KeyboardInterrupt: - logger.info("收到终止信号,开始关闭程序...") - finally: - tcp_server.stop() - logger.info("程序已退出") + asyncio.run(server.start()) + except KeyboardInterrupt: + logger.info("收到中断信号,正在关闭服务器...") + server.stop() except Exception as e: - logger.critical(f"程序启动失败: {str(e)}", exc_info=True) - exit(1) + logger.error(f"服务器运行错误: {e}") + server.stop() + finally: + logger.info("服务器进程结束") \ No newline at end of file