diff --git a/README.md b/README.md index 1646476..fabea81 100644 --- a/README.md +++ b/README.md @@ -21,7 +21,9 @@ 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. + +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. @@ -34,6 +36,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/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/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/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/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..d8da76b --- /dev/null +++ b/internal/policy/policy.go @@ -0,0 +1,151 @@ +// 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" +) + +// 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 { + // 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 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 { + return p.unlisted(startup.User) + } + + if !role.ReadOnly { + return wire.Enforcement{Handler: allow()} + } + + return wire.Enforcement{ + Prime: []string{readOnlyPrime}, + Handler: wire.HandlerFunc(readOnly), + } +} + +func (p Policy) unlisted(user string) wire.Enforcement { + if !p.FailClosed { + return wire.Enforcement{Handler: allow()} + } + + 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: "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 new file mode 100644 index 0000000..a99c1d8 --- /dev/null +++ b/internal/policy/policy_test.go @@ -0,0 +1,139 @@ +package policy_test + +import ( + "slices" + "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 TestReadOnlyRolePrimesAndDeniesWrites(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"}) + + 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 := deny(enf, "select id from t"); v.Deny { + t.Errorf("select: got denied (%q), want allowed", v.Message) + } + + 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)", + } + 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 TestWritableRoleAllowsEverythingAndDoesNotPrime(t *testing.T) { + t.Parallel() + + p := load(t, "roles:\n app_rw:\n agent: web\n") + enf := p.Resolve(wire.Startup{User: "app_rw"}) + + 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 TestUnlistedRoleHonorsFailMode(t *testing.T) { + t.Parallel() + + open := load(t, "roles: {}\n") + 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 := 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) + } +} + +func TestZeroPolicyAllowsEverything(t *testing.T) { + t.Parallel() + + var p policy.Policy + 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() + + p, err := policy.Parse([]byte(src)) + if err != nil { + t.Fatalf("Parse: %v", err) + } + + return p +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 8491db2..e3b2a1b 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) @@ -113,12 +114,23 @@ func (s Server) handle(ctx context.Context, client net.Conn) { } logger = logger.With("user", startup.User, "database", startup.Database, "application", startup.Application) + + 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(s.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/sqlscan/scan.go b/internal/sqlscan/scan.go new file mode 100644 index 0000000..d505ff8 --- /dev/null +++ b/internal/sqlscan/scan.go @@ -0,0 +1,185 @@ +package sqlscan + +import "strings" + +// 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 { + 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 = skipString(sql, i+1) + case c == '"': + 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([]token, 0, 8) + i++ + case isWordStart(c): + j := skipWord(sql, i+1) + cur = append(cur, token{kind: kindWord, text: strings.ToUpper(sql[i:j])}) + i = j + default: + i++ + } + } + stmts = append(stmts, cur) + + 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' && s[i] != '\r' { + 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 +} + +// 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 { + 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] + if k := strings.Index(s[j+1:], 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 +} + +// 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) || 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..130a066 --- /dev/null +++ b/internal/sqlscan/sqlscan.go @@ -0,0 +1,377 @@ +// 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 + +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 // COPY ... TO STDOUT + Insert + Update + Delete + Merge + Truncate + SelectInto + CopyFrom + CopyExternal // COPY ... TO a file or program + Call + DDL +) + +func (k Kind) String() string { + switch k { + case Unknown: + return "unknown statement" + case Empty: + return "empty statement" + 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 STDOUT" + case Insert: + return "INSERT" + case Update: + return "UPDATE" + case Delete: + return "DELETE" + case Merge: + 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: + 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. +func (k Kind) Mutating() bool { + switch k { + case Insert, Update, Delete, Merge, Truncate, SelectInto, CopyFrom, CopyExternal, Call, DDL, Unknown: + return true + case Empty, Select, Cursor, Set, Tx, Explain, CopyTo: + return false + } + + return true +} + +// 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(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 +} + +// Mutating reports whether any statement in sql can change data or schema. +func Mutating(sql string) bool { + for _, f := range Scan(sql) { + if f.Kind.Mutating() { + return true + } + } + + return false +} + +func classify(toks []token) Kind { + w, i := firstWord(toks) + if i < 0 { + return Empty + } + + switch w { + case "SELECT", "VALUES", "TABLE": + return selectOrInto(toks) + case "SHOW", "DECLARE", "FETCH", "MOVE", "CLOSE", "LISTEN", "UNLISTEN", "DISCARD", + "PREPARE", "DEALLOCATE", "EXECUTE", "NOTIFY", "CHECKPOINT": + 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(toks) + case "WITH": + return classifyWith(toks) + case "EXPLAIN": + return classifyExplain(toks) + case "CREATE", "ALTER", "DROP", "GRANT", "REVOKE", "COMMENT", + "REINDEX", "CLUSTER", "VACUUM", "ANALYZE", "REFRESH", "IMPORT", "LOAD", "SECURITY", "LOCK": + return DDL + } + + return Unknown +} + +// 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 + } + + 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: + } + } + + return CopyFrom // no clear direction: assume the writing form +} + +// classifyWith treats a WITH statement as mutating if any data-modifying command +// 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": + return Update + case "DELETE": + return Delete + case "MERGE": + return Merge + } + } + + return selectOrInto(toks) +} + +// classifyExplain returns EXPLAIN unless the plan runs with ANALYZE, which +// executes the underlying statement; then it classifies that statement. +func classifyExplain(toks []token) Kind { + if !containsWord(toks, "ANALYZE") { + return Explain // without ANALYZE the underlying statement never runs + } + + 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 new file mode 100644 index 0000000..0d09461 --- /dev/null +++ b/internal/sqlscan/sqlscan_test.go @@ -0,0 +1,164 @@ +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 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 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}}, + } + + 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 1 as x$a$`, + } + for _, sql := range readOnly { + if sqlscan.Mutating(sql) { + t.Errorf("Mutating(%q) = true, want false", sql) + } + } + + mutating := []string{ + `update t set x = 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`, + } + for _, sql := range mutating { + if !sqlscan.Mutating(sql) { + t.Errorf("Mutating(%q) = false, want true", sql) + } + } +} + +// 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() + + // 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 +} diff --git a/internal/wire/wire.go b/internal/wire/wire.go index ae4e301..3c29cf2 100644 --- a/internal/wire/wire.go +++ b/internal/wire/wire.go @@ -45,6 +45,30 @@ func (f HandlerFunc) Statement(stmt Statement) Verdict { return f(stmt) } +// 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) Enforcement +} + +type GuardFunc func(startup Startup) Enforcement + +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) Enforcement { + return Enforcement{Handler: HandlerFunc(func(Statement) Verdict { return Verdict{} })} +}) + type Dialect interface { NewSession(client, upstream net.Conn) Session } @@ -54,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 }