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) }