第6讲:SQL 解析器

前五讲我们实现了 MiniDB 的存储引擎和事务管理------数据可以持久化存储、支持索引、具备事务能力。但到目前为止,所有操作都是通过 Python API 调用的。用户不可能用 engine.insert(txn_id, page_id, slot_id, data) 这样的接口来操作数据库。

这一讲,我们要实现 SQL 解析器,让 MiniDB 能理解并执行 SQL 语句。


一、解析器的整体架构

1.1 三步走

复制代码
SQL 文本 → 词法分析(Lexer)→ Token 流 → 语法分析(Parser)→ AST → 语义分析 → 可执行计划

1.2 支持的 SQL 子集

MiniDB 第一期支持以下 SQL:

复制代码
-- DDL
CREATE TABLE users (id INT, name VARCHAR(100), age INT);
DROP TABLE users;

-- DML
INSERT INTO users VALUES (1, 'Alice', 30);
UPDATE users SET age = 31 WHERE id = 1;
DELETE FROM users WHERE id = 1;

-- 查询
SELECT * FROM users;
SELECT id, name FROM users WHERE age > 25;
SELECT COUNT(*) FROM users GROUP BY age;

二、词法分析器(Lexer)

2.1 Token 定义

复制代码
# sql/lexer.py
from enum import Enum, auto
from dataclasses import dataclass
from typing import List, Optional

class TokenType(Enum):
    # 关键字
    SELECT = auto()
    FROM = auto()
    WHERE = auto()
    INSERT = auto()
    INTO = auto()
    VALUES = auto()
    UPDATE = auto()
    SET = auto()
    DELETE = auto()
    CREATE = auto()
    TABLE = auto()
    DROP = auto()
    AND = auto()
    OR = auto()
    NOT = auto()
    IN = auto()
    IS = auto()
    NULL = auto()
    TRUE = auto()
    FALSE = auto()
    AS = auto()
    ON = auto()
    JOIN = auto()
    LEFT = auto()
    RIGHT = auto()
    ORDER = auto()
    BY = auto()
    ASC = auto()
    DESC = auto()
    LIMIT = auto()
    OFFSET = auto()
    GROUP = auto()
    HAVING = auto()
    DISTINCT = auto()
    COUNT = auto()
    SUM = auto()
    AVG = auto()
    MAX = auto()
    MIN = auto()
    EXISTS = auto()
    BETWEEN = auto()
    LIKE = auto()
    
    # 数据类型
    TYPE_INT = auto()
    TYPE_VARCHAR = auto()
    TYPE_FLOAT = auto()
    TYPE_BOOL = auto()
    
    # 标识符与字面量
    IDENTIFIER = auto()
    NUMBER = auto()
    STRING = auto()
    
    # 运算符
    PLUS = auto()      # +
    MINUS = auto()     # -
    STAR = auto()      # *
    SLASH = auto()     # /
    EQ = auto()        # =
    NEQ = auto()       # != 或 <>
    LT = auto()        # <
    GT = auto()        # >
    LE = auto()        # <=
    GE = auto()        # >=
    ASSIGN = auto()    # :=
    
    # 分隔符
    LPAREN = auto()    # (
    RPAREN = auto()    # )
    COMMA = auto()     # ,
    SEMICOLON = auto() # ;
    DOT = auto()       # .
    
    # 特殊
    EOF = auto()
    ERROR = auto()

@dataclass
class Token:
    type: TokenType
    value: str
    line: int
    column: int
    
    def __repr__(self):
        return f"Token({self.type.name}, '{self.value}', L{self.line}:C{self.column})"

2.2 词法分析器实现

