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

38
internal/config/config.go Normal file
View File

@@ -0,0 +1,38 @@
// Package config defines PawSQL's declarative runtime configuration.
package config
import "time"
// Config is the typed representation of a Barkfile.
type Config struct {
Listen string
TLS TLSConfig
Databases []DatabaseConfig
}
// TLSConfig identifies the certificate material used to terminate client TLS.
type TLSConfig struct {
CertFile string
KeyFile string
}
// DatabaseConfig describes one PostgreSQL database. Hostname is optional:
// clients without TLS SNI route by the database's configured Name. Exactly one
// of Upstream and Postgres must be configured.
type DatabaseConfig struct {
Name string
Hostname string
Upstream string
Postgres *PostgresConfig
}
// PostgresConfig declares a PawSQL-managed PostgreSQL container. IdleTimeout
// starts after the last proxied client session closes. TrafficIdleTimeout
// applies to open sessions with no proxied bytes. Zero disables either timeout.
type PostgresConfig struct {
Image string
Volume string
PasswordEnv string
IdleTimeout time.Duration
TrafficIdleTimeout time.Duration
}

390
internal/config/parser.go Normal file
View File

@@ -0,0 +1,390 @@
package config
import (
"fmt"
"os"
"strings"
"time"
"unicode"
)
func ParseFile(path string) (Config, error) {
contents, err := os.ReadFile(path)
if err != nil {
return Config{}, fmt.Errorf("read Barkfile: %w", err)
}
return Parse(contents)
}
// Parse parses Barkfile contents. Syntax deliberately covers only PawSQL's MVP.
func Parse(input []byte) (Config, error) {
tokens, err := lex(string(input))
if err != nil {
return Config{}, err
}
p := parser{tokens: tokens}
return p.parse()
}
type tokenKind uint8
const (
tokenWord tokenKind = iota
tokenOpenBrace
tokenCloseBrace
tokenNewline
tokenEOF
)
type token struct {
kind tokenKind
text string
line int
}
type parser struct {
tokens []token
index int
}
func (p *parser) parse() (Config, error) {
p.skipNewlines()
if err := p.expectWord("pawsql"); err != nil {
return Config{}, err
}
if err := p.expect(tokenOpenBrace, "{"); err != nil {
return Config{}, err
}
var cfg Config
for {
p.skipNewlines()
if p.current().kind == tokenCloseBrace {
p.index++
break
}
if p.current().kind == tokenEOF {
return Config{}, p.errorf("expected } to close pawsql block")
}
if p.current().kind != tokenWord {
return Config{}, p.errorf("expected directive")
}
switch p.current().text {
case "listen":
if cfg.Listen != "" {
return Config{}, p.errorf("listen may only be specified once")
}
p.index++
value, err := p.value("listen address")
if err != nil {
return Config{}, err
}
cfg.Listen = value
if err := p.endLine(); err != nil {
return Config{}, err
}
case "tls":
if cfg.TLS.CertFile != "" || cfg.TLS.KeyFile != "" {
return Config{}, p.errorf("tls may only be specified once")
}
p.index++
tlsConfig, err := p.parseTLS()
if err != nil {
return Config{}, err
}
cfg.TLS = tlsConfig
case "database":
p.index++
name, err := p.value("database name")
if err != nil {
return Config{}, err
}
database, err := p.parseDatabase(name)
if err != nil {
return Config{}, err
}
cfg.Databases = append(cfg.Databases, database)
default:
return Config{}, p.errorf("unknown directive %q", p.current().text)
}
}
p.skipNewlines()
if p.current().kind != tokenEOF {
return Config{}, p.errorf("unexpected content after pawsql block")
}
return cfg, nil
}
func (p *parser) parseTLS() (TLSConfig, error) {
if err := p.expect(tokenOpenBrace, "{"); err != nil {
return TLSConfig{}, err
}
var cfg TLSConfig
for {
p.skipNewlines()
if p.current().kind == tokenCloseBrace {
p.index++
return cfg, nil
}
if p.current().kind == tokenEOF {
return TLSConfig{}, p.errorf("expected } to close tls block")
}
name := p.current()
if name.kind != tokenWord {
return TLSConfig{}, p.errorf("expected tls directive")
}
p.index++
value, err := p.value("tls value")
if err != nil {
return TLSConfig{}, err
}
switch name.text {
case "cert":
if cfg.CertFile != "" {
return TLSConfig{}, fmt.Errorf("line %d: cert may only be specified once", name.line)
}
cfg.CertFile = value
case "key":
if cfg.KeyFile != "" {
return TLSConfig{}, fmt.Errorf("line %d: key may only be specified once", name.line)
}
cfg.KeyFile = value
default:
return TLSConfig{}, fmt.Errorf("line %d: unknown tls directive %q", name.line, name.text)
}
if err := p.endLine(); err != nil {
return TLSConfig{}, err
}
}
}
func (p *parser) parseDatabase(name string) (DatabaseConfig, error) {
if err := p.expect(tokenOpenBrace, "{"); err != nil {
return DatabaseConfig{}, err
}
database := DatabaseConfig{Name: name}
for {
p.skipNewlines()
if p.current().kind == tokenCloseBrace {
p.index++
return database, nil
}
if p.current().kind == tokenEOF {
return DatabaseConfig{}, p.errorf("expected } to close database block")
}
directive := p.current()
if directive.kind != tokenWord {
return DatabaseConfig{}, p.errorf("expected database directive")
}
p.index++
if directive.text == "postgres" {
if database.Postgres != nil {
return DatabaseConfig{}, fmt.Errorf("line %d: postgres may only be specified once", directive.line)
}
postgres, err := p.parsePostgres()
if err != nil {
return DatabaseConfig{}, err
}
database.Postgres = &postgres
continue
}
value, err := p.value("database value")
if err != nil {
return DatabaseConfig{}, err
}
switch directive.text {
case "hostname":
if database.Hostname != "" {
return DatabaseConfig{}, fmt.Errorf("line %d: hostname may only be specified once", directive.line)
}
database.Hostname = value
case "upstream":
if database.Upstream != "" {
return DatabaseConfig{}, fmt.Errorf("line %d: upstream may only be specified once", directive.line)
}
database.Upstream = value
default:
return DatabaseConfig{}, fmt.Errorf("line %d: unknown database directive %q", directive.line, directive.text)
}
if err := p.endLine(); err != nil {
return DatabaseConfig{}, err
}
}
}
func (p *parser) parsePostgres() (PostgresConfig, error) {
if err := p.expect(tokenOpenBrace, "{"); err != nil {
return PostgresConfig{}, err
}
var cfg PostgresConfig
idleTimeoutSet := false
trafficIdleTimeoutSet := false
for {
p.skipNewlines()
if p.current().kind == tokenCloseBrace {
p.index++
return cfg, nil
}
if p.current().kind == tokenEOF {
return PostgresConfig{}, p.errorf("expected } to close postgres block")
}
directive := p.current()
if directive.kind != tokenWord {
return PostgresConfig{}, p.errorf("expected postgres directive")
}
p.index++
value, err := p.value("postgres value")
if err != nil {
return PostgresConfig{}, err
}
switch directive.text {
case "image":
if cfg.Image != "" {
return PostgresConfig{}, fmt.Errorf("line %d: image may only be specified once", directive.line)
}
cfg.Image = value
case "volume":
if cfg.Volume != "" {
return PostgresConfig{}, fmt.Errorf("line %d: volume may only be specified once", directive.line)
}
cfg.Volume = value
case "password_env":
if cfg.PasswordEnv != "" {
return PostgresConfig{}, fmt.Errorf("line %d: password_env may only be specified once", directive.line)
}
cfg.PasswordEnv = value
case "idle_timeout":
if idleTimeoutSet {
return PostgresConfig{}, fmt.Errorf("line %d: idle_timeout may only be specified once", directive.line)
}
idleTimeoutSet = true
timeout, err := time.ParseDuration(value)
if err != nil {
return PostgresConfig{}, fmt.Errorf("line %d: invalid idle_timeout %q: %w", directive.line, value, err)
}
cfg.IdleTimeout = timeout
case "traffic_idle_timeout":
if trafficIdleTimeoutSet {
return PostgresConfig{}, fmt.Errorf("line %d: traffic_idle_timeout may only be specified once", directive.line)
}
trafficIdleTimeoutSet = true
timeout, err := time.ParseDuration(value)
if err != nil {
return PostgresConfig{}, fmt.Errorf("line %d: invalid traffic_idle_timeout %q: %w", directive.line, value, err)
}
cfg.TrafficIdleTimeout = timeout
default:
return PostgresConfig{}, fmt.Errorf("line %d: unknown postgres directive %q", directive.line, directive.text)
}
if err := p.endLine(); err != nil {
return PostgresConfig{}, err
}
}
}
func (p *parser) value(description string) (string, error) {
current := p.current()
if current.kind != tokenWord {
return "", p.errorf("expected %s", description)
}
p.index++
return current.text, nil
}
func (p *parser) endLine() error {
if p.current().kind == tokenNewline {
p.skipNewlines()
return nil
}
if p.current().kind == tokenCloseBrace || p.current().kind == tokenEOF {
return nil
}
return p.errorf("expected end of line")
}
func (p *parser) expectWord(word string) error {
if p.current().kind != tokenWord || p.current().text != word {
return p.errorf("expected %q", word)
}
p.index++
return nil
}
func (p *parser) expect(kind tokenKind, name string) error {
if p.current().kind != kind {
return p.errorf("expected %s", name)
}
p.index++
return nil
}
func (p *parser) skipNewlines() {
for p.current().kind == tokenNewline {
p.index++
}
}
func (p *parser) current() token { return p.tokens[p.index] }
func (p *parser) errorf(format string, args ...any) error {
return fmt.Errorf("line %d: %s", p.current().line, fmt.Sprintf(format, args...))
}
func lex(input string) ([]token, error) {
var tokens []token
line := 1
for index := 0; index < len(input); {
ch := input[index]
switch {
case ch == '#':
for index < len(input) && input[index] != '\n' {
index++
}
case ch == '\n':
tokens = append(tokens, token{kind: tokenNewline, line: line})
index++
line++
case unicode.IsSpace(rune(ch)):
index++
case ch == '{':
tokens = append(tokens, token{kind: tokenOpenBrace, text: "{", line: line})
index++
case ch == '}':
tokens = append(tokens, token{kind: tokenCloseBrace, text: "}", line: line})
index++
case ch == '"':
startLine := line
index++
var value strings.Builder
terminated := false
for index < len(input) {
if input[index] == '\n' {
return nil, fmt.Errorf("line %d: unterminated quoted string", startLine)
}
if input[index] == '"' {
index++
terminated = true
break
}
if input[index] == '\\' && index+1 < len(input) {
index++
value.WriteByte(input[index])
index++
continue
}
value.WriteByte(input[index])
index++
}
if !terminated {
return nil, fmt.Errorf("line %d: unterminated quoted string", startLine)
}
tokens = append(tokens, token{kind: tokenWord, text: value.String(), line: startLine})
default:
start := index
for index < len(input) && !unicode.IsSpace(rune(input[index])) && !strings.ContainsRune("{}#\"", rune(input[index])) {
index++
}
if start == index {
return nil, fmt.Errorf("line %d: unexpected character %q", line, input[index])
}
tokens = append(tokens, token{kind: tokenWord, text: input[start:index], line: line})
}
}
tokens = append(tokens, token{kind: tokenEOF, line: line})
return tokens, nil
}

