From 6a45b470bb7bd4542fb1f336cdeaf0d8e31ddf86 Mon Sep 17 00:00:00 2001 From: Jeffrey Aven Date: Tue, 6 Oct 2026 16:18:55 +1100 Subject: [PATCH 1/5] fix zero-row and null wire results Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- row.go | 13 ++++++++ row_test.go | 61 +++++++++++++++++++++++++++++++++++-- wire_test.go | 85 ++++++++++++++++++++++++++++++++++++++++++++++++++++ writer.go | 7 ----- 4 files changed, 157 insertions(+), 9 deletions(-) diff --git a/row.go b/row.go index 28eb0c1..2262a9b 100644 --- a/row.go +++ b/row.go @@ -135,6 +135,19 @@ func (column Column) Write(ctx context.Context, writer buffer.Writer, src interf return err } + if bb == nil { + if value, ok := src.(string); ok && value == "" { + bb = []byte{} + } + if value, ok := src.([]byte); ok && value != nil && len(value) == 0 { + bb = value + } + if bb == nil { + writer.AddInt32(-1) + return nil + } + } + writer.AddInt32(int32(len(bb))) writer.AddBytes(bb) diff --git a/row_test.go b/row_test.go index 1250866..bc2c09f 100644 --- a/row_test.go +++ b/row_test.go @@ -44,8 +44,10 @@ func readDataRowValues(t *testing.T, data []byte) [][]byte { values[i] = nil } else { val := make([]byte, length) - if _, err := r.Read(val); err != nil { - t.Fatalf("read value for col %d: %v", i, err) + if length > 0 { + if _, err := r.Read(val); err != nil { + t.Fatalf("read value for col %d: %v", i, err) + } } values[i] = val } @@ -209,6 +211,61 @@ func TestColumnWrite_TextBypass_NullHandling(t *testing.T) { } } +func TestColumnWrite_NullAndEmptyValues(t *testing.T) { + ctx := setTypeInfo(context.Background()) + + tests := []struct { + name string + format FormatCode + src interface{} + wantData []byte + wantIsNull bool + }{ + {name: "text nil interface", format: TextFormat, src: nil, wantIsNull: true}, + {name: "text nil byte slice", format: TextFormat, src: []byte(nil), wantIsNull: true}, + {name: "text empty string", format: TextFormat, src: "", wantData: []byte{}}, + {name: "text empty byte slice", format: TextFormat, src: []byte{}, wantData: []byte{}}, + {name: "binary nil interface", format: BinaryFormat, src: nil, wantIsNull: true}, + {name: "binary nil byte slice", format: BinaryFormat, src: []byte(nil), wantIsNull: true}, + {name: "binary empty string", format: BinaryFormat, src: "", wantData: []byte{}}, + {name: "binary empty byte slice", format: BinaryFormat, src: []byte{}, wantData: []byte{}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + col := Column{ + Name: "test", + Oid: oid.T_text, + Width: -1, + Format: tt.format, + } + + var buf bytes.Buffer + writer := buffer.NewWriter(&buf) + if err := (Columns{col}).Write(ctx, writer, []interface{}{tt.src}); err != nil { + t.Fatalf("Write error: %v", err) + } + + values := readDataRowValues(t, buf.Bytes()) + if len(values) != 1 { + t.Fatalf("got %d values, want 1", len(values)) + } + if tt.wantIsNull { + if values[0] != nil { + t.Errorf("got %q, want NULL", values[0]) + } + return + } + if values[0] == nil { + t.Fatal("got NULL, want a non-NULL value") + } + if !bytes.Equal(values[0], tt.wantData) { + t.Errorf("got %q, want %q", values[0], tt.wantData) + } + }) + } +} + func TestColumnWrite_NonString_UsesStandardEncoder(t *testing.T) { ctx := setTypeInfo(context.Background()) diff --git a/wire_test.go b/wire_test.go index a4bb745..8f988c8 100644 --- a/wire_test.go +++ b/wire_test.go @@ -11,6 +11,7 @@ import ( _ "github.com/lib/pq" "github.com/lib/pq/oid" "github.com/stackql/psql-wire/internal/mock" + "github.com/stackql/psql-wire/internal/types" "github.com/stackql/psql-wire/pkg/sqlbackend" "github.com/stackql/psql-wire/pkg/sqldata" ) @@ -101,6 +102,90 @@ func TestClientConnect(t *testing.T) { }) } +func TestZeroRowResultRetainsColumnMetadata(t *testing.T) { + t.Parallel() + + handler := func(ctx context.Context, query string, writer DataWriter) error { + if err := writer.Define(Columns{{ + Name: "dependency_name", + Oid: oid.T_text, + Width: -1, + Format: TextFormat, + }}); err != nil { + return err + } + return writer.Complete("", "SELECT 0") + } + + server, err := NewServer(SimpleQuery(handler)) + if err != nil { + t.Fatal(err) + } + address := TListenAndServe(t, server) + + t.Run("wire messages", func(t *testing.T) { + conn, err := net.Dial("tcp", address.String()) + if err != nil { + t.Fatal(err) + } + client := mock.NewClient(conn) + client.Handshake(t) + client.Authenticate(t) + client.ReadyForQuery(t) + + client.Start(types.ClientSimpleQuery) + client.AddString("SELECT dependency_name") + client.AddNullTerminate() + if err := client.End(); err != nil { + t.Fatal(err) + } + + for _, expected := range []types.ServerMessage{ + types.ServerRowDescription, + types.ServerCommandComplete, + types.ServerReady, + } { + got, _, err := client.ReadTypedMsg() + if err != nil { + t.Fatal(err) + } + if got != expected { + t.Fatalf("got message %q, want %q", got, expected) + } + } + client.Close(t) + }) + + t.Run("database sql metadata", func(t *testing.T) { + connstr := fmt.Sprintf("host=%s port=%d sslmode=disable", address.IP, address.Port) + conn, err := sql.Open("postgres", connstr) + if err != nil { + t.Fatal(err) + } + defer conn.Close() //nolint:errcheck + + rows, err := conn.Query("SELECT dependency_name") + if err != nil { + t.Fatal(err) + } + defer rows.Close() //nolint:errcheck + + columns, err := rows.Columns() + if err != nil { + t.Fatal(err) + } + if len(columns) != 1 || columns[0] != "dependency_name" { + t.Fatalf("got columns %v, want [dependency_name]", columns) + } + if rows.Next() { + t.Fatal("zero-row result unexpectedly returned a row") + } + if err := rows.Err(); err != nil { + t.Fatal(err) + } + }) +} + func TestServerWritingResult(t *testing.T) { t.Parallel() diff --git a/writer.go b/writer.go index a065dc0..16d0b19 100644 --- a/writer.go +++ b/writer.go @@ -106,13 +106,6 @@ func (writer *dataWriter) Complete(notices, description string) error { return ErrClosedWriter } - if writer.written == 0 && writer.columns != nil { - err := writer.Empty() - if err != nil { - return err - } - } - defer writer.close() if notices != "" { noticesComplete(writer.client, notices) From 7d866d30c50c24be6922fa386979db682dbb7881 Mon Sep 17 00:00:00 2001 From: Jeffrey Aven Date: Wed, 7 Oct 2026 10:01:49 +1100 Subject: [PATCH 2/5] narrow wire results to stream schema and zero rows Expose non-consuming stream schema, define it before iteration, and reuse described portal columns and result formats. Remove the previous NULL encoding changes and related tests. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- command.go | 96 ++++++++-------- docs/postgres_emulation.md | 17 ++- extended_query.go | 49 +++----- pkg/sqldata/sqldata.go | 23 +++- pkg/sqldata/sqldata_test.go | 72 ++++++++++++ prepared.go | 1 + row.go | 13 --- row_test.go | 61 +--------- schema_test.go | 218 ++++++++++++++++++++++++++++++++++++ 9 files changed, 395 insertions(+), 155 deletions(-) create mode 100644 pkg/sqldata/sqldata_test.go create mode 100644 schema_test.go diff --git a/command.go b/command.go index 0bbb8fa..85b8f2b 100644 --- a/command.go +++ b/command.go @@ -6,7 +6,6 @@ import ( "fmt" "io" - "github.com/lib/pq/oid" "github.com/stackql/psql-wire/codes" psqlerr "github.com/stackql/psql-wire/errors" "github.com/stackql/psql-wire/internal/buffer" @@ -191,7 +190,6 @@ func (srv *Server) handleCommand(ctx context.Context, conn SQLConnection, t type return nil } - func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) error { if srv.SimpleQuery == nil && srv.SQLBackendFactory == nil { ErrorCode(cn, NewErrUnimplementedMessageType(types.ClientSimpleQuery)) @@ -228,38 +226,18 @@ func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) erro ctx: ctx, client: cn, } - var headersWritten bool - for { - if rdr == nil { - dw.Complete("", "OK") - return readyForQuery(cn, types.ServerIdle) - } - res, err := rdr.Read() - if err != nil { - if errors.Is(err, io.EOF) { - notices := cn.GetDebugStr() - if res == nil { - dw.Complete(notices, "OK") - return readyForQuery(cn, types.ServerIdle) - } - if !headersWritten { - headersWritten = true - srv.writeSQLResultHeader(ctx, res, dw, nil) - } - srv.writeSQLResultRows(ctx, res, dw) - // TODO: add debug messages, configurably - dw.Complete(notices, "OK") - return readyForQuery(cn, types.ServerIdle) - } - ErrorCode(cn, err) - return readyForQuery(cn, types.ServerIdle) - } - if !headersWritten { - headersWritten = true - dw.Define(nil) + err = srv.writeSQLResultStream(rdr, dw, nil) + if err != nil { + if writeErr := ErrorCode(cn, err); writeErr != nil { + return writeErr } - srv.writeSQLResultRows(ctx, res, dw) + return readyForQuery(cn, types.ServerIdle) } + err = dw.Complete(cn.GetDebugStr(), "OK") + if err != nil { + return err + } + return readyForQuery(cn, types.ServerIdle) } } @@ -276,27 +254,47 @@ func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) erro return readyForQuery(cn, types.ServerIdle) } -func (srv *Server) writeSQLResultRows(ctx context.Context, res sqldata.ISQLResult, writer DataWriter) error { - for _, r := range res.GetRows() { - writer.Row(r.GetRowDataForPgWire()) +func (srv *Server) writeSQLResultStream( + stream sqldata.ISQLResultStream, + writer *dataWriter, + resultFormats []int16, +) error { + if stream == nil { + return nil + } + if columns := stream.GetColumns(); columns != nil && writer.columns == nil { + if err := writer.Define(sqlColumns(columns, resultFormats)); err != nil { + return err + } + } + for { + result, err := stream.Read() + if err != nil && !errors.Is(err, io.EOF) { + return err + } + if result != nil { + if writeErr := srv.writeSQLResult(result, writer, resultFormats); writeErr != nil { + return writeErr + } + } + if errors.Is(err, io.EOF) { + return nil + } } - return nil } -func (srv *Server) writeSQLResultHeader(ctx context.Context, res sqldata.ISQLResult, writer DataWriter, resultFormats []int16) error { - var colz Columns - for i, c := range res.GetColumns() { - colz = append(colz, - Column{ - Table: c.GetTableId(), - Name: c.GetName(), - Oid: oid.Oid(c.GetObjectID()), - Width: c.GetWidth(), - Format: resolveResultFormat(resultFormats, i), - }, - ) +func (srv *Server) writeSQLResult(result sqldata.ISQLResult, writer *dataWriter, resultFormats []int16) error { + if writer.columns == nil { + if err := writer.Define(sqlColumns(result.GetColumns(), resultFormats)); err != nil { + return err + } } - return writer.Define(colz) + for _, row := range result.GetRows() { + if err := writer.Row(row.GetRowDataForPgWire()); err != nil { + return err + } + } + return nil } func (srv *Server) handleConnClose(ctx context.Context) error { diff --git a/docs/postgres_emulation.md b/docs/postgres_emulation.md index 7c50238..0ced1e1 100644 --- a/docs/postgres_emulation.md +++ b/docs/postgres_emulation.md @@ -39,6 +39,22 @@ For `FormatCode=0` (text), string/`[]byte` values bypass `pgtype` encoding and w `resultFormats` from `Bind` are threaded through to `RowDescription` and data encoding, supporting per-column text/binary format selection per [protocol spec](https://www.postgresql.org/docs/16/protocol-message-formats.html). +### Stream schema and zero-row results + +`ISQLResultStream.GetColumns()` returns schema without reading, peeking, or waiting +for results. Custom stream implementations must add this method. Simple streams +return their result's columns (or nil for a nil result). Channel streams accept +schema at construction: `NewChannelSQLResultStream(columns)`. Existing zero-argument +callers remain valid; when their schema is nil, execution uses the first result's +columns. Supply schema at construction to preserve metadata even when the channel +closes without producing any results. + +Simple and extended execution define columns once before reading when schema is +available. Extended execution reuses a successful portal Describe's columns and +negotiated formats without emitting a duplicate `RowDescription`. Structured +zero-row results complete normally, not with `EmptyQueryResponse`; explicit +`DataWriter.Empty()` behavior is unchanged. Standard value encoding is unchanged. + ## What requires stackql-side implementation The `IExtendedQueryBackend` interface is fully wired. A `DefaultExtendedQueryBackend` delegates to `HandleSimpleQuery`, providing basic compatibility. For full fidelity, stackql implements: @@ -59,4 +75,3 @@ See [stackql core repository](https://github.com/stackql/stackql) for the backen - **Portal suspension** — `Execute` with `maxRows > 0` does not yet suspend and resume portals via `ServerPortalSuspended`. - **`FunctionCall` message** — deprecated in PostgreSQL, not implemented. - **Notification / `LISTEN`/`NOTIFY`** — async notification messages are not emitted. - diff --git a/extended_query.go b/extended_query.go index 9dd74c7..281bbf2 100644 --- a/extended_query.go +++ b/extended_query.go @@ -3,7 +3,6 @@ package wire import ( "context" "errors" - "io" "github.com/lib/pq/oid" "github.com/stackql/psql-wire/internal/buffer" @@ -239,8 +238,14 @@ func (srv *Server) handleDescribePortal(ctx context.Context, conn SQLConnection, } if columns != nil { - return writeRowDescriptionFromSQLColumns(ctx, conn, columns, portal.ResultFormats) + colz := sqlColumns(columns, portal.ResultFormats) + if err := colz.Define(ctx, conn); err != nil { + return err + } + portal.columns = colz + return nil } + portal.columns = nil return writeNoData(conn) } @@ -294,37 +299,15 @@ func (srv *Server) handleExecute(ctx context.Context, conn SQLConnection) error dw := &dataWriter{ ctx: ctx, client: conn, + columns: portal.columns, resultFormats: portal.ResultFormats, } - var headersWritten bool - for { - res, err := rdr.Read() - if err != nil { - if errors.Is(err, io.EOF) { - notices := conn.GetDebugStr() - if res == nil { - dw.Complete(notices, "OK") - return nil - } - if !headersWritten { - headersWritten = true - srv.writeSQLResultHeader(ctx, res, dw, portal.ResultFormats) - } - srv.writeSQLResultRows(ctx, res, dw) - dw.Complete(notices, "OK") - return nil - } - return extendedError(conn, err) - } - if !headersWritten { - headersWritten = true - // For extended query, we don't send RowDescription here if Describe already sent it. - // However, the dataWriter.Define will handle this correctly since columns may already be set. - dw.Define(nil) - } - srv.writeSQLResultRows(ctx, res, dw) + err = srv.writeSQLResultStream(rdr, dw, portal.ResultFormats) + if err != nil { + return extendedError(conn, err) } + return dw.Complete(conn.GetDebugStr(), "OK") } // handleClose handles the Close message ('C') of the extended query protocol. @@ -419,7 +402,11 @@ func writeParameterDescription(writer buffer.Writer, paramOIDs []uint32) error { } func writeRowDescriptionFromSQLColumns(ctx context.Context, writer buffer.Writer, columns []sqldata.ISQLColumn, resultFormats []int16) error { - var colz Columns + return sqlColumns(columns, resultFormats).Define(ctx, writer) +} + +func sqlColumns(columns []sqldata.ISQLColumn, resultFormats []int16) Columns { + colz := make(Columns, 0, len(columns)) for i, c := range columns { colz = append(colz, Column{ Table: c.GetTableId(), @@ -430,7 +417,7 @@ func writeRowDescriptionFromSQLColumns(ctx context.Context, writer buffer.Writer Format: resolveResultFormat(resultFormats, i), }) } - return colz.Define(ctx, writer) + return colz } // resolveResultFormat determines the format code for column i based on the diff --git a/pkg/sqldata/sqldata.go b/pkg/sqldata/sqldata.go index df67587..3f2e67d 100644 --- a/pkg/sqldata/sqldata.go +++ b/pkg/sqldata/sqldata.go @@ -15,6 +15,8 @@ type ISQLResult interface { } type ISQLResultStream interface { + // GetColumns returns schema without reading or waiting for a result. + GetColumns() []ISQLColumn Read() (ISQLResult, error) Write(ISQLResult) error Close() error @@ -27,6 +29,7 @@ type SimpleSQLResultStream struct { type ChannelSQLResultStream struct { res chan ISQLResult nextResultCached ISQLResult + columns []ISQLColumn } func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { @@ -35,12 +38,28 @@ func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { } } -func NewChannelSQLResultStream() ISQLResultStream { +func NewChannelSQLResultStream(columns ...[]ISQLColumn) ISQLResultStream { + var schema []ISQLColumn + if len(columns) > 0 { + schema = columns[0] + } return &ChannelSQLResultStream{ - res: make(chan ISQLResult, 1), + res: make(chan ISQLResult, 1), + columns: schema, } } +func (srs *SimpleSQLResultStream) GetColumns() []ISQLColumn { + if srs.res == nil { + return nil + } + return srs.res.GetColumns() +} + +func (srs *ChannelSQLResultStream) GetColumns() []ISQLColumn { + return srs.columns +} + func (srs *SimpleSQLResultStream) Read() (ISQLResult, error) { return srs.res, io.EOF } diff --git a/pkg/sqldata/sqldata_test.go b/pkg/sqldata/sqldata_test.go new file mode 100644 index 0000000..e23773e --- /dev/null +++ b/pkg/sqldata/sqldata_test.go @@ -0,0 +1,72 @@ +package sqldata + +import ( + "errors" + "io" + "reflect" + "testing" + "time" +) + +func TestSimpleSQLResultStreamGetColumns(t *testing.T) { + columns := []ISQLColumn{ + NewSQLColumn(NewSQLTable(0, ""), "id", 0, 23, 4, -1, "text"), + } + result := NewSQLResult(columns, 0, 0, nil) + stream := NewSimpleSQLResultStream(result) + for i := 0; i < 2; i++ { + if !reflect.DeepEqual(stream.GetColumns(), columns) { + t.Fatal("GetColumns did not return the result schema") + } + } + got, err := stream.Read() + if got != result || !errors.Is(err, io.EOF) { + t.Fatalf("Read after GetColumns = (%v, %v), want original result and EOF", got, err) + } + if NewSimpleSQLResultStream(nil).GetColumns() != nil { + t.Fatal("nil result should have nil schema") + } +} + +func TestChannelSQLResultStreamGetColumns(t *testing.T) { + columns := []ISQLColumn{ + NewSQLColumn(NewSQLTable(0, ""), "id", 0, 23, 4, -1, "text"), + } + for _, supplied := range []bool{false, true} { + t.Run(map[bool]string{false: "legacy", true: "supplied schema"}[supplied], func(t *testing.T) { + stream := NewChannelSQLResultStream() + var want []ISQLColumn + if supplied { + stream = NewChannelSQLResultStream(columns) + want = columns + } + done := make(chan []ISQLColumn, 1) + go func() { done <- stream.GetColumns() }() + select { + case got := <-done: + if !reflect.DeepEqual(got, want) { + t.Fatalf("GetColumns = %v, want %v", got, want) + } + case <-time.After(time.Second): + t.Fatal("GetColumns blocked on an empty channel") + } + + result := NewSQLResult(columns, 0, 0, nil) + if err := stream.Write(result); err != nil { + t.Fatal(err) + } + if err := stream.Close(); err != nil { + t.Fatal(err) + } + for i := 0; i < 2; i++ { + if !reflect.DeepEqual(stream.GetColumns(), want) { + t.Fatal("GetColumns changed the constructor schema") + } + } + got, err := stream.Read() + if got != result || !errors.Is(err, io.EOF) { + t.Fatalf("Read after GetColumns = (%v, %v), want original result and EOF", got, err) + } + }) + } +} diff --git a/prepared.go b/prepared.go index f3fcb75..08a82fa 100644 --- a/prepared.go +++ b/prepared.go @@ -14,6 +14,7 @@ type Portal struct { ParamFormats []int16 ParamValues [][]byte ResultFormats []int16 + columns Columns } // PreparedStatementCache stores prepared statements for a connection. diff --git a/row.go b/row.go index 2262a9b..28eb0c1 100644 --- a/row.go +++ b/row.go @@ -135,19 +135,6 @@ func (column Column) Write(ctx context.Context, writer buffer.Writer, src interf return err } - if bb == nil { - if value, ok := src.(string); ok && value == "" { - bb = []byte{} - } - if value, ok := src.([]byte); ok && value != nil && len(value) == 0 { - bb = value - } - if bb == nil { - writer.AddInt32(-1) - return nil - } - } - writer.AddInt32(int32(len(bb))) writer.AddBytes(bb) diff --git a/row_test.go b/row_test.go index bc2c09f..1250866 100644 --- a/row_test.go +++ b/row_test.go @@ -44,10 +44,8 @@ func readDataRowValues(t *testing.T, data []byte) [][]byte { values[i] = nil } else { val := make([]byte, length) - if length > 0 { - if _, err := r.Read(val); err != nil { - t.Fatalf("read value for col %d: %v", i, err) - } + if _, err := r.Read(val); err != nil { + t.Fatalf("read value for col %d: %v", i, err) } values[i] = val } @@ -211,61 +209,6 @@ func TestColumnWrite_TextBypass_NullHandling(t *testing.T) { } } -func TestColumnWrite_NullAndEmptyValues(t *testing.T) { - ctx := setTypeInfo(context.Background()) - - tests := []struct { - name string - format FormatCode - src interface{} - wantData []byte - wantIsNull bool - }{ - {name: "text nil interface", format: TextFormat, src: nil, wantIsNull: true}, - {name: "text nil byte slice", format: TextFormat, src: []byte(nil), wantIsNull: true}, - {name: "text empty string", format: TextFormat, src: "", wantData: []byte{}}, - {name: "text empty byte slice", format: TextFormat, src: []byte{}, wantData: []byte{}}, - {name: "binary nil interface", format: BinaryFormat, src: nil, wantIsNull: true}, - {name: "binary nil byte slice", format: BinaryFormat, src: []byte(nil), wantIsNull: true}, - {name: "binary empty string", format: BinaryFormat, src: "", wantData: []byte{}}, - {name: "binary empty byte slice", format: BinaryFormat, src: []byte{}, wantData: []byte{}}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - col := Column{ - Name: "test", - Oid: oid.T_text, - Width: -1, - Format: tt.format, - } - - var buf bytes.Buffer - writer := buffer.NewWriter(&buf) - if err := (Columns{col}).Write(ctx, writer, []interface{}{tt.src}); err != nil { - t.Fatalf("Write error: %v", err) - } - - values := readDataRowValues(t, buf.Bytes()) - if len(values) != 1 { - t.Fatalf("got %d values, want 1", len(values)) - } - if tt.wantIsNull { - if values[0] != nil { - t.Errorf("got %q, want NULL", values[0]) - } - return - } - if values[0] == nil { - t.Fatal("got NULL, want a non-NULL value") - } - if !bytes.Equal(values[0], tt.wantData) { - t.Errorf("got %q, want %q", values[0], tt.wantData) - } - }) - } -} - func TestColumnWrite_NonString_UsesStandardEncoder(t *testing.T) { ctx := setTypeInfo(context.Background()) diff --git a/schema_test.go b/schema_test.go new file mode 100644 index 0000000..b73a857 --- /dev/null +++ b/schema_test.go @@ -0,0 +1,218 @@ +package wire + +import ( + "bytes" + "context" + "encoding/binary" + "errors" + "io" + "testing" + + "github.com/stackql/psql-wire/internal/buffer" + "github.com/stackql/psql-wire/internal/mock" + "github.com/stackql/psql-wire/internal/types" + "github.com/stackql/psql-wire/pkg/sqlbackend" + "github.com/stackql/psql-wire/pkg/sqldata" +) + +func schemaTestColumns() []sqldata.ISQLColumn { + return []sqldata.ISQLColumn{ + sqldata.NewSQLColumn(sqldata.NewSQLTable(0, ""), "id", 0, 23, 4, -1, "text"), + } +} + +func schemaTestStream(kind string, columns []sqldata.ISQLColumn) sqldata.ISQLResultStream { + switch kind { + case "simple zero rows": + return sqldata.NewSimpleSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) + case "channel no results": + stream := sqldata.NewChannelSQLResultStream(columns) + _ = stream.Close() + return stream + default: + stream := sqldata.NewChannelSQLResultStream() + if kind == "channel rows" { + stream = sqldata.NewChannelSQLResultStream(columns) + } + go func() { + for _, value := range []int32{1, 2} { + _ = stream.Write(sqldata.NewSQLResult(columns, 0, 0, + []sqldata.ISQLRow{sqldata.NewSQLRow([]interface{}{value})})) + } + _ = stream.Close() + }() + return stream + } +} + +func expectSchema(t *testing.T, client *mock.Client, format FormatCode) { + t.Helper() + expectMsg(t, client, types.ServerRowDescription) + body := client.PeekMsg() + if len(body) != 23 || binary.BigEndian.Uint16(body[:2]) != 1 || + string(body[2:5]) != "id\x00" || binary.BigEndian.Uint32(body[11:15]) != 23 || + binary.BigEndian.Uint16(body[21:23]) != uint16(format) { + t.Fatalf("unexpected RowDescription: %x", body) + } +} + +func expectSchemaRows(t *testing.T, client *mock.Client, kind string, format FormatCode) { + t.Helper() + if kind != "channel rows" && kind != "legacy channel rows" { + return + } + for _, value := range []int32{1, 2} { + expectMsg(t, client, types.ServerDataRow) + data := []byte{byte('0' + value)} + if format == BinaryFormat { + data = make([]byte, 4) + binary.BigEndian.PutUint32(data, uint32(value)) + } + body := client.PeekMsg() + if len(body) != 6+len(data) || binary.BigEndian.Uint16(body[:2]) != 1 || + binary.BigEndian.Uint32(body[2:6]) != uint32(len(data)) || !bytes.Equal(body[6:], data) { + t.Fatalf("unexpected DataRow: %x", body) + } + } +} + +func TestSimpleQueryStreamSchema(t *testing.T) { + for _, kind := range []string{"simple zero rows", "channel no results", "channel rows", "legacy channel rows"} { + t.Run(kind, func(t *testing.T) { + callback := func(context.Context, string) (sqldata.ISQLResultStream, error) { + return schemaTestStream(kind, schemaTestColumns()), nil + } + server, err := NewServer(SQLBackendFactory(sqlbackend.NewSimpleSQLBackendFactory(callback))) + if err != nil { + t.Fatal(err) + } + client := connectAndHandshake(t, TListenAndServe(t, server)) + client.Start(types.ClientSimpleQuery) + client.AddString("SELECT id") + client.AddNullTerminate() + if err := client.End(); err != nil { + t.Fatal(err) + } + expectSchema(t, client, TextFormat) + expectSchemaRows(t, client, kind, TextFormat) + expectMsg(t, client, types.ServerCommandComplete) + expectReadyForQuery(t, client, types.ServerIdle) + client.Close(t) + }) + } +} + +type schemaTestBackend struct { + sqlbackend.ISQLBackend + sqlbackend.IExtendedQueryBackend + columns []sqldata.ISQLColumn +} + +func (backend *schemaTestBackend) HandleDescribePortal( + context.Context, string, string, string, []uint32, +) ([]sqldata.ISQLColumn, error) { + return backend.columns, nil +} + +type schemaTestBackendFactory struct { + backend sqlbackend.ISQLBackend +} + +func (factory *schemaTestBackendFactory) NewSQLBackend() (sqlbackend.ISQLBackend, error) { + return factory.backend, nil +} + +func TestExtendedQueryStreamSchema(t *testing.T) { + for _, described := range []bool{false, true} { + for _, format := range []FormatCode{TextFormat, BinaryFormat} { + for _, kind := range []string{"simple zero rows", "channel no results", "channel rows", "legacy channel rows"} { + name := kind + "/" + map[bool]string{false: "execute", true: "describe execute"}[described] + + "/" + map[FormatCode]string{TextFormat: "text", BinaryFormat: "binary"}[format] + t.Run(name, func(t *testing.T) { + columns := schemaTestColumns() + callback := func(context.Context, string) (sqldata.ISQLResultStream, error) { + return schemaTestStream(kind, columns), nil + } + simple := sqlbackend.NewSimpleSQLBackend(callback) + backend := &schemaTestBackend{ + ISQLBackend: simple, + IExtendedQueryBackend: sqlbackend.NewDefaultExtendedQueryBackend(simple), + columns: columns, + } + server, err := NewServer(SQLBackendFactory(&schemaTestBackendFactory{backend: backend})) + if err != nil { + t.Fatal(err) + } + client := connectAndHandshake(t, TListenAndServe(t, server)) + sendParse(t, client, "", "SELECT id", nil) + expectMsg(t, client, types.ServerParseComplete) + client.Start(types.ClientBind) + client.AddString("") + client.AddNullTerminate() + client.AddString("") + client.AddNullTerminate() + client.AddInt16(0) + client.AddInt16(0) + client.AddInt16(1) + client.AddInt16(int16(format)) + if err := client.End(); err != nil { + t.Fatal(err) + } + expectMsg(t, client, types.ServerBindComplete) + if described { + sendDescribePortal(t, client, "") + expectSchema(t, client, format) + } + sendExecute(t, client, "", 0) + if !described { + expectSchema(t, client, format) + } + expectSchemaRows(t, client, kind, format) + expectMsg(t, client, types.ServerCommandComplete) + sendSync(t, client) + expectReadyForQuery(t, client, types.ServerIdle) + client.Close(t) + }) + } + } + } +} + +func TestExplicitEmptyQueryResponse(t *testing.T) { + var output bytes.Buffer + writer := &dataWriter{ctx: context.Background(), client: buffer.NewWriter(&output), columns: Columns{}} + if err := writer.Empty(); err != nil { + t.Fatal(err) + } + if !bytes.Equal(output.Bytes(), []byte{'I', 0, 0, 0, 4}) { + t.Fatalf("unexpected empty query response: %x", output.Bytes()) + } + if err := writer.Complete("", "OK"); !errors.Is(err, ErrClosedWriter) { + t.Fatalf("Complete after Empty = %v, want ErrClosedWriter", err) + } +} + +type schemaOnlyStream struct { + columns []sqldata.ISQLColumn + output *bytes.Buffer + t *testing.T +} + +func (stream *schemaOnlyStream) GetColumns() []sqldata.ISQLColumn { return stream.columns } +func (stream *schemaOnlyStream) Write(sqldata.ISQLResult) error { return errors.New("not supported") } +func (stream *schemaOnlyStream) Close() error { return nil } +func (stream *schemaOnlyStream) Read() (sqldata.ISQLResult, error) { + if stream.output.Len() == 0 || stream.output.Bytes()[0] != byte(types.ServerRowDescription) { + stream.t.Fatal("stream was read before RowDescription was written") + } + return nil, io.EOF +} + +func TestStreamSchemaWrittenBeforeRead(t *testing.T) { + var output bytes.Buffer + stream := &schemaOnlyStream{columns: schemaTestColumns(), output: &output, t: t} + writer := &dataWriter{ctx: context.Background(), client: buffer.NewWriter(&output)} + if err := (&Server{}).writeSQLResultStream(stream, writer, nil); err != nil { + t.Fatal(err) + } +} From 20c3d93330b034fd47c38aa6fe760d2fcac28901 Mon Sep 17 00:00:00 2001 From: Jeffrey Aven Date: Wed, 7 Oct 2026 10:04:44 +1100 Subject: [PATCH 3/5] test execute schema after portal describe returns no data Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- schema_test.go | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/schema_test.go b/schema_test.go index b73a857..92bcc11 100644 --- a/schema_test.go +++ b/schema_test.go @@ -123,10 +123,10 @@ func (factory *schemaTestBackendFactory) NewSQLBackend() (sqlbackend.ISQLBackend } func TestExtendedQueryStreamSchema(t *testing.T) { - for _, described := range []bool{false, true} { + for _, description := range []string{"execute", "describe execute", "describe no data"} { for _, format := range []FormatCode{TextFormat, BinaryFormat} { for _, kind := range []string{"simple zero rows", "channel no results", "channel rows", "legacy channel rows"} { - name := kind + "/" + map[bool]string{false: "execute", true: "describe execute"}[described] + + name := kind + "/" + description + "/" + map[FormatCode]string{TextFormat: "text", BinaryFormat: "binary"}[format] t.Run(name, func(t *testing.T) { columns := schemaTestColumns() @@ -139,6 +139,9 @@ func TestExtendedQueryStreamSchema(t *testing.T) { IExtendedQueryBackend: sqlbackend.NewDefaultExtendedQueryBackend(simple), columns: columns, } + if description == "describe no data" { + backend.columns = nil + } server, err := NewServer(SQLBackendFactory(&schemaTestBackendFactory{backend: backend})) if err != nil { t.Fatal(err) @@ -159,12 +162,16 @@ func TestExtendedQueryStreamSchema(t *testing.T) { t.Fatal(err) } expectMsg(t, client, types.ServerBindComplete) - if described { + if description != "execute" { sendDescribePortal(t, client, "") - expectSchema(t, client, format) + if description == "describe no data" { + expectMsg(t, client, types.ServerNoData) + } else { + expectSchema(t, client, format) + } } sendExecute(t, client, "", 0) - if !described { + if description != "describe execute" { expectSchema(t, client, format) } expectSchemaRows(t, client, kind, format) From 1b2ac7430bb2905b8e0f02c22f7909c7071a8420 Mon Sep 17 00:00:00 2001 From: Jeffrey Aven Date: Wed, 7 Oct 2026 11:41:05 +1100 Subject: [PATCH 4/5] accept column provider for channel result streams Delegate schema lookup on demand without consuming results, preserving zero-argument callers. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- docs/postgres_emulation.md | 7 ++++-- pkg/sqldata/sqldata.go | 24 ++++++++++++------ pkg/sqldata/sqldata_test.go | 49 ++++++++++++++++++++++++++++++++++++- schema_test.go | 4 +-- 4 files changed, 71 insertions(+), 13 deletions(-) diff --git a/docs/postgres_emulation.md b/docs/postgres_emulation.md index 0ced1e1..d043124 100644 --- a/docs/postgres_emulation.md +++ b/docs/postgres_emulation.md @@ -44,9 +44,12 @@ For `FormatCode=0` (text), string/`[]byte` values bypass `pgtype` encoding and w `ISQLResultStream.GetColumns()` returns schema without reading, peeking, or waiting for results. Custom stream implementations must add this method. Simple streams return their result's columns (or nil for a nil result). Channel streams accept -schema at construction: `NewChannelSQLResultStream(columns)`. Existing zero-argument +an optional `ColumnProvider` handle: `NewChannelSQLResultStream(result)`. +`ISQLResult` already satisfies this interface. The getter delegates on every +call; providers must supply schema without consuming results. A nil provider +returns nil schema. Existing zero-argument callers remain valid; when their schema is nil, execution uses the first result's -columns. Supply schema at construction to preserve metadata even when the channel +columns. Supply a schema provider to preserve metadata even when the channel closes without producing any results. Simple and extended execution define columns once before reading when schema is diff --git a/pkg/sqldata/sqldata.go b/pkg/sqldata/sqldata.go index 3f2e67d..97d151b 100644 --- a/pkg/sqldata/sqldata.go +++ b/pkg/sqldata/sqldata.go @@ -6,6 +6,11 @@ import ( "io" ) +// ColumnProvider supplies result schema on demand without consuming results. +type ColumnProvider interface { + GetColumns() []ISQLColumn +} + type ISQLResult interface { GetColumns() []ISQLColumn GetRowsAffected() uint64 @@ -29,7 +34,7 @@ type SimpleSQLResultStream struct { type ChannelSQLResultStream struct { res chan ISQLResult nextResultCached ISQLResult - columns []ISQLColumn + provider ColumnProvider } func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { @@ -38,14 +43,14 @@ func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { } } -func NewChannelSQLResultStream(columns ...[]ISQLColumn) ISQLResultStream { - var schema []ISQLColumn - if len(columns) > 0 { - schema = columns[0] +func NewChannelSQLResultStream(providers ...ColumnProvider) ISQLResultStream { + var provider ColumnProvider + if len(providers) > 0 { + provider = providers[0] } return &ChannelSQLResultStream{ - res: make(chan ISQLResult, 1), - columns: schema, + res: make(chan ISQLResult, 1), + provider: provider, } } @@ -57,7 +62,10 @@ func (srs *SimpleSQLResultStream) GetColumns() []ISQLColumn { } func (srs *ChannelSQLResultStream) GetColumns() []ISQLColumn { - return srs.columns + if srs.provider == nil { + return nil + } + return srs.provider.GetColumns() } func (srs *SimpleSQLResultStream) Read() (ISQLResult, error) { diff --git a/pkg/sqldata/sqldata_test.go b/pkg/sqldata/sqldata_test.go index e23773e..62b2bbd 100644 --- a/pkg/sqldata/sqldata_test.go +++ b/pkg/sqldata/sqldata_test.go @@ -37,9 +37,10 @@ func TestChannelSQLResultStreamGetColumns(t *testing.T) { stream := NewChannelSQLResultStream() var want []ISQLColumn if supplied { - stream = NewChannelSQLResultStream(columns) + stream = NewChannelSQLResultStream(NewSQLResult(columns, 0, 0, nil)) want = columns } + done := make(chan []ISQLColumn, 1) go func() { done <- stream.GetColumns() }() select { @@ -70,3 +71,49 @@ func TestChannelSQLResultStreamGetColumns(t *testing.T) { }) } } + +type countingColumnProvider struct { + columns []ISQLColumn + calls int +} + +func (provider *countingColumnProvider) GetColumns() []ISQLColumn { + provider.calls++ + return provider.columns +} + +func TestChannelSQLResultStreamDelegatesGetColumns(t *testing.T) { + provider := &countingColumnProvider{} + stream := NewChannelSQLResultStream(provider) + if provider.calls != 0 { + t.Fatal("constructor called the column provider") + } + if stream.GetColumns() != nil || provider.calls != 1 { + t.Fatal("GetColumns did not delegate to the provider") + } + provider.columns = []ISQLColumn{ + NewSQLColumn(NewSQLTable(0, ""), "id", 0, 23, 4, -1, "text"), + } + result := NewSQLResult(provider.columns, 0, 0, nil) + if err := stream.Write(result); err != nil { + t.Fatal(err) + } + if err := stream.Close(); err != nil { + t.Fatal(err) + } + for i := 0; i < 2; i++ { + if !reflect.DeepEqual(stream.GetColumns(), provider.columns) { + t.Fatal("GetColumns returned a stale schema") + } + } + if provider.calls != 3 { + t.Fatalf("provider called %d times, want 3", provider.calls) + } + got, err := stream.Read() + if got != result || !errors.Is(err, io.EOF) { + t.Fatalf("Read after GetColumns = (%v, %v), want original result and EOF", got, err) + } + if NewChannelSQLResultStream(nil).GetColumns() != nil { + t.Fatal("nil provider should have nil schema") + } +} diff --git a/schema_test.go b/schema_test.go index 92bcc11..b69aba8 100644 --- a/schema_test.go +++ b/schema_test.go @@ -26,13 +26,13 @@ func schemaTestStream(kind string, columns []sqldata.ISQLColumn) sqldata.ISQLRes case "simple zero rows": return sqldata.NewSimpleSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) case "channel no results": - stream := sqldata.NewChannelSQLResultStream(columns) + stream := sqldata.NewChannelSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) _ = stream.Close() return stream default: stream := sqldata.NewChannelSQLResultStream() if kind == "channel rows" { - stream = sqldata.NewChannelSQLResultStream(columns) + stream = sqldata.NewChannelSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) } go func() { for _, value := range []int32{1, 2} { From 83031391569ab2583e584d697a17aa7ad277fa10 Mon Sep 17 00:00:00 2001 From: Jeffrey Aven Date: Wed, 7 Oct 2026 11:58:16 +1100 Subject: [PATCH 5/5] require unary schema provider and remove writer empty API Keep empty-statement responses in protocol handlers, separate from schema-aware structured results. Migrate channel callers and cover zero-column and empty-statement protocol behavior. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- command.go | 12 +++++ docs/postgres_emulation.md | 16 ++++--- extended_query.go | 4 ++ pkg/sqldata/sqldata.go | 6 +-- pkg/sqldata/sqldata_test.go | 15 +++---- schema_test.go | 90 ++++++++++++++++++++++++++++++++----- writer.go | 28 ------------ 7 files changed, 111 insertions(+), 60 deletions(-) diff --git a/command.go b/command.go index 85b8f2b..1be9b26 100644 --- a/command.go +++ b/command.go @@ -5,6 +5,7 @@ import ( "errors" "fmt" "io" + "strings" "github.com/stackql/psql-wire/codes" psqlerr "github.com/stackql/psql-wire/errors" @@ -203,6 +204,13 @@ func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) erro srv.logger.Debug("incoming query", zap.String("query", query)) + if isEmptyQuery(query) { + if err = emptyQuery(cn); err != nil { + return err + } + return readyForQuery(cn, types.ServerIdle) + } + if cn.HasSQLBackend() { qArr, err := cn.SplitCompoundQuery(query) if err != nil { @@ -254,6 +262,10 @@ func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) erro return readyForQuery(cn, types.ServerIdle) } +func isEmptyQuery(query string) bool { + return strings.Trim(query, " \t\r\n;") == "" +} + func (srv *Server) writeSQLResultStream( stream sqldata.ISQLResultStream, writer *dataWriter, diff --git a/docs/postgres_emulation.md b/docs/postgres_emulation.md index d043124..dee412c 100644 --- a/docs/postgres_emulation.md +++ b/docs/postgres_emulation.md @@ -44,19 +44,21 @@ For `FormatCode=0` (text), string/`[]byte` values bypass `pgtype` encoding and w `ISQLResultStream.GetColumns()` returns schema without reading, peeking, or waiting for results. Custom stream implementations must add this method. Simple streams return their result's columns (or nil for a nil result). Channel streams accept -an optional `ColumnProvider` handle: `NewChannelSQLResultStream(result)`. +one required `ColumnProvider` handle: `NewChannelSQLResultStream(result)`. `ISQLResult` already satisfies this interface. The getter delegates on every call; providers must supply schema without consuming results. A nil provider -returns nil schema. Existing zero-argument -callers remain valid; when their schema is nil, execution uses the first result's -columns. Supply a schema provider to preserve metadata even when the channel -closes without producing any results. +returns nil schema. Zero-argument construction is no longer supported; a +zero-column result must supply a provider returning an empty slice. Supply a +schema provider to preserve metadata even when the channel closes without +producing any results. Simple and extended execution define columns once before reading when schema is available. Extended execution reuses a successful portal Describe's columns and negotiated formats without emitting a duplicate `RowDescription`. Structured -zero-row results complete normally, not with `EmptyQueryResponse`; explicit -`DataWriter.Empty()` behavior is unchanged. Standard value encoding is unchanged. +zero-row results complete normally, not with `EmptyQueryResponse`. +`DataWriter.Empty()` has been removed. Actual empty statements are handled by +the simple/extended protocol handlers independently of result writers, using +`EmptyQueryResponse` rather than `CommandComplete`. Standard value encoding is unchanged. ## What requires stackql-side implementation diff --git a/extended_query.go b/extended_query.go index 281bbf2..eefea92 100644 --- a/extended_query.go +++ b/extended_query.go @@ -273,6 +273,10 @@ func (srv *Server) handleExecute(ctx context.Context, conn SQLConnection) error return extendedError(conn, errors.New("portal does not exist: "+portalName)) } + if isEmptyQuery(portal.Statement.Query) { + return emptyQuery(conn) + } + extBackend := conn.ExtendedBackend() if extBackend == nil { return commandComplete(conn, "OK") diff --git a/pkg/sqldata/sqldata.go b/pkg/sqldata/sqldata.go index 97d151b..965c58f 100644 --- a/pkg/sqldata/sqldata.go +++ b/pkg/sqldata/sqldata.go @@ -43,11 +43,7 @@ func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { } } -func NewChannelSQLResultStream(providers ...ColumnProvider) ISQLResultStream { - var provider ColumnProvider - if len(providers) > 0 { - provider = providers[0] - } +func NewChannelSQLResultStream(provider ColumnProvider) ISQLResultStream { return &ChannelSQLResultStream{ res: make(chan ISQLResult, 1), provider: provider, diff --git a/pkg/sqldata/sqldata_test.go b/pkg/sqldata/sqldata_test.go index 62b2bbd..e23a8b1 100644 --- a/pkg/sqldata/sqldata_test.go +++ b/pkg/sqldata/sqldata_test.go @@ -32,14 +32,9 @@ func TestChannelSQLResultStreamGetColumns(t *testing.T) { columns := []ISQLColumn{ NewSQLColumn(NewSQLTable(0, ""), "id", 0, 23, 4, -1, "text"), } - for _, supplied := range []bool{false, true} { - t.Run(map[bool]string{false: "legacy", true: "supplied schema"}[supplied], func(t *testing.T) { - stream := NewChannelSQLResultStream() - var want []ISQLColumn - if supplied { - stream = NewChannelSQLResultStream(NewSQLResult(columns, 0, 0, nil)) - want = columns - } + for _, want := range [][]ISQLColumn{columns, {}} { + t.Run(map[bool]string{false: "zero columns", true: "columns"}[len(want) > 0], func(t *testing.T) { + stream := NewChannelSQLResultStream(NewSQLResult(want, 0, 0, nil)) done := make(chan []ISQLColumn, 1) go func() { done <- stream.GetColumns() }() @@ -52,7 +47,7 @@ func TestChannelSQLResultStreamGetColumns(t *testing.T) { t.Fatal("GetColumns blocked on an empty channel") } - result := NewSQLResult(columns, 0, 0, nil) + result := NewSQLResult(want, 0, 0, nil) if err := stream.Write(result); err != nil { t.Fatal(err) } @@ -61,7 +56,7 @@ func TestChannelSQLResultStreamGetColumns(t *testing.T) { } for i := 0; i < 2; i++ { if !reflect.DeepEqual(stream.GetColumns(), want) { - t.Fatal("GetColumns changed the constructor schema") + t.Fatal("GetColumns changed the provider schema") } } got, err := stream.Read() diff --git a/schema_test.go b/schema_test.go index b69aba8..6c78bd5 100644 --- a/schema_test.go +++ b/schema_test.go @@ -30,10 +30,7 @@ func schemaTestStream(kind string, columns []sqldata.ISQLColumn) sqldata.ISQLRes _ = stream.Close() return stream default: - stream := sqldata.NewChannelSQLResultStream() - if kind == "channel rows" { - stream = sqldata.NewChannelSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) - } + stream := sqldata.NewChannelSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) go func() { for _, value := range []int32{1, 2} { _ = stream.Write(sqldata.NewSQLResult(columns, 0, 0, @@ -58,7 +55,7 @@ func expectSchema(t *testing.T, client *mock.Client, format FormatCode) { func expectSchemaRows(t *testing.T, client *mock.Client, kind string, format FormatCode) { t.Helper() - if kind != "channel rows" && kind != "legacy channel rows" { + if kind != "channel rows" { return } for _, value := range []int32{1, 2} { @@ -77,7 +74,7 @@ func expectSchemaRows(t *testing.T, client *mock.Client, kind string, format For } func TestSimpleQueryStreamSchema(t *testing.T) { - for _, kind := range []string{"simple zero rows", "channel no results", "channel rows", "legacy channel rows"} { + for _, kind := range []string{"simple zero rows", "channel no results", "channel rows"} { t.Run(kind, func(t *testing.T) { callback := func(context.Context, string) (sqldata.ISQLResultStream, error) { return schemaTestStream(kind, schemaTestColumns()), nil @@ -125,7 +122,7 @@ func (factory *schemaTestBackendFactory) NewSQLBackend() (sqlbackend.ISQLBackend func TestExtendedQueryStreamSchema(t *testing.T) { for _, description := range []string{"execute", "describe execute", "describe no data"} { for _, format := range []FormatCode{TextFormat, BinaryFormat} { - for _, kind := range []string{"simple zero rows", "channel no results", "channel rows", "legacy channel rows"} { + for _, kind := range []string{"simple zero rows", "channel no results", "channel rows"} { name := kind + "/" + description + "/" + map[FormatCode]string{TextFormat: "text", BinaryFormat: "binary"}[format] t.Run(name, func(t *testing.T) { @@ -187,15 +184,88 @@ func TestExtendedQueryStreamSchema(t *testing.T) { func TestExplicitEmptyQueryResponse(t *testing.T) { var output bytes.Buffer - writer := &dataWriter{ctx: context.Background(), client: buffer.NewWriter(&output), columns: Columns{}} - if err := writer.Empty(); err != nil { + if err := emptyQuery(buffer.NewWriter(&output)); err != nil { t.Fatal(err) } if !bytes.Equal(output.Bytes(), []byte{'I', 0, 0, 0, 4}) { t.Fatalf("unexpected empty query response: %x", output.Bytes()) } +} + +func TestEmptyStatementsWireProtocol(t *testing.T) { + for _, query := range []string{"", " \t\r\n", ";", " ; ; \n"} { + for _, path := range []string{"simple callback", "simple backend", "extended"} { + t.Run(path+"/"+query, func(t *testing.T) { + callback := func(context.Context, string) (sqldata.ISQLResultStream, error) { + return nil, errors.New("empty statement must not execute the backend") + } + option := SQLBackendFactory(sqlbackend.NewSimpleSQLBackendFactory(callback)) + if path == "simple callback" { + option = SimpleQuery(func(context.Context, string, DataWriter) error { + return errors.New("empty statement must not execute the callback") + }) + } + server, err := NewServer(option) + if err != nil { + t.Fatal(err) + } + client := connectAndHandshake(t, TListenAndServe(t, server)) + if path == "extended" { + sendParse(t, client, "", query, nil) + expectMsg(t, client, types.ServerParseComplete) + sendBind(t, client, "", "") + expectMsg(t, client, types.ServerBindComplete) + sendExecute(t, client, "", 0) + } else { + client.Start(types.ClientSimpleQuery) + client.AddString(query) + client.AddNullTerminate() + if err := client.End(); err != nil { + t.Fatal(err) + } + } + expectMsg(t, client, types.ServerEmptyQuery) + if len(client.PeekMsg()) != 0 { + t.Fatal("EmptyQueryResponse should have no payload") + } + if path == "extended" { + sendSync(t, client) + } + expectReadyForQuery(t, client, types.ServerIdle) + client.Close(t) + }) + } + } +} + +func TestZeroColumnResultCompletesNormally(t *testing.T) { + result := sqldata.NewSQLResult([]sqldata.ISQLColumn{}, 0, 0, nil) + stream := sqldata.NewChannelSQLResultStream(result) + if err := stream.Close(); err != nil { + t.Fatal(err) + } + var output bytes.Buffer + writer := &dataWriter{ctx: context.Background(), client: buffer.NewWriter(&output)} + if err := (&Server{}).writeSQLResultStream(stream, writer, nil); err != nil { + t.Fatal(err) + } + if err := writer.Complete("", "OK"); err != nil { + t.Fatal(err) + } + client := mock.NewReader(&output) + got, _, err := client.ReadTypedMsg() + if err != nil || got != types.ServerRowDescription || !bytes.Equal(client.PeekMsg(), []byte{0, 0}) { + t.Fatalf("expected zero-column RowDescription, got %q, %v", got, err) + } + got, _, err = client.ReadTypedMsg() + if err != nil || got != types.ServerCommandComplete { + t.Fatalf("expected CommandComplete, got %q, %v", got, err) + } + if _, _, err := client.ReadTypedMsg(); !errors.Is(err, io.EOF) { + t.Fatalf("unexpected additional message: %v", err) + } if err := writer.Complete("", "OK"); !errors.Is(err, ErrClosedWriter) { - t.Fatalf("Complete after Empty = %v, want ErrClosedWriter", err) + t.Fatalf("second Complete = %v, want ErrClosedWriter", err) } } diff --git a/writer.go b/writer.go index 16d0b19..9cac2ce 100644 --- a/writer.go +++ b/writer.go @@ -23,10 +23,6 @@ type DataWriter interface { // values are encoded as NULL values. Row([]interface{}) error - // Empty announces to the client a empty response and that no data rows should - // be expected. - Empty() error - // Complete announces to the client that the command has been completed and // no further data should be expected. Complete(notices string, description string) error @@ -40,10 +36,6 @@ var ErrColumnsDefined = errors.New("columns have already been defined") // yet been defined. var ErrUndefinedColumns = errors.New("columns have not been defined") -// ErrDataWritten is thrown when an empty result is attempted to be send to the -// client while data has already been written. -var ErrDataWritten = errors.New("data has already been written") - // ErrClosedWriter is thrown when the data writer has been closed var ErrClosedWriter = errors.New("closed writer") @@ -53,7 +45,6 @@ type dataWriter struct { ctx context.Context client buffer.Writer closed bool - written uint64 resultFormats []int16 // from Bind message; nil means all text } @@ -79,28 +70,9 @@ func (writer *dataWriter) Row(values []interface{}) error { return ErrUndefinedColumns } - writer.written++ - return writer.columns.Write(writer.ctx, writer.client, values) } -func (writer *dataWriter) Empty() error { - if writer.closed { - return ErrClosedWriter - } - - if writer.columns == nil { - return ErrUndefinedColumns - } - - if writer.written != 0 { - return ErrDataWritten - } - - defer writer.close() - return emptyQuery(writer.client) -} - func (writer *dataWriter) Complete(notices, description string) error { if writer.closed { return ErrClosedWriter