复制代码
class Lexer:
    """
    SQL 词法分析器
    
    将 SQL 文本拆分为 Token 流
    """
    
    # 关键字映射表
    KEYWORDS = {
        'SELECT': TokenType.SELECT, 'FROM': TokenType.FROM,
        'WHERE': TokenType.WHERE, 'INSERT': TokenType.INSERT,
        'INTO': TokenType.INTO, 'VALUES': TokenType.VALUES,
        'UPDATE': TokenType.UPDATE, 'SET': TokenType.SET,
        'DELETE': TokenType.DELETE, 'CREATE': TokenType.CREATE,
        'TABLE': TokenType.TABLE, 'DROP': TokenType.DROP,
        'AND': TokenType.AND, 'OR': TokenType.OR, 'NOT': TokenType.NOT,
        'IN': TokenType.IN, 'IS': TokenType.IS, 'NULL': TokenType.NULL,
        'TRUE': TokenType.TRUE, 'FALSE': TokenType.FALSE,
        'AS': TokenType.AS, 'ON': TokenType.ON, 'JOIN': TokenType.JOIN,
        'LEFT': TokenType.LEFT, 'RIGHT': TokenType.RIGHT,
        'ORDER': TokenType.ORDER, 'BY': TokenType.BY,
        'ASC': TokenType.ASC, 'DESC': TokenType.DESC,
        'LIMIT': TokenType.LIMIT, 'OFFSET': TokenType.OFFSET,
        'GROUP': TokenType.GROUP, 'HAVING': TokenType.HAVING,
        'DISTINCT': TokenType.DISTINCT,
        'COUNT': TokenType.COUNT, 'SUM': TokenType.SUM,
        'AVG': TokenType.AVG, 'MAX': TokenType.MAX, 'MIN': TokenType.MIN,
        'EXISTS': TokenType.EXISTS, 'BETWEEN': TokenType.BETWEEN,
        'LIKE': TokenType.LIKE,
        'INT': TokenType.TYPE_INT, 'INTEGER': TokenType.TYPE_INT,
        'VARCHAR': TokenType.TYPE_VARCHAR,
        'FLOAT': TokenType.TYPE_FLOAT, 'BOOL': TokenType.TYPE_BOOL,
        'BOOLEAN': TokenType.TYPE_BOOL,
    }
    
    # 单字符 Token
    SINGLE_CHAR_TOKENS = {
        '+': TokenType.PLUS, '-': TokenType.MINUS,
        '*': TokenType.STAR, '/': TokenType.SLASH,
        '(': TokenType.LPAREN, ')': TokenType.RPAREN,
        ',': TokenType.COMMA, ';': TokenType.SEMICOLON,
        '.': TokenType.DOT,
    }
    
    def __init__(self, text: str):
        self.text = text
        self.pos = 0
        self.line = 1
        self.column = 1
        self.tokens: List[Token] = []
    
    def tokenize(self) -> List[Token]:
        """执行词法分析"""
        while self.pos < len(self.text):
            char = self.text[self.pos]
            
            # 跳过空白
            if char in ' \t\r':
                self._advance()
            # 换行
            elif char == '\n':
                self.line += 1
                self.column = 1
                self.pos += 1
            # 单行注释 --
            elif char == '-' and self._peek() == '-':
                self._skip_line_comment()
            # 多行注释 /* */
            elif char == '/' and self._peek() == '*':
                self._skip_block_comment()
            # 数字
            elif char.isdigit() or (char == '-' and self._peek().isdigit()):
                self._read_number()
            # 字符串
            elif char in ("'", '"'):
                self._read_string()
            # 标识符或关键字
            elif char.isalpha() or char == '_':
                self._read_identifier()
            # 运算符
            elif char == '=':
                self._add_token(TokenType.EQ, '=')
            elif char == '!' and self._peek() == '=':
                self._add_token(TokenType.NEQ, '!=')
                self._advance()
            elif char == '<':
                if self._peek() == '=':
                    self._add_token(TokenType.LE, '<=')
                    self._advance()
                elif self._peek() == '>':
                    self._add_token(TokenType.NEQ, '<>')
                    self._advance()
                else:
                    self._add_token(TokenType.LT, '<')
            elif char == '>':
                if self._peek() == '=':
                    self._add_token(TokenType.GE, '>=')
                    self._advance()
                else:
                    self._add_token(TokenType.GT, '>')
            elif char == ':':
                self._add_token(TokenType.ASSIGN, ':=')
            elif char in self.SINGLE_CHAR_TOKENS:
                self._add_token(self.SINGLE_CHAR_TOKENS[char], char)
            else:
                self._add_token(TokenType.ERROR, char)
            
            self._advance()
        
        self.tokens.append(Token(TokenType.EOF, '', self.line, self.column))
        return self.tokens
    
    def _advance(self):
        """前进一个字符"""
        self.pos += 1
        self.column += 1
    
    def _peek(self, offset: int = 0) -> str:
        """查看当前位置之后的字符"""
        idx = self.pos + offset + 1
        return self.text[idx] if idx < len(self.text) else ''
    
    def _add_token(self, token_type: TokenType, value: str):
        """添加 Token"""
        self.tokens.append(Token(token_type, value, self.line, self.column))
    
    def _skip_line_comment(self):
        """跳过单行注释"""
        while self.pos < len(self.text) and self.text[self.pos] != '\n':
            self._advance()
    
    def _skip_block_comment(self):
        """跳过块注释"""
        self._advance()  # 跳过 /
        self._advance()  # 跳过 *
        while self.pos < len(self.text):
            if self.text[self.pos] == '*' and self._peek() == '/':
                self._advance()
                self._advance()
                return
            if self.text[self.pos] == '\n':
                self.line += 1
                self.column = 1
            self._advance()
    
    def _read_number(self):
        """读取数字"""
        start = self.pos
        is_float = False
        
        if self.text[self.pos] == '-':
            self._advance()
        
        while self.pos < len(self.text) and self.text[self.pos].isdigit():
            self._advance()
        
        if self.pos < len(self.text) and self.text[self.pos] == '.':
            is_float = True
            self._advance()
            while self.pos < len(self.text) and self.text[self.pos].isdigit():
                self._advance()
        
        value = self.text[start:self.pos]
        self._add_token(TokenType.NUMBER, value)
    
    def _read_string(self):
        """读取字符串"""
        quote = self.text[self.pos]
        self._advance()  # 跳过引号
        
        start = self.pos
        while self.pos < len(self.text):
            if self.text[self.pos] == quote:
                value = self.text[start:self.pos]
                self._add_token(TokenType.STRING, value)
                return
            if self.text[self.pos] == '\\':
                self._advance()  # 跳过转义
            self._advance()
        
        self._add_token(TokenType.ERROR, self.text[start:self.pos])
    
    def _read_identifier(self):
        """读取标识符或关键字"""
        start = self.pos
        while self.pos < len(self.text) and (self.text[self.pos].isalnum() or self.text[self.pos] == '_'):
            self._advance()
        
        value = self.text[start:self.pos]
        upper = value.upper()
        
        if upper in self.KEYWORDS:
            self._add_token(self.KEYWORDS[upper], value)
        else:
            self._add_token(TokenType.IDENTIFIER, value)

