feat: add PawSQL docs examples and image CI
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 2m28s

This commit is contained in:
2026-09-15 18:49:26 -04:00
commit 865f7c26c9
25 changed files with 2811 additions and 0 deletions

View File

@@ -0,0 +1,40 @@
// Package pgwire contains the minimal PostgreSQL wire framing PawSQL needs before proxying.
package pgwire
import (
"encoding/binary"
"errors"
"fmt"
"io"
)
const (
sslRequestLength uint32 = 8
sslRequestCode uint32 = 80877103
)
var ErrNotSSLRequest = errors.New("expected PostgreSQL SSLRequest")
// ReadSSLRequest validates the eight-byte PostgreSQL SSL negotiation request.
func ReadSSLRequest(reader io.Reader) error {
var request [8]byte
if _, err := io.ReadFull(reader, request[:]); err != nil {
return fmt.Errorf("read PostgreSQL SSLRequest: %w", err)
}
if length := binary.BigEndian.Uint32(request[0:4]); length != sslRequestLength {
return fmt.Errorf("%w: invalid length %d", ErrNotSSLRequest, length)
}
if code := binary.BigEndian.Uint32(request[4:8]); code != sslRequestCode {
return fmt.Errorf("%w: invalid request code %d", ErrNotSSLRequest, code)
}
return nil
}
// SSLRequest returns the exact client preamble used to request TLS from PostgreSQL.
// It exists to make protocol-level tests and clients unambiguous.
func SSLRequest() [8]byte {
var request [8]byte
binary.BigEndian.PutUint32(request[0:4], sslRequestLength)
binary.BigEndian.PutUint32(request[4:8], sslRequestCode)
return request
}

108
internal/pgwire/startup.go Normal file
View File

@@ -0,0 +1,108 @@
package pgwire
import (
"bytes"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
)
const (
startupProtocolVersion uint32 = 196608
maxStartupMessageSize = 64 << 10
)
var (
ErrInvalidStartupMessage = errors.New("invalid PostgreSQL StartupMessage")
ErrMissingDatabase = errors.New("PostgreSQL StartupMessage has no database")
)
// StartupMessage is the PostgreSQL protocol preamble sent after TLS negotiation.
// PawSQL reads it only when TLS SNI is unavailable for routing.
type StartupMessage struct {
raw []byte
Database string
}
// ReadStartupMessage reads the bounded protocol v3 StartupMessage and exposes
// its database parameter without inspecting later PostgreSQL traffic.
func ReadStartupMessage(reader io.Reader) (StartupMessage, error) {
var header [4]byte
if _, err := io.ReadFull(reader, header[:]); err != nil {
return StartupMessage{}, fmt.Errorf("read StartupMessage length: %w", err)
}
length := binary.BigEndian.Uint32(header[:])
if length < 8 || length > maxStartupMessageSize {
return StartupMessage{}, fmt.Errorf("%w: invalid length %d", ErrInvalidStartupMessage, length)
}
raw := make([]byte, length)
copy(raw, header[:])
if _, err := io.ReadFull(reader, raw[4:]); err != nil {
return StartupMessage{}, fmt.Errorf("read StartupMessage body: %w", err)
}
if version := binary.BigEndian.Uint32(raw[4:8]); version != startupProtocolVersion {
return StartupMessage{}, fmt.Errorf("%w: unsupported protocol version %d", ErrInvalidStartupMessage, version)
}
message := StartupMessage{raw: raw}
databaseSeen := false
terminated := false
for offset := 8; offset < len(raw); {
if raw[offset] == 0 {
if offset != len(raw)-1 {
return StartupMessage{}, fmt.Errorf("%w: unexpected parameter terminator", ErrInvalidStartupMessage)
}
terminated = true
break
}
keyStart := offset
keyEnd := bytes.IndexByte(raw[keyStart:], 0)
if keyEnd < 0 {
return StartupMessage{}, fmt.Errorf("%w: unterminated parameter name", ErrInvalidStartupMessage)
}
keyEnd += keyStart
valueStart := keyEnd + 1
valueEnd := bytes.IndexByte(raw[valueStart:], 0)
if valueEnd < 0 {
return StartupMessage{}, fmt.Errorf("%w: unterminated parameter value", ErrInvalidStartupMessage)
}
valueEnd += valueStart
if string(raw[keyStart:keyEnd]) == "database" {
if databaseSeen {
return StartupMessage{}, fmt.Errorf("%w: duplicate database parameter", ErrInvalidStartupMessage)
}
message.Database = string(raw[valueStart:valueEnd])
databaseSeen = true
}
offset = valueEnd + 1
}
if !terminated {
return StartupMessage{}, fmt.Errorf("%w: missing parameter terminator", ErrInvalidStartupMessage)
}
if !databaseSeen || message.Database == "" {
return StartupMessage{}, ErrMissingDatabase
}
return message, nil
}
// Bytes returns the exact StartupMessage bytes for replay to the backend.
func (m StartupMessage) Bytes() []byte {
return m.raw
}
// Replay returns a connection that yields prefix before reading from conn.
func Replay(conn net.Conn, prefix []byte) net.Conn {
return &replayConn{Conn: conn, reader: io.MultiReader(bytes.NewReader(prefix), conn)}
}
type replayConn struct {
net.Conn
reader io.Reader
}
func (c *replayConn) Read(buffer []byte) (int, error) {
return c.reader.Read(buffer)
}

