package api import ( "encoding/json" "fmt" "net/http" "strconv" "strings" "unicode" "github.com/restxlsx/restxlsx/internal/auth" "github.com/restxlsx/restxlsx/internal/excel" "github.com/restxlsx/restxlsx/internal/logger" ) type token struct { kind tokenKind value string } type tokenKind int const ( tokEOF tokenKind = iota tokIdent tokString tokNumber tokStar tokComma tokLParen tokRParen tokEq tokSemicolon tokKeyword ) type sqlStmt struct { kind string // SELECT, INSERT, UPDATE, DELETE table string columns []string values []string sets map[string]string whereCol string whereVal string } type SQLHandler struct { engine *excel.Engine log *logger.Logger } func NewSQLHandler(engine *excel.Engine, log *logger.Logger) *SQLHandler { return &SQLHandler{engine: engine, log: log} } func (h *SQLHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { var body struct { Query string `json:"query"` } if err := json.NewDecoder(r.Body).Decode(&body); err != nil { writeJSON(w, 400, map[string]string{"error": "invalid JSON: " + err.Error()}) return } stmt, err := parseSQL(body.Query) if err != nil { writeJSON(w, 400, map[string]string{"error": "parse error: " + err.Error()}) return } role := auth.RoleFromContext(r.Context()) switch stmt.kind { case "SELECT": case "INSERT", "UPDATE", "DELETE": if !role.CanWrite() { writeJSON(w, 403, map[string]string{"error": "forbidden: write access required"}) return } default: writeJSON(w, 400, map[string]string{"error": "unsupported statement: " + stmt.kind}) return } h.engine.Lock() result, err := h.execute(stmt) h.engine.Unlock() if err != nil { writeJSON(w, 400, map[string]string{"error": err.Error()}) return } writeJSON(w, 200, result) } func (h *SQLHandler) execute(stmt sqlStmt) (any, error) { table := h.engine.GetTable(stmt.table) if table == nil { return nil, fmt.Errorf("table %q not found", stmt.table) } switch stmt.kind { case "SELECT": return h.executeSelect(table, stmt) case "INSERT": return h.executeInsert(table, stmt) case "UPDATE": return h.executeUpdate(table, stmt) case "DELETE": return h.executeDelete(table, stmt) } return nil, fmt.Errorf("unsupported: %s", stmt.kind) } func (h *SQLHandler) executeSelect(table *excel.Table, stmt sqlStmt) (any, error) { var results []map[string]any for _, row := range table.Rows { if stmt.whereCol != "" { if fmt.Sprintf("%v", row[stmt.whereCol]) != stmt.whereVal { continue } } if len(stmt.columns) == 1 && stmt.columns[0] == "*" { results = append(results, row) } else { projected := make(map[string]any) for _, col := range stmt.columns { if v, ok := row[col]; ok { projected[col] = v } } results = append(results, projected) } } if stmt.whereCol != "" && len(results) == 1 { return results[0], nil } if results == nil { return []map[string]any{}, nil } return results, nil } func (h *SQLHandler) executeInsert(table *excel.Table, stmt sqlStmt) (any, error) { if len(stmt.columns) != len(stmt.values) { return nil, fmt.Errorf("column count (%d) != value count (%d)", len(stmt.columns), len(stmt.values)) } rec := make(map[string]any) for i, col := range stmt.columns { rec[col] = parseValue(stmt.values[i]) } idCol := h.engine.IDColumn(table.Name) if idCol != "" { if _, ok := rec[idCol]; !ok { rec[idCol] = len(table.Rows) + 1 } } table.Rows = append(table.Rows, rec) if err := h.engine.Flush(); err != nil { return nil, err } return rec, nil } func (h *SQLHandler) executeUpdate(table *excel.Table, stmt sqlStmt) (any, error) { var updated []map[string]any for i, row := range table.Rows { if stmt.whereCol != "" { if fmt.Sprintf("%v", row[stmt.whereCol]) != stmt.whereVal { continue } } for k, v := range stmt.sets { table.Rows[i][k] = parseValue(v) } updated = append(updated, table.Rows[i]) } if err := h.engine.Flush(); err != nil { return nil, err } if updated == nil { return []map[string]any{}, nil } return updated, nil } func (h *SQLHandler) executeDelete(table *excel.Table, stmt sqlStmt) (any, error) { var deleted []map[string]any kept := make([]map[string]any, 0, len(table.Rows)) for _, row := range table.Rows { if stmt.whereCol != "" { if fmt.Sprintf("%v", row[stmt.whereCol]) == stmt.whereVal { deleted = append(deleted, row) continue } } else { deleted = append(deleted, row) continue } kept = append(kept, row) } table.Rows = kept if err := h.engine.Flush(); err != nil { return nil, err } return map[string]any{"deleted": len(deleted)}, nil } // ─── SQL Parser ──────────────────────────────────────────────── type lexer struct { input string pos int } func (l *lexer) peek() byte { if l.pos >= len(l.input) { return 0 } return l.input[l.pos] } func (l *lexer) advance() byte { b := l.input[l.pos] l.pos++ return b } func (l *lexer) skipWS() { for l.pos < len(l.input) && (l.input[l.pos] == ' ' || l.input[l.pos] == '\t' || l.input[l.pos] == '\n') { l.pos++ } } var keywords = map[string]bool{ "select": true, "from": true, "where": true, "insert": true, "into": true, "values": true, "update": true, "set": true, "delete": true, } func lex(input string) ([]token, error) { l := &lexer{input: input} var toks []token for l.pos < len(l.input) { l.skipWS() if l.pos >= len(l.input) { break } c := l.peek() switch { case c == ',': toks = append(toks, token{tokComma, ","}) l.advance() case c == '(': toks = append(toks, token{tokLParen, "("}) l.advance() case c == ')': toks = append(toks, token{tokRParen, ")"}) l.advance() case c == ';': toks = append(toks, token{tokSemicolon, ";"}) l.advance() case c == '=': toks = append(toks, token{tokEq, "="}) l.advance() case c == '*': toks = append(toks, token{tokStar, "*"}) l.advance() case c == '\'' || c == '"': quote := l.advance() var buf strings.Builder for l.pos < len(l.input) && l.peek() != quote { buf.WriteByte(l.advance()) } if l.pos >= len(l.input) { return nil, fmt.Errorf("unterminated string literal") } l.advance() toks = append(toks, token{tokString, buf.String()}) case c >= '0' && c <= '9' || c == '-': var buf strings.Builder for l.pos < len(l.input) && (l.peek() >= '0' && l.peek() <= '9' || l.peek() == '.') { buf.WriteByte(l.advance()) } toks = append(toks, token{tokNumber, buf.String()}) case unicode.IsLetter(rune(c)) || c == '_': var buf strings.Builder for l.pos < len(l.input) && (unicode.IsLetter(rune(l.peek())) || unicode.IsDigit(rune(l.peek())) || l.peek() == '_') { buf.WriteByte(l.advance()) } word := buf.String() if keywords[strings.ToLower(word)] { toks = append(toks, token{tokKeyword, strings.ToUpper(word)}) } else { toks = append(toks, token{tokIdent, word}) } default: return nil, fmt.Errorf("unexpected character: %c", c) } } if len(toks) == 0 { return nil, fmt.Errorf("empty query") } toks = append(toks, token{tokEOF, ""}) return toks, nil } func parseSQL(input string) (sqlStmt, error) { toks, err := lex(input) if err != nil { return sqlStmt{}, err } pos := 0 next := func() token { if pos >= len(toks) { return token{tokEOF, ""} } t := toks[pos] pos++ return t } peek := func() token { if pos >= len(toks) { return token{tokEOF, ""} } return toks[pos] } expect := func(kind tokenKind, msg string) (token, error) { t := next() if t.kind != kind { return t, fmt.Errorf("%s: expected %d got %q", msg, kind, t.value) } return t, nil } first := next() if first.kind != tokKeyword { return sqlStmt{}, fmt.Errorf("expected keyword, got %q", first.value) } switch first.value { case "SELECT": return parseSelect(&pos, toks, next, peek, expect) case "INSERT": return parseInsert(&pos, toks, next, peek, expect) case "UPDATE": return parseUpdate(&pos, toks, next, peek, expect) case "DELETE": return parseDelete(&pos, toks, next, peek, expect) default: return sqlStmt{}, fmt.Errorf("unsupported statement: %s", first.value) } } func parseSelect(pos *int, toks []token, next, peek func() token, expect func(tokenKind, string) (token, error)) (sqlStmt, error) { stmt := sqlStmt{kind: "SELECT"} if peek().kind == tokStar { next() stmt.columns = []string{"*"} } else { for { t, err := expect(tokIdent, "select column") if err != nil { return stmt, err } stmt.columns = append(stmt.columns, t.value) if peek().kind != tokComma { break } next() } } if _, err := expect(tokKeyword, "FROM"); err != nil { return stmt, err } t, err := expect(tokIdent, "table name") if err != nil { return stmt, err } stmt.table = t.value if peek().kind == tokKeyword && strings.ToUpper(peek().value) == "WHERE" { next() t, err := expect(tokIdent, "where column") if err != nil { return stmt, err } stmt.whereCol = t.value if _, err := expect(tokEq, "="); err != nil { return stmt, err } v := next() if v.kind != tokString && v.kind != tokNumber { return stmt, fmt.Errorf("expected value in WHERE, got %q", v.value) } stmt.whereVal = v.value } return stmt, nil } func parseInsert(pos *int, toks []token, next, peek func() token, expect func(tokenKind, string) (token, error)) (sqlStmt, error) { stmt := sqlStmt{kind: "INSERT"} if _, err := expect(tokKeyword, "INTO"); err != nil { return stmt, err } t, err := expect(tokIdent, "table name") if err != nil { return stmt, err } stmt.table = t.value if _, err := expect(tokLParen, "("); err != nil { return stmt, err } for { t, err := expect(tokIdent, "column name") if err != nil { return stmt, err } stmt.columns = append(stmt.columns, t.value) if peek().kind != tokComma { break } next() } if _, err := expect(tokRParen, ")"); err != nil { return stmt, err } if _, err := expect(tokKeyword, "VALUES"); err != nil { return stmt, err } if _, err := expect(tokLParen, "("); err != nil { return stmt, err } for { v := next() if v.kind != tokString && v.kind != tokNumber { return stmt, fmt.Errorf("expected value, got %q", v.value) } stmt.values = append(stmt.values, v.value) if peek().kind != tokComma { break } next() } if _, err := expect(tokRParen, ")"); err != nil { return stmt, err } return stmt, nil } func parseUpdate(pos *int, toks []token, next, peek func() token, expect func(tokenKind, string) (token, error)) (sqlStmt, error) { stmt := sqlStmt{kind: "UPDATE", sets: make(map[string]string)} t, err := expect(tokIdent, "table name") if err != nil { return stmt, err } stmt.table = t.value if _, err := expect(tokKeyword, "SET"); err != nil { return stmt, err } for { t, err := expect(tokIdent, "column name") if err != nil { return stmt, err } col := t.value if _, err := expect(tokEq, "="); err != nil { return stmt, err } v := next() if v.kind != tokString && v.kind != tokNumber { return stmt, fmt.Errorf("expected value in SET, got %q", v.value) } stmt.sets[col] = v.value if peek().kind != tokComma { break } next() } if peek().kind == tokKeyword && strings.ToUpper(peek().value) == "WHERE" { next() t, err := expect(tokIdent, "where column") if err != nil { return stmt, err } stmt.whereCol = t.value if _, err := expect(tokEq, "="); err != nil { return stmt, err } v := next() if v.kind != tokString && v.kind != tokNumber { return stmt, fmt.Errorf("expected value in WHERE, got %q", v.value) } stmt.whereVal = v.value } return stmt, nil } func parseDelete(pos *int, toks []token, next, peek func() token, expect func(tokenKind, string) (token, error)) (sqlStmt, error) { stmt := sqlStmt{kind: "DELETE"} if _, err := expect(tokKeyword, "FROM"); err != nil { return stmt, err } t, err := expect(tokIdent, "table name") if err != nil { return stmt, err } stmt.table = t.value if peek().kind == tokKeyword && strings.ToUpper(peek().value) == "WHERE" { next() t, err := expect(tokIdent, "where column") if err != nil { return stmt, err } stmt.whereCol = t.value if _, err := expect(tokEq, "="); err != nil { return stmt, err } v := next() if v.kind != tokString && v.kind != tokNumber { return stmt, fmt.Errorf("expected value in WHERE, got %q", v.value) } stmt.whereVal = v.value } return stmt, nil } func parseValue(s string) any { if i, err := strconv.Atoi(s); err == nil { return i } if f, err := strconv.ParseFloat(s, 64); err == nil { return f } return s }