// Copyright 2021 Molecula Corp. All rights reserved. package sql2 import ( "io" "strings" ) // Parser represents a SQL parser. type Parser struct { s *Scanner pos Pos // current position tok Token // current token lit string // current literal value full bool // buffer full } // NewParser returns a new instance of Parser that reads from r. func NewParser(r io.Reader) *Parser { return &Parser{ s: NewScanner(r), } } // ParseExprString parses s into an expression. Returns nil if s is blank. func ParseExprString(s string) (Expr, error) { if s == "" { return nil, nil } return NewParser(strings.NewReader(s)).ParseExpr() } // MustParseExprString parses s into an expression. Panic on error. func MustParseExprString(s string) Expr { expr, err := ParseExprString(s) if err != nil { panic(err) } return expr } func (p *Parser) ParseStatement() (stmt Statement, err error) { switch tok := p.peek(); tok { case EOF: return nil, io.EOF case EXPLAIN: if stmt, err = p.parseExplainStatement(); err != nil { return stmt, err } default: if stmt, err = p.parseNonExplainStatement(); err != nil { return stmt, err } } // Read trailing semicolon or end of file. if tok := p.peek(); tok != EOF && tok != SEMI { return stmt, p.errorExpected(p.pos, p.tok, "semicolon or EOF") } p.scan() return stmt, nil } // parseExplain parses EXPLAIN [QUERY PLAN] STMT. func (p *Parser) parseExplainStatement() (_ *ExplainStatement, err error) { var tok Token // Parse initial "EXPLAIN" token. var stmt ExplainStatement stmt.Explain, tok, _ = p.scan() assert(tok == EXPLAIN) // Parse optional "QUERY PLAN" tokens. if p.peek() == QUERY { stmt.Query, _, _ = p.scan() if p.peek() != PLAN { return &stmt, p.errorExpected(p.pos, p.tok, "PLAN") } stmt.QueryPlan, _, _ = p.scan() } // Parse statement to be explained. if stmt.Stmt, err = p.parseNonExplainStatement(); err != nil { return &stmt, err } return &stmt, nil } // parseStmt parses all statement types. func (p *Parser) parseNonExplainStatement() (Statement, error) { switch p.peek() { case ANALYZE: return p.parseAnalyzeStatement() case ALTER: return p.parseAlterTableStatement() case BEGIN: return p.parseBeginStatement() case COMMIT, END: return p.parseCommitStatement() case ROLLBACK: return p.parseRollbackStatement() case SAVEPOINT: return p.parseSavepointStatement() case RELEASE: return p.parseReleaseStatement() case CREATE: return p.parseCreateStatement() case DROP: return p.parseDropStatement() case SELECT, VALUES: return p.parseSelectStatement(false, nil) case INSERT, REPLACE: return p.parseInsertStatement(nil) case UPDATE: return p.parseUpdateStatement(nil) case DELETE: return p.parseDeleteStatement(nil) case WITH: return p.parseWithStatement() default: return nil, p.errorExpected(p.pos, p.tok, "statement") } } // parseWithStatement is called only from parseNonExplainStatement as we don't // know what kind of statement we'll have after the CTEs (e.g. SELECT, INSERT, etc). func (p *Parser) parseWithStatement() (Statement, error) { withClause, err := p.parseWithClause() if err != nil { return nil, err } switch p.peek() { case SELECT, VALUES: return p.parseSelectStatement(false, withClause) case INSERT, REPLACE: return p.parseInsertStatement(withClause) case UPDATE: return p.parseUpdateStatement(withClause) case DELETE: return p.parseDeleteStatement(withClause) default: return nil, p.errorExpected(p.pos, p.tok, "SELECT, VALUES, INSERT, REPLACE, UPDATE, or DELETE") } } func (p *Parser) parseBeginStatement() (*BeginStatement, error) { assert(p.peek() == BEGIN) var stmt BeginStatement stmt.Begin, _, _ = p.scan() // Parse transaction type. switch p.peek() { case DEFERRED: stmt.Deferred, _, _ = p.scan() case IMMEDIATE: stmt.Immediate, _, _ = p.scan() case EXCLUSIVE: stmt.Exclusive, _, _ = p.scan() } // Parse optional TRANSCTION keyword. if p.peek() == TRANSACTION { stmt.Transaction, _, _ = p.scan() } return &stmt, nil } func (p *Parser) parseCommitStatement() (*CommitStatement, error) { assert(p.peek() == COMMIT || p.peek() == END) var stmt CommitStatement if p.peek() == COMMIT { stmt.Commit, _, _ = p.scan() } else { stmt.End, _, _ = p.scan() } if p.peek() == TRANSACTION { stmt.Transaction, _, _ = p.scan() } return &stmt, nil } func (p *Parser) parseRollbackStatement() (_ *RollbackStatement, err error) { assert(p.peek() == ROLLBACK) var stmt RollbackStatement stmt.Rollback, _, _ = p.scan() // Parse optional "TRANSACTION". if p.peek() == TRANSACTION { stmt.Transaction, _, _ = p.scan() } // Parse optional "TO SAVEPOINT savepoint-name" if p.peek() == TO { stmt.To, _, _ = p.scan() if p.peek() == SAVEPOINT { stmt.Savepoint, _, _ = p.scan() } if stmt.SavepointName, err = p.parseIdent("savepoint name"); err != nil { return &stmt, err } } return &stmt, nil } func (p *Parser) parseSavepointStatement() (_ *SavepointStatement, err error) { assert(p.peek() == SAVEPOINT) var stmt SavepointStatement stmt.Savepoint, _, _ = p.scan() if stmt.Name, err = p.parseIdent("savepoint name"); err != nil { return &stmt, err } return &stmt, nil } func (p *Parser) parseReleaseStatement() (_ *ReleaseStatement, err error) { assert(p.peek() == RELEASE) var stmt ReleaseStatement stmt.Release, _, _ = p.scan() if p.peek() == SAVEPOINT { stmt.Savepoint, _, _ = p.scan() } if stmt.Name, err = p.parseIdent("savepoint name"); err != nil { return &stmt, err } return &stmt, nil } func (p *Parser) parseCreateStatement() (Statement, error) { assert(p.peek() == CREATE) pos, tok, _ := p.scan() switch p.peek() { case TABLE: return p.parseCreateTableStatement(pos) case VIEW: return p.parseCreateViewStatement(pos) case INDEX, UNIQUE: return p.parseCreateIndexStatement(pos) case TRIGGER: return p.parseCreateTriggerStatement(pos) default: return nil, p.errorExpected(pos, tok, "TABLE, VIEW, INDEX, TRIGGER") } } func (p *Parser) parseDropStatement() (Statement, error) { assert(p.peek() == DROP) pos, tok, _ := p.scan() switch p.peek() { case TABLE: return p.parseDropTableStatement(pos) case VIEW: return p.parseDropViewStatement(pos) case INDEX: return p.parseDropIndexStatement(pos) case TRIGGER: return p.parseDropTriggerStatement(pos) default: return nil, p.errorExpected(pos, tok, "TABLE, VIEW, INDEX, or TRIGGER") } } func (p *Parser) parseCreateTableStatement(createPos Pos) (_ *CreateTableStatement, err error) { assert(p.peek() == TABLE) var stmt CreateTableStatement stmt.Create = createPos stmt.Table, _, _ = p.scan() // Parse optional "IF NOT EXISTS". if p.peek() == IF { stmt.If, _, _ = p.scan() pos, tok, _ := p.scan() if tok != NOT { return &stmt, p.errorExpected(pos, tok, "NOT") } stmt.IfNot = pos pos, tok, _ = p.scan() if tok != EXISTS { return &stmt, p.errorExpected(pos, tok, "EXISTS") } stmt.IfNotExists = pos } if stmt.Name, err = p.parseIdent("table name"); err != nil { return &stmt, err } // Parse either a column/constraint list or build table from "AS