diff --git a/README.md b/README.md index 5496d52..1646476 100644 --- a/README.md +++ b/README.md @@ -21,7 +21,9 @@ Prebuilt binaries are on the [releases page](https://github.com/mickamy/rollcall ## Status -Early development. `rollcall proxy` currently relays PostgreSQL connections unchanged; policy enforcement and the access ledger are being built on top of it. +Early development. `rollcall proxy` speaks the PostgreSQL protocol: it relays authentication untouched, sees every statement on both the simple and the extended query protocol, and can refuse one before it reaches the server. The policy that decides what to refuse and the access ledger are being built on top of it; today everything is allowed. + +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. ## Usage diff --git a/internal/cli/cli_test.go b/internal/cli/cli_test.go index 025e9fe..faee8a6 100644 --- a/internal/cli/cli_test.go +++ b/internal/cli/cli_test.go @@ -3,7 +3,10 @@ package cli_test import ( "bytes" "context" + "encoding/binary" "io" + "net" + "regexp" "strings" "sync" "testing" @@ -106,6 +109,28 @@ func TestRun(t *testing.T) { } } +func TestIsLoopback(t *testing.T) { + t.Parallel() + + tests := map[string]bool{ + "127.0.0.1:6432": true, + "[::1]:6432": true, + "localhost:6432": true, + "0.0.0.0:6432": false, + "[::]:6432": false, + ":6432": false, + "10.0.0.5:6432": false, + "db.internal:6432": false, + "garbage": false, + } + + for addr, want := range tests { + if got := cli.IsLoopback(addr); got != want { + t.Errorf("isLoopback(%q): got %v, want %v", addr, got, want) + } + } +} + func TestRunProxyStopsOnCancel(t *testing.T) { t.Parallel() @@ -139,6 +164,185 @@ func TestRunProxyStopsOnCancel(t *testing.T) { } } +func TestRunProxyServesPostgreSQL(t *testing.T) { + t.Parallel() + + upstream := startFakePostgres(t) + + 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"} + 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 (stderr: %q)", c, errOut.String()) + case <-time.After(5 * time.Second): + t.Fatal("proxy did not start listening") + } + + addr := regexp.MustCompile(`addr=(\S+)`).FindStringSubmatch(errOut.String()) + if addr == nil { + t.Fatalf("stderr: got %q, want the listening address", errOut.String()) + } + + var dialer net.Dialer + client, err := dialer.DialContext(ctx, "tcp", addr[1]) + if err != nil { + t.Fatalf("dial proxy: %v", err) + } + defer func() { _ = client.Close() }() + if err := client.SetDeadline(time.Now().Add(5 * time.Second)); err != nil { + t.Fatalf("set deadline: %v", err) + } + + writeAll(t, client, startupPacket("user", "alice", "database", "app")) + expectMessage(t, client, pgMessage('R', be32(0))) + expectMessage(t, client, pgMessage('Z', []byte("I"))) + + writeAll(t, client, pgMessage('Q', cstring("select 1"))) + expectMessage(t, client, pgMessage('C', cstring("SELECT 1"))) + expectMessage(t, client, pgMessage('Z', []byte("I"))) + + writeAll(t, client, pgMessage('X')) + if _, err := client.Read(make([]byte, 1)); err == nil { + t.Error("connection still open after Terminate") + } + + cancel() + + select { + case c := <-code: + if c != exit.OK { + t.Errorf("exit code: got %d, want %d (stderr: %q)", c, exit.OK, errOut.String()) + } + case <-time.After(5 * time.Second): + t.Fatal("proxy did not stop after cancel") + } + + for _, want := range []string{"msg=\"session opened\"", "user=alice", "database=app", "msg=\"session closed\""} { + if !strings.Contains(errOut.String(), want) { + t.Errorf("stderr: got %q, want it to contain %q", errOut.String(), want) + } + } +} + +// startFakePostgres serves trust authentication and answers every simple +// query with "SELECT 1" until the client terminates. +func startFakePostgres(t *testing.T) string { + t.Helper() + + var lc net.ListenConfig + ln, err := lc.Listen(t.Context(), "tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + t.Cleanup(func() { _ = ln.Close() }) + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + + go serveFakePostgres(conn) + } + }() + + return ln.Addr().String() +} + +func serveFakePostgres(conn net.Conn) { + defer func() { _ = conn.Close() }() + + length := make([]byte, 4) + if _, err := io.ReadFull(conn, length); err != nil { + return + } + if _, err := io.CopyN(io.Discard, conn, int64(binary.BigEndian.Uint32(length))-4); err != nil { + return + } + if _, err := conn.Write(bytes.Join([][]byte{pgMessage('R', be32(0)), pgMessage('Z', []byte("I"))}, nil)); err != nil { + return + } + + for { + header := make([]byte, 5) + if _, err := io.ReadFull(conn, header); err != nil { + return + } + if _, err := io.CopyN(io.Discard, conn, int64(binary.BigEndian.Uint32(header[1:]))-4); err != nil { + return + } + + switch header[0] { + case 'Q': + reply := bytes.Join([][]byte{pgMessage('C', cstring("SELECT 1")), pgMessage('Z', []byte("I"))}, nil) + if _, err := conn.Write(reply); err != nil { + return + } + case 'X': + return + } + } +} + +func be32(v uint32) []byte { + return binary.BigEndian.AppendUint32(nil, v) +} + +func cstring(s string) []byte { + return append([]byte(s), 0) +} + +func pgMessage(typ byte, parts ...[]byte) []byte { + body := bytes.Join(parts, nil) + out := []byte{typ} + out = binary.BigEndian.AppendUint32(out, uint32(len(body)+4)) //nolint:gosec // test payloads are tiny + + return append(out, body...) +} + +func startupPacket(params ...string) []byte { + var body []byte + body = binary.BigEndian.AppendUint32(body, 196608) + for _, p := range params { + body = append(body, cstring(p)...) + } + body = append(body, 0) + + out := binary.BigEndian.AppendUint32(nil, uint32(len(body)+4)) //nolint:gosec // test payloads are tiny + + return append(out, body...) +} + +func writeAll(t *testing.T, w io.Writer, b []byte) { + t.Helper() + + if _, err := w.Write(b); err != nil { + t.Fatalf("write: %v", err) + } +} + +func expectMessage(t *testing.T, r io.Reader, want []byte) { + t.Helper() + + got := make([]byte, len(want)) + if _, err := io.ReadFull(r, got); err != nil { + t.Fatalf("read: %v", err) + } + if !bytes.Equal(got, want) { + t.Fatalf("message: got %q, want %q", got, want) + } +} + // notifyWriter closes seen once the accumulated output contains want. type notifyWriter struct { mu sync.Mutex diff --git a/internal/cli/export_test.go b/internal/cli/export_test.go new file mode 100644 index 0000000..d3f0d5a --- /dev/null +++ b/internal/cli/export_test.go @@ -0,0 +1,3 @@ +package cli + +var IsLoopback = isLoopback diff --git a/internal/cli/proxy.go b/internal/cli/proxy.go index bca3001..afbc083 100644 --- a/internal/cli/proxy.go +++ b/internal/cli/proxy.go @@ -10,6 +10,7 @@ import ( "net" "github.com/mickamy/rollcall/internal/exit" + "github.com/mickamy/rollcall/internal/pg" "github.com/mickamy/rollcall/internal/proxy" ) @@ -52,8 +53,11 @@ func runProxy(ctx context.Context, args []string, std IO) int { 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) + } - srv := proxy.Server{Upstream: *upstream, Logger: logger} + srv := proxy.Server{Upstream: *upstream, Dialect: pg.Dialect{}, Logger: logger} if err := srv.Serve(ctx, ln); err != nil { return fail(std, err) } @@ -77,9 +81,25 @@ func validateAddr(flagName, addr string) error { return nil } +// isLoopback reports whether addr can only be reached from this host. +func isLoopback(addr string) bool { + host, _, err := net.SplitHostPort(addr) + if err != nil || host == "" { + return false + } + if host == "localhost" { + return true + } + + ip := net.ParseIP(host) + + return ip != nil && ip.IsLoopback() +} + func printProxyUsage(w io.Writer) { fmt.Fprintf(w, "Usage: %s proxy -upstream ADDR [-listen ADDR]\n\n", Name) - fmt.Fprint(w, "Accept database connections and relay them to the upstream database.\n") + fmt.Fprint(w, "Accept PostgreSQL connections and relay them to the upstream database.\n") + fmt.Fprint(w, "Both sides are plaintext; keep the listener on loopback or a pod-local network.\n") fmt.Fprint(w, "Stops when interrupted.\n\n") fmt.Fprint(w, "Flags:\n") fmt.Fprintf(w, " -upstream ADDR %s (required)\n", upstreamUsage) diff --git a/internal/pg/message.go b/internal/pg/message.go new file mode 100644 index 0000000..390b68d --- /dev/null +++ b/internal/pg/message.go @@ -0,0 +1,241 @@ +package pg + +import ( + "bufio" + "bytes" + "encoding/binary" + "errors" + "fmt" + "io" +) + +const ( + protocolVersion3 = 196608 + sslRequestCode = 80877103 + gssEncRequestCode = 80877104 + cancelRequestCode = 80877102 + + // maxStartupPacket matches PostgreSQL's own MaxStartupPacketLength. + maxStartupPacket = 10000 + maxAuthMessage = 1 << 20 + + // headerLen is the message type byte plus the length word. + headerLen = 5 + // lengthLen is the length word, counted inside the length itself. + lengthLen = 4 + + maxBodyLen = 1<<31 - 1 - lengthLen +) + +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' +) + +const ( + authOK = 0 + authSASLFinal = 12 +) + +const ( + fieldSeverity = 'S' + fieldSeverityText = 'V' + fieldCode = 'C' + fieldMessage = 'M' + fieldHint = 'H' + + sqlStateInsufficientPrivilege = "42501" + sqlStateProgramLimitExceeded = "54000" + sqlStateFeatureNotSupported = "0A000" +) + +var errMalformed = errors.New("malformed message") + +// readStartupPacket returns the whole untyped startup packet, length word included. +func readStartupPacket(r *bufio.Reader) ([]byte, error) { + var length [lengthLen]byte + if _, err := io.ReadFull(r, length[:]); err != nil { + return nil, fmt.Errorf("read startup length: %w", err) + } + + n := binary.BigEndian.Uint32(length[:]) + if n < 2*lengthLen || n > maxStartupPacket { + return nil, fmt.Errorf("%w: startup packet length %d", errMalformed, n) + } + + packet := make([]byte, n) + copy(packet, length[:]) + if _, err := io.ReadFull(r, packet[lengthLen:]); err != nil { + return nil, fmt.Errorf("read startup packet: %w", err) + } + + return packet, nil +} + +// readHeader returns a typed message's type and body length, leaving the body unread. +func readHeader(r *bufio.Reader) (typ byte, bodyLen uint32, err error) { + var header [headerLen]byte + if _, err := io.ReadFull(r, header[:]); err != nil { + return 0, 0, fmt.Errorf("read header: %w", err) + } + + n := binary.BigEndian.Uint32(header[1:]) + if n < lengthLen { + return 0, 0, fmt.Errorf("%w: %q with length %d", errMalformed, header[0], n) + } + + return header[0], n - lengthLen, nil +} + +func readMessage(r *bufio.Reader, maxBody uint32) (typ byte, body []byte, err error) { + typ, n, err := readHeader(r) + if err != nil { + return 0, nil, err + } + if n > maxBody { + return 0, nil, fmt.Errorf("%w: %q body of %d bytes exceeds %d", errMalformed, typ, n, maxBody) + } + + body = make([]byte, n) + if _, err := io.ReadFull(r, body); err != nil { + return 0, nil, fmt.Errorf("read %q body: %w", typ, err) + } + + return typ, body, nil +} + +// readBody reads a message body whose length the client chose, committing +// memory only as bytes actually arrive. +func readBody(r *bufio.Reader, n uint32) ([]byte, error) { + var buf bytes.Buffer + buf.Grow(int(min(n, 64<<10))) + if _, err := io.CopyN(&buf, r, int64(n)); err != nil { + return nil, fmt.Errorf("read body: %w", err) + } + + return buf.Bytes(), nil +} + +func discard(r *bufio.Reader, n uint32) error { + if _, err := io.CopyN(io.Discard, r, int64(n)); err != nil { + return fmt.Errorf("discard body: %w", err) + } + + return nil +} + +func writeMessage(w *bufio.Writer, typ byte, body []byte) error { + if len(body) > maxBodyLen { + return fmt.Errorf("%w: %q body of %d bytes", errMalformed, typ, len(body)) + } + + if err := writeHeader(w, typ, uint32(len(body))); err != nil { //nolint:gosec // bounded by the maxBodyLen check above + return err + } + if _, err := w.Write(body); err != nil { + return fmt.Errorf("write %q: %w", typ, err) + } + + return nil +} + +// stageHeader writes a message header into a growable buffer, for batching +// several messages before they are forwarded together. +func stageHeader(buf *bytes.Buffer, typ byte, bodyLen int) { + buf.WriteByte(typ) + var length [lengthLen]byte + binary.BigEndian.PutUint32(length[:], uint32(bodyLen+lengthLen)) //nolint:gosec // bodyLen is bounded by the caller + buf.Write(length[:]) +} + +func writeHeader(w *bufio.Writer, typ byte, bodyLen uint32) error { + var header [headerLen]byte + header[0] = typ + binary.BigEndian.PutUint32(header[1:], bodyLen+lengthLen) + + if _, err := w.Write(header[:]); err != nil { + return fmt.Errorf("write %q header: %w", typ, err) + } + + return nil +} + +// forward streams one message whose header has already been read from r to w +// without flushing. +func forward(w *bufio.Writer, typ byte, bodyLen uint32, r io.Reader) error { + if err := writeHeader(w, typ, bodyLen); err != nil { + return err + } + if _, err := io.CopyN(w, r, int64(bodyLen)); err != nil { + return fmt.Errorf("forward %q body: %w", typ, err) + } + + return nil +} + +func flush(w *bufio.Writer) error { + if err := w.Flush(); err != nil { + return fmt.Errorf("flush: %w", err) + } + + return nil +} + +func cstring(b []byte) (s string, rest []byte, err error) { + value, rest, ok := bytes.Cut(b, []byte{0}) + if !ok { + return "", nil, fmt.Errorf("%w: unterminated string", errMalformed) + } + + return string(value), rest, nil +} + +// errorMessage extracts the human-readable message field from an ErrorResponse body. +func errorMessage(body []byte) string { + for len(body) > 0 && body[0] != 0 { + field := body[0] + value, rest, err := cstring(body[1:]) + if err != nil { + return "" + } + if field == fieldMessage { + return value + } + body = rest + } + + return "" +} + +func errorResponse(code, message, hint string) []byte { + var b bytes.Buffer + writeField(&b, fieldSeverity, "ERROR") + writeField(&b, fieldSeverityText, "ERROR") + writeField(&b, fieldCode, code) + writeField(&b, fieldMessage, message) + if hint != "" { + writeField(&b, fieldHint, hint) + } + b.WriteByte(0) + + return b.Bytes() +} + +func writeField(b *bytes.Buffer, field byte, value string) { + b.WriteByte(field) + b.WriteString(value) + b.WriteByte(0) +} diff --git a/internal/pg/pg.go b/internal/pg/pg.go new file mode 100644 index 0000000..f1a7bee --- /dev/null +++ b/internal/pg/pg.go @@ -0,0 +1,751 @@ +// Package pg implements wire.Dialect for the PostgreSQL frontend/backend protocol, version 3. +package pg + +import ( + "bufio" + "bytes" + "encoding/binary" + "errors" + "fmt" + "io" + "net" + "sync" + + "github.com/mickamy/rollcall/internal/wire" +) + +const ( + defaultMaxStatement = 16 << 20 + + readBufferSize = 64 << 10 + writeBufferSize = 32 << 10 + + // maxPending bounds the extended-protocol messages held before a Sync so a + // batch can be rejected atomically; a batch larger than this is forwarded + // early (see stage) rather than buffered without limit. + maxPending = 8 << 20 + // maxQueue bounds outstanding responses. Reaching it blocks the frontend + // until the backend drains one, applying backpressure like the server does. + maxQueue = 512 + // smallBody bounds the messages Backend reads fully before taking the + // client write lock, so a stalled upstream cannot hold denials hostage. + smallBody = 32 << 10 +) + +// errDeniedAfterForward ends a session when a statement is denied after part of +// its batch already reached the upstream. Closing the connection makes the +// upstream roll back the implicit transaction instead of committing at Sync. +var errDeniedAfterForward = errors.New("statement denied after its batch was partially forwarded") + +type Dialect struct { + // MaxStatement caps the SQL text of a Query or Parse message; longer + // statements are denied without being forwarded. Zero means 16 MiB. + MaxStatement uint32 +} + +var _ wire.Dialect = (*Dialect)(nil) + +func (d Dialect) NewSession(client, upstream net.Conn) wire.Session { + maxStatement := d.MaxStatement + if maxStatement == 0 { + maxStatement = defaultMaxStatement + } + + s := &session{ + maxStatement: maxStatement, + cr: bufio.NewReaderSize(client, readBufferSize), + cw: bufio.NewWriterSize(client, writeBufferSize), + ur: bufio.NewReaderSize(upstream, readBufferSize), + uw: bufio.NewWriterSize(upstream, writeBufferSize), + small: make([]byte, smallBody), + } + s.drained = sync.NewCond(&s.mu) + + return s +} + +// slot is one client request awaiting its answer, in request order. A +// forwarded slot is settled by the upstream's ReadyForQuery; simple marks the +// simple query protocol, whose ReadyForQuery survives COPY. A synthesized slot +// (not forwarded) is written by the proxy once every slot ahead of it is +// answered. +type slot struct { + forwarded bool + simple bool + code string + message string + hint string + ready bool +} + +type session struct { + maxStatement uint32 + + // Frontend goroutine only. + cr *bufio.Reader + uw *bufio.Writer + pending bytes.Buffer + forwarded bool // part of the current batch was already sent upstream + denied bool // a statement in the current batch was denied + + // Backend goroutine only. + ur *bufio.Reader + small []byte + + // 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 +} + +func (s *session) Handshake() (wire.Startup, error) { + for { + packet, err := readStartupPacket(s.cr) + if err != nil { + return wire.Startup{}, err + } + + switch code := binary.BigEndian.Uint32(packet[lengthLen:]); code { + case sslRequestCode, gssEncRequestCode: + // The proxy speaks plaintext to its clients; refusing makes libpq fall back. + if err := s.writeClientRaw([]byte{'N'}); err != nil { + return wire.Startup{}, err + } + case cancelRequestCode: + if err := s.writeUpstreamRaw(packet); err != nil { + return wire.Startup{}, err + } + + return wire.Startup{}, wire.ErrNoSession + case protocolVersion3: + startup, err := parseStartup(packet[2*lengthLen:]) + if err != nil { + return wire.Startup{}, err + } + if err := s.writeUpstreamRaw(packet); err != nil { + return startup, err + } + + return startup, s.relayAuth() + default: + return wire.Startup{}, fmt.Errorf("unsupported protocol %d.%d", code>>16, code&0xffff) + } + } +} + +// 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 { + for { + if s.cr.Buffered() == 0 { + if err := flush(s.uw); err != nil { + return err + } + } + + typ, n, err := readHeader(s.cr) + if err != nil { + if errors.Is(err, io.EOF) { + return nil + } + + return fmt.Errorf("read client: %w", err) + } + + done, err := s.dispatch(h, typ, n) + if err != nil { + return err + } + if done { + return flush(s.uw) + } + } +} + +// Backend relays upstream messages and settles the request queue on every +// ReadyForQuery, writing any denials that were waiting their turn. +func (s *session) Backend() error { + for { + if s.ur.Buffered() == 0 { + if err := s.flushClient(); err != nil { + return err + } + } + + typ, n, err := readHeader(s.ur) + if err != nil { + if errors.Is(err, io.EOF) { + return nil + } + + return fmt.Errorf("read upstream: %w", err) + } + + switch typ { + case typeReadyForQuery: + err = s.readyForQuery(n) + case typeCopyInResponse: + // The server ignores the batch's Sync while reading COPY data, so + // the extended-protocol slot for it will never be answered; drop it. + if err = s.forwardToClient(typ, n); err == nil { + s.copyInStarted() + } + default: + err = s.forwardToClient(typ, n) + } + if err != nil { + return err + } + } +} + +// relayAuth passes the authentication exchange through untouched until the +// upstream reports ReadyForQuery or refuses the connection. +func (s *session) relayAuth() error { + for { + typ, body, err := readMessage(s.ur, maxAuthMessage) + if err != nil { + return fmt.Errorf("read upstream during handshake: %w", err) + } + if err := s.writeClient(typ, body); err != nil { + return err + } + + switch typ { + case typeReadyForQuery: + if len(body) != 1 { + return fmt.Errorf("%w: ReadyForQuery with %d byte body", errMalformed, len(body)) + } + s.setTx(body[0]) + + return nil + case typeErrorResponse: + return fmt.Errorf("%w: %s", wire.ErrRejected, errorMessage(body)) + case typeAuthentication: + if len(body) < lengthLen { + return fmt.Errorf("%w: Authentication with %d byte body", errMalformed, len(body)) + } + if !expectsResponse(binary.BigEndian.Uint32(body)) { + continue + } + + typ, body, err := readMessage(s.cr, maxAuthMessage) + if err != nil { + return fmt.Errorf("read client during handshake: %w", err) + } + if err := writeMessage(s.uw, typ, body); err != nil { + return err + } + if err := flush(s.uw); err != nil { + return err + } + } + } +} + +// dispatch routes one client message. It reports whether the session is done. +func (s *session) dispatch(h wire.Handler, typ byte, n uint32) (bool, error) { + switch typ { + case typeQuery: + return false, s.query(h, n) + case typeParse: + return false, s.parse(h, n) + case typeBind, typeDescribe, typeExecute, typeClose: + return false, s.buffer(typ, n) + case typeFlush: + return false, s.flushBatch(n) + case typeSync: + return false, s.sync(n) + case typeFunctionCall: + return false, s.functionCall(n) + case typeTerminate: + if err := s.forwardPending(); err != nil { + return false, err + } + + return true, forward(s.uw, typ, n, s.cr) + default: + // Copy sub-protocol (CopyData/CopyDone/CopyFail) and anything else that + // carries no SQL: forward immediately, after any buffered batch. + if err := s.forwardPending(); err != nil { + return false, err + } + + return false, forward(s.uw, typ, n, s.cr) + } +} + +func (s *session) query(h wire.Handler, n uint32) error { + if err := s.forwardPending(); err != nil { + return err + } + + if n > s.maxStatement { + if err := discard(s.cr, n); err != nil { + return err + } + + return s.respond(readySlot(s.tooLarge(n))) + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + + sql := string(bytes.TrimSuffix(body, []byte{0})) + if v := h.Statement(wire.Statement{SQL: sql}); v.Deny { + 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}) + + return writeMessage(s.uw, typeQuery, body) +} + +// parse handles the extended protocol's Parse message, the one place where SQL +// enters that path. A denial rejects the whole batch: earlier buffered messages +// are dropped and later ones ignored until Sync. +func (s *session) parse(h wire.Handler, n uint32) error { + if s.denied { + return discard(s.cr, n) + } + if n > s.maxStatement { + if err := discard(s.cr, n); err != nil { + return err + } + + return s.denyBatch(s.tooLarge(n)) + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + + sql, err := parseSQL(body) + if err != nil { + return err + } + + if v := h.Statement(wire.Statement{SQL: sql}); v.Deny { + return s.denyBatch(denial(v)) + } + + return s.stage(typeParse, body) +} + +func (s *session) buffer(typ byte, 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, typ, n, s.cr) + } + + body, err := readBody(s.cr, n) + if err != nil { + return err + } + + return s.stage(typ, body) +} + +// stage buffers one message until the batch's Sync, forwarding what is already +// buffered when the buffer would grow past maxPending. +func (s *session) stage(typ byte, body []byte) error { + if s.forwarded { + return writeMessage(s.uw, typ, body) + } + if s.pending.Len() > 0 && s.pending.Len()+headerLen+len(body) > maxPending { + if err := s.forwardPending(); err != nil { + return err + } + s.forwarded = true + + return writeMessage(s.uw, typ, body) + } + + stageHeader(&s.pending, typ, len(body)) + s.pending.Write(body) + + return nil +} + +func (s *session) flushBatch(n uint32) error { + if err := discard(s.cr, n); err != nil { + return err + } + if s.denied { + return nil + } + if err := s.forwardPending(); err != nil { + return err + } + s.forwarded = true + if err := writeMessage(s.uw, typeFlush, nil); err != nil { + return err + } + + return flush(s.uw) +} + +// sync ends a batch. A clean batch is forwarded together with its Sync; a +// denied batch (which never forwarded anything) sends only a Sync so the +// client's ReadyForQuery still comes from a real upstream response. +func (s *session) sync(n uint32) error { + if err := discard(s.cr, n); err != nil { + return err + } + + if !s.denied { + if err := s.forwardPending(); err != nil { + return err + } + } + + 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}) + if err := writeMessage(s.uw, typeSync, nil); err != nil { + return err + } + + return flush(s.uw) +} + +func (s *session) functionCall(n uint32) error { + if err := discard(s.cr, n); err != nil { + return err + } + if s.denied { + return nil + } + if err := s.forwardPending(); err != nil { + return err + } + + return s.respond(readySlot(slot{ + code: sqlStateFeatureNotSupported, + message: "the function call protocol is not supported through the proxy", + })) +} + +// denyBatch rejects the current batch. When nothing has been forwarded it is +// rejected atomically. When part of it already reached the upstream, the +// 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.resetPending() + + if s.forwarded { + s.emitError(sl) + + return errDeniedAfterForward + } + + return s.respond(sl) +} + +func (s *session) forwardPending() error { + if s.pending.Len() == 0 { + return nil + } + if _, err := s.uw.Write(s.pending.Bytes()); err != nil { + return fmt.Errorf("forward buffered batch: %w", err) + } + s.resetPending() + + return nil +} + +// resetPending clears the buffer, dropping an oversized backing array so a +// connection that saw one large batch does not hold that memory for its life. +func (s *session) resetPending() { + if s.pending.Len() > maxPending { + s.pending = bytes.Buffer{} + + return + } + s.pending.Reset() +} + +// respond enqueues a synthesized answer and writes it, and any answers ahead of +// it that are now unblocked, when nothing forwarded is still pending. +func (s *session) respond(sl slot) error { + s.mu.Lock() + defer s.mu.Unlock() + + s.enqueueLocked(sl) + + emitted := false + for len(s.queue) > 0 && !s.queue[0].forwarded { + if err := s.emit(s.queue[0]); err != nil { + return err + } + s.queue = s.queue[1:] + emitted = true + } + if emitted { + s.drained.Broadcast() + + return flush(s.cw) + } + + return nil +} + +func (s *session) enqueue(sl slot) { + s.mu.Lock() + defer s.mu.Unlock() + + s.enqueueLocked(sl) +} + +// enqueueLocked appends a slot, blocking while the queue is full so the client +// is backpressured instead of dropped. Callers hold mu. +func (s *session) enqueueLocked(sl slot) { + for len(s.queue) >= maxQueue { + s.drained.Wait() + } + s.queue = append(s.queue, sl) +} + +// emit writes a synthesized answer. Callers hold mu. +func (s *session) emit(sl slot) error { + if sl.code != "" { + if err := writeMessage(s.cw, typeErrorResponse, errorResponse(sl.code, sl.message, sl.hint)); err != nil { + return err + } + } + if sl.ready { + if err := writeMessage(s.cw, typeReadyForQuery, []byte{s.tx}); err != nil { + return err + } + } + + return nil +} + +// emitError writes a denial straight to the client, best effort, for the +// teardown path where the queue is about to be abandoned. +func (s *session) emitError(sl slot) { + s.mu.Lock() + defer s.mu.Unlock() + + _ = writeMessage(s.cw, typeErrorResponse, errorResponse(sl.code, sl.message, sl.hint)) + _ = flush(s.cw) +} + +func (s *session) readyForQuery(n uint32) error { + if n != 1 { + return fmt.Errorf("%w: ReadyForQuery with %d byte body", errMalformed, n) + } + + var status [1]byte + if _, err := io.ReadFull(s.ur, status[:]); err != nil { + return fmt.Errorf("read ReadyForQuery: %w", err) + } + + s.mu.Lock() + defer s.mu.Unlock() + + for len(s.queue) > 0 && !s.queue[0].forwarded { + if err := s.emit(s.queue[0]); err != nil { + return err + } + s.queue = s.queue[1:] + } + if len(s.queue) == 0 || !s.queue[0].forwarded { + return fmt.Errorf("%w: ReadyForQuery without a pending request", errMalformed) + } + + s.queue = s.queue[1:] + s.tx = status[0] + if err := writeMessage(s.cw, typeReadyForQuery, status[:]); err != nil { + return err + } + + for len(s.queue) > 0 && !s.queue[0].forwarded { + if err := s.emit(s.queue[0]); err != nil { + return err + } + s.queue = s.queue[1:] + } + + s.drained.Broadcast() + + return flush(s.cw) +} + +// copyInStarted drops the extended-protocol Sync slot for a COPY, whose Sync the +// server ignores while reading copy data. A simple query keeps its slot, since +// its ReadyForQuery still arrives when the copy completes. +func (s *session) copyInStarted() { + s.mu.Lock() + defer s.mu.Unlock() + + if len(s.queue) > 0 && s.queue[0].forwarded && !s.queue[0].simple { + s.queue = s.queue[1:] + s.drained.Broadcast() + } +} + +func (s *session) forwardToClient(typ byte, n uint32) error { + if n > smallBody { + s.mu.Lock() + defer s.mu.Unlock() + + return forward(s.cw, typ, n, s.ur) + } + + body := s.small[:n] + if _, err := io.ReadFull(s.ur, body); err != nil { + return fmt.Errorf("read %q body: %w", typ, err) + } + + s.mu.Lock() + defer s.mu.Unlock() + + return writeMessage(s.cw, typ, body) +} + +func (s *session) flushClient() error { + s.mu.Lock() + defer s.mu.Unlock() + + return flush(s.cw) +} + +func (s *session) setTx(status byte) { + s.mu.Lock() + defer s.mu.Unlock() + + s.tx = status +} + +func (s *session) writeClient(typ byte, body []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + + if err := writeMessage(s.cw, typ, body); err != nil { + return err + } + + return flush(s.cw) +} + +func (s *session) writeClientRaw(b []byte) error { + s.mu.Lock() + defer s.mu.Unlock() + + if _, err := s.cw.Write(b); err != nil { + return fmt.Errorf("write client: %w", err) + } + + return flush(s.cw) +} + +func (s *session) writeUpstreamRaw(b []byte) error { + if _, err := s.uw.Write(b); err != nil { + return fmt.Errorf("write upstream: %w", err) + } + + return flush(s.uw) +} + +func expectsResponse(authCode uint32) bool { + switch authCode { + case authOK, authSASLFinal: + return false + default: + return true + } +} + +func denial(v wire.Verdict) slot { + return slot{code: sqlStateInsufficientPrivilege, message: v.Message, hint: v.Hint} +} + +func (s *session) tooLarge(n uint32) slot { + return slot{ + code: sqlStateProgramLimitExceeded, + message: fmt.Sprintf("statement of %d bytes exceeds the %d byte limit", n, s.maxStatement), + } +} + +func readySlot(sl slot) slot { + sl.ready = true + + 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) + 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 sql, nil +} + +func parseStartup(body []byte) (wire.Startup, error) { + params := make(map[string]string) + for len(body) > 0 { + key, rest, err := cstring(body) + if err != nil { + return wire.Startup{}, fmt.Errorf("startup packet: %w", err) + } + if key == "" { + break + } + + value, rest, err := cstring(rest) + if err != nil { + return wire.Startup{}, fmt.Errorf("startup packet: %w", err) + } + + params[key] = value + body = rest + } + + user := params["user"] + if user == "" { + return wire.Startup{}, fmt.Errorf("%w: startup packet without user", errMalformed) + } + + database := params["database"] + if database == "" { + database = user + } + + return wire.Startup{ + User: user, + Database: database, + Application: params["application_name"], + Params: params, + }, nil +} diff --git a/internal/pg/pg_test.go b/internal/pg/pg_test.go new file mode 100644 index 0000000..4418213 --- /dev/null +++ b/internal/pg/pg_test.go @@ -0,0 +1,991 @@ +package pg_test + +import ( + "bytes" + "encoding/binary" + "errors" + "io" + "net" + "os" + "strings" + "testing" + "time" + + "github.com/mickamy/rollcall/internal/pg" + "github.com/mickamy/rollcall/internal/wire" +) + +const ( + timeout = 5 * time.Second + protocolVersion3 = 196608 + cancelRequest = 80877102 + sslRequest = 80877103 +) + +func TestHandshakeRelaysPasswordAuthentication(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + result := p.handshake() + startup := startupPacket("user", "alice", "database", "app", "application_name", "psql") + + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, startup); err != nil { + return err + } + if err := write(s, auth(3)); err != nil { + return err + } + if err := expectBytes(s, msg('p', cstr("secret"))); err != nil { + return err + } + + return write(s, + auth(0), + msg('S', cstr("server_version"), cstr("16")), + msg('K', be32(7), be32(9)), + msg('Z', []byte("I")), + ) + }) + + mustWrite(t, p.client, packet(sslRequest)) + if got := readN(t, p.client, 1); got[0] != 'N' { + t.Fatalf("SSLRequest answer: got %q, want 'N'", got) + } + mustWrite(t, p.client, startup) + + expectMsg(t, p.client, auth(3)) + mustWrite(t, p.client, msg('p', cstr("secret"))) + expectMsg(t, p.client, auth(0)) + expectMsg(t, p.client, msg('S', cstr("server_version"), cstr("16"))) + expectMsg(t, p.client, msg('K', be32(7), be32(9))) + expectMsg(t, p.client, msg('Z', []byte("I"))) + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + got := <-result + if got.err != nil { + t.Fatalf("Handshake: %v", got.err) + } + want := wire.Startup{User: "alice", Database: "app", Application: "psql"} + if got.startup.Params["application_name"] != "psql" { + t.Errorf("Params: got %v, want application_name=psql", got.startup.Params) + } + gotIdentity := [3]string{got.startup.User, got.startup.Database, got.startup.Application} + wantIdentity := [3]string{want.User, want.Database, want.Application} + if gotIdentity != wantIdentity { + t.Errorf("Handshake: got %v, want %v", gotIdentity, wantIdentity) + } +} + +func TestHandshakeRelaysSASLWithoutExpectingAReplyToTheFinalMessage(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + result := p.handshake() + + server := p.serve(func(s net.Conn) error { + if _, err := readStartup(s); err != nil { + return err + } + if err := write(s, auth(10, cstr("SCRAM-SHA-256"), []byte{0})); err != nil { + return err + } + if err := expectBytes(s, msg('p', cstr("SCRAM-SHA-256"), be32(3), []byte("n,,"))); err != nil { + return err + } + if err := write(s, auth(11, []byte("r=nonce"))); err != nil { + return err + } + if err := expectBytes(s, msg('p', []byte("c=biws"))); err != nil { + return err + } + + return write(s, auth(12, []byte("v=proof")), auth(0), msg('Z', []byte("I"))) + }) + + mustWrite(t, p.client, startupPacket("user", "alice")) + expectMsg(t, p.client, auth(10, cstr("SCRAM-SHA-256"), []byte{0})) + mustWrite(t, p.client, msg('p', cstr("SCRAM-SHA-256"), be32(3), []byte("n,,"))) + expectMsg(t, p.client, auth(11, []byte("r=nonce"))) + mustWrite(t, p.client, msg('p', []byte("c=biws"))) + expectMsg(t, p.client, auth(12, []byte("v=proof"))) + expectMsg(t, p.client, auth(0)) + expectMsg(t, p.client, msg('Z', []byte("I"))) + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + if got := <-result; got.err != nil { + t.Fatalf("Handshake: %v", got.err) + } +} + +func TestHandshakeDefaultsDatabaseToUser(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + got := p.connect(t, "user", "alice") + + if got.Database != "alice" { + t.Errorf("Database: got %q, want %q", got.Database, "alice") + } +} + +func TestHandshakeForwardsCancelRequest(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + result := p.handshake() + cancel := packet(cancelRequest, be32(1234), be32(5678)) + + server := p.serve(func(s net.Conn) error { + return expectBytes(s, cancel) + }) + + mustWrite(t, p.client, cancel) + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + if got := <-result; !errors.Is(got.err, wire.ErrNoSession) { + t.Errorf("Handshake: got %v, want %v", got.err, wire.ErrNoSession) + } +} + +func TestHandshakeReportsRejection(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + result := p.handshake() + rejection := msg('E', + []byte("SFATAL\x00"), []byte("C28P01\x00"), []byte("Mpassword authentication failed\x00"), []byte{0}, + ) + + server := p.serve(func(s net.Conn) error { + if _, err := readStartup(s); err != nil { + return err + } + + return write(s, rejection) + }) + + mustWrite(t, p.client, startupPacket("user", "alice")) + expectMsg(t, p.client, rejection) + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + got := <-result + if !errors.Is(got.err, wire.ErrRejected) { + t.Errorf("Handshake: got %v, want %v", got.err, wire.ErrRejected) + } + if !strings.Contains(got.err.Error(), "password authentication failed") { + t.Errorf("Handshake: got %v, want it to carry the upstream message", got.err) + } + if got.startup.User != "alice" { + t.Errorf("Handshake: got user %q with the rejection, want %q", got.startup.User, "alice") + } +} + +func TestHandshakeRejectsUnknownProtocol(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + result := p.handshake() + + mustWrite(t, p.client, packet(131072, cstr("user"), cstr("alice"), []byte{0})) + + if got := <-result; got.err == nil || !strings.Contains(got.err.Error(), "unsupported protocol 2.0") { + t.Errorf("Handshake: got %v, want an unsupported protocol error", got.err) + } +} + +func TestFrontendForwardsAllowedQueriesAndTerminate(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + + var seen []string + done := p.frontend(wire.HandlerFunc(func(stmt wire.Statement) wire.Verdict { + seen = append(seen, stmt.SQL) + + return wire.Verdict{} + })) + + query := msg('Q', cstr("select 1")) + mustWrite(t, p.client, query) + expectMsg(t, p.upstream, query) + + mustWrite(t, p.client, msg('X')) + expectMsg(t, p.upstream, msg('X')) + + if err := <-done; err != nil { + t.Fatalf("Frontend: %v", err) + } + if len(seen) != 1 || seen[0] != "select 1" { + t.Errorf("handler saw %q, want [\"select 1\"]", seen) + } +} + +func TestFrontendDeniesQueriesWithoutForwarding(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + + done := p.frontend(wire.HandlerFunc(func(wire.Statement) wire.Verdict { + return wire.Verdict{Deny: true, Message: "writes are not allowed", Hint: "ask for approval"} + })) + + mustWrite(t, p.client, msg('Q', cstr("delete from orders"))) + + typ, body := readMsgT(t, p.client) + if typ != 'E' { + t.Fatalf("first reply: got %q, want ErrorResponse", typ) + } + fields := errorFields(body) + want := map[byte]string{'S': "ERROR", 'C': "42501", 'M': "writes are not allowed", 'H': "ask for approval"} + for k, v := range want { + if fields[k] != v { + t.Errorf("ErrorResponse field %q: got %q, want %q", k, fields[k], v) + } + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + expectNothing(t, p.upstream) + + mustWrite(t, p.client, msg('X')) + expectMsg(t, p.upstream, msg('X')) + if err := <-done; err != nil { + t.Fatalf("Frontend: %v", err) + } +} + +func TestFrontendDeniesOversizedQueriesAndStaysInSync(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{MaxStatement: 8}) + p.connect(t, "user", "alice") + done := p.frontend(allow) + + mustWrite(t, p.client, msg('Q', cstr("select 1"))) + typ, body := readMsgT(t, p.client) + if typ != 'E' || !strings.Contains(errorFields(body)['M'], "exceeds the 8 byte limit") { + t.Fatalf("oversized query reply: got %q %q, want a limit error", typ, body) + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + + short := msg('Q', cstr("x")) + mustWrite(t, p.client, short) + expectMsg(t, p.upstream, short) + + mustWrite(t, p.client, msg('X')) + expectMsg(t, p.upstream, msg('X')) + if err := <-done; err != nil { + t.Fatalf("Frontend: %v", err) + } +} + +func TestFrontendForwardsCopyDataImmediately(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + done := p.frontend(wire.HandlerFunc(func(wire.Statement) wire.Verdict { + t.Error("handler called for a message that carries no SQL") + + return wire.Verdict{} + })) + + for _, m := range [][]byte{msg('d', []byte("1\ttwo\n")), msg('c')} { + mustWrite(t, p.client, m) + expectMsg(t, p.upstream, m) + } + + mustWrite(t, p.client, msg('X')) + expectMsg(t, p.upstream, msg('X')) + if err := <-done; err != nil { + t.Fatalf("Frontend: %v", err) + } +} + +func TestFrontendRejectsMalformedHeader(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + done := p.frontend(allow) + + mustWrite(t, p.client, []byte{'Q', 0, 0, 0, 2}) + + if err := <-done; err == nil || !strings.Contains(err.Error(), "malformed") { + t.Errorf("Frontend: got %v, want a malformed message error", err) + } +} + +func TestBackendForwardsAndTracksTransactionStatus(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(denyContaining("update")) + + begin := msg('Q', cstr("begin")) + replies := [][]byte{msg('C', cstr("BEGIN")), msg('Z', []byte("T"))} + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, begin); err != nil { + return err + } + + return write(s, replies...) + }) + + mustWrite(t, p.client, begin) + for _, want := range replies { + expectMsg(t, p.client, want) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + mustWrite(t, p.client, msg('Q', cstr("update t set x = 1"))) + if typ, _ := readMsgT(t, p.client); typ != 'E' { + t.Fatalf("denied query reply: got %q, want ErrorResponse", typ) + } + expectMsg(t, p.client, msg('Z', []byte("T"))) + + _ = p.upstream.Close() + if err := <-backend; err != nil { + t.Errorf("Backend: got %v, want nil after upstream EOF", err) + } + _ = p.client.Close() + <-frontend +} + +func TestBackendRejectsUnexpectedReadyForQuery(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + mustWrite(t, p.upstream, msg('Z', []byte("I"))) + + if err := <-backend; err == nil || !strings.Contains(err.Error(), "without a pending request") { + t.Errorf("Backend: got %v, want an unexpected ReadyForQuery error", err) + } +} + +func TestFrontendHoldsDenialsUntilEarlierQueriesAreAnswered(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(denyContaining("delete")) + + first := msg('Q', cstr("select 1")) + second := msg('Q', cstr("delete from orders")) + replies := [][]byte{msg('C', cstr("SELECT 1")), msg('Z', []byte("I"))} + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, first); err != nil { + return err + } + + return write(s, replies...) + }) + + mustWrite(t, p.client, append(first, second...)) + + for _, want := range replies { + expectMsg(t, p.client, want) + } + if typ, body := readMsgT(t, p.client); typ != 'E' || errorFields(body)['M'] != "denied" { + t.Fatalf("after the first query's replies: got %q %q, want the denial", typ, body) + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + expectNothing(t, p.upstream) + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend +} + +func TestFrontendForwardsAllowedBatchAtSync(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + + var seen []string + frontend := p.frontend(wire.HandlerFunc(func(stmt wire.Statement) wire.Verdict { + seen = append(seen, stmt.SQL) + + return wire.Verdict{} + })) + + batch := extendedBatch("select $1") + replies := [][]byte{msg('1'), msg('2'), msg('C', cstr("SELECT 1")), msg('Z', []byte("I"))} + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, bytes.Join(batch, nil)); err != nil { + return err + } + + return write(s, replies...) + }) + + mustWrite(t, p.client, bytes.Join(batch, nil)) + for _, want := range replies { + expectMsg(t, p.client, want) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + if len(seen) != 1 || seen[0] != "select $1" { + t.Errorf("handler saw %q, want [\"select $1\"]", seen) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend +} + +func TestFrontendRejectsAWholeBatchWhenOneStatementIsDenied(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(denyContaining("delete")) + + // An allowed statement precedes a denied one in the same batch. Nothing but + // a Sync must reach the upstream, so the allowed statement never executes. + batch := bytes.Join([][]byte{ + pMsg("select 1"), bindMsg, execMsg, + pMsg("delete from orders"), 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) + + typ, body := readMsgT(t, p.client) + if typ != 'E' || errorFields(body)['C'] != "42501" { + t.Fatalf("batch denial: got %q %q, want ErrorResponse 42501", typ, body) + } + 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 +} + +func TestFrontendDeniesParseThenAcceptsTheNextBatch(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(denyContaining("delete")) + + allowed := extendedBatch("select 1") + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, syncMsg); err != nil { // the denied batch forwards only its Sync + return err + } + if err := write(s, msg('Z', []byte("I"))); err != nil { + return err + } + if err := expectBytes(s, bytes.Join(allowed, nil)); err != nil { + return err + } + + return write(s, msg('1'), msg('2'), msg('C', cstr("SELECT 1")), msg('Z', []byte("I"))) + }) + + mustWrite(t, p.client, bytes.Join(extendedBatch("delete from orders"), nil)) + if typ, body := readMsgT(t, p.client); typ != 'E' || errorFields(body)['C'] != "42501" { + t.Fatalf("Parse denial: got %q %q, want ErrorResponse 42501", typ, body) + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + + mustWrite(t, p.client, bytes.Join(allowed, nil)) + for _, want := range [][]byte{msg('1'), msg('2'), msg('C', cstr("SELECT 1")), msg('Z', []byte("I"))} { + expectMsg(t, p.client, want) + } + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend +} + +func TestFrontendRejectsFunctionCalls(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + done := p.frontend(allow) + + mustWrite(t, p.client, msg('F', be32(1), be16(0), be16(0), be16(0))) + + typ, body := readMsgT(t, p.client) + if typ != 'E' || errorFields(body)['C'] != "0A000" { + t.Fatalf("FunctionCall reply: got %q %q, want ErrorResponse 0A000", typ, body) + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + expectNothing(t, p.upstream) + + mustWrite(t, p.client, msg('X')) + expectMsg(t, p.upstream, msg('X')) + if err := <-done; err != nil { + t.Fatalf("Frontend: %v", err) + } +} + +func TestExtendedCopyInDropsTheIgnoredSyncSlot(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(denyContaining("delete")) + + // Extended COPY FROM STDIN: the server ignores this batch's Sync while + // reading copy data and answers a later Sync instead, so two Syncs yield + // one ReadyForQuery. The proxy must not leave a phantom slot behind. + copyBatch := bytes.Join([][]byte{pMsg("copy t from stdin"), bindMsg, execMsg, syncMsg}, nil) + copyIn := msg('G', []byte{0}, be16(1), be16(0)) + data := bytes.Join([][]byte{msg('d', []byte("1\n")), msg('c'), syncMsg}, nil) + + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, copyBatch); err != nil { + return err + } + if err := write(s, msg('1'), msg('2'), copyIn); err != nil { // ParseComplete, BindComplete, CopyInResponse; no Z + return err + } + if err := expectBytes(s, data); err != nil { + return err + } + + return write(s, msg('C', cstr("COPY 1")), msg('Z', []byte("I"))) // one Z for the second Sync + }) + + mustWrite(t, p.client, copyBatch) + for _, want := range [][]byte{msg('1'), msg('2'), copyIn} { + expectMsg(t, p.client, want) + } + mustWrite(t, p.client, data) + expectMsg(t, p.client, msg('C', cstr("COPY 1"))) + expectMsg(t, p.client, msg('Z', []byte("I"))) + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + // The queue is balanced: a following denial is reported at once, not stuck + // behind a phantom slot. + mustWrite(t, p.client, msg('Q', cstr("delete from x"))) + if typ, body := readMsgT(t, p.client); typ != 'E' || errorFields(body)['C'] != "42501" { + t.Fatalf("denial after COPY: got %q %q, want ErrorResponse 42501", typ, body) + } + expectMsg(t, p.client, msg('Z', []byte("I"))) + + _ = p.client.Close() + _ = p.upstream.Close() + <-frontend + <-backend +} + +func TestSimpleCopyInKeepsItsSlot(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(allow) + + // A simple-protocol COPY has one ReadyForQuery at the end, so its slot must + // survive the CopyInResponse. + q := msg('Q', cstr("copy t from stdin")) + copyIn := msg('G', []byte{0}, be16(1), be16(0)) + data := bytes.Join([][]byte{msg('d', []byte("1\n")), msg('c')}, nil) + + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, q); err != nil { + return err + } + if err := write(s, copyIn); err != nil { + return err + } + if err := expectBytes(s, data); err != nil { + return err + } + + return write(s, msg('C', cstr("COPY 1")), msg('Z', []byte("I"))) + }) + + mustWrite(t, p.client, q) + expectMsg(t, p.client, copyIn) + mustWrite(t, p.client, data) + expectMsg(t, p.client, msg('C', cstr("COPY 1"))) + 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 +} + +func TestDenyAfterEarlyForwardTearsDownTheSession(t *testing.T) { + t.Parallel() + + p := newPipes(t, pg.Dialect{}) + p.connect(t, "user", "alice") + backend := p.backend() + frontend := p.frontend(denyContaining("delete")) + + // An explicit Flush forwards the first statement, then a denied Parse + // arrives. The proxy reports the denial and tears down, so the upstream + // rolls back instead of committing at a plain Sync. + first := bytes.Join([][]byte{pMsg("select 1"), bindMsg, execMsg, msg('H')}, nil) + server := p.serve(func(s net.Conn) error { + if err := expectBytes(s, first); err != nil { + return err + } + + return write(s, msg('1'), msg('2'), msg('C', cstr("SELECT 1"))) + }) + + mustWrite(t, p.client, first) + for _, want := range [][]byte{msg('1'), msg('2'), msg('C', cstr("SELECT 1"))} { + expectMsg(t, p.client, want) + } + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + + mustWrite(t, p.client, pMsg("delete from orders")) + if typ, body := readMsgT(t, p.client); typ != 'E' || errorFields(body)['C'] != "42501" { + t.Fatalf("denial: got %q %q, want ErrorResponse 42501", typ, body) + } + + if err := <-frontend; err == nil { + t.Error("Frontend: got nil, want a teardown error after denying a partially forwarded batch") + } + _ = p.client.Close() + _ = p.upstream.Close() + <-backend +} + +// pipes connects a session to a fake client and a fake upstream over net.Pipe. +type pipes struct { + sess wire.Session + client net.Conn + upstream net.Conn +} + +type handshakeResult struct { + startup wire.Startup + err error +} + +func newPipes(t *testing.T, d pg.Dialect) pipes { + t.Helper() + + clientSide, clientPeer := net.Pipe() + upstreamSide, upstreamPeer := net.Pipe() + for _, c := range []net.Conn{clientSide, clientPeer, upstreamSide, upstreamPeer} { + if err := c.SetDeadline(time.Now().Add(timeout)); err != nil { + t.Fatalf("set deadline: %v", err) + } + t.Cleanup(func() { _ = c.Close() }) + } + + return pipes{sess: d.NewSession(clientSide, upstreamSide), client: clientPeer, upstream: upstreamPeer} +} + +func (p pipes) handshake() <-chan handshakeResult { + result := make(chan handshakeResult, 1) + go func() { + startup, err := p.sess.Handshake() + result <- handshakeResult{startup: startup, err: err} + }() + + return result +} + +func (p pipes) frontend(h wire.Handler) <-chan error { + done := make(chan error, 1) + go func() { done <- p.sess.Frontend(h) }() + + return done +} + +func (p pipes) backend() <-chan error { + done := make(chan error, 1) + go func() { done <- p.sess.Backend() }() + + return done +} + +// serve runs script against the upstream end on its own goroutine. +func (p pipes) serve(script func(s net.Conn) error) <-chan error { + done := make(chan error, 1) + go func() { done <- script(p.upstream) }() + + return done +} + +// connect completes a trust-authenticated handshake for the given startup parameters. +func (p pipes) connect(t *testing.T, params ...string) wire.Startup { + t.Helper() + + result := p.handshake() + server := p.serve(func(s net.Conn) error { + if _, err := readStartup(s); err != nil { + return err + } + + return write(s, auth(0), msg('Z', []byte("I"))) + }) + + mustWrite(t, p.client, startupPacket(params...)) + expectMsg(t, p.client, auth(0)) + expectMsg(t, p.client, msg('Z', []byte("I"))) + + if err := <-server; err != nil { + t.Fatalf("fake server: %v", err) + } + got := <-result + if got.err != nil { + t.Fatalf("Handshake: %v", got.err) + } + + return got.startup +} + +var allow = wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }) + +func denyContaining(word string) wire.Handler { + return wire.HandlerFunc(func(stmt wire.Statement) wire.Verdict { + if strings.Contains(stmt.SQL, word) { + return wire.Verdict{Deny: true, Message: "denied"} + } + + return wire.Verdict{} + }) +} + +var ( + bindMsg = msg('B', cstr(""), cstr(""), be16(0), be16(0), be16(0)) + execMsg = msg('E', cstr(""), be32(0)) + syncMsg = msg('S') +) + +func pMsg(sql string) []byte { + return msg('P', cstr(""), cstr(sql), be16(0)) +} + +// extendedBatch is Parse, Bind, Execute, Sync for an unnamed statement. +func extendedBatch(sql string) [][]byte { + return [][]byte{pMsg(sql), bindMsg, execMsg, syncMsg} +} + +func be16(v uint16) []byte { + return binary.BigEndian.AppendUint16(nil, v) +} + +func be32(v uint32) []byte { + return binary.BigEndian.AppendUint32(nil, v) +} + +func cstr(s string) []byte { + return append([]byte(s), 0) +} + +func msg(typ byte, parts ...[]byte) []byte { + body := bytes.Join(parts, nil) + out := []byte{typ} + out = binary.BigEndian.AppendUint32(out, uint32(len(body)+4)) //nolint:gosec // test payloads are tiny + + return append(out, body...) +} + +func auth(code uint32, parts ...[]byte) []byte { + return msg('R', append([][]byte{be32(code)}, parts...)...) +} + +// packet builds an untyped startup-phase packet: length, code, payload. +func packet(code uint32, parts ...[]byte) []byte { + body := bytes.Join(parts, nil) + out := binary.BigEndian.AppendUint32(nil, uint32(len(body)+8)) //nolint:gosec // test payloads are tiny + out = binary.BigEndian.AppendUint32(out, code) + + return append(out, body...) +} + +func startupPacket(params ...string) []byte { + parts := make([][]byte, 0, len(params)+1) + for _, p := range params { + parts = append(parts, cstr(p)) + } + parts = append(parts, []byte{0}) + + return packet(protocolVersion3, parts...) +} + +func write(w io.Writer, msgs ...[]byte) error { + for _, m := range msgs { + if _, err := w.Write(m); err != nil { + return err //nolint:wrapcheck // test helper surfaces the raw pipe error + } + } + + return nil +} + +func mustWrite(t *testing.T, w io.Writer, b []byte) { + t.Helper() + + if _, err := w.Write(b); err != nil { + t.Fatalf("write: %v", err) + } +} + +func readN(t *testing.T, r io.Reader, n int) []byte { + t.Helper() + + buf := make([]byte, n) + if _, err := io.ReadFull(r, buf); err != nil { + t.Fatalf("read %d bytes: %v", n, err) + } + + return buf +} + +// readStartup reads one untyped startup-phase packet and returns it whole. +func readStartup(r io.Reader) ([]byte, error) { + length := make([]byte, 4) + if _, err := io.ReadFull(r, length); err != nil { + return nil, err //nolint:wrapcheck // test helper surfaces the raw pipe error + } + + out := make([]byte, binary.BigEndian.Uint32(length)) + copy(out, length) + if _, err := io.ReadFull(r, out[4:]); err != nil { + return nil, err //nolint:wrapcheck // test helper surfaces the raw pipe error + } + + return out, nil +} + +// readMsg reads one typed message and returns it whole, header included. +func readMsg(r io.Reader) ([]byte, error) { + header := make([]byte, 5) + if _, err := io.ReadFull(r, header); err != nil { + return nil, err //nolint:wrapcheck // test helper surfaces the raw pipe error + } + + n := binary.BigEndian.Uint32(header[1:]) + out := make([]byte, 5+int(n)-4) + copy(out, header) + if _, err := io.ReadFull(r, out[5:]); err != nil { + return nil, err //nolint:wrapcheck // test helper surfaces the raw pipe error + } + + return out, nil +} + +func readMsgT(t *testing.T, r io.Reader) (typ byte, body []byte) { + t.Helper() + + m, err := readMsg(r) + if err != nil { + t.Fatalf("read message: %v", err) + } + + return m[0], m[5:] +} + +func expectMsg(t *testing.T, r io.Reader, want []byte) { + t.Helper() + + got, err := readMsg(r) + if err != nil { + t.Fatalf("read message: %v", err) + } + if !bytes.Equal(got, want) { + t.Fatalf("message: got %q, want %q", got, want) + } +} + +func expectBytes(r io.Reader, want []byte) error { + got := make([]byte, len(want)) + if _, err := io.ReadFull(r, got); err != nil { + return err //nolint:wrapcheck // test helper surfaces the raw pipe error + } + if !bytes.Equal(got, want) { + return errors.New("unexpected bytes: " + string(got)) + } + + return nil +} + +func expectNothing(t *testing.T, c net.Conn) { + t.Helper() + + if err := c.SetReadDeadline(time.Now().Add(100 * time.Millisecond)); err != nil { + t.Fatalf("set read deadline: %v", err) + } + defer func() { _ = c.SetReadDeadline(time.Now().Add(timeout)) }() + + buf := make([]byte, 1) + if _, err := c.Read(buf); !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("read: got %v, want nothing to arrive", err) + } +} + +func errorFields(body []byte) map[byte]string { + fields := make(map[byte]string) + for len(body) > 0 && body[0] != 0 { + end := bytes.IndexByte(body[1:], 0) + if end < 0 { + break + } + fields[body[0]] = string(body[1 : 1+end]) + body = body[2+end:] + } + + return fields +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 0ec7a26..8491db2 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -9,21 +9,34 @@ import ( "net" "sync" "time" + + "github.com/mickamy/rollcall/internal/wire" ) const ( - dialTimeout = 10 * time.Second - minAcceptDelay = 5 * time.Millisecond - maxAcceptDelay = time.Second + dialTimeout = 10 * time.Second + handshakeTimeout = 30 * time.Second + minAcceptDelay = 5 * time.Millisecond + maxAcceptDelay = time.Second ) -// Server relays every accepted connection to Upstream byte for byte. +// Server accepts client connections and drives each one through Dialect +// against a connection to Upstream. type Server struct { Upstream string - Logger *slog.Logger + Dialect wire.Dialect + // Handler decides each statement; nil allows everything. + Handler wire.Handler + Logger *slog.Logger } func (s Server) Serve(ctx context.Context, ln net.Listener) error { + if s.Dialect == nil { + return errors.New("proxy: Dialect is required") + } + if s.Handler == nil { + s.Handler = wire.HandlerFunc(func(wire.Statement) wire.Verdict { return wire.Verdict{} }) + } if s.Logger == nil { s.Logger = slog.New(slog.DiscardHandler) } @@ -83,12 +96,35 @@ func (s Server) handle(ctx context.Context, client net.Conn) { }) defer stop() - logger.Debug("session opened") + sess := s.Dialect.NewSession(client, upstream) + startup, err := handshake(sess, client, upstream) + if err != nil { + switch { + case ctx.Err() != nil: + case errors.Is(err, wire.ErrNoSession): + logger.Debug("out-of-band request relayed") + case errors.Is(err, wire.ErrRejected): + logger.Info("session rejected", "user", startup.User, "database", startup.Database, "error", err) + default: + logger.Warn("handshake", "error", err) + } + + return + } + + logger = logger.With("user", startup.User, "database", startup.Database, "application", startup.Application) + logger.Info("session opened") var wg sync.WaitGroup var toUpstream, toClient error - wg.Go(func() { toUpstream = relay(upstream, client) }) - wg.Go(func() { toClient = relay(client, upstream) }) + wg.Go(func() { + toUpstream = sess.Frontend(s.Handler) + closeWrite(upstream) + }) + wg.Go(func() { + toClient = sess.Backend() + closeWrite(client) + }) wg.Wait() for _, err := range []error{toUpstream, toClient} { @@ -97,29 +133,44 @@ func (s Server) handle(ctx context.Context, client net.Conn) { } } - logger.Debug("session closed") + logger.Info("session closed") +} + +// handshake bounds the handshake with a deadline on both connections and lifts +// it again once the session is established. +func handshake(sess wire.Session, conns ...net.Conn) (wire.Startup, error) { + deadline := time.Now().Add(handshakeTimeout) + for _, c := range conns { + _ = c.SetDeadline(deadline) + } + + startup, err := sess.Handshake() + + for _, c := range conns { + _ = c.SetDeadline(time.Time{}) + } + + if err != nil { + return startup, fmt.Errorf("handshake: %w", err) + } + + return startup, nil } type closeWriter interface { CloseWrite() error } -// relay copies src to dst until src is done, then passes the end of stream on -// to dst as a half-close so that dst's owner can still drain the other direction. -func relay(dst, src net.Conn) error { - _, err := io.Copy(dst, src) - - if cw, ok := dst.(closeWriter); ok { +// closeWrite passes the end of one direction on as a half-close so that the +// peer can still drain the other direction. +func closeWrite(c net.Conn) { + if cw, ok := c.(closeWriter); ok { _ = cw.CloseWrite() - } else { - _ = dst.Close() - } - if err != nil { - return fmt.Errorf("copy: %w", err) + return } - return nil + _ = c.Close() } func abnormal(err error) bool { diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index 4b4c043..8b8b447 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -3,6 +3,7 @@ package proxy_test import ( "context" "errors" + "fmt" "io" "net" "strings" @@ -10,6 +11,7 @@ import ( "time" "github.com/mickamy/rollcall/internal/proxy" + "github.com/mickamy/rollcall/internal/wire" ) const ( @@ -99,7 +101,7 @@ func TestServeReturnsAcceptError(t *testing.T) { ln := listen(t) errc := make(chan error, 1) go func() { - errc <- proxy.Server{Upstream: unreachable}.Serve(t.Context(), ln) + errc <- proxy.Server{Upstream: unreachable, Dialect: rawDialect{}}.Serve(t.Context(), ln) }() _ = ln.Close() @@ -114,6 +116,51 @@ func TestServeReturnsAcceptError(t *testing.T) { } } +func TestServeRequiresDialect(t *testing.T) { + t.Parallel() + + ln := listen(t) + defer func() { _ = ln.Close() }() + + err := proxy.Server{Upstream: unreachable}.Serve(t.Context(), ln) + if err == nil || !strings.Contains(err.Error(), "Dialect is required") { + t.Errorf("Serve: got %v, want a missing Dialect error", err) + } +} + +// rawDialect relays bytes without interpreting them, so the proxy's lifecycle +// can be tested against a plain echo server. +type rawDialect struct{} + +func (rawDialect) NewSession(client, upstream net.Conn) wire.Session { + return rawSession{client: client, upstream: upstream} +} + +type rawSession struct { + client net.Conn + upstream net.Conn +} + +func (rawSession) Handshake() (wire.Startup, error) { + return wire.Startup{User: "raw"}, nil +} + +func (s rawSession) Frontend(wire.Handler) error { + if _, err := io.Copy(s.upstream, s.client); err != nil { + return fmt.Errorf("frontend: %w", err) + } + + return nil +} + +func (s rawSession) Backend() error { + if _, err := io.Copy(s.client, s.upstream); err != nil { + return fmt.Errorf("backend: %w", err) + } + + return nil +} + func startProxy(t *testing.T, upstream string) (addr string, cancel context.CancelFunc, wait func() error) { t.Helper() @@ -123,7 +170,7 @@ func startProxy(t *testing.T, upstream string) (addr string, cancel context.Canc errc := make(chan error, 1) go func() { - errc <- proxy.Server{Upstream: upstream}.Serve(ctx, ln) + errc <- proxy.Server{Upstream: upstream, Dialect: rawDialect{}}.Serve(ctx, ln) }() wait = func() error { diff --git a/internal/wire/wire.go b/internal/wire/wire.go new file mode 100644 index 0000000..ae4e301 --- /dev/null +++ b/internal/wire/wire.go @@ -0,0 +1,59 @@ +// Package wire holds the dialect-neutral types that the proxy, policy, and +// ledger share. Database-specific packages implement Dialect and Session and +// never leak their own message types through this boundary. +package wire + +import ( + "errors" + "net" +) + +var ( + // ErrNoSession reports a connection that carried an out-of-band request + // (such as a cancel request) and is finished without starting a session. + ErrNoSession = errors.New("connection did not start a session") + // ErrRejected reports that the upstream refused the connection during the handshake. + ErrRejected = errors.New("upstream rejected the connection") +) + +// Startup is what the client declared about itself when it connected. +type Startup struct { + User string + Database string + Application string + Params map[string]string +} + +type Statement struct { + SQL string +} + +// Verdict is the zero value to let a statement through. +type Verdict struct { + Deny bool + Message string + Hint string +} + +type Handler interface { + Statement(stmt Statement) Verdict +} + +type HandlerFunc func(stmt Statement) Verdict + +func (f HandlerFunc) Statement(stmt Statement) Verdict { + return f(stmt) +} + +type Dialect interface { + NewSession(client, upstream net.Conn) Session +} + +// Session drives one client connection and its upstream connection. +// Handshake runs first and alone; Frontend and Backend then run concurrently +// until their side of the conversation ends. +type Session interface { + Handshake() (Startup, error) + Frontend(h Handler) error + Backend() error +}