Skip to content

Instantly share code, notes, and snippets.

@spraints
Last active June 27, 2026 14:07
Show Gist options
  • Select an option

  • Save spraints/2527a589757f0ea33f9f2995ac776cd0 to your computer and use it in GitHub Desktop.

Select an option

Save spraints/2527a589757f0ea33f9f2995ac776cd0 to your computer and use it in GitHub Desktop.
A really simple postgres protocol analyzer
test.crt
test.key

Simple postgres protocol analyzer

Synopsis:

(Assume postgres is running on localhost.)

$ go run . &

$ psql "host=127.0.0.1 port=5555 sslmode=disable user=... password=..."
(protocol details spew)

postgres=> select 1;
(more protocol details spew)

pg-mitm

A man-in-the-middle proxy for PostgreSQL connections that logs protocol messages.

Usage

go run . [flags]
Flag Default Description
--listen 127.0.0.1:5555 Address the proxy listens on
--dial 127.0.0.1:5432 Address of the upstream PostgreSQL server
--ssl-cert (none) Path to an SSL certificate for the proxy's frontend (auto-generated if path given but file missing)
--ssl-key (none) Path to an SSL key for the proxy's frontend (auto-generated if path given but file missing)
--tls-mode allow How the proxy handles TLS toward the backend: disable, allow, require, or insecure
--debug-tls / --tls-debug false Log TLS message details (very noisy)
--pgx false Use pgx to parse messages instead of the built-in parser
--debug-addr 127.0.0.1:8967 Address to serve the metrics/debug HTTP endpoint on

Example with Docker

1. Start PostgreSQL

docker compose up -d

This starts PostgreSQL on 127.0.0.1:5432 with user postgres and password secret.

2. Start the proxy

go run . --listen 127.0.0.1:5555 --dial 127.0.0.1:5432

The proxy listens on port 5555 and forwards to the database on port 5432. Protocol messages are logged to stderr.

3. Connect through the proxy

psql "host=127.0.0.1 port=5555 sslmode=disable user=postgres password=secret dbname=postgres"

Queries you run in psql will be logged by the proxy.

Metrics

While the proxy is running, visit http://127.0.0.1:8967/metrics to see connection metrics.

Instructions for Adding PostgreSQL Message Types

Overview

This document describes how to add implementations for PostgreSQL protocol message types in messages.go.

Reference

Message types are defined in the PostgreSQL documentation: https://www.postgresql.org/docs/current/protocol-message-formats.html

Implementation Guidelines

1. Switch Statement Ordering

In the getMessageFormatter() function, add new message type cases in alphabetical order:

  • All uppercase letters before lowercase letters
  • Within each group (uppercase/lowercase), maintain alphabetical order

Example ordering:

case 'K':  // uppercase
case 'Q':
case 'R':
case 'S':
case 'T':
case 'X':
case 'Z':
case 'p':  // lowercase

2. Type Definition Ordering

Define the message type structs in the same order as they appear in the switch statement.

Each message type struct follows this pattern:

type messageName struct {
    want int
    have []byte
}

3. Format Method

Each message type must implement the messageFormatter interface by providing a Format([]byte) string method.

The Format method should:

  1. Accumulate incoming data: m.have = append(m.have, data...)
  2. Check if all data has arrived: if len(m.have) < m.want { return "..." }
  3. Validate the message length
  4. Parse and format the message content according to the protocol specification
  5. Handle error cases gracefully (invalid data, unexpected trailing bytes, etc.)

4. Common Patterns

  • Use decodeString() for null-terminated strings
  • Use byteOrder.UintXX() for reading multi-byte integers
  • Use strings.Builder for constructing complex output
  • Report hex dumps for raw/invalid data: %x
  • Quote strings in output: %q

Example Implementation

See existing message types in messages.go for examples:

  • Simple messages: terminate (no payload), readyForQuery (single byte)
  • String messages: query, parameterStatus
  • Complex messages: rowDescription, authenticationRequest
