commit 865f7c26c9e48ea291217c4a9a911dddfaacdf91 Author: Shaun Campbell Date: Tue Sep 15 18:49:26 2026 -0400 feat: add PawSQL docs examples and image CI diff --git a/.gitea/workflows/main-image.yml b/.gitea/workflows/main-image.yml new file mode 100644 index 0000000..2bc2e7c --- /dev/null +++ b/.gitea/workflows/main-image.yml @@ -0,0 +1,56 @@ +name: Build and Push Image + +on: + push: + branches: + - main + workflow_dispatch: + +env: + REGISTRY_HOST: registry.campbellwireless.net + IMAGE_NAME: ${{ github.repository }} + +jobs: + docker-build-and-push: + runs-on: ubuntu-latest + + steps: + - name: Checkout + env: + REPO_URL: ${{ github.server_url }}/${{ github.repository }}.git + run: | + set -eux + git init . + git remote add origin "$REPO_URL" + auth="$(printf '%s' '${{ github.actor }}:${{ secrets.GITHUB_TOKEN }}' | base64 | tr -d '\n')" + git config --local "http.${{ github.server_url }}/.extraheader" "AUTHORIZATION: basic $auth" + git fetch --prune --no-recurse-submodules origin +refs/heads/*:refs/remotes/origin/* +refs/tags/*:refs/tags/* + git checkout --detach "${{ github.sha }}" + + - name: Compute image tags + id: image + run: | + echo "sha_short=$(echo '${{ github.sha }}' | cut -c1-12)" >> "$GITHUB_OUTPUT" + + - name: Setup Docker Buildx + uses: docker/setup-buildx-action@v3 + + - name: Log in to Gitea registry + uses: docker/login-action@v3 + with: + registry: ${{ env.REGISTRY_HOST }} + username: ${{ secrets.REGISTRY_USERNAME }} + password: ${{ secrets.REGISTRY_PASSWORD }} + + - name: Build and push image + uses: docker/build-push-action@v6 + with: + context: . + file: ./Dockerfile + push: true + tags: | + ${{ env.REGISTRY_HOST }}/${{ env.IMAGE_NAME }}:latest + ${{ env.REGISTRY_HOST }}/${{ env.IMAGE_NAME }}:${{ steps.image.outputs.sha_short }} + labels: | + org.opencontainers.image.source=${{ github.server_url }}/${{ github.repository }} + org.opencontainers.image.revision=${{ github.sha }} diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..2cc413a --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +/Barkfile +/pawsql diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..40fc0d8 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,11 @@ +FROM golang:1.24-alpine AS build +WORKDIR /src +COPY go.mod ./ +COPY cmd ./cmd +COPY internal ./internal +RUN CGO_ENABLED=0 go build -trimpath -ldflags='-s -w' -o /pawsql ./cmd/pawsql + +FROM gcr.io/distroless/static-debian12:nonroot +COPY --from=build /pawsql /usr/local/bin/pawsql +ENTRYPOINT ["/usr/local/bin/pawsql"] +CMD ["--config", "/etc/pawsql/Barkfile"] diff --git a/README.md b/README.md new file mode 100644 index 0000000..69f3588 --- /dev/null +++ b/README.md @@ -0,0 +1,101 @@ +# PawSQL + +PawSQL is a TLS-terminating PostgreSQL router. It accepts PostgreSQL clients on one address, chooses a configured route from the TLS Server Name Indication (SNI) or database name, and proxies the PostgreSQL stream to an external or PawSQL-managed PostgreSQL server. + +## Prerequisites + +- Go 1.24 or later to build and run PawSQL natively. +- Docker Engine and a usable `docker` CLI to build the PawSQL image. PawSQL also needs them in its own execution environment when it manages PostgreSQL containers. +- A TLS certificate and private key readable by PawSQL. The certificate must cover every hostname clients use for SNI routing. +- Docker Engine access for each `postgres` route. Managed database images are limited to `postgres:16`, `postgres:17`, and `postgres:18`. + +## Build, configure, and run + +Build a native binary: + +```sh +go build -o pawsql ./cmd/pawsql +``` + +Create a `Barkfile` and validate it before starting: + +```sh +./pawsql validate --config Barkfile +./pawsql --config Barkfile +``` + +`--config` defaults to `Barkfile`. The listener, TLS material, and at least one database route are required. + +To build and run the PawSQL container image for routes reachable from that container: + +```sh +docker build -t pawsql . +docker run --rm --publish 5432:5432 \ + --volume "$PWD/Barkfile:/etc/pawsql/Barkfile:ro" \ + --volume "$PWD/tls:/etc/pawsql/tls:ro" \ + pawsql +``` + +The provided image contains only PawSQL and is suitable for external `upstream` routes. Managed PostgreSQL routes require native PawSQL or a custom image that supplies a Docker CLI and access to the Docker Engine, typically through the Docker socket. + +## Barkfile + +A Barkfile has one `pawsql` block. Each `database` has a unique name and exactly one route type: an `upstream` external PostgreSQL address or a `postgres` managed database. + +```text +pawsql { + listen :5432 + + tls { + cert /etc/pawsql/tls/fullchain.pem + key /etc/pawsql/tls/privkey.pem + } + + database reporting { + hostname reports.db.example.com + upstream reporting.internal:5432 + } + + database application { + hostname app.db.example.com + postgres { + image postgres:17 + volume pawsql-application-data + password_env APPLICATION_POSTGRES_PASSWORD + idle_timeout 10m + traffic_idle_timeout 1h + } + } +} +``` + +`listen` is PawSQL's TCP address. `cert` and `key` identify the client-facing TLS certificate and key. `hostname` is optional; it is used only for SNI routing. `upstream` is the address of an existing PostgreSQL server. + +For a managed `postgres` route, `image`, `volume`, and `password_env` are required. On first use, PawSQL reads the named environment variable to create the database container and configures the database and PostgreSQL user with the route's database name. The named Docker volume preserves its data. Set the password environment variable in PawSQL's environment, not in the Barkfile. + +See [`examples/Barkfile`](examples/Barkfile) and its accompanying [`examples/docker-compose.yml`](examples/docker-compose.yml) for a two-route external PostgreSQL example with locally generated development certificates: + +```sh +cd examples +docker compose up --build +``` + +## Routing and TLS + +PawSQL requires PostgreSQL's SSL negotiation and terminates client TLS before proxying PostgreSQL bytes to the selected upstream. + +- **With SNI:** PawSQL uses the TLS server name to select an exact configured `hostname` match. Hostname matching is case-insensitive and ignores a trailing dot. An unknown SNI name is rejected; PawSQL does not fall back to a database-name route when SNI is present. +- **Without SNI:** After TLS is established, PawSQL reads the PostgreSQL startup message and selects the route whose `database` name exactly matches the requested PostgreSQL database. This makes a route without `hostname` usable by non-SNI clients. + +Use a certificate trusted by clients and containing the SNI hostname they present. Clients that do not send SNI must request the configured database route name. + +## Managed PostgreSQL lifecycle + +Managed PostgreSQL is lazy: PawSQL creates or starts its `pawsql-` container only when a client selects that route, waits for PostgreSQL to accept connections, then proxies the session. PawSQL stops managed containers but does not remove their data volumes. + +Two optional Go-duration controls govern stopping a managed container; `0` disables either control: + +- `idle_timeout` starts only after the last proxied client session closes. When that countdown expires, PawSQL stops the managed container. +- `traffic_idle_timeout` starts for an open session and resets whenever PawSQL proxies bytes in either direction. When it expires, PawSQL stops the managed container even though sessions remain open. + +`traffic_idle_timeout` is intentionally aggressive: it terminates open but silent sessions. Do not enable it for workloads that keep idle connections, transactions, listeners, or connection-pool sessions alive unless that interruption is acceptable. diff --git a/cmd/pawsql/main.go b/cmd/pawsql/main.go new file mode 100644 index 0000000..a8b3575 --- /dev/null +++ b/cmd/pawsql/main.go @@ -0,0 +1,103 @@ +package main + +import ( + "crypto/tls" + "errors" + "flag" + "fmt" + "log/slog" + "net" + "os" + "os/signal" + "syscall" + + "github.com/barkstack/pawsql/internal/config" + "github.com/barkstack/pawsql/internal/postgres" + "github.com/barkstack/pawsql/internal/router" + "github.com/barkstack/pawsql/internal/server" +) + +func main() { + logger := slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: slog.LevelInfo})) + if err := run(os.Args[1:], logger); err != nil { + logger.Error("pawsql failed", "error", err) + os.Exit(1) + } +} + +func run(args []string, logger *slog.Logger) error { + validateOnly, configPath, err := parseArguments(args) + if err != nil { + return err + } + cfg, err := config.ParseFile(configPath) + if err != nil { + return err + } + if err := cfg.Validate(); err != nil { + return fmt.Errorf("invalid configuration: %w", err) + } + certificate, err := tls.LoadX509KeyPair(cfg.TLS.CertFile, cfg.TLS.KeyFile) + if err != nil { + return fmt.Errorf("load TLS certificate and key: %w", err) + } + if validateOnly { + logger.Info("configuration is valid", "config", configPath) + return nil + } + + staticResolver, err := router.NewStaticResolver(cfg.Databases) + if err != nil { + return fmt.Errorf("build route resolver: %w", err) + } + resolver := postgres.NewResolver(staticResolver, cfg.Databases, postgres.NewProvisioner(logger)) + routingServer, err := server.New(&tls.Config{Certificates: []tls.Certificate{certificate}}, resolver, logger) + if err != nil { + return err + } + listener, err := net.Listen("tcp", cfg.Listen) + if err != nil { + return fmt.Errorf("listen on %s: %w", cfg.Listen, err) + } + logger.Info("pawsql listening", "address", listener.Addr().String()) + + signals := make(chan os.Signal, 1) + signal.Notify(signals, os.Interrupt, syscall.SIGTERM) + defer signal.Stop(signals) + serveErrors := make(chan error, 1) + go func() { serveErrors <- routingServer.Serve(listener) }() + + select { + case received := <-signals: + logger.Info("shutdown signal received", "signal", received.String()) + if err := routingServer.Shutdown(); err != nil && !errors.Is(err, net.ErrClosed) { + return fmt.Errorf("stop listener: %w", err) + } + if err := <-serveErrors; err != nil { + return err + } + routingServer.Wait() + logger.Info("pawsql shutdown complete") + return nil + case err := <-serveErrors: + return err + } +} + +func parseArguments(args []string) (validateOnly bool, configPath string, err error) { + if len(args) > 0 && args[0] == "validate" { + validateOnly = true + args = args[1:] + } + flags := flag.NewFlagSet("pawsql", flag.ContinueOnError) + flags.SetOutput(os.Stderr) + configPath = "Barkfile" + flags.StringVar(&configPath, "config", configPath, "path to Barkfile") + if err := flags.Parse(args); err != nil { + return false, "", err + } + if flags.NArg() != 0 { + return false, "", fmt.Errorf("unexpected arguments: %v", flags.Args()) + } + return validateOnly, configPath, nil +} diff --git a/examples/Barkfile b/examples/Barkfile new file mode 100644 index 0000000..4e4e6df --- /dev/null +++ b/examples/Barkfile @@ -0,0 +1,18 @@ +pawsql { + listen :5432 + + tls { + cert /etc/pawsql/tls/fullchain.pem + key /etc/pawsql/tls/privkey.pem + } + + database bark_bistro { + hostname bistro.pawsql.test + upstream bark-bistro:5432 + } + + database wagging_tail { + hostname tail.pawsql.test + upstream wagging-tail:5432 + } +} diff --git a/examples/docker-compose.yml b/examples/docker-compose.yml new file mode 100644 index 0000000..a511461 --- /dev/null +++ b/examples/docker-compose.yml @@ -0,0 +1,49 @@ +services: + certificates: + image: alpine/openssl:latest + command: + - sh + - -ec + - >- + openssl req -x509 -newkey rsa:2048 -nodes -days 7 + -keyout /certs/privkey.pem -out /certs/fullchain.pem + -subj /CN=*.pawsql.test + volumes: + - certificates:/certs + + pawsql: + build: .. + depends_on: + certificates: + condition: service_completed_successfully + bark-bistro: + condition: service_started + wagging-tail: + condition: service_started + ports: + - "5432:5432" + volumes: + - ./Barkfile:/etc/pawsql/Barkfile:ro + - certificates:/etc/pawsql/tls:ro + networks: + default: + aliases: + - bistro.pawsql.test + - tail.pawsql.test + + bark-bistro: + image: postgres:17-alpine + environment: + POSTGRES_DB: bark_bistro + POSTGRES_USER: bark_bistro + POSTGRES_PASSWORD: example-only + + wagging-tail: + image: postgres:17-alpine + environment: + POSTGRES_DB: wagging_tail + POSTGRES_USER: wagging_tail + POSTGRES_PASSWORD: example-only + +volumes: + certificates: diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..21d973e --- /dev/null +++ b/go.mod @@ -0,0 +1,3 @@ +module github.com/barkstack/pawsql + +go 1.24 diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..e86011c --- /dev/null +++ b/internal/config/config.go @@ -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 +} diff --git a/internal/config/parser.go b/internal/config/parser.go new file mode 100644 index 0000000..74ac418 --- /dev/null +++ b/internal/config/parser.go @@ -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 +} diff --git a/internal/config/parser_test.go b/internal/config/parser_test.go new file mode 100644 index 0000000..c0080eb --- /dev/null +++ b/internal/config/parser_test.go @@ -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") + } +} diff --git a/internal/config/validate.go b/internal/config/validate.go new file mode 100644 index 0000000..999c9b2 --- /dev/null +++ b/internal/config/validate.go @@ -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 +} diff --git a/internal/pgwire/sslrequest.go b/internal/pgwire/sslrequest.go new file mode 100644 index 0000000..2696735 --- /dev/null +++ b/internal/pgwire/sslrequest.go @@ -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 +} diff --git a/internal/pgwire/startup.go b/internal/pgwire/startup.go new file mode 100644 index 0000000..cdf374c --- /dev/null +++ b/internal/pgwire/startup.go @@ -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) +} diff --git a/internal/pgwire/startup_test.go b/internal/pgwire/startup_test.go new file mode 100644 index 0000000..889472e --- /dev/null +++ b/internal/pgwire/startup_test.go @@ -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 } diff --git a/internal/postgres/provisioner.go b/internal/postgres/provisioner.go new file mode 100644 index 0000000..809641e --- /dev/null +++ b/internal/postgres/provisioner.go @@ -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" +} diff --git a/internal/postgres/provisioner_integration_test.go b/internal/postgres/provisioner_integration_test.go new file mode 100644 index 0000000..df237c8 --- /dev/null +++ b/internal/postgres/provisioner_integration_test.go @@ -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) + } +} diff --git a/internal/postgres/provisioner_test.go b/internal/postgres/provisioner_test.go new file mode 100644 index 0000000..80f0413 --- /dev/null +++ b/internal/postgres/provisioner_test.go @@ -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) + } + } +} diff --git a/internal/postgres/resolver.go b/internal/postgres/resolver.go new file mode 100644 index 0000000..4afef0c --- /dev/null +++ b/internal/postgres/resolver.go @@ -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) + } +} diff --git a/internal/postgres/resolver_test.go b/internal/postgres/resolver_test.go new file mode 100644 index 0000000..3266e4d --- /dev/null +++ b/internal/postgres/resolver_test.go @@ -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 +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go new file mode 100644 index 0000000..e245593 --- /dev/null +++ b/internal/proxy/proxy.go @@ -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) diff --git a/internal/router/routes.go b/internal/router/routes.go new file mode 100644 index 0000000..9c6bf40 --- /dev/null +++ b/internal/router/routes.go @@ -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 +} diff --git a/internal/router/routes_test.go b/internal/router/routes_test.go new file mode 100644 index 0000000..e3e1ef9 --- /dev/null +++ b/internal/router/routes_test.go @@ -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) + } +} diff --git a/internal/server/server.go b/internal/server/server.go new file mode 100644 index 0000000..7b37f73 --- /dev/null +++ b/internal/server/server.go @@ -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, ".")) +} diff --git a/internal/server/server_test.go b/internal/server/server_test.go new file mode 100644 index 0000000..e83714c --- /dev/null +++ b/internal/server/server_test.go @@ -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}} +}