三、语法分析器(Parser)

3.1 AST 节点定义

复制代码
# sql/ast.py
from dataclasses import dataclass, field
from typing import List, Optional, Any
from enum import Enum, auto

class StatementType(Enum):
    CREATE_TABLE = auto()
    DROP_TABLE = auto()
    INSERT = auto()
    SELECT = auto()
    UPDATE = auto()
    DELETE = auto()

class ExpressionType(Enum):
    COLUMN_REF = auto()
    LITERAL = auto()
    BINARY_OP = auto()
    UNARY_OP = auto()
    FUNCTION_CALL = auto()
    STAR = auto()

@dataclass
class ColumnDef:
    """列定义"""
    name: str
    col_type: str
    length: int = 0
    nullable: bool = True
    primary_key: bool = False

@dataclass
class Expression:
    """表达式"""
    expr_type: ExpressionType
    value: Any = None
    left: 'Expression' = None
    right: 'Expression' = None
    operator: str = ''
    args: List['Expression'] = field(default_factory=list)

@dataclass
class SelectStatement:
    """SELECT 语句"""
    columns: List[Expression]
    from_table: str
    where_clause: Optional[Expression] = None
    group_by: List[str] = field(default_factory=list)
    having: Optional[Expression] = None
    order_by: List[tuple] = field(default_factory=list)  # (column, asc)
    limit: Optional[int] = None
    offset: int = 0
    distinct: bool = False

@dataclass
class InsertStatement:
    """INSERT 语句"""
    table_name: str
    columns: List[str] = field(default_factory=list)
    values: List[List[Any]] = field(default_factory=list)

@dataclass
class UpdateStatement:
    """UPDATE 语句"""
    table_name: str
    assignments: List[tuple] = field(default_factory=list)  # (column, value)
    where_clause: Optional[Expression] = None

@dataclass
class DeleteStatement:
    """DELETE 语句"""
    table_name: str
    where_clause: Optional[Expression] = None

@dataclass
class CreateTableStatement:
    """CREATE TABLE 语句"""
    table_name: str
    columns: List[ColumnDef] = field(default_factory=list)