View File

@@ -0,0 +1,254 @@
package config
import (
"errors"
"strings"
"testing"
"time"
)
func TestParseValidConfiguration(t *testing.T) {
cfg, err := Parse([]byte(`pawsql {
listen :5432
tls {
cert "./certs/fullchain.pem"
key ./certs/privkey.pem
}
database gramps {
hostname Gramps.PawSQL.Barkstack.Dev
upstream 192.168.27.10:5432
}
}`))
if err != nil {
t.Fatalf("Parse() error = %v", err)
}
if cfg.Listen != ":5432" {
t.Errorf("Listen = %q", cfg.Listen)
}
if cfg.TLS.CertFile != "./certs/fullchain.pem" || cfg.TLS.KeyFile != "./certs/privkey.pem" {
t.Errorf("TLS = %#v", cfg.TLS)
}
if len(cfg.Databases) != 1 {
t.Fatalf("databases = %d", len(cfg.Databases))
}
if got := cfg.Databases[0]; got.Name != "gramps" || got.Hostname != "Gramps.PawSQL.Barkstack.Dev" || got.Upstream != "192.168.27.10:5432" {
t.Errorf("database = %#v", got)
}
if err := cfg.Validate(); err != nil {
t.Fatalf("Validate() error = %v", err)
}
}
func TestParseCommentsAndMultipleDatabases(t *testing.T) {
cfg, err := Parse([]byte(`# external comment
pawsql {
listen :5432 # client port
tls { cert cert.pem # inline
key key.pem }
database one { hostname one.pawsql.test
upstream postgres-one:5432 }
# route another application
database two { hostname two.pawsql.test
upstream 127.0.0.1:55432 }
}`))
if err != nil {
t.Fatalf("Parse() error = %v", err)
}
if len(cfg.Databases) != 2 {
t.Fatalf("databases = %d, want 2", len(cfg.Databases))
}
if cfg.Databases[1].Name != "two" || cfg.Databases[1].Upstream != "127.0.0.1:55432" {
t.Errorf("second database = %#v", cfg.Databases[1])
}
}
func TestValidateMissingFieldsAndDuplicateHostname(t *testing.T) {
cfg, err := Parse([]byte(`pawsql {
listen :5432
tls { cert cert.pem }
database one {
hostname FOO.pawsql.test
}
database two {
hostname foo.pawsql.test
upstream postgres-two:5432
}
}`))
if err != nil {
t.Fatalf("Parse() error = %v", err)
}
err = cfg.Validate()
if err == nil {
t.Fatal("Validate() error = nil")
}
for _, want := range []string{"TLS private key is required", "upstream or postgres is required", "duplicate hostname"} {
if !strings.Contains(err.Error(), want) {
t.Errorf("Validate() error = %q, missing %q", err, want)
}
}
}
func TestValidateRejectsMalformedUpstreamAddress(t *testing.T) {
cfg := Config{
Listen: ":5432",
TLS: TLSConfig{CertFile: "cert.pem", KeyFile: "key.pem"},
Databases: []DatabaseConfig{{
Name: "foo",
Hostname: "foo.pawsql.test",
Upstream: "postgres-foo",
}},
}
err := cfg.Validate()
if err == nil || !strings.Contains(err.Error(), "missing port") {
t.Fatalf("Validate() error = %v, want malformed upstream address", err)
}
}
func TestValidateAllowsDatabaseWithoutHostname(t *testing.T) {
cfg := Config{
Listen: ":5432",
TLS: TLSConfig{CertFile: "cert.pem", KeyFile: "key.pem"},
Databases: []DatabaseConfig{{
Name: "analytics",
Upstream: "postgres-foo:5432",
}},
}
if err := cfg.Validate(); err != nil {
t.Fatalf("Validate() error = %v", err)
}
}
func TestValidateRejectsDuplicateDatabaseName(t *testing.T) {
cfg := Config{
Listen: ":5432",
TLS: TLSConfig{CertFile: "cert.pem", KeyFile: "key.pem"},
Databases: []DatabaseConfig{
{Name: "analytics", Upstream: "postgres-foo:5432"},
{Name: "analytics", Upstream: "postgres-bar:5432"},
},
}
err := cfg.Validate()
if err == nil || !strings.Contains(err.Error(), "duplicate database name") {
t.Fatalf("Validate() error = %v, want duplicate database name", err)
}
}
func TestParsePostgresContainer(t *testing.T) {
cfg, err := Parse([]byte(`pawsql {
listen :5432
tls {
cert cert.pem
key key.pem
}
database analytics {
postgres {
image postgres:18
volume analytics-data
password_env ANALYTICS_POSTGRES_PASSWORD
idle_timeout 15m
traffic_idle_timeout 1h
}
}
}`))
if err != nil {
t.Fatal(err)
}
database := cfg.Databases[0]
if database.Postgres == nil || database.Postgres.Image != "postgres:18" || database.Postgres.Volume != "analytics-data" || database.Postgres.PasswordEnv != "ANALYTICS_POSTGRES_PASSWORD" || database.Postgres.IdleTimeout != 15*time.Minute || database.Postgres.TrafficIdleTimeout != time.Hour {
t.Errorf("Postgres = %#v", database.Postgres)
}
if err := cfg.Validate(); err != nil {
t.Fatalf("Validate() error = %v", err)
}
}
func TestParseRejectsInvalidPostgresIdleTimeout(t *testing.T) {
_, err := Parse([]byte(`pawsql {
listen :5432
tls {
cert cert.pem
key key.pem
}
database analytics {
postgres {
image postgres:18
volume analytics-data
password_env ANALYTICS_POSTGRES_PASSWORD
idle_timeout whenever
}
}
}`))
if err == nil || !strings.Contains(err.Error(), "invalid idle_timeout") {
t.Fatalf("Parse() error = %v, want invalid idle_timeout", err)
}
}
func TestParseRejectsInvalidPostgresTrafficIdleTimeout(t *testing.T) {
_, err := Parse([]byte(`pawsql {
listen :5432
tls {
cert cert.pem
key key.pem
}
database analytics {
postgres {
image postgres:18
volume analytics-data
password_env ANALYTICS_POSTGRES_PASSWORD
traffic_idle_timeout whenever
}
}
}`))
if err == nil || !strings.Contains(err.Error(), "invalid traffic_idle_timeout") {
t.Fatalf("Parse() error = %v, want invalid traffic_idle_timeout", err)
}
}
func TestValidateAllowsSupportedPostgresImages(t *testing.T) {
for _, image := range []string{"postgres:16", "postgres:17", "postgres:18"} {
cfg := Config{
Listen: ":5432",
TLS: TLSConfig{CertFile: "cert.pem", KeyFile: "key.pem"},
Databases: []DatabaseConfig{{
Name: "analytics",
Postgres: &PostgresConfig{Image: image, Volume: "analytics-data", PasswordEnv: "ANALYTICS_POSTGRES_PASSWORD"},
}},
}
if err := cfg.Validate(); err != nil {
t.Errorf("Validate() image %q error = %v", image, err)
}
}
}
func TestValidateRejectsUnsupportedPostgresImage(t *testing.T) {
cfg := Config{
Listen: ":5432",
TLS: TLSConfig{CertFile: "cert.pem", KeyFile: "key.pem"},
Databases: []DatabaseConfig{{
Name: "analytics",
Postgres: &PostgresConfig{Image: "postgres:15", Volume: "analytics-data", PasswordEnv: "ANALYTICS_POSTGRES_PASSWORD"},
}},
}
err := cfg.Validate()
if err == nil || !strings.Contains(err.Error(), "postgres:16") {
t.Fatalf("Validate() error = %v, want unsupported PostgreSQL image", err)
}
}
func TestParseMalformedBlocksReportLine(t *testing.T) {
_, err := Parse([]byte("pawsql {\n tls {\n cert cert.pem\n"))
if err == nil {
t.Fatal("Parse() error = nil")
}
if !strings.Contains(err.Error(), "line 4") {
t.Errorf("error = %q, want line number", err)
}
_, err = Parse([]byte("pawsql {\n listen :5432 unexpected\n}"))
if err == nil {
t.Fatal("Parse() error = nil for trailing directive")
}
if errors.Is(err, nil) {
t.Fatal("unexpected nil error")
}
}

