From c9345b9bc80230dec4b08bbd30019d6828b181a0 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 26 Aug 2026 22:31:08 +0900 Subject: [PATCH 1/5] feat: classify SQL statements for read-only enforcement --- internal/sqlscan/scan.go | 150 ++++++++++++++++++++++ internal/sqlscan/sqlscan.go | 214 +++++++++++++++++++++++++++++++ internal/sqlscan/sqlscan_test.go | 111 ++++++++++++++++ 3 files changed, 475 insertions(+) create mode 100644 internal/sqlscan/scan.go create mode 100644 internal/sqlscan/sqlscan.go create mode 100644 internal/sqlscan/sqlscan_test.go diff --git a/internal/sqlscan/scan.go b/internal/sqlscan/scan.go new file mode 100644 index 0000000..5a571f0 --- /dev/null +++ b/internal/sqlscan/scan.go @@ -0,0 +1,150 @@ +package sqlscan + +import "strings" + +// statements splits sql into top-level statements, each a slice of the +// uppercased unquoted words it contains. Strings, comments, dollar-quoted +// bodies, and quoted identifiers are skipped so their contents never look like +// keywords; numbers and punctuation are dropped. +func statements(sql string) [][]string { + var stmts [][]string + cur := make([]string, 0, 8) + i, n := 0, len(sql) + + for i < n { + c := sql[i] + switch { + case isSpace(c): + i++ + case c == '-' && i+1 < n && sql[i+1] == '-': + i = skipLineComment(sql, i+2) + case c == '/' && i+1 < n && sql[i+1] == '*': + i = skipBlockComment(sql, i+2) + case c == '\'': + i = skipQuoted(sql, i+1, '\'') + case c == '"': + i = skipQuoted(sql, i+1, '"') + case c == '$': + if end, ok := skipDollarQuote(sql, i); ok { + i = end + } else { + i = skipWord(sql, i+1) // parameter such as $1 + } + case c == ';': + stmts = append(stmts, cur) + cur = make([]string, 0, 8) + i++ + case isWordStart(c): + j := skipWord(sql, i+1) + cur = append(cur, strings.ToUpper(sql[i:j])) + i = j + default: + i++ + } + } + stmts = append(stmts, cur) + + return stmts +} + +func skipLineComment(s string, i int) int { + for i < len(s) && s[i] != '\n' { + i++ + } + + return i +} + +func skipBlockComment(s string, i int) int { + depth := 1 + for i < len(s) && depth > 0 { + switch { + case i+1 < len(s) && s[i] == '/' && s[i+1] == '*': + depth++ + i += 2 + case i+1 < len(s) && s[i] == '*' && s[i+1] == '/': + depth-- + i += 2 + default: + i++ + } + } + + return i +} + +// skipQuoted skips a string literal or quoted identifier that has already had +// its opening quote consumed, honoring the doubled-quote escape. +func skipQuoted(s string, i int, quote byte) int { + for i < len(s) { + if s[i] == quote { + if i+1 < len(s) && s[i+1] == quote { + i += 2 + + continue + } + + return i + 1 + } + i++ + } + + return i +} + +// skipDollarQuote skips a $tag$...$tag$ body. It reports false when the dollar +// does not open a valid tag (for example a $1 parameter), leaving it to the +// caller. +func skipDollarQuote(s string, i int) (int, bool) { + j := i + 1 + for j < len(s) && isTagPart(s[j]) { + j++ + } + if j >= len(s) || s[j] != '$' { + return 0, false + } + if j > i+1 && isDigit(s[i+1]) { + return 0, false // tags do not start with a digit; this is a parameter + } + + tag := s[i : j+1] + rest := s[j+1:] + if k := strings.Index(rest, tag); k >= 0 { + return j + 1 + k + len(tag), true + } + + return len(s), true // unterminated: consume the remainder +} + +func skipWord(s string, i int) int { + for i < len(s) && isWordPart(s[i]) { + i++ + } + + return i +} + +func isSpace(c byte) bool { + switch c { + case ' ', '\t', '\n', '\r', '\f', '\v': + return true + default: + return false + } +} + +func isWordStart(c byte) bool { + return c == '_' || (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || c >= 0x80 +} + +func isWordPart(c byte) bool { + return isWordStart(c) || isDigit(c) +} + +func isTagPart(c byte) bool { + return c == '_' || (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || isDigit(c) +} + +func isDigit(c byte) bool { + return c >= '0' && c <= '9' +} diff --git a/internal/sqlscan/sqlscan.go b/internal/sqlscan/sqlscan.go new file mode 100644 index 0000000..26995a6 --- /dev/null +++ b/internal/sqlscan/sqlscan.go @@ -0,0 +1,214 @@ +// Package sqlscan classifies SQL statements well enough to enforce a read-only +// policy. It is a lexical classifier, not a parser: it tokenizes far enough to +// skip strings, comments, dollar-quoted bodies, and quoted identifiers so that +// keywords inside them are never mistaken for commands, then classifies each +// top-level statement by its leading keyword. +package sqlscan + +import "slices" + +type Kind int + +const ( + Unknown Kind = iota + Empty + Select + Cursor // DECLARE/FETCH/MOVE/CLOSE, SHOW, and other read-only reads + Set // SET/RESET + Tx // BEGIN/COMMIT/ROLLBACK/SAVEPOINT + Explain + CopyTo + Insert + Update + Delete + Merge + Truncate + CopyFrom + Call + DDL +) + +func (k Kind) String() string { + switch k { + case Unknown: + return "unknown statement" + case Empty: + return "empty" + case Select: + return "SELECT" + case Cursor: + return "cursor operation" + case Set: + return "SET" + case Tx: + return "transaction control" + case Explain: + return "EXPLAIN" + case CopyTo: + return "COPY TO" + case Insert: + return "INSERT" + case Update: + return "UPDATE" + case Delete: + return "DELETE" + case Merge: + return "MERGE" + case Truncate: + return "TRUNCATE" + case CopyFrom: + return "COPY FROM" + case Call: + return "CALL" + case DDL: + return "DDL" + } + + return "unknown statement" +} + +// Mutating reports whether a statement of this kind can change data or schema. +// Unknown is treated as mutating so an unrecognized statement fails closed under +// a read-only policy. +func (k Kind) Mutating() bool { + switch k { + case Insert, Update, Delete, Merge, Truncate, CopyFrom, Call, DDL, Unknown: + return true + case Empty, Select, Cursor, Set, Tx, Explain, CopyTo: + return false + } + + return true +} + +// Classify returns the kind of each top-level statement in sql. Empty input, or +// input that is only comments and whitespace, yields a single Empty. +func Classify(sql string) []Kind { + kinds := make([]Kind, 0, 1) + for _, words := range statements(sql) { + kinds = append(kinds, classify(words)) + } + if len(kinds) == 0 { + return []Kind{Empty} + } + + return kinds +} + +// Mutating reports whether any statement in sql can change data or schema. +func Mutating(sql string) bool { + for _, k := range Classify(sql) { + if k.Mutating() { + return true + } + } + + return false +} + +func classify(words []string) Kind { + if len(words) == 0 { + return Empty + } + + switch words[0] { + case "SELECT", "VALUES", "TABLE": + return Select + case "SHOW", "DECLARE", "FETCH", "MOVE", "CLOSE", "LISTEN", "UNLISTEN", "DISCARD": + return Cursor + case "SET", "RESET": + return Set + case "BEGIN", "START", "COMMIT", "END", "ROLLBACK", "SAVEPOINT", "RELEASE", "ABORT": + return Tx + case "INSERT": + return Insert + case "UPDATE": + return Update + case "DELETE": + return Delete + case "MERGE": + return Merge + case "TRUNCATE": + return Truncate + case "CALL", "DO": + return Call + case "COPY": + return classifyCopy(words) + case "WITH": + return classifyWith(words) + case "EXPLAIN": + return classifyExplain(words) + case "CREATE", "ALTER", "DROP", "GRANT", "REVOKE", "COMMENT", + "REINDEX", "CLUSTER", "VACUUM", "ANALYZE", "REFRESH", "IMPORT", + "CHECKPOINT", "LOAD", "SECURITY", "LOCK": + return DDL + default: + return Unknown + } +} + +// classifyCopy distinguishes COPY ... FROM (a write) from COPY ... TO and the +// COPY (query) TO form (both reads). +func classifyCopy(words []string) Kind { + for _, w := range words[1:] { + if w == "SELECT" || w == "WITH" || w == "VALUES" { + return CopyTo // only the query form, which is always TO + } + } + for _, w := range words[1:] { + switch w { + case "FROM": + return CopyFrom + case "TO": + return CopyTo + } + } + + return CopyFrom // no clear direction: assume the writing form +} + +// classifyWith treats a WITH statement as mutating if any data-modifying command +// appears anywhere in it, since a CTE such as WITH d AS (DELETE ...) executes. +func classifyWith(words []string) Kind { + for _, w := range words[1:] { + switch w { + case "INSERT": + return Insert + case "UPDATE": + return Update + case "DELETE": + return Delete + case "MERGE": + return Merge + } + } + + return Select +} + +// classifyExplain returns EXPLAIN unless the plan is run with ANALYZE, which +// executes the underlying statement; then it classifies that statement. +func classifyExplain(words []string) Kind { + if !slices.Contains(words[1:], "ANALYZE") { + return Explain + } + + for _, w := range words[1:] { + switch w { + case "SELECT", "VALUES", "TABLE": + return Select + case "INSERT": + return Insert + case "UPDATE": + return Update + case "DELETE": + return Delete + case "MERGE": + return Merge + case "WITH": + return classifyWith(words) + } + } + + return Explain +} diff --git a/internal/sqlscan/sqlscan_test.go b/internal/sqlscan/sqlscan_test.go new file mode 100644 index 0000000..c2ba289 --- /dev/null +++ b/internal/sqlscan/sqlscan_test.go @@ -0,0 +1,111 @@ +package sqlscan_test + +import ( + "testing" + + "github.com/mickamy/rollcall/internal/sqlscan" +) + +func TestClassify(t *testing.T) { + t.Parallel() + + tests := map[string]struct { + sql string + want []sqlscan.Kind + }{ + "select": {"select 1", []sqlscan.Kind{sqlscan.Select}}, + "select with cte": {"with t as (select 1) select * from t", []sqlscan.Kind{sqlscan.Select}}, + "values": {"values (1),(2)", []sqlscan.Kind{sqlscan.Select}}, + "insert": {"insert into t values (1)", []sqlscan.Kind{sqlscan.Insert}}, + "update": {"update t set x = 1", []sqlscan.Kind{sqlscan.Update}}, + "delete": {"delete from t", []sqlscan.Kind{sqlscan.Delete}}, + "truncate": {"truncate t", []sqlscan.Kind{sqlscan.Truncate}}, + "create": {"create table t (id int)", []sqlscan.Kind{sqlscan.DDL}}, + "drop": {"drop table t", []sqlscan.Kind{sqlscan.DDL}}, + "set": {"set search_path to public", []sqlscan.Kind{sqlscan.Set}}, + "begin": {"begin", []sqlscan.Kind{sqlscan.Tx}}, + "show": {"show all", []sqlscan.Kind{sqlscan.Cursor}}, + "call": {"call do_work()", []sqlscan.Kind{sqlscan.Call}}, + "empty": {"", []sqlscan.Kind{sqlscan.Empty}}, + "comment only": {"-- just a comment\n", []sqlscan.Kind{sqlscan.Empty}}, + "leading comment": {"/* c */ select 1", []sqlscan.Kind{sqlscan.Select}}, + "leading paren select": {"(select 1) union (select 2)", []sqlscan.Kind{sqlscan.Select}}, + "data modifying cte": {"with d as (delete from t returning *) select * from d", []sqlscan.Kind{sqlscan.Delete}}, + "insert cte": { + "with i as (insert into t values (1) returning id) select * from i", + []sqlscan.Kind{sqlscan.Insert}, + }, + "explain select": {"explain select 1", []sqlscan.Kind{sqlscan.Explain}}, + "explain update": {"explain update t set x = 1", []sqlscan.Kind{sqlscan.Explain}}, + "explain analyze update": {"explain analyze update t set x = 1", []sqlscan.Kind{sqlscan.Update}}, + "explain analyze select": {"explain (analyze, buffers) select 1", []sqlscan.Kind{sqlscan.Select}}, + "copy from stdin": {"copy t from stdin", []sqlscan.Kind{sqlscan.CopyFrom}}, + "copy to stdout": {"copy t to stdout", []sqlscan.Kind{sqlscan.CopyTo}}, + "copy query to": {"copy (select * from t) to stdout", []sqlscan.Kind{sqlscan.CopyTo}}, + "multi statement": {"select 1; update t set x = 1", []sqlscan.Kind{sqlscan.Select, sqlscan.Update}}, + "trailing semicolon": {"select 1;", []sqlscan.Kind{sqlscan.Select, sqlscan.Empty}}, + "unknown": {"frobnicate t", []sqlscan.Kind{sqlscan.Unknown}}, + } + + for name, tt := range tests { + t.Run(name, func(t *testing.T) { + t.Parallel() + + got := sqlscan.Classify(tt.sql) + if len(got) != len(tt.want) { + t.Fatalf("Classify(%q) = %v, want %v", tt.sql, got, tt.want) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("Classify(%q)[%d] = %v, want %v", tt.sql, i, got[i], tt.want[i]) + } + } + }) + } +} + +func TestMutatingIgnoresKeywordsInsideLiteralsAndIdentifiers(t *testing.T) { + t.Parallel() + + readOnly := []string{ + `select * from orders where status = 'delete pending'`, + `select "update" from t`, + `select 'insert' || 'update' as note`, + `select $$ delete from t $$ as body`, + `select $tag$ update t $tag$`, + `select col from t -- update t set x = 1`, + `select col /* insert into t */ from t`, + `select delete_flag, update_ts from t`, + `select current_user`, + } + for _, sql := range readOnly { + if sqlscan.Mutating(sql) { + t.Errorf("Mutating(%q) = true, want false", sql) + } + } + + mutating := []string{ + `update t set x = 1`, + `insert into t values (1)`, + `with d as (delete from t returning *) select * from d`, + `explain analyze delete from t`, + `copy t from stdin`, + `select 1; drop table t`, + `truncate t`, + } + for _, sql := range mutating { + if !sqlscan.Mutating(sql) { + t.Errorf("Mutating(%q) = false, want true", sql) + } + } +} + +func TestClassifyDollarQuotedFunctionBody(t *testing.T) { + t.Parallel() + + sql := `do $$ begin delete from t; end $$` + got := sqlscan.Classify(sql) + if len(got) != 1 || got[0] != sqlscan.Call { + t.Errorf("Classify(%q) = %v, want [CALL] (the DELETE is inside the body)", sql, got) + } +} From ae0f31288714b45526a16c9efe002689f68f995e Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 26 Aug 2026 22:31:08 +0900 Subject: [PATCH 2/5] feat: resolve a per-role policy guard from a YAML file --- go.mod | 2 + go.sum | 4 ++ internal/policy/export_test.go | 3 + internal/policy/policy.go | 125 +++++++++++++++++++++++++++++++++ internal/policy/policy_test.go | 114 ++++++++++++++++++++++++++++++ internal/wire/wire.go | 17 +++++ 6 files changed, 265 insertions(+) create mode 100644 go.sum create mode 100644 internal/policy/export_test.go create mode 100644 internal/policy/policy.go create mode 100644 internal/policy/policy_test.go diff --git a/go.mod b/go.mod index 77d6ddc..6044575 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,5 @@ module github.com/mickamy/rollcall go 1.27.0 + +require gopkg.in/yaml.v3 v3.0.1 diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..a62c313 --- /dev/null +++ b/go.sum @@ -0,0 +1,4 @@ +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/policy/export_test.go b/internal/policy/export_test.go new file mode 100644 index 0000000..5f2ab19 --- /dev/null +++ b/internal/policy/export_test.go @@ -0,0 +1,3 @@ +package policy + +var Parse = parse diff --git a/internal/policy/policy.go b/internal/policy/policy.go new file mode 100644 index 0000000..34cb081 --- /dev/null +++ b/internal/policy/policy.go @@ -0,0 +1,125 @@ +// Package policy resolves a per-session guard from a YAML file: it maps a +// connection's database role to an agent, a purpose, and the statements that +// role may run. +package policy + +import ( + "bytes" + "fmt" + "os" + + "gopkg.in/yaml.v3" + + "github.com/mickamy/rollcall/internal/sqlscan" + "github.com/mickamy/rollcall/internal/wire" +) + +// Policy decides what each database role may do. The zero value denies nothing; +// load a file to enforce rules. +type Policy struct { + // FailClosed denies connections whose role is not listed. The default, + // fail-open, lets unlisted roles through unchanged. + FailClosed bool + Roles map[string]Role +} + +// Role is the set of rules bound to one database role. +type Role struct { + Agent string + Purpose string + ReadOnly bool +} + +var _ wire.Guard = (*Policy)(nil) + +type file struct { + Fail string `yaml:"fail"` + Roles map[string]role `yaml:"roles"` +} + +type role struct { + Agent string `yaml:"agent"` + Purpose string `yaml:"purpose"` + ReadOnly bool `yaml:"read_only"` +} + +// Load reads and validates a policy file. +func Load(path string) (Policy, error) { + data, err := os.ReadFile(path) + if err != nil { + return Policy{}, fmt.Errorf("read policy: %w", err) + } + + return parse(data) +} + +func parse(data []byte) (Policy, error) { + var f file + dec := yaml.NewDecoder(bytes.NewReader(data)) + dec.KnownFields(true) + if err := dec.Decode(&f); err != nil { + return Policy{}, fmt.Errorf("parse policy: %w", err) + } + + failClosed, err := parseFail(f.Fail) + if err != nil { + return Policy{}, err + } + + p := Policy{FailClosed: failClosed, Roles: make(map[string]Role, len(f.Roles))} + for name, r := range f.Roles { + p.Roles[name] = Role(r) + } + + return p, nil +} + +func parseFail(s string) (bool, error) { + switch s { + case "", "open": + return false, nil + case "closed": + return true, nil + default: + return false, fmt.Errorf("policy: fail must be \"open\" or \"closed\", got %q", s) + } +} + +// Resolve returns the handler for a session, binding the rules of the role that +// matches the connection's database user. +func (p Policy) Resolve(startup wire.Startup) wire.Handler { + role, ok := p.Roles[startup.User] + if !ok { + if p.FailClosed { + return wire.HandlerFunc(func(wire.Statement) wire.Verdict { + return wire.Verdict{ + Deny: true, + Message: fmt.Sprintf("no policy for role %q", startup.User), + Hint: "add the role to the policy, or connect as a configured role", + } + }) + } + + return wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }) + } + + return wire.HandlerFunc(role.evaluate) +} + +func (r Role) evaluate(stmt wire.Statement) wire.Verdict { + if !r.ReadOnly { + return wire.Verdict{} + } + + for _, kind := range sqlscan.Classify(stmt.SQL) { + if kind.Mutating() { + return wire.Verdict{ + Deny: true, + Message: fmt.Sprintf("%s is not allowed: this connection is read-only", kind), + Hint: "run the statement as a role that permits writes, or request approval", + } + } + } + + return wire.Verdict{} +} diff --git a/internal/policy/policy_test.go b/internal/policy/policy_test.go new file mode 100644 index 0000000..6f28d17 --- /dev/null +++ b/internal/policy/policy_test.go @@ -0,0 +1,114 @@ +package policy_test + +import ( + "strings" + "testing" + + "github.com/mickamy/rollcall/internal/policy" + "github.com/mickamy/rollcall/internal/wire" +) + +func TestParseReadsRolesAndFailMode(t *testing.T) { + t.Parallel() + + p := load(t, ` +fail: closed +roles: + agent_ops: + agent: claude-ops + purpose: incident-investigation + read_only: true + app_rw: + agent: web +`) + + if !p.FailClosed { + t.Error("FailClosed: got false, want true") + } + if got := p.Roles["agent_ops"]; got.Agent != "claude-ops" || !got.ReadOnly { + t.Errorf("agent_ops: got %+v", got) + } + if got := p.Roles["app_rw"]; got.ReadOnly { + t.Errorf("app_rw: got read-only, want writable") + } +} + +func TestParseRejectsUnknownFieldsAndFailValues(t *testing.T) { + t.Parallel() + + if _, err := policy.Parse([]byte("fail: sideways")); err == nil { + t.Error("bad fail value: got nil error") + } + if _, err := policy.Parse([]byte("roles:\n r:\n reed_only: true")); err == nil { + t.Error("unknown field: got nil error") + } +} + +func TestResolveReadOnlyDeniesWrites(t *testing.T) { + t.Parallel() + + p := load(t, ` +roles: + agent_ops: + read_only: true +`) + h := p.Resolve(wire.Startup{User: "agent_ops"}) + + if v := h.Statement(wire.Statement{SQL: "select * from t"}); v.Deny { + t.Errorf("select: got denied (%q), want allowed", v.Message) + } + + v := h.Statement(wire.Statement{SQL: "update t set x = 1"}) + if !v.Deny { + t.Fatal("update: got allowed, want denied") + } + if !strings.Contains(v.Message, "UPDATE") || !strings.Contains(v.Message, "read-only") { + t.Errorf("update denial message: got %q", v.Message) + } +} + +func TestResolveWritableRoleAllowsEverything(t *testing.T) { + t.Parallel() + + p := load(t, "roles:\n app_rw:\n agent: web\n") + h := p.Resolve(wire.Startup{User: "app_rw"}) + + if v := h.Statement(wire.Statement{SQL: "delete from t"}); v.Deny { + t.Errorf("delete for a writable role: got denied (%q), want allowed", v.Message) + } +} + +func TestResolveUnknownRoleHonorsFailMode(t *testing.T) { + t.Parallel() + + open := load(t, "roles: {}\n") + if v := open.Resolve(wire.Startup{User: "ghost"}).Statement(wire.Statement{SQL: "delete from t"}); v.Deny { + t.Errorf("fail-open unknown role: got denied (%q), want allowed", v.Message) + } + + closed := load(t, "fail: closed\nroles: {}\n") + v := closed.Resolve(wire.Startup{User: "ghost"}).Statement(wire.Statement{SQL: "select 1"}) + if !v.Deny || !strings.Contains(v.Message, "ghost") { + t.Errorf("fail-closed unknown role: got %+v, want a denial naming the role", v) + } +} + +func TestZeroPolicyAllowsEverything(t *testing.T) { + t.Parallel() + + var p policy.Policy + if v := p.Resolve(wire.Startup{User: "anyone"}).Statement(wire.Statement{SQL: "drop table t"}); v.Deny { + t.Errorf("zero policy: got denied (%q), want allowed", v.Message) + } +} + +func load(t *testing.T, src string) policy.Policy { + t.Helper() + + p, err := policy.Parse([]byte(src)) + if err != nil { + t.Fatalf("Parse: %v", err) + } + + return p +} diff --git a/internal/wire/wire.go b/internal/wire/wire.go index ae4e301..7885c87 100644 --- a/internal/wire/wire.go +++ b/internal/wire/wire.go @@ -45,6 +45,23 @@ func (f HandlerFunc) Statement(stmt Statement) Verdict { return f(stmt) } +// Guard resolves the Handler for a session from the client's identity. Resolve +// runs once, after the handshake. +type Guard interface { + Resolve(startup Startup) Handler +} + +type GuardFunc func(startup Startup) Handler + +func (f GuardFunc) Resolve(startup Startup) Handler { + return f(startup) +} + +// AllowAll is a Guard whose sessions permit every statement. +var AllowAll Guard = GuardFunc(func(Startup) Handler { + return HandlerFunc(func(Statement) Verdict { return Verdict{} }) +}) + type Dialect interface { NewSession(client, upstream net.Conn) Session } From 012e73dd224192a1d048a5e8d8d0d7ebdd99de0e Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Wed, 26 Aug 2026 22:31:08 +0900 Subject: [PATCH 3/5] feat: enforce a policy per session in the proxy command --- README.md | 18 +++++++++++++++++- internal/cli/cli_test.go | 5 +++++ internal/cli/proxy.go | 16 +++++++++++++++- internal/proxy/proxy.go | 15 +++++++++------ 4 files changed, 46 insertions(+), 8 deletions(-) diff --git a/README.md b/README.md index 1646476..713d731 100644 --- a/README.md +++ b/README.md @@ -21,7 +21,7 @@ Prebuilt binaries are on the [releases page](https://github.com/mickamy/rollcall ## Status -Early development. `rollcall proxy` speaks the PostgreSQL protocol: it relays authentication untouched, sees every statement on both the simple and the extended query protocol, and can refuse one before it reaches the server. The policy that decides what to refuse and the access ledger are being built on top of it; today everything is allowed. +Early development. `rollcall proxy` speaks the PostgreSQL protocol: it relays authentication untouched, sees every statement on both the simple and the extended query protocol, and refuses one before it reaches the server. A policy file maps each database role to an agent and can make it read-only; without `-policy`, every statement is allowed. The access ledger is next. The proxy speaks plaintext on both sides and answers `SSLRequest` with `N`, so `sslmode=prefer` clients fall back to plaintext. Keep the listener on loopback or a pod-local network until TLS lands. @@ -34,6 +34,22 @@ rollcall proxy -upstream 127.0.0.1:5432 # listens on 127.0.0.1:6432 PGHOST=127.0.0.1 PGPORT=6432 psql -U agent_claude_ops prod ``` +Enforce a read-only role with a policy file: + +```yaml +# policy.yaml +fail: closed +roles: + agent_ops: + agent: claude-ops + purpose: incident-investigation + read_only: true +``` + +```sh +rollcall proxy -upstream 127.0.0.1:5432 -policy policy.yaml +``` + ```sh rollcall proxy -h rollcall help diff --git a/internal/cli/cli_test.go b/internal/cli/cli_test.go index faee8a6..1eebbab 100644 --- a/internal/cli/cli_test.go +++ b/internal/cli/cli_test.go @@ -85,6 +85,11 @@ func TestRun(t *testing.T) { wantCode: exit.Error, wantErr: cli.Name + ": listen tcp", }, + "proxy with a missing policy file": { + args: []string{"proxy", "-upstream", "127.0.0.1:5432", "-policy", "/no/such/policy.yaml"}, + wantCode: exit.Error, + wantErr: cli.Name + ": read policy:", + }, } for name, tt := range tests { diff --git a/internal/cli/proxy.go b/internal/cli/proxy.go index afbc083..fcec168 100644 --- a/internal/cli/proxy.go +++ b/internal/cli/proxy.go @@ -11,19 +11,23 @@ import ( "github.com/mickamy/rollcall/internal/exit" "github.com/mickamy/rollcall/internal/pg" + "github.com/mickamy/rollcall/internal/policy" "github.com/mickamy/rollcall/internal/proxy" + "github.com/mickamy/rollcall/internal/wire" ) const ( defaultListen = "127.0.0.1:6432" listenUsage = "address to accept client connections on" upstreamUsage = "address of the upstream database" + policyUsage = "path to a policy file; without one every statement is allowed" ) func runProxy(ctx context.Context, args []string, std IO) int { fs := newFlagSet("proxy", std.Err, printProxyUsage) listen := fs.String("listen", defaultListen, listenUsage) upstream := fs.String("upstream", "", upstreamUsage) + policyPath := fs.String("policy", "", policyUsage) if err := fs.Parse(args); err != nil { if errors.Is(err, flag.ErrHelp) { return exit.OK @@ -44,6 +48,15 @@ func runProxy(ctx context.Context, args []string, std IO) int { return fail(std, err) } + guard := wire.AllowAll + if *policyPath != "" { + p, err := policy.Load(*policyPath) + if err != nil { + return fail(std, err) + } + guard = p + } + var lc net.ListenConfig ln, err := lc.Listen(ctx, "tcp", *listen) if err != nil { @@ -57,7 +70,7 @@ func runProxy(ctx context.Context, args []string, std IO) int { logger.Warn("listening outside loopback: clients and the upstream are served in plaintext", "addr", *listen) } - srv := proxy.Server{Upstream: *upstream, Dialect: pg.Dialect{}, Logger: logger} + srv := proxy.Server{Upstream: *upstream, Dialect: pg.Dialect{}, Guard: guard, Logger: logger} if err := srv.Serve(ctx, ln); err != nil { return fail(std, err) } @@ -104,4 +117,5 @@ func printProxyUsage(w io.Writer) { fmt.Fprint(w, "Flags:\n") fmt.Fprintf(w, " -upstream ADDR %s (required)\n", upstreamUsage) fmt.Fprintf(w, " -listen ADDR %s (default %q)\n", listenUsage, defaultListen) + fmt.Fprintf(w, " -policy PATH %s\n", policyUsage) } diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 8491db2..b4b32c5 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -25,17 +25,18 @@ const ( type Server struct { Upstream string Dialect wire.Dialect - // Handler decides each statement; nil allows everything. - Handler wire.Handler - Logger *slog.Logger + // Guard resolves the per-session handler from the client's identity; nil + // allows everything. + Guard wire.Guard + Logger *slog.Logger } func (s Server) Serve(ctx context.Context, ln net.Listener) error { if s.Dialect == nil { return errors.New("proxy: Dialect is required") } - if s.Handler == nil { - s.Handler = wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }) + if s.Guard == nil { + s.Guard = wire.AllowAll } if s.Logger == nil { s.Logger = slog.New(slog.DiscardHandler) @@ -115,10 +116,12 @@ func (s Server) handle(ctx context.Context, client net.Conn) { logger = logger.With("user", startup.User, "database", startup.Database, "application", startup.Application) logger.Info("session opened") + handler := s.Guard.Resolve(startup) + var wg sync.WaitGroup var toUpstream, toClient error wg.Go(func() { - toUpstream = sess.Frontend(s.Handler) + toUpstream = sess.Frontend(handler) closeWrite(upstream) }) wg.Go(func() { From c2161be707a6154b0f23c42cdb5a78defa722930 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 08:55:42 +0900 Subject: [PATCH 4/5] fix: harden the SQL classifier against comment, dollar-quote, COPY, and EXPLAIN bypasses --- internal/sqlscan/scan.go | 71 ++++++-- internal/sqlscan/sqlscan.go | 297 ++++++++++++++++++++++++------- internal/sqlscan/sqlscan_test.go | 115 ++++++++---- 3 files changed, 367 insertions(+), 116 deletions(-) diff --git a/internal/sqlscan/scan.go b/internal/sqlscan/scan.go index 5a571f0..d505ff8 100644 --- a/internal/sqlscan/scan.go +++ b/internal/sqlscan/scan.go @@ -2,13 +2,29 @@ package sqlscan import "strings" -// statements splits sql into top-level statements, each a slice of the -// uppercased unquoted words it contains. Strings, comments, dollar-quoted -// bodies, and quoted identifiers are skipped so their contents never look like -// keywords; numbers and punctuation are dropped. -func statements(sql string) [][]string { - var stmts [][]string - cur := make([]string, 0, 8) +// tokKind marks the few token shapes classification needs. Strings, comments, +// dollar-quoted bodies, numbers, and other punctuation are dropped entirely. +type tokKind uint8 + +const ( + kindWord tokKind = iota // an unquoted word, stored uppercased (a keyword candidate) + kindIdent // a quoted identifier, stored lowercased (never a keyword) + kindOpen // ( + kindClose // ) +) + +type token struct { + kind tokKind + text string +} + +// tokenize splits sql into top-level statements, each a slice of tokens. It +// skips string literals, line and block comments, dollar-quoted bodies, and the +// contents of quoted identifiers so that keywords hidden in them are never +// mistaken for commands. +func tokenize(sql string) [][]token { + var stmts [][]token + cur := make([]token, 0, 8) i, n := 0, len(sql) for i < n { @@ -21,22 +37,30 @@ func statements(sql string) [][]string { case c == '/' && i+1 < n && sql[i+1] == '*': i = skipBlockComment(sql, i+2) case c == '\'': - i = skipQuoted(sql, i+1, '\'') + i = skipString(sql, i+1) case c == '"': - i = skipQuoted(sql, i+1, '"') + j := skipString(sql, i+1) + cur = append(cur, token{kind: kindIdent, text: quotedText(sql, i, j)}) + i = j case c == '$': if end, ok := skipDollarQuote(sql, i); ok { i = end } else { i = skipWord(sql, i+1) // parameter such as $1 } + case c == '(': + cur = append(cur, token{kind: kindOpen}) + i++ + case c == ')': + cur = append(cur, token{kind: kindClose}) + i++ case c == ';': stmts = append(stmts, cur) - cur = make([]string, 0, 8) + cur = make([]token, 0, 8) i++ case isWordStart(c): j := skipWord(sql, i+1) - cur = append(cur, strings.ToUpper(sql[i:j])) + cur = append(cur, token{kind: kindWord, text: strings.ToUpper(sql[i:j])}) i = j default: i++ @@ -47,8 +71,16 @@ func statements(sql string) [][]string { return stmts } +// quotedText returns the lowercased content of a quoted identifier spanning +// s[start:end] (quotes included), so a quoted GUC name folds to compare. +func quotedText(s string, start, end int) string { + inner := s[start+1 : end-1] + + return strings.ToLower(strings.ReplaceAll(inner, `""`, `"`)) +} + func skipLineComment(s string, i int) int { - for i < len(s) && s[i] != '\n' { + for i < len(s) && s[i] != '\n' && s[i] != '\r' { i++ } @@ -73,9 +105,11 @@ func skipBlockComment(s string, i int) int { return i } -// skipQuoted skips a string literal or quoted identifier that has already had -// its opening quote consumed, honoring the doubled-quote escape. -func skipQuoted(s string, i int, quote byte) int { +// skipString skips a string literal or quoted identifier whose opening quote at +// i-1 has been consumed, honoring the doubled-quote escape, and returns the +// index just past the closing quote. +func skipString(s string, i int) int { + quote := s[i-1] for i < len(s) { if s[i] == quote { if i+1 < len(s) && s[i+1] == quote { @@ -108,8 +142,7 @@ func skipDollarQuote(s string, i int) (int, bool) { } tag := s[i : j+1] - rest := s[j+1:] - if k := strings.Index(rest, tag); k >= 0 { + if k := strings.Index(s[j+1:], tag); k >= 0 { return j + 1 + k + len(tag), true } @@ -137,8 +170,10 @@ func isWordStart(c byte) bool { return c == '_' || (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || c >= 0x80 } +// isWordPart includes '$' so that identifiers such as x$a$ read as one word and +// their inner '$' is not taken to start a dollar-quoted string. func isWordPart(c byte) bool { - return isWordStart(c) || isDigit(c) + return isWordStart(c) || isDigit(c) || c == '$' } func isTagPart(c byte) bool { diff --git a/internal/sqlscan/sqlscan.go b/internal/sqlscan/sqlscan.go index 26995a6..130a066 100644 --- a/internal/sqlscan/sqlscan.go +++ b/internal/sqlscan/sqlscan.go @@ -1,12 +1,14 @@ -// Package sqlscan classifies SQL statements well enough to enforce a read-only +// Package sqlscan classifies SQL statements well enough to support a read-only // policy. It is a lexical classifier, not a parser: it tokenizes far enough to // skip strings, comments, dollar-quoted bodies, and quoted identifiers so that // keywords inside them are never mistaken for commands, then classifies each // top-level statement by its leading keyword. +// +// It cannot see writes performed through functions (for example nextval or a +// volatile user function), so it is one layer of a read-only defense whose +// guarantee comes from the server-side read-only transaction mode. package sqlscan -import "slices" - type Kind int const ( @@ -17,13 +19,15 @@ const ( Set // SET/RESET Tx // BEGIN/COMMIT/ROLLBACK/SAVEPOINT Explain - CopyTo + CopyTo // COPY ... TO STDOUT Insert Update Delete Merge Truncate + SelectInto CopyFrom + CopyExternal // COPY ... TO a file or program Call DDL ) @@ -33,7 +37,7 @@ func (k Kind) String() string { case Unknown: return "unknown statement" case Empty: - return "empty" + return "empty statement" case Select: return "SELECT" case Cursor: @@ -45,7 +49,7 @@ func (k Kind) String() string { case Explain: return "EXPLAIN" case CopyTo: - return "COPY TO" + return "COPY TO STDOUT" case Insert: return "INSERT" case Update: @@ -56,8 +60,12 @@ func (k Kind) String() string { return "MERGE" case Truncate: return "TRUNCATE" + case SelectInto: + return "SELECT INTO" case CopyFrom: return "COPY FROM" + case CopyExternal: + return "COPY TO a file or program" case Call: return "CALL" case DDL: @@ -68,11 +76,10 @@ func (k Kind) String() string { } // Mutating reports whether a statement of this kind can change data or schema. -// Unknown is treated as mutating so an unrecognized statement fails closed under -// a read-only policy. +// Unknown is treated as mutating so an unrecognized statement fails closed. func (k Kind) Mutating() bool { switch k { - case Insert, Update, Delete, Merge, Truncate, CopyFrom, Call, DDL, Unknown: + case Insert, Update, Delete, Merge, Truncate, SelectInto, CopyFrom, CopyExternal, Call, DDL, Unknown: return true case Empty, Select, Cursor, Set, Tx, Explain, CopyTo: return false @@ -81,15 +88,36 @@ func (k Kind) Mutating() bool { return true } -// Classify returns the kind of each top-level statement in sql. Empty input, or -// input that is only comments and whitespace, yields a single Empty. -func Classify(sql string) []Kind { - kinds := make([]Kind, 0, 1) - for _, words := range statements(sql) { - kinds = append(kinds, classify(words)) +// Finding describes one top-level statement. +type Finding struct { + Kind Kind + // DisablesReadOnly is true when the statement could return the connection to + // read-write, for example SET default_transaction_read_only = off or + // BEGIN ... READ WRITE. Such statements must be refused on a read-only + // connection or the server-side read-only mode can be turned off. + DisablesReadOnly bool +} + +// Scan classifies each top-level statement in sql. +func Scan(sql string) []Finding { + stmts := tokenize(sql) + out := make([]Finding, 0, len(stmts)) + for _, toks := range stmts { + out = append(out, Finding{Kind: classify(toks), DisablesReadOnly: disablesReadOnly(toks)}) } - if len(kinds) == 0 { - return []Kind{Empty} + if len(out) == 0 { + return []Finding{{Kind: Empty}} + } + + return out +} + +// Classify returns the kind of each top-level statement in sql. +func Classify(sql string) []Kind { + findings := Scan(sql) + kinds := make([]Kind, len(findings)) + for i, f := range findings { + kinds[i] = f.Kind } return kinds @@ -97,8 +125,8 @@ func Classify(sql string) []Kind { // Mutating reports whether any statement in sql can change data or schema. func Mutating(sql string) bool { - for _, k := range Classify(sql) { - if k.Mutating() { + for _, f := range Scan(sql) { + if f.Kind.Mutating() { return true } } @@ -106,15 +134,17 @@ func Mutating(sql string) bool { return false } -func classify(words []string) Kind { - if len(words) == 0 { +func classify(toks []token) Kind { + w, i := firstWord(toks) + if i < 0 { return Empty } - switch words[0] { + switch w { case "SELECT", "VALUES", "TABLE": - return Select - case "SHOW", "DECLARE", "FETCH", "MOVE", "CLOSE", "LISTEN", "UNLISTEN", "DISCARD": + return selectOrInto(toks) + case "SHOW", "DECLARE", "FETCH", "MOVE", "CLOSE", "LISTEN", "UNLISTEN", "DISCARD", + "PREPARE", "DEALLOCATE", "EXECUTE", "NOTIFY", "CHECKPOINT": return Cursor case "SET", "RESET": return Set @@ -133,34 +163,55 @@ func classify(words []string) Kind { case "CALL", "DO": return Call case "COPY": - return classifyCopy(words) + return classifyCopy(toks) case "WITH": - return classifyWith(words) + return classifyWith(toks) case "EXPLAIN": - return classifyExplain(words) + return classifyExplain(toks) case "CREATE", "ALTER", "DROP", "GRANT", "REVOKE", "COMMENT", - "REINDEX", "CLUSTER", "VACUUM", "ANALYZE", "REFRESH", "IMPORT", - "CHECKPOINT", "LOAD", "SECURITY", "LOCK": + "REINDEX", "CLUSTER", "VACUUM", "ANALYZE", "REFRESH", "IMPORT", "LOAD", "SECURITY", "LOCK": return DDL - default: - return Unknown } + + return Unknown } -// classifyCopy distinguishes COPY ... FROM (a write) from COPY ... TO and the -// COPY (query) TO form (both reads). -func classifyCopy(words []string) Kind { - for _, w := range words[1:] { - if w == "SELECT" || w == "WITH" || w == "VALUES" { - return CopyTo // only the query form, which is always TO - } +// selectOrInto classifies a SELECT, upgrading it to SelectInto when a top-level +// INTO creates a table. +func selectOrInto(toks []token) Kind { + if hasWordAtTopLevel(toks, "INTO") { + return SelectInto } - for _, w := range words[1:] { - switch w { - case "FROM": - return CopyFrom - case "TO": - return CopyTo + + return Select +} + +// classifyCopy finds the first FROM or TO at the top level. COPY ... FROM writes +// a table; COPY ... TO STDOUT reads; COPY ... TO a file or program writes on the +// server, so only TO STDOUT is treated as a read. +func classifyCopy(toks []token) Kind { + depth := 0 + for i := 1; i < len(toks); i++ { + switch toks[i].kind { + case kindOpen: + depth++ + case kindClose: + depth-- + case kindWord: + if depth != 0 { + continue + } + switch toks[i].text { + case "FROM": + return CopyFrom + case "TO": + if nextWord(toks, i) == "STDOUT" { + return CopyTo + } + + return CopyExternal + } + case kindIdent: } } @@ -168,10 +219,13 @@ func classifyCopy(words []string) Kind { } // classifyWith treats a WITH statement as mutating if any data-modifying command -// appears anywhere in it, since a CTE such as WITH d AS (DELETE ...) executes. -func classifyWith(words []string) Kind { - for _, w := range words[1:] { - switch w { +// appears in it, since a CTE such as WITH d AS (DELETE ...) executes. +func classifyWith(toks []token) Kind { + for _, t := range toks[1:] { + if t.kind != kindWord { + continue + } + switch t.text { case "INSERT": return Insert case "UPDATE": @@ -183,32 +237,141 @@ func classifyWith(words []string) Kind { } } - return Select + return selectOrInto(toks) } -// classifyExplain returns EXPLAIN unless the plan is run with ANALYZE, which +// classifyExplain returns EXPLAIN unless the plan runs with ANALYZE, which // executes the underlying statement; then it classifies that statement. -func classifyExplain(words []string) Kind { - if !slices.Contains(words[1:], "ANALYZE") { - return Explain +func classifyExplain(toks []token) Kind { + if !containsWord(toks, "ANALYZE") { + return Explain // without ANALYZE the underlying statement never runs } - for _, w := range words[1:] { - switch w { - case "SELECT", "VALUES", "TABLE": - return Select - case "INSERT": - return Insert - case "UPDATE": - return Update - case "DELETE": - return Delete - case "MERGE": - return Merge - case "WITH": - return classifyWith(words) + for i := 1; i < len(toks); i++ { + if toks[i].kind != kindWord { + continue + } + if isExplainOption(toks[i].text) { + continue } + + return classify(toks[i:]) } return Explain } + +func isExplainOption(w string) bool { + switch w { + case "ANALYZE", "VERBOSE", "COSTS", "SETTINGS", "GENERIC_PLAN", "BUFFERS", "WAL", + "TIMING", "SUMMARY", "MEMORY", "SERIALIZE", "FORMAT", + "ON", "OFF", "TRUE", "FALSE", "TEXT", "JSON", "XML", "YAML", "NONE": + return true + } + + return false +} + +// disablesReadOnly reports whether the statement could turn off read-only mode. +func disablesReadOnly(toks []token) bool { + w, _ := firstWord(toks) + switch w { + case "SET", "RESET", "BEGIN", "START", "SELECT", "WITH": + default: + return false + } + + for i, t := range toks { + // READ WRITE, as in SET TRANSACTION READ WRITE or BEGIN ... READ WRITE. + if t.kind == kindWord && t.text == "READ" && nextWord(toks, i) == "WRITE" { + return true + } + // Any reference to the read-only GUCs, quoted or not. + if folds(t, "default_transaction_read_only") || folds(t, "transaction_read_only") { + return true + } + // set_config('...transaction_read_only', ...) reached through SELECT. + if t.kind == kindWord && t.text == "SET_CONFIG" { + return true + } + // RESET ALL clears every session setting, including read-only mode. + if t.kind == kindWord && t.text == "RESET" && nextWord(toks, i) == "ALL" { + return true + } + } + + return false +} + +func folds(t token, name string) bool { + switch t.kind { + case kindWord: + return t.text == upper(name) + case kindIdent: + return t.text == name + case kindOpen, kindClose: + return false + } + + return false +} + +func firstWord(toks []token) (string, int) { + for i, t := range toks { + if t.kind == kindWord { + return t.text, i + } + } + + return "", -1 +} + +func nextWord(toks []token, i int) string { + for j := i + 1; j < len(toks); j++ { + if toks[j].kind == kindWord { + return toks[j].text + } + } + + return "" +} + +func containsWord(toks []token, w string) bool { + for _, t := range toks { + if t.kind == kindWord && t.text == w { + return true + } + } + + return false +} + +func hasWordAtTopLevel(toks []token, w string) bool { + depth := 0 + for _, t := range toks { + switch t.kind { + case kindOpen: + depth++ + case kindClose: + depth-- + case kindWord: + if depth == 0 && t.text == w { + return true + } + case kindIdent: + } + } + + return false +} + +func upper(s string) string { + b := []byte(s) + for i := range b { + if b[i] >= 'a' && b[i] <= 'z' { + b[i] -= 'a' - 'A' + } + } + + return string(b) +} diff --git a/internal/sqlscan/sqlscan_test.go b/internal/sqlscan/sqlscan_test.go index c2ba289..0d09461 100644 --- a/internal/sqlscan/sqlscan_test.go +++ b/internal/sqlscan/sqlscan_test.go @@ -13,35 +13,35 @@ func TestClassify(t *testing.T) { sql string want []sqlscan.Kind }{ - "select": {"select 1", []sqlscan.Kind{sqlscan.Select}}, - "select with cte": {"with t as (select 1) select * from t", []sqlscan.Kind{sqlscan.Select}}, - "values": {"values (1),(2)", []sqlscan.Kind{sqlscan.Select}}, - "insert": {"insert into t values (1)", []sqlscan.Kind{sqlscan.Insert}}, - "update": {"update t set x = 1", []sqlscan.Kind{sqlscan.Update}}, - "delete": {"delete from t", []sqlscan.Kind{sqlscan.Delete}}, - "truncate": {"truncate t", []sqlscan.Kind{sqlscan.Truncate}}, - "create": {"create table t (id int)", []sqlscan.Kind{sqlscan.DDL}}, - "drop": {"drop table t", []sqlscan.Kind{sqlscan.DDL}}, - "set": {"set search_path to public", []sqlscan.Kind{sqlscan.Set}}, - "begin": {"begin", []sqlscan.Kind{sqlscan.Tx}}, - "show": {"show all", []sqlscan.Kind{sqlscan.Cursor}}, - "call": {"call do_work()", []sqlscan.Kind{sqlscan.Call}}, - "empty": {"", []sqlscan.Kind{sqlscan.Empty}}, - "comment only": {"-- just a comment\n", []sqlscan.Kind{sqlscan.Empty}}, - "leading comment": {"/* c */ select 1", []sqlscan.Kind{sqlscan.Select}}, - "leading paren select": {"(select 1) union (select 2)", []sqlscan.Kind{sqlscan.Select}}, - "data modifying cte": {"with d as (delete from t returning *) select * from d", []sqlscan.Kind{sqlscan.Delete}}, - "insert cte": { - "with i as (insert into t values (1) returning id) select * from i", - []sqlscan.Kind{sqlscan.Insert}, - }, + "select": {"select 1", []sqlscan.Kind{sqlscan.Select}}, + "select cte": {"with t as (select 1) select * from t", []sqlscan.Kind{sqlscan.Select}}, + "values": {"values (1),(2)", []sqlscan.Kind{sqlscan.Select}}, + "insert": {"insert into t values (1)", []sqlscan.Kind{sqlscan.Insert}}, + "update": {"update t set x = 1", []sqlscan.Kind{sqlscan.Update}}, + "delete": {"delete from t", []sqlscan.Kind{sqlscan.Delete}}, + "truncate": {"truncate t", []sqlscan.Kind{sqlscan.Truncate}}, + "create": {"create table t (id int)", []sqlscan.Kind{sqlscan.DDL}}, + "set": {"set search_path to public", []sqlscan.Kind{sqlscan.Set}}, + "begin": {"begin", []sqlscan.Kind{sqlscan.Tx}}, + "show": {"show all", []sqlscan.Kind{sqlscan.Cursor}}, + "call": {"call do_work()", []sqlscan.Kind{sqlscan.Call}}, + "empty": {"", []sqlscan.Kind{sqlscan.Empty}}, + "comment only": {"-- just a comment\n", []sqlscan.Kind{sqlscan.Empty}}, + "leading comment": {"/* c */ select 1", []sqlscan.Kind{sqlscan.Select}}, + "leading paren select": {"(select 1) union (select 2)", []sqlscan.Kind{sqlscan.Select}}, + "data modifying cte": {"with d as (delete from t returning *) select * from d", []sqlscan.Kind{sqlscan.Delete}}, "explain select": {"explain select 1", []sqlscan.Kind{sqlscan.Explain}}, "explain update": {"explain update t set x = 1", []sqlscan.Kind{sqlscan.Explain}}, "explain analyze update": {"explain analyze update t set x = 1", []sqlscan.Kind{sqlscan.Update}}, "explain analyze select": {"explain (analyze, buffers) select 1", []sqlscan.Kind{sqlscan.Select}}, + "explain analyze ctas": {"explain analyze create table t as select 1", []sqlscan.Kind{sqlscan.DDL}}, + "select into": {"select id into newt from t", []sqlscan.Kind{sqlscan.SelectInto}}, "copy from stdin": {"copy t from stdin", []sqlscan.Kind{sqlscan.CopyFrom}}, + "copy from with": {"copy t from stdin with (format csv)", []sqlscan.Kind{sqlscan.CopyFrom}}, "copy to stdout": {"copy t to stdout", []sqlscan.Kind{sqlscan.CopyTo}}, - "copy query to": {"copy (select * from t) to stdout", []sqlscan.Kind{sqlscan.CopyTo}}, + "copy to file": {"copy t to '/tmp/x.csv'", []sqlscan.Kind{sqlscan.CopyExternal}}, + "copy to program": {"copy t to program 'cat'", []sqlscan.Kind{sqlscan.CopyExternal}}, + "copy query to stdout": {"copy (select * from t) to stdout", []sqlscan.Kind{sqlscan.CopyTo}}, "multi statement": {"select 1; update t set x = 1", []sqlscan.Kind{sqlscan.Select, sqlscan.Update}}, "trailing semicolon": {"select 1;", []sqlscan.Kind{sqlscan.Select, sqlscan.Empty}}, "unknown": {"frobnicate t", []sqlscan.Kind{sqlscan.Unknown}}, @@ -76,7 +76,7 @@ func TestMutatingIgnoresKeywordsInsideLiteralsAndIdentifiers(t *testing.T) { `select col from t -- update t set x = 1`, `select col /* insert into t */ from t`, `select delete_flag, update_ts from t`, - `select current_user`, + `select 1 as x$a$`, } for _, sql := range readOnly { if sqlscan.Mutating(sql) { @@ -86,12 +86,12 @@ func TestMutatingIgnoresKeywordsInsideLiteralsAndIdentifiers(t *testing.T) { mutating := []string{ `update t set x = 1`, - `insert into t values (1)`, `with d as (delete from t returning *) select * from d`, `explain analyze delete from t`, `copy t from stdin`, + `copy t to program 'cat'`, + `select id into newt from t`, `select 1; drop table t`, - `truncate t`, } for _, sql := range mutating { if !sqlscan.Mutating(sql) { @@ -100,12 +100,65 @@ func TestMutatingIgnoresKeywordsInsideLiteralsAndIdentifiers(t *testing.T) { } } -func TestClassifyDollarQuotedFunctionBody(t *testing.T) { +// These are the injection tricks a client could use to hide a write from a +// lexical scanner; each must be seen for what it really is. +func TestClassifyResistsHidingStatements(t *testing.T) { t.Parallel() - sql := `do $$ begin delete from t; end $$` - got := sqlscan.Classify(sql) - if len(got) != 1 || got[0] != sqlscan.Call { - t.Errorf("Classify(%q) = %v, want [CALL] (the DELETE is inside the body)", sql, got) + // A carriage-return ends a line comment in PostgreSQL, so the INSERT is real. + if !sqlscan.Mutating("select 1 --c\rinsert into t values (1)") { + t.Error("CR line-comment: hidden INSERT not detected") } + // x$a$ is one identifier, not the start of a dollar-quoted body, so the + // INSERT between the fake tags is a real statement. + sql := "select 1 as x$a$; insert into t values (1); select $a$ $a$" + if !sqlscan.Mutating(sql) { + t.Errorf("dollar-in-identifier: hidden INSERT not detected in %q", sql) + } +} + +func TestDisablesReadOnly(t *testing.T) { + t.Parallel() + + disabling := []string{ + `set default_transaction_read_only = off`, + `SET default_transaction_read_only TO false`, + `set "default_transaction_read_only" = off`, + `set transaction read write`, + `begin read write`, + `start transaction read write`, + `set session characteristics as transaction read write`, + `reset default_transaction_read_only`, + `reset all`, + `select set_config('default_transaction_read_only', 'off', false)`, + } + for _, sql := range disabling { + if !disables(sql) { + t.Errorf("DisablesReadOnly(%q) = false, want true", sql) + } + } + + safe := []string{ + `set search_path to public`, + `set time zone 'UTC'`, + `begin`, + `begin transaction isolation level serializable`, + `select 1`, + `set role readonly`, + } + for _, sql := range safe { + if disables(sql) { + t.Errorf("DisablesReadOnly(%q) = true, want false", sql) + } + } +} + +func disables(sql string) bool { + for _, f := range sqlscan.Scan(sql) { + if f.DisablesReadOnly { + return true + } + } + + return false } From b5e41ea57a600b9d1ff36710d1a1d096785480a9 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 08:56:01 +0900 Subject: [PATCH 5/5] feat: enforce read-only with a server-side transaction mode and blocked escape hatches --- README.md | 2 + internal/pg/pg.go | 29 +++++++++++++++ internal/policy/policy.go | 68 +++++++++++++++++++++++----------- internal/policy/policy_test.go | 65 ++++++++++++++++++++++---------- internal/proxy/proxy.go | 15 ++++++-- internal/proxy/proxy_test.go | 4 ++ internal/wire/wire.go | 25 +++++++++---- 7 files changed, 157 insertions(+), 51 deletions(-) diff --git a/README.md b/README.md index 713d731..fabea81 100644 --- a/README.md +++ b/README.md @@ -23,6 +23,8 @@ Prebuilt binaries are on the [releases page](https://github.com/mickamy/rollcall Early development. `rollcall proxy` speaks the PostgreSQL protocol: it relays authentication untouched, sees every statement on both the simple and the extended query protocol, and refuses one before it reaches the server. A policy file maps each database role to an agent and can make it read-only; without `-policy`, every statement is allowed. The access ledger is next. +A read-only role is enforced in depth: the session is set read-only on the server (so writes through functions such as `nextval` are refused too), the proxy blocks attempts to turn that off, and obvious writes are refused early with a clear message. A lexical proxy cannot fully sandbox a role that already holds write privileges; for the strongest guarantee, also grant that database role only `SELECT`. + The proxy speaks plaintext on both sides and answers `SSLRequest` with `N`, so `sslmode=prefer` clients fall back to plaintext. Keep the listener on loopback or a pod-local network until TLS lands. ## Usage diff --git a/internal/pg/pg.go b/internal/pg/pg.go index f1a7bee..465dfda 100644 --- a/internal/pg/pg.go +++ b/internal/pg/pg.go @@ -137,6 +137,35 @@ func (s *session) Handshake() (wire.Startup, error) { } } +// Prime runs one statement on the upstream and consumes its response, before +// Frontend and Backend start, so it needs no locking. It fails if the upstream +// reports an error, so a failed read-only setup tears the session down. +func (s *session) Prime(sql string) error { + if err := writeMessage(s.uw, typeQuery, append([]byte(sql), 0)); err != nil { + return err + } + if err := flush(s.uw); err != nil { + return err + } + + for { + typ, body, err := readMessage(s.ur, maxAuthMessage) + if err != nil { + return fmt.Errorf("read upstream during prime: %w", err) + } + switch typ { + case typeErrorResponse: + return fmt.Errorf("prime %q: %s", sql, errorMessage(body)) + case typeReadyForQuery: + if len(body) == 1 { + s.tx = body[0] + } + + return nil + } + } +} + // Frontend relays client messages, consulting h for every statement. Extended // messages are held until their Sync so a batch can be accepted or rejected as // a unit. Output is flushed whenever the client has nothing more buffered. diff --git a/internal/policy/policy.go b/internal/policy/policy.go index 34cb081..d8da76b 100644 --- a/internal/policy/policy.go +++ b/internal/policy/policy.go @@ -14,6 +14,11 @@ import ( "github.com/mickamy/rollcall/internal/wire" ) +// readOnlyPrime is run on the upstream for a read-only role. It makes the server +// refuse every write, including writes performed through functions such as +// nextval, which the statement classifier cannot see. +const readOnlyPrime = "SET default_transaction_read_only = on" + // Policy decides what each database role may do. The zero value denies nothing; // load a file to enforce rules. type Policy struct { @@ -85,41 +90,62 @@ func parseFail(s string) (bool, error) { } } -// Resolve returns the handler for a session, binding the rules of the role that -// matches the connection's database user. -func (p Policy) Resolve(startup wire.Startup) wire.Handler { +// Resolve returns the enforcement for a session, binding the rules of the role +// that matches the connection's database user. +func (p Policy) Resolve(startup wire.Startup) wire.Enforcement { role, ok := p.Roles[startup.User] if !ok { - if p.FailClosed { - return wire.HandlerFunc(func(wire.Statement) wire.Verdict { - return wire.Verdict{ - Deny: true, - Message: fmt.Sprintf("no policy for role %q", startup.User), - Hint: "add the role to the policy, or connect as a configured role", - } - }) - } + return p.unlisted(startup.User) + } - return wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }) + if !role.ReadOnly { + return wire.Enforcement{Handler: allow()} } - return wire.HandlerFunc(role.evaluate) + return wire.Enforcement{ + Prime: []string{readOnlyPrime}, + Handler: wire.HandlerFunc(readOnly), + } } -func (r Role) evaluate(stmt wire.Statement) wire.Verdict { - if !r.ReadOnly { - return wire.Verdict{} +func (p Policy) unlisted(user string) wire.Enforcement { + if !p.FailClosed { + return wire.Enforcement{Handler: allow()} } - for _, kind := range sqlscan.Classify(stmt.SQL) { - if kind.Mutating() { + return wire.Enforcement{Handler: wire.HandlerFunc(func(wire.Statement) wire.Verdict { + return wire.Verdict{ + Deny: true, + Message: fmt.Sprintf("no policy for role %q", user), + Hint: "add the role to the policy, or connect as a configured role", + } + })} +} + +// readOnly denies statements that write or that could turn read-only mode off. +// The server-side read-only transaction is the real guarantee; this gives a +// clear, early refusal and stops the client from disabling it. +func readOnly(stmt wire.Statement) wire.Verdict { + for _, f := range sqlscan.Scan(stmt.SQL) { + if f.DisablesReadOnly { return wire.Verdict{ Deny: true, - Message: fmt.Sprintf("%s is not allowed: this connection is read-only", kind), - Hint: "run the statement as a role that permits writes, or request approval", + Message: "changing this connection to read-write is not allowed", + Hint: "this connection is read-only; use a role that permits writes", + } + } + if f.Kind.Mutating() { + return wire.Verdict{ + Deny: true, + Message: fmt.Sprintf("%s is not allowed: this connection is read-only", f.Kind), + Hint: "use a role that permits writes, or request approval", } } } return wire.Verdict{} } + +func allow() wire.Handler { + return wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }) +} diff --git a/internal/policy/policy_test.go b/internal/policy/policy_test.go index 6f28d17..a99c1d8 100644 --- a/internal/policy/policy_test.go +++ b/internal/policy/policy_test.go @@ -1,6 +1,7 @@ package policy_test import ( + "slices" "strings" "testing" @@ -44,50 +45,70 @@ func TestParseRejectsUnknownFieldsAndFailValues(t *testing.T) { } } -func TestResolveReadOnlyDeniesWrites(t *testing.T) { +func TestReadOnlyRolePrimesAndDeniesWrites(t *testing.T) { t.Parallel() - p := load(t, ` -roles: - agent_ops: - read_only: true -`) - h := p.Resolve(wire.Startup{User: "agent_ops"}) + p := load(t, "roles:\n agent_ops:\n read_only: true\n") + enf := p.Resolve(wire.Startup{User: "agent_ops"}) + + if !slices.Contains(enf.Prime, "SET default_transaction_read_only = on") { + t.Errorf("Prime: got %v, want it to set the read-only transaction default", enf.Prime) + } - if v := h.Statement(wire.Statement{SQL: "select * from t"}); v.Deny { + if v := deny(enf, "select id from t"); v.Deny { t.Errorf("select: got denied (%q), want allowed", v.Message) } - v := h.Statement(wire.Statement{SQL: "update t set x = 1"}) - if !v.Deny { - t.Fatal("update: got allowed, want denied") + v := deny(enf, "update t set x = 1") + if !v.Deny || !strings.Contains(v.Message, "UPDATE") { + t.Errorf("update: got %+v, want a read-only denial naming UPDATE", v) + } +} + +func TestReadOnlyRoleDeniesEscapeHatches(t *testing.T) { + t.Parallel() + + p := load(t, "roles:\n agent_ops:\n read_only: true\n") + enf := p.Resolve(wire.Startup{User: "agent_ops"}) + + hatches := []string{ + "set default_transaction_read_only = off", + `set "default_transaction_read_only" = off`, + "begin read write", + "select set_config('default_transaction_read_only','off',false)", } - if !strings.Contains(v.Message, "UPDATE") || !strings.Contains(v.Message, "read-only") { - t.Errorf("update denial message: got %q", v.Message) + for _, sql := range hatches { + v := deny(enf, sql) + if !v.Deny || !strings.Contains(v.Message, "read-write") { + t.Errorf("escape hatch %q: got %+v, want a read-write denial", sql, v) + } } } -func TestResolveWritableRoleAllowsEverything(t *testing.T) { +func TestWritableRoleAllowsEverythingAndDoesNotPrime(t *testing.T) { t.Parallel() p := load(t, "roles:\n app_rw:\n agent: web\n") - h := p.Resolve(wire.Startup{User: "app_rw"}) + enf := p.Resolve(wire.Startup{User: "app_rw"}) - if v := h.Statement(wire.Statement{SQL: "delete from t"}); v.Deny { + if len(enf.Prime) != 0 { + t.Errorf("Prime: got %v, want none for a writable role", enf.Prime) + } + if v := deny(enf, "delete from t"); v.Deny { t.Errorf("delete for a writable role: got denied (%q), want allowed", v.Message) } } -func TestResolveUnknownRoleHonorsFailMode(t *testing.T) { +func TestUnlistedRoleHonorsFailMode(t *testing.T) { t.Parallel() open := load(t, "roles: {}\n") - if v := open.Resolve(wire.Startup{User: "ghost"}).Statement(wire.Statement{SQL: "delete from t"}); v.Deny { + if v := deny(open.Resolve(wire.Startup{User: "ghost"}), "delete from t"); v.Deny { t.Errorf("fail-open unknown role: got denied (%q), want allowed", v.Message) } closed := load(t, "fail: closed\nroles: {}\n") - v := closed.Resolve(wire.Startup{User: "ghost"}).Statement(wire.Statement{SQL: "select 1"}) + v := deny(closed.Resolve(wire.Startup{User: "ghost"}), "select 1") if !v.Deny || !strings.Contains(v.Message, "ghost") { t.Errorf("fail-closed unknown role: got %+v, want a denial naming the role", v) } @@ -97,11 +118,15 @@ func TestZeroPolicyAllowsEverything(t *testing.T) { t.Parallel() var p policy.Policy - if v := p.Resolve(wire.Startup{User: "anyone"}).Statement(wire.Statement{SQL: "drop table t"}); v.Deny { + if v := deny(p.Resolve(wire.Startup{User: "anyone"}), "drop table t"); v.Deny { t.Errorf("zero policy: got denied (%q), want allowed", v.Message) } } +func deny(enf wire.Enforcement, sql string) wire.Verdict { + return enf.Handler.Statement(wire.Statement{SQL: sql}) +} + func load(t *testing.T, src string) policy.Policy { t.Helper() diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index b4b32c5..e3b2a1b 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -114,14 +114,23 @@ func (s Server) handle(ctx context.Context, client net.Conn) { } logger = logger.With("user", startup.User, "database", startup.Database, "application", startup.Application) - logger.Info("session opened") - handler := s.Guard.Resolve(startup) + enforcement := s.Guard.Resolve(startup) + for _, stmt := range enforcement.Prime { + if err := sess.Prime(stmt); err != nil { + if ctx.Err() == nil { + logger.Warn("prime", "error", err) + } + + return + } + } + logger.Info("session opened") var wg sync.WaitGroup var toUpstream, toClient error wg.Go(func() { - toUpstream = sess.Frontend(handler) + toUpstream = sess.Frontend(enforcement.Handler) closeWrite(upstream) }) wg.Go(func() { diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index 8b8b447..18bf5fb 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -145,6 +145,10 @@ func (rawSession) Handshake() (wire.Startup, error) { return wire.Startup{User: "raw"}, nil } +func (rawSession) Prime(string) error { + return nil +} + func (s rawSession) Frontend(wire.Handler) error { if _, err := io.Copy(s.upstream, s.client); err != nil { return fmt.Errorf("frontend: %w", err) diff --git a/internal/wire/wire.go b/internal/wire/wire.go index 7885c87..3c29cf2 100644 --- a/internal/wire/wire.go +++ b/internal/wire/wire.go @@ -45,21 +45,28 @@ func (f HandlerFunc) Statement(stmt Statement) Verdict { return f(stmt) } -// Guard resolves the Handler for a session from the client's identity. Resolve -// runs once, after the handshake. +// Enforcement is how a session is guarded: statements to run on the upstream +// before the relay starts, and the handler that judges each client statement. +type Enforcement struct { + Prime []string + Handler Handler +} + +// Guard resolves the enforcement for a session from the client's identity. +// Resolve runs once, after the handshake. type Guard interface { - Resolve(startup Startup) Handler + Resolve(startup Startup) Enforcement } -type GuardFunc func(startup Startup) Handler +type GuardFunc func(startup Startup) Enforcement -func (f GuardFunc) Resolve(startup Startup) Handler { +func (f GuardFunc) Resolve(startup Startup) Enforcement { return f(startup) } // AllowAll is a Guard whose sessions permit every statement. -var AllowAll Guard = GuardFunc(func(Startup) Handler { - return HandlerFunc(func(Statement) Verdict { return Verdict{} }) +var AllowAll Guard = GuardFunc(func(Startup) Enforcement { + return Enforcement{Handler: HandlerFunc(func(Statement) Verdict { return Verdict{} })} }) type Dialect interface { @@ -71,6 +78,10 @@ type Dialect interface { // until their side of the conversation ends. type Session interface { Handshake() (Startup, error) + // Prime runs one statement on the upstream and consumes its result before + // the relay starts, for session setup such as enabling read-only mode. It + // runs after Handshake and before Frontend and Backend. + Prime(sql string) error Frontend(h Handler) error Backend() error }