@dataclass
class DropTableStatement:
    """DROP TABLE 语句"""
    table_name: str

3.2 语法分析器实现

复制代码
# sql/parser.py
from .lexer import Lexer, Token, TokenType
from .ast import *

class ParseError(Exception):
    """语法错误"""
    def __init__(self, message: str, token: Token = None):
        if token:
            message = f"L{token.line}:C{token.column} - {message}"
        super().__init__(message)

class Parser:
    """
    SQL 语法分析器
    
    使用递归下降分析法
    """
    
    def __init__(self, tokens: List[Token]):
        self.tokens = tokens
        self.pos = 0
    
    def parse(self) -> StatementType:
        """解析入口"""
        if self._match(TokenType.CREATE):
            return self._parse_create()
        elif self._match(TokenType.DROP):
            return self._parse_drop()
        elif self._match(TokenType.INSERT):
            return self._parse_insert()
        elif self._match(TokenType.SELECT):
            return self._parse_select()
        elif self._match(TokenType.UPDATE):
            return self._parse_update()
        elif self._match(TokenType.DELETE):
            return self._parse_delete()
        else:
            raise ParseError("期望 SQL 语句", self._peek())
    
    def _parse_create(self) -> CreateTableStatement:
        """CREATE TABLE 语句"""
        self._expect(TokenType.TABLE)
        table_name = self._expect(TokenType.IDENTIFIER).value
        
        self._expect(TokenType.LPAREN)
        columns = []
        
        while not self._check(TokenType.RPAREN):
            col_name = self._expect(TokenType.IDENTIFIER).value
            col_type = self._expect([
                TokenType.TYPE_INT, TokenType.TYPE_VARCHAR,
                TokenType.TYPE_FLOAT, TokenType.TYPE_BOOL
            ]).value
            
            col_def = ColumnDef(name=col_name, col_type=col_type.upper())
            
            # VARCHAR 长度
            if col_type.upper() == 'VARCHAR':
                self._expect(TokenType.LPAREN)
                col_def.length = int(self._expect(TokenType.NUMBER).value)
                self._expect(TokenType.RPAREN)
            
            # 可选的约束
            if self._match(TokenType.NOT):
                self._expect(TokenType.NULL)
                col_def.nullable = False
            
            if self._match(TokenType.PRIMARY):
                self._expect(TokenType.KEY)
                col_def.primary_key = True
            
            columns.append(col_def)
            
            if not self._check(TokenType.RPAREN):
                self._expect(TokenType.COMMA)
        
        self._expect(TokenType.RPAREN)
        self._expect(TokenType.SEMICOLON)
        
        return CreateTableStatement(table_name=table_name, columns=columns)
    
    def _parse_drop(self) -> DropTableStatement:
        """DROP TABLE 语句"""
        self._expect(TokenType.TABLE)
        table_name = self._expect(TokenType.IDENTIFIER).value
        self._expect(TokenType.SEMICOLON)
        return DropTableStatement(table_name=table_name)
    
    def _parse_insert(self) -> InsertStatement:
        """INSERT 语句"""
        self._expect(TokenType.INTO)
        table_name = self._expect(TokenType.IDENTIFIER).value
        
        # 可选的列列表
        columns = []
        if self._check(TokenType.LPAREN):
            self._expect(TokenType.LPAREN)
            while not self._check(TokenType.RPAREN):
                columns.append(self._expect(TokenType.IDENTIFIER).value)
                if not self._check(TokenType.RPAREN):
                    self._expect(TokenType.COMMA)
            self._expect(TokenType.RPAREN)
        
        self._expect(TokenType.VALUES)
        self._expect(TokenType.LPAREN)
        
        values = []
        while not self._check(TokenType.RPAREN):
            values.append(self._parse_literal())
            if not self._check(TokenType.RPAREN):
                self._expect(TokenType.COMMA)
        
        self._expect(TokenType.RPAREN)
        self._expect(TokenType.SEMICOLON)
        
        return InsertStatement(
            table_name=table_name,
            columns=columns,
            values=[values]
        )
    
    def _parse_select(self) -> SelectStatement:
        """SELECT 语句"""
        distinct = self._match(TokenType.DISTINCT)
        
        # SELECT 列
        columns = []
        while True:
            if self._match(TokenType.STAR):
                columns.append(Expression(expr_type=ExpressionType.STAR))
            else:
                columns.append(self._parse_expression())
            
            if not self._match(TokenType.COMMA):
                break
        
        self._expect(TokenType.FROM)
        from_table = self._expect(TokenType.IDENTIFIER).value
        
        # WHERE
        where = None
        if self._match(TokenType.WHERE):
            where = self._parse_expression()
        
        # GROUP BY
        group_by = []
        if self._match(TokenType.GROUP):
            self._expect(TokenType.BY)
            while True:
                group_by.append(self._expect(TokenType.IDENTIFIER).value)
                if not self._match(TokenType.COMMA):
                    break
        
        # HAVING
        having = None
        if self._match(TokenType.HAVING):
            having = self._parse_expression()
        
        # ORDER BY
        order_by = []
        if self._match(TokenType.ORDER):
            self._expect(TokenType.BY)
            while True:
                col = self._expect(TokenType.IDENTIFIER).value
                asc = True
                if self._match(TokenType.ASC):
                    asc = True
                elif self._match(TokenType.DESC):
                    asc = False
                order_by.append((col, asc))
                if not self._match(TokenType.COMMA):
                    break
        
        # LIMIT
        limit = None
        if self._match(TokenType.LIMIT):
            limit = int(self._expect(TokenType.NUMBER).value)
        
        # OFFSET
        offset = 0
        if self._match(TokenType.OFFSET):
            offset = int(self._expect(TokenType.NUMBER).value)
        
        self._expect(TokenType.SEMICOLON)
        
        return SelectStatement(
            columns=columns, from_table=from_table,
            where_clause=where, group_by=group_by,
            having=having, order_by=order_by,
            limit=limit, offset=offset, distinct=distinct
        )
    
    def _parse_update(self) -> UpdateStatement:
        """UPDATE 语句"""
        table_name = self._expect(TokenType.IDENTIFIER).value
        self._expect(TokenType.SET)
        
        assignments = []
        while True:
            col = self._expect(TokenType.IDENTIFIER).value
            self._expect(TokenType.EQ)
            val = self._parse_literal()
            assignments.append((col, val))
            if not self._match(TokenType.COMMA):
                break
        
        where = None
        if self._match(TokenType.WHERE):
            where = self._parse_expression()
        
        self._expect(TokenType.SEMICOLON)
        
        return UpdateStatement(
            table_name=table_name,
            assignments=assignments,
            where_clause=where
        )
    
    def _parse_delete(self) -> DeleteStatement:
        """DELETE 语句"""
        self._expect(TokenType.FROM)
        table_name = self._expect(TokenType.IDENTIFIER).value
        
        where = None
        if self._match(TokenType.WHERE):
            where = self._parse_expression()
        
        self._expect(TokenType.SEMICOLON)
        
        return DeleteStatement(table_name=table_name, where_clause=where)
    
    def _parse_expression(self) -> Expression:
        """解析表达式(支持优先级)"""
        return self._parse_or()
    
    def _parse_or(self) -> Expression:
        left = self._parse_and()
        while self._match(TokenType.OR):
            right = self._parse_and()
            left = Expression(ExpressionType.BINARY_OP, operator='OR', left=left, right=right)
        return left
    
    def _parse_and(self) -> Expression:
        left = self._parse_comparison()
        while self._match(TokenType.AND):
            right = self._parse_comparison()
            left = Expression(ExpressionType.BINARY_OP, operator='AND', left=left, right=right)
        return left
    
    def _parse_comparison(self) -> Expression:
        left = self._parse_term()
        
        if self._match(TokenType.EQ):
            right = self._parse_term()
            return Expression(ExpressionType.BINARY_OP, operator='=', left=left, right=right)
        elif self._match(TokenType.NEQ):
            right = self._parse_term()
            return Expression(ExpressionType.BINARY_OP, operator='!=', left=left, right=right)
        elif self._match(TokenType.LT):
            right = self._parse_term()
            return Expression(ExpressionType.BINARY_OP, operator='<', left=left, right=right)
        elif self._match(TokenType.GT):
            right = self._parse_term()
            return Expression(ExpressionType.BINARY_OP, operator='>', left=left, right=right)
        elif self._match(TokenType.LE):
            right = self._parse_term()
            return Expression(ExpressionType.BINARY_OP, operator='<=', left=left, right=right)
        elif self._match(TokenType.GE):
            right = self._parse_term()
            return Expression(ExpressionType.BINARY_OP, operator='>=', left=left, right=right)
        
        return left
    
    def _parse_term(self) -> Expression:
        left = self._parse_factor()
        
        while self._check(TokenType.PLUS) or self._check(TokenType.MINUS):
            if self._match(TokenType.PLUS):
                right = self._parse_factor()
                left = Expression(ExpressionType.BINARY_OP, operator='+', left=left, right=right)
            elif self._match(TokenType.MINUS):
                right = self._parse_factor()
                left = Expression(ExpressionType.BINARY_OP, operator='-', left=left, right=right)
        
        return left
    
    def _parse_factor(self) -> Expression:
        left = self._parse_unary()
        
        while self._check(TokenType.STAR) or self._check(TokenType.SLASH):
            if self._match(TokenType.STAR):
                right = self._parse_unary()
                left = Expression(ExpressionType.BINARY_OP, operator='*', left=left, right=right)
            elif self._match(TokenType.SLASH):
                right = self._parse_unary()
                left = Expression(ExpressionType.BINARY_OP, operator='/', left=left, right=right)
        
        return left
    
    def _parse_unary(self) -> Expression:
        if self._match(TokenType.MINUS):
            operand = self._parse_primary()
            return Expression(ExpressionType.UNARY_OP, operator='-', right=operand)
        elif self._match(TokenType.NOT):
            operand = self._parse_primary()
            return Expression(ExpressionType.UNARY_OP, operator='NOT', right=operand)
        return self._parse_primary()
    
    def _parse_primary(self) -> Expression:
        # 括号表达式
        if self._match(TokenType.LPAREN):
            expr = self._parse_expression()
            self._expect(TokenType.RPAREN)
            return expr
        
        # 函数调用
        if self._check(TokenType.IDENTIFIER) and self._peek(1).type == TokenType.LPAREN:
            name = self._expect(TokenType.IDENTIFIER).value
            self._expect(TokenType.LPAREN)
            args = []
            if not self._check(TokenType.RPAREN):
                args.append(self._parse_expression())
                while self._match(TokenType.COMMA):
                    args.append(self._parse_expression())
            self._expect(TokenType.RPAREN)
            return Expression(ExpressionType.FUNCTION_CALL, value=name, args=args)
        
        # 列引用
        if self._check(TokenType.IDENTIFIER):
            name = self._expect(TokenType.IDENTIFIER).value
            return Expression(ExpressionType.COLUMN_REF, value=name)
        
        # 字面量
        return Expression(ExpressionType.LITERAL, value=self._parse_literal())
    
    def _parse_literal(self) -> Any:
        """解析字面量"""
        if self._match(TokenType.NUMBER):
            val = self.tokens[self.pos - 1].value
            return int(val) if '.' not in val else float(val)
        elif self._match(TokenType.STRING):
            return self.tokens[self.pos - 1].value
        elif self._match(TokenType.TRUE):
            return True
        elif self._match(TokenType.FALSE):
            return False
        elif self._match(TokenType.NULL):
            return None
        raise ParseError("期望字面量", self._peek())
    
    def _match(self, *types) -> bool:
        """匹配并消费 Token"""
        if self._check(*types):
            self.pos += 1
            return True
        return False
    
    def _check(self, *types) -> bool:
        """检查当前 Token 类型"""
        if self.pos >= len(self.tokens):
            return False
        return self.tokens[self.pos].type in types
    
    def _expect(self, *types) -> Token:
        """期望指定类型的 Token,否则报错"""
        if self._check(*types):
            token = self.tokens[self.pos]
            self.pos += 1
            return token
        expected = ', '.join(t.name for t in types)
        raise ParseError(f"期望 {expected}", self._peek())
    
    def _peek(self, offset: int = 0) -> Token:
        """查看当前 Token"""
        idx = self.pos + offset
        return self.tokens[idx] if idx < len(self.tokens) else self.tokens[-1]

