diff --git a/go-backend/internal/http/handler/control_plane.go b/go-backend/internal/http/handler/control_plane.go index 0c3cf5a..0e3d084 100644 --- a/go-backend/internal/http/handler/control_plane.go +++ b/go-backend/internal/http/handler/control_plane.go @@ -889,11 +889,11 @@ func firstPortFromRange(portRange string) int { func (h *Handler) listChainNodesForTunnel(tunnelID int64) ([]chainNodeRecord, error) { rows, err := h.repo.DB().Query(` - SELECT ct.chain_type, COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy + SELECT CAST(ct.chain_type AS INTEGER), COALESCE(ct.inx, 0), ct.node_id, COALESCE(ct.port, 0), n.name, ct.protocol, ct.strategy FROM chain_tunnel ct LEFT JOIN node n ON n.id = ct.node_id WHERE ct.tunnel_id = ? - ORDER BY ct.chain_type ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC + ORDER BY CAST(ct.chain_type AS INTEGER) ASC, COALESCE(ct.inx, 0) ASC, ct.id ASC `, tunnelID) if err != nil { return nil, err diff --git a/go-backend/internal/http/handler/federation.go b/go-backend/internal/http/handler/federation.go index 1825716..4997ca2 100644 --- a/go-backend/internal/http/handler/federation.go +++ b/go-backend/internal/http/handler/federation.go @@ -799,7 +799,7 @@ func (h *Handler) federationTunnelCreate(w http.ResponseWriter, r *http.Request) return } - _, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, 1, ?, ?, 'fifo', 0, ?)`, + _, err = tx.Exec(`INSERT INTO chain_tunnel (tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES (?, '1', ?, ?, 'fifo', 0, ?)`, tunnelID, share.NodeID, req.RemotePort, diff --git a/go-backend/internal/http/handler/mutations.go b/go-backend/internal/http/handler/mutations.go index 4ee0407..78f7cd9 100644 --- a/go-backend/internal/http/handler/mutations.go +++ b/go-backend/internal/http/handler/mutations.go @@ -2846,7 +2846,7 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac if nodeID <= 0 { continue } - _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 1, ?, NULL, NULL, 0, ?)`, + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '1', ?, NULL, NULL, 0, ?)`, tunnelID, nodeID, defaultString(asString(n["protocol"]), "tls")) if err != nil { return err @@ -2865,7 +2865,7 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac return pickErr } } - _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 3, ?, ?, ?, 0, ?)`, + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '3', ?, ?, ?, 0, ?)`, tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), defaultString(asString(n["protocol"]), "tls")) if err != nil { return err @@ -2886,7 +2886,7 @@ func replaceTunnelChainsTx(tx *store.Tx, tunnelID int64, req map[string]interfac return pickErr } } - _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, 2, ?, ?, ?, ?, ?)`, + _, err := tx.Exec(`INSERT INTO chain_tunnel(tunnel_id, chain_type, node_id, port, strategy, inx, protocol) VALUES(?, '2', ?, ?, ?, ?, ?)`, tunnelID, nodeID, port, defaultString(asString(n["strategy"]), "round"), i+1, defaultString(asString(n["protocol"]), "tls")) if err != nil { return err @@ -2972,7 +2972,7 @@ func (h *Handler) batchForwardStatus(ids []int64, status int) (int, int) { } func (h *Handler) tunnelEntryNodeIDs(tunnelID int64) ([]int64, error) { - rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = 1 ORDER BY inx ASC, id ASC`, tunnelID) + rows, err := h.repo.DB().Query(`SELECT node_id FROM chain_tunnel WHERE tunnel_id = ? AND chain_type = '1' ORDER BY inx ASC, id ASC`, tunnelID) if err != nil { return nil, err } diff --git a/go-backend/internal/store/db.go b/go-backend/internal/store/db.go index 50107c3..a597eef 100644 --- a/go-backend/internal/store/db.go +++ b/go-backend/internal/store/db.go @@ -4,7 +4,7 @@ package store import ( "database/sql" - "fmt" + "strconv" "strings" ) @@ -98,7 +98,7 @@ func (db *DB) Begin() (*Tx, error) { func (db *DB) ExecReturningID(query string, args ...any) (int64, error) { q := db.rewrite(query) if db.dialect == DialectPostgres { - q = strings.TrimRight(q, "; \t\n") + " RETURNING id" + q = ensureReturningID(q) var id int64 if err := db.raw.QueryRow(q, args...).Scan(&id); err != nil { return 0, err @@ -143,7 +143,7 @@ func (tx *Tx) Rollback() error { return tx.raw.Rollback() } func (tx *Tx) ExecReturningID(query string, args ...any) (int64, error) { q := rewriteQuery(tx.dialect, query) if tx.dialect == DialectPostgres { - q = strings.TrimRight(q, "; \t\n") + " RETURNING id" + q = ensureReturningID(q) var id int64 if err := tx.raw.QueryRow(q, args...).Scan(&id); err != nil { return 0, err @@ -174,35 +174,15 @@ func rewriteQuery(dialect Dialect, query string) string { func rewriteUserIdentifier(query string) string { var buf strings.Builder buf.Grow(len(query) + 16) - inSingle := false - inDouble := false i := 0 for i < len(query) { - ch := query[i] - if ch == '\'' && !inDouble { - if inSingle && i+1 < len(query) && query[i+1] == '\'' { - buf.WriteByte(ch) - buf.WriteByte(query[i+1]) - i += 2 - continue - } - inSingle = !inSingle - buf.WriteByte(ch) - i++ - continue - } - if ch == '"' && !inSingle { - inDouble = !inDouble - buf.WriteByte(ch) - i++ - continue - } - if inSingle || inDouble { - buf.WriteByte(ch) - i++ + if end, ok := skipSQLProtectedSegment(query, i); ok { + buf.WriteString(query[i:end]) + i = end continue } + ch := query[i] if isIdentifierChar(ch) { j := i + 1 for j < len(query) && isIdentifierChar(query[j]) { @@ -238,39 +218,43 @@ func isIdentifierChar(ch byte) bool { } func rewriteInsertOrIgnore(query string) string { - upper := strings.ToUpper(query) - idx := strings.Index(upper, "INSERT OR IGNORE INTO") - if idx < 0 { + start, end, ok := findKeywordSequenceOutside(query, []string{"INSERT", "OR", "IGNORE", "INTO"}, 0) + if !ok { return query } - prefix := query[:idx] - suffix := query[idx+len("INSERT OR IGNORE INTO"):] - result := prefix + "INSERT INTO" + suffix - trimmed := strings.TrimRight(result, "; \t\n") - return trimmed + " ON CONFLICT DO NOTHING" + rewritten := query[:start] + "INSERT INTO" + query[end:] + rewritten = strings.TrimRight(rewritten, "; \t\n") + + insertIntoEnd := start + len("INSERT INTO") + if _, _, hasOnConflict := findKeywordSequenceOutside(rewritten, []string{"ON", "CONFLICT"}, insertIntoEnd); hasOnConflict { + return rewritten + } + + if retStart, _, hasReturning := findKeywordSequenceOutside(rewritten, []string{"RETURNING"}, insertIntoEnd); hasReturning { + prefix := strings.TrimRight(rewritten[:retStart], " \t\n") + suffix := strings.TrimLeft(rewritten[retStart:], " \t\n") + return prefix + " ON CONFLICT DO NOTHING " + suffix + } + + return rewritten + " ON CONFLICT DO NOTHING" } func rewritePlaceholders(query string) string { var buf strings.Builder buf.Grow(len(query) + 16) n := 1 - inString := false for i := 0; i < len(query); i++ { - ch := query[i] - if ch == '\'' { - if inString && i+1 < len(query) && query[i+1] == '\'' { - buf.WriteByte(ch) - buf.WriteByte(query[i+1]) - i++ - continue - } - inString = !inString - buf.WriteByte(ch) + if end, ok := skipSQLProtectedSegment(query, i); ok { + buf.WriteString(query[i:end]) + i = end - 1 continue } - if ch == '?' && !inString { - buf.WriteString(fmt.Sprintf("$%d", n)) + + ch := query[i] + if ch == '?' { + buf.WriteByte('$') + buf.WriteString(strconv.Itoa(n)) n++ continue } @@ -278,3 +262,205 @@ func rewritePlaceholders(query string) string { } return buf.String() } + +func ensureReturningID(query string) string { + trimmed := strings.TrimRight(query, "; \t\n") + if _, _, ok := findKeywordSequenceOutside(trimmed, []string{"RETURNING"}, 0); ok { + return trimmed + } + return trimmed + " RETURNING id" +} + +func findKeywordSequenceOutside(query string, keywords []string, from int) (int, int, bool) { + if len(keywords) == 0 { + return 0, 0, false + } + if from < 0 { + from = 0 + } + if from >= len(query) { + return 0, 0, false + } + + matched := 0 + seqStart := -1 + + for i := from; i < len(query); { + if end, ok := skipSQLProtectedSegment(query, i); ok { + i = end + continue + } + + ch := query[i] + if isIdentifierChar(ch) { + j := i + 1 + for j < len(query) && isIdentifierChar(query[j]) { + j++ + } + tok := query[i:j] + + if strings.EqualFold(tok, keywords[matched]) { + if matched == 0 { + seqStart = i + } + matched++ + if matched == len(keywords) { + return seqStart, j, true + } + } else if strings.EqualFold(tok, keywords[0]) { + seqStart = i + matched = 1 + } else { + matched = 0 + seqStart = -1 + } + + i = j + continue + } + + if !isSQLSpace(ch) { + matched = 0 + seqStart = -1 + } + i++ + } + + return 0, 0, false +} + +func skipSQLProtectedSegment(query string, i int) (int, bool) { + if i < 0 || i >= len(query) { + return 0, false + } + + switch query[i] { + case '\'': + return skipSingleQuotedLiteral(query, i), true + case '"': + return skipDoubleQuotedIdentifier(query, i), true + case '-': + if i+1 < len(query) && query[i+1] == '-' { + return skipLineComment(query, i), true + } + case '/': + if i+1 < len(query) && query[i+1] == '*' { + return skipBlockComment(query, i), true + } + case '$': + if end, ok := skipDollarQuotedLiteral(query, i); ok { + return end, true + } + } + + return 0, false +} + +func skipSingleQuotedLiteral(query string, i int) int { + for j := i + 1; j < len(query); j++ { + if query[j] != '\'' { + continue + } + if j+1 < len(query) && query[j+1] == '\'' { + j++ + continue + } + return j + 1 + } + return len(query) +} + +func skipDoubleQuotedIdentifier(query string, i int) int { + for j := i + 1; j < len(query); j++ { + if query[j] != '"' { + continue + } + if j+1 < len(query) && query[j+1] == '"' { + j++ + continue + } + return j + 1 + } + return len(query) +} + +func skipLineComment(query string, i int) int { + for j := i + 2; j < len(query); j++ { + if query[j] == '\n' { + return j + } + } + return len(query) +} + +func skipBlockComment(query string, i int) int { + depth := 1 + for j := i + 2; j < len(query)-1; j++ { + if query[j] == '/' && query[j+1] == '*' { + depth++ + j++ + continue + } + if query[j] == '*' && query[j+1] == '/' { + depth-- + j++ + if depth == 0 { + return j + 1 + } + } + } + return len(query) +} + +func skipDollarQuotedLiteral(query string, i int) (int, bool) { + if i < 0 || i >= len(query) || query[i] != '$' { + return 0, false + } + + if i+1 >= len(query) { + return 0, false + } + + var endTag int + if query[i+1] == '$' { + endTag = i + 1 + } else { + if !isDollarTagStart(query[i+1]) { + return 0, false + } + j := i + 2 + for j < len(query) && isDollarTagChar(query[j]) { + j++ + } + if j >= len(query) || query[j] != '$' { + return 0, false + } + endTag = j + } + + tag := query[i : endTag+1] + if closeIdx := strings.Index(query[endTag+1:], tag); closeIdx >= 0 { + return endTag + 1 + closeIdx + len(tag), true + } + return len(query), true +} + +func isDollarTagStart(ch byte) bool { + return ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') +} + +func isDollarTagChar(ch byte) bool { + if isDollarTagStart(ch) { + return true + } + return ch >= '0' && ch <= '9' +} + +func isSQLSpace(ch byte) bool { + switch ch { + case ' ', '\t', '\n', '\r', '\f': + return true + default: + return false + } +} diff --git a/go-backend/internal/store/db_test.go b/go-backend/internal/store/db_test.go new file mode 100644 index 0000000..4a3f580 --- /dev/null +++ b/go-backend/internal/store/db_test.go @@ -0,0 +1,116 @@ +package store + +import "testing" + +func TestRewritePlaceholdersSkipsProtectedSegments(t *testing.T) { + q := `SELECT ?, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ? +FROM t /* block ? */ WHERE id = ?` + got := rewritePlaceholders(q) + want := `SELECT $1, '?', "id?", $$body ? $$, $tag$X?$tag$, col -- comment ? +FROM t /* block ? */ WHERE id = $2` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewriteInsertOrIgnoreBasic(t *testing.T) { + q := `INSERT OR IGNORE INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?)` + got := rewriteInsertOrIgnore(q) + want := `INSERT INTO user_group_user(user_group_id, user_id, created_time) VALUES(?, ?, ?) ON CONFLICT DO NOTHING` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewriteInsertOrIgnoreBeforeReturning(t *testing.T) { + q := `INSERT OR IGNORE INTO x(a) VALUES(?) RETURNING id` + got := rewriteInsertOrIgnore(q) + want := `INSERT INTO x(a) VALUES(?) ON CONFLICT DO NOTHING RETURNING id` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewriteInsertOrIgnoreNotDuplicatingOnConflict(t *testing.T) { + q := `INSERT OR IGNORE INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a` + got := rewriteInsertOrIgnore(q) + want := `INSERT INTO x(a) VALUES(?) ON CONFLICT(a) DO UPDATE SET a=excluded.a` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestEnsureReturningID(t *testing.T) { + if got := ensureReturningID(`INSERT INTO x(a) VALUES($1)`); got != `INSERT INTO x(a) VALUES($1) RETURNING id` { + t.Fatalf("missing RETURNING append: %s", got) + } + if got := ensureReturningID(`INSERT INTO x(a) VALUES($1) RETURNING other_id`); got != `INSERT INTO x(a) VALUES($1) RETURNING other_id` { + t.Fatalf("RETURNING should not be duplicated: %s", got) + } +} + +func TestRewriteUserIdentifierSafety(t *testing.T) { + q := `SELECT user, user_id, 'user', "user", note FROM user -- user +WHERE owner='user'` + got := rewriteUserIdentifier(q) + want := `SELECT "user", user_id, 'user', "user", note FROM "user" -- user +WHERE owner='user'` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewriteQueryPostgresPipeline(t *testing.T) { + q := `INSERT OR IGNORE INTO user(name, note) VALUES(?, '?')` + got := rewriteQuery(DialectPostgres, q) + want := `INSERT INTO "user"(name, note) VALUES($1, '?') ON CONFLICT DO NOTHING` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewriteInsertOrIgnoreSkipsStringLiteral(t *testing.T) { + q := `SELECT 'INSERT OR IGNORE INTO t(a) VALUES(?)' AS q` + got := rewriteInsertOrIgnore(q) + if got != q { + t.Fatalf("string literal should stay unchanged\nwant: %s\ngot: %s", q, got) + } +} + +func TestRewriteInsertOrIgnoreSkipsCommentedKeyword(t *testing.T) { + q := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?) +INSERT OR IGNORE INTO real_t(a) VALUES(?)` + got := rewriteInsertOrIgnore(q) + want := `-- INSERT OR IGNORE INTO ignored(a) VALUES(?) +INSERT INTO real_t(a) VALUES(?) ON CONFLICT DO NOTHING` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewritePlaceholdersSkipsNestedBlockComment(t *testing.T) { + q := `SELECT ? /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = ?` + got := rewritePlaceholders(q) + want := `SELECT $1 /* outer ? /* inner ? */ still_outer ? */ FROM t WHERE id = $2` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewritePlaceholdersSkipsUnterminatedBlockComment(t *testing.T) { + q := `SELECT ? /* unterminated ? comment` + got := rewritePlaceholders(q) + want := `SELECT $1 /* unterminated ? comment` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} + +func TestRewriteUserIdentifierSkipsDollarQuotedAndComment(t *testing.T) { + q := `SELECT user, $$user ?$$ AS body, col FROM user /* user */ -- user` + got := rewriteUserIdentifier(q) + want := `SELECT "user", $$user ?$$ AS body, col FROM "user" /* user */ -- user` + if got != want { + t.Fatalf("unexpected rewrite\nwant: %s\ngot: %s", want, got) + } +} diff --git a/go-backend/internal/store/sqlite/repository.go b/go-backend/internal/store/sqlite/repository.go index c5588d5..3a6232b 100644 --- a/go-backend/internal/store/sqlite/repository.go +++ b/go-backend/internal/store/sqlite/repository.go @@ -897,9 +897,9 @@ func (r *Repository) ListTunnels() ([]map[string]interface{}, error) { } chainRows, err := r.db.Query(` - SELECT tunnel_id, chain_type, node_id, protocol, strategy, COALESCE(inx, 0) + SELECT tunnel_id, CAST(chain_type AS INTEGER), node_id, protocol, strategy, COALESCE(inx, 0) FROM chain_tunnel - ORDER BY tunnel_id ASC, chain_type ASC, inx ASC, id ASC + ORDER BY tunnel_id ASC, CAST(chain_type AS INTEGER) ASC, inx ASC, id ASC `) if err != nil { return nil, err