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:
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, "."))
|
||||
}
|
||||
Reference in New Issue
Block a user