Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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
Expand Down
126 changes: 126 additions & 0 deletions internal/cli/cli_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,11 @@ import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"io"
"net"
"os"
"path/filepath"
"regexp"
"strings"
"sync"
Expand All @@ -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) {
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand All @@ -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)
}
Expand Down
35 changes: 34 additions & 1 deletion internal/cli/proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -21,13 +23,15 @@ 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 {
fs := newFlagSet("proxy", std.Err, printProxyUsage)
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
Expand Down Expand Up @@ -57,14 +61,32 @@ func runProxy(ctx context.Context, args []string, std IO) int {
guard = p
}

logger := slog.New(slog.NewTextHandler(std.Err, nil))

if *ledgerPath != "" {
prev, err := ledger.LastHash(*ledgerPath)
if err != nil {
return fail(std, err)
}

f, err := os.OpenFile(*ledgerPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0o600)
if err != nil {
return fail(std, fmt.Errorf("open ledger: %w", err))
}
defer func() { _ = f.Close() }()

sink := ledger.NewSink(f, ledger.Options{Prev: prev, Key: ledgerKey(), Logger: logger})
defer sink.Close()
guard = ledger.Guard{Inner: guard, Sink: sink}
}

var lc net.ListenConfig
ln, err := lc.Listen(ctx, "tcp", *listen)
if err != nil {
return fail(std, err)
}
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)
Expand All @@ -78,6 +100,16 @@ func runProxy(ctx context.Context, args []string, std IO) int {
return exit.OK
}

// ledgerKey returns the optional key that makes the ledger chain an HMAC,
// from ROLLCALL_LEDGER_KEY. Without it the chain is a plain SHA-256.
func ledgerKey() []byte {
if v := os.Getenv("ROLLCALL_LEDGER_KEY"); v != "" {
return []byte(v)
}

return nil
}

func validateAddr(flagName, addr string) error {
if addr == "" {
return fmt.Errorf("%s is required", flagName)
Expand Down Expand Up @@ -118,4 +150,5 @@ func printProxyUsage(w io.Writer) {
fmt.Fprintf(w, " -upstream ADDR %s (required)\n", upstreamUsage)
fmt.Fprintf(w, " -listen ADDR %s (default %q)\n", listenUsage, defaultListen)
fmt.Fprintf(w, " -policy PATH %s\n", policyUsage)
fmt.Fprintf(w, " -ledger PATH %s\n", ledgerUsage)
}
7 changes: 7 additions & 0 deletions internal/ledger/clock.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
package ledger

import "time"

func defaultNow() string {
return time.Now().UTC().Format(time.RFC3339Nano)
}
Loading
Loading