第9讲:网络层与客户端
前面八讲我们实现了一个功能完整的单机数据库引擎——存储、索引、事务、SQL解析、查询执行、优化器都有了。但到目前为止,MiniDB 只能通过 Python API 本地调用。
这一讲,我们要给 MiniDB 加上网络层,让它变成一个真正的数据库服务器——支持远程连接、并发客户端、交互式查询。
一、整体架构
┌──────────────┐ TCP/IP ┌──────────────────┐ │ Client App │ ◄────────────► │ MiniDB Server │ │ (psql-like) │ │ │ └──────────────┘ │ Connection Pool │ │ Session Manager │ │ Query Processor │ │ Storage Engine │ └──────────────────┘1.1 通信协议设计
MiniDB 使用简单的文本协议,类似 PostgreSQL 的简化版:
客户端 → 服务器: SQL 语句字符串 + "\n" 服务器 → 客户端: JSON 格式的结果 成功响应: {"status": "success", "columns": ["id", "name"], "rows": [[1, "Alice"], [2, "Bob"]], "affected": 0} 错误响应: {"status": "error", "message": "语法错误: ..."}二、协议实现
2.1 消息定义
# network/protocol.py import json import struct from typing import List, Optional, Any from dataclasses import dataclass, asdict from enum import Enum class MessageType(Enum): QUERY = 'Q' # 查询请求 PREPARE = 'P' # 准备语句 EXECUTE = 'E' # 执行准备好的语句 CLOSE = 'C' # 关闭 TERMINATE = 'X' # 终止连接 @dataclass class QueryRequest: """查询请求""" type: str = 'Q' sql: str = '' @classmethod def from_bytes(cls, data: bytes) -> 'QueryRequest': return cls(type='Q', sql=data.decode('utf-8').strip()) @dataclass class QueryResult: """查询结果""" status: str = 'success' # success | error columns: List[str] = None # 列名 rows: List[List[Any]] = None # 数据行 affected: int = 0 # 影响的行数 message: str = '' # 错误消息 def to_json(self) -> str: return json.dumps(asdict(self)) + '\n' class ProtocolHandler: """协议处理器""" CHUNK_SIZE = 4096 @staticmethod def encode_result(result: QueryResult) -> bytes: """编码结果为字节流""" return result.to_json().encode('utf-8') @staticmethod def decode_request(data: bytes) -> QueryRequest: """解码请求""" return QueryRequest.from_bytes(data) @staticmethod def read_message(socket) -> Optional[str]: """从socket读取一条完整消息""" chunks = [] while True: chunk = socket.recv(ProtocolHandler.CHUNK_SIZE) if not chunk: return None chunks.append(chunk) # 检查是否收到完整的SQL语句(以换行结尾) if b'\n' in chunk: break data = b''.join(chunks) return data.decode('utf-8').strip()三、会话管理
3.1 会话状态
# network/session.py import uuid import threading from datetime import datetime from typing import Dict, Optional class Session: """客户端会话""" def __init__(self, session_id: str = None): self.session_id = session_id or str(uuid.uuid4()) self.created_at = datetime.now() self.last_active = datetime.now() self.transaction_id: Optional[int] = None self.prepared_statements: Dict[str, str] = {} # name -> sql self.settings: Dict[str, str] = { 'search_path': 'public', 'timezone': 'UTC', } def touch(self): """更新活动时间""" self.last_active = datetime.now() def prepare_statement(self, name: str, sql: str): """准备语句""" self.prepared_statements[name] = sql def get_prepared_statement(self, name: str) -> Optional[str]: """获取已准备的语句""" return self.prepared_statements.get(name) def close(self): """关闭会话""" self.prepared_statements.clear() class SessionManager: """会话管理器""" def __init__(self, session_timeout: int = 3600): self.sessions: Dict[str, Session] = {} self.timeout = session_timeout self.lock = threading.Lock() # 启动清理线程 self._start_cleanup_thread() def create_session(self) -> Session: """创建新会话""" with self.lock: session = Session() self.sessions[session.session_id] = session return session def get_session(self, session_id: str) -> Optional[Session]: """获取会话""" with self.lock: session = self.sessions.get(session_id) if session: session.touch() return session def remove_session(self, session_id: str): """移除会话""" with self.lock: if session_id in self.sessions: self.sessions[session_id].close() del self.sessions[session_id] def cleanup_expired(self): """清理过期会话""" now = datetime.now() with self.lock: expired = [ sid for sid, sess in self.sessions.items() if (now - sess.last_active).seconds > self.timeout ] for sid in expired: self.sessions[sid].close() del self.sessions[sid] def _start_cleanup_thread(self): """启动清理线程""" def cleanup_loop(): import time while True: time.sleep(300) # 每5分钟清理一次 self.cleanup_expired() thread = threading.Thread(target=cleanup_loop, daemon=True) thread.start()四、连接池
# network/pool.py import socket import threading from queue import Queue, Empty from typing import Optional class Connection: """数据库连接""" def __init__(self, host: str, port: int): self.host = host self.port = port self.socket: Optional[socket.socket] = None self.in_use = False def connect(self): """建立连接""" self.socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) self.socket.connect((self.host, self.port)) self.in_use = False def send_query(self, sql: str) -> str: """发送查询""" if not self.socket: raise ConnectionError("Connection not established") # 发送SQL self.socket.sendall((sql + '\n').encode('utf-8')) # 接收结果 from .protocol import ProtocolHandler result = ProtocolHandler.read_message(self.socket) return result def close(self): """关闭连接""" if self.socket: try: self.socket.close() except: pass self.socket = None self.in_use = False class ConnectionPool: """连接池""" def __init__(self, host: str, port: int, min_size: int = 5, max_size: int = 20): self.host = host self.port = port self.min_size = min_size self.max_size = max_size self._pool: Queue = Queue() self._active_count = 0 self.lock = threading.Lock() # 初始化最小连接数 self._initialize_pool() def _initialize_pool(self): """初始化连接池""" for _ in range(self.min_size): conn = self._create_connection() self._pool.put(conn) def _create_connection(self) -> Connection: """创建新连接""" conn = Connection(self.host, self.port) conn.connect() with self.lock: self._active_count += 1 return conn def get_connection(self, timeout: int = 30) -> Connection: """获取连接""" try: conn = self._pool.get(timeout=timeout) conn.in_use = True return conn except Empty: # 如果没有可用连接,创建新连接(不超过最大限制) with self.lock: if self._active_count < self.max_size: conn = self._create_connection() conn.in_use = True return conn else: raise TimeoutError("No available connections") def return_connection(self, conn: Connection): """归还连接""" conn.in_use = False self._pool.put(conn) def close_all(self): """关闭所有连接""" while not self._pool.empty(): try: conn = self._pool.get_nowait() conn.close() except Empty: break五、数据库服务器
5.1 服务器实现
# network/server.py import socket import threading import signal import sys from typing import Optional from .protocol import ProtocolHandler, QueryRequest, QueryResult from .session import SessionManager class MiniDBServer: """MiniDB 数据库服务器""" def __init__(self, host: str = 'localhost', port: int = 54321): self.host = host self.port = port self.server_socket: Optional[socket.socket] = None self.running = False self.session_manager = SessionManager() # 数据库引擎(由外部注入) self.engine = None self.parser = None self.executor = None # 客户端连接处理线程 self.client_threads = [] def initialize(self, engine, parser, executor): """初始化数据库引擎""" self.engine = engine self.parser = parser self.executor = executor def start(self): """启动服务器""" 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(128) self.running = True print(f"🚀 MiniDB Server started on {self.host}:{self.port}") print(f" Press Ctrl+C to stop") # 注册信号处理器 signal.signal(signal.SIGINT, self._signal_handler) try: while self.running: client_socket, address = self.server_socket.accept() print(f"📡 New connection from {address}") # 为每个客户端创建独立线程 thread = threading.Thread( target=self._handle_client, args=(client_socket, address), daemon=True ) thread.start() self.client_threads.append(thread) except KeyboardInterrupt: self.stop() def stop(self): """停止服务器""" print("\n🛑 Stopping server...") self.running = False if self.server_socket: self.server_socket.close() # 等待所有客户端线程结束 for thread in self.client_threads: thread.join(timeout=5) print("✅ Server stopped") def _handle_client(self, client_socket: socket.socket, address: tuple): """处理单个客户端连接""" session = self.session_manager.create_session() try: while self.running: # 读取客户端消息 message = ProtocolHandler.read_message(client_socket) if message is None: break # 客户端断开连接 # 处理查询 result = self._process_query(message, session) # 发送结果 response = ProtocolHandler.encode_result(result) client_socket.sendall(response) except Exception as e: print(f"❌ Error handling client {address}: {e}") finally: self.session_manager.remove_session(session.session_id) client_socket.close() print(f"📡 Connection closed from {address}") def _process_query(self, sql: str, session) -> QueryResult: """处理SQL查询""" try: # 1. 词法分析 from sql.lexer import Lexer lexer = Lexer(sql) tokens = lexer.tokenize() # 2. 语法分析 from sql.parser import Parser parser = Parser(tokens) ast = parser.parse() # 3. 查询优化 if hasattr(self, 'optimizer'): ast = self.optimizer.optimize(ast) # 4. 执行查询 result = self.executor.execute(ast) # 5. 格式化结果 return QueryResult( status='success', columns=result.get('columns', []), rows=result.get('rows', []), affected=result.get('affected', 0) ) except Exception as e: return QueryResult( status='error', message=str(e) ) def _signal_handler(self, sig, frame): """信号处理器""" self.stop() sys.exit(0)5.2 命令行客户端
# network/client.py import socket import sys import readline # 提供命令行编辑功能 from typing import Optional from .protocol import ProtocolHandler, QueryResult from .pool import ConnectionPool class MiniDBClient: """MiniDB 命令行客户端""" def __init__(self, host: str = 'localhost', port: int = 54321): self.host = host self.port = port self.socket: Optional[socket.socket] = None self.connected = False def connect(self): """连接到服务器""" try: self.socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) self.socket.connect((self.host, self.port)) self.connected = True print(f"🔗 Connected to MiniDB at {self.host}:{self.port}") return True except Exception as e: print(f"❌ Connection failed: {e}") return False def disconnect(self): """断开连接""" if self.socket: self.socket.close() self.connected = False print("🔗 Disconnected") def execute(self, sql: str) -> Optional[QueryResult]: """执行SQL""" if not self.connected: print("❌ Not connected") return None try: # 发送SQL self.socket.sendall((sql + '\n').encode('utf-8')) # 接收结果 result_str = ProtocolHandler.read_message(self.socket) if result_str is None: print("❌ Connection lost") self.connected = False return None # 解析结果 import json result_dict = json.loads(result_str) return QueryResult(**result_dict) except Exception as e: print(f"❌ Error: {e}") return None def run_interactive(self): """运行交互式模式""" if not self.connect(): return print("MiniDB Interactive Shell") print("Type 'exit' or 'quit' to quit") print("Type 'help' for help") print() history_file = '.minidb_history' try: readline.read_history_file(history_file) except FileNotFoundError: pass while True: try: sql = input('minidb> ').strip() if sql.lower() in ('exit', 'quit'): break if sql.lower() == 'help': self._print_help() continue if not sql: continue # 支持多行输入 while not sql.endswith(';'): line = input(' -> ').strip() sql += ' ' + line result = self.execute(sql) if result: self._display_result(result) except KeyboardInterrupt: print() continue except EOFError: break # 保存历史记录 readline.write_history_file(history_file) self.disconnect() def _display_result(self, result: QueryResult): """显示查询结果""" if result.status == 'error': print(f"❌ {result.message}") return if result.columns: # 打印列头 headers = ' | '.join(result.columns) separator = '-' * len(headers) print(headers) print(separator) # 打印数据行 if result.rows: for row in result.rows: formatted = ' | '.join(str(v) if v is not None else 'NULL' for v in row) print(formatted) print(f"\n({len(result.rows) if result.rows else 0} rows)") if result.affected > 0: print(f"Affected rows: {result.affected}") def _print_help(self): """打印帮助信息""" print(""" MiniDB Commands: SQL statements ending with ';' exit, quit - Exit the shell help - Show this help Example SQL: CREATE TABLE users (id INT, name VARCHAR(100), age INT); INSERT INTO users VALUES (1, 'Alice', 30); SELECT * FROM users; SELECT name, age FROM users WHERE age > 25; UPDATE users SET age = 31 WHERE id = 1; DELETE FROM users WHERE id = 1; DROP TABLE users; """) def main(): """主入口""" import argparse parser = argparse.ArgumentParser(description='MiniDB Client') parser.add_argument('-H', '--host', default='localhost', help='Server host') parser.add_argument('-P', '--port', type=int, default=54321, help='Server port') parser.add_argument('-c', '--command', help='Execute single command and exit') args = parser.parse_args() client = MiniDBClient(args.host, args.port) if args.command: # 执行单条命令模式 if client.connect(): result = client.execute(args.command) if result: client._display_result(result) client.disconnect() else: # 交互式模式 client.run_interactive() if __name__ == '__main__': main()六、集成测试
6.1 端到端测试
# tests/test_network.py import unittest import threading import time import socket import json from network.server import MiniDBServer from network.protocol import ProtocolHandler, QueryResult class TestNetworkLayer(unittest.TestCase): """网络层测试""" @classmethod def setUpClass(cls): """启动测试服务器""" cls.server = MiniDBServer('localhost', 15432) # 模拟数据库引擎 class MockEngine: def execute(self, sql): return { 'columns': ['id', 'name'], 'rows': [[1, 'Alice'], [2, 'Bob']], 'affected': 0 } cls.server.initialize(MockEngine(), None, None) # 在单独线程启动服务器 cls.server_thread = threading.Thread(target=cls.server.start, daemon=True) cls.server_thread.start() time.sleep(0.5) # 等待服务器启动 @classmethod def tearDownClass(cls): """停止测试服务器""" cls.server.stop() def test_connect_and_query(self): """测试连接和查询""" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) sock.connect(('localhost', 15432)) # 发送查询 sock.sendall(b'SELECT * FROM users;\n') # 接收结果 result_str = ProtocolHandler.read_message(sock) self.assertIsNotNone(result_str) result = json.loads(result_str) self.assertEqual(result['status'], 'success') self.assertIn('columns', result) self.assertIn('rows', result) sock.close() def test_multiple_queries(self): """测试多次查询""" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) sock.connect(('localhost', 15432)) for i in range(5): sock.sendall(f'SELECT {i};\n'.encode()) result_str = ProtocolHandler.read_message(sock) self.assertIsNotNone(result_str) sock.close() def test_concurrent_clients(self): """测试并发客户端""" def client_task(): sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) sock.connect(('localhost', 15432)) for _ in range(3): sock.sendall(b'SELECT 1;\n') result = ProtocolHandler.read_message(sock) self.assertIsNotNone(result) sock.close() threads = [] for _ in range(10): t = threading.Thread(target=client_task) threads.append(t) t.start() for t in threads: t.join() def test_invalid_sql(self): """测试无效SQL""" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) sock.connect(('localhost', 15432)) sock.sendall(b'INVALID SQL;\n') result_str = ProtocolHandler.read_message(sock) result = json.loads(result_str) self.assertEqual(result['status'], 'error') sock.close() def run_integration_test(): """运行集成测试""" print("=" * 60) print("🌐 网络层集成测试") print("=" * 60) # 1. 启动服务器 print("\n📡 启动服务器...") server = MiniDBServer('localhost', 15433) class SimpleEngine: def execute(self, sql): if sql.startswith('SELECT'): return { 'columns': ['result'], 'rows': [[f"Executed: {sql[:20]}..."]], 'affected': 0 } return {'columns': [], 'rows': [], 'affected': 1} server.initialize(SimpleEngine(), None, None) server_thread = threading.Thread(target=server.start, daemon=True) server_thread.start() time.sleep(0.5) # 2. 启动客户端 print("\n💻 启动客户端...") from network.client import MiniDBClient client = MiniDBClient('localhost', 15433) if client.connect(): # 执行几个查询 queries = [ "SELECT * FROM users;", "CREATE TABLE test (id INT);", "INSERT INTO test VALUES (1);", ] for sql in queries: print(f"\n📝 SQL: {sql}") result = client.execute(sql) if result: client._display_result(result) client.disconnect() # 3. 停止服务器 print("\n🛑 停止服务器...") server.stop() print("\n✅ 集成测试完成") if __name__ == '__main__': # 运行单元测试 unittest.main(argv=['first-arg-is-ignored'], exit=False) # 运行集成测试 run_integration_test()七、总结
这一讲给 MiniDB 加上了网络层:
通信协议:简单高效的文本协议,JSON 格式传输结果
会话管理:每个客户端独立会话,支持准备语句和设置
连接池:复用数据库连接,提高并发性能
数据库服务器:多线程处理客户端请求
命令行客户端:交互式 shell,支持历史记录
现在 MiniDB 是一个真正的数据库系统了:
# 启动服务器 $ python -m minidb.server --port 54321 # 另一个终端,连接并使用 $ python -m minidb.client -H localhost -P 54321 minidb> CREATE TABLE users (id INT, name VARCHAR(100), age INT); minidb> INSERT INTO users VALUES (1, 'Alice', 30); minidb> SELECT * FROM users; id | name | age ----------------- 1 | Alice | 30 (1 rows)下一讲将是最后一讲——性能优化与生产部署,让 MiniDB 真正能跑在生产环境。
