Przeglądaj źródła

fix PostgreSQL wire consistency and compatibility

Danilo Fragoso 2 tygodni temu
rodzic
commit
dc341a8a79

+ 3 - 2
pkg/analyzer/analyzer.go

@@ -576,8 +576,9 @@ func (a *Analyzer) analyzeCreateTable(stmt *parser.CreateTableStmt) error {
 		}
 	}
 
-	// Add to catalog
-	return a.catalog.CreateTable(tableInfo)
+	// Analysis is validation-only. Publishing the table before the durable
+	// schema write can leave a phantom catalog entry when that write fails.
+	return nil
 }
 
 // analyzeDropTable analyzes a DROP TABLE statement.

+ 154 - 15
pkg/executor/executor.go

@@ -29,9 +29,10 @@ type Executor struct {
 	currentDatabase   string                         // current database alias (default is "main")
 
 	// Transaction state
-	inTransaction bool
-	savepoints    []string     // stack of savepoint names
-	txLog         []txLogEntry // transaction log for rollback
+	inTransaction      bool
+	savepoints         []string // stack of savepoint names
+	savepointPositions []int
+	txLog              []txLogEntry // transaction log for rollback
 
 	// Subquery context for correlated subqueries
 	outerRow storage.Row
@@ -132,6 +133,10 @@ func (e *Executor) SyncCatalog() error {
 
 // Execute executes a SQL statement.
 func (e *Executor) Execute(stmt parser.Statement) (*Result, error) {
+	if _, beginning := stmt.(*parser.BeginStmt); !beginning && !e.inTransaction {
+		e.schema.LockStatement()
+		defer e.schema.UnlockStatement()
+	}
 	e.subqueryCache = make(map[*parser.SelectStmt]*Result)
 	e.correlatedAggCache = make(map[*parser.SelectStmt]*correlatedAggCache)
 	defer func() {
@@ -2144,7 +2149,18 @@ func (e *Executor) executeInsert(stmt *parser.InsertStmt) (*Result, error) {
 			}
 			rows = append(rows, row)
 		}
-		count, err := e.table.InsertBulk(tableName, rows)
+		var count int
+		if e.inTransaction {
+			for _, row := range rows {
+				if err := e.table.Insert(tableName, row); err != nil {
+					return nil, err
+				}
+				e.txLog = append(e.txLog, txLogEntry{operation: "INSERT", table: tableName, key: fmt.Sprintf("%v", row[schema.PrimaryKey])})
+				count++
+			}
+		} else {
+			count, err = e.table.InsertBulk(tableName, rows)
+		}
 		if err != nil {
 			return nil, err
 		}
@@ -2183,6 +2199,46 @@ func (e *Executor) executeInsert(stmt *parser.InsertStmt) (*Result, error) {
 
 		err := e.table.Insert(tableName, row)
 		if err != nil {
+			if strings.Contains(err.Error(), "duplicate") && (stmt.ConflictDoNothing || len(stmt.ConflictUpdate) > 0) {
+				if stmt.ConflictDoNothing {
+					continue
+				}
+				if len(stmt.ConflictTarget) > 0 && !containsFold(stmt.ConflictTarget, schema.PrimaryKey) {
+					return nil, fmt.Errorf("ON CONFLICT target must include primary key %s", schema.PrimaryKey)
+				}
+				pkValue := row[schema.PrimaryKey]
+				var oldRows []storage.Row
+				if e.inTransaction {
+					oldRows, _ = e.table.Select(tableName, func(existing storage.Row) bool {
+						return fmt.Sprintf("%v", existing[schema.PrimaryKey]) == fmt.Sprintf("%v", pkValue)
+					})
+				}
+				updated, updateErr := e.table.UpdateFunc(tableName, func(existing storage.Row) (storage.Row, error) {
+					context := e.addTableAlias(existing, tableName)
+					updates := make(storage.Row)
+					for _, assignment := range stmt.ConflictUpdate {
+						value, evalErr := e.evalExpr(assignment.Value, context)
+						if evalErr != nil {
+							return nil, evalErr
+						}
+						updates[assignment.Column] = value
+					}
+					return updates, nil
+				}, func(existing storage.Row) bool {
+					return fmt.Sprintf("%v", existing[schema.PrimaryKey]) == fmt.Sprintf("%v", pkValue)
+				})
+				if updateErr != nil {
+					return nil, updateErr
+				}
+				if updated != 1 {
+					return nil, fmt.Errorf("ON CONFLICT row disappeared during update")
+				}
+				if e.inTransaction && len(oldRows) == 1 {
+					e.txLog = append(e.txLog, txLogEntry{operation: "UPDATE", table: tableName, key: fmt.Sprintf("%v", pkValue), oldData: oldRows[0]})
+				}
+				count++
+				continue
+			}
 			// Handle conflict based on OnConflict action
 			if strings.Contains(err.Error(), "duplicate") {
 				switch stmt.OnConflict {
@@ -2213,6 +2269,9 @@ func (e *Executor) executeInsert(stmt *parser.InsertStmt) (*Result, error) {
 				return nil, err
 			}
 		}
+		if e.inTransaction {
+			e.txLog = append(e.txLog, txLogEntry{operation: "INSERT", table: tableName, key: fmt.Sprintf("%v", row[schema.PrimaryKey])})
+		}
 		count++
 	}
 
@@ -2221,9 +2280,22 @@ func (e *Executor) executeInsert(stmt *parser.InsertStmt) (*Result, error) {
 	return result, nil
 }
 
+func containsFold(values []string, target string) bool {
+	for _, value := range values {
+		if strings.EqualFold(value, target) {
+			return true
+		}
+	}
+	return false
+}
+
 // executeUpdate executes an UPDATE statement.
 func (e *Executor) executeUpdate(stmt *parser.UpdateStmt) (*Result, error) {
 	tableName := stmt.Table.Name
+	schema, err := e.schema.GetSchema(tableName)
+	if err != nil {
+		return nil, err
+	}
 
 	// Build filter
 	var filter func(storage.Row) bool
@@ -2250,10 +2322,20 @@ func (e *Executor) executeUpdate(stmt *parser.UpdateStmt) (*Result, error) {
 		return updates, nil
 	}
 
+	var oldRows []storage.Row
+	if e.inTransaction {
+		oldRows, err = e.table.Select(tableName, filter)
+		if err != nil {
+			return nil, err
+		}
+	}
 	count, err := e.table.UpdateFunc(tableName, updateFn, filter)
 	if err != nil {
 		return nil, err
 	}
+	for i := 0; e.inTransaction && i < count && i < len(oldRows); i++ {
+		e.txLog = append(e.txLog, txLogEntry{operation: "UPDATE", table: tableName, key: fmt.Sprintf("%v", oldRows[i][schema.PrimaryKey]), oldData: oldRows[i]})
+	}
 
 	result := NewResult("UPDATE")
 	result.SetRowCount(count)
@@ -2263,6 +2345,10 @@ func (e *Executor) executeUpdate(stmt *parser.UpdateStmt) (*Result, error) {
 // executeDelete executes a DELETE statement.
 func (e *Executor) executeDelete(stmt *parser.DeleteStmt) (*Result, error) {
 	tableName := stmt.Table.Name
+	schema, err := e.schema.GetSchema(tableName)
+	if err != nil {
+		return nil, err
+	}
 
 	// Build filter
 	var filter func(storage.Row) bool
@@ -2276,10 +2362,20 @@ func (e *Executor) executeDelete(stmt *parser.DeleteStmt) (*Result, error) {
 		}
 	}
 
+	var oldRows []storage.Row
+	if e.inTransaction {
+		oldRows, err = e.table.Select(tableName, filter)
+		if err != nil {
+			return nil, err
+		}
+	}
 	count, err := e.table.Delete(tableName, filter)
 	if err != nil {
 		return nil, err
 	}
+	for i := 0; e.inTransaction && i < count && i < len(oldRows); i++ {
+		e.txLog = append(e.txLog, txLogEntry{operation: "DELETE", table: tableName, key: fmt.Sprintf("%v", oldRows[i][schema.PrimaryKey]), oldData: oldRows[i]})
+	}
 
 	result := NewResult("DELETE")
 	result.SetRowCount(count)
@@ -2344,11 +2440,15 @@ func (e *Executor) executeCreateTable(stmt *parser.CreateTableStmt) (*Result, er
 	}
 
 	if err := e.schema.CreateTable(schema); err != nil {
+		if stmt.IfNotExists && strings.Contains(err.Error(), "table already exists") {
+			return NewResult("CREATE TABLE"), nil
+		}
 		return nil, err
 	}
 
-	// Update analyzer catalog
-	e.catalog.CreateTable(schema.ToAnalyzerTableInfo())
+	if err := e.SyncCatalog(); err != nil {
+		return nil, err
+	}
 
 	result := NewResult("CREATE TABLE")
 	return result, nil
@@ -2385,8 +2485,9 @@ func (e *Executor) executeDropTable(stmt *parser.DropTableStmt) (*Result, error)
 			return nil, err
 		}
 
-		// Update analyzer catalog
-		e.catalog.DropTable(tableRef.Name)
+	}
+	if err := e.SyncCatalog(); err != nil {
+		return nil, err
 	}
 
 	result := NewResult("DROP TABLE")
@@ -2436,6 +2537,9 @@ func (e *Executor) executeCreateIndex(stmt *parser.CreateIndexStmt) (*Result, er
 	}
 
 	if err := e.schema.CreateIndex(index); err != nil {
+		if stmt.IfNotExists && strings.Contains(err.Error(), "index already exists") {
+			return NewResult("CREATE INDEX"), nil
+		}
 		return nil, err
 	}
 
@@ -2575,6 +2679,15 @@ func (e *Executor) executeAlterTable(stmt *parser.AlterTableStmt) (*Result, erro
 
 // executeAlterTableAddColumn adds a column to a table.
 func (e *Executor) executeAlterTableAddColumn(table string, action *parser.AddColumnAction) (*Result, error) {
+	if action.IfNotExists {
+		schema, err := e.schema.GetSchema(table)
+		if err != nil {
+			return nil, err
+		}
+		if _, exists := schema.GetColumn(action.Column.Name); exists {
+			return NewResult("ALTER TABLE"), nil
+		}
+	}
 	col := storage.Column{
 		Name:     action.Column.Name,
 		Type:     action.Column.Type.Name,
@@ -2598,6 +2711,9 @@ func (e *Executor) executeAlterTableAddColumn(table string, action *parser.AddCo
 	}
 
 	if err := e.schema.AddColumn(table, col); err != nil {
+		if action.IfNotExists && strings.Contains(err.Error(), "column already exists") {
+			return NewResult("ALTER TABLE"), nil
+		}
 		return nil, err
 	}
 
@@ -2657,7 +2773,9 @@ func (e *Executor) executeBegin(stmt *parser.BeginStmt) (*Result, error) {
 
 	e.inTransaction = true
 	e.savepoints = nil
+	e.savepointPositions = nil
 	e.txLog = nil
+	e.schema.BeginTransaction()
 
 	result := NewResult("BEGIN")
 	return result, nil
@@ -2672,7 +2790,9 @@ func (e *Executor) executeCommit(stmt *parser.CommitStmt) (*Result, error) {
 	// Clear transaction state
 	e.inTransaction = false
 	e.savepoints = nil
+	e.savepointPositions = nil
 	e.txLog = nil
+	e.schema.EndTransaction()
 
 	result := NewResult("COMMIT")
 	return result, nil
@@ -2690,18 +2810,25 @@ func (e *Executor) executeRollback(stmt *parser.RollbackStmt) (*Result, error) {
 	}
 
 	// Full rollback - undo all operations in reverse order
+	var rollbackErr error
 	for i := len(e.txLog) - 1; i >= 0; i-- {
 		entry := e.txLog[i]
 		if err := e.undoOperation(entry); err != nil {
-			// Log error but continue with rollback
-			continue
+			if rollbackErr == nil {
+				rollbackErr = err
+			}
 		}
 	}
 
 	// Clear transaction state
 	e.inTransaction = false
 	e.savepoints = nil
+	e.savepointPositions = nil
 	e.txLog = nil
+	e.schema.EndTransaction()
+	if rollbackErr != nil {
+		return nil, fmt.Errorf("rollback failed: %w", rollbackErr)
+	}
 
 	result := NewResult("ROLLBACK")
 	return result, nil
@@ -2713,10 +2840,12 @@ func (e *Executor) executeSavepoint(stmt *parser.SavepointStmt) (*Result, error)
 		// SQLite allows SAVEPOINT outside transaction (starts implicit transaction)
 		e.inTransaction = true
 		e.txLog = nil
+		e.schema.BeginTransaction()
 	}
 
 	// Add savepoint marker
 	e.savepoints = append(e.savepoints, stmt.Name)
+	e.savepointPositions = append(e.savepointPositions, len(e.txLog))
 
 	result := NewResult("SAVEPOINT")
 	return result, nil
@@ -2733,6 +2862,7 @@ func (e *Executor) executeRelease(stmt *parser.ReleaseStmt) (*Result, error) {
 	for i := len(e.savepoints) - 1; i >= 0; i-- {
 		if e.savepoints[i] == stmt.Name {
 			e.savepoints = e.savepoints[:i]
+			e.savepointPositions = e.savepointPositions[:i]
 			found = true
 			break
 		}
@@ -2828,25 +2958,34 @@ func (e *Executor) rollbackToSavepoint(name string) (*Result, error) {
 		return nil, fmt.Errorf("no such savepoint: %s", name)
 	}
 
-	// Count operations to undo (operations after the savepoint)
-	// For simplicity, we track savepoint positions by counting log entries
-	// In a real implementation, we'd track log positions per savepoint
-
 	// Undo operations in reverse order
-	for i := len(e.txLog) - 1; i >= 0; i-- {
+	logPosition := e.savepointPositions[savepointIdx]
+	for i := len(e.txLog) - 1; i >= logPosition; i-- {
 		entry := e.txLog[i]
 		if err := e.undoOperation(entry); err != nil {
 			continue
 		}
 	}
+	e.txLog = e.txLog[:logPosition]
 
 	// Remove savepoints after the target
 	e.savepoints = e.savepoints[:savepointIdx+1]
+	e.savepointPositions = e.savepointPositions[:savepointIdx+1]
 
 	result := NewResult("ROLLBACK")
 	return result, nil
 }
 
+// RollbackActive rolls back an open transaction, such as when its client
+// disconnects before sending COMMIT or ROLLBACK.
+func (e *Executor) RollbackActive() error {
+	if !e.inTransaction {
+		return nil
+	}
+	_, err := e.executeRollback(&parser.RollbackStmt{})
+	return err
+}
+
 // undoOperation reverses a single operation.
 func (e *Executor) undoOperation(entry txLogEntry) error {
 	switch entry.operation {

+ 8 - 8
pkg/executor/executor_test.go

@@ -1138,15 +1138,10 @@ func TestTransactions(t *testing.T) {
 			t.Error("expected inTransaction to be false after ROLLBACK")
 		}
 
-		// Verify data was NOT committed (rollback currently doesn't undo changes due to PizzaKV limitations)
-		// This is a known limitation - the transaction log is built but rollback doesn't restore state
+		// Verify data was not committed.
 		checkResult, _ := execSQL(exec, "SELECT * FROM tx_test WHERE id = 2")
-		// Note: In the current implementation, rollback doesn't actually undo changes
-		// This test documents current behavior
-		if checkResult.RowCount == 0 {
-			t.Log("ROLLBACK successfully prevented data persistence (ideal)")
-		} else {
-			t.Log("ROLLBACK did not undo changes (current limitation)")
+		if checkResult.RowCount != 0 {
+			t.Errorf("expected no rows after rollback, got %d", checkResult.RowCount)
 		}
 	})
 
@@ -1198,6 +1193,11 @@ func TestTransactions(t *testing.T) {
 		if !exec.inTransaction {
 			t.Error("expected to still be in transaction after ROLLBACK TO")
 		}
+		before, _ := execSQL(exec, "SELECT * FROM tx_test WHERE id = 10")
+		after, _ := execSQL(exec, "SELECT * FROM tx_test WHERE id = 11")
+		if before.RowCount != 1 || after.RowCount != 0 {
+			t.Errorf("unexpected savepoint rollback rows: before=%d after=%d", before.RowCount, after.RowCount)
+		}
 
 		execSQL(exec, "ROLLBACK")
 	})

+ 14 - 4
pkg/lexer/token.go

@@ -160,6 +160,8 @@ const (
 	TokenTIME
 	TokenTIMESTAMP
 	TokenDATETIME
+	TokenJSON
+	TokenJSONB
 
 	// Transaction keywords
 	TokenBEGIN
@@ -185,6 +187,9 @@ const (
 	TokenIGNORE
 	TokenFAIL
 	TokenABORT
+	TokenCONFLICT
+	TokenDO
+	TokenNOTHING
 )
 
 var keywords = map[string]TokenType{
@@ -309,6 +314,8 @@ var keywords = map[string]TokenType{
 	"TIME":      TokenTIME,
 	"TIMESTAMP": TokenTIMESTAMP,
 	"DATETIME":  TokenDATETIME,
+	"JSON":      TokenJSON,
+	"JSONB":     TokenJSONB,
 
 	// Transactions
 	"BEGIN":       TokenBEGIN,
@@ -330,10 +337,13 @@ var keywords = map[string]TokenType{
 	"REINDEX": TokenREINDEX,
 
 	// Conflict resolution
-	"REPLACE": TokenREPLACE,
-	"IGNORE":  TokenIGNORE,
-	"FAIL":    TokenFAIL,
-	"ABORT":   TokenABORT,
+	"REPLACE":  TokenREPLACE,
+	"IGNORE":   TokenIGNORE,
+	"FAIL":     TokenFAIL,
+	"ABORT":    TokenABORT,
+	"CONFLICT": TokenCONFLICT,
+	"DO":       TokenDO,
+	"NOTHING":  TokenNOTHING,
 }
 
 // LookupKeyword returns the token type for an identifier.

+ 10 - 6
pkg/parser/ast.go

@@ -114,11 +114,14 @@ const (
 
 // InsertStmt represents an INSERT statement.
 type InsertStmt struct {
-	Table      *TableRef
-	Columns    []string
-	Values     [][]Expr
-	Select     *SelectStmt    // INSERT ... SELECT
-	OnConflict ConflictAction // OR REPLACE/IGNORE/etc.
+	Table             *TableRef
+	Columns           []string
+	Values            [][]Expr
+	Select            *SelectStmt    // INSERT ... SELECT
+	OnConflict        ConflictAction // OR REPLACE/IGNORE/etc.
+	ConflictTarget    []string
+	ConflictUpdate    []Assignment
+	ConflictDoNothing bool
 }
 
 func (s *InsertStmt) node()     {}
@@ -278,7 +281,8 @@ type AlterAction interface {
 
 // AddColumnAction represents ADD COLUMN action.
 type AddColumnAction struct {
-	Column *ColumnDef
+	Column      *ColumnDef
+	IfNotExists bool
 }
 
 func (a *AddColumnAction) node()        {}

+ 75 - 2
pkg/parser/parser.go

@@ -34,6 +34,9 @@ func (p *Parser) Parse() (Statement, error) {
 	if p.curTokenIs(lexer.TokenSemicolon) {
 		p.nextToken()
 	}
+	if !p.curTokenIs(lexer.TokenEOF) {
+		return nil, p.curError("unexpected trailing token: " + p.curToken.Type.String())
+	}
 
 	return stmt, nil
 }
@@ -779,6 +782,62 @@ func (p *Parser) parseInsert() (*InsertStmt, error) {
 		return nil, p.curError("expected VALUES or SELECT")
 	}
 
+	if p.curTokenIs(lexer.TokenON) {
+		p.nextToken()
+		if !p.curTokenIs(lexer.TokenCONFLICT) {
+			return nil, p.curError("expected CONFLICT after ON")
+		}
+		p.nextToken()
+		if p.curTokenIs(lexer.TokenLParen) {
+			p.nextToken()
+			target, err := p.parseIdentList()
+			if err != nil {
+				return nil, err
+			}
+			stmt.ConflictTarget = target
+			if !p.curTokenIs(lexer.TokenRParen) {
+				return nil, p.curError("expected )")
+			}
+			p.nextToken()
+		}
+		if !p.curTokenIs(lexer.TokenDO) {
+			return nil, p.curError("expected DO after ON CONFLICT")
+		}
+		p.nextToken()
+		if p.curTokenIs(lexer.TokenNOTHING) {
+			stmt.ConflictDoNothing = true
+			p.nextToken()
+		} else if p.curTokenIs(lexer.TokenUPDATE) {
+			p.nextToken()
+			if !p.curTokenIs(lexer.TokenSET) {
+				return nil, p.curError("expected SET after DO UPDATE")
+			}
+			p.nextToken()
+			for {
+				if !p.curTokenIs(lexer.TokenIdent) {
+					return nil, p.curError("expected column name")
+				}
+				column := p.curToken.Literal
+				p.nextToken()
+				if !p.curTokenIs(lexer.TokenEq) {
+					return nil, p.curError("expected =")
+				}
+				p.nextToken()
+				value, err := p.parseExpr()
+				if err != nil {
+					return nil, err
+				}
+				stmt.ConflictUpdate = append(stmt.ConflictUpdate, Assignment{Column: column, Value: value})
+				if !p.curTokenIs(lexer.TokenComma) {
+					break
+				}
+				p.nextToken()
+			}
+		} else {
+			return nil, p.curError("expected NOTHING or UPDATE after DO")
+		}
+	}
+
 	return stmt, nil
 }
 
@@ -1075,7 +1134,7 @@ func (p *Parser) isDataTypeKeyword() bool {
 		lexer.TokenNUMERIC, lexer.TokenDECIMAL,
 		lexer.TokenTEXT, lexer.TokenVARCHAR, lexer.TokenCHAR, lexer.TokenCHARACTER, lexer.TokenCLOB, lexer.TokenNCHAR, lexer.TokenNVARCHAR,
 		lexer.TokenBLOB, lexer.TokenBOOLEAN,
-		lexer.TokenDATE, lexer.TokenTIME, lexer.TokenTIMESTAMP, lexer.TokenDATETIME:
+		lexer.TokenDATE, lexer.TokenTIME, lexer.TokenTIMESTAMP, lexer.TokenDATETIME, lexer.TokenJSON, lexer.TokenJSONB:
 		return true
 	}
 	return false
@@ -1522,13 +1581,27 @@ func (p *Parser) parseAlterTableAdd(stmt *AlterTableStmt) (*AlterTableStmt, erro
 		p.nextToken()
 	}
 
+	ifNotExists := false
+	if p.curTokenIs(lexer.TokenIF) {
+		ifNotExists = true
+		p.nextToken()
+		if !p.curTokenIs(lexer.TokenNOT) {
+			return nil, p.curError("expected NOT after IF")
+		}
+		p.nextToken()
+		if !p.curTokenIs(lexer.TokenEXISTS) {
+			return nil, p.curError("expected EXISTS after IF NOT")
+		}
+		p.nextToken()
+	}
+
 	// Parse column definition
 	col, err := p.parseColumnDef()
 	if err != nil {
 		return nil, err
 	}
 
-	stmt.Action = &AddColumnAction{Column: col}
+	stmt.Action = &AddColumnAction{Column: col, IfNotExists: ifNotExists}
 	return stmt, nil
 }
 

+ 29 - 0
pkg/parser/parser_test.go

@@ -768,6 +768,35 @@ func TestParseMultiple(t *testing.T) {
 	}
 }
 
+func TestParseRejectsTrailingReturning(t *testing.T) {
+	l := lexer.New("INSERT INTO users (id) VALUES (1) RETURNING id")
+	if _, err := New(l).Parse(); err == nil {
+		t.Fatal("expected RETURNING to be rejected before execution")
+	}
+}
+
+func TestParsePostgresCompatibilityClauses(t *testing.T) {
+	t.Run("alter add column if not exists", func(t *testing.T) {
+		stmt := parse(t, "ALTER TABLE users ADD COLUMN IF NOT EXISTS revision INTEGER DEFAULT 0")
+		action := stmt.(*AlterTableStmt).Action.(*AddColumnAction)
+		if !action.IfNotExists || action.Column.Name != "revision" {
+			t.Fatalf("unexpected action: %#v", action)
+		}
+	})
+
+	t.Run("on conflict do update", func(t *testing.T) {
+		stmt := parse(t, "INSERT INTO users (id, count) VALUES (1, 1) ON CONFLICT (id) DO UPDATE SET count = users.count + 1")
+		insert := stmt.(*InsertStmt)
+		if len(insert.ConflictTarget) != 1 || insert.ConflictTarget[0] != "id" || len(insert.ConflictUpdate) != 1 {
+			t.Fatalf("unexpected conflict clause: %#v", insert)
+		}
+	})
+
+	t.Run("jsonb cast", func(t *testing.T) {
+		parse(t, "SELECT CAST('{}' AS JSONB)")
+	})
+}
+
 // Phase 4: PRAGMA and EXPLAIN tests
 
 func TestParsePragmaTableInfo(t *testing.T) {

+ 650 - 36
pkg/pgserver/connection.go

@@ -8,6 +8,9 @@ import (
 	"io"
 	"log"
 	"net"
+	"regexp"
+	"sort"
+	"strconv"
 	"strings"
 
 	"github.com/danfragoso/pizzasql-next/pkg/executor"
@@ -18,34 +21,56 @@ import (
 
 // Connection represents a client connection
 type Connection struct {
-	conn      net.Conn
-	reader    *bufio.Reader
-	writer    *bufio.Writer
-	executor  *executor.Executor
-	schema    *storage.SchemaManager
-	dbManager *storage.DatabaseManager
-	database  string
-	params    map[string]string
-	txStatus  byte
-	quiet     bool // Disable query logging
+	conn           net.Conn
+	reader         *bufio.Reader
+	writer         *bufio.Writer
+	executor       *executor.Executor
+	schema         *storage.SchemaManager
+	dbManager      *storage.DatabaseManager
+	database       string
+	params         map[string]string
+	txStatus       byte
+	quiet          bool // Disable query logging
+	statements     map[string]*preparedStatement
+	portals        map[string]*portal
+	extendedFailed bool
+}
+
+type preparedStatement struct {
+	query     string
+	paramOIDs []int32
+}
+
+type portal struct {
+	query    string
+	executed bool
 }
 
 // NewConnection creates a new connection handler
 func NewConnection(conn net.Conn, dbManager *storage.DatabaseManager, quiet bool) *Connection {
 	return &Connection{
-		conn:      conn,
-		reader:    bufio.NewReader(conn),
-		writer:    bufio.NewWriter(conn),
-		dbManager: dbManager,
-		params:    make(map[string]string),
-		txStatus:  TxStatusIdle,
-		quiet:     quiet,
+		conn:       conn,
+		reader:     bufio.NewReader(conn),
+		writer:     bufio.NewWriter(conn),
+		dbManager:  dbManager,
+		params:     make(map[string]string),
+		statements: make(map[string]*preparedStatement),
+		portals:    make(map[string]*portal),
+		txStatus:   TxStatusIdle,
+		quiet:      quiet,
 	}
 }
 
 // Handle processes the connection
 func (c *Connection) Handle() error {
-	defer c.conn.Close()
+	defer func() {
+		if c.executor != nil {
+			if err := c.executor.RollbackActive(); err != nil && !c.quiet {
+				log.Printf("failed to roll back disconnected transaction: %v", err)
+			}
+		}
+		c.conn.Close()
+	}()
 
 	// First, check for SSL request (sent before startup message)
 	// SSL request is 8 bytes: length(4) + code(4) where code = 80877103
@@ -187,9 +212,20 @@ func (c *Connection) handleMessage(msg *Message) error {
 		log.Printf("Client requested termination")
 		return io.EOF
 
-	case MsgParse, MsgBind, MsgDescribe, MsgExecute, MsgSync, MsgClose:
-		// Extended query protocol - not implemented yet
-		c.sendError("ERROR", ErrCodeFeatureNotSupported, "Extended query protocol not yet supported")
+	case MsgParse:
+		return c.handleParse(msg)
+	case MsgBind:
+		return c.handleBind(msg)
+	case MsgDescribe:
+		return c.handleDescribe(msg)
+	case MsgExecute:
+		return c.handleExecute(msg)
+	case MsgClose:
+		return c.handleClose(msg)
+	case MsgFlush:
+		return c.writer.Flush()
+	case MsgSync:
+		c.extendedFailed = false
 		return c.sendReadyForQuery()
 
 	default:
@@ -219,46 +255,622 @@ func (c *Connection) handleQuery(msg *Message) error {
 
 	// Handle special PostgreSQL system queries that drivers send
 	sqlUpper := strings.ToUpper(strings.TrimSpace(sql))
+	singleStatement := isSingleStatementSQL(sql)
 
 	// lib/pq and other drivers query these for connection validation
-	if strings.Contains(sqlUpper, "SELECT VERSION()") {
+	if singleStatement && strings.Contains(sqlUpper, "SELECT VERSION()") {
 		// Return a fake PostgreSQL version
 		return c.handleVersionQuery()
 	}
 
-	if strings.Contains(sqlUpper, "SELECT CURRENT_USER") {
+	if singleStatement && strings.Contains(sqlUpper, "SELECT CURRENT_USER") {
 		// Return the current user
 		return c.handleCurrentUserQuery()
 	}
 
-	if strings.Contains(sqlUpper, "SHOW") && (strings.Contains(sqlUpper, "SERVER_VERSION") ||
+	if singleStatement && strings.Contains(sqlUpper, "SHOW") && (strings.Contains(sqlUpper, "SERVER_VERSION") ||
 		strings.Contains(sqlUpper, "SERVER_ENCODING") ||
 		strings.Contains(sqlUpper, "CLIENT_ENCODING")) {
 		// Handle SHOW commands
 		return c.handleShowCommand(sqlUpper)
 	}
+	if result, handled, err := c.catalogResult(sql); handled {
+		if err != nil {
+			c.sendError("ERROR", ErrCodeInternalError, err.Error())
+			return c.sendReadyForQuery()
+		}
+		if err := c.sendTabularResult(result, "SELECT"); err != nil {
+			return err
+		}
+		return c.sendReadyForQuery()
+	}
 
-	// Execute query
+	// Parse the complete batch before executing its first statement. This keeps
+	// unsupported trailing clauses from turning into committed writes.
 	l := lexer.New(sql)
 	p := parser.New(l)
-	stmt, err := p.Parse()
+	stmts, err := p.ParseMultiple()
 	if err != nil {
 		c.sendError("ERROR", ErrCodeSyntaxError, fmt.Sprintf("Syntax error: %v", err))
 		return c.sendReadyForQuery()
 	}
 
+	for _, stmt := range stmts {
+		if c.txStatus == TxStatusFailed {
+			if _, rollback := stmt.(*parser.RollbackStmt); !rollback {
+				c.sendError("ERROR", ErrCodeTransactionAborted, "current transaction is aborted, commands ignored until end of transaction block")
+				return c.sendReadyForQuery()
+			}
+		}
+		result, err := c.executor.Execute(stmt)
+		if err != nil {
+			if c.txStatus == TxStatusInBlock {
+				c.txStatus = TxStatusFailed
+			}
+			c.sendError("ERROR", ErrCodeInternalError, fmt.Sprintf("Execution error: %v", err))
+			return c.sendReadyForQuery()
+		}
+		if err := c.sendResult(result, stmt); err != nil {
+			return err
+		}
+	}
+
+	return c.sendReadyForQuery()
+}
+
+func (c *Connection) handleParse(msg *Message) error {
+	if c.extendedFailed {
+		return nil
+	}
+	name, pos, err := readCString(msg.Data, 0)
+	if err != nil {
+		return c.failExtended(ErrCodeProtocolViolation, err)
+	}
+	query, pos, err := readCString(msg.Data, pos)
+	if err != nil || pos+2 > len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Parse message"))
+	}
+	count := int(binary.BigEndian.Uint16(msg.Data[pos : pos+2]))
+	pos += 2
+	if pos+count*4 != len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Parse parameter list"))
+	}
+	oids := make([]int32, count)
+	for i := range oids {
+		oids[i] = int32(binary.BigEndian.Uint32(msg.Data[pos : pos+4]))
+		pos += 4
+	}
+	c.statements[name] = &preparedStatement{query: query, paramOIDs: oids}
+	return c.writeMessage(MsgParseComplete, nil)
+}
+
+func (c *Connection) handleBind(msg *Message) error {
+	if c.extendedFailed {
+		return nil
+	}
+	portalName, pos, err := readCString(msg.Data, 0)
+	if err != nil {
+		return c.failExtended(ErrCodeProtocolViolation, err)
+	}
+	statementName, pos, err := readCString(msg.Data, pos)
+	if err != nil {
+		return c.failExtended(ErrCodeProtocolViolation, err)
+	}
+	statement, ok := c.statements[statementName]
+	if !ok {
+		return c.failExtended(ErrCodeInvalidParameter, fmt.Errorf("prepared statement %q does not exist", statementName))
+	}
+	formats, pos, err := readInt16List(msg.Data, pos)
+	if err != nil || pos+2 > len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Bind format list"))
+	}
+	paramCount := int(binary.BigEndian.Uint16(msg.Data[pos : pos+2]))
+	pos += 2
+	params := make([]boundParameter, paramCount)
+	for i := range params {
+		if pos+4 > len(msg.Data) {
+			return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Bind parameter"))
+		}
+		length := int32(binary.BigEndian.Uint32(msg.Data[pos : pos+4]))
+		pos += 4
+		params[i].oid = parameterOID(statement.paramOIDs, i)
+		params[i].format = parameterFormat(formats, i)
+		if length == -1 {
+			params[i].null = true
+			continue
+		}
+		if length < 0 || pos+int(length) > len(msg.Data) {
+			return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Bind parameter length"))
+		}
+		params[i].value = append([]byte(nil), msg.Data[pos:pos+int(length)]...)
+		pos += int(length)
+	}
+	_, pos, err = readInt16List(msg.Data, pos) // result formats; text output is currently used
+	if err != nil || pos != len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Bind result format list"))
+	}
+	query, err := bindQuery(statement.query, params)
+	if err != nil {
+		return c.failExtended(ErrCodeInvalidParameter, err)
+	}
+	c.portals[portalName] = &portal{query: query}
+	return c.writeMessage(MsgBindComplete, nil)
+}
+
+func (c *Connection) handleDescribe(msg *Message) error {
+	if c.extendedFailed {
+		return nil
+	}
+	if len(msg.Data) < 2 {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Describe message"))
+	}
+	name, pos, err := readCString(msg.Data, 1)
+	if err != nil || pos != len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Describe name"))
+	}
+	switch msg.Data[0] {
+	case 'S':
+		statement, ok := c.statements[name]
+		if !ok {
+			return c.failExtended(ErrCodeInvalidParameter, fmt.Errorf("prepared statement %q does not exist", name))
+		}
+		mb := NewMessageBuilder()
+		mb.WriteInt16(int16(len(statement.paramOIDs)))
+		for _, oid := range statement.paramOIDs {
+			mb.WriteInt32(oid)
+		}
+		if err := c.writeMessage(MsgParameterDescription, mb.Bytes()); err != nil {
+			return err
+		}
+	case 'P':
+		if _, ok := c.portals[name]; !ok {
+			return c.failExtended(ErrCodeInvalidParameter, fmt.Errorf("portal %q does not exist", name))
+		}
+	default:
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Describe target"))
+	}
+	// Execute sends the row description once the bound statement has been
+	// analyzed, avoiding side effects during Describe.
+	return c.writeMessage(MsgNoData, nil)
+}
+
+func (c *Connection) handleExecute(msg *Message) error {
+	if c.extendedFailed {
+		return nil
+	}
+	name, pos, err := readCString(msg.Data, 0)
+	if err != nil || pos+4 != len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Execute message"))
+	}
+	portal, ok := c.portals[name]
+	if !ok {
+		return c.failExtended(ErrCodeInvalidParameter, fmt.Errorf("portal %q does not exist", name))
+	}
+	if portal.executed {
+		return c.failExtended(ErrCodeFeatureNotSupported, fmt.Errorf("portal can only be executed once"))
+	}
+	portal.executed = true
+	if result, handled, catalogErr := c.catalogResult(portal.query); handled {
+		if catalogErr != nil {
+			return c.failExtended(ErrCodeInternalError, catalogErr)
+		}
+		return c.sendTabularResult(result, "SELECT")
+	}
+	l := lexer.New(portal.query)
+	stmt, err := parser.New(l).Parse()
+	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 {
-		c.sendError("ERROR", ErrCodeInternalError, fmt.Sprintf("Execution error: %v", err))
-		return c.sendReadyForQuery()
+		if c.txStatus == TxStatusInBlock {
+			c.txStatus = TxStatusFailed
+		}
+		return c.failExtended(ErrCodeInternalError, fmt.Errorf("execution error: %w", err))
+	}
+	return c.sendResult(result, stmt)
+}
+
+var catalogFilterPattern = regexp.MustCompile(`(?i)\b(table_name|tablename|table_schema|schemaname|constraint_name|indexname)\s*=\s*'((?:''|[^'])*)'`)
+
+func (c *Connection) catalogResult(sql string) (*executor.Result, bool, error) {
+	if !isSingleStatementSQL(sql) {
+		return nil, false, nil
+	}
+	upper := strings.ToUpper(sql)
+	var source string
+	for _, candidate := range []string{
+		"INFORMATION_SCHEMA.TABLES", "INFORMATION_SCHEMA.COLUMNS",
+		"INFORMATION_SCHEMA.TABLE_CONSTRAINTS", "INFORMATION_SCHEMA.KEY_COLUMN_USAGE",
+		"PG_TABLES", "PG_INDEXES",
+	} {
+		if strings.Contains(upper, "FROM "+candidate) {
+			source = candidate
+			break
+		}
+	}
+	if source == "" {
+		return nil, false, nil
+	}
+	if c.txStatus == TxStatusIdle {
+		c.schema.LockStatement()
+		defer c.schema.UnlockStatement()
 	}
 
-	// Send result based on statement type
-	if err := c.sendResult(result, stmt); err != nil {
+	tables, err := c.schema.ListTables()
+	if err != nil {
+		return nil, true, err
+	}
+	rows := make([]map[string]interface{}, 0)
+	switch source {
+	case "INFORMATION_SCHEMA.TABLES":
+		for _, table := range tables {
+			rows = append(rows, map[string]interface{}{
+				"table_catalog": c.database, "table_schema": "public", "table_name": table, "table_type": "BASE TABLE",
+			})
+		}
+	case "PG_TABLES":
+		for _, table := range tables {
+			indexes, _ := c.schema.ListTableIndexes(table)
+			rows = append(rows, map[string]interface{}{
+				"schemaname": "public", "tablename": table, "tableowner": c.params["user"], "tablespace": nil,
+				"hasindexes": len(indexes) > 0, "hasrules": false, "hastriggers": false, "rowsecurity": false,
+			})
+		}
+	case "INFORMATION_SCHEMA.COLUMNS":
+		for _, table := range tables {
+			schema, schemaErr := c.schema.GetSchema(table)
+			if schemaErr != nil {
+				continue
+			}
+			for i, column := range schema.Columns {
+				rows = append(rows, map[string]interface{}{
+					"table_catalog": c.database, "table_schema": "public", "table_name": table,
+					"column_name": column.Name, "ordinal_position": int64(i + 1), "column_default": column.Default,
+					"is_nullable": yesNo(column.Nullable), "data_type": strings.ToLower(column.Type),
+				})
+			}
+		}
+	case "INFORMATION_SCHEMA.TABLE_CONSTRAINTS", "INFORMATION_SCHEMA.KEY_COLUMN_USAGE":
+		for _, table := range tables {
+			schema, schemaErr := c.schema.GetSchema(table)
+			if schemaErr != nil || schema.PrimaryKey == "" || schema.PrimaryKey == "_rowid_" {
+				continue
+			}
+			row := map[string]interface{}{
+				"constraint_catalog": c.database, "constraint_schema": "public", "constraint_name": table + "_pkey",
+				"table_catalog": c.database, "table_schema": "public", "table_name": table,
+			}
+			if source == "INFORMATION_SCHEMA.TABLE_CONSTRAINTS" {
+				row["constraint_type"] = "PRIMARY KEY"
+				row["is_deferrable"] = "NO"
+				row["initially_deferred"] = "NO"
+			} else {
+				row["column_name"] = schema.PrimaryKey
+				row["ordinal_position"] = int64(1)
+			}
+			rows = append(rows, row)
+		}
+	case "PG_INDEXES":
+		for _, table := range tables {
+			indexes, _ := c.schema.ListTableIndexes(table)
+			for _, index := range indexes {
+				columns := make([]string, len(index.Columns))
+				for i, column := range index.Columns {
+					columns[i] = column.Name
+				}
+				unique := ""
+				if index.Unique {
+					unique = "UNIQUE "
+				}
+				rows = append(rows, map[string]interface{}{
+					"schemaname": "public", "tablename": table, "indexname": index.Name, "tablespace": nil,
+					"indexdef": fmt.Sprintf("CREATE %sINDEX %s ON %s (%s)", unique, index.Name, table, strings.Join(columns, ", ")),
+				})
+			}
+		}
+	}
+
+	for _, match := range catalogFilterPattern.FindAllStringSubmatch(sql, -1) {
+		column, expected := strings.ToLower(match[1]), strings.ReplaceAll(match[2], "''", "'")
+		filtered := rows[:0]
+		for _, row := range rows {
+			if value, ok := row[column]; ok && strings.EqualFold(fmt.Sprintf("%v", value), expected) {
+				filtered = append(filtered, row)
+			}
+		}
+		rows = filtered
+	}
+
+	columns := catalogProjection(sql, rows)
+	result := executor.NewResult("SELECT")
+	if len(columns) == 1 && columns[0] == "count(*)" {
+		result.AddColumnWithType("count", "INTEGER")
+		result.AddRow(int64(len(rows)))
+		return result, true, nil
+	}
+	for _, column := range columns {
+		columnType := "TEXT"
+		if column == "ordinal_position" {
+			columnType = "INTEGER"
+		} else if strings.HasPrefix(column, "has") || column == "rowsecurity" {
+			columnType = "BOOLEAN"
+		}
+		result.AddColumnWithType(column, columnType)
+	}
+	for _, row := range rows {
+		values := make([]interface{}, len(columns))
+		for i, column := range columns {
+			values[i] = row[column]
+		}
+		result.AddRow(values...)
+	}
+	return result, true, nil
+}
+
+func isSingleStatementSQL(sql string) bool {
+	trimmed := strings.TrimSpace(sql)
+	if strings.HasSuffix(trimmed, ";") {
+		trimmed = strings.TrimSpace(strings.TrimSuffix(trimmed, ";"))
+	}
+	inString, inIdent := false, false
+	for i := 0; i < len(trimmed); i++ {
+		switch trimmed[i] {
+		case '\'':
+			if !inIdent {
+				if inString && i+1 < len(trimmed) && trimmed[i+1] == '\'' {
+					i++
+					continue
+				}
+				inString = !inString
+			}
+		case '"':
+			if !inString {
+				inIdent = !inIdent
+			}
+		case ';':
+			if !inString && !inIdent {
+				return false
+			}
+		}
+	}
+	return true
+}
+
+func catalogProjection(sql string, rows []map[string]interface{}) []string {
+	upper := strings.ToUpper(sql)
+	selectPos, fromPos := strings.Index(upper, "SELECT"), strings.Index(upper, " FROM ")
+	if selectPos < 0 || fromPos < 0 || fromPos <= selectPos+6 {
+		return nil
+	}
+	projection := strings.TrimSpace(sql[selectPos+6 : fromPos])
+	projection = strings.TrimSpace(strings.TrimPrefix(strings.ToUpper(projection), "DISTINCT "))
+	if projection == "*" && len(rows) > 0 {
+		columns := make([]string, 0, len(rows[0]))
+		for column := range rows[0] {
+			columns = append(columns, column)
+		}
+		sort.Strings(columns)
+		return columns
+	}
+	parts := strings.Split(projection, ",")
+	columns := make([]string, 0, len(parts))
+	for _, part := range parts {
+		column := strings.TrimSpace(part)
+		if index := strings.Index(strings.ToUpper(column), " AS "); index >= 0 {
+			column = strings.TrimSpace(column[:index])
+		}
+		if index := strings.LastIndex(column, "."); index >= 0 {
+			column = column[index+1:]
+		}
+		columns = append(columns, strings.ToLower(strings.Trim(column, `"`)))
+	}
+	return columns
+}
+
+func yesNo(value bool) string {
+	if value {
+		return "YES"
+	}
+	return "NO"
+}
+
+func (c *Connection) sendTabularResult(result *executor.Result, tag string) error {
+	if err := c.sendRowDescription(result.Columns, result.ColumnTypes); err != nil {
 		return err
 	}
+	for _, row := range result.Rows {
+		if err := c.sendDataRow(row, result.Columns); err != nil {
+			return err
+		}
+	}
+	return c.sendCommandComplete(fmt.Sprintf("%s %d", tag, len(result.Rows)))
+}
 
-	return c.sendReadyForQuery()
+func (c *Connection) handleClose(msg *Message) error {
+	if c.extendedFailed {
+		return nil
+	}
+	if len(msg.Data) < 2 {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Close message"))
+	}
+	name, pos, err := readCString(msg.Data, 1)
+	if err != nil || pos != len(msg.Data) {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Close name"))
+	}
+	if msg.Data[0] == 'S' {
+		delete(c.statements, name)
+	} else if msg.Data[0] == 'P' {
+		delete(c.portals, name)
+	} else {
+		return c.failExtended(ErrCodeProtocolViolation, fmt.Errorf("invalid Close target"))
+	}
+	return c.writeMessage(MsgCloseComplete, nil)
+}
+
+func (c *Connection) failExtended(code string, err error) error {
+	c.extendedFailed = true
+	return c.sendError("ERROR", code, err.Error())
+}
+
+type boundParameter struct {
+	value  []byte
+	oid    int32
+	format int16
+	null   bool
+}
+
+func readCString(data []byte, pos int) (string, int, error) {
+	if pos < 0 || pos >= len(data) {
+		return "", pos, fmt.Errorf("missing null-terminated string")
+	}
+	end := bytes.IndexByte(data[pos:], 0)
+	if end < 0 {
+		return "", pos, fmt.Errorf("unterminated string")
+	}
+	return string(data[pos : pos+end]), pos + end + 1, nil
+}
+
+func readInt16List(data []byte, pos int) ([]int16, int, error) {
+	if pos+2 > len(data) {
+		return nil, pos, fmt.Errorf("missing list length")
+	}
+	count := int(binary.BigEndian.Uint16(data[pos : pos+2]))
+	pos += 2
+	if pos+count*2 > len(data) {
+		return nil, pos, fmt.Errorf("truncated list")
+	}
+	result := make([]int16, count)
+	for i := range result {
+		result[i] = int16(binary.BigEndian.Uint16(data[pos : pos+2]))
+		pos += 2
+	}
+	return result, pos, nil
+}
+
+func parameterOID(oids []int32, index int) int32 {
+	if index < len(oids) {
+		return oids[index]
+	}
+	return 0
+}
+
+func parameterFormat(formats []int16, index int) int16 {
+	if len(formats) == 1 {
+		return formats[0]
+	}
+	if index < len(formats) {
+		return formats[index]
+	}
+	return 0
+}
+
+func bindQuery(query string, params []boundParameter) (string, error) {
+	var result strings.Builder
+	inString, inIdent := false, false
+	for i := 0; i < len(query); {
+		ch := query[i]
+		if ch == '\'' && !inIdent {
+			result.WriteByte(ch)
+			if inString && i+1 < len(query) && query[i+1] == '\'' {
+				result.WriteByte(query[i+1])
+				i += 2
+				continue
+			}
+			inString = !inString
+			i++
+			continue
+		}
+		if ch == '"' && !inString {
+			inIdent = !inIdent
+			result.WriteByte(ch)
+			i++
+			continue
+		}
+		if ch == '$' && !inString && !inIdent && i+1 < len(query) && query[i+1] >= '0' && query[i+1] <= '9' {
+			end := i + 1
+			for end < len(query) && query[end] >= '0' && query[end] <= '9' {
+				end++
+			}
+			n, _ := strconv.Atoi(query[i+1 : end])
+			if n < 1 || n > len(params) {
+				return "", fmt.Errorf("parameter $%d was not provided", n)
+			}
+			literal, err := parameterLiteral(params[n-1])
+			if err != nil {
+				return "", fmt.Errorf("parameter $%d: %w", n, err)
+			}
+			result.WriteString(literal)
+			i = end
+			continue
+		}
+		result.WriteByte(ch)
+		i++
+	}
+	return result.String(), nil
+}
+
+func parameterLiteral(param boundParameter) (string, error) {
+	if param.null {
+		return "NULL", nil
+	}
+	if param.format == 1 {
+		switch param.oid {
+		case 16:
+			if len(param.value) != 1 {
+				return "", fmt.Errorf("invalid binary boolean")
+			}
+			if param.value[0] == 0 {
+				return "FALSE", nil
+			}
+			return "TRUE", nil
+		case 21:
+			if len(param.value) != 2 {
+				return "", fmt.Errorf("invalid binary int2")
+			}
+			return strconv.FormatInt(int64(int16(binary.BigEndian.Uint16(param.value))), 10), nil
+		case 23:
+			if len(param.value) != 4 {
+				return "", fmt.Errorf("invalid binary int4")
+			}
+			return strconv.FormatInt(int64(int32(binary.BigEndian.Uint32(param.value))), 10), nil
+		case 20:
+			if len(param.value) != 8 {
+				return "", fmt.Errorf("invalid binary int8")
+			}
+			return strconv.FormatInt(int64(binary.BigEndian.Uint64(param.value)), 10), nil
+		default:
+			return "", fmt.Errorf("binary format is unsupported for OID %d", param.oid)
+		}
+	}
+	value := string(param.value)
+	switch param.oid {
+	case 0:
+		if strings.EqualFold(value, "true") || strings.EqualFold(value, "false") {
+			return strings.ToUpper(value), nil
+		}
+		if _, err := strconv.ParseFloat(value, 64); err == nil && value != "" {
+			return value, nil
+		}
+		return "'" + strings.ReplaceAll(value, "'", "''") + "'", nil
+	case 16:
+		if strings.EqualFold(value, "true") || value == "1" || value == "t" {
+			return "TRUE", nil
+		}
+		return "FALSE", nil
+	case 20, 21, 23, 26, 700, 701, 1700:
+		if _, err := strconv.ParseFloat(value, 64); err != nil {
+			return "", fmt.Errorf("invalid numeric value")
+		}
+		return value, nil
+	default:
+		return "'" + strings.ReplaceAll(value, "'", "''") + "'", nil
+	}
 }
 
 // sendResult sends query results
@@ -298,6 +910,8 @@ func (c *Connection) getCommandTag(stmt parser.Statement, result *executor.Resul
 		return "CREATE INDEX"
 	case *parser.DropIndexStmt:
 		return "DROP INDEX"
+	case *parser.AlterTableStmt:
+		return "ALTER TABLE"
 	case *parser.InsertStmt:
 		return fmt.Sprintf("INSERT 0 %d", result.RowsAffected)
 	case *parser.UpdateStmt:
@@ -344,7 +958,7 @@ func (c *Connection) sendBackendKeyData(processID, secretKey int32) error {
 // sendReadyForQuery sends ready for query message
 func (c *Connection) sendReadyForQuery() error {
 	mb := NewMessageBuilder()
-	mb.WriteByte(c.txStatus)
+	mb.AppendByte(c.txStatus)
 	return c.writeMessage(MsgReadyForQuery, mb.Bytes())
 }
 
@@ -405,13 +1019,13 @@ func (c *Connection) sendCommandComplete(tag string) error {
 // sendError sends an error response
 func (c *Connection) sendError(severity, code, message string) error {
 	mb := NewMessageBuilder()
-	mb.WriteByte(ErrorFieldSeverity)
+	mb.AppendByte(ErrorFieldSeverity)
 	mb.WriteString(severity)
-	mb.WriteByte(ErrorFieldCode)
+	mb.AppendByte(ErrorFieldCode)
 	mb.WriteString(code)
-	mb.WriteByte(ErrorFieldMessage)
+	mb.AppendByte(ErrorFieldMessage)
 	mb.WriteString(message)
-	mb.WriteByte(0) // Terminator
+	mb.AppendByte(0) // Terminator
 
 	return c.writeMessage(MsgErrorResponse, mb.Bytes())
 }

+ 28 - 0
pkg/pgserver/connection_test.go

@@ -0,0 +1,28 @@
+package pgserver
+
+import "testing"
+
+func TestBindQuery(t *testing.T) {
+	query, err := bindQuery(
+		"SELECT $1, $2, $3, $4, '$5'",
+		[]boundParameter{
+			{value: []byte("O'Reilly")},
+			{value: []byte("42")},
+			{null: true},
+			{value: []byte("true")},
+		},
+	)
+	if err != nil {
+		t.Fatal(err)
+	}
+	want := "SELECT 'O''Reilly', 42, NULL, TRUE, '$5'"
+	if query != want {
+		t.Fatalf("bound query = %q, want %q", query, want)
+	}
+}
+
+func TestBindQueryRequiresEveryParameter(t *testing.T) {
+	if _, err := bindQuery("SELECT $2", []boundParameter{{value: []byte("one")}}); err == nil {
+		t.Fatal("expected missing parameter error")
+	}
+}

+ 3 - 2
pkg/pgserver/protocol.go

@@ -81,6 +81,7 @@ const (
 	ErrCodeConnectionFailure   = "08006"
 	ErrCodeProtocolViolation   = "08P01"
 	ErrCodeFeatureNotSupported = "0A000"
+	ErrCodeTransactionAborted  = "25P02"
 )
 
 // Message represents a PostgreSQL protocol message
@@ -214,8 +215,8 @@ func NewMessageBuilder() *MessageBuilder {
 	}
 }
 
-// WriteByte writes a single byte
-func (mb *MessageBuilder) WriteByte(b byte) {
+// AppendByte appends a single byte.
+func (mb *MessageBuilder) AppendByte(b byte) {
 	mb.data = append(mb.data, b)
 }
 

+ 32 - 2
pkg/storage/kv.go

@@ -2,6 +2,7 @@ package storage
 
 import (
 	"bufio"
+	"errors"
 	"fmt"
 	"net"
 	"strings"
@@ -308,6 +309,35 @@ func (p *KVPool) WithClient(fn func(*KVClient) error) error {
 	if err != nil {
 		return err
 	}
-	defer p.Put(client)
-	return fn(client)
+	if err := fn(client); err != nil {
+		if !isConnectionError(err) {
+			p.Put(client)
+			return err
+		}
+		// A timeout or short response can leave an acknowledgement buffered on
+		// this connection. Never let the next request consume that response.
+		client.Close()
+		p.mu.Lock()
+		if !p.closed {
+			select {
+			case p.pool <- nil:
+			default:
+			}
+		}
+		p.mu.Unlock()
+		return err
+	}
+	p.Put(client)
+	return nil
+}
+
+func isConnectionError(err error) bool {
+	var netErr net.Error
+	if errors.As(err, &netErr) {
+		return true
+	}
+	message := err.Error()
+	return strings.Contains(message, "write command failed:") ||
+		strings.Contains(message, "flush failed:") ||
+		strings.Contains(message, "read response failed:")
 }

+ 34 - 5
pkg/storage/schema.go

@@ -53,8 +53,22 @@ type SchemaManager struct {
 	rowIDInitialized map[string]bool
 	version          uint64
 	mu               sync.RWMutex
+	txMu             sync.RWMutex
 }
 
+// BeginTransaction prevents other connections from observing intermediate
+// changes until this connection commits or rolls back.
+func (m *SchemaManager) BeginTransaction() { m.txMu.Lock() }
+
+// EndTransaction releases the database transaction lock.
+func (m *SchemaManager) EndTransaction() { m.txMu.Unlock() }
+
+// LockStatement serializes a non-transactional statement with transactions.
+func (m *SchemaManager) LockStatement() { m.txMu.RLock() }
+
+// UnlockStatement releases a non-transactional statement lock.
+func (m *SchemaManager) UnlockStatement() { m.txMu.RUnlock() }
+
 // NewSchemaManager creates a new schema manager.
 func NewSchemaManager(pool *KVPool, database string) *SchemaManager {
 	return &SchemaManager{
@@ -118,7 +132,9 @@ func (m *SchemaManager) CreateTable(schema *Schema) error {
 		return fmt.Errorf("table already exists: %s", schema.Name)
 	}
 
-	// Set creation time
+	// Keep the cached schema private so callers cannot mutate a published
+	// catalog snapshot after this operation returns.
+	schema = cloneSchema(schema)
 	schema.CreatedAt = time.Now()
 
 	// Determine primary key if not set
@@ -229,7 +245,7 @@ func (m *SchemaManager) GetSchema(name string) (*Schema, error) {
 	m.mu.RLock()
 	if schema, ok := m.cache[strings.ToLower(name)]; ok {
 		m.mu.RUnlock()
-		return schema, nil
+		return cloneSchema(schema), nil
 	}
 	m.mu.RUnlock()
 
@@ -238,7 +254,7 @@ func (m *SchemaManager) GetSchema(name string) (*Schema, error) {
 
 	// Double-check after acquiring write lock
 	if schema, ok := m.cache[strings.ToLower(name)]; ok {
-		return schema, nil
+		return cloneSchema(schema), nil
 	}
 
 	key := m.schemaKey(name)
@@ -261,7 +277,16 @@ func (m *SchemaManager) GetSchema(name string) (*Schema, error) {
 	}
 
 	m.cache[strings.ToLower(name)] = &schema
-	return &schema, nil
+	return cloneSchema(&schema), nil
+}
+
+func cloneSchema(schema *Schema) *Schema {
+	if schema == nil {
+		return nil
+	}
+	cloned := *schema
+	cloned.Columns = append([]Column(nil), schema.Columns...)
+	return &cloned
 }
 
 // TableExists checks if a table exists.
@@ -767,7 +792,8 @@ func (m *SchemaManager) AddColumn(table string, column Column) error {
 		}
 	}
 
-	// Add column
+	// Publish a fresh snapshot instead of mutating readers' shared pointer.
+	schema = cloneSchema(schema)
 	schema.Columns = append(schema.Columns, column)
 
 	// Update schema
@@ -804,6 +830,7 @@ func (m *SchemaManager) DropColumn(table, columnName string) error {
 		return fmt.Errorf("column not found: %s", columnName)
 	}
 
+	schema = cloneSchema(schema)
 	schema.Columns = newColumns
 
 	// Update schema
@@ -827,6 +854,7 @@ func (m *SchemaManager) RenameTable(oldName, newName string) error {
 		return fmt.Errorf("table already exists: %s", newName)
 	}
 
+	schema = cloneSchema(schema)
 	// Update schema name
 	schema.Name = newName
 
@@ -896,6 +924,7 @@ func (m *SchemaManager) RenameColumn(table, oldName, newName string) error {
 		}
 	}
 
+	schema = cloneSchema(schema)
 	// Find and rename column
 	found := false
 	for i, col := range schema.Columns {