diff --git a/CHANGELOG.md b/CHANGELOG.md index 2e1f58f..80e3f16 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -18,19 +18,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 Postgres 16 primary and two streaming replicas over real physical replication, with dedicated replication slots, a seeded test schema, and a standby bootstrap via `pg_basebackup`. `make up` produces a working cluster - and `make smoke` asserts that a row written to the primary is replayed and - served by both replicas. Rationale recorded in + and `make smoke` asserts replication. Rationale in `docs/adr/0001-dev-cluster-replication.md`. - Transparent proxy (Phase 2): `internal/proxy` accepts client connections, - refuses SSL/GSS negotiation with `N`, forwards the startup and authentication - exchange to the primary untouched, and pipes bytes bidirectionally with - graceful, leak-free shutdown. `cmd/pgpilot` now runs the proxy (`-listen`, - `-primary`). An integration test (`make itest`) asserts that psql through the - proxy behaves identically to psql direct. Rationale recorded in + refuses SSL/GSS negotiation with `N`, and pipes bytes bidirectionally with + graceful, leak-free shutdown. An integration test (`make itest`) asserts that + psql through the proxy behaves identically to psql direct. Rationale in `docs/adr/0002-transparent-proxy-ssl-refusal.md`. - Protocol codec (Phase 3): `internal/protocol` decodes the wire messages pgpilot routes on via `jackc/pgx/v5/pgproto3`, with round-trip tests and a - panic-safe, fuzzed decoder. The proxy now relays messages frame-by-frame and + panic-safe, fuzzed decoder. The proxy relays messages frame-by-frame and tracks each session's transaction status from ReadyForQuery. `cmd/pgpilot` gains a `-log-level` flag. Rationale in `docs/adr/0003-message-aware-relay.md`. @@ -38,16 +35,24 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 of backend connections — configurable max size, acquire timeout, idle timeout with a background reaper, and per-connection health checks — that applies backpressure instead of queueing acquirers without bound. -- SCRAM authentication and pooling foundation (Phase 4b): `internal/scram` - implements SCRAM-SHA-256 for both the client and server roles; - `internal/backend` opens and authenticates pgpilot's own connections to the - primary with SCRAM, resets them for reuse (`ROLLBACK`, then `DISCARD ALL`), - and manages one connection pool per `(user, database)`; `internal/config` - loads pgpilot's JSON configuration (backend, users, pool sizing). The SCRAM - client is validated against real PostgreSQL. Wiring the client-facing SCRAM - authenticator and session pooling into the proxy is the remaining half of this - phase. Rationale recorded in - `docs/adr/0004-auth-termination-and-pooling.md`. +- SCRAM authentication and session pooling (Phase 4b): pgpilot now terminates + authentication and pools backend connections. + - `internal/scram` implements SCRAM-SHA-256 for both the client and server + roles; `internal/backend` opens and authenticates connections to the primary + with SCRAM, resets them for reuse (`ROLLBACK`, then `DISCARD ALL`), and + manages one pool per `(user, database)`; `internal/config` loads pgpilot's + JSON configuration. + - The proxy authenticates each client with SCRAM-SHA-256, acquires a pooled + backend for the client's `(user, database)`, replays the backend's startup + parameters, relays the session while tracking transaction status, intercepts + the client's Terminate, and resets and returns the backend to the pool on a + clean disconnect (discarding it otherwise). + - `cmd/pgpilot` now runs from a `-config ` (replacing `-primary`); see + `pgpilot.example.json`. + - SCRAM is validated against real PostgreSQL (backend) and real psql + (client); `make itest` asserts psql through pgpilot matches psql direct. + Transaction pooling and feature detection follow in Phase 4c. Rationale in + `docs/adr/0004-auth-termination-and-pooling.md`. ### Dependencies diff --git a/README.md b/README.md index 61ee43c..72f8620 100644 --- a/README.md +++ b/README.md @@ -45,28 +45,28 @@ The goal is to do one thing — correct, observable read/write routing — well. ## Status Early development, built in phases (see the roadmap). Not production-ready yet. -Today pgpilot is a **protocol-aware transparent proxy**: it forwards every -connection to a single primary, but it now frames and decodes the wire protocol -and tracks each session's transaction status. Read/write routing and fencing -arrive in later phases. +Today pgpilot is an **authenticating connection pooler**: it verifies each +client with SCRAM-SHA-256, hands it a pooled, SCRAM-authenticated connection to a +single primary, and relays the session. Read/write routing across replicas and +LSN fencing arrive in later phases. ## Roadmap -| Phase | Focus | Status | -| ----: | ------------------------------------------------------------ | ------ | -| 0 | Repo hygiene, CI, licensing | done | -| 1 | Dev cluster: primary + 2 streaming replicas (docker-compose) | done | -| 2 | Transparent proxy (byte-level passthrough) | done | -| 3 | Protocol codec (typed frontend/backend messages) | done | -| 4 | Connection pooling (session + transaction) | next | -| 5 | Query classification (read vs. write via pg_query) | | -| 6 | Replica registry, health polling, circuit breakers | | -| 7 | LSN fencing | | -| 8 | Routing policy engine | | -| 9 | Observability (Prometheus, structured logs, pprof) | | -| 10 | Fault-injection harness | | -| 11 | Benchmarks vs. direct connection and pgbouncer | | -| 12 | Docs and the v0.1.0 release | | +| Phase | Focus | Status | +| ----: | ------------------------------------------------------------ | ----------- | +| 0 | Repo hygiene, CI, licensing | done | +| 1 | Dev cluster: primary + 2 streaming replicas (docker-compose) | done | +| 2 | Transparent proxy (byte-level passthrough) | done | +| 3 | Protocol codec (typed frontend/backend messages) | done | +| 4 | Connection pooling (session + transaction) | in progress | +| 5 | Query classification (read vs. write via pg_query) | | +| 6 | Replica registry, health polling, circuit breakers | | +| 7 | LSN fencing | | +| 8 | Routing policy engine | | +| 9 | Observability (Prometheus, structured logs, pprof) | | +| 10 | Fault-injection harness | | +| 11 | Benchmarks vs. direct connection and pgbouncer | | +| 12 | Docs and the v0.1.0 release | | ## Technology @@ -85,7 +85,7 @@ make test # run tests with the race detector make lint # run golangci-lint make up # bring up the local primary + replica cluster make smoke # assert the cluster replicates (run after `make up`) -make itest # assert psql through the proxy matches psql direct +make itest # assert psql through the pooler matches psql direct make down # tear the cluster down make bench # run benchmarks ``` @@ -107,66 +107,55 @@ Postgres 16 primary and two streaming replicas. The primary is configured for physical streaming replication with a dedicated replication slot per standby. Each replica bootstraps by cloning the primary with `pg_basebackup` on its first start, then streams WAL to stay current. The -seeded schema lives in `docker/primary/initdb/`. - -```sh -make up # primary + 2 replicas, waits until healthy -make smoke # a write on the primary is served by both replicas -psql -h localhost -p 55432 -U pgpilot pgpilot # connect to the primary -make down # stop and delete the cluster's volumes -``` - -The smoke test (`test/smoke`) writes a unique row to the primary, waits for each -replica to replay past the write's LSN, and asserts the row is readable there — -the same read-your-writes invariant pgpilot will enforce automatically once LSN -fencing lands in Phase 7. - -The design decisions behind the cluster are recorded in +seeded schema lives in `docker/primary/initdb/`. The design decisions are +recorded in [`docs/adr/0001-dev-cluster-replication.md`](docs/adr/0001-dev-cluster-replication.md). ## Running the proxy -pgpilot accepts client connections and forwards them to a single upstream -primary. It refuses TLS (replying `N` to `SSLRequest`) and passes the startup -and authentication exchange through untouched, so any auth method the backend -uses — including SCRAM — works end-to-end. +pgpilot reads a JSON config file (see [`pgpilot.example.json`](pgpilot.example.json)): + +```json +{ + "listen": "127.0.0.1:6432", + "primary": "127.0.0.1:55432", + "users": [{"name": "pgpilot", "password": "pgpilot"}], + "pool": {"max_size": 10, "max_waiters": 100, "acquire_timeout": "5s", "idle_timeout": "5m"} +} +``` ```sh -make up # backends on 55432–55434 +make up # backends on 55432–55434 make build && ./bin/pgpilot \ - -listen 127.0.0.1:6432 \ - -primary 127.0.0.1:55432 \ - -log-level debug # proxy on 6432 -> primary on 55432 + -config pgpilot.example.json \ + -log-level debug -# Connect through the proxy exactly as you would connect directly: +# Connect through pgpilot exactly as you would connect directly: psql "host=localhost port=6432 dbname=pgpilot user=pgpilot sslmode=prefer" ``` +pgpilot **verifies the client** with SCRAM-SHA-256 against the password in the +config, then hands the session a pooled connection it opened to the primary +(also with SCRAM). Connections are keyed by `(user, database)`, reset with +`DISCARD ALL` between clients, and drawn from a bounded pool that applies +backpressure. `make itest` asserts a session through pgpilot is byte-for-byte +equivalent to a direct one. + Because TLS is refused for now, clients must permit a cleartext connection to pgpilot; libpq's default `sslmode=prefer` falls back automatically, whereas -`sslmode=require` will fail until TLS termination arrives (see -[`docs/adr/0002-transparent-proxy-ssl-refusal.md`](docs/adr/0002-transparent-proxy-ssl-refusal.md)). -`make itest` asserts that a session through the proxy is byte-for-byte -equivalent to a direct one. +`sslmode=require` fails until TLS termination lands. Cancel requests are not yet +supported. The auth and pooling design is recorded in +[`docs/adr/0004-auth-termination-and-pooling.md`](docs/adr/0004-auth-termination-and-pooling.md). ### Protocol awareness -Rather than piping opaque bytes, pgpilot frames every message on the wire and -decodes the ones it routes on, using `jackc/pgx`'s `pgproto3` (see -[`internal/protocol`](internal/protocol)). It tracks each session's transaction -status from the backend's ReadyForQuery indicator (`I`/`T`/`E`) — the signal the -router will use to pin a session to one backend for the life of a transaction. -Run with `-log-level debug` to watch the transitions: - -``` -level=DEBUG msg="transaction status" session=1 status=idle -level=DEBUG msg="transaction status" session=1 status="in transaction" -level=DEBUG msg="transaction status" session=1 status=idle -``` - -The message decoder is fuzzed and hardened against malformed input; the design -is recorded in -[`docs/adr/0003-message-aware-relay.md`](docs/adr/0003-message-aware-relay.md). +pgpilot frames every message on the wire and decodes the ones it routes on, +using `jackc/pgx`'s `pgproto3` (see [`internal/protocol`](internal/protocol)). It +tracks each session's transaction status from the backend's ReadyForQuery +indicator (`I`/`T`/`E`) — the signal the router will use to pin a session to one +backend for the life of a transaction. Run with `-log-level debug` to watch the +transitions. The message decoder is fuzzed and hardened against malformed input +(see [`docs/adr/0003-message-aware-relay.md`](docs/adr/0003-message-aware-relay.md)). ## License diff --git a/cmd/pgpilot/main.go b/cmd/pgpilot/main.go index 3ba7b11..179742f 100644 --- a/cmd/pgpilot/main.go +++ b/cmd/pgpilot/main.go @@ -1,8 +1,9 @@ // Command pgpilot is a transparent, LSN-fencing PostgreSQL connection router. // -// At this phase it is a transparent proxy: it forwards every connection to a -// single upstream primary. Routing, pooling, and fencing arrive in later -// phases. See the roadmap in README.md. +// At this phase it is an authenticating connection pooler: it verifies each +// client with SCRAM-SHA-256 and hands it a pooled connection to a single +// primary. Read/write routing and fencing arrive in later phases. See the +// roadmap in README.md. package main import ( @@ -14,8 +15,9 @@ import ( "os" "os/signal" "syscall" - "time" + "github.com/sachhg/pgpilot/internal/backend" + "github.com/sachhg/pgpilot/internal/config" "github.com/sachhg/pgpilot/internal/proxy" ) @@ -31,8 +33,7 @@ func main() { func run(args []string) error { fs := flag.NewFlagSet("pgpilot", flag.ContinueOnError) - listen := fs.String("listen", "127.0.0.1:6432", "address to accept client connections on") - primary := fs.String("primary", "127.0.0.1:55432", "address of the upstream PostgreSQL primary") + configPath := fs.String("config", "pgpilot.json", "path to the JSON config file") logLevel := fs.String("log-level", "info", "log level: debug, info, warn, or error") showVersion := fs.Bool("version", false, "print version and exit") if err := fs.Parse(args); err != nil { @@ -51,18 +52,34 @@ func run(args []string) error { return fmt.Errorf("invalid -log-level %q: %w", *logLevel, err) } logger := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: level})) - srv := proxy.New(proxy.Config{ - ListenAddr: *listen, - UpstreamAddr: *primary, - DialTimeout: 5 * time.Second, - Logger: logger, + + cfg, err := config.Load(*configPath) + if err != nil { + return err + } + users := make(map[string]string, len(cfg.Users)) + for _, u := range cfg.Users { + users[u.Name] = u.Password + } + mgr := backend.NewManager(cfg.Primary, users, backend.PoolConfig{ + MaxSize: cfg.Pool.MaxSize, + MaxWaiters: cfg.Pool.MaxWaiters, + AcquireTimeout: cfg.Pool.AcquireTimeout.Std(), + IdleTimeout: cfg.Pool.IdleTimeout.Std(), }) + defer mgr.Close() + srv := proxy.New(proxy.Config{ + ListenAddr: cfg.Listen, + Users: cfg, + Manager: mgr, + Logger: logger, + }) addr, err := srv.Listen() if err != nil { return err } - logger.Info("pgpilot listening", "addr", addr.String(), "primary", *primary, "version", version) + logger.Info("pgpilot listening", "addr", addr.String(), "primary", cfg.Primary, "version", version) ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) defer stop() diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index c26a2ae..8bd6215 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -1,8 +1,7 @@ -// Package proxy implements a transparent PostgreSQL proxy: it accepts client -// connections, refuses TLS negotiation, forwards the startup and authentication -// exchange to an upstream backend untouched, and pipes bytes in both directions -// for the life of each session. It does not yet parse or route queries; that -// arrives in later phases. +// Package proxy is pgpilot's client-facing server. It accepts client +// connections, refuses TLS, authenticates each client with SCRAM-SHA-256, +// acquires a pooled backend connection for the client's (user, database), and +// relays protocol messages between them for the life of the session. package proxy import ( @@ -15,6 +14,8 @@ import ( "sync/atomic" "time" + "github.com/sachhg/pgpilot/internal/backend" + "github.com/sachhg/pgpilot/internal/config" "github.com/sachhg/pgpilot/internal/protocol" ) @@ -22,20 +23,19 @@ import ( type Config struct { // ListenAddr is the TCP address the proxy accepts client connections on. ListenAddr string - // UpstreamAddr is the PostgreSQL backend every session is forwarded to. - UpstreamAddr string - // DialTimeout bounds a single upstream dial. Zero selects a default. - DialTimeout time.Duration + // Users holds the credentials pgpilot verifies clients against. + Users *config.Config + // Manager supplies pooled, authenticated backend connections. + Manager *backend.Manager // Logger receives structured logs. Nil selects slog.Default. Logger *slog.Logger } -// Server is a transparent PostgreSQL proxy. The zero value is not usable; -// construct a Server with New, call Listen, then Serve. +// Server is pgpilot's client-facing proxy. Construct it with New, then call +// Listen and Serve. type Server struct { - cfg Config - log *slog.Logger - dialer *net.Dialer + cfg Config + log *slog.Logger sessions atomic.Uint64 wg sync.WaitGroup @@ -50,15 +50,7 @@ func New(cfg Config) *Server { if logger == nil { logger = slog.Default() } - timeout := cfg.DialTimeout - if timeout <= 0 { - timeout = 5 * time.Second - } - return &Server{ - cfg: cfg, - log: logger, - dialer: &net.Dialer{Timeout: timeout}, - } + return &Server{cfg: cfg, log: logger} } // Listen binds the configured listen address and returns the resolved address. @@ -74,8 +66,7 @@ func (s *Server) Listen() (net.Addr, error) { return ln.Addr(), nil } -// Addr returns the address the proxy is listening on, or nil if Listen has not -// been called. +// Addr returns the address the proxy is listening on, or nil before Listen. func (s *Server) Addr() net.Addr { s.mu.Lock() defer s.mu.Unlock() @@ -86,8 +77,7 @@ func (s *Server) Addr() net.Addr { } // Serve accepts connections until ctx is cancelled, then drains in-flight -// sessions and returns. Because every session is bound to ctx, no goroutine -// outlives Serve. +// sessions and returns. func (s *Server) Serve(ctx context.Context) error { s.mu.Lock() ln := s.listener @@ -98,17 +88,17 @@ func (s *Server) Serve(ctx context.Context) error { go func() { <-ctx.Done() - _ = ln.Close() // unblock Accept + _ = ln.Close() }() for { client, err := ln.Accept() if err != nil { if ctx.Err() != nil { - break // shutting down + break } s.log.Warn("accept failed", "error", err) - time.Sleep(10 * time.Millisecond) // avoid a hot loop on transient errors + time.Sleep(10 * time.Millisecond) continue } s.wg.Add(1) @@ -129,11 +119,11 @@ func (s *Server) handle(ctx context.Context, client net.Conn) { log.Info("session opened") sess := &session{ - client: client, - upstream: s.cfg.UpstreamAddr, - dialer: s.dialer, - log: log, - tracker: &protocol.TxTracker{}, + client: client, + cfg: s.cfg.Users, + manager: s.cfg.Manager, + log: log, + tracker: &protocol.TxTracker{}, } if err := sess.serve(ctx); err != nil { log.Warn("session closed", "error", err) diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index a193a98..47c7a71 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -2,16 +2,23 @@ package proxy import ( "bytes" - "context" "encoding/binary" + "fmt" "io" "log/slog" "net" - "runtime" "testing" "time" + + "github.com/jackc/pgx/v5/pgproto3" + + "github.com/sachhg/pgpilot/internal/scram" ) +func discardLogger() *slog.Logger { + return slog.New(slog.NewTextHandler(io.Discard, nil)) +} + // framedMessage builds a length-prefixed protocol message (type, length, body). func framedMessage(msgType byte, body []byte) []byte { out := make([]byte, 5+len(body)) @@ -21,174 +28,153 @@ func framedMessage(msgType byte, body []byte) []byte { return out } -func discardLogger() *slog.Logger { - return slog.New(slog.NewTextHandler(io.Discard, nil)) -} +// scramClientHandshake plays a SCRAM-SHA-256 client against pgpilot's server side +// over conn, expecting to reach AuthenticationOk. +func scramClientHandshake(conn net.Conn, password string) error { + fe := pgproto3.NewFrontend(conn, conn) -// fakeUpstream is a stand-in PostgreSQL backend: it reads one startup packet, -// records it, then echoes everything else back. -type fakeUpstream struct { - ln net.Listener - startup chan []byte -} + msg, err := fe.Receive() + if err != nil { + return err + } + if _, ok := msg.(*pgproto3.AuthenticationSASL); !ok { + return fmt.Errorf("expected AuthenticationSASL, got %T", msg) + } + client, err := scram.NewClient(password) + if err != nil { + return err + } + fe.Send(&pgproto3.SASLInitialResponse{AuthMechanism: scramMechanism, Data: []byte(client.FirstMessage())}) + if err := fe.Flush(); err != nil { + return err + } -func newFakeUpstream(t *testing.T) *fakeUpstream { - t.Helper() - ln, err := net.Listen("tcp", "127.0.0.1:0") + msg, err = fe.Receive() if err != nil { - t.Fatalf("listen: %v", err) + return err + } + cont, ok := msg.(*pgproto3.AuthenticationSASLContinue) + if !ok { + return fmt.Errorf("expected AuthenticationSASLContinue, got %T", msg) + } + final, err := client.HandleServerFirst(string(cont.Data)) + if err != nil { + return err + } + fe.Send(&pgproto3.SASLResponse{Data: []byte(final)}) + if err := fe.Flush(); err != nil { + return err } - fu := &fakeUpstream{ln: ln, startup: make(chan []byte, 4)} - go fu.accept() - t.Cleanup(func() { _ = ln.Close() }) - return fu -} -func (fu *fakeUpstream) addr() string { return fu.ln.Addr().String() } - -func (fu *fakeUpstream) accept() { - for { - c, err := fu.ln.Accept() - if err != nil { - return - } - go func(c net.Conn) { - defer func() { _ = c.Close() }() - pkt, err := readStartupPacket(c) - if err != nil { - return - } - select { - case fu.startup <- pkt.raw: - default: - } - _, _ = io.Copy(c, c) // echo the rest - }(c) + msg, err = fe.Receive() + if err != nil { + return err + } + fin, ok := msg.(*pgproto3.AuthenticationSASLFinal) + if !ok { + return fmt.Errorf("expected AuthenticationSASLFinal, got %T", msg) + } + if err := client.HandleServerFinal(string(fin.Data)); err != nil { + return err } -} -func startServer(t *testing.T, upstream string) (srv *Server, addr string, cancel context.CancelFunc, done chan struct{}) { - t.Helper() - srv = New(Config{ListenAddr: "127.0.0.1:0", UpstreamAddr: upstream, Logger: discardLogger()}) - a, err := srv.Listen() + msg, err = fe.Receive() if err != nil { - t.Fatalf("listen: %v", err) + return err } - ctx, c := context.WithCancel(context.Background()) - done = make(chan struct{}) - go func() { - _ = srv.Serve(ctx) - close(done) - }() - return srv, a.String(), c, done + if _, ok := msg.(*pgproto3.AuthenticationOk); !ok { + return fmt.Errorf("expected AuthenticationOk, got %T", msg) + } + return nil } -func TestServer_RefusesSSLForwardsStartupAndPipes(t *testing.T) { - fu := newFakeUpstream(t) - _, addr, cancel, _ := startServer(t, fu.addr()) - defer cancel() +func TestAuthenticateClient_Succeeds(t *testing.T) { + clientEnd, serverEnd := net.Pipe() + defer func() { _ = clientEnd.Close() }() + defer func() { _ = serverEnd.Close() }() + _ = clientEnd.SetDeadline(time.Now().Add(3 * time.Second)) + _ = serverEnd.SetDeadline(time.Now().Add(3 * time.Second)) - c, err := net.Dial("tcp", addr) - if err != nil { - t.Fatalf("dial proxy: %v", err) - } - defer func() { _ = c.Close() }() - _ = c.SetDeadline(time.Now().Add(3 * time.Second)) + sess := &session{client: serverEnd, log: discardLogger()} + srvErr := make(chan error, 1) + go func() { + srvErr <- sess.authenticateClient(pgproto3.NewBackend(serverEnd, serverEnd), "pencil") + }() - // SSLRequest is refused with 'N'. - if _, err := c.Write(buildStartup(sslRequestCode, nil)); err != nil { - t.Fatalf("write SSLRequest: %v", err) + if err := scramClientHandshake(clientEnd, "pencil"); err != nil { + t.Fatalf("client handshake: %v", err) } - buf := make([]byte, 1) - if _, err := io.ReadFull(c, buf); err != nil { - t.Fatalf("read refusal: %v", err) - } - if buf[0] != sslNotSupported { - t.Fatalf("refusal = %q, want %q", buf[0], byte(sslNotSupported)) + if err := <-srvErr; err != nil { + t.Fatalf("authenticateClient: %v", err) } +} - // The StartupMessage is forwarded to the upstream untouched. - startup := buildStartup(protocolVersion3, []byte("user\x00pgpilot\x00\x00")) - if _, err := c.Write(startup); err != nil { - t.Fatalf("write startup: %v", err) - } - select { - case got := <-fu.startup: - if !bytes.Equal(got, startup) { - t.Fatalf("upstream received %x, want %x", got, startup) - } - case <-time.After(3 * time.Second): - t.Fatal("upstream never received the startup packet") +func TestAuthenticateClient_WrongPassword(t *testing.T) { + clientEnd, serverEnd := net.Pipe() + defer func() { _ = clientEnd.Close() }() + defer func() { _ = serverEnd.Close() }() + _ = clientEnd.SetDeadline(time.Now().Add(3 * time.Second)) + _ = serverEnd.SetDeadline(time.Now().Add(3 * time.Second)) + + sess := &session{client: serverEnd, log: discardLogger()} + srvErr := make(chan error, 1) + go func() { + srvErr <- sess.authenticateClient(pgproto3.NewBackend(serverEnd, serverEnd), "correct") + }() + + _ = scramClientHandshake(clientEnd, "wrong") + if err := <-srvErr; err == nil { + t.Fatal("authenticateClient accepted a wrong password") } +} + +func TestRelayFrontend_InterceptsTerminate(t *testing.T) { + query := framedMessage('Q', []byte("SELECT 1\x00")) + term := framedMessage('X', nil) + src := bytes.NewReader(append(append([]byte{}, query...), term...)) - // A framed protocol message flows through the message-aware relay in both - // directions and comes back byte-for-byte (the fake upstream echoes). - msg := framedMessage('Q', []byte("SELECT 1\x00")) - if _, err := c.Write(msg); err != nil { - t.Fatalf("write message: %v", err) + var dst bytes.Buffer + terminated, err := relayFrontend(&dst, src) + if err != nil { + t.Fatalf("relayFrontend: %v", err) } - echo := make([]byte, len(msg)) - if _, err := io.ReadFull(c, echo); err != nil { - t.Fatalf("read echo: %v", err) + if !terminated { + t.Error("terminated = false, want true") } - if !bytes.Equal(echo, msg) { - t.Fatalf("echo = %x, want %x", echo, msg) + if !bytes.Equal(dst.Bytes(), query) { + t.Errorf("forwarded %x, want only the query %x (Terminate must not be forwarded)", dst.Bytes(), query) } } -func TestServer_GracefulShutdownClosesSessionsNoLeak(t *testing.T) { - fu := newFakeUpstream(t) - baseline := runtime.NumGoroutine() - srv, addr, cancel, done := startServer(t, fu.addr()) - _ = srv - - var clients []net.Conn - for i := 0; i < 3; i++ { - c, err := net.Dial("tcp", addr) - if err != nil { - t.Fatalf("dial %d: %v", i, err) - } - _ = c.SetDeadline(time.Now().Add(3 * time.Second)) - if _, err := c.Write(buildStartup(protocolVersion3, []byte("user\x00pgpilot\x00\x00"))); err != nil { - t.Fatalf("write startup %d: %v", i, err) - } - clients = append(clients, c) - } - time.Sleep(100 * time.Millisecond) // let sessions establish - - cancel() - select { - case <-done: - case <-time.After(3 * time.Second): - t.Fatal("Serve did not return after cancel") - } - - // Every client connection should have been closed by the proxy. - for i, c := range clients { - _ = c.SetReadDeadline(time.Now().Add(2 * time.Second)) - var b [1]byte - if _, err := c.Read(b[:]); err == nil { - t.Errorf("client %d still open after shutdown", i) - } - _ = c.Close() - } - - // Session goroutines should be gone (allow slack and a moment to settle). - if !eventually(2*time.Second, func() bool { - return runtime.NumGoroutine() <= baseline+2 - }) { - t.Errorf("goroutines did not settle: have %d, baseline %d", runtime.NumGoroutine(), baseline) +func TestRelayFrontend_CloseWithoutTerminate(t *testing.T) { + query := framedMessage('Q', []byte("x\x00")) + var dst bytes.Buffer + terminated, err := relayFrontend(&dst, bytes.NewReader(query)) + if err != nil { + t.Fatalf("relayFrontend: %v", err) + } + if terminated { + t.Error("terminated = true on EOF without Terminate, want false") + } + if !bytes.Equal(dst.Bytes(), query) { + t.Error("query was not forwarded") } } -func eventually(d time.Duration, cond func() bool) bool { - deadline := time.Now().Add(d) - for time.Now().Before(deadline) { - if cond() { - return true - } - runtime.GC() - time.Sleep(20 * time.Millisecond) +func TestParseStartupParams(t *testing.T) { + sm := &pgproto3.StartupMessage{ + ProtocolVersion: 196608, + Parameters: map[string]string{"user": "alice", "database": "shop"}, + } + wire, err := sm.Encode(nil) + if err != nil { + t.Fatal(err) + } + params, err := parseStartupParams(startupPacket{raw: wire, code: 196608}) + if err != nil { + t.Fatalf("parseStartupParams: %v", err) + } + if params["user"] != "alice" || params["database"] != "shop" { + t.Errorf("params = %v", params) } - return cond() } diff --git a/internal/proxy/session.go b/internal/proxy/session.go index 0e5bfb2..a03d57b 100644 --- a/internal/proxy/session.go +++ b/internal/proxy/session.go @@ -2,34 +2,45 @@ package proxy import ( "context" + "crypto/rand" + "encoding/binary" "errors" "fmt" "io" "log/slog" "net" + "os" "sync" + "time" + "github.com/jackc/pgx/v5/pgproto3" + + "github.com/sachhg/pgpilot/internal/backend" + "github.com/sachhg/pgpilot/internal/config" "github.com/sachhg/pgpilot/internal/protocol" + "github.com/sachhg/pgpilot/internal/scram" ) -// sslNotSupported is the single byte a server sends to refuse an SSLRequest or -// GSSENCRequest, telling the client to continue in cleartext. -const sslNotSupported = 'N' +const ( + // sslNotSupported is the single byte a server sends to refuse an SSLRequest + // or GSSENCRequest, telling the client to continue in cleartext. + sslNotSupported = 'N' + scramMechanism = "SCRAM-SHA-256" + resetTimeout = 5 * time.Second +) -// session proxies one client connection: it refuses encryption negotiation, -// forwards the client's startup packet to the upstream untouched, then relays -// protocol messages in both directions, tracking transaction status as it goes. +// session authenticates one client, acquires a pooled backend for its +// (user, database), and relays protocol messages between them, returning the +// backend to the pool when the client disconnects cleanly. type session struct { - client net.Conn - upstream string - dialer *net.Dialer - log *slog.Logger - tracker *protocol.TxTracker + client net.Conn + cfg *config.Config + manager *backend.Manager + log *slog.Logger + tracker *protocol.TxTracker } -// serve runs the session to completion. When ctx is cancelled (server -// shutdown), both connections are closed, which unblocks the relay, so no -// goroutine outlives serve. +// serve runs the session to completion. func (s *session) serve(ctx context.Context) error { defer func() { _ = s.client.Close() }() @@ -37,42 +48,205 @@ func (s *session) serve(ctx context.Context) error { if err != nil { return fmt.Errorf("startup negotiation: %w", err) } + if pkt.isCancelRequest() { + // Cancel requests cannot yet be mapped to a pooled backend; ignore. + s.log.Debug("ignoring unsupported cancel request") + return nil + } + + params, err := parseStartupParams(pkt) + if err != nil { + return s.reject("08P01", "invalid startup packet") + } + user := params["user"] + if user == "" { + return s.reject("08P01", "startup packet has no user") + } + database := params["database"] + if database == "" { + database = user + } + + u, ok := s.cfg.User(user) + if !ok { + return s.reject("28P01", fmt.Sprintf("password authentication failed for user %q", user)) + } + + clientBackend := pgproto3.NewBackend(s.client, s.client) + if err := s.authenticateClient(clientBackend, u.Password); err != nil { + return err + } + + be, err := s.manager.Acquire(ctx, user, database) + if err != nil { + return s.reject("53300", "could not acquire a backend connection") + } + reuse := false + defer s.releaseBackend(user, database, be, &reuse) + + if err := s.completeStartup(clientBackend, be); err != nil { + return fmt.Errorf("complete startup: %w", err) + } + + reuse = s.relayPooled(ctx, be) + return nil +} + +// authenticateClient runs the server side of the SCRAM-SHA-256 exchange against +// the client, sending AuthenticationOk on success. On failure it sends a FATAL +// ErrorResponse and returns an error. +func (s *session) authenticateClient(be *pgproto3.Backend, password string) error { + be.Send(&pgproto3.AuthenticationSASL{AuthMechanisms: []string{scramMechanism}}) + if err := be.Flush(); err != nil { + return fmt.Errorf("send AuthenticationSASL: %w", err) + } - up, err := s.dialer.DialContext(ctx, "tcp", s.upstream) + srv, err := scram.NewServer(password) if err != nil { - return fmt.Errorf("dial upstream %s: %w", s.upstream, err) + return err } - // Close both ends when ctx is cancelled, and stop the watcher - // deterministically when the session ends on its own. - ctx, cancel := context.WithCancel(ctx) + if err := be.SetAuthType(pgproto3.AuthTypeSASL); err != nil { + return err + } + msg, err := be.Receive() + if err != nil { + return fmt.Errorf("receive SASL initial response: %w", err) + } + initial, ok := msg.(*pgproto3.SASLInitialResponse) + if !ok { + return fmt.Errorf("expected SASLInitialResponse, got %T", msg) + } + if initial.AuthMechanism != scramMechanism { + return s.reject("28P01", "unsupported SASL mechanism") + } + serverFirst, err := srv.HandleClientFirst(string(initial.Data)) + if err != nil { + return s.reject("28P01", "SCRAM negotiation failed") + } + be.Send(&pgproto3.AuthenticationSASLContinue{Data: []byte(serverFirst)}) + if err := be.Flush(); err != nil { + return err + } + + if err := be.SetAuthType(pgproto3.AuthTypeSASLContinue); err != nil { + return err + } + msg, err = be.Receive() + if err != nil { + return fmt.Errorf("receive SASL response: %w", err) + } + resp, ok := msg.(*pgproto3.SASLResponse) + if !ok { + return fmt.Errorf("expected SASLResponse, got %T", msg) + } + serverFinal, err := srv.HandleClientFinal(string(resp.Data)) + if err != nil { + return s.reject("28P01", "password authentication failed") + } + be.Send(&pgproto3.AuthenticationSASLFinal{Data: []byte(serverFinal)}) + be.Send(&pgproto3.AuthenticationOk{}) + if err := be.Flush(); err != nil { + return err + } + return nil +} + +// completeStartup finishes the client's startup after authentication: it replays +// the backend's parameters, sends a synthesized BackendKeyData, and reports the +// session ready. +func (s *session) completeStartup(be *pgproto3.Backend, conn *backend.Conn) error { + for name, value := range conn.Params() { + be.Send(&pgproto3.ParameterStatus{Name: name, Value: value}) + } + pid, key, err := randomKeyData() + if err != nil { + return err + } + be.Send(&pgproto3.BackendKeyData{ProcessID: pid, SecretKey: key}) + be.Send(&pgproto3.ReadyForQuery{TxStatus: byte(protocol.StatusIdle)}) + return be.Flush() +} + +// relayPooled relays messages between the client and the backend until the +// client disconnects, and reports whether the backend is clean enough to reuse. +// The client's Terminate is intercepted, not forwarded, so the backend survives +// for the next client. +func (s *session) relayPooled(ctx context.Context, be *backend.Conn) (clean bool) { + backendConn := be.NetConn() + + // On server shutdown, close the client and interrupt the backend read so + // both relay directions unwind; on normal completion the watcher does + // nothing. + watchDone := make(chan struct{}) var watcher sync.WaitGroup watcher.Add(1) go func() { defer watcher.Done() - <-ctx.Done() - _ = s.client.Close() - _ = up.Close() + select { + case <-ctx.Done(): + _ = s.client.Close() + _ = backendConn.SetReadDeadline(time.Now()) + case <-watchDone: + } }() defer func() { - cancel() + close(watchDone) watcher.Wait() }() - // Forward the client's startup packet to the primary untouched. - if _, err := up.Write(pkt.raw); err != nil { - return fmt.Errorf("forward startup packet: %w", err) + beDone := make(chan error, 1) + go func() { + err := protocol.Relay(s.client, backendConn, s.trackBackend) + _ = s.client.Close() // unblock the frontend relay if the backend ends first + beDone <- err + }() + + terminated, feErr := relayFrontend(backendConn, s.client) + + select { + case <-beDone: + // The backend relay ended on its own (backend closed, errored, or we are + // shutting down): the connection is not reusable. + return false + default: + _ = backendConn.SetReadDeadline(time.Now()) // interrupt the idle backend read + <-beDone + _ = backendConn.SetReadDeadline(time.Time{}) } - if err := s.relay(s.client, up); err != nil { - return fmt.Errorf("relay: %w", err) + return terminated && feErr == nil && ctx.Err() == nil +} + +// releaseBackend returns the backend to its pool if it can be reset for reuse, +// discarding it otherwise. It is deferred with a pointer to the reuse decision. +func (s *session) releaseBackend(user, database string, be *backend.Conn, reuse *bool) { + if !*reuse { + s.manager.Discard(user, database, be) + return + } + ctx, cancel := context.WithTimeout(context.Background(), resetTimeout) + defer cancel() + if err := be.Reset(ctx); err != nil { + s.log.Debug("discarding backend that failed to reset", "error", err) + s.manager.Discard(user, database, be) + return + } + s.manager.Release(user, database, be) +} + +// trackBackend updates the session's transaction status from ReadyForQuery. +func (s *session) trackBackend(msgType byte, body []byte) error { + if msgType == protocol.MsgReadyForQuery { + if st, ok := protocol.ParseReadyForQuery(body); ok && s.tracker.Update(st) { + s.log.Debug("transaction status", "status", st.String()) + } } return nil } // negotiateStartup answers any SSLRequest/GSSENCRequest with a refusal until the -// client sends a real startup (or cancel) packet, which it returns unmodified -// for forwarding. +// client sends a real startup (or cancel) packet, which it returns. func (s *session) negotiateStartup() (startupPacket, error) { for { pkt, err := readStartupPacket(s.client) @@ -89,56 +263,60 @@ func (s *session) negotiateStartup() (startupPacket, error) { } } -// relay copies protocol messages in both directions until either side closes. -// Bytes are forwarded verbatim, so the relay stays transparent even for messages -// it does not interpret; backend messages are additionally decoded far enough to -// track transaction status. Each direction half-closes on EOF so both drain. -func (s *session) relay(client, up net.Conn) error { - var wg sync.WaitGroup - errc := make(chan error, 2) - wg.Add(2) - - go func() { // frontend: client -> upstream - defer wg.Done() - err := protocol.Relay(up, client, nil) - if hc, ok := up.(halfCloser); ok { - _ = hc.CloseWrite() +// reject sends a FATAL ErrorResponse to the client and returns an error. +func (s *session) reject(code, message string) error { + e := &pgproto3.ErrorResponse{Severity: "FATAL", SeverityUnlocalized: "FATAL", Code: code, Message: message} + if buf, err := e.Encode(nil); err == nil { + _, _ = s.client.Write(buf) + } + return fmt.Errorf("proxy: rejected client: %s %s", code, message) +} + +// relayFrontend forwards frontend messages from src to dst until the client +// sends Terminate (which is not forwarded, so the backend can be reused) or the +// connection ends. It reports whether the client terminated cleanly. +func relayFrontend(dst io.Writer, src io.Reader) (terminated bool, err error) { + var header [5]byte + for { + if _, err := io.ReadFull(src, header[:]); err != nil { + if errors.Is(err, io.EOF) || isClosedConnErr(err) || errors.Is(err, os.ErrDeadlineExceeded) { + return false, nil + } + return false, err } - errc <- err - }() - go func() { // backend: upstream -> client - defer wg.Done() - err := protocol.Relay(client, up, s.trackBackend) - if hc, ok := client.(halfCloser); ok { - _ = hc.CloseWrite() + if header[0] == protocol.MsgTerminate { + return true, nil } - errc <- err - }() - - wg.Wait() - close(errc) - for err := range errc { - if err != nil && !isClosedConnErr(err) && !errors.Is(err, io.ErrUnexpectedEOF) { - return err + length := binary.BigEndian.Uint32(header[1:5]) + if length < 4 { + return false, fmt.Errorf("proxy: frontend message length %d below minimum", length) + } + if _, err := dst.Write(header[:]); err != nil { + return false, err + } + if _, err := io.CopyN(dst, src, int64(length-4)); err != nil { + return false, err } } - return nil } -// trackBackend updates the session's transaction status from ReadyForQuery. -func (s *session) trackBackend(msgType byte, body []byte) error { - if msgType == protocol.MsgReadyForQuery { - if st, ok := protocol.ParseReadyForQuery(body); ok && s.tracker.Update(st) { - s.log.Debug("transaction status", "status", st.String()) - } +// parseStartupParams decodes the StartupMessage parameters (user, database, ...). +func parseStartupParams(pkt startupPacket) (map[string]string, error) { + var sm pgproto3.StartupMessage + if err := sm.Decode(pkt.raw[4:]); err != nil { + return nil, err } - return nil + return sm.Parameters, nil } -// halfCloser is implemented by *net.TCPConn: it lets one direction signal EOF -// without tearing down the other, so each side can drain fully. -type halfCloser interface { - CloseWrite() error +// randomKeyData generates a synthesized BackendKeyData for the client. pgpilot +// does not yet support cancellation, so these values are not mapped to a backend. +func randomKeyData() (pid, key uint32, err error) { + var b [8]byte + if _, err := rand.Read(b[:]); err != nil { + return 0, 0, err + } + return binary.BigEndian.Uint32(b[0:4]), binary.BigEndian.Uint32(b[4:8]), nil } // isClosedConnErr reports whether err is the "use of closed network connection" diff --git a/internal/proxy/startup.go b/internal/proxy/startup.go index 0af1b35..078683f 100644 --- a/internal/proxy/startup.go +++ b/internal/proxy/startup.go @@ -15,6 +15,7 @@ const ( // message a client sends, in place of a real protocol version. sslRequestCode = 80877103 // request a TLS-encrypted connection gssEncRequestCode = 80877104 // request a GSSAPI-encrypted connection + cancelRequestCode = 80877102 // request cancellation of an in-flight query ) // startupPacket is a raw, length-prefixed startup-phase message from a client. @@ -27,6 +28,7 @@ type startupPacket struct { func (p startupPacket) isSSLRequest() bool { return p.code == sslRequestCode } func (p startupPacket) isGSSEncRequest() bool { return p.code == gssEncRequestCode } +func (p startupPacket) isCancelRequest() bool { return p.code == cancelRequestCode } // readStartupPacket reads a single length-prefixed startup-phase message. The // four-byte length prefix counts itself; the four bytes after it are either a diff --git a/pgpilot.example.json b/pgpilot.example.json new file mode 100644 index 0000000..4bcb545 --- /dev/null +++ b/pgpilot.example.json @@ -0,0 +1,13 @@ +{ + "listen": "127.0.0.1:6432", + "primary": "127.0.0.1:55432", + "users": [ + {"name": "pgpilot", "password": "pgpilot"} + ], + "pool": { + "max_size": 10, + "max_waiters": 100, + "acquire_timeout": "5s", + "idle_timeout": "5m" + } +} diff --git a/test/integration/proxy_integration_test.go b/test/integration/proxy_integration_test.go index d79aaad..e2383ff 100644 --- a/test/integration/proxy_integration_test.go +++ b/test/integration/proxy_integration_test.go @@ -17,39 +17,45 @@ import ( "testing" "time" + "github.com/sachhg/pgpilot/internal/backend" + "github.com/sachhg/pgpilot/internal/config" "github.com/sachhg/pgpilot/internal/proxy" ) const ( - pgUser = "pgpilot" - pgDB = "pgpilot" - // primaryHostPort is where the dev-cluster primary is published on the host; - // the proxy dials it as its upstream. + pgUser = "pgpilot" + pgPass = "pgpilot" + pgDB = "pgpilot" primaryHostPort = "127.0.0.1:55432" ) -// TestProxy_PsqlThroughMatchesDirect asserts the Phase 2 invariant: psql routed -// through the transparent proxy behaves identically to psql connected directly -// to the primary. The proxy runs in-process on the host; psql runs inside the -// primary container and reaches the proxy back on the host via -// host.docker.internal (Docker Desktop). +// TestProxy_PsqlThroughMatchesDirect asserts that psql, authenticating to +// pgpilot with SCRAM-SHA-256 and served by a pooled backend, behaves identically +// to psql connected directly to the primary — validating the SCRAM server +// against real psql and the end-to-end session-pooling path. Requires `make up`. func TestProxy_PsqlThroughMatchesDirect(t *testing.T) { if _, err := exec.LookPath("docker"); err != nil { t.Skip("docker not on PATH; skipping proxy integration test") } compose := composeFile(t) - - // Precondition: the cluster must be up. if _, err := runPsql(compose, "127.0.0.1", 5432, "SELECT 1;"); err != nil { t.Fatalf("primary not reachable; run `make up` first: %v", err) } - // Start the transparent proxy in-process, forwarding to the primary. + cfg := &config.Config{ + Listen: "0.0.0.0:0", + Primary: primaryHostPort, + Users: []config.User{{Name: pgUser, Password: pgPass}}, + } + mgr := backend.NewManager(cfg.Primary, map[string]string{pgUser: pgPass}, + backend.PoolConfig{MaxSize: 5, AcquireTimeout: 5 * time.Second, IdleTimeout: time.Minute}) + defer mgr.Close() + srv := proxy.New(proxy.Config{ - ListenAddr: "0.0.0.0:0", - UpstreamAddr: primaryHostPort, - DialTimeout: 5 * time.Second, - Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), + ListenAddr: cfg.Listen, + Users: cfg, + Manager: mgr, + Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), }) addr, err := srv.Listen() if err != nil { @@ -69,7 +75,6 @@ func TestProxy_PsqlThroughMatchesDirect(t *testing.T) { t.Error("proxy Serve did not return after cancel") } }) - proxyPort := addr.(*net.TCPAddr).Port queries := []string{ @@ -78,8 +83,8 @@ func TestProxy_PsqlThroughMatchesDirect(t *testing.T) { "SELECT email FROM accounts ORDER BY id;", "SELECT current_database();", "SHOW server_version_num;", - "SELECT id, token FROM replication_probe ORDER BY id;", - "SELECT 1 AS a; SELECT 2 AS b;", // multi-statement simple query + "SELECT 1 AS a; SELECT 2 AS b;", + "BEGIN; SELECT count(*) FROM accounts; COMMIT;", } for _, q := range queries { direct, err := runPsql(compose, "127.0.0.1", 5432, q) @@ -96,13 +101,20 @@ func TestProxy_PsqlThroughMatchesDirect(t *testing.T) { } t.Logf("identical: %q -> %q", q, via) } + + // Several sequential sessions must reuse the pooled backend without error. + for i := 0; i < 3; i++ { + if _, err := runPsql(compose, "host.docker.internal", proxyPort, "SELECT 1;"); err != nil { + t.Fatalf("reuse session %d failed: %v", i, err) + } + } } // runPsql runs a query from inside the primary container against the given host // and port, returning psql's trimmed tuples-only output. func runPsql(compose, host string, port int, sql string) (string, error) { cmd := exec.Command("docker", "compose", "-f", compose, "exec", "-T", - "-e", "PGPASSWORD="+pgUser, "primary", + "-e", "PGPASSWORD="+pgPass, "primary", "psql", "-h", host, "-p", strconv.Itoa(port), "-U", pgUser, "-d", pgDB, "-tAX", "-c", sql) var stdout, stderr bytes.Buffer