115
internal/config/validate.go Normal file
View File

@@ -0,0 +1,115 @@
package config
import (
"errors"
"fmt"
"net"
"strconv"
"strings"
)
// NormalizeHostname returns the canonical lookup form for a DNS hostname.
func NormalizeHostname(hostname string) string {
return strings.TrimSuffix(strings.ToLower(strings.TrimSpace(hostname)), ".")
}
// Validate verifies configuration invariants that do not require opening files.
func (c Config) Validate() error {
var errs []error
if strings.TrimSpace(c.Listen) == "" {
errs = append(errs, errors.New("listen address is required"))
} else if err := validateListenAddress(c.Listen); err != nil {
errs = append(errs, fmt.Errorf("listen address %q: %w", c.Listen, err))
}
if strings.TrimSpace(c.TLS.CertFile) == "" {
errs = append(errs, errors.New("TLS certificate is required"))
}
if strings.TrimSpace(c.TLS.KeyFile) == "" {
errs = append(errs, errors.New("TLS private key is required"))
}
if len(c.Databases) == 0 {
errs = append(errs, errors.New("at least one database route is required"))
}
seenHostnames := make(map[string]string, len(c.Databases))
seenNames := make(map[string]struct{}, len(c.Databases))
for _, database := range c.Databases {
name := strings.TrimSpace(database.Name)
if name == "" {
errs = append(errs, errors.New("database name is required"))
} else if _, exists := seenNames[name]; exists {
errs = append(errs, fmt.Errorf("duplicate database name %q", name))
} else {
seenNames[name] = struct{}{}
}
hostname := NormalizeHostname(database.Hostname)
if hostname != "" {
if existing, ok := seenHostnames[hostname]; ok {
errs = append(errs, fmt.Errorf("duplicate hostname %q for databases %q and %q", hostname, existing, name))
} else {
seenHostnames[hostname] = name
}
}
if database.Postgres != nil {
if database.Upstream != "" {
errs = append(errs, fmt.Errorf("database %q: upstream and postgres cannot both be configured", name))
}
if err := validatePostgres(*database.Postgres); err != nil {
errs = append(errs, fmt.Errorf("database %q: postgres: %w", name, err))
}
} else if strings.TrimSpace(database.Upstream) == "" {
errs = append(errs, fmt.Errorf("database %q: upstream or postgres is required", name))
} else if err := validateAddress(database.Upstream); err != nil {
errs = append(errs, fmt.Errorf("database %q: upstream %q: %w", name, database.Upstream, err))
}
}
return errors.Join(errs...)
}
func validatePostgres(postgres PostgresConfig) error {
switch postgres.Image {
case "postgres:16", "postgres:17", "postgres:18":
default:
return fmt.Errorf("image %q must be postgres:16, postgres:17, or postgres:18", postgres.Image)
}
if strings.TrimSpace(postgres.Volume) == "" {
return errors.New("volume is required")
}
if strings.TrimSpace(postgres.PasswordEnv) == "" {
return errors.New("password_env is required")
}
if postgres.IdleTimeout < 0 {
return errors.New("idle_timeout cannot be negative")
}
if postgres.TrafficIdleTimeout < 0 {
return errors.New("traffic_idle_timeout cannot be negative")
}
return nil
}
func validateListenAddress(address string) error {
_, port, err := net.SplitHostPort(address)
if err != nil {
return err
}
portNumber, err := strconv.ParseUint(port, 10, 16)
if err != nil || portNumber == 0 {
return errors.New("port must be between 1 and 65535")
}
return nil
}
func validateAddress(address string) error {
host, port, err := net.SplitHostPort(address)
if err != nil {
return err
}
if strings.TrimSpace(host) == "" {
return errors.New("host is required")
}
portNumber, err := strconv.ParseUint(port, 10, 16)
if err != nil || portNumber == 0 {
return errors.New("port must be between 1 and 65535")
}
return nil
}

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 }

View File

