上一讲我们实现了SQL解析器,MiniDB能把SQL语句转换成AST。但AST只是一棵语法树,它告诉了我们"要做什么",却没有说"怎么做"。
这一讲,我们要实现查询执行引擎——把AST变成实际的数据操作,从表中读取数据、过滤、聚合、排序,最后返回结果。
一、火山模型(Volcano Iterator Model)
1.1 核心思想
火山模型是数据库执行引擎中最经典的架构。它的核心是一个迭代器接口:
class Iterator: def init(self): # 初始化 def next(self): # 返回下一行,没有则返回None def close(self): # 清理资源每个算子都是一个迭代器,算子之间通过next()方法串联:
SELECT name, age FROM users WHERE age > 25 ORDER BY age ↓ ┌─────────────────┐ │ SortIterator │ ← 排序 └────────┬────────┘ │ ┌────────▼────────┐ │ FilterIterator │ ← 过滤 age > 25 └────────┬────────┘ │ ┌────────▼────────┐ │ ProjectIterator │ ← 投影 name, age └────────┬────────┘ │ ┌────────▼────────┐ │ SeqScanIterator │ ← 全表扫描 └─────────────────┘执行过程:顶层算子不断调用next(),下层算子向上返回数据,像火山喷发一样逐行流动。
二、迭代器接口定义
# sql/executor/iterator.py from abc import ABC, abstractmethod from typing import Any, Dict, List, Optional, Iterator as PyIterator class Tuple: """一行数据""" def __init__(self, values: Dict[str, Any], schema: List[str]): self.values = values self.schema = schema def get(self, column: str) -> Any: return self.values.get(column) def set(self, column: str, value: Any): self.values[column] = value def copy(self) -> 'Tuple': return Tuple(self.values.copy(), self.schema.copy()) def __repr__(self): cols = ', '.join(f"{k}={v}" for k, v in self.values.items()) return f"Tuple({cols})" class Executor(ABC): """执行器基类(火山模型)""" def __init__(self, children: List['Executor'] = None): self.children = children or [] self.opened = False def open(self): """初始化执行器""" self.opened = True for child in self.children: child.open() @abstractmethod def next(self) -> Optional[Tuple]: """返回下一行数据,没有则返回None""" pass def close(self): """清理资源""" for child in self.children: child.close() self.opened = False def __iter__(self) -> PyIterator[Tuple]: self.open() return self def __next__(self) -> Tuple: row = self.next() if row is None: self.close() raise StopIteration return row三、物理算子实现
3.1 全表扫描
class SeqScanExecutor(Executor): """全表扫描执行器""" def __init__(self, table_name: str, catalog, storage): super().__init__() self.table_name = table_name self.catalog = catalog self.storage = storage self.schema = None self.page_ids: List[int] = [] self.current_page_idx = 0 self.current_slot_idx = 0 self.current_page = None def open(self): super().open() # 获取表的元数据 self.schema = self.catalog.get_table_schema(self.table_name) # 获取表的所有数据页 self.page_ids = self.catalog.get_table_pages(self.table_name) def next(self) -> Optional[Tuple]: while self.current_page_idx < len(self.page_ids): if self.current_page is None: # 加载下一页 page_id = self.page_ids[self.current_page_idx] self.current_page = self.storage.read_page(page_id) self.current_slot_idx = 0 # 读取当前slot while self.current_slot_idx < self.current_page.num_records: record = self.current_page.read_record(self.current_slot_idx) self.current_slot_idx += 1 if record is not None: # 反序列化为Tuple return self._deserialize_tuple(record) # 当前页读完,移到下一页 self.current_page_idx += 1 self.current_page = None return None def _deserialize_tuple(self, record: bytes) -> Tuple: """将二进制记录反序列化为Tuple""" values = {} offset = 0 for col in self.schema.columns: if col.col_type == 'INT': import struct values[col.name] = struct.unpack_from('!i', record, offset)[0] offset += 4 elif col.col_type == 'VARCHAR': length = struct.unpack_from('!H', record, offset)[0] offset += 2 values[col.name] = record[offset:offset + length].decode('utf-8') offset += length elif col.col_type == 'FLOAT': import struct values[col.name] = struct.unpack_from('!f', record, offset)[0] offset += 4 elif col.col_type == 'BOOL': values[col.name] = bool(record[offset]) offset += 1 return Tuple(values, [c.name for c in self.schema.columns])3.2 过滤
class FilterExecutor(Executor): """过滤执行器(WHERE子句)""" def __init__(self, child: Executor, condition): super().__init__([child]) self.condition = condition def next(self) -> Optional[Tuple]: while True: row = self.children[0].next() if row is None: return None if self._evaluate_condition(row, self.condition): return row def _evaluate_condition(self, row: Tuple, expr) -> bool: """递归求值条件表达式""" from sql.ast import ExpressionType if expr.expr_type == ExpressionType.BINARY_OP: left = self._evaluate_condition(row, expr.left) right = self._evaluate_condition(row, expr.right) if expr.operator == 'AND': return left and right elif expr.operator == 'OR': return left or right elif expr.operator == '=': return left == right elif expr.operator == '!=': return left != right elif expr.operator == '<': return left < right elif expr.operator == '>': return left > right elif expr.operator == '<=': return left <= right elif expr.operator == '>=': return left >= right elif expr.operator == '+': return left + right elif expr.operator == '-': return left - right elif expr.operator == '*': return left * right elif expr.operator == '/': return left / right elif expr.expr_type == ExpressionType.UNARY_OP: operand = self._evaluate_condition(row, expr.right) if expr.operator == 'NOT': return not operand elif expr.operator == '-': return -operand elif expr.expr_type == ExpressionType.COLUMN_REF: return row.get(expr.value) elif expr.expr_type == ExpressionType.LITERAL: return expr.value elif expr.expr_type == ExpressionType.FUNCTION_CALL: args = [self._evaluate_condition(row, arg) for arg in expr.args] return self._call_function(expr.value, args) return True def _call_function(self, name: str, args: List) -> Any: """调用内置函数""" name = name.upper() if name == 'COUNT': # COUNT在聚合算子中处理 return len(args) elif name == 'UPPER': return str(args[0]).upper() if args else '' elif name == 'LOWER': return str(args[0]).lower() if args else '' elif name == 'LENGTH': return len(str(args[0])) if args else 0 return args[0] if args else None3.3 投影
class ProjectExecutor(Executor): """投影执行器(SELECT子句)""" def __init__(self, child: Executor, columns: List): super().__init__([child]) self.columns = columns def next(self) -> Optional[Tuple]: row = self.children[0].next() if row is None: return None # 如果是 SELECT *,返回所有列 if len(self.columns) == 1 and self.columns[0].expr_type.name == 'STAR': return row # 投影指定的列 new_values = {} for col_expr in self.columns: if col_expr.expr_type.name == 'COLUMN_REF': col_name = col_expr.value new_values[col_name] = row.get(col_name) else: # 表达式列 from sql.executor.iterator import FilterExecutor evaluator = FilterExecutor(None, None) value = evaluator._evaluate_condition(row, col_expr) new_values[col_expr.value] = value return Tuple(new_values, list(new_values.keys()))3.4 排序
class SortExecutor(Executor): """排序执行器(ORDER BY子句)""" def __init__(self, child: Executor, order_by: List[tuple]): super().__init__([child]) self.order_by = order_by # [(column, asc), ...] self.sorted_rows: List[Tuple] = [] self.current_idx = 0 def open(self): super().open() # 读取所有数据到内存 rows = [] while True: row = self.children[0].next() if row is None: break rows.append(row.copy()) # 排序 self.sorted_rows = sorted(rows, key=self._sort_key) if self.order_by and not self.order_by[0][1]: self.sorted_rows.reverse() def next(self) -> Optional[Tuple]: if self.current_idx >= len(self.sorted_rows): return None row = self.sorted_rows[self.current_idx] self.current_idx += 1 return row def _sort_key(self, row: Tuple) -> tuple: """生成排序键""" keys = [] for col_name, _ in self.order_by: keys.append(row.get(col_name)) return tuple(keys)3.5 聚合
class AggregateExecutor(Executor): """聚合执行器(GROUP BY + 聚合函数)""" def __init__(self, child: Executor, group_by: List[str], agg_funcs: List[tuple]): super().__init__([child]) self.group_by = group_by self.agg_funcs = agg_funcs # [(func_name, column), ...] self.aggregated_rows: List[Tuple] = [] self.current_idx = 0 def open(self): super().open() # 分组聚合 groups = {} while True: row = self.children[0].next() if row is None: break # 生成分组键 if self.group_by: group_key = tuple(row.get(col) for col in self.group_by) else: group_key = 'all' if group_key not in groups: groups[group_key] = [] groups[group_key].append(row) # 计算聚合结果 for group_key, rows in groups.items(): values = {} # GROUP BY 列 if self.group_by: for i, col in enumerate(self.group_by): values[col] = group_key[i] # 聚合函数 for func_name, col_name in self.agg_funcs: col_values = [r.get(col_name) for r in rows if r.get(col_name) is not None] if func_name.upper() == 'COUNT': values[f'{func_name}({col_name})'] = len(col_values) elif func_name.upper() == 'SUM': values[f'{func_name}({col_name})'] = sum(col_values) elif func_name.upper() == 'AVG': values[f'{func_name}({col_name})'] = sum(col_values) / len(col_values) if col_values else 0 elif func_name.upper() == 'MAX': values[f'{func_name}({col_name})'] = max(col_values) if col_values else None elif func_name.upper() == 'MIN': values[f'{func_name}({col_name})'] = min(col_values) if col_values else None self.aggregated_rows.append(Tuple(values, list(values.keys()))) def next(self) -> Optional[Tuple]: if self.current_idx >= len(self.aggregated_rows): return None row = self.aggregated_rows[self.current_idx] self.current_idx += 1 return row3.6 限制与偏移
class LimitExecutor(Executor): """限制执行器(LIMIT/OFFSET子句)""" def __init__(self, child: Executor, limit: int = None, offset: int = 0): super().__init__([child]) self.limit = limit self.offset = offset self.returned = 0 self.skipped = 0 def next(self) -> Optional[Tuple]: # 跳过OFFSET行 while self.skipped < self.offset: row = self.children[0].next() if row is None: return None self.skipped += 1 # 检查LIMIT if self.limit is not None and self.returned >= self.limit: return None row = self.children[0].next() if row is not None: self.returned += 1 return row四、执行计划构建器
class PlanBuilder: """执行计划构建器:将AST转换为执行器树""" def __init__(self, catalog, storage): self.catalog = catalog self.storage = storage def build(self, statement) -> Executor: """根据语句类型构建执行计划""" from sql.ast import ( SelectStatement, InsertStatement, UpdateStatement, DeleteStatement, CreateTableStatement, DropTableStatement ) if isinstance(statement, SelectStatement): return self._build_select(statement) elif isinstance(statement, InsertStatement): return self._build_insert(statement) elif isinstance(statement, UpdateStatement): return self._build_update(statement) elif isinstance(statement, DeleteStatement): return self._build_delete(statement) elif isinstance(statement, CreateTableStatement): return self._build_create_table(statement) elif isinstance(statement, DropTableStatement): return self._build_drop_table(statement) else: raise ValueError(f"不支持的语句类型: {type(statement)}") def _build_select(self, stmt: SelectStatement) -> Executor: """构建SELECT执行计划""" # 1. 全表扫描(最底层) executor = SeqScanExecutor(stmt.from_table, self.catalog, self.storage) # 2. 过滤(WHERE) if stmt.where_clause: executor = FilterExecutor(executor, stmt.where_clause) # 3. 聚合(GROUP BY) if stmt.group_by: # 从SELECT列中提取聚合函数 agg_funcs = self._extract_agg_funcs(stmt.columns) executor = AggregateExecutor(executor, stmt.group_by, agg_funcs) # 4. 排序(ORDER BY) if stmt.order_by: executor = SortExecutor(executor, stmt.order_by) # 5. 限制(LIMIT/OFFSET) if stmt.limit is not None or stmt.offset > 0: executor = LimitExecutor(executor, stmt.limit, stmt.offset) # 6. 投影(SELECT列) executor = ProjectExecutor(executor, stmt.columns) return executor def _build_insert(self, stmt: InsertStatement) -> Executor: """构建INSERT执行计划""" return InsertExecutor(stmt.table_name, stmt.columns, stmt.values, self.catalog, self.storage) def _build_update(self, stmt: UpdateStatement) -> Executor: """构建UPDATE执行计划""" return UpdateExecutor(stmt.table_name, stmt.assignments, stmt.where_clause, self.catalog, self.storage) def _build_delete(self, stmt: DeleteStatement) -> Executor: """构建DELETE执行计划""" return DeleteExecutor(stmt.table_name, stmt.where_clause, self.catalog, self.storage) def _build_create_table(self, stmt: CreateTableStatement) -> Executor: """构建CREATE TABLE执行计划""" return DDLExecutor('CREATE_TABLE', stmt, self.catalog) def _build_drop_table(self, stmt: DropTableStatement) -> Executor: """构建DROP TABLE执行计划""" return DDLExecutor('DROP_TABLE', stmt, self.catalog) def _extract_agg_funcs(self, columns: List) -> List[tuple]: """从SELECT列中提取聚合函数""" from sql.ast import ExpressionType agg_funcs = [] for col in columns: if col.expr_type == ExpressionType.FUNCTION_CALL: name = col.value.upper() if name in ('COUNT', 'SUM', 'AVG', 'MAX', 'MIN'): col_name = col.args[0].value if col.args else '*' agg_funcs.append((name, col_name)) return agg_funcs五、DML执行器
class InsertExecutor(Executor): """INSERT执行器""" def __init__(self, table_name, columns, values, catalog, storage): super().__init__() self.table_name = table_name self.columns = columns self.values = values self.catalog = catalog self.storage = storage self.done = False def next(self) -> Optional[Tuple]: if self.done: return None self.done = True schema = self.catalog.get_table_schema(self.table_name) for value_row in self.values: # 构建列名到值的映射 if self.columns: col_map = dict(zip(self.columns, value_row)) else: col_map = dict(zip([c.name for c in schema.columns], value_row)) # 序列化为二进制 record = self._serialize_record(col_map, schema) # 写入存储 page_id = self.catalog.get_insert_page(self.table_name) slot_id = self.storage.write_record(page_id, record) # 更新索引 self.catalog.update_indexes(self.table_name, col_map, page_id, slot_id) return Tuple({'affected_rows': len(self.values)}, ['affected_rows']) def _serialize_record(self, col_map: dict, schema) -> bytes: """将记录序列化为二进制""" import struct data = bytearray() for col in schema.columns: value = col_map.get(col.name) if col.col_type == 'INT': data += struct.pack('!i', value or 0) elif col.col_type == 'VARCHAR': encoded = str(value or '').encode('utf-8') data += struct.pack('!H', len(encoded)) data += encoded elif col.col_type == 'FLOAT': data += struct.pack('!f', float(value or 0)) elif col.col_type == 'BOOL': data += bytes([1 if value else 0]) return bytes(data) class UpdateExecutor(Executor): """UPDATE执行器""" def __init__(self, table_name, assignments, where_clause, catalog, storage): super().__init__() self.table_name = table_name self.assignments = assignments self.where_clause = where_clause self.catalog = catalog self.storage = storage self.done = False def next(self) -> Optional[Tuple]: if self.done: return None self.done = True # 先用SeqScan + Filter找到要更新的行 scan = SeqScanExecutor(self.table_name, self.catalog, self.storage) filter_op = FilterExecutor(scan, self.where_clause) if self.where_clause else scan updated = 0 filter_op.open() while True: row = filter_op.next() if row is None: break # 应用更新 for col_name, new_value in self.assignments: row.set(col_name, new_value) # 写回存储 # ... (简化处理) updated += 1 filter_op.close() return Tuple({'affected_rows': updated}, ['affected_rows']) class DeleteExecutor(Executor): """DELETE执行器""" def __init__(self, table_name, where_clause, catalog, storage): super().__init__() self.table_name = table_name self.where_clause = where_clause self.catalog = catalog self.storage = storage self.done = False def next(self) -> Optional[Tuple]: if self.done: return None self.done = True scan = SeqScanExecutor(self.table_name, self.catalog, self.storage) filter_op = FilterExecutor(scan, self.where_clause) if self.where_clause else scan deleted = 0 filter_op.open() while True: row = filter_op.next() if row is None: break # 标记删除 # ... (简化处理) deleted += 1 filter_op.close() return Tuple({'affected_rows': deleted}, ['affected_rows']) class DDLExecutor(Executor): """DDL执行器(CREATE/DROP TABLE)""" def __init__(self, action: str, statement, catalog): super().__init__() self.action = action self.statement = statement self.catalog = catalog self.done = False def next(self) -> Optional[Tuple]: if self.done: return None self.done = True if self.action == 'CREATE_TABLE': self.catalog.create_table(self.statement.table_name, self.statement.columns) return Tuple({'message': f'Table {self.statement.table_name} created'}, ['message']) elif self.action == 'DROP_TABLE': self.catalog.drop_table(self.statement.table_name) return Tuple({'message': f'Table {self.statement.table_name} dropped'}, ['message']) return None六、完整演示
def test_query_executor(): print("=" * 60) print("⚡ 查询执行引擎测试") print("=" * 60) # 模拟Catalog和Storage class MockCatalog: def __init__(self): self.tables = {} def create_table(self, name, columns): self.tables[name] = { 'columns': columns, 'pages': [], 'data': [] } def get_table_schema(self, name): return self.tables[name] def get_table_pages(self, name): return self.tables[name]['pages'] def get_insert_page(self, name): return 0 class MockStorage: def read_page(self, page_id): return None catalog = MockCatalog() storage = MockStorage() # 创建测试表 catalog.create_table('users', [ type('Column', (), {'name': 'id', 'col_type': 'INT'}), type('Column', (), {'name': 'name', 'col_type': 'VARCHAR'}), type('Column', (), {'name': 'age', 'col_type': 'INT'}), ]) # 模拟数据 mock_data = [ Tuple({'id': 1, 'name': 'Alice', 'age': 30}, ['id', 'name', 'age']), Tuple({'id': 2, 'name': 'Bob', 'age': 25}, ['id', 'name', 'age']), Tuple({'id': 3, 'name': 'Charlie', 'age': 35}, ['id', 'name', 'age']), Tuple({'id': 4, 'name': 'Diana', 'age': 28}, ['id', 'name', 'age']), ] # 1. 全表扫描 print("\n📄 1. 全表扫描:") from sql.ast import Expression, ExpressionType star_expr = Expression(expr_type=ExpressionType.STAR) plan = ProjectExecutor( MockSeqScan(mock_data), [star_expr] ) plan.open() while True: row = plan.next() if row is None: break print(f" {row}") plan.close() # 2. 过滤 print("\n🔍 2. 过滤 (age > 27):") from sql.ast import Expression, ExpressionType age_col = Expression(expr_type=ExpressionType.COLUMN_REF, value='age') literal_27 = Expression(expr_type=ExpressionType.LITERAL, value=27) condition = Expression(expr_type=ExpressionType.BINARY_OP, operator='>', left=age_col, right=literal_27) plan = ProjectExecutor( FilterExecutor(MockSeqScan(mock_data), condition), [star_expr] ) plan.open() while True: row = plan.next() if row is None: break print(f" {row}") plan.close() # 3. 排序 print("\n📊 3. 排序 (ORDER BY age DESC):") plan = ProjectExecutor( SortExecutor(MockSeqScan(mock_data), [('age', False)]), [star_expr] ) plan.open() while True: row = plan.next() if row is None: break print(f" {row}") plan.close() # 4. 投影 + 过滤 + 排序 print("\n🎯 4. 组合查询 (SELECT name, age WHERE age > 25 ORDER BY name):") name_expr = Expression(expr_type=ExpressionType.COLUMN_REF, value='name') age_expr = Expression(expr_type=ExpressionType.COLUMN_REF, value='age') plan = ProjectExecutor( SortExecutor( FilterExecutor(MockSeqScan(mock_data), condition), [('name', True)] ), [name_expr, age_expr] ) plan.open() while True: row = plan.next() if row is None: break print(f" {row}") plan.close() class MockSeqScan(Executor): """模拟全表扫描""" def __init__(self, data: List[Tuple]): super().__init__() self.data = data self.idx = 0 def next(self) -> Optional[Tuple]: if self.idx >= len(self.data): return None row = self.data[self.idx] self.idx += 1 return row.copy() if __name__ == "__main__": test_query_executor()七、总结
这一讲实现了MiniDB的查询执行引擎:
火山模型:迭代器接口,算子之间通过
next()串联物理算子:SeqScan、Filter、Project、Sort、Aggregate、Limit
执行计划构建:将AST转换为执行器树
DML执行:INSERT/UPDATE/DELETE的实现
现在MiniDB已经能执行完整的SQL查询了。下一讲将实现查询优化器,让查询执行得更高效——比如选择合适的索引、调整连接顺序等。