View File

@@ -0,0 +1,76 @@
package pgwire
import (
"bytes"
"encoding/binary"
"errors"
"io"
"net"
"testing"
"time"
)
func TestStartupMessageReadsDatabase(t *testing.T) {
raw := startupBytes("analytics", "ruckstack")
message, err := ReadStartupMessage(bytes.NewReader(raw))
if err != nil {
t.Fatal(err)
}
if message.Database != "analytics" {
t.Fatalf("Database = %q", message.Database)
}
}
func TestReadStartupMessageRejectsMissingDatabase(t *testing.T) {
_, err := ReadStartupMessage(bytes.NewReader(startupBytes("", "ruckstack")))
if !errors.Is(err, ErrMissingDatabase) {
t.Errorf("ReadStartupMessage() error = %v, want ErrMissingDatabase", err)
}
}
func TestReplayForwardsStartupAndUnderlyingStream(t *testing.T) {
reader, writer := io.Pipe()
defer reader.Close()
go func() {
_, _ = writer.Write([]byte("tail"))
_ = writer.Close()
}()
connection := &readOnlyConn{Reader: reader}
replayed := Replay(connection, []byte("startup"))
got, err := io.ReadAll(replayed)
if err != nil {
t.Fatal(err)
}
if string(got) != "startuptail" {
t.Errorf("Replay() = %q", got)
}
}
func startupBytes(database, user string) []byte {
body := make([]byte, 4)
binary.BigEndian.PutUint32(body, startupProtocolVersion)
if database != "" {
body = append(body, "database"...)
body = append(body, 0)
body = append(body, database...)
body = append(body, 0)
}
body = append(body, "user"...)
body = append(body, 0)
body = append(body, user...)
body = append(body, 0, 0)
message := make([]byte, 4, 4+len(body))
message = append(message, body...)
binary.BigEndian.PutUint32(message, uint32(len(message)))
return message
}
type readOnlyConn struct{ io.Reader }
func (c *readOnlyConn) Write([]byte) (int, error) { return 0, io.ErrClosedPipe }
func (c *readOnlyConn) Close() error { return nil }
func (c *readOnlyConn) LocalAddr() net.Addr { return nil }
func (c *readOnlyConn) RemoteAddr() net.Addr { return nil }
func (c *readOnlyConn) SetDeadline(time.Time) error { return nil }
func (c *readOnlyConn) SetReadDeadline(time.Time) error { return nil }
func (c *readOnlyConn) SetWriteDeadline(time.Time) error { return nil }