549 lines
12 KiB
Go
549 lines
12 KiB
Go
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
|
|
}
|