@@ -0,0 +1,217 @@
// Package postgres provisions PawSQL-managed PostgreSQL containers through the Docker CLI.
package postgres
import (
"bytes"
"context"
"fmt"
"log/slog"
"net"
"os"
"os/exec"
"strings"
"time"
"github.com/barkstack/pawsql/internal/config"
)
const managedDatabaseLabel = "io.barkstack.pawsql.database"
// Provisioner ensures configured PostgreSQL containers exist and exposes their
// loopback-published PostgreSQL address to PawSQL's static router.
type Provisioner struct {
DockerPath string
Logger *slog.Logger
}
// NewProvisioner creates a Docker CLI-backed provisioner.
func NewProvisioner(logger *slog.Logger) *Provisioner {
if logger == nil {
logger = slog.Default()
}
return &Provisioner{DockerPath: "docker", Logger: logger}
}
// Ensure provisions each PostgreSQL-backed database and returns configuration
// with its discovered upstream addresses. Existing containers and volumes are
// adopted only when they carry PawSQL's matching ownership label.
func (p *Provisioner) Ensure(ctx context.Context, cfg config.Config) (config.Config, error) {
resolved := cfg
resolved.Databases = append([]config.DatabaseConfig(nil), cfg.Databases...)
for index := range resolved.Databases {
database := &resolved.Databases[index]
if database.Postgres == nil {
continue
}
address, err := p.EnsureDatabase(ctx, database.Name, *database.Postgres)
if err != nil {
return config.Config{}, fmt.Errorf("provision database %q: %w", database.Name, err)
}
database.Upstream = address
}
return resolved, nil
}
// EnsureDatabase creates or starts one managed PostgreSQL database and waits
// until its loopback-published address accepts connections.
func (p *Provisioner) EnsureDatabase(ctx context.Context, database string, postgres config.PostgresConfig) (string, error) {
name := containerName(database)
if _, err := p.run(ctx, nil, "volume", "create", postgres.Volume); err != nil {
return "", fmt.Errorf("create volume %q: %w", postgres.Volume, err)
}
label, exists, err := p.containerLabel(ctx, name)
if err != nil {
return "", err
}
if !exists {
password, ok := os.LookupEnv(postgres.PasswordEnv)
if !ok || password == "" {
return "", fmt.Errorf("environment variable %q is required to create the container", postgres.PasswordEnv)
}
p.Logger.Info("creating PostgreSQL container", "database", database, "container", name, "image", postgres.Image, "volume", postgres.Volume)
mount := postgresDataMount(postgres.Image)
_, err := p.run(ctx, []string{"POSTGRES_PASSWORD=" + password},
"container", "create",
"--name", name,
"--label", managedDatabaseLabel+"="+database,
"--env", "POSTGRES_DB="+database,
"--env", "POSTGRES_USER="+database,
"--env", "POSTGRES_PASSWORD",
"--volume", postgres.Volume+":"+mount,
"--publish", "127.0.0.1::5432",
postgres.Image,
)
if err != nil {
return "", fmt.Errorf("create container %q: %w", name, err)
}
} else if label != database {
return "", fmt.Errorf("container %q belongs to %q, not PawSQL database %q", name, label, database)
}
running, err := p.containerRunning(ctx, name)
if err != nil {
return "", err
}
if !running {
p.Logger.Info("starting PostgreSQL container", "database", database, "container", name)
if _, err := p.run(ctx, nil, "container", "start", name); err != nil {
return "", fmt.Errorf("start container %q: %w", name, err)
}
}
port, err := p.hostPort(ctx, name)
if err != nil {
return "", err
}
address := net.JoinHostPort("127.0.0.1", port)
if err := waitForPostgres(ctx, address); err != nil {
return "", fmt.Errorf("wait for PostgreSQL container %q: %w", name, err)
}
return address, nil
}
// StopDatabase stops a managed PostgreSQL container without removing its data volume.
func (p *Provisioner) StopDatabase(ctx context.Context, database string) error {
name := containerName(database)
label, exists, err := p.containerLabel(ctx, name)
if err != nil {
return err
}
if !exists {
return nil
}
if label != database {
return fmt.Errorf("container %q belongs to %q, not PawSQL database %q", name, label, database)
}
running, err := p.containerRunning(ctx, name)
if err != nil {
return err
}
if !running {
return nil
}
p.Logger.Info("stopping idle PostgreSQL container", "database", database, "container", name)
if _, err := p.run(ctx, nil, "container", "stop", name); err != nil {
return fmt.Errorf("stop container %q: %w", name, err)
}
return nil
}
func waitForPostgres(ctx context.Context, address string) error {
deadline := time.NewTimer(30 * time.Second)
defer deadline.Stop()
for {
connection, err := (&net.Dialer{Timeout: time.Second}).DialContext(ctx, "tcp", address)
if err == nil {
return connection.Close()
}
retry := time.NewTimer(200 * time.Millisecond)
select {
case <-ctx.Done():
retry.Stop()
return ctx.Err()
case <-deadline.C:
retry.Stop()
return err
case <-retry.C:
}
}
}
func (p *Provisioner) containerLabel(ctx context.Context, name string) (label string, exists bool, err error) {
output, err := p.run(ctx, nil, "container", "inspect", "--format", "{{ index .Config.Labels \""+managedDatabaseLabel+"\" }}", name)
if err != nil {
if strings.Contains(err.Error(), "No such container") {
return "", false, nil
}
return "", false, fmt.Errorf("inspect container %q: %w", name, err)
}
return strings.TrimSpace(output), true, nil
}
func (p *Provisioner) containerRunning(ctx context.Context, name string) (bool, error) {
output, err := p.run(ctx, nil, "container", "inspect", "--format", "{{.State.Running}}", name)
if err != nil {
return false, fmt.Errorf("inspect container state %q: %w", name, err)
}
return strings.TrimSpace(output) == "true", nil
}
func (p *Provisioner) hostPort(ctx context.Context, name string) (string, error) {
output, err := p.run(ctx, nil, "container", "inspect", "--format", "{{(index (index .NetworkSettings.Ports \"5432/tcp\") 0).HostPort}}", name)
if err != nil {
return "", fmt.Errorf("inspect PostgreSQL port %q: %w", name, err)
}
port := strings.TrimSpace(output)
if port == "" {
return "", fmt.Errorf("container %q has no published PostgreSQL port", name)
}
return port, nil
}
func (p *Provisioner) run(ctx context.Context, environment []string, args ...string) (string, error) {
path := p.DockerPath
if path == "" {
path = "docker"
}
command := exec.CommandContext(ctx, path, args...)
command.Env = append(os.Environ(), environment...)
var stdout, stderr bytes.Buffer
command.Stdout = &stdout
command.Stderr = &stderr
if err := command.Run(); err != nil {
return "", fmt.Errorf("docker %s: %w: %s", strings.Join(args, " "), err, strings.TrimSpace(stderr.String()))
}
return stdout.String(), nil
}
func containerName(database string) string {
return "pawsql-" + database
}
func postgresDataMount(image string) string {
if image == "postgres:18" {
return "/var/lib/postgresql"
}
return "/var/lib/postgresql/data"
}

View File

@@ -0,0 +1,65 @@
//go:build integration
package postgres
import (
"context"
"net"
"testing"
"time"
"github.com/barkstack/pawsql/internal/config"
)
func TestProvisionerCreatesPersistentPostgres18(t *testing.T) {
const (
database = "pawsql_integration_test"
volume = "pawsql-integration-test-data"
envName = "PAWSQL_INTEGRATION_POSTGRES_PASSWORD"
)
t.Setenv(envName, "integration-test-password")
provisioner := NewProvisioner(nil)
defer func() {
_, _ = provisioner.run(context.Background(), nil, "container", "rm", "--force", containerName(database))
_, _ = provisioner.run(context.Background(), nil, "volume", "rm", "--force", volume)
}()
cfg := config.Config{Databases: []config.DatabaseConfig{{
Name: database,
Postgres: &config.PostgresConfig{
Image: "postgres:18",
Volume: volume,
PasswordEnv: envName,
},
}}}
resolved, err := provisioner.Ensure(context.Background(), cfg)
if err != nil {
t.Fatal(err)
}
address := resolved.Databases[0].Upstream
waitForAddress(t, address)
if err := provisioner.StopDatabase(context.Background(), database); err != nil {
t.Fatal(err)
}
resolved, err = provisioner.Ensure(context.Background(), cfg)
if err != nil {
t.Fatal(err)
}
waitForAddress(t, resolved.Databases[0].Upstream)
}
func waitForAddress(t *testing.T, address string) {
t.Helper()
deadline := time.Now().Add(90 * time.Second)
for {
connection, err := net.DialTimeout("tcp", address, time.Second)
if err == nil {
_ = connection.Close()
return
}
if time.Now().After(deadline) {
t.Fatalf("PostgreSQL at %s did not accept connections: %v", address, err)
}
time.Sleep(250 * time.Millisecond)
}
}

View File

@@ -0,0 +1,19 @@
package postgres
import "testing"
func TestPostgresDataMount(t *testing.T) {
tests := []struct {
image string
want string
}{
{"postgres:16", "/var/lib/postgresql/data"},
{"postgres:17", "/var/lib/postgresql/data"},
{"postgres:18", "/var/lib/postgresql"},
}
for _, test := range tests {
if got := postgresDataMount(test.image); got != test.want {
t.Errorf("postgresDataMount(%q) = %q, want %q", test.image, got, test.want)
}
}
}

View File

@@ -0,0 +1,253 @@
package postgres
import (
"context"
"log/slog"
"sync"
"time"
"github.com/barkstack/pawsql/internal/config"
"github.com/barkstack/pawsql/internal/router"
)
// DatabaseController starts and stops managed PostgreSQL databases.
type DatabaseController interface {
EnsureDatabase(context.Context, string, config.PostgresConfig) (string, error)
StopDatabase(context.Context, string) error
}
// Resolver lazily starts managed PostgreSQL containers when a route is selected.
// Concurrent connections to one database share a single start or stop operation.
type Resolver struct {
routes router.BackendResolver
controller DatabaseController
managed map[string]config.PostgresConfig
locksMu sync.Mutex
locks map[string]*sync.Mutex
statesMu sync.Mutex
states map[string]*leaseState
}
type leaseState struct {
active int
idleGeneration uint64
idleTimer *time.Timer
trafficGeneration uint64
trafficTimer *time.Timer
lastActivity time.Time
clientBytes uint64
backendBytes uint64
}
// TrafficStats is a point-in-time aggregate for one managed database.
type TrafficStats struct {
ActiveConnections int
ClientBytes uint64
BackendBytes uint64
LastActivity time.Time
}
// NewResolver wraps static route matching with lazy PostgreSQL provisioning.
func NewResolver(routes router.BackendResolver, databases []config.DatabaseConfig, controller DatabaseController) *Resolver {
managed := make(map[string]config.PostgresConfig)
for _, database := range databases {
if database.Postgres != nil {
managed[database.Name] = *database.Postgres
}
}
return &Resolver{
routes: routes, controller: controller, managed: managed,
locks: make(map[string]*sync.Mutex), states: make(map[string]*leaseState),
}
}
// Resolve selects a hostname route and starts its PostgreSQL container if needed.
func (r *Resolver) Resolve(ctx context.Context, hostname string) (router.Backend, error) {
backend, err := r.routes.Resolve(ctx, hostname)
if err != nil {
return router.Backend{}, err
}
return r.ensure(ctx, backend)
}
// ResolveDatabase selects a database route and starts its PostgreSQL container if needed.
func (r *Resolver) ResolveDatabase(ctx context.Context, database string) (router.Backend, error) {
backend, err := r.routes.ResolveDatabase(ctx, database)
if err != nil {
return router.Backend{}, err
}
return r.ensure(ctx, backend)
}
func (r *Resolver) ensure(ctx context.Context, backend router.Backend) (router.Backend, error) {
postgres, managed := r.managed[backend.DatabaseName]
if !managed {
return backend, nil
}
lock := r.lockFor(backend.DatabaseName)
lock.Lock()
defer lock.Unlock()
address, err := r.controller.EnsureDatabase(ctx, backend.DatabaseName, postgres)
if err != nil {
return router.Backend{}, err
}
r.acquire(backend.DatabaseName, postgres)
backend.Address = address
return backend, nil
}
func (r *Resolver) lockFor(database string) *sync.Mutex {
r.locksMu.Lock()
defer r.locksMu.Unlock()
lock := r.locks[database]
if lock == nil {
lock = &sync.Mutex{}
r.locks[database] = lock
}
return lock
}
// ReleaseConnection records a proxied managed-database session ending. Once the
// final session closes, an idle_timeout countdown begins.
func (r *Resolver) ReleaseConnection(database string) {
postgres, managed := r.managed[database]
if !managed {
return
}
r.statesMu.Lock()
state := r.states[database]
if state == nil || state.active == 0 {
r.statesMu.Unlock()
return
}
state.active--
if state.active == 0 {
state.trafficGeneration++
if state.trafficTimer != nil {
state.trafficTimer.Stop()
state.trafficTimer = nil
}
if postgres.IdleTimeout > 0 {
state.idleGeneration++
generation := state.idleGeneration
state.idleTimer = time.AfterFunc(postgres.IdleTimeout, func() {
r.stopAfterConnectionIdle(database, generation)
})
}
}
r.statesMu.Unlock()
}
// RecordTraffic meters proxied bytes and resets the traffic-idle countdown.
func (r *Resolver) RecordTraffic(database string, clientToBackend bool, bytes int64) {
postgres, managed := r.managed[database]
if !managed || bytes <= 0 {
return
}
r.statesMu.Lock()
state := r.stateFor(database)
if clientToBackend {
state.clientBytes += uint64(bytes)
} else {
state.backendBytes += uint64(bytes)
}
state.lastActivity = time.Now()
if state.active > 0 && postgres.TrafficIdleTimeout > 0 {
r.scheduleTrafficIdleLocked(database, state, postgres.TrafficIdleTimeout)
}
r.statesMu.Unlock()
}
// TrafficStats reports aggregate traffic collected since PawSQL started.
func (r *Resolver) TrafficStats(database string) (TrafficStats, bool) {
if _, managed := r.managed[database]; !managed {
return TrafficStats{}, false
}
r.statesMu.Lock()
defer r.statesMu.Unlock()
state := r.states[database]
if state == nil {
return TrafficStats{}, true
}
return TrafficStats{
ActiveConnections: state.active,
ClientBytes: state.clientBytes,
BackendBytes: state.backendBytes,
LastActivity: state.lastActivity,
}, true
}
func (r *Resolver) acquire(database string, postgres config.PostgresConfig) {
r.statesMu.Lock()
defer r.statesMu.Unlock()
state := r.stateFor(database)
state.active++
state.lastActivity = time.Now()
state.idleGeneration++
if state.idleTimer != nil {
state.idleTimer.Stop()
state.idleTimer = nil
}
if postgres.TrafficIdleTimeout > 0 {
r.scheduleTrafficIdleLocked(database, state, postgres.TrafficIdleTimeout)
}
}
func (r *Resolver) stateFor(database string) *leaseState {
state := r.states[database]
if state == nil {
state = &leaseState{}
r.states[database] = state
}
return state
}
func (r *Resolver) scheduleTrafficIdleLocked(database string, state *leaseState, timeout time.Duration) {
state.trafficGeneration++
generation := state.trafficGeneration
if state.trafficTimer != nil {
state.trafficTimer.Stop()
}
state.trafficTimer = time.AfterFunc(timeout, func() {
r.stopAfterTrafficIdle(database, generation)
})
}
func (r *Resolver) stopAfterConnectionIdle(database string, generation uint64) {
r.stopIfIdle(database, generation, false)
}
func (r *Resolver) stopAfterTrafficIdle(database string, generation uint64) {
r.stopIfIdle(database, generation, true)
}
func (r *Resolver) stopIfIdle(database string, generation uint64, traffic bool) {
lock := r.lockFor(database)
lock.Lock()
defer lock.Unlock()
r.statesMu.Lock()
state := r.states[database]
if state == nil {
r.statesMu.Unlock()
return
}
if traffic {
if state.active == 0 || state.trafficGeneration != generation {
r.statesMu.Unlock()
return
}
state.trafficTimer = nil
} else {
if state.active != 0 || state.idleGeneration != generation {
r.statesMu.Unlock()
return
}
state.idleTimer = nil
}
r.statesMu.Unlock()
if err := r.controller.StopDatabase(context.Background(), database); err != nil {
slog.Error("stop idle PostgreSQL container", "database", database, "traffic_idle", traffic, "error", err)
}
}

View File

@@ -0,0 +1,198 @@
package postgres
import (
"context"
"sync"
"testing"
"time"
"github.com/barkstack/pawsql/internal/config"
"github.com/barkstack/pawsql/internal/router"
)
func TestResolverDefersManagedDatabaseStartup(t *testing.T) {
routes, err := router.NewStaticResolver([]config.DatabaseConfig{
{Name: "external", Upstream: "192.0.2.1:5432"},
{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}},
})
if err != nil {
t.Fatal(err)
}
ensurer := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}}
resolver := NewResolver(routes, []config.DatabaseConfig{
{Name: "external", Upstream: "192.0.2.1:5432"},
{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}},
}, ensurer)
backend, err := resolver.ResolveDatabase(context.Background(), "external")
if err != nil {
t.Fatal(err)
}
if backend.Address != "192.0.2.1:5432" || ensurer.calls != 0 {
t.Errorf("external backend = %#v, calls = %d", backend, ensurer.calls)
}
backend, err = resolver.ResolveDatabase(context.Background(), "managed")
if err != nil {
t.Fatal(err)
}
if backend.Address != "127.0.0.1:55432" || ensurer.calls != 1 {
t.Errorf("managed backend = %#v, calls = %d", backend, ensurer.calls)
}
}
func TestResolverSerializesManagedStartup(t *testing.T) {
routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}}})
if err != nil {
t.Fatal(err)
}
ensurer := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, gate: make(chan struct{}), started: make(chan struct{})}
resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD"}}}, ensurer)
var wg sync.WaitGroup
for range 2 {
wg.Add(1)
go func() {
defer wg.Done()
_, _ = resolver.ResolveDatabase(context.Background(), "managed")
}()
}
<-ensurer.started
close(ensurer.gate)
wg.Wait()
if ensurer.calls != 2 {
t.Errorf("EnsureDatabase calls = %d, want 2 route resolutions", ensurer.calls)
}
if ensurer.maxActive != 1 {
t.Errorf("concurrent EnsureDatabase calls = %d, want 1", ensurer.maxActive)
}
}
func TestResolverStopsDatabaseAfterIdleTimeout(t *testing.T) {
routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 20 * time.Millisecond}}})
if err != nil {
t.Fatal(err)
}
controller := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, stopped: make(chan string, 1)}
resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 20 * time.Millisecond}}}, controller)
if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil {
t.Fatal(err)
}
resolver.ReleaseConnection("managed")
select {
case database := <-controller.stopped:
if database != "managed" {
t.Errorf("stopped database = %q", database)
}
case <-time.After(time.Second):
t.Fatal("managed database was not stopped after idle timeout")
}
}
func TestResolverCancelsIdleStopForNewConnection(t *testing.T) {
routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 40 * time.Millisecond}}})
if err != nil {
t.Fatal(err)
}
controller := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, stopped: make(chan string, 1)}
resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: &config.PostgresConfig{Image: "postgres:18", Volume: "managed-data", PasswordEnv: "MANAGED_PASSWORD", IdleTimeout: 40 * time.Millisecond}}}, controller)
if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil {
t.Fatal(err)
}
resolver.ReleaseConnection("managed")
if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil {
t.Fatal(err)
}
select {
case database := <-controller.stopped:
t.Fatalf("stopped active database %q", database)
case <-time.After(80 * time.Millisecond):
}
resolver.ReleaseConnection("managed")
select {
case <-controller.stopped:
case <-time.After(time.Second):
t.Fatal("managed database was not stopped after final lease release")
}
}
func TestResolverMetersTrafficAndStopsTrafficIdleSession(t *testing.T) {
postgres := &config.PostgresConfig{
Image: "postgres:18",
Volume: "managed-data",
PasswordEnv: "MANAGED_PASSWORD",
TrafficIdleTimeout: 100 * time.Millisecond,
}
routes, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "managed", Postgres: postgres}})
if err != nil {
t.Fatal(err)
}
controller := &fakeEnsurer{addresses: map[string]string{"managed": "127.0.0.1:55432"}, stopped: make(chan string, 1)}
resolver := NewResolver(routes, []config.DatabaseConfig{{Name: "managed", Postgres: postgres}}, controller)
if _, err := resolver.ResolveDatabase(context.Background(), "managed"); err != nil {
t.Fatal(err)
}
time.Sleep(70 * time.Millisecond)
resolver.RecordTraffic("managed", true, 7)
resolver.RecordTraffic("managed", false, 11)
stats, ok := resolver.TrafficStats("managed")
if !ok || stats.ActiveConnections != 1 || stats.ClientBytes != 7 || stats.BackendBytes != 11 || stats.LastActivity.IsZero() {
t.Errorf("TrafficStats() = %#v, %t", stats, ok)
}
select {
case database := <-controller.stopped:
t.Fatalf("stopped database %q despite recent traffic", database)
case <-time.After(50 * time.Millisecond):
}
select {
case database := <-controller.stopped:
if database != "managed" {
t.Errorf("stopped database = %q", database)
}
case <-time.After(time.Second):
t.Fatal("managed database was not stopped after traffic idle timeout")
}
}
type fakeEnsurer struct {
mu sync.Mutex
addresses map[string]string
calls int
active int
maxActive int
gate chan struct{}
started chan struct{}
startedOnce sync.Once
stopped chan string
}
func (f *fakeEnsurer) EnsureDatabase(_ context.Context, database string, _ config.PostgresConfig) (string, error) {
f.mu.Lock()
f.calls++
f.active++
if f.active > f.maxActive {
f.maxActive = f.active
}
if f.started != nil {
f.startedOnce.Do(func() { close(f.started) })
}
f.mu.Unlock()
if f.gate != nil {
<-f.gate
}
f.mu.Lock()
f.active--
address := f.addresses[database]
f.mu.Unlock()
return address, nil
}
func (f *fakeEnsurer) StopDatabase(_ context.Context, database string) error {
if f.stopped != nil {
f.stopped <- database
}
return nil
}

82
internal/proxy/proxy.go Normal file
View File

@@ -0,0 +1,82 @@
// Package proxy transports an opaque bidirectional byte stream.
package proxy
import (
"errors"
"io"
"net"
)
type copyResult struct {
err error
dst net.Conn
}
type closeWriter interface {
CloseWrite() error
}
// TrafficObserver receives each positive byte count as it is copied.
// clientToBackend identifies the PostgreSQL client-to-server direction.
type TrafficObserver func(clientToBackend bool, bytes int64)
// Bidirectional copies bytes in both directions. TCP peers retain half-close
// semantics; a TLS client is closed when its upstream direction has ended,
// because crypto/tls exposes no CloseWrite operation.
func Bidirectional(client, backend net.Conn) error {
return BidirectionalWithTraffic(client, backend, nil)
}
// BidirectionalWithTraffic copies bytes in both directions and reports every
// positive read to observer when one is provided.
func BidirectionalWithTraffic(client, backend net.Conn, observer TrafficObserver) error {
results := make(chan copyResult, 2)
copyStream := func(dst, src net.Conn, clientToBackend bool) {
_, err := io.Copy(dst, trafficReader{Reader: src, clientToBackend: clientToBackend, observer: observer})
results <- copyResult{err: err, dst: dst}
}
go copyStream(backend, client, true)
go copyStream(client, backend, false)
first := <-results
if writer, ok := first.dst.(closeWriter); ok {
_ = writer.CloseWrite()
} else {
// TLS cannot be half-closed. Closing unblocks the opposite copy and
// prevents an idle peer from retaining a handler goroutine.
_ = client.Close()
}
second := <-results
if writer, ok := second.dst.(closeWriter); ok {
_ = writer.CloseWrite()
}
return combine(first.err, second.err)
}
type trafficReader struct {
io.Reader
clientToBackend bool
observer TrafficObserver
}
func (r trafficReader) Read(buffer []byte) (int, error) {
bytes, err := r.Reader.Read(buffer)
if bytes > 0 && r.observer != nil {
r.observer(r.clientToBackend, int64(bytes))
}
return bytes, err
}
func combine(errs ...error) error {
var relevant []error
for _, err := range errs {
if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) {
relevant = append(relevant, err)
}
}
return errors.Join(relevant...)
}
// Compile-time assertion: a TCP connection has practical half-close support.
var _ closeWriter = (*net.TCPConn)(nil)

89
internal/router/routes.go Normal file
View File

@@ -0,0 +1,89 @@
// Package router resolves incoming SNI names to PostgreSQL backends.
package router
import (
"context"
"errors"
"fmt"
"strings"
"github.com/barkstack/pawsql/internal/config"
)
var (
ErrUnknownHostname = errors.New("unknown database hostname")
ErrUnknownDatabase = errors.New("unknown configured database")
)
// Backend is a resolved PostgreSQL destination.
type Backend struct {
DatabaseName string
Address string
}
// BackendResolver permits future lifecycle-aware backend discovery.
type BackendResolver interface {
Resolve(context.Context, string) (Backend, error)
ResolveDatabase(context.Context, string) (Backend, error)
}
// ConnectionLeaseManager releases a backend lease when a proxied client session ends.
// Resolvers that do not manage lifecycle state may ignore this optional interface.
type ConnectionLeaseManager interface {
ReleaseConnection(string)
}
// TrafficMeter receives byte counts flowing through a proxied backend session.
// clientToBackend identifies the PostgreSQL client-to-server direction.
type TrafficMeter interface {
RecordTraffic(database string, clientToBackend bool, bytes int64)
}
// StaticResolver resolves routes loaded from a Barkfile.
type StaticResolver struct {
hostnameRoutes map[string]Backend
databaseRoutes map[string]Backend
}
// NewStaticResolver builds a resolver from typed configuration.
func NewStaticResolver(databases []config.DatabaseConfig) (*StaticResolver, error) {
hostnameRoutes := make(map[string]Backend, len(databases))
databaseRoutes := make(map[string]Backend, len(databases))
for _, database := range databases {
if database.Name == "" {
return nil, errors.New("database name is required")
}
if _, exists := databaseRoutes[database.Name]; exists {
return nil, fmt.Errorf("duplicate database name %q", database.Name)
}
backend := Backend{DatabaseName: database.Name, Address: database.Upstream}
databaseRoutes[database.Name] = backend
hostname := config.NormalizeHostname(database.Hostname)
if hostname == "" {
continue
}
if _, exists := hostnameRoutes[hostname]; exists {
return nil, fmt.Errorf("duplicate hostname %q", hostname)
}
hostnameRoutes[hostname] = backend
}
return &StaticResolver{hostnameRoutes: hostnameRoutes, databaseRoutes: databaseRoutes}, nil
}
// Resolve returns the exact case-insensitive hostname match, never a default route.
func (r *StaticResolver) Resolve(_ context.Context, hostname string) (Backend, error) {
backend, ok := r.hostnameRoutes[strings.ToLower(strings.TrimSuffix(strings.TrimSpace(hostname), "."))]
if !ok {
return Backend{}, fmt.Errorf("%w: %s", ErrUnknownHostname, hostname)
}
return backend, nil
}
// ResolveDatabase returns the exact configured PostgreSQL database match.
func (r *StaticResolver) ResolveDatabase(_ context.Context, database string) (Backend, error) {
backend, ok := r.databaseRoutes[database]
if !ok {
return Backend{}, fmt.Errorf("%w: %s", ErrUnknownDatabase, database)
}
return backend, nil
}

View File

@@ -0,0 +1,64 @@
package router
import (
"context"
"errors"
"testing"
"github.com/barkstack/pawsql/internal/config"
)
func TestStaticResolverRoutesCaseInsensitiveHostnames(t *testing.T) {
resolver, err := NewStaticResolver([]config.DatabaseConfig{
{Name: "foo", Hostname: "foo.pawsql.barkstack.dev", Upstream: "postgres-foo:5432"},
{Name: "bar", Hostname: "bar.pawsql.barkstack.dev", Upstream: "postgres-bar:5432"},
})
if err != nil {
t.Fatalf("NewStaticResolver() error = %v", err)
}
for _, test := range []struct{ hostname, name, address string }{
{"FOO.pawsql.barkstack.dev", "foo", "postgres-foo:5432"},
{"bar.pawsql.barkstack.dev", "bar", "postgres-bar:5432"},
} {
backend, err := resolver.Resolve(context.Background(), test.hostname)
if err != nil {
t.Errorf("Resolve(%q) error = %v", test.hostname, err)
continue
}
if backend.DatabaseName != test.name || backend.Address != test.address {
t.Errorf("Resolve(%q) = %#v", test.hostname, backend)
}
}
}
func TestStaticResolverRejectsUnknownHostname(t *testing.T) {
resolver, err := NewStaticResolver([]config.DatabaseConfig{{Name: "foo", Hostname: "foo.pawsql.barkstack.dev", Upstream: "postgres-foo:5432"}})
if err != nil {
t.Fatal(err)
}
_, err = resolver.Resolve(context.Background(), "unknown.pawsql.barkstack.dev")
if !errors.Is(err, ErrUnknownHostname) {
t.Errorf("Resolve() error = %v, want ErrUnknownHostname", err)
}
}
func TestStaticResolverRoutesConfiguredDatabase(t *testing.T) {
resolver, err := NewStaticResolver([]config.DatabaseConfig{
{Name: "analytics", Upstream: "postgres-shared:5432"},
{Name: "app", Hostname: "app.pawsql.barkstack.dev", Upstream: "postgres-shared:5432"},
})
if err != nil {
t.Fatal(err)
}
backend, err := resolver.ResolveDatabase(context.Background(), "analytics")
if err != nil {
t.Fatal(err)
}
if backend.DatabaseName != "analytics" || backend.Address != "postgres-shared:5432" {
t.Errorf("ResolveDatabase() = %#v", backend)
}
_, err = resolver.ResolveDatabase(context.Background(), "unknown")
if !errors.Is(err, ErrUnknownDatabase) {
t.Errorf("ResolveDatabase() error = %v, want ErrUnknownDatabase", err)
}
}

199
internal/server/server.go Normal file
View File

@@ -0,0 +1,199 @@
// Package server accepts PostgreSQL TLS sessions and routes them by SNI.
package server
import (
"context"
"crypto/tls"
"errors"
"fmt"
"log/slog"
"net"
"strings"
"sync"
"time"
"github.com/barkstack/pawsql/internal/pgwire"
"github.com/barkstack/pawsql/internal/proxy"
"github.com/barkstack/pawsql/internal/router"
)
const (
defaultHandshakeTimeout = 15 * time.Second
defaultResolveTimeout = 60 * time.Second
defaultDialTimeout = 10 * time.Second
)
// Server terminates client TLS and transparently proxies PostgreSQL bytes.
type Server struct {
TLSConfig *tls.Config
Resolver router.BackendResolver
Logger *slog.Logger
HandshakeTimeout time.Duration
ResolveTimeout time.Duration
DialTimeout time.Duration
mu sync.Mutex
listener net.Listener
handlers sync.WaitGroup
}
// New validates server dependencies and applies safe protocol defaults.
func New(tlsConfig *tls.Config, resolver router.BackendResolver, logger *slog.Logger) (*Server, error) {
if tlsConfig == nil {
return nil, errors.New("TLS configuration is required")
}
if resolver == nil {
return nil, errors.New("backend resolver is required")
}
if logger == nil {
logger = slog.Default()
}
copy := tlsConfig.Clone()
if copy.MinVersion == 0 {
copy.MinVersion = tls.VersionTLS12
}
return &Server{TLSConfig: copy, Resolver: resolver, Logger: logger, HandshakeTimeout: defaultHandshakeTimeout, ResolveTimeout: defaultResolveTimeout, DialTimeout: defaultDialTimeout}, nil
}
// Serve accepts connections until Shutdown closes its listener.
func (s *Server) Serve(listener net.Listener) error {
s.mu.Lock()
if s.listener != nil {
s.mu.Unlock()
return errors.New("server is already serving")
}
s.listener = listener
s.mu.Unlock()
for {
connection, err := listener.Accept()
if err != nil {
if errors.Is(err, net.ErrClosed) {
return nil
}
if temporary, ok := err.(interface{ Temporary() bool }); ok && temporary.Temporary() {
s.Logger.Warn("temporary accept failure", "error", err)
continue
}
return fmt.Errorf("accept connection: %w", err)
}
s.handlers.Add(1)
go func() {
defer s.handlers.Done()
s.handle(connection)
}()
}
}
// Shutdown stops accepting new connections. Existing sessions continue until they close.
func (s *Server) Shutdown() error {
s.mu.Lock()
listener := s.listener
s.mu.Unlock()
if listener == nil {
return nil
}
return listener.Close()
}
// Wait waits for existing routed connections to finish.
func (s *Server) Wait() { s.handlers.Wait() }
func (s *Server) handle(connection net.Conn) {
started := time.Now()
remote := connection.RemoteAddr().String()
var sni string
var backend router.Backend
result := "rejected"
defer func() {
_ = connection.Close()
s.Logger.Info("connection finished", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address, "duration", time.Since(started), "result", result)
}()
handshakeTimeout := s.HandshakeTimeout
if handshakeTimeout <= 0 {
handshakeTimeout = defaultHandshakeTimeout
}
_ = connection.SetDeadline(time.Now().Add(handshakeTimeout))
if err := pgwire.ReadSSLRequest(connection); err != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_ssl_request", "error", err)
return
}
if _, err := connection.Write([]byte{'S'}); err != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "ssl_response_failed", "error", err)
return
}
tlsConnection := tls.Server(connection, s.TLSConfig)
if err := tlsConnection.Handshake(); err != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "tls_handshake_failed", "error", err)
return
}
_ = tlsConnection.SetDeadline(time.Time{})
sni = normalizeHostname(tlsConnection.ConnectionState().ServerName)
var (
proxyClient net.Conn = tlsConnection
err error
)
resolveTimeout := s.ResolveTimeout
if resolveTimeout <= 0 {
resolveTimeout = defaultResolveTimeout
}
resolveContext, cancelResolve := context.WithTimeout(context.Background(), resolveTimeout)
defer cancelResolve()
if sni != "" {
backend, err = s.Resolver.Resolve(resolveContext, sni)
if err != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "sni", sni, "result", "unknown_hostname", "error", err)
return
}
} else {
startup, startupErr := pgwire.ReadStartupMessage(tlsConnection)
if startupErr != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_startup_message", "error", startupErr)
return
}
backend, err = s.Resolver.ResolveDatabase(resolveContext, startup.Database)
if err != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "result", "unknown_database", "error", err)
return
}
proxyClient = pgwire.Replay(tlsConnection, startup.Bytes())
}
if leases, ok := s.Resolver.(router.ConnectionLeaseManager); ok {
defer leases.ReleaseConnection(backend.DatabaseName)
}
dialTimeout := s.DialTimeout
if dialTimeout <= 0 {
dialTimeout = defaultDialTimeout
}
ctx, cancel := context.WithTimeout(context.Background(), dialTimeout)
upstream, err := (&net.Dialer{}).DialContext(ctx, "tcp", backend.Address)
cancel()
if err != nil {
s.Logger.Warn("connection rejected", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address, "result", "upstream_connect_failed", "error", err)
return
}
defer upstream.Close()
s.Logger.Info("connection routed", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address)
var proxyErr error
if meter, ok := s.Resolver.(router.TrafficMeter); ok {
proxyErr = proxy.BidirectionalWithTraffic(proxyClient, upstream, func(clientToBackend bool, bytes int64) {
meter.RecordTraffic(backend.DatabaseName, clientToBackend, bytes)
})
} else {
proxyErr = proxy.Bidirectional(proxyClient, upstream)
}
if proxyErr != nil {
result = "proxy_error"
s.Logger.Warn("connection proxy failed", "remote_address", remote, "sni", sni, "database", backend.DatabaseName, "upstream", backend.Address, "error", proxyErr)
return
}
result = "closed"
}
func normalizeHostname(hostname string) string {
return strings.ToLower(strings.TrimSuffix(hostname, "."))
}

View File

@@ -0,0 +1,261 @@
package server
import (
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/binary"
"encoding/pem"
"io"
"math/big"
"net"
"testing"
"time"
"github.com/barkstack/pawsql/internal/config"
"github.com/barkstack/pawsql/internal/pgwire"
"github.com/barkstack/pawsql/internal/router"
)
func TestServerHandlesPostgreSQLSSLRequestBeforeTLS(t *testing.T) {
backendListener := listen(t)
defer backendListener.Close()
backendReceived := make(chan []byte, 1)
go func() {
connection, err := backendListener.Accept()
if err != nil {
return
}
defer connection.Close()
payload := make([]byte, 4)
if _, err := io.ReadFull(connection, payload); err == nil {
backendReceived <- payload
_, _ = connection.Write([]byte("backend"))
}
}()
staticResolver, err := router.NewStaticResolver([]config.DatabaseConfig{{Name: "foo", Hostname: "foo.pawsql.test", Upstream: backendListener.Addr().String()}})
if err != nil {
t.Fatal(err)
}
resolver := &releaseTrackingResolver{BackendResolver: staticResolver, released: make(chan string, 1), traffic: make(chan trafficEvent, 2)}
routingServer, err := New(testCertificate(t), resolver, nil)
if err != nil {
t.Fatal(err)
}
listener := listen(t)
defer func() {
_ = routingServer.Shutdown()
routingServer.Wait()
_ = listener.Close()
}()
go func() { _ = routingServer.Serve(listener) }()
connection, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer connection.Close()
request := pgwire.SSLRequest()
if _, err := connection.Write(request[:]); err != nil {
t.Fatal(err)
}
response := make([]byte, 1)
if _, err := io.ReadFull(connection, response); err != nil {
t.Fatal(err)
}
if response[0] != 'S' {
t.Fatalf("SSL response = %q, want S", response)
}
client := tls.Client(connection, &tls.Config{ServerName: "FOO.pawsql.test", InsecureSkipVerify: true}) // test certificate is ephemeral
if err := client.Handshake(); err != nil {
t.Fatalf("TLS handshake: %v", err)
}
if _, err := client.Write([]byte("ping")); err != nil {
t.Fatal(err)
}
response = make([]byte, len("backend"))
if _, err := io.ReadFull(client, response); err != nil {
t.Fatal(err)
}
if string(response) != "backend" {
t.Fatalf("proxied response = %q", response)
}
if err := client.Close(); err != nil {
t.Fatal(err)
}
select {
case payload := <-backendReceived:
if string(payload) != "ping" {
t.Errorf("backend payload = %q", payload)
}
case <-time.After(time.Second):
t.Fatal("backend did not receive decrypted PostgreSQL bytes")
}
select {
case database := <-resolver.released:
if database != "foo" {
t.Errorf("released database = %q, want foo", database)
}
case <-time.After(time.Second):
t.Fatal("server did not release backend lifecycle lease")
}
var clientBytes, backendBytes int64
for range 2 {
select {
case event := <-resolver.traffic:
if event.clientToBackend {
clientBytes += event.bytes
} else {
backendBytes += event.bytes
}
case <-time.After(time.Second):
t.Fatal("server did not report proxied traffic")
}
}
if clientBytes != int64(len("ping")) || backendBytes != int64(len("backend")) {
t.Errorf("proxied traffic = client %d, backend %d", clientBytes, backendBytes)
}
}
func TestServerRoutesNoSNIByConfiguredDatabase(t *testing.T) {
backendListener := listen(t)
defer backendListener.Close()
backendDatabase := make(chan string, 1)
go func() {
connection, err := backendListener.Accept()
if err != nil {
return
}
defer connection.Close()
startup, err := pgwire.ReadStartupMessage(connection)
if err != nil {
return
}
backendDatabase <- startup.Database
_, _ = connection.Write([]byte("backend"))
}()
resolver, err := router.NewStaticResolver([]config.DatabaseConfig{{
Name: "analytics",
Upstream: backendListener.Addr().String(),
}})
if err != nil {
t.Fatal(err)
}
routingServer, err := New(testCertificate(t), resolver, nil)
if err != nil {
t.Fatal(err)
}
listener := listen(t)
defer func() {
_ = routingServer.Shutdown()
routingServer.Wait()
_ = listener.Close()
}()
go func() { _ = routingServer.Serve(listener) }()
connection, err := net.Dial("tcp", listener.Addr().String())
if err != nil {
t.Fatal(err)
}
defer connection.Close()
request := pgwire.SSLRequest()
if _, err := connection.Write(request[:]); err != nil {
t.Fatal(err)
}
response := make([]byte, 1)
if _, err := io.ReadFull(connection, response); err != nil || response[0] != 'S' {
t.Fatalf("SSL response = %q, error = %v", response, err)
}
client := tls.Client(connection, &tls.Config{InsecureSkipVerify: true}) // no ServerName intentionally omits SNI
if err := client.Handshake(); err != nil {
t.Fatal(err)
}
if _, err := client.Write(startupMessage("analytics")); err != nil {
t.Fatal(err)
}
response = make([]byte, len("backend"))
if _, err := io.ReadFull(client, response); err != nil {
t.Fatal(err)
}
if string(response) != "backend" {
t.Fatalf("proxied response = %q", response)
}
select {
case database := <-backendDatabase:
if database != "analytics" {
t.Errorf("backend database = %q, want analytics", database)
}
case <-time.After(time.Second):
t.Fatal("backend did not receive original StartupMessage")
}
}
func startupMessage(database string) []byte {
body := make([]byte, 4)
binary.BigEndian.PutUint32(body, 196608)
body = append(body, "user"...)
body = append(body, 0)
body = append(body, "ruckstack"...)
body = append(body, 0)
body = append(body, "database"...)
body = append(body, 0)
body = append(body, database...)
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 trafficEvent struct {
clientToBackend bool
bytes int64
}
type releaseTrackingResolver struct {
router.BackendResolver
released chan string
traffic chan trafficEvent
}
func (r *releaseTrackingResolver) RecordTraffic(_ string, clientToBackend bool, bytes int64) {
r.traffic <- trafficEvent{clientToBackend: clientToBackend, bytes: bytes}
}
func (r *releaseTrackingResolver) ReleaseConnection(database string) {
r.released <- database
}
func listen(t *testing.T) net.Listener {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
return listener
}
func testCertificate(t *testing.T) *tls.Config {
t.Helper()
key, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
t.Fatal(err)
}
certificateTemplate := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "pawsql test"}, NotBefore: time.Now().Add(-time.Minute), NotAfter: time.Now().Add(time.Hour), DNSNames: []string{"*.pawsql.test"}, KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}}
certificateDER, err := x509.CreateCertificate(rand.Reader, certificateTemplate, certificateTemplate, &key.PublicKey, key)
if err != nil {
t.Fatal(err)
}
certificatePEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificateDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
certificate, err := tls.X509KeyPair(certificatePEM, keyPEM)
if err != nil {
t.Fatal(err)
}
return &tls.Config{Certificates: []tls.Certificate{certificate}}
}