diff --git a/README.md b/README.md index fabea81..398227a 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. +With `-ledger PATH`, every statement is recorded to a JSON-lines file: which agent ran it, when, its kind, a fingerprint with literals removed (so no literal is stored), the decision, and how many rows it returned. The records are hash-chained so tampering shows. + 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. @@ -49,7 +51,7 @@ roles: ``` ```sh -rollcall proxy -upstream 127.0.0.1:5432 -policy policy.yaml +rollcall proxy -upstream 127.0.0.1:5432 -policy policy.yaml -ledger ledger.jsonl ``` ```sh diff --git a/internal/cli/cli_test.go b/internal/cli/cli_test.go index 1eebbab..18a44a3 100644 --- a/internal/cli/cli_test.go +++ b/internal/cli/cli_test.go @@ -4,8 +4,11 @@ import ( "bytes" "context" "encoding/binary" + "encoding/json" "io" "net" + "os" + "path/filepath" "regexp" "strings" "sync" @@ -14,6 +17,7 @@ import ( "github.com/mickamy/rollcall/internal/cli" "github.com/mickamy/rollcall/internal/exit" + "github.com/mickamy/rollcall/internal/ledger" ) func TestRun(t *testing.T) { @@ -238,6 +242,114 @@ func TestRunProxyServesPostgreSQL(t *testing.T) { } } +func TestRunProxyLedgerRecordsPreparedExecutions(t *testing.T) { + t.Parallel() + + upstream := startFakePostgres(t) + ledgerPath := filepath.Join(t.TempDir(), "ledger.jsonl") + + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + + errOut := newNotifyWriter("msg=listening") + code := make(chan int, 1) + go func() { + args := []string{"proxy", "-upstream", upstream, "-listen", "127.0.0.1:0", "-ledger", ledgerPath} + code <- cli.Run(ctx, args, cli.IO{In: strings.NewReader(""), Out: io.Discard, Err: errOut}) + }() + + select { + case <-errOut.seen: + case c := <-code: + t.Fatalf("proxy exited with %d before listening", c) + case <-time.After(5 * time.Second): + t.Fatal("proxy did not start listening") + } + addr := regexp.MustCompile(`addr=(\S+)`).FindStringSubmatch(errOut.String())[1] + + var dialer net.Dialer + client, err := dialer.DialContext(ctx, "tcp", addr) + if err != nil { + t.Fatalf("dial: %v", err) + } + defer func() { _ = client.Close() }() + _ = client.SetDeadline(time.Now().Add(5 * time.Second)) + + writeAll(t, client, startupPacket("user", "agent", "database", "app")) + expectMessage(t, client, pgMessage('R', be32(0))) + expectMessage(t, client, pgMessage('Z', []byte("I"))) + + // Prepare once, then re-execute in a second batch that carries no Parse. + parse := pgMessage('P', cstring("ps1"), cstring("select id from t where id > $1"), be16(0)) + bind := pgMessage('B', cstring(""), cstring("ps1"), be16(0), be16(0), be16(0)) + exec := pgMessage('E', cstring(""), be32(0)) + sync := pgMessage('S') + writeAll(t, client, bytes.Join([][]byte{parse, bind, exec, sync}, nil)) + readUntilReadyForQuery(t, client) + writeAll(t, client, bytes.Join([][]byte{bind, exec, sync}, nil)) + readUntilReadyForQuery(t, client) + + writeAll(t, client, pgMessage('X')) + _, _ = client.Read(make([]byte, 1)) + cancel() + <-code + + records := readLedger(t, ledgerPath) + if len(records) != 2 { + t.Fatalf("ledger records: got %d, want 2 (prepare + reuse)", len(records)) + } + for i, r := range records { + if r.Kind != "SELECT" || r.Decision != "allowed" || r.Rows != 1 { + t.Errorf("record %d: got %+v, want an allowed SELECT with 1 row", i, r) + } + if strings.Contains(r.Fingerprint, "1") { + t.Errorf("record %d fingerprint kept a literal: %q", i, r.Fingerprint) + } + } + if records[1].PrevHash != records[0].Hash { + t.Error("ledger records are not chained") + } +} + +func readUntilReadyForQuery(t *testing.T, r io.Reader) { + t.Helper() + + for { + header := make([]byte, 5) + if _, err := io.ReadFull(r, header); err != nil { + t.Fatalf("read: %v", err) + } + if _, err := io.CopyN(io.Discard, r, int64(binary.BigEndian.Uint32(header[1:]))-4); err != nil { + t.Fatalf("read body: %v", err) + } + if header[0] == 'Z' { + return + } + } +} + +func readLedger(t *testing.T, path string) []ledger.Record { + t.Helper() + + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read ledger: %v", err) + } + var out []ledger.Record + for line := range strings.SplitSeq(strings.TrimSpace(string(data)), "\n") { + if line == "" { + continue + } + var rec ledger.Record + if err := json.Unmarshal([]byte(line), &rec); err != nil { + t.Fatalf("unmarshal %q: %v", line, err) + } + out = append(out, rec) + } + + return out +} + // startFakePostgres serves trust authentication and answers every simple // query with "SELECT 1" until the client terminates. func startFakePostgres(t *testing.T) string { @@ -293,6 +405,16 @@ func serveFakePostgres(conn net.Conn) { if _, err := conn.Write(reply); err != nil { return } + case 'S': // Sync: answer the batch with a one-row result set + reply := bytes.Join([][]byte{ + pgMessage('T', be16(1), cstring("id"), be32(0), be16(0), be32(23), be16(4), be32(0xffffffff), be16(0)), + pgMessage('D', be16(1), be32(1), []byte("7")), + pgMessage('C', cstring("SELECT 1")), + pgMessage('Z', []byte("I")), + }, nil) + if _, err := conn.Write(reply); err != nil { + return + } case 'X': return } @@ -303,6 +425,10 @@ func be32(v uint32) []byte { return binary.BigEndian.AppendUint32(nil, v) } +func be16(v uint16) []byte { + return binary.BigEndian.AppendUint16(nil, v) +} + func cstring(s string) []byte { return append([]byte(s), 0) } diff --git a/internal/cli/proxy.go b/internal/cli/proxy.go index fcec168..3e07df9 100644 --- a/internal/cli/proxy.go +++ b/internal/cli/proxy.go @@ -8,8 +8,10 @@ import ( "io" "log/slog" "net" + "os" "github.com/mickamy/rollcall/internal/exit" + "github.com/mickamy/rollcall/internal/ledger" "github.com/mickamy/rollcall/internal/pg" "github.com/mickamy/rollcall/internal/policy" "github.com/mickamy/rollcall/internal/proxy" @@ -21,6 +23,7 @@ const ( 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" + ledgerUsage = "path to append the access ledger to as JSON lines; without one nothing is recorded" ) func runProxy(ctx context.Context, args []string, std IO) int { @@ -28,6 +31,7 @@ func runProxy(ctx context.Context, args []string, std IO) int { listen := fs.String("listen", defaultListen, listenUsage) upstream := fs.String("upstream", "", upstreamUsage) policyPath := fs.String("policy", "", policyUsage) + ledgerPath := fs.String("ledger", "", ledgerUsage) if err := fs.Parse(args); err != nil { if errors.Is(err, flag.ErrHelp) { return exit.OK @@ -57,6 +61,25 @@ func runProxy(ctx context.Context, args []string, std IO) int { guard = p } + logger := slog.New(slog.NewTextHandler(std.Err, nil)) + + if *ledgerPath != "" { + prev, err := ledger.LastHash(*ledgerPath) + if err != nil { + return fail(std, err) + } + + f, err := os.OpenFile(*ledgerPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0o600) + if err != nil { + return fail(std, fmt.Errorf("open ledger: %w", err)) + } + defer func() { _ = f.Close() }() + + sink := ledger.NewSink(f, ledger.Options{Prev: prev, Key: ledgerKey(), Logger: logger}) + defer sink.Close() + guard = ledger.Guard{Inner: guard, Sink: sink} + } + var lc net.ListenConfig ln, err := lc.Listen(ctx, "tcp", *listen) if err != nil { @@ -64,7 +87,6 @@ func runProxy(ctx context.Context, args []string, std IO) int { } defer func() { _ = ln.Close() }() - logger := slog.New(slog.NewTextHandler(std.Err, nil)) logger.Info("listening", "addr", ln.Addr().String(), "upstream", *upstream) if !isLoopback(*listen) { logger.Warn("listening outside loopback: clients and the upstream are served in plaintext", "addr", *listen) @@ -78,6 +100,16 @@ func runProxy(ctx context.Context, args []string, std IO) int { return exit.OK } +// ledgerKey returns the optional key that makes the ledger chain an HMAC, +// from ROLLCALL_LEDGER_KEY. Without it the chain is a plain SHA-256. +func ledgerKey() []byte { + if v := os.Getenv("ROLLCALL_LEDGER_KEY"); v != "" { + return []byte(v) + } + + return nil +} + func validateAddr(flagName, addr string) error { if addr == "" { return fmt.Errorf("%s is required", flagName) @@ -118,4 +150,5 @@ func printProxyUsage(w io.Writer) { 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) + fmt.Fprintf(w, " -ledger PATH %s\n", ledgerUsage) } diff --git a/internal/ledger/clock.go b/internal/ledger/clock.go new file mode 100644 index 0000000..ba4804a --- /dev/null +++ b/internal/ledger/clock.go @@ -0,0 +1,7 @@ +package ledger + +import "time" + +func defaultNow() string { + return time.Now().UTC().Format(time.RFC3339Nano) +} diff --git a/internal/ledger/ledger_test.go b/internal/ledger/ledger_test.go new file mode 100644 index 0000000..f6f1800 --- /dev/null +++ b/internal/ledger/ledger_test.go @@ -0,0 +1,290 @@ +package ledger_test + +import ( + "bufio" + "bytes" + "encoding/json" + "os" + "path/filepath" + "strings" + "testing" + + "github.com/mickamy/rollcall/internal/ledger" + "github.com/mickamy/rollcall/internal/wire" +) + +func TestGuardRecordsStatements(t *testing.T) { + t.Parallel() + + var buf bytes.Buffer + inner := wire.GuardFunc(func(s wire.Startup) wire.Enforcement { + return wire.Enforcement{ + Principal: wire.Principal{ + Agent: "claude-ops", Purpose: "incident", User: s.User, Database: s.Database, + }, + Handler: wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }), + } + }) + sink := ledger.NewSink(&buf, ledger.Options{}) + g := ledger.Guard{Inner: inner, Sink: sink} + + enf := g.Resolve(wire.Startup{User: "agent_ops", Database: "prod"}) + + // An allowed SELECT that returns three rows. + res := enf.Recorder.Begin("select id from orders where email = 'a@b.com'", wire.Allowed) + res.Columns([]string{"id"}) + res.Complete("SELECT 3") + res.Done() + + // A denied UPDATE. + enf.Recorder.Begin("update orders set x = 1", wire.Denied).Done() + + sink.Close() + recs := decode(t, &buf) + if len(recs) != 2 { + t.Fatalf("records: got %d, want 2", len(recs)) + } + + first := recs[0] + gotPrincipal := [4]string{first.Agent, first.Purpose, first.User, first.Database} + if gotPrincipal != [4]string{"claude-ops", "incident", "agent_ops", "prod"} { + t.Errorf("first record principal: got %v", gotPrincipal) + } + if first.Kind != "SELECT" || first.Rows != 3 { + t.Errorf("first record: got kind=%q rows=%d, want SELECT/3", first.Kind, first.Rows) + } + if first.Fingerprint != "SELECT ID FROM ORDERS WHERE EMAIL = ?" { + t.Errorf("first fingerprint: got %q", first.Fingerprint) + } + if strings.Contains(first.Fingerprint, "a@b.com") { + t.Error("fingerprint leaked a literal") + } + if recs[1].Decision != wire.Denied || recs[1].Kind != "UPDATE" { + t.Errorf("second record: got decision=%q kind=%q", recs[1].Decision, recs[1].Kind) + } +} + +func TestSinkChainsHashes(t *testing.T) { + t.Parallel() + + var buf bytes.Buffer + s := ledger.NewSink(&buf, ledger.Options{}) + for range 3 { + s.Write(ledger.Record{User: "u", Kind: "SELECT"}) + } + s.Close() + + recs := decode(t, &buf) + if recs[0].PrevHash != "" { + t.Errorf("first PrevHash: got %q, want empty", recs[0].PrevHash) + } + for i := 1; i < len(recs); i++ { + if recs[i].PrevHash != recs[i-1].Hash { + t.Errorf("record %d PrevHash %q != previous Hash %q", i, recs[i].PrevHash, recs[i-1].Hash) + } + if recs[i].Hash == "" || recs[i].Hash == recs[i-1].Hash { + t.Errorf("record %d Hash %q not distinct", i, recs[i].Hash) + } + } +} + +func decode(t *testing.T, buf *bytes.Buffer) []ledger.Record { + t.Helper() + + var out []ledger.Record + sc := bufio.NewScanner(buf) + for sc.Scan() { + var rec ledger.Record + if err := json.Unmarshal(sc.Bytes(), &rec); err != nil { + t.Fatalf("unmarshal %q: %v", sc.Text(), err) + } + out = append(out, rec) + } + + return out +} + +func TestSinkResumesChainWithPrevAndKey(t *testing.T) { + t.Parallel() + + var first bytes.Buffer + s1 := ledger.NewSink(&first, ledger.Options{Key: []byte("secret")}) + s1.Write(ledger.Record{User: "u", Kind: "SELECT"}) + s1.Write(ledger.Record{User: "u", Kind: "SELECT"}) + s1.Close() + firstRecs := decode(t, &first) + last := firstRecs[len(firstRecs)-1].Hash + + // A new sink seeded with the last hash continues the same chain. + var second bytes.Buffer + s2 := ledger.NewSink(&second, ledger.Options{Prev: last, Key: []byte("secret")}) + s2.Write(ledger.Record{User: "u", Kind: "DELETE"}) + s2.Close() + next := decode(t, &second) + + if next[0].PrevHash != last { + t.Errorf("resumed PrevHash: got %q, want %q", next[0].PrevHash, last) + } + if next[0].Hash == "" { + t.Error("resumed record has no hash") + } +} + +func TestFingerprintIsBounded(t *testing.T) { + t.Parallel() + + var buf bytes.Buffer + inner := wire.GuardFunc(func(wire.Startup) wire.Enforcement { + return wire.Enforcement{Handler: wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} })} + }) + sink := ledger.NewSink(&buf, ledger.Options{}) + enf := ledger.Guard{Inner: inner, Sink: sink}.Resolve(wire.Startup{User: "u"}) + + huge := "select " + strings.Repeat("col_a, ", 5000) + "col_z from t" + enf.Recorder.Begin(huge, wire.Allowed).Done() + sink.Close() + + rec := decode(t, &buf)[0] + if len(rec.Fingerprint) > 4200 { + t.Errorf("fingerprint length %d not bounded", len(rec.Fingerprint)) + } + if !strings.Contains(rec.Fingerprint, "#") { + t.Error("truncated fingerprint has no hash suffix") + } +} + +func TestRecordsCarryKeyID(t *testing.T) { + t.Parallel() + + var keyed, plain bytes.Buffer + k := ledger.NewSink(&keyed, ledger.Options{Key: []byte("secret")}) + k.Write(ledger.Record{User: "u", Kind: "SELECT"}) + k.Close() + p := ledger.NewSink(&plain, ledger.Options{}) + p.Write(ledger.Record{User: "u", Kind: "SELECT"}) + p.Close() + + if id := decode(t, &keyed)[0].KeyID; id == "" { + t.Error("keyed record has no key_id") + } + if id := decode(t, &plain)[0].KeyID; id != "" { + t.Errorf("unkeyed record key_id: got %q, want empty", id) + } +} + +func TestLastHashReadsTheTail(t *testing.T) { + t.Parallel() + + path := filepath.Join(t.TempDir(), "ledger.jsonl") + f, err := os.Create(path) + if err != nil { + t.Fatal(err) + } + sink := ledger.NewSink(f, ledger.Options{}) + for range 5 { + sink.Write(ledger.Record{User: "u", Kind: "SELECT"}) + } + sink.Close() + last := sink.LastHash() + _ = f.Close() + + got, err := ledger.LastHash(path) + if err != nil { + t.Fatalf("LastHash: %v", err) + } + if got != last { + t.Errorf("LastHash: got %q, want %q", got, last) + } + + // A partial trailing write must fall back to the last complete record. + appendString(t, path, `{"partial":`) + got, err = ledger.LastHash(path) + if err != nil { + t.Fatalf("LastHash after partial write: %v", err) + } + if got != last { + t.Errorf("LastHash ignored a partial line: got %q, want %q", got, last) + } +} + +func TestLastHashMissingFile(t *testing.T) { + t.Parallel() + + got, err := ledger.LastHash(filepath.Join(t.TempDir(), "absent.jsonl")) + if err != nil || got != "" { + t.Errorf("LastHash on a missing file: got %q, %v; want empty, nil", got, err) + } +} + +func appendString(t *testing.T, path, s string) { + t.Helper() + + f, err := os.OpenFile(path, os.O_APPEND|os.O_WRONLY, 0o600) + if err != nil { + t.Fatalf("open: %v", err) + } + defer func() { _ = f.Close() }() + if _, err := f.WriteString(s); err != nil { + t.Fatalf("write: %v", err) + } +} + +// TestChainResumesAcrossAKeyChange mirrors what the CLI does on restart: read +// the file's last hash, then continue the chain with a possibly different key. +func TestChainResumesAcrossAKeyChange(t *testing.T) { + t.Parallel() + + path := filepath.Join(t.TempDir(), "ledger.jsonl") + + writeWith := func(key []byte, kind string) { + prev, err := ledger.LastHash(path) + if err != nil { + t.Fatalf("LastHash: %v", err) + } + f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0o600) + if err != nil { + t.Fatal(err) + } + s := ledger.NewSink(f, ledger.Options{Prev: prev, Key: key}) + s.Write(ledger.Record{User: "u", Kind: kind}) + s.Close() + _ = f.Close() + } + + writeWith([]byte("alpha"), "SELECT") + writeWith([]byte("alpha"), "SELECT") + writeWith([]byte("beta"), "DELETE") // restart with a different key + + f, err := os.Open(path) + if err != nil { + t.Fatal(err) + } + defer func() { _ = f.Close() }() + recs := decode(t, bufferOf(t, f)) + + if len(recs) != 3 { + t.Fatalf("records: got %d, want 3", len(recs)) + } + for i := 1; i < len(recs); i++ { + if recs[i].PrevHash != recs[i-1].Hash { + t.Errorf("record %d does not link to the previous", i) + } + } + if recs[0].KeyID == recs[2].KeyID { + t.Error("key_id did not change when the key changed") + } + if recs[0].KeyID != recs[1].KeyID { + t.Error("key_id changed while the key stayed the same") + } +} + +func bufferOf(t *testing.T, f *os.File) *bytes.Buffer { + t.Helper() + + data, err := os.ReadFile(f.Name()) + if err != nil { + t.Fatal(err) + } + + return bytes.NewBuffer(data) +} diff --git a/internal/ledger/record.go b/internal/ledger/record.go new file mode 100644 index 0000000..4a0fa1f --- /dev/null +++ b/internal/ledger/record.go @@ -0,0 +1,27 @@ +// Package ledger records which agent ran which statement, when, and to what +// effect. Records carry no statement literals and no result values: SQL is +// stored as a fingerprint, and the ledger is chained so tampering shows. +package ledger + +import "github.com/mickamy/rollcall/internal/wire" + +// Record is one ledger entry: one statement handled by one session. +type Record struct { + Time string `json:"time"` + Agent string `json:"agent,omitempty"` + Purpose string `json:"purpose,omitempty"` + User string `json:"user"` + Database string `json:"database"` + Application string `json:"application,omitempty"` + Kind string `json:"kind"` + Fingerprint string `json:"fingerprint"` + Decision wire.Decision `json:"decision"` + Rows int `json:"rows"` + // KeyID names the chain key (a prefix of its hash, empty when unkeyed), so a + // verifier knows which key each record was signed with and can see where the + // key changed across restarts. + KeyID string `json:"key_id,omitempty"` + // PrevHash and Hash chain the records; Hash covers this record and PrevHash. + PrevHash string `json:"prev_hash"` + Hash string `json:"hash"` +} diff --git a/internal/ledger/recorder.go b/internal/ledger/recorder.go new file mode 100644 index 0000000..0542b3e --- /dev/null +++ b/internal/ledger/recorder.go @@ -0,0 +1,128 @@ +package ledger + +import ( + "crypto/sha256" + "encoding/hex" + "strings" + + "github.com/mickamy/rollcall/internal/sqlscan" + "github.com/mickamy/rollcall/internal/wire" +) + +// maxFingerprint bounds a record's fingerprint so a connected agent cannot grow +// the ledger a statement at a time. A longer fingerprint is truncated and given +// a hash suffix, so distinct large statements still differ. +const maxFingerprint = 4096 + +func fingerprint(sql string) string { + fp := sqlscan.Fingerprint(sql) + if len(fp) <= maxFingerprint { + return fp + } + + sum := sha256.Sum256([]byte(fp)) + + return fp[:maxFingerprint] + "…#" + hex.EncodeToString(sum[:8]) +} + +// Guard wraps another guard, recording each session's statements to Sink. +type Guard struct { + Inner wire.Guard + Sink *Sink +} + +var _ wire.Guard = Guard{} + +func (g Guard) Resolve(startup wire.Startup) wire.Enforcement { + enf := g.Inner.Resolve(startup) + enf.Recorder = recorder{principal: enf.Principal, sink: g.Sink} + + return enf +} + +// now is overridable in tests; the ledger stamps records with the wall clock. +var now = defaultNow + +// statementKind names the kinds in sql. A single statement gives its kind; a +// multi-statement query lists each distinct kind in order, so a +// "SELECT 1; DELETE FROM t" is not recorded as a plain SELECT. +func statementKind(sql string) string { + kinds := sqlscan.Classify(sql) + if len(kinds) == 0 { + return sqlscan.Empty.String() + } + + seen := make(map[string]bool, len(kinds)) + names := make([]string, 0, len(kinds)) + for _, k := range kinds { + if k == sqlscan.Empty { + continue + } + name := k.String() + if seen[name] { + continue + } + seen[name] = true + names = append(names, name) + } + if len(names) == 0 { + return sqlscan.Empty.String() + } + + return strings.Join(names, ", ") +} + +type recorder struct { + principal wire.Principal + sink *Sink +} + +var _ wire.Recorder = recorder{} + +func (r recorder) Begin(sql string, decision wire.Decision) wire.Result { + kind := statementKind(sql) + + rec := Record{ + Time: now(), + Agent: r.principal.Agent, + Purpose: r.principal.Purpose, + User: r.principal.User, + Database: r.principal.Database, + Application: r.principal.Application, + Kind: kind, + Fingerprint: fingerprint(sql), + Decision: decision, + } + + return &result{rec: rec, sink: r.sink} +} + +// result accumulates a statement's outcome and writes the record when the +// statement finishes. +type result struct { + rec Record + sink *Sink + done bool +} + +var _ wire.Result = (*result)(nil) + +// Columns captures nothing yet; subject extraction is a later step. +func (result) Columns([]string) []int { + return nil +} + +func (result) Row([][]byte) {} + +func (r *result) Complete(tag string) { + r.rec.Rows += rowsFromTag(tag) +} + +func (r *result) Done() { + if r.done { + return + } + r.done = true + + r.sink.Write(r.rec) +} diff --git a/internal/ledger/resume.go b/internal/ledger/resume.go new file mode 100644 index 0000000..38ff085 --- /dev/null +++ b/internal/ledger/resume.go @@ -0,0 +1,77 @@ +package ledger + +import ( + "bytes" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "slices" +) + +// tailWindow bounds how far back LastHash reads. Records are far smaller than +// this (fingerprints are capped), so the last complete line is always within it. +const tailWindow = 1 << 20 + +// LastHash returns the Hash of the last record in the ledger file at path, so a +// new Sink can continue the chain across restarts. It reads only the tail of the +// file. A missing or empty file yields "". +func LastHash(path string) (string, error) { + f, err := os.Open(path) + if errors.Is(err, os.ErrNotExist) { + return "", nil + } + if err != nil { + return "", fmt.Errorf("open ledger: %w", err) + } + defer func() { _ = f.Close() }() + + info, err := f.Stat() + if err != nil { + return "", fmt.Errorf("stat ledger: %w", err) + } + + size := info.Size() + if size == 0 { + return "", nil + } + + start, length := int64(0), size + if size > tailWindow { + start, length = size-tailWindow, tailWindow + } + buf := make([]byte, length) + if _, err := f.ReadAt(buf, start); err != nil && !errors.Is(err, io.EOF) { + return "", fmt.Errorf("read ledger tail: %w", err) + } + + return lastHashInTail(buf, start > 0) +} + +// lastHashInTail returns the hash of the last complete record in tail. windowed +// reports whether the tail was cut from a larger file, in which case the first +// line may be incomplete and is not trusted. +func lastHashInTail(tail []byte, windowed bool) (string, error) { + lines := bytes.Split(tail, []byte{'\n'}) + for i, raw := range slices.Backward(lines) { + line := bytes.TrimSpace(raw) + if len(line) == 0 { + continue + } + + var rec struct { + Hash string `json:"hash"` + } + if err := json.Unmarshal(line, &rec); err == nil && rec.Hash != "" { + return rec.Hash, nil + } + // A partial trailing write, or a line with no hash; try the record before + // it, unless the tail window may have cut this first line short. + if i == 0 && windowed { + return "", fmt.Errorf("ledger: last record exceeds %d bytes; cannot resume", tailWindow) + } + } + + return "", nil +} diff --git a/internal/ledger/sink.go b/internal/ledger/sink.go new file mode 100644 index 0000000..5d58f29 --- /dev/null +++ b/internal/ledger/sink.go @@ -0,0 +1,150 @@ +package ledger + +import ( + "crypto/hmac" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "hash" + "log/slog" + "sync" +) + +// Sink appends records as JSON lines, chaining each record's hash to the +// previous so the log is tamper-evident. Records are queued to a single writer +// goroutine and never dropped, so disk latency does not block a session's +// response path until the queue fills, after which Write blocks (backpressure) +// rather than dropping a record. +type Sink struct { + ch chan Record + done chan struct{} + w writer + key []byte + keyID string + logger *slog.Logger + + mu sync.Mutex + prev string +} + +type writer interface { + Write(p []byte) (int, error) +} + +// Options configure a Sink. +type Options struct { + // Prev is the hash of the last record already in the output, so appending + // across restarts keeps one unbroken chain. Empty starts a new chain. + Prev string + // Key, when set, makes the chain an HMAC so it cannot be recomputed without + // the key. Empty falls back to a plain SHA-256 chain. + Key []byte + Logger *slog.Logger +} + +const queueDepth = 1024 + +// NewSink writes records to w and starts its writer goroutine. Close it to flush. +func NewSink(w writer, opts Options) *Sink { + logger := opts.Logger + if logger == nil { + logger = slog.New(slog.DiscardHandler) + } + + s := &Sink{ + ch: make(chan Record, queueDepth), + done: make(chan struct{}), + w: w, + key: opts.Key, + keyID: keyID(opts.Key), + logger: logger, + prev: opts.Prev, + } + go s.run() + + return s +} + +// Write queues one record. It returns at once; the writer goroutine chains and +// appends it. Records are chained in the order Write is called. +func (s *Sink) Write(rec Record) { + s.ch <- rec +} + +// Close stops accepting records and waits for the queue to drain. +func (s *Sink) Close() { + close(s.ch) + <-s.done +} + +// LastHash reports the hash of the most recently written record, for a caller +// that wants to seed a later Sink. Safe to call after Close. +func (s *Sink) LastHash() string { + s.mu.Lock() + defer s.mu.Unlock() + + return s.prev +} + +func (s *Sink) run() { + defer close(s.done) + + for rec := range s.ch { + rec.KeyID = s.keyID + rec.PrevHash = s.prev + rec.Hash = s.hashRecord(rec) + + line, err := json.Marshal(rec) + if err != nil { + s.logger.Error("ledger marshal", "error", err) + + continue + } + if _, err := s.w.Write(append(line, '\n')); err != nil { + s.logger.Error("ledger write", "error", err) + + continue + } + + s.setPrev(rec.Hash) + } +} + +// keyID is a short, non-reversible name for a chain key: a prefix of its hash, +// or empty when unkeyed. +func keyID(key []byte) string { + if len(key) == 0 { + return "" + } + + m := hmac.New(sha256.New, key) + _, _ = m.Write([]byte("rollcall/ledger/key-id")) + + return hex.EncodeToString(m.Sum(nil)[:8]) +} + +// hashRecord hashes the record with its Hash field cleared, so the chain covers +// order and content. With a key it is an HMAC; without, a plain SHA-256. +func (s *Sink) hashRecord(rec Record) string { + rec.Hash = "" + data, err := json.Marshal(rec) + if err != nil { + return "" + } + + var h hash.Hash + if len(s.key) > 0 { + h = hmac.New(sha256.New, s.key) + } else { + h = sha256.New() + } + _, _ = h.Write(data) + + return hex.EncodeToString(h.Sum(nil)) +} + +func (s *Sink) setPrev(h string) { + s.mu.Lock() + s.prev = h + s.mu.Unlock() +} diff --git a/internal/ledger/tag.go b/internal/ledger/tag.go new file mode 100644 index 0000000..39cc7a1 --- /dev/null +++ b/internal/ledger/tag.go @@ -0,0 +1,23 @@ +package ledger + +import ( + "strconv" + "strings" +) + +// rowsFromTag reads the affected-row count from a CommandComplete tag such as +// "SELECT 5", "INSERT 0 3", "UPDATE 2", or "COPY 10"; it returns 0 when the tag +// carries no count. +func rowsFromTag(tag string) int { + fields := strings.Fields(tag) + if len(fields) == 0 { + return 0 + } + + n, err := strconv.Atoi(fields[len(fields)-1]) + if err != nil { + return 0 + } + + return n +} diff --git a/internal/pg/message.go b/internal/pg/message.go index 390b68d..81edace 100644 --- a/internal/pg/message.go +++ b/internal/pg/message.go @@ -28,21 +28,26 @@ const ( ) const ( - typeQuery = 'Q' - typeParse = 'P' - typeBind = 'B' - typeDescribe = 'D' - typeExecute = 'E' - typeClose = 'C' - typeFlush = 'H' - typeSync = 'S' - typeFunctionCall = 'F' - typeTerminate = 'X' - typePassword = 'p' - typeAuthentication = 'R' - typeErrorResponse = 'E' - typeReadyForQuery = 'Z' - typeCopyInResponse = 'G' + typeQuery = 'Q' + typeParse = 'P' + typeBind = 'B' + typeDescribe = 'D' + typeExecute = 'E' + typeClose = 'C' + typeFlush = 'H' + typeSync = 'S' + typeFunctionCall = 'F' + typeTerminate = 'X' + typePassword = 'p' + typeAuthentication = 'R' + typeErrorResponse = 'E' + typeReadyForQuery = 'Z' + typeCopyInResponse = 'G' + typeRowDescription = 'T' + typeDataRow = 'D' + typeCommandComplete = 'C' + typeEmptyQuery = 'I' + typePortalSuspended = 's' ) const ( @@ -194,6 +199,73 @@ func flush(w *bufio.Writer) error { return nil } +// columnNames reads the field names from a RowDescription body. +func columnNames(body []byte) []string { + if len(body) < 2 { + return nil + } + count := int(binary.BigEndian.Uint16(body)) + body = body[2:] + + names := make([]string, 0, count) + for range count { + name, rest, err := cstring(body) + if err != nil { + break + } + names = append(names, name) + if len(rest) < 18 { // tableOID(4) col(2) typeOID(4) len(2) mod(4) format(2) + break + } + body = rest[18:] + } + + return names +} + +// rowValues reads the values of the given column indices from a DataRow body. +// A null column yields a nil slice. +func rowValues(body []byte, capture []int) [][]byte { + if len(body) < 2 { + return nil + } + count := int(binary.BigEndian.Uint16(body)) + body = body[2:] + + fields := make([][]byte, 0, count) + for range count { + if len(body) < 4 { + break + } + n := int32(binary.BigEndian.Uint32(body)) //nolint:gosec // length is a signed int32 by protocol + body = body[4:] + if n < 0 { + fields = append(fields, nil) + + continue + } + if len(body) < int(n) { + break + } + fields = append(fields, body[:n]) + body = body[n:] + } + + out := make([][]byte, 0, len(capture)) + for _, c := range capture { + if c >= 0 && c < len(fields) { + out = append(out, fields[c]) + } + } + + return out +} + +// commandTag reads the tag from a CommandComplete body, dropping its terminator. +func commandTag(body []byte) string { + return string(bytes.TrimSuffix(body, []byte{0})) +} + func cstring(b []byte) (s string, rest []byte, err error) { value, rest, ok := bytes.Cut(b, []byte{0}) if !ok { diff --git a/internal/pg/pg.go b/internal/pg/pg.go index 465dfda..009654f 100644 --- a/internal/pg/pg.go +++ b/internal/pg/pg.go @@ -27,6 +27,10 @@ const ( // maxQueue bounds outstanding responses. Reaching it blocks the frontend // until the backend drains one, applying backpressure like the server does. maxQueue = 512 + // maxPreparedCount and maxPreparedBytes bound the prepared-statement bookkeeping + // so a client cannot exhaust proxy memory by preparing statements without end. + maxPreparedCount = 8192 + maxPreparedBytes = 16 << 20 // smallBody bounds the messages Backend reads fully before taking the // client write lock, so a stalled upstream cannot hold denials hostage. smallBody = 32 << 10 @@ -58,6 +62,8 @@ func (d Dialect) NewSession(client, upstream net.Conn) wire.Session { ur: bufio.NewReaderSize(upstream, readBufferSize), uw: bufio.NewWriterSize(upstream, writeBufferSize), small: make([]byte, smallBody), + prepared: make(map[string]string), + portals: make(map[string]string), } s.drained = sync.NewCond(&s.mu) @@ -72,6 +78,8 @@ func (d Dialect) NewSession(client, upstream net.Conn) wire.Session { type slot struct { forwarded bool simple bool + result wire.Result + execCount int // extended-protocol Executes whose records this slot finalizes code string message string hint string @@ -82,24 +90,32 @@ type session struct { maxStatement uint32 // Frontend goroutine only. - cr *bufio.Reader - uw *bufio.Writer - pending bytes.Buffer - forwarded bool // part of the current batch was already sent upstream - denied bool // a statement in the current batch was denied + cr *bufio.Reader + uw *bufio.Writer + recorder wire.Recorder + pending bytes.Buffer + prepared map[string]string // prepared statement name -> SQL + portals map[string]string // portal name -> prepared statement name + pendingExecs []string // SQL (or "" when unresolved) of each Execute in the current batch + preparedSize int // total bytes of SQL held in prepared, bounded by maxPreparedBytes + forwarded bool // part of the current batch was already sent upstream + denied bool // a statement in the current batch was denied // Backend goroutine only. - ur *bufio.Reader - small []byte + ur *bufio.Reader + small []byte + curCap []int // columns the recorder asked to capture for the current result set + execIdx int // index of the extended-protocol Execute currently answering // mu serializes writes to the client, which Frontend (denials) and Backend // (forwarded messages) both perform, and guards tx and queue alongside them. // drained wakes the frontend when the backend frees a queue slot. - mu sync.Mutex - drained *sync.Cond - cw *bufio.Writer - tx byte - queue []slot + mu sync.Mutex + drained *sync.Cond + cw *bufio.Writer + tx byte + queue []slot + execQueue []wire.Result // extended-protocol Execute records awaiting their results } func (s *session) Handshake() (wire.Startup, error) { @@ -169,7 +185,8 @@ func (s *session) Prime(sql string) error { // 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. -func (s *session) Frontend(h wire.Handler) error { +func (s *session) Frontend(h wire.Handler, rec wire.Recorder) error { + s.recorder = rec for { if s.cr.Buffered() == 0 { if err := flush(s.uw); err != nil { @@ -197,8 +214,17 @@ func (s *session) Frontend(h wire.Handler) error { } // Backend relays upstream messages and settles the request queue on every -// ReadyForQuery, writing any denials that were waiting their turn. +// ReadyForQuery, writing any denials that were waiting their turn. On exit it +// finalizes records for statements still in flight, so a session that ends +// before a ReadyForQuery still records what it sent. func (s *session) Backend() error { + err := s.backend() + s.finishInFlight() + + return err +} + +func (s *session) backend() error { for { if s.ur.Buffered() == 0 { if err := s.flushClient(); err != nil { @@ -224,6 +250,14 @@ func (s *session) Backend() error { if err = s.forwardToClient(typ, n); err == nil { s.copyInStarted() } + case typeRowDescription: + err = s.describeResult(n) + case typeDataRow: + err = s.rowResult(n) + case typeCommandComplete: + err = s.completeResult(n) + case typeEmptyQuery, typePortalSuspended: + err = s.endResult(typ, n) default: err = s.forwardToClient(typ, n) } @@ -284,7 +318,13 @@ func (s *session) dispatch(h wire.Handler, typ byte, n uint32) (bool, error) { return false, s.query(h, n) case typeParse: return false, s.parse(h, n) - case typeBind, typeDescribe, typeExecute, typeClose: + case typeBind: + return false, s.bind(n) + case typeExecute: + return false, s.execute(n) + case typeClose: + return false, s.closePrepared(n) + case typeDescribe: return false, s.buffer(typ, n) case typeFlush: return false, s.flushBatch(n) @@ -318,6 +358,7 @@ func (s *session) query(h wire.Handler, n uint32) error { if err := discard(s.cr, n); err != nil { return err } + s.recordDenied("") return s.respond(readySlot(s.tooLarge(n))) } @@ -329,13 +370,15 @@ func (s *session) query(h wire.Handler, n uint32) error { sql := string(bytes.TrimSuffix(body, []byte{0})) if v := h.Statement(wire.Statement{SQL: sql}); v.Deny { + s.recordDenied(sql) + return s.respond(readySlot(denial(v))) } // Enqueue before the statement can reach the upstream, so Backend never sees // the reply before the slot that accounts for it. A simple query's // ReadyForQuery survives COPY, so mark it. - s.enqueue(slot{forwarded: true, simple: true}) + s.enqueue(slot{forwarded: true, simple: true, result: s.begin(sql, wire.Allowed)}) return writeMessage(s.uw, typeQuery, body) } @@ -351,6 +394,8 @@ func (s *session) parse(h wire.Handler, n uint32) error { if err := discard(s.cr, n); err != nil { return err } + s.recordDenied("") + s.pendingExecs = nil return s.denyBatch(s.tooLarge(n)) } @@ -360,18 +405,154 @@ func (s *session) parse(h wire.Handler, n uint32) error { return err } - sql, err := parseSQL(body) + name, sql, err := parseNameSQL(body) if err != nil { return err } if v := h.Statement(wire.Statement{SQL: sql}); v.Deny { + s.recordDenied(sql) + s.pendingExecs = nil // the batch is rejected; earlier statements never run + return s.denyBatch(denial(v)) } + // Remember the prepared statement so its later Executes, which carry no SQL, + // can be attributed even across batches (driver statement caches reuse it). + s.storePrepared(name, sql) return s.stage(typeParse, body) } +// bind records which prepared statement a portal is bound to, then stages the +// message unchanged. +func (s *session) bind(n uint32) error { + if s.denied { + return discard(s.cr, n) + } + if s.forwarded || int64(n) > maxPending { + if err := s.forwardPending(); err != nil { + return err + } + s.forwarded = true + + return forward(s.uw, typeBind, n, s.cr) // too large to inspect; skip the portal map + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + if portal, stmt, ok := parseBind(body); ok && s.recorder != nil { + if _, exists := s.portals[portal]; exists || len(s.portals) < maxPreparedCount { + s.portals[portal] = stmt + } + } + + return s.stage(typeBind, body) +} + +// execute stages an Execute and, when its portal resolves to a known prepared +// statement, queues that statement's SQL to be recorded at Sync. +func (s *session) execute(n uint32) error { + if s.denied { + return discard(s.cr, n) + } + if s.forwarded || int64(n) > maxPending { + if err := s.forwardPending(); err != nil { + return err + } + s.forwarded = true + s.stageExec("") // forwarded without inspection; a placeholder keeps results aligned + + return forward(s.uw, typeExecute, n, s.cr) + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + s.stageExec(s.resolveExec(body)) + + return s.stage(typeExecute, body) +} + +// resolveExec returns the SQL an Execute runs, or "" when its portal is unknown. +func (s *session) resolveExec(body []byte) string { + portal, ok := parsePortal(body) + if !ok { + return "" + } + stmt, ok := s.portals[portal] + if !ok { + return "" + } + + return s.prepared[stmt] +} + +// stageExec records one Execute in the current batch, keeping a placeholder for +// every Execute so its result set maps to the right record. Only tracked when a +// recorder is set. +func (s *session) stageExec(sql string) { + if s.recorder == nil { + return + } + s.pendingExecs = append(s.pendingExecs, sql) +} + +// storePrepared records a prepared statement's SQL for later Execute attribution, +// bounded so the map cannot grow without limit. Tracked only when recording. +func (s *session) storePrepared(name, sql string) { + if s.recorder == nil { + return + } + if old, ok := s.prepared[name]; ok { + s.preparedSize -= len(old) + } + if len(s.prepared) >= maxPreparedCount || s.preparedSize+len(sql) > maxPreparedBytes { + delete(s.prepared, name) // at the limit: stop tracking this name rather than grow + return + } + s.prepared[name] = sql + s.preparedSize += len(sql) +} + +// closePrepared handles a Close message, dropping the named statement or portal +// from the maps so the server's release is mirrored, then stages the message. +func (s *session) closePrepared(n uint32) error { + if s.denied { + return discard(s.cr, n) + } + if s.forwarded || int64(n) > maxPending { + if err := s.forwardPending(); err != nil { + return err + } + s.forwarded = true + + return forward(s.uw, typeClose, n, s.cr) + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + if s.recorder != nil && len(body) > 1 { + if name, _, err := cstring(body[1:]); err == nil { + switch body[0] { + case 'S': + if sql, ok := s.prepared[name]; ok { + s.preparedSize -= len(sql) + delete(s.prepared, name) + } + case 'P': + delete(s.portals, name) + } + } + } + + return s.stage(typeClose, body) +} + func (s *session) buffer(typ byte, n uint32) error { if s.denied { return discard(s.cr, n) @@ -446,12 +627,17 @@ func (s *session) sync(n uint32) error { } } + answer := slot{forwarded: true} + if !s.denied { + answer.execCount = s.enqueueExecs() + } + s.pendingExecs = nil s.denied = false s.forwarded = false // Enqueue before the Sync reaches the upstream, so Backend never sees the // ReadyForQuery before the slot that accounts for it. - s.enqueue(slot{forwarded: true}) + s.enqueue(answer) if err := writeMessage(s.uw, typeSync, nil); err != nil { return err } @@ -459,6 +645,27 @@ func (s *session) sync(n uint32) error { return flush(s.uw) } +// enqueueExecs starts a record for each Execute in the batch and queues it for +// the backend to fill with its result. It returns how many records it queued. +func (s *session) enqueueExecs() int { + if len(s.pendingExecs) == 0 { + return 0 + } + + s.mu.Lock() + defer s.mu.Unlock() + + for _, sql := range s.pendingExecs { + var r wire.Result + if sql != "" { + r = s.begin(sql, wire.Allowed) + } + s.execQueue = append(s.execQueue, r) // nil keeps an unrecorded Execute in position + } + + return len(s.pendingExecs) +} + func (s *session) functionCall(n uint32) error { if err := discard(s.cr, n); err != nil { return err @@ -481,6 +688,7 @@ func (s *session) functionCall(n uint32) error { // session is torn down so the upstream rolls back rather than committing at Sync. func (s *session) denyBatch(sl slot) error { s.denied = true + s.pendingExecs = nil s.resetPending() if s.forwarded { @@ -592,6 +800,17 @@ func (s *session) readyForQuery(n uint32) error { return fmt.Errorf("read ReadyForQuery: %w", err) } + // finishing records are Done after the client's ReadyForQuery is flushed and + // the lock is released, so a full ledger queue never stalls the response path. + var finishing []wire.Result + defer func() { + for _, r := range finishing { + if r != nil { + r.Done() + } + } + }() + s.mu.Lock() defer s.mu.Unlock() @@ -605,6 +824,15 @@ func (s *session) readyForQuery(n uint32) error { return fmt.Errorf("%w: ReadyForQuery without a pending request", errMalformed) } + if r := s.queue[0].result; r != nil { + finishing = append(finishing, r) // the simple query's record + } + if c := s.queue[0].execCount; c > 0 && c <= len(s.execQueue) { + finishing = append(finishing, s.execQueue[:c:c]...) // this batch's Execute records + s.execQueue = s.execQueue[c:] + } + s.execIdx = 0 + s.curCap = nil s.queue = s.queue[1:] s.tx = status[0] if err := writeMessage(s.cw, typeReadyForQuery, status[:]); err != nil { @@ -655,6 +883,185 @@ func (s *session) forwardToClient(typ byte, n uint32) error { return writeMessage(s.cw, typ, body) } +// begin starts a ledger record for an allowed statement, returning the result +// to attach to its slot, or nil when nothing records. +func (s *session) begin(sql string, decision wire.Decision) wire.Result { + if s.recorder == nil { + return nil + } + + return s.recorder.Begin(sql, decision) +} + +// recordDenied writes a ledger record for a statement the proxy refused, which +// has no result to observe. +func (s *session) recordDenied(sql string) { + if r := s.begin(sql, wire.Denied); r != nil { + r.Done() + } +} + +// describeResult forwards a RowDescription and tells the current result which +// columns to capture. +func (s *session) describeResult(n uint32) error { + s.curCap = nil + if !s.recording() { + return s.forwardToClient(typeRowDescription, n) // nothing records; stream it + } + + body, err := readBody(s.ur, n) + if err != nil { + return err + } + + res, err := s.forwardResult(typeRowDescription, body) + if err != nil { + return err + } + if res != nil { + s.curCap = res.Columns(columnNames(body)) + } + + return nil +} + +// rowResult forwards a DataRow, capturing the requested columns when the result +// is being recorded. +func (s *session) rowResult(n uint32) error { + if len(s.curCap) == 0 { + return s.forwardToClient(typeDataRow, n) // stream without materializing + } + + body, err := readBody(s.ur, n) + if err != nil { + return err + } + res, err := s.forwardResult(typeDataRow, body) + if err != nil { + return err + } + if res != nil { + res.Row(rowValues(body, s.curCap)) + } + + return nil +} + +// completeResult forwards a CommandComplete, adds its row count to the result, +// and, in the extended protocol, advances to the next Execute's record. +func (s *session) completeResult(n uint32) error { + if !s.recording() { + if err := s.forwardToClient(typeCommandComplete, n); err != nil { + return err + } + s.advanceExec() + + return nil + } + + body, err := readBody(s.ur, n) + if err != nil { + return err + } + + res, err := s.forwardResult(typeCommandComplete, body) + if err != nil { + return err + } + if res != nil { + res.Complete(commandTag(body)) + } + s.curCap = nil + s.advanceExec() + + return nil +} + +// endResult forwards a message that ends a result set without a row count +// (EmptyQueryResponse, PortalSuspended) and advances the extended-protocol index. +func (s *session) endResult(typ byte, n uint32) error { + if err := s.forwardToClient(typ, n); err != nil { + return err + } + s.curCap = nil + s.advanceExec() + + return nil +} + +// advanceExec moves to the next queued Execute record, once the current result +// set ends. It has no effect for a simple query, whose one record spans every +// result set and is finalized at ReadyForQuery. +func (s *session) advanceExec() { + s.mu.Lock() + defer s.mu.Unlock() + + if len(s.queue) > 0 && s.queue[0].forwarded && !s.queue[0].simple { + s.execIdx++ + } +} + +// forwardResult writes one already-read message to the client and returns the +// record currently being filled: the simple query's own record, or the extended +// protocol's current Execute record. +func (s *session) forwardResult(typ byte, body []byte) (wire.Result, error) { + s.mu.Lock() + defer s.mu.Unlock() + + res := s.targetLocked() + if err := writeMessage(s.cw, typ, body); err != nil { + return nil, err + } + + return res, nil +} + +// targetLocked returns the record the current upstream result feeds. Callers hold mu. +func (s *session) targetLocked() wire.Result { + if len(s.queue) == 0 || !s.queue[0].forwarded { + return nil + } + if s.queue[0].simple { + return s.queue[0].result + } + if s.execIdx < s.queue[0].execCount && s.execIdx < len(s.execQueue) { + return s.execQueue[s.execIdx] + } + + return nil +} + +// recording reports whether the current upstream result has a record to fill, +// so results are only materialized when the ledger needs them. +func (s *session) recording() bool { + s.mu.Lock() + defer s.mu.Unlock() + + return s.targetLocked() != nil +} + +// finishInFlight finalizes the records of forwarded requests still queued when +// the session ends, so their records are written even without a ReadyForQuery. +func (s *session) finishInFlight() { + s.mu.Lock() + var pending []wire.Result + for _, sl := range s.queue { + if sl.forwarded && sl.result != nil { + pending = append(pending, sl.result) + } + } + pending = append(pending, s.execQueue...) + s.queue = nil + s.execQueue = nil + s.mu.Unlock() + + for _, r := range pending { + if r != nil { + r.Done() + } + } +} + func (s *session) flushClient() error { s.mu.Lock() defer s.mu.Unlock() @@ -725,20 +1132,44 @@ func readySlot(sl slot) slot { return sl } -// parseSQL returns the query text from a Parse message body, which is the -// statement name followed by the SQL, both null-terminated. -func parseSQL(body []byte) (string, error) { - _, rest, err := cstring(body) +// parseNameSQL returns the prepared statement name and query text from a Parse +// message body: the name and the SQL, both null-terminated. +func parseNameSQL(body []byte) (name, sql string, err error) { + name, rest, err := cstring(body) if err != nil { - return "", fmt.Errorf("parse message: %w", err) + return "", "", fmt.Errorf("parse message: %w", err) } - sql, _, err := cstring(rest) + sql, _, err = cstring(rest) + if err != nil { + return "", "", fmt.Errorf("parse message: %w", err) + } + + return name, sql, nil +} + +// parseBind returns the portal and prepared statement names from a Bind body. +func parseBind(body []byte) (portal, stmt string, ok bool) { + portal, rest, err := cstring(body) + if err != nil { + return "", "", false + } + stmt, _, err = cstring(rest) + if err != nil { + return "", "", false + } + + return portal, stmt, true +} + +// parsePortal returns the portal name from an Execute body. +func parsePortal(body []byte) (string, bool) { + portal, _, err := cstring(body) if err != nil { - return "", fmt.Errorf("parse message: %w", err) + return "", false } - return sql, nil + return portal, true } func parseStartup(body []byte) (wire.Startup, error) { diff --git a/internal/pg/pg_test.go b/internal/pg/pg_test.go index 4418213..ebe3510 100644 --- a/internal/pg/pg_test.go +++ b/internal/pg/pg_test.go @@ -708,6 +708,529 @@ func TestDenyAfterEarlyForwardTearsDownTheSession(t *testing.T) { <-backend } +type capture struct { + sql string + decision wire.Decision + rows int + done bool +} + +type capturingRecorder struct{ records *[]*capture } + +func (r capturingRecorder) Begin(sql string, d wire.Decision) wire.Result { + c := &capture{sql: sql, decision: d} + *r.records = append(*r.records, c) + + return c +} + +func (capture) Columns([]string) []int { return nil } +func (capture) Row([][]byte) {} +func (c *capture) Complete(tag string) { c.rows += tagRows(tag) } +func (c *capture) Done() { c.done = true } + +func tagRows(tag string) int { + fields := splitSpace(tag) + if len(fields) == 0 { + return 0 + } + + n := 0 + for _, c := range fields[len(fields)-1] { + if c < '0' || c > '9' { + return 0 + } + n = n*10 + int(c-'0') + } + + return n +} + +func splitSpace(s string) []string { + var out []string + cur := "" + for _, c := range s { + if c == ' ' { + if cur != "" { + out = append(out, cur) + cur = "" + } + + continue + } + cur += string(c) + } + if cur != "" { + out = append(out, cur) + } + + return out +} + +func TestFrontendRecordsSimpleQueryOutcome(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + query := msg('Q', cstr("select id from t")) + reply := [][]byte{ + msg('T', be16(1), cstr("id"), be32(0), be16(0), be32(23), be16(4), be32(0xffffffff), be16(0)), + msg('D', be16(1), be32(1), []byte("7")), + msg('D', be16(1), be32(1), []byte("8")), + msg('C', cstr("SELECT 2")), + msg('Z', []byte("I")), + } + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, query); err != nil { + return err + } + + return write(s, reply...) + }) + + mustWrite(t, p.client, query) + for _, want := range reply { + expectMsg(t, p.client, want) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 1 { + t.Fatalf("records: got %d, want 1", len(records)) + } + rec := records[0] + if rec.sql != "select id from t" || rec.decision != wire.Allowed || rec.rows != 2 || !rec.done { + t.Errorf("record: got %+v, want allowed select with 2 rows, done", rec) + } +} + +func TestFrontendRecordsDenial(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + var records []*capture + done := p.frontendRec(denyContaining("delete"), capturingRecorder{records: &records}) + + mustWrite(t, p.client, msg('Q', cstr("delete from t"))) + if typ, _ := readMsgT(t, p.client); typ != 'E' { + t.Fatal("expected ErrorResponse") + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + + mustWrite(t, p.client, msg('X')) + expectMsg(t, p.upstream, msg('X')) + <-done + + if len(records) != 1 || records[0].decision != wire.Denied || !records[0].done { + t.Errorf("records: got %+v, want one denied+done record", records) + } +} + +func TestFrontendRecordsExtendedSingleStatement(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + batch := bytes.Join([][]byte{pMsg("select id from t"), bindMsg, execMsg, syncMsg}, nil) + reply := [][]byte{ + msg('1'), msg('2'), + msg('T', be16(1), cstr("id"), be32(0), be16(0), be32(23), be16(4), be32(0xffffffff), be16(0)), + msg('D', be16(1), be32(1), []byte("7")), + msg('C', cstr("SELECT 1")), + msg('Z', []byte("I")), + } + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, batch); err != nil { + return err + } + + return write(s, reply...) + }) + + mustWrite(t, p.client, batch) + for _, want := range reply { + expectMsg(t, p.client, want) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 1 { + t.Fatalf("records: got %d, want 1", len(records)) + } + if r := records[0]; r.sql != "select id from t" || r.decision != wire.Allowed || r.rows != 1 || !r.done { + t.Errorf("record: got %+v, want allowed select id from t with 1 row", r) + } +} + +func TestFrontendRecordsOnlyTheDenialOfARejectedBatch(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(denyContaining("delete"), capturingRecorder{records: &records}) + + // An allowed Parse precedes a denied one; the batch is rejected, so only the + // denial is recorded and the earlier statement, which never ran, is not. + batch := bytes.Join([][]byte{ + pMsg("select 1"), bindMsg, execMsg, + pMsg("delete from t"), bindMsg, execMsg, + syncMsg, + }, nil) + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, syncMsg); err != nil { + return err + } + + return write(s, msg('Z', []byte("I"))) + }) + + mustWrite(t, p.client, batch) + if typ, _ := readMsgT(t, p.client); typ != 'E' { + t.Fatal("expected ErrorResponse for the denied batch") + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 1 { + t.Fatalf("records: got %d (%+v), want 1", len(records), records) + } + if r := records[0]; r.sql != "delete from t" || r.decision != wire.Denied { + t.Errorf("record: got %+v, want the delete recorded as denied", r) + } +} + +func namedParse(name, sql string) []byte { + return msg('P', cstr(name), cstr(sql), be16(0)) +} + +func namedBind(portal, stmt string) []byte { + return msg('B', cstr(portal), cstr(stmt), be16(0), be16(0), be16(0)) +} + +// TestFrontendRecordsReExecutionOfAPreparedStatement is the driver case: the +// statement is prepared once, then re-executed in a later batch with no Parse. +// Both executions must be recorded. +func TestFrontendRecordsReExecutionOfAPreparedStatement(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + prepare := bytes.Join([][]byte{namedParse("s1", "select id from t"), namedBind("", "s1"), execMsg, syncMsg}, nil) + reuse := bytes.Join([][]byte{namedBind("", "s1"), execMsg, syncMsg}, nil) + result := func(withParse bool) [][]byte { + out := [][]byte{} + if withParse { + out = append(out, msg('1')) + } + out = append(out, + msg('2'), + msg('T', be16(1), cstr("id"), be32(0), be16(0), be32(23), be16(4), be32(0xffffffff), be16(0)), + msg('D', be16(1), be32(1), []byte("7")), + msg('C', cstr("SELECT 1")), + msg('Z', []byte("I")), + ) + + return out + } + + server := p.serve(func(sc net.Conn) error { + if err := expectBytes(sc, prepare); err != nil { + return err + } + if err := write(sc, result(true)...); err != nil { + return err + } + if err := expectBytes(sc, reuse); err != nil { + return err + } + + return write(sc, result(false)...) + }) + + mustWrite(t, p.client, prepare) + for _, w := range result(true) { + expectMsg(t, p.client, w) + } + mustWrite(t, p.client, reuse) + for _, w := range result(false) { + expectMsg(t, p.client, w) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 2 { + t.Fatalf("records: got %d, want 2 (prepare+reuse)", len(records)) + } + for i, r := range records { + if r.sql != "select id from t" || r.decision != wire.Allowed || r.rows != 1 || !r.done { + t.Errorf("record %d: got %+v, want allowed select id from t, 1 row", i, r) + } + } +} + +// TestFrontendDoesNotRecordAPrepareOnlyBatch checks that preparing a statement +// without executing it is not recorded as an execution. +func TestFrontendDoesNotRecordAPrepareOnlyBatch(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + batch := bytes.Join([][]byte{namedParse("s1", "select id from t"), msg('D', []byte("S"), cstr("s1")), syncMsg}, nil) + reply := [][]byte{ + msg('1'), + msg('t', be16(0)), + msg('T', be16(1), cstr("id"), be32(0), be16(0), be32(23), be16(4), be32(0xffffffff), be16(0)), + msg('Z', []byte("I")), + } + server := p.serve(func(sc net.Conn) error { + if err := expectBytes(sc, batch); err != nil { + return err + } + + return write(sc, reply...) + }) + + mustWrite(t, p.client, batch) + for _, w := range reply { + expectMsg(t, p.client, w) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 0 { + t.Errorf("records: got %d (%+v), want 0 for a prepare-only batch", len(records), records) + } +} + +func execPortal(portal string) []byte { + return msg('E', cstr(portal), be32(0)) +} + +// TestExecuteOfAnUnknownPortalIsNotMisattributed guards against resolving an +// unknown portal to the unnamed prepared statement. +func TestExecuteOfAnUnknownPortalIsNotMisattributed(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + prepare := bytes.Join([][]byte{namedParse("", "select id from t"), namedBind("", ""), execMsg, syncMsg}, nil) + ghost := bytes.Join([][]byte{execPortal("ghost"), syncMsg}, nil) + commandZ := [][]byte{msg('C', cstr("SELECT 1")), msg('Z', []byte("I"))} + + server := p.serve(func(sc net.Conn) error { + if err := expectBytes(sc, prepare); err != nil { + return err + } + if err := write(sc, commandZ...); err != nil { + return err + } + if err := expectBytes(sc, ghost); err != nil { + return err + } + + return write(sc, msg('C', cstr("SELECT 0")), msg('Z', []byte("I"))) + }) + + mustWrite(t, p.client, prepare) + for _, w := range commandZ { + expectMsg(t, p.client, w) + } + mustWrite(t, p.client, ghost) + expectMsg(t, p.client, msg('C', cstr("SELECT 0"))) + expectMsg(t, p.client, msg('Z', []byte("I"))) + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 1 { + t.Fatalf("records: got %d (%+v), want 1 (the unknown-portal Execute must not record)", len(records), records) + } + if records[0].sql != "select id from t" { + t.Errorf("record: got sql %q, want the prepared statement", records[0].sql) + } +} + +// TestUnrecordedExecuteDoesNotStealTheNextResult checks that a placeholder keeps +// an unrecorded Execute's result set from being credited to a later record. +func TestUnrecordedExecuteDoesNotStealTheNextResult(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + // One batch: an Execute of an unbound portal (unresolved), then an Execute of + // a freshly prepared statement. The server returns a result set for each. + batch := bytes.Join([][]byte{ + namedBind("p1", "nope"), execPortal("p1"), + namedParse("s2", "select v from t"), namedBind("p2", "s2"), execPortal("p2"), + syncMsg, + }, nil) + reply := [][]byte{ + msg('C', cstr("SELECT 5")), // for the unresolved Execute + msg('C', cstr("SELECT 3")), // for the recorded Execute + msg('Z', []byte("I")), + } + server := p.serve(func(sc net.Conn) error { + if err := expectBytes(sc, batch); err != nil { + return err + } + + return write(sc, reply...) + }) + + mustWrite(t, p.client, batch) + for _, w := range reply { + expectMsg(t, p.client, w) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 1 { + t.Fatalf("records: got %d, want 1", len(records)) + } + if records[0].sql != "select v from t" || records[0].rows != 3 { + t.Errorf("record: got %+v, want select v from t with 3 rows (not the unresolved Execute's 5)", records[0]) + } +} + +// TestCloseForgetsThePreparedStatement checks that after a Close, re-executing +// the statement's old name is no longer attributed to it. +func TestCloseForgetsThePreparedStatement(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var records []*capture + frontend := p.frontendRec(allow, capturingRecorder{records: &records}) + + // Prepare s1, then close it; a later Bind/Execute of s1 must not record. + prepare := bytes.Join([][]byte{namedParse("s1", "select id from t"), namedBind("", "s1"), execMsg, syncMsg}, nil) + closeIt := bytes.Join([][]byte{msg('C', []byte("S"), cstr("s1")), syncMsg}, nil) + reuse := bytes.Join([][]byte{namedBind("", "s1"), execMsg, syncMsg}, nil) + commandZ := [][]byte{msg('C', cstr("SELECT 1")), msg('Z', []byte("I"))} + + server := p.serve(func(sc net.Conn) error { + if err := expectBytes(sc, prepare); err != nil { + return err + } + if err := write(sc, commandZ...); err != nil { + return err + } + if err := expectBytes(sc, closeIt); err != nil { + return err + } + if err := write(sc, msg('3'), msg('Z', []byte("I"))); err != nil { // CloseComplete, ReadyForQuery + return err + } + if err := expectBytes(sc, reuse); err != nil { + return err + } + + return write(sc, commandZ...) + }) + + mustWrite(t, p.client, prepare) + for _, w := range commandZ { + expectMsg(t, p.client, w) + } + mustWrite(t, p.client, closeIt) + expectMsg(t, p.client, msg('3')) + expectMsg(t, p.client, msg('Z', []byte("I"))) + mustWrite(t, p.client, reuse) + for _, w := range commandZ { + expectMsg(t, p.client, w) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend + + if len(records) != 1 { + t.Errorf("records: got %d, want 1 (reuse after Close must not record)", len(records)) + } +} + // pipes connects a session to a fake client and a fake upstream over net.Pipe. type pipes struct { sess wire.Session @@ -746,8 +1269,12 @@ func (p pipes) handshake() <-chan handshakeResult { } func (p pipes) frontend(h wire.Handler) <-chan error { + return p.frontendRec(h, nil) +} + +func (p pipes) frontendRec(h wire.Handler, rec wire.Recorder) <-chan error { done := make(chan error, 1) - go func() { done <- p.sess.Frontend(h) }() + go func() { done <- p.sess.Frontend(h, rec) }() return done } diff --git a/internal/policy/policy.go b/internal/policy/policy.go index d8da76b..9ee72c1 100644 --- a/internal/policy/policy.go +++ b/internal/policy/policy.go @@ -95,28 +95,37 @@ func parseFail(s string) (bool, error) { func (p Policy) Resolve(startup wire.Startup) wire.Enforcement { role, ok := p.Roles[startup.User] if !ok { - return p.unlisted(startup.User) + return p.unlisted(startup) } + principal := wire.Principal{ + Agent: role.Agent, + Purpose: role.Purpose, + User: startup.User, + Database: startup.Database, + Application: startup.Application, + } if !role.ReadOnly { - return wire.Enforcement{Handler: allow()} + return wire.Enforcement{Principal: principal, Handler: allow()} } return wire.Enforcement{ - Prime: []string{readOnlyPrime}, - Handler: wire.HandlerFunc(readOnly), + Principal: principal, + Prime: []string{readOnlyPrime}, + Handler: wire.HandlerFunc(readOnly), } } -func (p Policy) unlisted(user string) wire.Enforcement { +func (p Policy) unlisted(startup wire.Startup) wire.Enforcement { + principal := wire.Principal{User: startup.User, Database: startup.Database, Application: startup.Application} if !p.FailClosed { - return wire.Enforcement{Handler: allow()} + return wire.Enforcement{Principal: principal, Handler: allow()} } - return wire.Enforcement{Handler: wire.HandlerFunc(func(wire.Statement) wire.Verdict { + return wire.Enforcement{Principal: principal, Handler: wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{ Deny: true, - Message: fmt.Sprintf("no policy for role %q", user), + Message: fmt.Sprintf("no policy for role %q", startup.User), Hint: "add the role to the policy, or connect as a configured role", } })} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index e3b2a1b..63ceaf7 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -130,7 +130,7 @@ func (s Server) handle(ctx context.Context, client net.Conn) { var wg sync.WaitGroup var toUpstream, toClient error wg.Go(func() { - toUpstream = sess.Frontend(enforcement.Handler) + toUpstream = sess.Frontend(enforcement.Handler, enforcement.Recorder) closeWrite(upstream) }) wg.Go(func() { diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index 18bf5fb..57c1723 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -149,7 +149,7 @@ func (rawSession) Prime(string) error { return nil } -func (s rawSession) Frontend(wire.Handler) error { +func (s rawSession) Frontend(wire.Handler, wire.Recorder) error { if _, err := io.Copy(s.upstream, s.client); err != nil { return fmt.Errorf("frontend: %w", err) } diff --git a/internal/sqlscan/fingerprint.go b/internal/sqlscan/fingerprint.go new file mode 100644 index 0000000..3d9cbc4 --- /dev/null +++ b/internal/sqlscan/fingerprint.go @@ -0,0 +1,105 @@ +package sqlscan + +import "strings" + +// Fingerprint returns sql with its literal values removed, so statements of the +// same shape share a fingerprint and no literal (which may be personal data) is +// kept. String, escape-string, dollar-quoted, and numeric literals become '?'; +// comments are dropped; whitespace is collapsed to single spaces. Identifiers +// and keywords are preserved (uppercased for words, lowercased for quoted +// identifiers). +func Fingerprint(sql string) string { + var b strings.Builder + space := false + write := func(tok string) { + if space && b.Len() > 0 { + b.WriteByte(' ') + } + space = false + b.WriteString(tok) + } + + for i, n := 0, len(sql); i < n; { + c := sql[i] + switch { + case isSpace(c): + space = true + i++ + case isComment(sql, i): + i = skipComment(sql, i) + space = true + default: + var tok string + i, tok = fingerprintToken(sql, i) + write(tok) + } + } + + return b.String() +} + +func isComment(s string, i int) bool { + if i+1 >= len(s) { + return false + } + + return (s[i] == '-' && s[i+1] == '-') || (s[i] == '/' && s[i+1] == '*') +} + +func skipComment(s string, i int) int { + if s[i] == '-' { + return skipLineComment(s, i+2) + } + + return skipBlockComment(s, i+2) +} + +// fingerprintToken consumes one non-space, non-comment token and returns the +// next index and the text to emit for it. +func fingerprintToken(sql string, i int) (int, string) { + n := len(sql) + c := sql[i] + switch { + case (c == 'e' || c == 'E') && i+1 < n && sql[i+1] == '\'': + return skipEString(sql, i+2), "?" + case c == '\'': + return skipString(sql, i+1), "?" + case c == '"': + j := skipString(sql, i+1) + + return j, `"` + quotedText(sql, i, j) + `"` + case c == '$': + if end, ok := skipDollarQuote(sql, i); ok { + return end, "?" + } + + return skipWord(sql, i+1), "?" + case isDigit(c): + return skipNumber(sql, i), "?" + case isWordStart(c): + j := skipWord(sql, i+1) + + return j, strings.ToUpper(sql[i:j]) + default: + return i + 1, string(c) + } +} + +// skipNumber consumes a numeric literal, including hex/binary/octal prefixes +// (0x1F), digit-group underscores (1_000), decimals, and scientific notation +// (1.5e-9), so no fragment survives as an identifier. +func skipNumber(s string, i int) int { + for i < len(s) { + c := s[i] + switch { + case isWordPart(c) || c == '.': + i++ + case (c == '+' || c == '-') && i > 0 && (s[i-1] == 'e' || s[i-1] == 'E'): + i++ + default: + return i + } + } + + return i +} diff --git a/internal/sqlscan/scan.go b/internal/sqlscan/scan.go index d505ff8..a1d1d26 100644 --- a/internal/sqlscan/scan.go +++ b/internal/sqlscan/scan.go @@ -36,6 +36,8 @@ func tokenize(sql string) [][]token { i = skipLineComment(sql, i+2) case c == '/' && i+1 < n && sql[i+1] == '*': i = skipBlockComment(sql, i+2) + case (c == 'e' || c == 'E') && i+1 < n && sql[i+1] == '\'': + i = skipEString(sql, i+2) case c == '\'': i = skipString(sql, i+1) case c == '"': @@ -129,6 +131,30 @@ func skipString(s string, i int) int { // 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. +// skipEString skips a PostgreSQL escape string E'...', whose opening quote at +// i-1 has been consumed. A backslash escapes the next byte, and ” still stands +// for a quote. +func skipEString(s string, i int) int { + for i < len(s) { + switch s[i] { + case '\\': + i += 2 + case '\'': + if i+1 < len(s) && s[i+1] == '\'' { + i += 2 + + continue + } + + return i + 1 + default: + i++ + } + } + + return i +} + func skipDollarQuote(s string, i int) (int, bool) { j := i + 1 for j < len(s) && isTagPart(s[j]) { diff --git a/internal/sqlscan/sqlscan_test.go b/internal/sqlscan/sqlscan_test.go index 0d09461..f565ca6 100644 --- a/internal/sqlscan/sqlscan_test.go +++ b/internal/sqlscan/sqlscan_test.go @@ -1,6 +1,7 @@ package sqlscan_test import ( + "strings" "testing" "github.com/mickamy/rollcall/internal/sqlscan" @@ -162,3 +163,37 @@ func disables(sql string) bool { return false } + +func TestFingerprint(t *testing.T) { + t.Parallel() + + tests := map[string]string{ + "select * from t where id = 42": "SELECT * FROM T WHERE ID = ?", + "select id from t where email = 'a@b.com'": "SELECT ID FROM T WHERE EMAIL = ?", + "insert into t values (1, 'x'), (2, 'y')": "INSERT INTO T VALUES (?, ?), (?, ?)", + "select col from t -- note\n": "SELECT COL FROM T", + `select "MixedCase" from t where x = $1`: `SELECT "mixedcase" FROM T WHERE X = ?`, + "select $$body$$ as b": "SELECT ? AS B", + "select 3.14, 1e9 from t": "SELECT ?, ? FROM T", + } + for sql, want := range tests { + if got := sqlscan.Fingerprint(sql); got != want { + t.Errorf("Fingerprint(%q) = %q, want %q", sql, got, want) + } + } +} + +func TestEscapeStringsDoNotLeak(t *testing.T) { + t.Parallel() + + // In an E'' string a backslash escapes the quote, so the words inside must + // not surface as identifiers in the fingerprint or flip classification. + sql := `select id from t where name = E'O\'Brien said delete from t'` + fp := sqlscan.Fingerprint(sql) + if strings.Contains(fp, "BRIEN") || strings.Contains(fp, "DELETE") { + t.Errorf("Fingerprint leaked E-string content: %q", fp) + } + if sqlscan.Mutating(sql) { + t.Errorf("Mutating(%q) = true; the DELETE is inside an E-string", sql) + } +} diff --git a/internal/wire/wire.go b/internal/wire/wire.go index 3c29cf2..3c6af89 100644 --- a/internal/wire/wire.go +++ b/internal/wire/wire.go @@ -45,11 +45,50 @@ 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. +// Principal is who a session's statements are attributed to. +type Principal struct { + Agent string + Purpose string + User string + Database string + Application string +} + +// Decision is what the proxy did with a statement. +type Decision string + +const ( + Allowed Decision = "allowed" + Denied Decision = "denied" +) + +// Recorder records the statements of one session for the access ledger. A nil +// Recorder records nothing. +type Recorder interface { + // Begin starts a record for a statement and its decision, returning a + // Result to receive the statement's result, or nil to record no result. + Begin(sql string, decision Decision) Result +} + +// Result receives one statement's result from the session's backend goroutine, +// in order: Columns for each result set (returning which columns to capture), +// Row for each row (with the captured columns only), Complete for each result +// set's command tag, and Done once when the statement finishes. +type Result interface { + Columns(names []string) (capture []int) + Row(captured [][]byte) + Complete(tag string) + Done() +} + +// Enforcement is how a session is guarded: who it is attributed to, statements +// to run on the upstream before the relay starts, the handler that judges each +// client statement, and an optional recorder for the ledger. type Enforcement struct { - Prime []string - Handler Handler + Principal Principal + Prime []string + Handler Handler + Recorder Recorder } // Guard resolves the enforcement for a session from the client's identity. @@ -64,9 +103,12 @@ 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{} })} +// AllowAll is a Guard whose sessions permit every statement and record nothing. +var AllowAll Guard = GuardFunc(func(s Startup) Enforcement { + return Enforcement{ + Principal: Principal{User: s.User, Database: s.Database, Application: s.Application}, + Handler: HandlerFunc(func(Statement) Verdict { return Verdict{} }), + } }) type Dialect interface { @@ -82,6 +124,6 @@ type Session interface { // 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 + Frontend(h Handler, rec Recorder) error Backend() error }