From 13603fae18b32ef37a9dd3d3a6cb5bb5343fa15c Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 09:25:43 +0900 Subject: [PATCH 01/14] feat: fingerprint SQL with its literals removed --- internal/sqlscan/fingerprint.go | 69 ++++++++++++++++++++++++++++++++ internal/sqlscan/sqlscan_test.go | 19 +++++++++ 2 files changed, 88 insertions(+) create mode 100644 internal/sqlscan/fingerprint.go diff --git a/internal/sqlscan/fingerprint.go b/internal/sqlscan/fingerprint.go new file mode 100644 index 0000000..f8ce1fc --- /dev/null +++ b/internal/sqlscan/fingerprint.go @@ -0,0 +1,69 @@ +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, 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 + i, n := 0, len(sql) + space := false + + emit := func(s string) { + if space && b.Len() > 0 { + b.WriteByte(' ') + } + space = false + b.WriteString(s) + } + + for i < n { + c := sql[i] + switch { + case isSpace(c): + space = true + i++ + case c == '-' && i+1 < n && sql[i+1] == '-': + i = skipLineComment(sql, i+2) + space = true + case c == '/' && i+1 < n && sql[i+1] == '*': + i = skipBlockComment(sql, i+2) + space = true + case c == '\'': + i = skipString(sql, i+1) + emit("?") + case c == '"': + j := skipString(sql, i+1) + emit(`"` + quotedText(sql, i, j) + `"`) + i = j + case c == '$': + if end, ok := skipDollarQuote(sql, i); ok { + i = end + emit("?") + } else { + i = skipWord(sql, i+1) + emit("?") // parameter such as $1 + } + case c >= '0' && c <= '9': + for i < n && (isDigit(sql[i]) || sql[i] == '.' || sql[i] == 'e' || sql[i] == 'E') { + i++ + } + emit("?") + case isWordStart(c): + j := skipWord(sql, i+1) + emit(strings.ToUpper(sql[i:j])) + i = j + case c == '(' || c == ')' || c == ',' || c == ';' || c == '*': + emit(string(c)) + i++ + default: + emit(string(c)) + i++ + } + } + + return b.String() +} diff --git a/internal/sqlscan/sqlscan_test.go b/internal/sqlscan/sqlscan_test.go index 0d09461..bc8c066 100644 --- a/internal/sqlscan/sqlscan_test.go +++ b/internal/sqlscan/sqlscan_test.go @@ -162,3 +162,22 @@ 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) + } + } +} From ebb1e9d9ced1090010d36972134ac01d8f42d8c4 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 09:25:43 +0900 Subject: [PATCH 02/14] feat: record each statement to a hash-chained access ledger --- internal/ledger/clock.go | 7 ++ internal/ledger/export_test.go | 9 +++ internal/ledger/ledger_test.go | 104 +++++++++++++++++++++++++ internal/ledger/record.go | 23 ++++++ internal/ledger/recorder.go | 98 ++++++++++++++++++++++++ internal/ledger/sink.go | 58 ++++++++++++++ internal/ledger/tag.go | 23 ++++++ internal/pg/message.go | 100 ++++++++++++++++++++---- internal/pg/pg.go | 122 ++++++++++++++++++++++++++++- internal/pg/pg_test.go | 136 ++++++++++++++++++++++++++++++++- internal/proxy/proxy.go | 2 +- internal/proxy/proxy_test.go | 2 +- internal/wire/wire.go | 58 ++++++++++++-- 13 files changed, 712 insertions(+), 30 deletions(-) create mode 100644 internal/ledger/clock.go create mode 100644 internal/ledger/export_test.go create mode 100644 internal/ledger/ledger_test.go create mode 100644 internal/ledger/record.go create mode 100644 internal/ledger/recorder.go create mode 100644 internal/ledger/sink.go create mode 100644 internal/ledger/tag.go 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/export_test.go b/internal/ledger/export_test.go new file mode 100644 index 0000000..ec16ed1 --- /dev/null +++ b/internal/ledger/export_test.go @@ -0,0 +1,9 @@ +package ledger + +// SetNow overrides the record clock for tests and returns a restore function. +func SetNow(t func() string) func() { + prev := now + now = t + + return func() { now = prev } +} diff --git a/internal/ledger/ledger_test.go b/internal/ledger/ledger_test.go new file mode 100644 index 0000000..fce795b --- /dev/null +++ b/internal/ledger/ledger_test.go @@ -0,0 +1,104 @@ +package ledger_test + +import ( + "bufio" + "bytes" + "encoding/json" + "strings" + "testing" + + "github.com/mickamy/rollcall/internal/ledger" + "github.com/mickamy/rollcall/internal/wire" +) + +func TestGuardRecordsStatements(t *testing.T) { + t.Parallel() + + defer ledger.SetNow(func() string { return "2026-08-27T00:00:00Z" })() + + 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{} }), + } + }) + g := ledger.Guard{Inner: inner, Sink: ledger.NewSink(&buf)} + + 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() + + 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) + for range 3 { + if err := s.Write(ledger.Record{User: "u", Kind: "SELECT"}); err != nil { + t.Fatalf("Write: %v", err) + } + } + + 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 +} diff --git a/internal/ledger/record.go b/internal/ledger/record.go new file mode 100644 index 0000000..0281ca2 --- /dev/null +++ b/internal/ledger/record.go @@ -0,0 +1,23 @@ +// 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"` + // 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..9ea95be --- /dev/null +++ b/internal/ledger/recorder.go @@ -0,0 +1,98 @@ +package ledger + +import ( + "log/slog" + + "github.com/mickamy/rollcall/internal/sqlscan" + "github.com/mickamy/rollcall/internal/wire" +) + +// Guard wraps another guard, recording each session's statements to Sink. +type Guard struct { + Inner wire.Guard + Sink *Sink + Logger *slog.Logger +} + +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, logger: g.logger()} + + return enf +} + +func (g Guard) logger() *slog.Logger { + if g.Logger == nil { + return slog.New(slog.DiscardHandler) + } + + return g.Logger +} + +// now is overridable in tests; the ledger stamps records with the wall clock. +var now = defaultNow + +type recorder struct { + principal wire.Principal + sink *Sink + logger *slog.Logger +} + +var _ wire.Recorder = recorder{} + +func (r recorder) Begin(sql string, decision wire.Decision) wire.Result { + kinds := sqlscan.Classify(sql) + kind := "empty" + if len(kinds) > 0 { + kind = kinds[0].String() + } + + 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: sqlscan.Fingerprint(sql), + Decision: decision, + } + + return &result{rec: rec, sink: r.sink, logger: r.logger} +} + +// result accumulates a statement's outcome and writes the record when the +// statement finishes. +type result struct { + rec Record + sink *Sink + logger *slog.Logger + 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 + + if err := r.sink.Write(r.rec); err != nil { + r.logger.Error("ledger write", "error", err) + } +} diff --git a/internal/ledger/sink.go b/internal/ledger/sink.go new file mode 100644 index 0000000..f5ca8b2 --- /dev/null +++ b/internal/ledger/sink.go @@ -0,0 +1,58 @@ +package ledger + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "sync" +) + +// Sink appends records to an output as JSON lines, chaining each record's hash +// to the previous so the log is tamper-evident. It is safe for concurrent use. +type Sink struct { + mu sync.Mutex + w io.Writer + prev string +} + +// NewSink writes records to w. +func NewSink(w io.Writer) *Sink { + return &Sink{w: w} +} + +// Write chains and appends one record. It fills in PrevHash and Hash. +func (s *Sink) Write(rec Record) error { + s.mu.Lock() + defer s.mu.Unlock() + + rec.PrevHash = s.prev + rec.Hash = hashRecord(rec) + + line, err := json.Marshal(rec) + if err != nil { + return fmt.Errorf("marshal record: %w", err) + } + if _, err := s.w.Write(append(line, '\n')); err != nil { + return fmt.Errorf("write record: %w", err) + } + + s.prev = rec.Hash + + return nil +} + +// hashRecord hashes the record with its Hash field cleared, over PrevHash, so +// the chain covers order and content. +func hashRecord(rec Record) string { + rec.Hash = "" + data, err := json.Marshal(rec) + if err != nil { + return "" + } + + sum := sha256.Sum256(data) + + return hex.EncodeToString(sum[:]) +} 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..6f015f4 100644 --- a/internal/pg/message.go +++ b/internal/pg/message.go @@ -28,21 +28,24 @@ 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' ) const ( @@ -194,6 +197,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..459feca 100644 --- a/internal/pg/pg.go +++ b/internal/pg/pg.go @@ -72,6 +72,7 @@ func (d Dialect) NewSession(client, upstream net.Conn) wire.Session { type slot struct { forwarded bool simple bool + result wire.Result code string message string hint string @@ -84,13 +85,16 @@ type session struct { // Frontend goroutine only. cr *bufio.Reader uw *bufio.Writer + recorder wire.Recorder pending bytes.Buffer 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 + curResult wire.Result // the result currently streaming from the upstream + curCap []int // columns the recorder asked to capture for it // mu serializes writes to the client, which Frontend (denials) and Backend // (forwarded messages) both perform, and guards tx and queue alongside them. @@ -169,7 +173,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 { @@ -224,6 +229,12 @@ 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) default: err = s.forwardToClient(typ, n) } @@ -329,13 +340,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.record(sql, wire.Denied) + 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) } @@ -366,8 +379,11 @@ func (s *session) parse(h wire.Handler, n uint32) error { } if v := h.Statement(wire.Statement{SQL: sql}); v.Deny { + s.record(sql, wire.Denied) + return s.denyBatch(denial(v)) } + s.record(sql, wire.Allowed) return s.stage(typeParse, body) } @@ -605,6 +621,11 @@ 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 { + r.Done() + } + s.curResult = nil + s.curCap = nil s.queue = s.queue[1:] s.tx = status[0] if err := writeMessage(s.cw, typeReadyForQuery, status[:]); err != nil { @@ -655,6 +676,99 @@ 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) +} + +// record writes a ledger record for a statement with no result to observe, such +// as a denial or an extended-protocol Parse. +func (s *session) record(sql string, decision wire.Decision) { + if r := s.begin(sql, decision); r != nil { + r.Done() + } +} + +// describeResult forwards a RowDescription and tells the current result which +// columns to capture. +func (s *session) describeResult(n uint32) error { + body, err := readBody(s.ur, n) + if err != nil { + return err + } + + res, err := s.forwardResult(typeRowDescription, body) + if err != nil { + return err + } + s.curResult = res + s.curCap = nil + 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 s.curResult == nil || len(s.curCap) == 0 { + return s.forwardToClient(typeDataRow, n) // stream without materializing + } + + body, err := readBody(s.ur, n) + if err != nil { + return err + } + if _, err := s.forwardResult(typeDataRow, body); err != nil { + return err + } + s.curResult.Row(rowValues(body, s.curCap)) + + return nil +} + +// completeResult forwards a CommandComplete and adds its row count to the result. +func (s *session) completeResult(n uint32) error { + 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)) + } + + return nil +} + +// forwardResult writes one already-read message to the client and returns the +// result attached to the request currently being answered, if any. +func (s *session) forwardResult(typ byte, body []byte) (wire.Result, error) { + s.mu.Lock() + defer s.mu.Unlock() + + var res wire.Result + if len(s.queue) > 0 && s.queue[0].forwarded { + res = s.queue[0].result + } + if err := writeMessage(s.cw, typ, body); err != nil { + return nil, err + } + + return res, nil +} + func (s *session) flushClient() error { s.mu.Lock() defer s.mu.Unlock() diff --git a/internal/pg/pg_test.go b/internal/pg/pg_test.go index 4418213..bc430d2 100644 --- a/internal/pg/pg_test.go +++ b/internal/pg/pg_test.go @@ -708,6 +708,136 @@ 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) + } +} + // pipes connects a session to a fake client and a fake upstream over net.Pipe. type pipes struct { sess wire.Session @@ -746,8 +876,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/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/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 } From 2010fc419ab2ab08f7538af45863789385f8a71b Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 09:25:43 +0900 Subject: [PATCH 03/14] feat: attribute records to the policy principal and add the -ledger flag --- README.md | 4 +++- internal/cli/proxy.go | 17 ++++++++++++++++- internal/policy/policy.go | 25 +++++++++++++++++-------- 3 files changed, 36 insertions(+), 10 deletions(-) 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/proxy.go b/internal/cli/proxy.go index fcec168..bff1489 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,17 @@ func runProxy(ctx context.Context, args []string, std IO) int { guard = p } + logger := slog.New(slog.NewTextHandler(std.Err, nil)) + + if *ledgerPath != "" { + 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() }() + guard = ledger.Guard{Inner: guard, Sink: ledger.NewSink(f), Logger: logger} + } + var lc net.ListenConfig ln, err := lc.Listen(ctx, "tcp", *listen) if err != nil { @@ -64,7 +79,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) @@ -118,4 +132,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/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", } })} From 30dfc8541e87dc071de28347de8160db550bd230 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 16:48:03 +0900 Subject: [PATCH 04/14] fix: mask escape-string literals and simplify the fingerprint scanner --- internal/sqlscan/fingerprint.go | 109 +++++++++++++++++++------------ internal/sqlscan/scan.go | 26 ++++++++ internal/sqlscan/sqlscan_test.go | 16 +++++ 3 files changed, 109 insertions(+), 42 deletions(-) diff --git a/internal/sqlscan/fingerprint.go b/internal/sqlscan/fingerprint.go index f8ce1fc..66b4bb3 100644 --- a/internal/sqlscan/fingerprint.go +++ b/internal/sqlscan/fingerprint.go @@ -4,66 +4,91 @@ 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, 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). +// 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 - i, n := 0, len(sql) space := false - - emit := func(s string) { + write := func(tok string) { if space && b.Len() > 0 { b.WriteByte(' ') } space = false - b.WriteString(s) + b.WriteString(tok) } - for i < n { + for i, n := 0, len(sql); i < n; { c := sql[i] switch { case isSpace(c): space = true i++ - case c == '-' && i+1 < n && sql[i+1] == '-': - i = skipLineComment(sql, i+2) - space = true - case c == '/' && i+1 < n && sql[i+1] == '*': - i = skipBlockComment(sql, i+2) + case isComment(sql, i): + i = skipComment(sql, i) space = true - case c == '\'': - i = skipString(sql, i+1) - emit("?") - case c == '"': - j := skipString(sql, i+1) - emit(`"` + quotedText(sql, i, j) + `"`) - i = j - case c == '$': - if end, ok := skipDollarQuote(sql, i); ok { - i = end - emit("?") - } else { - i = skipWord(sql, i+1) - emit("?") // parameter such as $1 - } - case c >= '0' && c <= '9': - for i < n && (isDigit(sql[i]) || sql[i] == '.' || sql[i] == 'e' || sql[i] == 'E') { - i++ - } - emit("?") - case isWordStart(c): - j := skipWord(sql, i+1) - emit(strings.ToUpper(sql[i:j])) - i = j - case c == '(' || c == ')' || c == ',' || c == ';' || c == '*': - emit(string(c)) - i++ default: - emit(string(c)) - i++ + 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) + } +} + +func skipNumber(s string, i int) int { + for i < len(s) && (isDigit(s[i]) || s[i] == '.' || s[i] == 'e' || s[i] == 'E') { + 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 bc8c066..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" @@ -181,3 +182,18 @@ func TestFingerprint(t *testing.T) { } } } + +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) + } +} From a4f1f7b0b5f69ce8c51df40580ecd6045a0384d6 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 16:48:03 +0900 Subject: [PATCH 05/14] fix: write the ledger off the response path, key and resume its chain --- internal/cli/proxy.go | 20 +++++- internal/ledger/ledger_test.go | 37 ++++++++-- internal/ledger/recorder.go | 64 ++++++++++------- internal/ledger/resume.go | 45 ++++++++++++ internal/ledger/sink.go | 127 ++++++++++++++++++++++++++------- 5 files changed, 235 insertions(+), 58 deletions(-) create mode 100644 internal/ledger/resume.go diff --git a/internal/cli/proxy.go b/internal/cli/proxy.go index bff1489..3e07df9 100644 --- a/internal/cli/proxy.go +++ b/internal/cli/proxy.go @@ -64,12 +64,20 @@ func runProxy(ctx context.Context, args []string, std IO) int { 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() }() - guard = ledger.Guard{Inner: guard, Sink: ledger.NewSink(f), Logger: logger} + + 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 @@ -92,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) diff --git a/internal/ledger/ledger_test.go b/internal/ledger/ledger_test.go index fce795b..3aa45db 100644 --- a/internal/ledger/ledger_test.go +++ b/internal/ledger/ledger_test.go @@ -25,7 +25,8 @@ func TestGuardRecordsStatements(t *testing.T) { Handler: wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }), } }) - g := ledger.Guard{Inner: inner, Sink: ledger.NewSink(&buf)} + sink := ledger.NewSink(&buf, ledger.Options{}) + g := ledger.Guard{Inner: inner, Sink: sink} enf := g.Resolve(wire.Startup{User: "agent_ops", Database: "prod"}) @@ -38,6 +39,7 @@ func TestGuardRecordsStatements(t *testing.T) { // 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)) @@ -66,12 +68,11 @@ func TestSinkChainsHashes(t *testing.T) { t.Parallel() var buf bytes.Buffer - s := ledger.NewSink(&buf) + s := ledger.NewSink(&buf, ledger.Options{}) for range 3 { - if err := s.Write(ledger.Record{User: "u", Kind: "SELECT"}); err != nil { - t.Fatalf("Write: %v", err) - } + s.Write(ledger.Record{User: "u", Kind: "SELECT"}) } + s.Close() recs := decode(t, &buf) if recs[0].PrevHash != "" { @@ -102,3 +103,29 @@ func decode(t *testing.T, buf *bytes.Buffer) []ledger.Record { 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") + } +} diff --git a/internal/ledger/recorder.go b/internal/ledger/recorder.go index 9ea95be..5cb18c8 100644 --- a/internal/ledger/recorder.go +++ b/internal/ledger/recorder.go @@ -1,7 +1,7 @@ package ledger import ( - "log/slog" + "strings" "github.com/mickamy/rollcall/internal/sqlscan" "github.com/mickamy/rollcall/internal/wire" @@ -9,45 +9,60 @@ import ( // Guard wraps another guard, recording each session's statements to Sink. type Guard struct { - Inner wire.Guard - Sink *Sink - Logger *slog.Logger + 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, logger: g.logger()} + enf.Recorder = recorder{principal: enf.Principal, sink: g.Sink} return enf } -func (g Guard) logger() *slog.Logger { - if g.Logger == nil { - return slog.New(slog.DiscardHandler) +// 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() } - return g.Logger -} + 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() + } -// now is overridable in tests; the ledger stamps records with the wall clock. -var now = defaultNow + return strings.Join(names, ", ") +} type recorder struct { principal wire.Principal sink *Sink - logger *slog.Logger } var _ wire.Recorder = recorder{} func (r recorder) Begin(sql string, decision wire.Decision) wire.Result { - kinds := sqlscan.Classify(sql) - kind := "empty" - if len(kinds) > 0 { - kind = kinds[0].String() - } + kind := statementKind(sql) rec := Record{ Time: now(), @@ -61,16 +76,15 @@ func (r recorder) Begin(sql string, decision wire.Decision) wire.Result { Decision: decision, } - return &result{rec: rec, sink: r.sink, logger: r.logger} + 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 - logger *slog.Logger - done bool + rec Record + sink *Sink + done bool } var _ wire.Result = (*result)(nil) @@ -92,7 +106,5 @@ func (r *result) Done() { } r.done = true - if err := r.sink.Write(r.rec); err != nil { - r.logger.Error("ledger write", "error", err) - } + r.sink.Write(r.rec) } diff --git a/internal/ledger/resume.go b/internal/ledger/resume.go new file mode 100644 index 0000000..7128c3b --- /dev/null +++ b/internal/ledger/resume.go @@ -0,0 +1,45 @@ +package ledger + +import ( + "bufio" + "encoding/json" + "errors" + "fmt" + "io" + "os" +) + +// LastHash returns the Hash of the last record in the ledger file at path, so a +// new Sink can continue the chain across restarts. A missing 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() }() + + last := "" + sc := bufio.NewScanner(f) + sc.Buffer(make([]byte, 0, 64*1024), 16*1024*1024) + for sc.Scan() { + line := sc.Bytes() + if len(line) == 0 { + continue + } + var rec struct { + Hash string `json:"hash"` + } + if err := json.Unmarshal(line, &rec); err != nil { + return "", fmt.Errorf("parse ledger tail: %w", err) + } + last = rec.Hash + } + if err := sc.Err(); err != nil && !errors.Is(err, io.EOF) { + return "", fmt.Errorf("read ledger: %w", err) + } + + return last, nil +} diff --git a/internal/ledger/sink.go b/internal/ledger/sink.go index f5ca8b2..e6b6300 100644 --- a/internal/ledger/sink.go +++ b/internal/ledger/sink.go @@ -1,58 +1,133 @@ package ledger import ( + "crypto/hmac" "crypto/sha256" "encoding/hex" "encoding/json" - "fmt" - "io" + "hash" + "log/slog" "sync" ) -// Sink appends records to an output as JSON lines, chaining each record's hash -// to the previous so the log is tamper-evident. It is safe for concurrent use. +// Sink appends records as JSON lines, chaining each record's hash to the +// previous so the log is tamper-evident. Writes are queued to a single writer +// goroutine, so a slow disk never blocks a session's response path or serializes +// sessions against each other. type Sink struct { + ch chan Record + done chan struct{} + w writer + key []byte + logger *slog.Logger + mu sync.Mutex - w io.Writer prev string } -// NewSink writes records to w. -func NewSink(w io.Writer) *Sink { - return &Sink{w: w} +type writer interface { + Write(p []byte) (int, error) } -// Write chains and appends one record. It fills in PrevHash and Hash. -func (s *Sink) Write(rec Record) error { - s.mu.Lock() - defer s.mu.Unlock() +// 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 +} - rec.PrevHash = s.prev - rec.Hash = hashRecord(rec) +const queueDepth = 1024 - line, err := json.Marshal(rec) - if err != nil { - return fmt.Errorf("marshal record: %w", err) +// 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) } - if _, err := s.w.Write(append(line, '\n')); err != nil { - return fmt.Errorf("write record: %w", err) + + s := &Sink{ + ch: make(chan Record, queueDepth), + done: make(chan struct{}), + w: w, + key: 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.PrevHash = s.prev + rec.Hash = s.hashRecord(rec) + + line, err := json.Marshal(rec) + if err != nil { + s.logger.Error("ledger marshal", "error", err) - s.prev = rec.Hash + continue + } + if _, err := s.w.Write(append(line, '\n')); err != nil { + s.logger.Error("ledger write", "error", err) - return nil + continue + } + + s.setPrev(rec.Hash) + } } -// hashRecord hashes the record with its Hash field cleared, over PrevHash, so -// the chain covers order and content. -func hashRecord(rec Record) string { +// 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 "" } - sum := sha256.Sum256(data) + 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)) +} - return hex.EncodeToString(sum[:]) +func (s *Sink) setPrev(h string) { + s.mu.Lock() + s.prev = h + s.mu.Unlock() } From 9eeb11125fc0bba630f4d0235d1173cddcbff657 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Thu, 27 Aug 2026 16:48:03 +0900 Subject: [PATCH 06/14] fix: record extended-protocol decisions accurately and finalize on teardown --- internal/pg/pg.go | 83 ++++++++++++++++++++++++++++++++++--- internal/pg/pg_test.go | 94 ++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 171 insertions(+), 6 deletions(-) diff --git a/internal/pg/pg.go b/internal/pg/pg.go index 459feca..bb13142 100644 --- a/internal/pg/pg.go +++ b/internal/pg/pg.go @@ -87,8 +87,9 @@ type session struct { uw *bufio.Writer recorder wire.Recorder pending bytes.Buffer - forwarded bool // part of the current batch was already sent upstream - denied bool // a statement in the current batch was denied + batch []string // SQL of the allowed Parses staged in the current batch, awaiting Sync + 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 @@ -202,8 +203,17 @@ func (s *session) Frontend(h wire.Handler, rec wire.Recorder) 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 { @@ -364,6 +374,8 @@ func (s *session) parse(h wire.Handler, n uint32) error { if err := discard(s.cr, n); err != nil { return err } + s.record("", wire.Denied) + s.batch = nil return s.denyBatch(s.tooLarge(n)) } @@ -380,10 +392,13 @@ func (s *session) parse(h wire.Handler, n uint32) error { if v := h.Statement(wire.Statement{SQL: sql}); v.Deny { s.record(sql, wire.Denied) + s.batch = nil // the batch is rejected; earlier allowed statements never run return s.denyBatch(denial(v)) } - s.record(sql, wire.Allowed) + // Defer recording until Sync: a later denial rejects the whole batch, and an + // allowed statement is only real once its batch reaches Sync. + s.batch = append(s.batch, sql) return s.stage(typeParse, body) } @@ -462,12 +477,17 @@ func (s *session) sync(n uint32) error { } } + answer := slot{forwarded: true} + if !s.denied { + answer.result = s.recordBatch() + } + s.batch = 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 } @@ -475,6 +495,21 @@ func (s *session) sync(n uint32) error { return flush(s.uw) } +// recordBatch records the allowed statements accumulated in the current batch. +// A single-statement batch returns a result so its row count is captured; a +// multi-statement batch records each statement now, since results cannot be +// split between them yet. +func (s *session) recordBatch() wire.Result { + if len(s.batch) == 1 { + return s.begin(s.batch[0], wire.Allowed) + } + for _, sql := range s.batch { + s.record(sql, wire.Allowed) + } + + return nil +} + func (s *session) functionCall(n uint32) error { if err := discard(s.cr, n); err != nil { return err @@ -697,6 +732,12 @@ func (s *session) record(sql string, decision wire.Decision) { // describeResult forwards a RowDescription and tells the current result which // columns to capture. func (s *session) describeResult(n uint32) error { + s.curResult = nil + 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 @@ -707,7 +748,6 @@ func (s *session) describeResult(n uint32) error { return err } s.curResult = res - s.curCap = nil if res != nil { s.curCap = res.Columns(columnNames(body)) } @@ -736,6 +776,10 @@ func (s *session) rowResult(n uint32) error { // completeResult forwards a CommandComplete and adds its row count to the result. func (s *session) completeResult(n uint32) error { + if !s.recording() { + return s.forwardToClient(typeCommandComplete, n) // nothing records; stream it + } + body, err := readBody(s.ur, n) if err != nil { return err @@ -769,6 +813,33 @@ func (s *session) forwardResult(typ byte, body []byte) (wire.Result, error) { return res, nil } +// recording reports whether the request currently being answered has a result +// to record, so results are only materialized when the ledger needs them. +func (s *session) recording() bool { + s.mu.Lock() + defer s.mu.Unlock() + + return len(s.queue) > 0 && s.queue[0].forwarded && s.queue[0].result != 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) + } + } + s.queue = nil + s.mu.Unlock() + + for _, r := range pending { + r.Done() + } +} + func (s *session) flushClient() error { s.mu.Lock() defer s.mu.Unlock() diff --git a/internal/pg/pg_test.go b/internal/pg/pg_test.go index bc430d2..ce16a76 100644 --- a/internal/pg/pg_test.go +++ b/internal/pg/pg_test.go @@ -838,6 +838,100 @@ func TestFrontendRecordsDenial(t *testing.T) { } } +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) + } +} + // pipes connects a session to a fake client and a fake upstream over net.Pipe. type pipes struct { sess wire.Session From 95f2c15fd826a9fae1ceb7843ab66530faa0b532 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 11:15:26 +0900 Subject: [PATCH 07/14] fix: consume hex and grouped numeric literals fully in the fingerprint --- internal/sqlscan/fingerprint.go | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/internal/sqlscan/fingerprint.go b/internal/sqlscan/fingerprint.go index 66b4bb3..3d9cbc4 100644 --- a/internal/sqlscan/fingerprint.go +++ b/internal/sqlscan/fingerprint.go @@ -85,9 +85,20 @@ func fingerprintToken(sql string, i int) (int, string) { } } +// 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) && (isDigit(s[i]) || s[i] == '.' || s[i] == 'e' || s[i] == 'E') { - i++ + 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 From 330df69ca6e4505bb6d4fc5e036868434dd78aae Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 11:15:26 +0900 Subject: [PATCH 08/14] fix: bound the fingerprint, key the chain, and read only the ledger tail --- internal/ledger/ledger_test.go | 161 +++++++++++++++++++++++++++++++++ internal/ledger/record.go | 4 + internal/ledger/recorder.go | 20 +++- internal/ledger/resume.go | 59 +++++++++--- internal/ledger/sink.go | 22 ++++- 5 files changed, 249 insertions(+), 17 deletions(-) diff --git a/internal/ledger/ledger_test.go b/internal/ledger/ledger_test.go index 3aa45db..aee95b3 100644 --- a/internal/ledger/ledger_test.go +++ b/internal/ledger/ledger_test.go @@ -4,6 +4,8 @@ import ( "bufio" "bytes" "encoding/json" + "os" + "path/filepath" "strings" "testing" @@ -129,3 +131,162 @@ func TestSinkResumesChainWithPrevAndKey(t *testing.T) { 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 index 0281ca2..4a0fa1f 100644 --- a/internal/ledger/record.go +++ b/internal/ledger/record.go @@ -17,6 +17,10 @@ type Record struct { 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 index 5cb18c8..0542b3e 100644 --- a/internal/ledger/recorder.go +++ b/internal/ledger/recorder.go @@ -1,12 +1,30 @@ 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 @@ -72,7 +90,7 @@ func (r recorder) Begin(sql string, decision wire.Decision) wire.Result { Database: r.principal.Database, Application: r.principal.Application, Kind: kind, - Fingerprint: sqlscan.Fingerprint(sql), + Fingerprint: fingerprint(sql), Decision: decision, } diff --git a/internal/ledger/resume.go b/internal/ledger/resume.go index 7128c3b..1c5b170 100644 --- a/internal/ledger/resume.go +++ b/internal/ledger/resume.go @@ -1,16 +1,22 @@ package ledger import ( - "bufio" + "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. A missing file yields "". +// 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) { @@ -21,25 +27,52 @@ func LastHash(path string) (string, error) { } defer func() { _ = f.Close() }() - last := "" - sc := bufio.NewScanner(f) - sc.Buffer(make([]byte, 0, 64*1024), 16*1024*1024) - for sc.Scan() { - line := sc.Bytes() + 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 { - return "", fmt.Errorf("parse ledger tail: %w", err) + if i == 0 && windowed { + return "", fmt.Errorf("ledger: last record exceeds %d bytes; cannot resume", tailWindow) + } + + continue // a partial trailing write; try the record before it } - last = rec.Hash - } - if err := sc.Err(); err != nil && !errors.Is(err, io.EOF) { - return "", fmt.Errorf("read ledger: %w", err) + + return rec.Hash, nil } - return last, nil + return "", nil } diff --git a/internal/ledger/sink.go b/internal/ledger/sink.go index e6b6300..a079883 100644 --- a/internal/ledger/sink.go +++ b/internal/ledger/sink.go @@ -11,14 +11,16 @@ import ( ) // Sink appends records as JSON lines, chaining each record's hash to the -// previous so the log is tamper-evident. Writes are queued to a single writer -// goroutine, so a slow disk never blocks a session's response path or serializes -// sessions against each other. +// 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 @@ -54,6 +56,7 @@ func NewSink(w writer, opts Options) *Sink { done: make(chan struct{}), w: w, key: opts.Key, + keyID: keyID(opts.Key), logger: logger, prev: opts.Prev, } @@ -87,6 +90,7 @@ 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) @@ -106,6 +110,18 @@ func (s *Sink) run() { } } +// 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 "" + } + + sum := sha256.Sum256(key) + + return hex.EncodeToString(sum[:6]) +} + // 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 { From f27515cd8731924f358d05924d9a3fc315547f5f Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 11:15:26 +0900 Subject: [PATCH 09/14] fix: finalize a record after releasing the session lock --- internal/pg/pg.go | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/internal/pg/pg.go b/internal/pg/pg.go index bb13142..a409f48 100644 --- a/internal/pg/pg.go +++ b/internal/pg/pg.go @@ -643,6 +643,15 @@ func (s *session) readyForQuery(n uint32) error { return fmt.Errorf("read ReadyForQuery: %w", err) } + // finishing is recorded 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() { + if finishing != nil { + finishing.Done() + } + }() + s.mu.Lock() defer s.mu.Unlock() @@ -656,9 +665,7 @@ 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 { - r.Done() - } + finishing = s.queue[0].result s.curResult = nil s.curCap = nil s.queue = s.queue[1:] From 5b97d1140dbb9ecc7a364e70576648ada4173340 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 11:49:50 +0900 Subject: [PATCH 10/14] fix: record oversized simple queries, key the ledger id, and skip empty-hash tails --- internal/ledger/resume.go | 15 +++++++-------- internal/ledger/sink.go | 5 +++-- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/internal/ledger/resume.go b/internal/ledger/resume.go index 1c5b170..38ff085 100644 --- a/internal/ledger/resume.go +++ b/internal/ledger/resume.go @@ -63,15 +63,14 @@ func lastHashInTail(tail []byte, windowed bool) (string, error) { var rec struct { Hash string `json:"hash"` } - if err := json.Unmarshal(line, &rec); err != nil { - if i == 0 && windowed { - return "", fmt.Errorf("ledger: last record exceeds %d bytes; cannot resume", tailWindow) - } - - continue // a partial trailing write; try the record before it + 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 rec.Hash, nil } return "", nil diff --git a/internal/ledger/sink.go b/internal/ledger/sink.go index a079883..5d58f29 100644 --- a/internal/ledger/sink.go +++ b/internal/ledger/sink.go @@ -117,9 +117,10 @@ func keyID(key []byte) string { return "" } - sum := sha256.Sum256(key) + m := hmac.New(sha256.New, key) + _, _ = m.Write([]byte("rollcall/ledger/key-id")) - return hex.EncodeToString(sum[:6]) + return hex.EncodeToString(m.Sum(nil)[:8]) } // hashRecord hashes the record with its Hash field cleared, so the chain covers From 3e3ae37c269d252616d0fafbe8bd10b03ee0952c Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 11:49:50 +0900 Subject: [PATCH 11/14] feat: record prepared-statement executions, including re-execution without Parse --- internal/pg/message.go | 2 + internal/pg/pg.go | 296 +++++++++++++++++++++++++++++++---------- internal/pg/pg_test.go | 125 +++++++++++++++++ 3 files changed, 351 insertions(+), 72 deletions(-) diff --git a/internal/pg/message.go b/internal/pg/message.go index 6f015f4..81edace 100644 --- a/internal/pg/message.go +++ b/internal/pg/message.go @@ -46,6 +46,8 @@ const ( typeRowDescription = 'T' typeDataRow = 'D' typeCommandComplete = 'C' + typeEmptyQuery = 'I' + typePortalSuspended = 's' ) const ( diff --git a/internal/pg/pg.go b/internal/pg/pg.go index a409f48..e7ef608 100644 --- a/internal/pg/pg.go +++ b/internal/pg/pg.go @@ -58,6 +58,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) @@ -73,6 +75,7 @@ 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 @@ -83,28 +86,31 @@ type session struct { maxStatement uint32 // Frontend goroutine only. - cr *bufio.Reader - uw *bufio.Writer - recorder wire.Recorder - pending bytes.Buffer - batch []string // SQL of the allowed Parses staged in the current batch, awaiting Sync - 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 of the Executes staged in the current batch, awaiting Sync + 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 - curResult wire.Result // the result currently streaming from the upstream - curCap []int // columns the recorder asked to capture for it + 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) { @@ -245,6 +251,8 @@ func (s *session) backend() error { err = s.rowResult(n) case typeCommandComplete: err = s.completeResult(n) + case typeEmptyQuery, typePortalSuspended: + err = s.endResult(typ, n) default: err = s.forwardToClient(typ, n) } @@ -305,7 +313,11 @@ 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 typeDescribe, typeClose: return false, s.buffer(typ, n) case typeFlush: return false, s.flushBatch(n) @@ -339,6 +351,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))) } @@ -350,7 +363,7 @@ 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.record(sql, wire.Denied) + s.recordDenied(sql) return s.respond(readySlot(denial(v))) } @@ -374,8 +387,8 @@ func (s *session) parse(h wire.Handler, n uint32) error { if err := discard(s.cr, n); err != nil { return err } - s.record("", wire.Denied) - s.batch = nil + s.recordDenied("") + s.pendingExecs = nil return s.denyBatch(s.tooLarge(n)) } @@ -385,24 +398,78 @@ 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.record(sql, wire.Denied) - s.batch = nil // the batch is rejected; earlier allowed statements never run + s.recordDenied(sql) + s.pendingExecs = nil // the batch is rejected; earlier statements never run return s.denyBatch(denial(v)) } - // Defer recording until Sync: a later denial rejects the whole batch, and an - // allowed statement is only real once its batch reaches Sync. - s.batch = append(s.batch, sql) + // Remember the prepared statement so its later Executes, which carry no SQL, + // can be attributed even across batches (driver statement caches reuse it). + s.prepared[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.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 + + return forward(s.uw, typeExecute, n, s.cr) + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + if portal, ok := parsePortal(body); ok { + if sql, ok := s.prepared[s.portals[portal]]; ok { + s.pendingExecs = append(s.pendingExecs, sql) + } + } + + return s.stage(typeExecute, body) +} + func (s *session) buffer(typ byte, n uint32) error { if s.denied { return discard(s.cr, n) @@ -479,9 +546,9 @@ func (s *session) sync(n uint32) error { answer := slot{forwarded: true} if !s.denied { - answer.result = s.recordBatch() + answer.execCount = s.enqueueExecs() } - s.batch = nil + s.pendingExecs = nil s.denied = false s.forwarded = false @@ -495,19 +562,25 @@ func (s *session) sync(n uint32) error { return flush(s.uw) } -// recordBatch records the allowed statements accumulated in the current batch. -// A single-statement batch returns a result so its row count is captured; a -// multi-statement batch records each statement now, since results cannot be -// split between them yet. -func (s *session) recordBatch() wire.Result { - if len(s.batch) == 1 { - return s.begin(s.batch[0], wire.Allowed) +// 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 } - for _, sql := range s.batch { - s.record(sql, wire.Allowed) + + s.mu.Lock() + defer s.mu.Unlock() + + count := 0 + for _, sql := range s.pendingExecs { + if r := s.begin(sql, wire.Allowed); r != nil { + s.execQueue = append(s.execQueue, r) + count++ + } } - return nil + return count } func (s *session) functionCall(n uint32) error { @@ -532,6 +605,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 { @@ -643,12 +717,12 @@ func (s *session) readyForQuery(n uint32) error { return fmt.Errorf("read ReadyForQuery: %w", err) } - // finishing is recorded 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 + // 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() { - if finishing != nil { - finishing.Done() + for _, r := range finishing { + r.Done() } }() @@ -665,8 +739,14 @@ func (s *session) readyForQuery(n uint32) error { return fmt.Errorf("%w: ReadyForQuery without a pending request", errMalformed) } - finishing = s.queue[0].result - s.curResult = nil + 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] @@ -728,10 +808,10 @@ func (s *session) begin(sql string, decision wire.Decision) wire.Result { return s.recorder.Begin(sql, decision) } -// record writes a ledger record for a statement with no result to observe, such -// as a denial or an extended-protocol Parse. -func (s *session) record(sql string, decision wire.Decision) { - if r := s.begin(sql, decision); r != nil { +// 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() } } @@ -739,7 +819,6 @@ func (s *session) record(sql string, decision wire.Decision) { // describeResult forwards a RowDescription and tells the current result which // columns to capture. func (s *session) describeResult(n uint32) error { - s.curResult = nil s.curCap = nil if !s.recording() { return s.forwardToClient(typeRowDescription, n) // nothing records; stream it @@ -754,7 +833,6 @@ func (s *session) describeResult(n uint32) error { if err != nil { return err } - s.curResult = res if res != nil { s.curCap = res.Columns(columnNames(body)) } @@ -765,7 +843,7 @@ func (s *session) describeResult(n uint32) error { // rowResult forwards a DataRow, capturing the requested columns when the result // is being recorded. func (s *session) rowResult(n uint32) error { - if s.curResult == nil || len(s.curCap) == 0 { + if len(s.curCap) == 0 { return s.forwardToClient(typeDataRow, n) // stream without materializing } @@ -773,18 +851,27 @@ func (s *session) rowResult(n uint32) error { if err != nil { return err } - if _, err := s.forwardResult(typeDataRow, body); err != nil { + res, err := s.forwardResult(typeDataRow, body) + if err != nil { return err } - s.curResult.Row(rowValues(body, s.curCap)) + if res != nil { + res.Row(rowValues(body, s.curCap)) + } return nil } -// completeResult forwards a CommandComplete and adds its row count to the result. +// 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() { - return s.forwardToClient(typeCommandComplete, n) // nothing records; stream it + if err := s.forwardToClient(typeCommandComplete, n); err != nil { + return err + } + s.advanceExec() + + return nil } body, err := readBody(s.ur, n) @@ -799,20 +886,44 @@ func (s *session) completeResult(n uint32) error { 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 -// result attached to the request currently being answered, if any. +// 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() - var res wire.Result - if len(s.queue) > 0 && s.queue[0].forwarded { - res = s.queue[0].result - } + res := s.targetLocked() if err := writeMessage(s.cw, typ, body); err != nil { return nil, err } @@ -820,13 +931,28 @@ func (s *session) forwardResult(typ byte, body []byte) (wire.Result, error) { return res, nil } -// recording reports whether the request currently being answered has a result -// to record, so results are only materialized when the ledger needs them. +// 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 < 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 len(s.queue) > 0 && s.queue[0].forwarded && s.queue[0].result != nil + return s.targetLocked() != nil } // finishInFlight finalizes the records of forwarded requests still queued when @@ -839,7 +965,9 @@ func (s *session) finishInFlight() { pending = append(pending, sl.result) } } + pending = append(pending, s.execQueue...) s.queue = nil + s.execQueue = nil s.mu.Unlock() for _, r := range pending { @@ -917,20 +1045,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) + } + + 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 "", fmt.Errorf("parse message: %w", err) + return "", "", false } - sql, _, err := cstring(rest) + 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 ce16a76..7637fc7 100644 --- a/internal/pg/pg_test.go +++ b/internal/pg/pg_test.go @@ -932,6 +932,131 @@ func TestFrontendRecordsOnlyTheDenialOfARejectedBatch(t *testing.T) { } } +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) + } +} + // pipes connects a session to a fake client and a fake upstream over net.Pipe. type pipes struct { sess wire.Session From abfc4c7f8aa9dd6784ce57f331c2e9d5104a64be Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 11:49:50 +0900 Subject: [PATCH 12/14] test: cover the ledger end to end for prepared-statement re-execution --- internal/cli/cli_test.go | 126 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 126 insertions(+) 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) } From f00d1596fd01c07602ff0fe92fbddac648741d74 Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 14:23:22 +0900 Subject: [PATCH 13/14] fix: align Execute results per batch, resolve portals safely, and bound the prepared-statement maps --- internal/pg/pg.go | 121 ++++++++++++++++++++++++---- internal/pg/pg_test.go | 174 +++++++++++++++++++++++++++++++++++++++++ 2 files changed, 278 insertions(+), 17 deletions(-) diff --git a/internal/pg/pg.go b/internal/pg/pg.go index e7ef608..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 @@ -92,7 +96,8 @@ type session struct { pending bytes.Buffer prepared map[string]string // prepared statement name -> SQL portals map[string]string // portal name -> prepared statement name - pendingExecs []string // SQL of the Executes staged in the current batch, awaiting Sync + 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 @@ -317,7 +322,9 @@ func (s *session) dispatch(h wire.Handler, typ byte, n uint32) (bool, error) { return false, s.bind(n) case typeExecute: return false, s.execute(n) - case typeDescribe, typeClose: + case typeClose: + return false, s.closePrepared(n) + case typeDescribe: return false, s.buffer(typ, n) case typeFlush: return false, s.flushBatch(n) @@ -411,7 +418,7 @@ func (s *session) parse(h wire.Handler, n uint32) error { } // Remember the prepared statement so its later Executes, which carry no SQL, // can be attributed even across batches (driver statement caches reuse it). - s.prepared[name] = sql + s.storePrepared(name, sql) return s.stage(typeParse, body) } @@ -435,8 +442,10 @@ func (s *session) bind(n uint32) error { if err != nil { return err } - if portal, stmt, ok := parseBind(body); ok { - s.portals[portal] = stmt + 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) @@ -453,6 +462,7 @@ func (s *session) execute(n uint32) error { return err } s.forwarded = true + s.stageExec("") // forwarded without inspection; a placeholder keeps results aligned return forward(s.uw, typeExecute, n, s.cr) } @@ -461,13 +471,86 @@ func (s *session) execute(n uint32) error { if err != nil { return err } - if portal, ok := parsePortal(body); ok { - if sql, ok := s.prepared[s.portals[portal]]; ok { - s.pendingExecs = append(s.pendingExecs, sql) + 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) } - return s.stage(typeExecute, body) + 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 { @@ -572,15 +655,15 @@ func (s *session) enqueueExecs() int { s.mu.Lock() defer s.mu.Unlock() - count := 0 for _, sql := range s.pendingExecs { - if r := s.begin(sql, wire.Allowed); r != nil { - s.execQueue = append(s.execQueue, r) - count++ + 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 count + return len(s.pendingExecs) } func (s *session) functionCall(n uint32) error { @@ -722,7 +805,9 @@ func (s *session) readyForQuery(n uint32) error { var finishing []wire.Result defer func() { for _, r := range finishing { - r.Done() + if r != nil { + r.Done() + } } }() @@ -939,7 +1024,7 @@ func (s *session) targetLocked() wire.Result { if s.queue[0].simple { return s.queue[0].result } - if s.execIdx < len(s.execQueue) { + if s.execIdx < s.queue[0].execCount && s.execIdx < len(s.execQueue) { return s.execQueue[s.execIdx] } @@ -971,7 +1056,9 @@ func (s *session) finishInFlight() { s.mu.Unlock() for _, r := range pending { - r.Done() + if r != nil { + r.Done() + } } } diff --git a/internal/pg/pg_test.go b/internal/pg/pg_test.go index 7637fc7..ebe3510 100644 --- a/internal/pg/pg_test.go +++ b/internal/pg/pg_test.go @@ -1057,6 +1057,180 @@ func TestFrontendDoesNotRecordAPrepareOnlyBatch(t *testing.T) { } } +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 From 12d62b4ecd1a34033c553b2787b230c43d28930a Mon Sep 17 00:00:00 2001 From: Tetsuro Mikami Date: Mon, 31 Aug 2026 14:25:53 +0900 Subject: [PATCH 14/14] test: drop the shared clock override that raced with parallel ledger tests --- internal/ledger/export_test.go | 9 --------- internal/ledger/ledger_test.go | 2 -- 2 files changed, 11 deletions(-) delete mode 100644 internal/ledger/export_test.go diff --git a/internal/ledger/export_test.go b/internal/ledger/export_test.go deleted file mode 100644 index ec16ed1..0000000 --- a/internal/ledger/export_test.go +++ /dev/null @@ -1,9 +0,0 @@ -package ledger - -// SetNow overrides the record clock for tests and returns a restore function. -func SetNow(t func() string) func() { - prev := now - now = t - - return func() { now = prev } -} diff --git a/internal/ledger/ledger_test.go b/internal/ledger/ledger_test.go index aee95b3..f6f1800 100644 --- a/internal/ledger/ledger_test.go +++ b/internal/ledger/ledger_test.go @@ -16,8 +16,6 @@ import ( func TestGuardRecordsStatements(t *testing.T) { t.Parallel() - defer ledger.SetNow(func() string { return "2026-08-27T00:00:00Z" })() - var buf bytes.Buffer inner := wire.GuardFunc(func(s wire.Startup) wire.Enforcement { return wire.Enforcement{