// Copyright 2021 Molecula Corp. All rights reserved. package parser import ( "fmt" "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.parseAlterStatement() case BULK: return p.parseBulkInsertStatement() case CREATE: return p.parseCreateStatement() case DROP: return p.parseDropStatement() case SELECT: return p.parseSelectStatement(false, nil) case INSERT, REPLACE: return p.parseInsertStatement(nil) case UPDATE: return p.parseUpdateStatement(nil) case DELETE: return p.parseDeleteStatement() // case WITH: // return p.parseWithStatement() case SHOW: return p.parseShowStatement() 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) parseShowStatement() (Statement, error) { assert(p.peek() == SHOW) show, _, _ := p.scan() switch p.peek() { case DATABASES: return p.parseShowDatabasesStatement(show) case TABLES: return p.parseShowTablesStatement(show) case COLUMNS: return p.parseShowColumnsStatement(show) case CREATE: return p.parseShowCreateStatement(show) default: return nil, p.errorExpected(p.pos, p.tok, "DATABASES, TABLES, COLUMNS or CREATE") } } func (p *Parser) parseShowDatabasesStatement(showPos Pos) (*ShowDatabasesStatement, error) { switch p.peek() { case DATABASES: var stmt ShowDatabasesStatement stmt.Show = showPos stmt.Databases, _, _ = p.scan() return &stmt, nil default: return nil, p.errorExpected(p.pos, p.tok, "DATABASES") } } func (p *Parser) parseShowTablesStatement(showPos Pos) (*ShowTablesStatement, error) { switch p.peek() { case TABLES: var stmt ShowTablesStatement stmt.Show = showPos stmt.Tables, _, _ = p.scan() return &stmt, nil default: return nil, p.errorExpected(p.pos, p.tok, "TABLES") } } func (p *Parser) parseShowColumnsStatement(showPos Pos) (_ *ShowColumnsStatement, err error) { assert(p.peek() == COLUMNS) columns, _, _ := p.scan() var stmt ShowColumnsStatement stmt.Show = showPos stmt.Columns = columns switch p.peek() { case FROM: stmt.From, _, _ = p.scan() if stmt.TableName, err = p.parseIdent("table name"); err != nil { return &stmt, err } return &stmt, nil default: return nil, p.errorExpected(p.pos, p.tok, "FROM") } } func (p *Parser) parseShowCreateStatement(showPos Pos) (Statement, error) { assert(p.peek() == CREATE) create, _, _ := p.scan() switch p.peek() { case TABLE: return p.parseShowCreateTableStatement(showPos, create) default: return nil, p.errorExpected(p.pos, p.tok, "TABLES") } } func (p *Parser) parseShowCreateTableStatement(showPos Pos, createPos Pos) (_ *ShowCreateTableStatement, err error) { assert(p.peek() == TABLE) table, _, _ := p.scan() var stmt ShowCreateTableStatement stmt.Show = showPos stmt.Create = createPos stmt.Table = table if stmt.TableName, err = p.parseIdent("table name"); err != nil { return &stmt, err } return &stmt, nil } /*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 DATABASE: return p.parseCreateDatabaseStatement(pos) case TABLE: return p.parseCreateTableStatement(pos) case VIEW: return p.parseCreateViewStatement(pos) /*case INDEX, UNIQUE: return p.parseCreateIndexStatement(pos)*/ case FUNCTION: return p.parseCreateFunctionStatement(pos) default: return nil, p.errorExpected(pos, tok, "DATABASE, TABLE, VIEW or FUNCTION") } } func (p *Parser) parseAlterStatement() (Statement, error) { assert(p.peek() == ALTER) pos, tok, _ := p.scan() switch p.peek() { case DATABASE: return p.parseAlterDatabaseStatement(pos) case TABLE: return p.parseAlterTableStatement(pos) case VIEW: return p.parseAlterViewStatement(pos) default: return nil, p.errorExpected(pos, tok, "DATABASE, TABLE or VIEW") } } func (p *Parser) parseDropStatement() (Statement, error) { assert(p.peek() == DROP) pos, tok, _ := p.scan() switch p.peek() { case DATABASE: return p.parseDropDatabaseStatement(pos) case TABLE: return p.parseDropTableStatement(pos) case VIEW: return p.parseDropViewStatement(pos) /* case INDEX: return p.parseDropIndexStatement(pos)*/ case FUNCTION: return p.parseDropFunctionStatement(pos) default: return nil, p.errorExpected(pos, tok, "DATABASE, TABLE, VIEW or FUNCTION") } } func (p *Parser) parseCreateDatabaseStatement(createPos Pos) (_ *CreateDatabaseStatement, err error) { assert(p.peek() == DATABASE) var stmt CreateDatabaseStatement stmt.Create = createPos stmt.Database, _, _ = 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("database name"); err != nil { return &stmt, err } switch p.peek() { case WITH: stmt.With, _, _ = p.scan() // look for database options if stmt.Options, err = p.parseDatabaseOptions(); err != nil { return &stmt, err } if len(stmt.Options) == 0 { return &stmt, p.errorExpected(stmt.With, p.peek(), "at least one option after WITH") } } return &stmt, nil } func (p *Parser) parseDatabaseOptions() (_ []DatabaseOption, err error) { if !isDatabaseOptionStartToken(p.peek()) { return nil, nil } var a []DatabaseOption for { if !isDatabaseOptionStartToken(p.peek()) { return a, nil } cons, err := p.parseDatabaseOption() if cons != nil { a = append(a, cons) } if err != nil { return a, err } } } func (p *Parser) parseDatabaseOption() (_ DatabaseOption, err error) { assert(isDatabaseOptionStartToken(p.peek())) // Parse database options. switch p.peek() { case UNITS: return p.parseUnitsOption() default: assert(p.peek() == COMMENT) return p.parseCommentOption() } } func (p *Parser) parseUnitsOption() (_ *UnitsOption, err error) { assert(p.peek() == UNITS) var opt UnitsOption opt.Units, _, _ = p.scan() if isLiteralToken(p.peek()) { opt.Expr = p.mustParseLiteral() } else { return &opt, p.errorExpected(p.pos, p.tok, "literal") } return &opt, nil } 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