Jelajahi Sumber

harden PostgreSQL wire protocol

Danilo Fragoso 2 minggu lalu
induk
melakukan
395f0de408

+ 42 - 7
pkg/pgserver/connection.go

@@ -237,7 +237,11 @@ func (c *Connection) handleMessage(msg *Message) error {
 
 // handleQuery processes a simple query
 func (c *Connection) handleQuery(msg *Message) error {
-	// Parse query string (null-terminated)
+	// A Query message carries a single NUL-terminated query string.
+	if len(msg.Data) == 0 || msg.Data[len(msg.Data)-1] != 0 {
+		c.sendError("ERROR", ErrCodeProtocolViolation, "invalid Query message: missing null terminator")
+		return c.sendReadyForQuery()
+	}
 	sql := string(msg.Data[:len(msg.Data)-1])
 
 	if !c.quiet {
@@ -257,6 +261,15 @@ func (c *Connection) handleQuery(msg *Message) error {
 	sqlUpper := strings.ToUpper(strings.TrimSpace(sql))
 	singleStatement := isSingleStatementSQL(sql)
 
+	// In a failed transaction only ROLLBACK (or ROLLBACK TO SAVEPOINT) is
+	// accepted. Gate before honoring the special driver queries below so they
+	// return 25P02 instead of a result. Multi-statement batches are gated
+	// per statement in the execution loop below.
+	if c.txStatus == TxStatusFailed && singleStatement && !isRollbackStatement(sql) {
+		c.sendError("ERROR", ErrCodeTransactionAborted, "current transaction is aborted, commands ignored until end of transaction block")
+		return c.sendReadyForQuery()
+	}
+
 	// lib/pq and other drivers query these for connection validation
 	if singleStatement && strings.Contains(sqlUpper, "SELECT VERSION()") {
 		// Return a fake PostgreSQL version
@@ -450,6 +463,12 @@ func (c *Connection) handleExecute(msg *Message) error {
 		return c.failExtended(ErrCodeFeatureNotSupported, fmt.Errorf("portal can only be executed once"))
 	}
 	portal.executed = true
+	// Check the transaction state before catalog emulation. Catalog queries do
+	// not all parse as regular PizzaSQL statements, but must still return 25P02
+	// while the transaction is aborted.
+	if c.txStatus == TxStatusFailed && !isRollbackStatement(portal.query) {
+		return c.failExtended(ErrCodeTransactionAborted, fmt.Errorf("current transaction is aborted, commands ignored until end of transaction block"))
+	}
 	if result, handled, catalogErr := c.catalogResult(portal.query); handled {
 		if catalogErr != nil {
 			return c.failExtended(ErrCodeInternalError, catalogErr)
@@ -461,11 +480,6 @@ func (c *Connection) handleExecute(msg *Message) error {
 	if err != nil {
 		return c.failExtended(ErrCodeSyntaxError, fmt.Errorf("syntax error: %w", err))
 	}
-	if c.txStatus == TxStatusFailed {
-		if _, rollback := stmt.(*parser.RollbackStmt); !rollback {
-			return c.failExtended(ErrCodeTransactionAborted, fmt.Errorf("current transaction is aborted, commands ignored until end of transaction block"))
-		}
-	}
 	result, err := c.executor.Execute(stmt)
 	if err != nil {
 		if c.txStatus == TxStatusInBlock {
@@ -642,6 +656,17 @@ func isSingleStatementSQL(sql string) bool {
 	return true
 }
 
+// isRollbackStatement reports whether sql parses as a single ROLLBACK
+// statement, including ROLLBACK TO SAVEPOINT.
+func isRollbackStatement(sql string) bool {
+	stmt, err := parser.New(lexer.New(sql)).Parse()
+	if err != nil {
+		return false
+	}
+	_, ok := stmt.(*parser.RollbackStmt)
+	return ok
+}
+
 func catalogProjection(sql string, rows []map[string]interface{}) []string {
 	upper := strings.ToUpper(sql)
 	selectPos, fromPos := strings.Index(upper, "SELECT"), strings.Index(upper, " FROM ")
@@ -901,7 +926,7 @@ func (c *Connection) sendResult(result *executor.Result, stmt parser.Statement)
 
 // getCommandTag returns the command completion tag
 func (c *Connection) getCommandTag(stmt parser.Statement, result *executor.Result) string {
-	switch stmt.(type) {
+	switch s := stmt.(type) {
 	case *parser.CreateTableStmt:
 		return "CREATE TABLE"
 	case *parser.DropTableStmt:
@@ -925,8 +950,18 @@ func (c *Connection) getCommandTag(stmt parser.Statement, result *executor.Resul
 		c.txStatus = TxStatusIdle
 		return "COMMIT"
 	case *parser.RollbackStmt:
+		if s.Savepoint != "" {
+			c.txStatus = TxStatusInBlock
+			return "ROLLBACK"
+		}
 		c.txStatus = TxStatusIdle
 		return "ROLLBACK"
+	case *parser.SavepointStmt:
+		c.txStatus = TxStatusInBlock
+		return "SAVEPOINT"
+	case *parser.ReleaseStmt:
+		c.txStatus = TxStatusInBlock
+		return "RELEASE"
 	default:
 		return "OK"
 	}

+ 205 - 1
pkg/pgserver/connection_test.go

@@ -1,6 +1,15 @@
 package pgserver
 
-import "testing"
+import (
+	"bufio"
+	"bytes"
+	"net"
+	"testing"
+
+	"github.com/danfragoso/pizzasql-next/pkg/executor"
+	"github.com/danfragoso/pizzasql-next/pkg/lexer"
+	"github.com/danfragoso/pizzasql-next/pkg/parser"
+)
 
 func TestBindQuery(t *testing.T) {
 	query, err := bindQuery(
@@ -26,3 +35,198 @@ func TestBindQueryRequiresEveryParameter(t *testing.T) {
 		t.Fatal("expected missing parameter error")
 	}
 }
+
+func parseStmt(t *testing.T, sql string) parser.Statement {
+	t.Helper()
+	l := lexer.New(sql)
+	p := parser.New(l)
+	stmt, err := p.Parse()
+	if err != nil {
+		t.Fatalf("parse %q: %v", sql, err)
+	}
+	return stmt
+}
+
+func TestGetCommandTagSavepointKeepsTransaction(t *testing.T) {
+	c := &Connection{txStatus: TxStatusIdle}
+	tag := c.getCommandTag(parseStmt(t, "SAVEPOINT sp1"), executor.NewResult("SAVEPOINT"))
+	if tag != "SAVEPOINT" {
+		t.Fatalf("tag = %q, want %q", tag, "SAVEPOINT")
+	}
+	if c.txStatus != TxStatusInBlock {
+		t.Fatalf("txStatus = %c, want %c", c.txStatus, TxStatusInBlock)
+	}
+}
+
+func TestGetCommandTagReleaseKeepsTransaction(t *testing.T) {
+	c := &Connection{txStatus: TxStatusInBlock}
+	tag := c.getCommandTag(parseStmt(t, "RELEASE SAVEPOINT sp1"), executor.NewResult("RELEASE"))
+	if tag != "RELEASE" {
+		t.Fatalf("tag = %q, want %q", tag, "RELEASE")
+	}
+	if c.txStatus != TxStatusInBlock {
+		t.Fatalf("txStatus = %c, want %c", c.txStatus, TxStatusInBlock)
+	}
+}
+
+func TestGetCommandTagRollbackToSavepointKeepsTransaction(t *testing.T) {
+	c := &Connection{txStatus: TxStatusFailed}
+	tag := c.getCommandTag(parseStmt(t, "ROLLBACK TO SAVEPOINT sp1"), executor.NewResult("ROLLBACK"))
+	if tag != "ROLLBACK" {
+		t.Fatalf("tag = %q, want %q", tag, "ROLLBACK")
+	}
+	if c.txStatus != TxStatusInBlock {
+		t.Fatalf("txStatus = %c, want %c", c.txStatus, TxStatusInBlock)
+	}
+}
+
+func TestGetCommandTagFullRollbackEndsTransaction(t *testing.T) {
+	c := &Connection{txStatus: TxStatusFailed}
+	tag := c.getCommandTag(parseStmt(t, "ROLLBACK"), executor.NewResult("ROLLBACK"))
+	if tag != "ROLLBACK" {
+		t.Fatalf("tag = %q, want %q", tag, "ROLLBACK")
+	}
+	if c.txStatus != TxStatusIdle {
+		t.Fatalf("txStatus = %c, want %c", c.txStatus, TxStatusIdle)
+	}
+}
+
+func newTestConnection(t *testing.T) (*Connection, net.Conn) {
+	t.Helper()
+	server, client := net.Pipe()
+	t.Cleanup(func() {
+		server.Close()
+		client.Close()
+	})
+	c := &Connection{
+		conn:       server,
+		reader:     bufio.NewReader(server),
+		writer:     bufio.NewWriter(server),
+		params:     map[string]string{"user": "tester"},
+		statements: make(map[string]*preparedStatement),
+		portals:    make(map[string]*portal),
+		txStatus:   TxStatusIdle,
+		quiet:      true,
+	}
+	return c, client
+}
+
+func runQuery(t *testing.T, c *Connection, client net.Conn, msg *Message) []*Message {
+	t.Helper()
+	done := make(chan error, 1)
+	go func() { done <- c.handleQuery(msg) }()
+
+	reader := bufio.NewReader(client)
+	var msgs []*Message
+	for {
+		m, err := ReadMessage(reader)
+		if err != nil {
+			t.Fatalf("read response: %v", err)
+		}
+		msgs = append(msgs, m)
+		if m.Type == MsgReadyForQuery {
+			break
+		}
+	}
+	if err := <-done; err != nil {
+		t.Fatalf("handleQuery: %v", err)
+	}
+	return msgs
+}
+
+func errorCode(msg *Message) string {
+	if msg == nil || msg.Type != MsgErrorResponse {
+		return ""
+	}
+	data := msg.Data
+	for i := 0; i < len(data); {
+		if data[i] == 0 {
+			break
+		}
+		field := data[i]
+		i++
+		end := bytes.IndexByte(data[i:], 0)
+		if end < 0 {
+			break
+		}
+		value := string(data[i : i+end])
+		if field == ErrorFieldCode {
+			return value
+		}
+		i += end + 1
+	}
+	return ""
+}
+
+func TestHandleQueryEmptyPayloadReturnsProtocolError(t *testing.T) {
+	c, client := newTestConnection(t)
+	msgs := runQuery(t, c, client, &Message{Type: MsgQuery, Data: []byte{}})
+	if len(msgs) != 2 {
+		t.Fatalf("got %d messages, want 2", len(msgs))
+	}
+	if code := errorCode(msgs[0]); code != ErrCodeProtocolViolation {
+		t.Fatalf("error code = %q, want %q", code, ErrCodeProtocolViolation)
+	}
+}
+
+func TestHandleQueryMissingNullTerminatorReturnsProtocolError(t *testing.T) {
+	c, client := newTestConnection(t)
+	msgs := runQuery(t, c, client, &Message{Type: MsgQuery, Data: []byte("SELECT 1")})
+	if len(msgs) != 2 {
+		t.Fatalf("got %d messages, want 2", len(msgs))
+	}
+	if code := errorCode(msgs[0]); code != ErrCodeProtocolViolation {
+		t.Fatalf("error code = %q, want %q", code, ErrCodeProtocolViolation)
+	}
+}
+
+func TestHandleQueryEmptyQueryReturnsEmptyQueryResponse(t *testing.T) {
+	for _, data := range [][]byte{{0}, []byte(";\x00"), []byte("   \x00")} {
+		c, client := newTestConnection(t)
+		msgs := runQuery(t, c, client, &Message{Type: MsgQuery, Data: data})
+		if len(msgs) != 2 {
+			t.Fatalf("payload %q: got %d messages, want 2", data, len(msgs))
+		}
+		if msgs[0].Type != MsgEmptyQueryResponse {
+			t.Fatalf("payload %q: first message type = %c, want %c", data, msgs[0].Type, MsgEmptyQueryResponse)
+		}
+	}
+}
+
+func TestHandleQueryFailedTransactionInterceptsSpecialQueries(t *testing.T) {
+	queries := []string{
+		"SELECT version()",
+		"SELECT current_user",
+		"SHOW server_version",
+		"SELECT * FROM information_schema.tables",
+	}
+	for _, q := range queries {
+		c, client := newTestConnection(t)
+		c.txStatus = TxStatusFailed
+		msgs := runQuery(t, c, client, &Message{Type: MsgQuery, Data: []byte(q + "\x00")})
+		if len(msgs) != 2 {
+			t.Fatalf("query %q: got %d messages, want 2", q, len(msgs))
+		}
+		if code := errorCode(msgs[0]); code != ErrCodeTransactionAborted {
+			t.Fatalf("query %q: error code = %q, want %q", q, code, ErrCodeTransactionAborted)
+		}
+	}
+}
+
+func TestIsRollbackStatement(t *testing.T) {
+	for _, sql := range []string{"ROLLBACK", "ROLLBACK TO SAVEPOINT sp1", "rollback;", "  rollback  ;"} {
+		if !isRollbackStatement(sql) {
+			t.Errorf("isRollbackStatement(%q) = false, want true", sql)
+		}
+	}
+	for _, sql := range []string{"SELECT version()", "SHOW server_version", "SAVEPOINT sp1"} {
+		if isRollbackStatement(sql) {
+			t.Errorf("isRollbackStatement(%q) = true, want false", sql)
+		}
+	}
+	for _, sql := range []string{"ROLLBACKBOGUS", "ROLLBACK; SELECT 1"} {
+		if isRollbackStatement(sql) {
+			t.Errorf("isRollbackStatement(%q) = true, want false", sql)
+		}
+	}
+}

+ 15 - 0
pkg/pgserver/protocol.go

@@ -6,6 +6,15 @@ import (
 	"io"
 )
 
+// Maximum message sizes enforced before any allocation, to bound memory use
+// against malformed or hostile clients. Regular protocol messages (queries,
+// bind parameters, etc.) are capped at 16 MiB; startup messages are much
+// smaller and capped at 1 MiB.
+const (
+	MaxMessageSize        = 16 * 1024 * 1024 // 16 MiB
+	MaxStartupMessageSize = 1 * 1024 * 1024  // 1 MiB
+)
+
 // Message type constants (first byte of message)
 const (
 	// Frontend (client) messages
@@ -128,6 +137,9 @@ func ReadMessage(r io.Reader) (*Message, error) {
 	if length < 4 {
 		return nil, fmt.Errorf("invalid message length: %d", length)
 	}
+	if length > MaxMessageSize {
+		return nil, fmt.Errorf("message length %d exceeds maximum %d", length, MaxMessageSize)
+	}
 
 	// Read message data
 	data := make([]byte, length-4)
@@ -152,6 +164,9 @@ func ReadStartupMessage(r io.Reader) (map[string]string, error) {
 	if length < 8 {
 		return nil, fmt.Errorf("invalid startup message length: %d", length)
 	}
+	if length > MaxStartupMessageSize {
+		return nil, fmt.Errorf("startup message length %d exceeds maximum %d", length, MaxStartupMessageSize)
+	}
 
 	// Read protocol version
 	var version uint32

+ 75 - 0
pkg/pgserver/protocol_test.go

@@ -0,0 +1,75 @@
+package pgserver
+
+import (
+	"bytes"
+	"encoding/binary"
+	"testing"
+)
+
+func TestReadMessageRejectsOversizedMessage(t *testing.T) {
+	var buf bytes.Buffer
+	buf.WriteByte('Q')
+	if err := binary.Write(&buf, binary.BigEndian, uint32(MaxMessageSize+1)); err != nil {
+		t.Fatal(err)
+	}
+	if _, err := ReadMessage(&buf); err == nil {
+		t.Fatal("expected error for message larger than MaxMessageSize")
+	}
+}
+
+func TestReadMessageAcceptsNormalMessage(t *testing.T) {
+	payload := []byte("SELECT 1\x00")
+	var buf bytes.Buffer
+	buf.WriteByte(MsgQuery)
+	if err := binary.Write(&buf, binary.BigEndian, uint32(len(payload)+4)); err != nil {
+		t.Fatal(err)
+	}
+	buf.Write(payload)
+
+	msg, err := ReadMessage(&buf)
+	if err != nil {
+		t.Fatalf("ReadMessage: %v", err)
+	}
+	if msg.Type != MsgQuery {
+		t.Fatalf("type = %c, want %c", msg.Type, MsgQuery)
+	}
+	if !bytes.Equal(msg.Data, payload) {
+		t.Fatalf("data = %q, want %q", msg.Data, payload)
+	}
+}
+
+func TestReadStartupMessageRejectsOversizedMessage(t *testing.T) {
+	var buf bytes.Buffer
+	if err := binary.Write(&buf, binary.BigEndian, uint32(MaxStartupMessageSize+1)); err != nil {
+		t.Fatal(err)
+	}
+	if _, err := ReadStartupMessage(&buf); err == nil {
+		t.Fatal("expected error for startup message larger than MaxStartupMessageSize")
+	}
+}
+
+func TestReadStartupMessageAcceptsNormalMessage(t *testing.T) {
+	// Build a minimal startup payload: protocol version + user param.
+	var payload bytes.Buffer
+	if err := binary.Write(&payload, binary.BigEndian, uint32(196608)); err != nil { // 3.0
+		t.Fatal(err)
+	}
+	payload.WriteString("user\x00tester\x00\x00")
+
+	var buf bytes.Buffer
+	if err := binary.Write(&buf, binary.BigEndian, uint32(payload.Len()+4)); err != nil {
+		t.Fatal(err)
+	}
+	buf.Write(payload.Bytes())
+
+	params, err := ReadStartupMessage(&buf)
+	if err != nil {
+		t.Fatalf("ReadStartupMessage: %v", err)
+	}
+	if params["user"] != "tester" {
+		t.Fatalf("user = %q, want %q", params["user"], "tester")
+	}
+	if params["protocol_version"] != "196608" {
+		t.Fatalf("protocol_version = %q, want %q", params["protocol_version"], "196608")
+	}
+}