diff --git a/internal/pgwire/sslrequest.go b/internal/pgwire/sslrequest.go index 2696735..4bb22a8 100644 --- a/internal/pgwire/sslrequest.go +++ b/internal/pgwire/sslrequest.go @@ -6,6 +6,8 @@ import ( "errors" "fmt" "io" + "net" + "time" ) const ( @@ -13,7 +15,10 @@ const ( sslRequestCode uint32 = 80877103 ) -var ErrNotSSLRequest = errors.New("expected PostgreSQL SSLRequest") +var ( + ErrNotSSLRequest = errors.New("expected PostgreSQL SSLRequest") + ErrNotNegotiation = errors.New("not a PostgreSQL SSL or GSSENC negotiation request") +) // ReadSSLRequest validates the eight-byte PostgreSQL SSL negotiation request. func ReadSSLRequest(reader io.Reader) error { @@ -30,6 +35,42 @@ func ReadSSLRequest(reader io.Reader) error { return nil } +// NegotiatePlainPostgreSQL handles a plaintext server's first protocol exchange +// and returns the connection positioned at the StartupMessage. When the client +// opens with an SSLRequest or GSSENCRequest it answers 'N', the standard +// PostgreSQL decline, so the client retries without encryption. +func NegotiatePlainPostgreSQL(connection net.Conn) (net.Conn, error) { + var preamble [8]byte + if err := connection.SetReadDeadline(time.Now().Add(15 * time.Second)); err != nil { + return connection, fmt.Errorf("set negotiation deadline: %w", err) + } + if _, err := io.ReadFull(connection, preamble[:]); err != nil { + return connection, fmt.Errorf("read PostgreSQL negotiation preamble: %w", err) + } + if err := connection.SetReadDeadline(time.Time{}); err != nil { + return connection, fmt.Errorf("clear negotiation deadline: %w", err) + } + length := binary.BigEndian.Uint32(preamble[0:4]) + code := binary.BigEndian.Uint32(preamble[4:8]) + if length == sslRequestLength && (code == sslRequestCode || code == gssEncryptionCode) { + if err := declineSSL(connection); err != nil { + return connection, err + } + return connection, nil + } + if length < 8 || length > maxStartupMessageSize { + return connection, fmt.Errorf("%w: invalid length %d", ErrInvalidStartupMessage, length) + } + return Replay(connection, preamble[:]), nil +} + +func declineSSL(connection net.Conn) error { + if _, err := connection.Write([]byte{'N'}); err != nil { + return fmt.Errorf("write PostgreSQL SSL decline: %w", err) + } + 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 { diff --git a/internal/pgwire/startup.go b/internal/pgwire/startup.go index cdf374c..30350e8 100644 --- a/internal/pgwire/startup.go +++ b/internal/pgwire/startup.go @@ -11,7 +11,8 @@ import ( const ( startupProtocolVersion uint32 = 196608 - maxStartupMessageSize = 64 << 10 + maxStartupMessageSize uint32 = 64 << 10 + gssEncryptionCode uint32 = 80877104 ) var ( diff --git a/internal/server/server.go b/internal/server/server.go index 4267611..9663e84 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -145,6 +145,14 @@ func (s *Server) handle(connection net.Conn) { resolveContext, cancelResolve := context.WithTimeout(context.Background(), resolveTimeout) defer cancelResolve() + if s.TLSConfig == nil { + negotiated, negotiateErr := pgwire.NegotiatePlainPostgreSQL(connection) + if negotiateErr != nil { + s.Logger.Warn("connection rejected", "remote_address", remote, "result", "invalid_negotiation", "error", negotiateErr) + return + } + proxyClient = negotiated + } if sni != "" { backend, err = s.Resolver.Resolve(resolveContext, sni) if err != nil { diff --git a/internal/server/server_test.go b/internal/server/server_test.go index d87988f..3323a82 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -257,6 +257,79 @@ func TestServerRoutesPlaintextConnectionsByDatabase(t *testing.T) { } } +func TestServerDeclinesSSLNegotiationInPlaintextMode(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")) + }() + + staticResolver, err := router.NewStaticResolver([]config.DatabaseConfig{{ + Name: "analytics", + Upstream: backendListener.Addr().String(), + }}) + if err != nil { + t.Fatal(err) + } + routingServer, err := New(nil, staticResolver, 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] != 'N' { + t.Fatalf("SSL response = %q, want N", response) + } + if _, err := connection.Write(startupMessage("analytics")); err != nil { + t.Fatal(err) + } + response = make([]byte, len("backend")) + if _, err := io.ReadFull(connection, 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 the StartupMessage after SSL decline") + } +} + func startupMessage(database string) []byte { body := make([]byte, 4) binary.BigEndian.PutUint32(body, 196608)