package main
import (
"log"
"net"
"net/http"
)
func serveDebug(addr string, m *metrics) {
l, err := net.Listen("tcp", addr)
if err != nil {
panic(err)
}
defer l.Close()
log.Printf("debug server listening on %v", l.Addr())
mux := http.NewServeMux()
mux.Handle("/metrics", m.Handler())
s := &http.Server{Handler: mux}
s.Serve(l)
}
services:
postgres:
image: postgres:17
environment:
POSTGRES_USER: postgres
POSTGRES_PASSWORD: secret
POSTGRES_DB: postgres
ports:
- "127.0.0.1:5432:5432"
module pgmitm
go 1.25.0
require (
github.com/jackc/pgx/v5 v5.7.6
github.com/spf13/pflag v1.0.10
go.withmatt.com/metrics v0.6.0
)
require (
github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
golang.org/x/crypto v0.46.0 // indirect
golang.org/x/sys v0.39.0 // indirect
golang.org/x/text v0.32.0 // indirect
)
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM=
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.7.6 h1:rWQc5FwZSPX58r1OQmkuaNicxdmExaEz5A2DO2hUuTk=
github.com/jackc/pgx/v5 v5.7.6/go.mod h1:aruU7o91Tc2q2cFp5h4uP3f6ztExVpyVv88Xl/8Vl8M=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/spf13/pflag v1.0.10 h1:4EBh2KAYBwaONj6b2Ye1GiHfwjqyROoF4RwYO+vPwFk=
github.com/spf13/pflag v1.0.10/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk=
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
go.withmatt.com/metrics v0.6.0 h1:KbaV069e38ife9VnZXAn/YLjt0AgTD2HalJqPzdv4CU=
go.withmatt.com/metrics v0.6.0/go.mod h1:UcQWoQtd/XaG8G0uG7EdASihWoQ4boXBbzge7XrhXIM=
golang.org/x/crypto v0.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sys v0.39.0 h1:CvCKL8MeisomCi6qNZ+wbb0DN9E5AATixKsvNtMoMFk=
golang.org/x/sys v0.39.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/text v0.32.0 h1:ZD01bjUt1FQ9WJ0ClOL5vxgxOI/sVCNgX1YtKwcY0mU=
golang.org/x/text v0.32.0/go.mod h1:o/rUWzghvpD5TXrTIBuJU77MTaN0ljMWE47kxGJQ7jY=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
#!/bin/bash
set -eux
# load data directly to postgres
PGPASSWORD=secret pgbench -h 127.0.0.1 -p 5432 -U postgres -i postgres
# run queries through the proxy
PGPASSWORD=secret pgbench -h 127.0.0.1 -p 5555 -U postgres -c 5 -T 30 postgres
// Synopsis:
//
// $ go run . --listen 127.0.0.1:5555 --dial 127.0.0.1:5432
//
// $ psql "host=127.0.0.1 port=5555 sslmode=disable user=xx password=yy"
package main
import (
"bytes"
"context"
"crypto/tls"
"fmt"
"io"
"log"
"net"
"os"
"os/signal"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"github.com/spf13/pflag"
)
func main() {
if err := mainImpl(); err != nil {
log.Printf("fatal error: %v", err)
os.Exit(1)
}
}
func mainImpl() error {
ctx, cancel := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
defer cancel()
listenAddr := pflag.String("listen", "127.0.0.1:5555", "address to listen on")
upstreamAddr := pflag.String("dial", "127.0.0.1:5432", "address to forward connections to")
sslCertFile := pflag.String("ssl-cert", "", "path to an SSL certificate to use for the server (will be generated if the path is present but the file is not)")
sslKeyFile := pflag.String("ssl-key", "", "path to an SSL key to use for the server (will be generated if the path is present but the file is not)")
beTLSModeStr := pflag.String("tls-mode", "allow", "what to do about TLS on the backend (disable, require, allow, insecure)")
tlsDebug := pflag.Bool("debug-tls", false, "show information about TLS messages (very noisy)")
pflag.BoolVar(tlsDebug, "tls-debug", false, "show information about TLS messages (very noisy)")
pgx := pflag.Bool("pgx", false, "use pgx to parse messages")
debugAddr := pflag.String("debug-addr", "127.0.0.1:8967", "address to serve metrics on")
pflag.Parse()
m := newMetrics()
go serveDebug(*debugAddr, m)
var frontendTLS *tls.Config
if *sslCertFile != "" && *sslKeyFile != "" {
if c, err := initSSL(*sslCertFile, *sslKeyFile); err != nil {
return err
} else {
log.Printf("TLS enabled")
frontendTLS = c
}
}
beTLS, ok := configureBackendTLS(*beTLSModeStr, strings.Split(*upstreamAddr, ":")[0])
if !ok {
return fmt.Errorf("invalid tls-mode %q", *beTLSModeStr)
}
l, err := net.Listen("tcp", *listenAddr)
if err != nil {
return fmt.Errorf("%s: error listening: %w", *listenAddr, err)
}
log.Printf("%s: listening", *listenAddr)
var closing int32
go func(ctx context.Context, l net.Listener) {
<-ctx.Done()
log.Printf("received shutdown signal")
atomic.StoreInt32(&closing, 1)
l.Close()
}(ctx, l)
var wg sync.WaitGroup
for {
conn, err := l.Accept()
if err != nil {
if atomic.LoadInt32(&closing) == 0 {
log.Printf("%s: shutting down because of accept error: %v", *listenAddr, err)
}
break
}
wg.Go(func() {
logger := log.New(log.Writer(), "["+conn.RemoteAddr().String()+"] ", log.Flags())
logger.Println("accepted connection")
if *pgx {
servepgx(conn, *upstreamAddr, frontendTLS, beTLS, logger)
} else {
serve(conn, *upstreamAddr, frontendTLS, beTLS, *tlsDebug, m, logger)
}
})
}
cancel()
wg.Wait()
return nil
}
func serve(conn net.Conn, pgAddr string, feTLS *tls.Config, beTLS backendTLSPolicy, tlsDebug bool, m *metrics, l *log.Logger) {
defer conn.Close()
defer l.Printf("closing connection")
bg, err := net.Dial("tcp", pgAddr)
if err != nil {
l.Printf("%s: error dialing backend: %v", pgAddr, err)
return
}
defer bg.Close()
handles, err := negotiate(conn, bg, feTLS, beTLS, tlsDebug, l)
if err != nil {
l.Printf("%s: error negotiating transport: %v", pgAddr, err)
return
}
connMetrics := m.connMetrics(conn)
defer connMetrics.closed()
var wg sync.WaitGroup
wg.Go(func() { serveClient(handles.ClientReader, handles.ServerWriter, connMetrics, l) })
wg.Go(func() { serveServer(handles.ServerReader, handles.ClientWriter, connMetrics, l) })
wg.Wait()
}
type negotiated struct {
ClientReader, ServerReader io.ReadCloser
ClientWriter, ServerWriter io.WriteCloser
}
var pgSSLRequest = []byte{0, 0, 0, 8, 0x04, 0xd2, 0x16, 0x2f}
func negotiate(frontend, backend net.Conn, feTLS *tls.Config, beTLS backendTLSPolicy, tlsDebug bool, l *log.Logger) (negotiated, error) {
c, err := negotiateClient(frontend, feTLS, tlsDebug, l)
if err != nil {
return negotiated{}, err
}
s, err := negotiateServer(backend, beTLS, tlsDebug, l)
if err != nil {
return negotiated{}, err
}
return negotiated{
ClientReader: c,
ClientWriter: c,
ServerReader: s,
ServerWriter: s,
}, nil
}
func negotiateClient(conn net.Conn, feTLS *tls.Config, tlsDebug bool, l *log.Logger) (io.ReadWriteCloser, error) {
preamble := make([]byte, len(pgSSLRequest))
n, err := conn.Read(preamble)
if err != nil {
return nil, fmt.Errorf("error getting startup message: %v", err)
}
if !bytes.Equal(pgSSLRequest, preamble[:n]) {
return prepended(preamble[:n], conn), nil
}
if feTLS == nil {
l.Printf("--- SSL requested by client, declining ['N'] ---")
if _, err := conn.Write([]byte("N")); err != nil {
return nil, err
}
return conn, nil
}
l.Printf("--- performing SSL handshake with client ['S', tls-handshake] ---")
if _, err := conn.Write([]byte("S")); err != nil {
return nil, err
}
if tlsDebug {
return tls.Server(&tlsServerSpy{conn, l}, feTLS), nil
}
return tls.Server(conn, feTLS), nil
}
func negotiateServer(conn net.Conn, beTLS backendTLSPolicy, tlsDebug bool, l *log.Logger) (io.ReadWriteCloser, error) {
if !beTLS.Request {
return conn, nil
}
l.Printf("server handshake: send MSG '%c' (%x) (raw = %x)", pgSSLRequest[0], pgSSLRequest[0], pgSSLRequest)
if _, err := conn.Write(pgSSLRequest); err != nil {
return nil, err
}
var resp [1]byte
n, err := conn.Read(resp[:])
if err != nil {
return nil, err
}
if n != 1 {
return nil, fmt.Errorf("error reading response to SSLRequest")
}
l.Printf("server handshake: recv MSG '%c' (%x) (raw = %x)", resp[0], resp[0], resp)
switch resp[0] {
case 'N':
if beTLS.Require {
return nil, fmt.Errorf("TLS is required but backend does not allow it")
}
l.Println("--- server declined SSL ---")
return conn, nil
case 'S':
l.Println("--- performing SSL handshake with server ---")
if tlsDebug {
return tls.Client(&tlsClientSpy{conn, l}, beTLS.Config), nil
}
return tls.Client(conn, beTLS.Config), nil
default:
return nil, fmt.Errorf("invalid response to SSLRequest %c(%02x)", resp[0], resp[0])
}
}
func prepended(preamble []byte, conn net.Conn) io.ReadWriteCloser {
return &prepend{
r: io.MultiReader(bytes.NewReader(preamble), conn),
w: conn,
c: conn,
}
}
type prepend struct {
r io.Reader
w io.Writer
c io.Closer
}
func (p *prepend) Read(buf []byte) (int, error) { return p.r.Read(buf) }
func (p *prepend) Write(buf []byte) (int, error) { return p.w.Write(buf) }
func (p *prepend) Close() error { return p.c.Close() }
func serveClient(client io.ReadCloser, server io.WriteCloser, m messageObserver, l *log.Logger) {
defer client.Close()
defer server.Close()
l = log.New(l.Writer(), l.Prefix()+"client->server ", l.Flags())
processMessages(client, server, m, l, frontend{}, 1)
// Wait a moment before closing connections.
time.Sleep(100 * time.Millisecond)
}
func serveServer(server io.ReadCloser, client io.WriteCloser, m messageObserver, l *log.Logger) {
defer server.Close()
defer client.Close()
l = log.New(l.Writer(), l.Prefix()+"server->client ", l.Flags())
processMessages(server, client, m, l, backend{}, 0)
// Wait a moment before closing connections.
time.Sleep(100 * time.Millisecond)
}
package main
import (
"fmt"
"strings"
)
type frontend struct{}
func (frontend) getMessageFormatter(msgType byte, msgLength int, obs messageObserver) messageFormatter {
// Just the "F" messages from https://www.postgresql.org/docs/current/protocol-message-formats.html
switch msgType {
case 0:
return &startup{want: msgLength}
case 'B':
return &bind{want: msgLength}
case 'C':
return &closeMsg{want: msgLength}
case 'D':
return &describe{want: msgLength}
case 'E':
return &execute{want: msgLength, obs: obs}
case 'P':
return &parse{want: msgLength}
case 'Q':
return &query{want: msgLength, obs: obs}
case 'S':
return &syncMsg{want: msgLength}
case 'X':
return &terminate{want: msgLength}
case 'p':
return &gssResponse{want: msgLength}
default:
panic(fmt.Sprintf("missing frontend message formatter for '%c' (%d)", msgType, msgType))
}
}
type backend struct{}
func (backend) getMessageFormatter(msgType byte, msgLength int, obs messageObserver) messageFormatter {
// Just the "B" messages from https://www.postgresql.org/docs/current/protocol-message-formats.html
switch msgType {
case '1':
return &parseComplete{want: msgLength}
case '2':
return &bindComplete{want: msgLength}
case '3':
return &closeComplete{want: msgLength}
case 'C':
return &commandComplete{want: msgLength}
case 'D':
return &dataRow{want: msgLength}
case 'E':
return &errorResponse{want: msgLength}
case 'K':
return &backendKey{want: msgLength}
case 'N':
return &noticeResponse{want: msgLength}
case 'R':
return &authenticationRequest{want: msgLength}
case 'S':
return &parameterStatus{want: msgLength}
case 'T':
return &rowDescription{want: msgLength}
case 'Z':
return &readyForQuery{want: msgLength, obs: obs}
case 't':
return &parameterDescription{want: msgLength}
default:
panic(fmt.Sprintf("missing backend message formatter '%c' (%d)", msgType, msgType))
}
}
type messageObserver interface {
serverReadyForQuery()
clientQuery()
clientExecute()
}
type messageFormatter interface {
Format([]byte) string
}
type startup struct {
want int
have []byte
}
func (m *startup) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return fmt.Sprintf("... (have=%d, want=%d)", len(m.have), m.want)
}
if m.want < 5 {
return fmt.Sprintf("StartupMessage too short raw=%x", m.have)
}
var sb strings.Builder
protoVersion := byteOrder.Uint32(m.have)
fmt.Fprintf(&sb, "StartupMessage proto=%d", protoVersion)
if protoVersion != 196610 {
fmt.Fprintf(&sb, "(!)")
}
p := m.have[4:]
for {
if len(p) == 1 && p[0] == 0 {
break
}
if len(p) < 1 {
fmt.Fprintf(&sb, " (unexpected end of parameters!)")
break
}
str, rest, ok := decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (unexpected end of parameters remaining=%x)", p)
break
}
fmt.Fprintf(&sb, " %s=", str)
p = rest
str, rest, ok = decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (unexpected end of parameters remaining=%x)", p)
break
}
fmt.Fprintf(&sb, "%q", str)
p = rest
}
return sb.String()
}
type bind struct {
want int
have []byte
}
func (m *bind) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 4 {
return fmt.Sprintf("Bind message too short raw=%x", m.have)
}
var sb strings.Builder
fmt.Fprintf(&sb, "Bind")
// Portal name
portalName, rest, ok := decodeString(m.have)
if !ok {
fmt.Fprintf(&sb, " (error decoding portal name) raw=%x", m.have)
return sb.String()
}
if len(portalName) == 0 {
fmt.Fprintf(&sb, " dstportal=<unnamed>")
} else {
fmt.Fprintf(&sb, " dstportal=%q", portalName)
}
p := rest
// Statement name
stmtName, rest, ok := decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (error decoding statement name) rest=%x", p)
return sb.String()
}
if len(stmtName) == 0 {
fmt.Fprintf(&sb, " stmt=<unnamed>")
} else {
fmt.Fprintf(&sb, " stmt=%q", stmtName)
}
p = rest
// Parameter format codes
if len(p) < 2 {
fmt.Fprintf(&sb, " (missing param format count) rest=%x", p)
return sb.String()
}
numParamFormats := byteOrder.Uint16(p[0:2])
p = p[2:]
if len(p) < int(numParamFormats)*2 {
fmt.Fprintf(&sb, " (incomplete param formats) rest=%x", p)
return sb.String()
}
paramFormats := make([]uint16, numParamFormats)
for i := 0; i < int(numParamFormats); i++ {
paramFormats[i] = byteOrder.Uint16(p[i*2 : (i+1)*2])
}
p = p[numParamFormats*2:]
// Parameter values
if len(p) < 2 {
fmt.Fprintf(&sb, " (missing param value count) rest=%x", p)
return sb.String()
}
numParams := byteOrder.Uint16(p[0:2])
p = p[2:]
fmt.Fprintf(&sb, " params=%d", numParams)
for i := 0; i < int(numParams); i++ {
if len(p) < 4 {
fmt.Fprintf(&sb, " (incomplete param %d length) rest=%x", i, p)
return sb.String()
}
paramLen := int32(byteOrder.Uint32(p[0:4]))
p = p[4:]
switch {
case paramLen == -1:
fmt.Fprintf(&sb, " NULL")
case paramLen < 0:
fmt.Fprintf(&sb, " (invalid param length %d)", paramLen)
return sb.String()
default:
if len(p) < int(paramLen) {
fmt.Fprintf(&sb, " (incomplete param %d data) rest=%x", i, p)
return sb.String()
}
paramData := p[:paramLen]
p = p[paramLen:]
if isPrintable(paramData) {
fmt.Fprintf(&sb, " %q", string(paramData))
} else {
fmt.Fprintf(&sb, " [%x]", paramData)
}
}
}
// Result format codes
if len(p) < 2 {
fmt.Fprintf(&sb, " (missing result format count) rest=%x", p)
return sb.String()
}
numResultFormats := byteOrder.Uint16(p[0:2])
p = p[2:]
if len(p) < int(numResultFormats)*2 {
fmt.Fprintf(&sb, " (incomplete result formats) rest=%x", p)
return sb.String()
}
if numResultFormats > 0 {
fmt.Fprintf(&sb, " result_fmts=[")
for i := 0; i < int(numResultFormats); i++ {
if i > 0 {
fmt.Fprintf(&sb, ",")
}
rf := byteOrder.Uint16(p[i*2 : (i+1)*2])
switch rf {
case 0:
fmt.Fprint(&sb, "text")
case 1:
fmt.Fprint(&sb, "binary")
default:
fmt.Fprintf(&sb, "invalid(%d)", rf)
}
}
fmt.Fprintf(&sb, "]")
p = p[numResultFormats*2:]
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
type closeMsg struct {
want int
have []byte
}
func (m *closeMsg) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 2 {
return fmt.Sprintf("Close message too short raw=%x", m.have)
}
var sb strings.Builder
fmt.Fprintf(&sb, "Close")
objType := m.have[0]
switch objType {
case 'S':
fmt.Fprintf(&sb, " statement")
case 'P':
fmt.Fprintf(&sb, " portal")
default:
fmt.Fprintf(&sb, " ??(%x %c)", objType, objType)
}
name, rest, ok := decodeString(m.have[1:])
if !ok {
fmt.Fprintf(&sb, " (error decoding name) rest=%x", m.have[1:])
return sb.String()
}
if len(name) == 0 {
fmt.Fprintf(&sb, " <unnamed>")
} else {
fmt.Fprintf(&sb, " %q", name)
}
if len(rest) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", rest)
}
return sb.String()
}
type describe struct {
want int
have []byte
}
func (m *describe) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 2 {
return fmt.Sprintf("Describe message too short raw=%x", m.have)
}
var sb strings.Builder
fmt.Fprintf(&sb, "Describe")
objType := m.have[0]
switch objType {
case 'S':
fmt.Fprintf(&sb, " statement")
case 'P':
fmt.Fprintf(&sb, " portal")
default:
fmt.Fprintf(&sb, " ??(%x %c)", objType, objType)
}
name, rest, ok := decodeString(m.have[1:])
if !ok {
fmt.Fprintf(&sb, " (error decoding name) rest=%x", m.have[1:])
return sb.String()
}
if len(name) == 0 {
fmt.Fprintf(&sb, " <unnamed>")
} else {
fmt.Fprintf(&sb, " %q", name)
}
if len(rest) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", rest)
}
return sb.String()
}
type execute struct {
want int
have []byte
obs messageObserver
}
func (m *execute) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 5 {
return fmt.Sprintf("Execute message too short raw=%x", m.have)
}
// Portal name
portalName, rest, ok := decodeString(m.have)
if !ok {
return fmt.Sprintf("Execute (error decoding portal name) raw=%x", m.have)
}
if len(rest) < 4 {
return fmt.Sprintf("Execute (missing max rows) rest=%x", rest)
}
maxRows := byteOrder.Uint32(rest[0:4])
rest = rest[4:]
var sb strings.Builder
m.obs.clientExecute()
fmt.Fprintf(&sb, "Execute")
if len(portalName) == 0 {
fmt.Fprintf(&sb, " portal=<unnamed>")
} else {
fmt.Fprintf(&sb, " portal=%q", portalName)
}
if maxRows == 0 {
fmt.Fprintf(&sb, " maxrows=unlimited")
} else {
fmt.Fprintf(&sb, " maxrows=%d", maxRows)
}
if len(rest) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", rest)
}
return sb.String()
}
type parse struct {
want int
have []byte
}
func (m *parse) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 3 {
return fmt.Sprintf("Parse message too short raw=%x", m.have)
}
var sb strings.Builder
fmt.Fprintf(&sb, "Parse")
// Statement name
stmtName, rest, ok := decodeString(m.have)
if !ok {
fmt.Fprintf(&sb, " (error decoding statement name) raw=%x", m.have)
return sb.String()
}
if len(stmtName) == 0 {
fmt.Fprintf(&sb, " stmt=<unnamed>")
} else {
fmt.Fprintf(&sb, " stmt=%q", stmtName)
}
p := rest
// Query string
queryStr, rest, ok := decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (error decoding query) rest=%x", p)
return sb.String()
}
fmt.Fprintf(&sb, " query=%q", queryStr)
p = rest
// Number of parameter types
if len(p) < 2 {
fmt.Fprintf(&sb, " (missing param count) rest=%x", p)
return sb.String()
}
numParams := byteOrder.Uint16(p[0:2])
p = p[2:]
fmt.Fprintf(&sb, " params=%d", numParams)
if numParams > 0 {
// Parameter type OIDs
if len(p) < int(numParams)*4 {
fmt.Fprintf(&sb, " (incomplete param types, need %d bytes) rest=%x", numParams*4, p)
return sb.String()
}
fmt.Fprintf(&sb, " types=[")
for i := 0; i < int(numParams); i++ {
if i > 0 {
fmt.Fprintf(&sb, ",")
}
typeOID := byteOrder.Uint32(p[i*4 : (i+1)*4])
if typeOID == 0 {
fmt.Fprintf(&sb, "?")
} else {
fmt.Fprintf(&sb, "%d", typeOID)
}
}
fmt.Fprintf(&sb, "]")
p = p[numParams*4:]
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
type parseComplete struct {
want int
have []byte
}
func (m *parseComplete) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want != 0 {
return fmt.Sprintf("ParseComplete (unexpected payload: %x)", m.have)
}
return "ParseComplete"
}
type bindComplete struct {
want int
have []byte
}
func (m *bindComplete) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want != 0 {
return fmt.Sprintf("BindComplete (unexpected payload: %x)", m.have)
}
return "BindComplete"
}
type closeComplete struct {
want int
have []byte
}
func (m *closeComplete) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want != 0 {
return fmt.Sprintf("CloseComplete (unexpected payload: %x)", m.have)
}
return "CloseComplete"
}
type commandComplete struct {
want int
have []byte
}
func (m *commandComplete) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 1 {
return fmt.Sprintf("CommandComplete message too short raw=%x", m.have)
}
tag, rest, ok := decodeString(m.have)
if !ok {
return fmt.Sprintf("CommandComplete invalid tag raw=%x", m.have)
}
if len(rest) > 0 {
return fmt.Sprintf("CommandComplete %q (unexpected trailing data: %x)", tag, rest)
}
return fmt.Sprintf("CommandComplete %q", tag)
}
type dataRow struct {
want int
have []byte
}
func (m *dataRow) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 2 {
return fmt.Sprintf("DataRow message too short raw=%x", m.have)
}
numCols := byteOrder.Uint16(m.have)
var sb strings.Builder
fmt.Fprintf(&sb, "DataRow cols=%d", numCols)
p := m.have[2:]
for i := 0; i < int(numCols); i++ {
if len(p) < 4 {
fmt.Fprintf(&sb, " (incomplete column %d, need length field) rest=%x", i, p)
break
}
colLen := int32(byteOrder.Uint32(p[0:4]))
p = p[4:]
if colLen == -1 {
fmt.Fprintf(&sb, " NULL")
continue
}
if colLen < 0 {
fmt.Fprintf(&sb, " (invalid length %d) rest=%x", colLen, p)
break
}
if len(p) < int(colLen) {
fmt.Fprintf(&sb, " (incomplete column %d, need %d bytes) rest=%x", i, colLen, p)
break
}
colData := p[:colLen]
p = p[colLen:]
// Try to display as string if it looks printable
if isPrintable(colData) {
fmt.Fprintf(&sb, " %q", string(colData))
} else {
fmt.Fprintf(&sb, " [%x]", colData)
}
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
func isPrintable(data []byte) bool {
for _, b := range data {
if b < 32 && b != '\t' && b != '\n' && b != '\r' {
return false
}
if b > 126 {
return false
}
}
return true
}
type errorResponse struct {
want int
have []byte
}
func (m *errorResponse) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 1 {
return fmt.Sprintf("ErrorResponse message too short raw=%x", m.have)
}
var sb strings.Builder
fmt.Fprintf(&sb, "ErrorResponse")
p := m.have
for {
if len(p) == 0 {
fmt.Fprintf(&sb, " (missing terminator)")
break
}
fieldType := p[0]
p = p[1:]
if fieldType == 0 {
// Terminator
break
}
fieldValue, rest, ok := decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (error decoding field %c) rest=%x", fieldType, p)
break
}
p = rest
switch fieldType {
case 'S':
fmt.Fprintf(&sb, " severity=%q", fieldValue)
case 'V':
fmt.Fprintf(&sb, " severity_nonlocal=%q", fieldValue)
case 'C':
fmt.Fprintf(&sb, " code=%q", fieldValue)
case 'M':
fmt.Fprintf(&sb, " message=%q", fieldValue)
case 'D':
fmt.Fprintf(&sb, " detail=%q", fieldValue)
case 'H':
fmt.Fprintf(&sb, " hint=%q", fieldValue)
case 'P':
fmt.Fprintf(&sb, " position=%q", fieldValue)
case 'p':
fmt.Fprintf(&sb, " internal_position=%q", fieldValue)
case 'q':
fmt.Fprintf(&sb, " internal_query=%q", fieldValue)
case 'W':
fmt.Fprintf(&sb, " where=%q", fieldValue)
case 's':
fmt.Fprintf(&sb, " schema=%q", fieldValue)
case 't':
fmt.Fprintf(&sb, " table=%q", fieldValue)
case 'c':
fmt.Fprintf(&sb, " column=%q", fieldValue)
case 'd':
fmt.Fprintf(&sb, " datatype=%q", fieldValue)
case 'n':
fmt.Fprintf(&sb, " constraint=%q", fieldValue)
case 'F':
fmt.Fprintf(&sb, " file=%q", fieldValue)
case 'L':
fmt.Fprintf(&sb, " line=%q", fieldValue)
case 'R':
fmt.Fprintf(&sb, " routine=%q", fieldValue)
default:
fmt.Fprintf(&sb, " %x=%q", fieldType, fieldValue)
}
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
type backendKey struct {
want int
have []byte
}
func (m *backendKey) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 4 {
return fmt.Sprintf("BackendKeyData too short raw=%x", m.have)
}
pid := byteOrder.Uint32(m.have)
return fmt.Sprintf("BackendKeyData pid=%d secret=%x", pid, m.have[4:])
}
type noticeResponse struct {
want int
have []byte
}
func (m *noticeResponse) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 1 {
return fmt.Sprintf("NoticeResponse message too short raw=%x", m.have)
}
var sb strings.Builder
fmt.Fprintf(&sb, "NoticeResponse")
p := m.have
for {
if len(p) == 0 {
fmt.Fprintf(&sb, " (missing terminator)")
break
}
fieldType := p[0]
p = p[1:]
if fieldType == 0 {
break
}
fieldValue, rest, ok := decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (error decoding field %c) rest=%x", fieldType, p)
break
}
p = rest
switch fieldType {
case 'S':
fmt.Fprintf(&sb, " severity=%q", fieldValue)
case 'V':
fmt.Fprintf(&sb, " severity_nonlocal=%q", fieldValue)
case 'C':
fmt.Fprintf(&sb, " code=%q", fieldValue)
case 'M':
fmt.Fprintf(&sb, " message=%q", fieldValue)
case 'D':
fmt.Fprintf(&sb, " detail=%q", fieldValue)
case 'H':
fmt.Fprintf(&sb, " hint=%q", fieldValue)
case 'P':
fmt.Fprintf(&sb, " position=%q", fieldValue)
case 'p':
fmt.Fprintf(&sb, " internal_position=%q", fieldValue)
case 'q':
fmt.Fprintf(&sb, " internal_query=%q", fieldValue)
case 'W':
fmt.Fprintf(&sb, " where=%q", fieldValue)
case 's':
fmt.Fprintf(&sb, " schema=%q", fieldValue)
case 't':
fmt.Fprintf(&sb, " table=%q", fieldValue)
case 'c':
fmt.Fprintf(&sb, " column=%q", fieldValue)
case 'd':
fmt.Fprintf(&sb, " datatype=%q", fieldValue)
case 'n':
fmt.Fprintf(&sb, " constraint=%q", fieldValue)
case 'F':
fmt.Fprintf(&sb, " file=%q", fieldValue)
case 'L':
fmt.Fprintf(&sb, " line=%q", fieldValue)
case 'R':
fmt.Fprintf(&sb, " routine=%q", fieldValue)
default:
fmt.Fprintf(&sb, " %x=%q", fieldType, fieldValue)
}
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
type query struct {
want int
have []byte
obs messageObserver
}
func (m *query) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 1 {
return fmt.Sprintf("Query message too short raw=%x", m.have)
}
queryString, rest, ok := decodeString(m.have)
if !ok {
return fmt.Sprintf("Query invalid string raw=%x", m.have)
}
if len(rest) > 0 {
return fmt.Sprintf("Query %q (unexpected trailing data: %x)", queryString, rest)
}
m.obs.clientQuery()
return fmt.Sprintf("Query %q", queryString)
}
type syncMsg struct {
want int
have []byte
}
func (m *syncMsg) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want != 0 {
return fmt.Sprintf("Sync (unexpected payload: %x)", m.have)
}
return "Sync"
}
type authenticationRequest struct {
want int
have []byte
}
func (m *authenticationRequest) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 4 {
return fmt.Sprintf("Authentication message too short raw=%x", m.have)
}
code := byteOrder.Uint32(m.have)
switch {
case m.want == 4 && code == 0:
return "AuthenticationOk"
case m.want == 4 && code == 2:
return "AuthenticationKerberosV5"
case m.want == 4 && code == 3:
return "AuthenticationCleartextPassword"
case m.want == 8 && code == 5:
return fmt.Sprintf("AuthenticationMD5Password salt=%x", m.have[4:])
case m.want == 4 && code == 7:
return "AuthenticationGSS"
case code == 8:
return fmt.Sprintf("AuthenticationGSSContinue data=%x", m.have[4:])
case m.want == 4 && code == 9:
return "AuthenticationSSPI"
case code == 10:
name, _, ok := decodeString(m.have[4:])
if ok {
return fmt.Sprintf("AuthenticationSASL name=%q", name)
} else {
return fmt.Sprintf("AuthenticationSASL invalid name=%q", m.have[4:])
}
case code == 11:
return fmt.Sprintf("AuthenticationSASLContinue data=%x", m.have[4:])
case code == 12:
return fmt.Sprintf("AuthenticationSASLFinal data=%x", m.have[4:])
default:
return fmt.Sprintf("unrecognized Authentication message len=%d code=%d payload=%x",
m.want, code, m.have[4:])
}
}
type parameterStatus struct {
want int
have []byte
}
func (m *parameterStatus) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 2 {
return fmt.Sprintf("ParameterStatus message too short raw=%x", m.have)
}
name, rest, ok := decodeString(m.have)
if !ok {
return fmt.Sprintf("ParameterStatus message invalid name raw=%x", m.have)
}
value, rest, ok := decodeString(rest)
if !ok {
return fmt.Sprintf("ParameterStatus %s= invalid value raw=%x", name, rest)
}
return fmt.Sprintf("ParameterStatus %s=%q", name, value)
}
type rowDescription struct {
want int
have []byte
}
func (m *rowDescription) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 2 {
return fmt.Sprintf("RowDescription message too short raw=%x", m.have)
}
numFields := byteOrder.Uint16(m.have)
var sb strings.Builder
fmt.Fprintf(&sb, "RowDescription fields=%d", numFields)
p := m.have[2:]
for i := 0; i < int(numFields); i++ {
// Field name
name, rest, ok := decodeString(p)
if !ok {
fmt.Fprintf(&sb, " (error decoding field %d name)", i)
break
}
fmt.Fprintf(&sb, " [%q", name)
p = rest
// Need at least 18 bytes for the remaining fields
if len(p) < 18 {
fmt.Fprintf(&sb, " incomplete] rest=%x", p)
break
}
tableOID := byteOrder.Uint32(p[0:4])
colAttrNum := byteOrder.Uint16(p[4:6])
typeOID := byteOrder.Uint32(p[6:10])
typeSize := int16(byteOrder.Uint16(p[10:12]))
typeMod := byteOrder.Uint32(p[12:16])
formatCode := byteOrder.Uint16(p[16:18])
fmt.Fprintf(&sb, " table=%d col=%d typeOID=%d", tableOID, colAttrNum, typeOID)
if typeSize < 0 {
fmt.Fprintf(&sb, " size=var(%d)", typeSize)
} else {
fmt.Fprintf(&sb, " size=%d", typeSize)
}
fmt.Fprintf(&sb, " mod=%08x", typeMod)
switch formatCode {
case 0:
fmt.Fprintf(&sb, " fmt=text")
case 1:
fmt.Fprintf(&sb, " fmt=binary")
default:
fmt.Fprintf(&sb, " fmt=!!%d!!", formatCode)
}
fmt.Fprintf(&sb, "]")
p = p[18:]
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
type terminate struct {
want int
have []byte
}
func (m *terminate) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want != 0 {
return fmt.Sprintf("Terminate (unexpected payload: %x)", m.have)
}
return "Terminate"
}
type readyForQuery struct {
want int
have []byte
obs messageObserver
}
func (m *readyForQuery) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 1 {
return fmt.Sprintf("ReadyForQuery message too short raw=%x", m.have)
}
status := m.have[0]
m.obs.serverReadyForQuery()
switch status {
case 'I':
return "ReadyForQuery (idle)"
case 'T':
return "ReadyForQuery (in a transaction block)"
case 'E':
return "ReadyForQuery (in a vailed transaction block)"
default:
return fmt.Sprintf("ReadyForQuery status=unknown(%c/%x)", status, status)
}
}
type parameterDescription struct {
want int
have []byte
}
func (m *parameterDescription) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
if m.want < 2 {
return fmt.Sprintf("ParameterDescription message too short raw=%x", m.have)
}
numParams := byteOrder.Uint16(m.have)
var sb strings.Builder
fmt.Fprintf(&sb, "ParameterDescription params=%d", numParams)
p := m.have[2:]
if numParams > 0 {
if len(p) < int(numParams)*4 {
fmt.Fprintf(&sb, " (incomplete, need %d bytes) rest=%x", numParams*4, p)
return sb.String()
}
fmt.Fprintf(&sb, " types=[")
for i := 0; i < int(numParams); i++ {
if i > 0 {
fmt.Fprintf(&sb, ",")
}
typeOID := byteOrder.Uint32(p[i*4 : (i+1)*4])
fmt.Fprintf(&sb, "%d", typeOID)
}
fmt.Fprintf(&sb, "]")
p = p[numParams*4:]
}
if len(p) > 0 {
fmt.Fprintf(&sb, " (unexpected trailing data: %x)", p)
}
return sb.String()
}
type gssResponse struct {
want int
have []byte
}
func (m *gssResponse) Format(data []byte) string {
m.have = append(m.have, data...)
if len(m.have) < m.want {
return "..."
}
return fmt.Sprintf("GSSResponse data=%x", m.have)
}
package main
import (
"net"
"net/http"
"sync"
"time"
mm "go.withmatt.com/metrics"
"go.withmatt.com/metrics/promhttp"
)
type metrics struct {
queriesStarted *mm.Uint64Vec
readyForQuery *mm.Uint64
durations *mm.HistogramVec
}
func newMetrics() *metrics {
return &metrics{
queriesStarted: mm.NewUint64Vec("queries_started", "query_type"),
readyForQuery: mm.NewUint64("ready_for_query"),
durations: mm.NewHistogramVec("queries", "query_type"),
}
}
func (m *metrics) Handler() http.Handler {
return promhttp.Handler()
}
func (m *metrics) connMetrics(_ net.Conn) *connMetrics {
return &connMetrics{
m: m,
}
}
type connMetrics struct {
m *metrics
mu sync.Mutex
pendingQueryType string
pendingQueryStart time.Time
}
func (c *connMetrics) clientExecute() {
c.m.queriesStarted.WithLabelValues("extended").Inc()
c.mu.Lock()
c.pendingQueryType = "extended"
c.pendingQueryStart = time.Now()
c.mu.Unlock()
}
func (c *connMetrics) clientQuery() {
c.m.queriesStarted.WithLabelValues("simple").Inc()
c.mu.Lock()
c.pendingQueryType = "simple"
c.pendingQueryStart = time.Now()
c.mu.Unlock()
}
func (c *connMetrics) serverReadyForQuery() {
c.m.readyForQuery.Inc()
c.mu.Lock()
qt := c.pendingQueryType
st := c.pendingQueryStart
c.pendingQueryType = ""
c.mu.Unlock()
if qt != "" {
c.m.durations.WithLabelValues(qt).Observe(time.Since(st).Seconds())
}
}
func (c *connMetrics) closed() {
}
package main
import (
"crypto/tls"
"log"
"net"
"sync"
"github.com/jackc/pgx/v5/pgproto3"
)
func servepgx(conn net.Conn, pgAddr string, feTLS *tls.Config, beTLS backendTLSPolicy, l *log.Logger) {
defer conn.Close()
defer l.Printf("closing connection")
cLog := log.New(l.Writer(), l.Prefix()+"client ", l.Flags())
sLog := log.New(l.Writer(), l.Prefix()+"server ", l.Flags())
client := pgproto3.NewBackend(conn, conn)
clientStart, err := client.ReceiveStartupMessage()
if err != nil {
cLog.Printf("error reading client startup message: %v", err)
return
}
logPgxMessage(cLog, "recv", clientStart)
if _, ok := clientStart.(*pgproto3.SSLRequest); ok {
switch feTLS {
case nil:
cLog.Printf("--- SSL requested by client, declining ---")
if _, err := conn.Write([]byte{'N'}); err != nil {
cLog.Printf("client write error: %v", err)
return
}
default:
cLog.Printf("--- performing SSL handshake with client ---")
if _, err := conn.Write([]byte{'S'}); err != nil {
cLog.Printf("client write error: %v", err)
return
}
tls := tls.Server(conn, feTLS)
client = pgproto3.NewBackend(tls, tls)
client.Trace(&lw{cLog}, pgproto3.TracerOptions{SuppressTimestamps: true})
}
clientStart, err = client.ReceiveStartupMessage()
if err != nil {
cLog.Printf("error reading client startup message: %v", err)
return
}
logPgxMessage(cLog, "recv", clientStart)
}
bg, err := net.Dial("tcp", pgAddr)
if err != nil {
sLog.Printf("%s: error dialing backend: %v", pgAddr, err)
return
}
defer bg.Close()
server := pgproto3.NewFrontend(bg, bg)
if beTLS.Request {
sLog.Printf("server handshake: try TLS")
logPgxMessage(sLog, "send", &pgproto3.SSLRequest{})
server.Send(&pgproto3.SSLRequest{})
if err := server.Flush(); err != nil {
sLog.Printf("error writing to server: %v", err)
return
}
var sslResp [1]byte
n, err := bg.Read(sslResp[:])
if err != nil {
sLog.Printf("error reading from server: %v", err)
return
}
if n != 1 {
sLog.Printf("empty response from server: %v", err)
return
}
sLog.Printf("server handshake: SSLRequest response: '%c' (%x)", sslResp[0], sslResp[0])
switch sslResp[0] {
case 'N':
if beTLS.Require {
sLog.Printf("TLS is required but backend does not allow it")
return
}
case 'S':
sLog.Printf("--- performing SSL handshake with server ---")
tls := tls.Client(conn, beTLS.Config)
server = pgproto3.NewFrontend(tls, tls)
default:
sLog.Printf("illegal response from backend")
return
}
}
logPgxMessage(sLog, "send", clientStart)
server.Send(clientStart)
if err := server.Flush(); err != nil {
sLog.Printf("error sending server startup message: %v", err)
return
}
var wg sync.WaitGroup
wg.Go(func() {
for {
msg, err := client.Receive()
if err != nil {
cLog.Printf("error receiving next message from client: %v", err)
return
}
logPgxMessage(cLog, "recv", msg)
server.Send(msg)
if err := server.Flush(); err != nil {
sLog.Printf("error forwarding message to server: %v", err)
return
}
}
})
wg.Go(func() {
for {
msg, err := server.Receive()
if err != nil {
sLog.Printf("error receiving next message from server: %v", err)
return
}
logPgxMessage(sLog, "recv", msg)
client.Send(msg)
if err := client.Flush(); err != nil {
cLog.Printf("error forwarding message to client: %v", err)
return
}
}
})
wg.Wait()
}
func logPgxMessage(l *log.Logger, s string, msg pgproto3.Message) {
l.Printf("%s: %#v", s, msg)
}
type lw struct {
l *log.Logger
}
func (l *lw) Write(data []byte) (int, error) {
l.l.Printf("%s", data)
return len(data), nil
}
package main
import (
"bytes"
"encoding/binary"
"io"
"log"
)
var byteOrder = binary.BigEndian
type side interface {
getMessageFormatter(msgType byte, msgLen int, obs messageObserver) messageFormatter
}
// processMessages does some introspection of messages received from r before
// forwarding them to w.
//
// hi should be 1 for client->server and 0 for server->client.
func processMessages(r io.Reader, w io.Writer, m messageObserver, l *log.Logger, side side, hi int) {
buf := make([]byte, 1024*1024)
const maxHeaderSize = 5
var hdr [maxHeaderSize]byte
bodyRemaining := 0
var pretty messageFormatter
for {
n, err := r.Read(buf)
if err != nil {
l.Printf("error reading data: %v", err)
return
}
for p := buf[:n]; len(p) > 0; {
switch {
case bodyRemaining == 0:
// we need more header bytes
copied := copy(hdr[hi:], p)
hi += copied
p = p[copied:]
if hi == maxHeaderSize {
// fall through
msgType := hdr[0]
length := int(byteOrder.Uint32(hdr[1:]))
l.Printf("MSG '%c' (%x) (len=%d)", msgType, msgType, length-4)
// SSL mode is []uint16{1234,5678} or []byte{0x04, 0xd2, 0x16, 0x2f}.
if msgType == 0 && length == 8 && bytes.Equal(p[:4], []byte{0x04, 0xd2, 0x16, 0x2f}) {
l.Printf("SSL mode is not supported")
return
}
bodyRemaining = length - 4
hi = 0
pretty = side.getMessageFormatter(msgType, bodyRemaining, m)
if bodyRemaining == 0 {
l.Printf(" %s", pretty.Format(nil))
}
}
case bodyRemaining > len(p):
l.Printf(" %s", pretty.Format(p))
bodyRemaining -= len(p)
p = nil
case bodyRemaining > 0:
l.Printf(" %s", pretty.Format(p[:bodyRemaining]))
p = p[bodyRemaining:]
bodyRemaining = 0
default:
l.Printf("invalid state, bailing")
}
}
l.Printf("write(%d)", n)
if _, err := w.Write(buf[:n]); err != nil {
l.Printf("error writing data: %v", err)
return
}
}
}
// decodeString decodes a Postgres 'String' atom.
//
// > A null-terminated string (C-style string). There is no specific length
// > limitation on strings. If s is specified it is the exact value that will
// > appear, otherwise the value is variable. Eg. String, String("user").
//
// returns (string, rest, true) if found, (nil, data, false) if not.
func decodeString(data []byte) ([]byte, []byte, bool) {
i := bytes.Index(data, []byte{0})
if i == -1 {
return nil, data, false
}
return data[:i], data[i+1:], true
}
package main
import (
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"fmt"
"math/big"
"net"
"os"
"time"
)
func initSSL(certPath, keyPath string) (*tls.Config, error) {
certData, certErr := os.ReadFile(certPath)
keyData, keyErr := os.ReadFile(keyPath)
switch {
case os.IsNotExist(certErr) && os.IsNotExist(keyErr):
// fall through to the code after the switch, which generates a new cert and key.
case certErr != nil:
return nil, fmt.Errorf("%s: %v", certPath, certErr)
case keyErr != nil:
return nil, fmt.Errorf("%s: %v", keyPath, keyErr)
default:
cert, err := tls.X509KeyPair(certData, keyData)
if err != nil {
return nil, err
}
return &tls.Config{
Certificates: []tls.Certificate{cert},
}, nil
}
// Generate self-signed certificate
priv, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, fmt.Errorf("error generating private key: %w", err)
}
notBefore := time.Now()
notAfter := notBefore.Add(365 * 24 * time.Hour) // Valid for 1 year
serialNumber, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
if err != nil {
return nil, fmt.Errorf("error generating serial number: %w", err)
}
template := x509.Certificate{
SerialNumber: serialNumber,
Subject: pkix.Name{
Organization: []string{"PostgreSQL MITM Proxy"},
CommonName: "localhost",
},
NotBefore: notBefore,
NotAfter: notAfter,
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
DNSNames: []string{"localhost"},
IPAddresses: []net.IP{net.ParseIP("127.0.0.1")},
}
derBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &priv.PublicKey, priv)
if err != nil {
return nil, fmt.Errorf("error creating certificate: %w", err)
}
// Save certificate
certOut, err := os.Create(certPath)
if err != nil {
return nil, fmt.Errorf("error creating cert file: %w", err)
}
defer certOut.Close()
if err := pem.Encode(certOut, &pem.Block{Type: "CERTIFICATE", Bytes: derBytes}); err != nil {
return nil, fmt.Errorf("error writing cert: %w", err)
}
// Save private key
keyOut, err := os.Create(keyPath)
if err != nil {
return nil, fmt.Errorf("error creating key file: %w", err)
}
defer keyOut.Close()
privBytes, err := x509.MarshalPKCS8PrivateKey(priv)
if err != nil {
return nil, fmt.Errorf("error marshaling private key: %w", err)
}
if err := pem.Encode(keyOut, &pem.Block{Type: "PRIVATE KEY", Bytes: privBytes}); err != nil {
return nil, fmt.Errorf("error writing key: %w", err)
}
// Load the newly created certificate
cert, err := tls.LoadX509KeyPair(certPath, keyPath)
if err != nil {
return nil, fmt.Errorf("error loading generated certificate: %w", err)
}
return &tls.Config{
Certificates: []tls.Certificate{cert},
}, nil
}
type backendTLSPolicy struct {
Request bool
Require bool
Config *tls.Config
}
func configureBackendTLS(mode, servername string) (backendTLSPolicy, bool) {
switch mode {
case "disable":
return backendTLSPolicy{}, true
case "allow":
return backendTLSPolicy{Request: true, Config: strictBETLSConfig(servername)}, true
case "require":
return backendTLSPolicy{Request: true, Require: true, Config: strictBETLSConfig(servername)}, true
case "insecure":
return backendTLSPolicy{Request: true, Config: insecureBETLSConfig()}, true
default:
return backendTLSPolicy{}, false
}
}
func strictBETLSConfig(servername string) *tls.Config {
return &tls.Config{
ServerName: servername,
}
}
func insecureBETLSConfig() *tls.Config {
return &tls.Config{
InsecureSkipVerify: true,
}
}
#!/bin/sh
for m in "sslmode=prefer" "sslmode=require" "sslmode=disable"; do
echo "==== $m ===="
psql "host=127.0.0.1 port=5555 user=postgres password=xxx $m" -c 'select * from hi'
done
package main
import (
"log"
"net"
"time"
)
type tlsServerSpy struct {
conn net.Conn
l *log.Logger
}
func (t *tlsServerSpy) Read(b []byte) (n int, err error) {
n, err = t.conn.Read(b)
switch {
case err != nil:
t.l.Printf("tls.read => error: %v", err)
case n == 0:
t.l.Printf("tls.read => empty")
case n < 10:
t.l.Printf("tls.read => (n=%d) %v", n, b[:n])
default:
t.l.Printf("tls.read => (n=%d) %v...", n, b[:10])
}
return n, err
}
func (t *tlsServerSpy) Write(b []byte) (n int, err error) {
n, err = t.conn.Write(b)
switch {
case err != nil:
t.l.Printf("tls.write => error: %v", err)
case n == 0:
t.l.Printf("tls.write => empty")
case n < 10:
t.l.Printf("tls.write => (n=%d) %v", n, b[:n])
default:
t.l.Printf("tls.write => (n=%d) %v...", n, b[:10])
}
return n, err
}
func (t *tlsServerSpy) Close() error { return t.conn.Close() }
func (t *tlsServerSpy) LocalAddr() net.Addr { return t.conn.LocalAddr() }
func (t *tlsServerSpy) RemoteAddr() net.Addr { return t.conn.RemoteAddr() }
func (t *tlsServerSpy) SetDeadline(dl time.Time) error { return t.conn.SetDeadline(dl) }
func (t *tlsServerSpy) SetReadDeadline(dl time.Time) error { return t.conn.SetReadDeadline(dl) }
func (t *tlsServerSpy) SetWriteDeadline(dl time.Time) error { return t.conn.SetWriteDeadline(dl) }
type tlsClientSpy = tlsServerSpy
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment