feat: add PawSQL docs examples and image CI
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 2m28s
All checks were successful
Build and Push Image / docker-build-and-push (push) Successful in 2m28s
This commit is contained in:
38
internal/config/config.go
Normal file
38
internal/config/config.go
Normal 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
390
internal/config/parser.go
Normal 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
|
||||
}
|
||||
254
internal/config/parser_test.go
Normal file
254
internal/config/parser_test.go
Normal 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
115
internal/config/validate.go
Normal 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
|
||||
}
|
||||
40
internal/pgwire/sslrequest.go
Normal file
40
internal/pgwire/sslrequest.go
Normal 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
108
internal/pgwire/startup.go
Normal 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)
|
||||
}
|
||||
76
internal/pgwire/startup_test.go
Normal file
76
internal/pgwire/startup_test.go
Normal 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 }
|
||||
217
internal/postgres/provisioner.go
Normal file
217
internal/postgres/provisioner.go
Normal 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"
|
||||
}
|
||||
65
internal/postgres/provisioner_integration_test.go
Normal file
65
internal/postgres/provisioner_integration_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
19
internal/postgres/provisioner_test.go
Normal file
19
internal/postgres/provisioner_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
253
internal/postgres/resolver.go
Normal file
253
internal/postgres/resolver.go
Normal 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)
|
||||
}
|
||||
}
|
||||
198
internal/postgres/resolver_test.go
Normal file
198
internal/postgres/resolver_test.go
Normal 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
82
internal/proxy/proxy.go
Normal 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
89
internal/router/routes.go
Normal 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
|
||||
}
|
||||
64
internal/router/routes_test.go
Normal file
64
internal/router/routes_test.go
Normal 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
199
internal/server/server.go
Normal 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, "."))
|
||||
}
|
||||
261
internal/server/server_test.go
Normal file
261
internal/server/server_test.go
Normal 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}}
|
||||
}
|
||||
Reference in New Issue
Block a user