110 lines
3.3 KiB
Go
110 lines
3.3 KiB
Go
package pgwire
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/binary"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
)
|
|
|
|
const (
|
|
startupProtocolVersion uint32 = 196608
|
|
maxStartupMessageSize uint32 = 64 << 10
|
|
gssEncryptionCode uint32 = 80877104
|
|
)
|
|
|
|
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)
|
|
}
|