四、完整演示

复制代码
def test_sql_parser():
    print("=" * 60)
    print("📝 SQL 解析器测试")
    print("=" * 60)
    
    test_cases = [
        "CREATE TABLE users (id INT, name VARCHAR(100), age INT);",
        "DROP TABLE users;",
        "INSERT INTO users VALUES (1, 'Alice', 30);",
        "SELECT * FROM users;",
        "SELECT id, name FROM users WHERE age > 25;",
        "UPDATE users SET age = 31 WHERE id = 1;",
        "DELETE FROM users WHERE id = 1;",
        "SELECT COUNT(*) FROM users GROUP BY age;",
        "SELECT name, age FROM users WHERE age >= 18 AND age <= 60 ORDER BY age DESC LIMIT 10;",
    ]
    
    for sql in test_cases:
        print(f"\n📄 SQL: {sql}")
        try:
            lexer = Lexer(sql)
            tokens = lexer.tokenize()
            print(f"   Token: {[t.type.name for t in tokens[:10]]}...")
            
            parser = Parser(tokens)
            ast = parser.parse()
            print(f"   AST: {ast}")
        except Exception as e:
            print(f"   ❌ 错误: {e}")

def test_complex_expressions():
    print("\n" + "=" * 60)
    print("🔢 复杂表达式测试")
    print("=" * 60)
    
    expressions = [
        "SELECT a + b * c FROM t;",
        "SELECT (a + b) * c FROM t;",
        "SELECT a > 10 AND b < 20 OR c = 30 FROM t;",
        "SELECT NOT finished AND age BETWEEN 18 AND 60 FROM t;",
    ]
    
    for sql in expressions:
        print(f"\n📄 SQL: {sql}")
        try:
            lexer = Lexer(sql)
            tokens = lexer.tokenize()
            parser = Parser(tokens)
            ast = parser.parse()
            print(f"   ✅ 解析成功")
        except Exception as e:
            print(f"   ❌ 错误: {e}")

