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:
40
internal/pgwire/sslrequest.go
Normal file
40
internal/pgwire/sslrequest.go
Normal file
@@ -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
|
||||
}
|
||||
108
internal/pgwire/startup.go
Normal file
108
internal/pgwire/startup.go
Normal file
@@ -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)
|
||||
}
|
||||
76
internal/pgwire/startup_test.go
Normal file
76
internal/pgwire/startup_test.go
Normal file
@@ -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 }
|
||||
Reference in New Issue
Block a user