def test_error_handling():
    print("\n" + "=" * 60)
    print("⚠️  错误处理测试")
    print("=" * 60)
    
    error_cases = [
        "SELECT FROM users;",           # 缺少列
        "CREATE TABLE (id INT);",       # 缺少表名
        "INSERT INTO VALUES (1);",      # 缺少表名
        "SELECT * FORM users;",         # FROM 拼写错误
    ]
    
    for sql in error_cases:
        print(f"\n📄 SQL: {sql}")
        try:
            lexer = Lexer(sql)
            tokens = lexer.tokenize()
            parser = Parser(tokens)
            ast = parser.parse()
            print(f"   ✅ 意外成功: {ast}")
        except ParseError as e:
            print(f"   ✅ 正确捕获错误: {e}")
        except Exception as e:
            print(f"   ❌ 其他错误: {e}")

if __name__ == "__main__":
    test_sql_parser()
    test_complex_expressions()
    test_error_handling()

五、总结

这一讲实现了 MiniDB 的 SQL 解析器:

  • 词法分析器:将 SQL 文本拆分为 Token 流,支持关键字、标识符、数字、字符串、运算符、注释

  • 语法分析器:递归下降分析法,支持 CREATE/DROP/INSERT/SELECT/UPDATE/DELETE

  • 表达式解析:支持算术运算、比较运算、逻辑运算、函数调用、括号分组

  • AST 生成:将 SQL 转换为结构化语法树

现在 MiniDB 能理解 SQL 语句了。下一讲将实现 查询执行引擎,把这些 AST 转化为实际的数据操作。

相关推荐
weixin_440784111 小时前
【OkHttp实现原理】
android·java·okhttp
a3535413821 小时前
C++项目如何架构优化
java·c++·架构
吴老弟i1 小时前
Mac 本地安装 MySQL 学习指南
数据库
AC赳赳老秦1 小时前
官方技术文档聚合实践:用 OpenClaw 批量抓取开源项目文档,构建离线可检索技术知识库
java·运维·服务器·python·信息可视化·deepseek·openclaw
breeze jiang1 小时前
Next.js App Router 全栈实战:从 SPA 的 SEO 痛点到服务端组件与 Hydration 水合机制
开发语言·javascript·ecmascript
程序员老陆1 小时前
Qt::WidgetAttribute 常用属性详解
开发语言·qt
略略略咯咯1 小时前
ActiveMQ
java·activemq·java-activemq
万亿少女的梦1681 小时前
基于Spring Boot、Java与MySQL的同城宠物服务系统设计与实现
java·spring boot·mysql·系统设计·协同过滤
不才不才不不才1 小时前
Spring 源码系列(20): @SpringBootApplication 三注解拆解
java·后端·spring