diff --git a/command.go b/command.go index 0bbb8fa..1be9b26 100644 --- a/command.go +++ b/command.go @@ -5,8 +5,8 @@ import ( "errors" "fmt" "io" + "strings" - "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 +191,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)) @@ -205,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 { @@ -228,38 +234,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 +262,51 @@ 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 isEmptyQuery(query string) bool { + return strings.Trim(query, " \t\r\n;") == "" +} + +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..dee412c 100644 --- a/docs/postgres_emulation.md +++ b/docs/postgres_emulation.md @@ -39,6 +39,27 @@ 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 +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. 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`. +`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 The `IExtendedQueryBackend` interface is fully wired. A `DefaultExtendedQueryBackend` delegates to `HandleSimpleQuery`, providing basic compatibility. For full fidelity, stackql implements: @@ -59,4 +80,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..eefea92 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) } @@ -268,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") @@ -294,37 +303,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 +406,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 +421,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..965c58f 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 @@ -15,6 +20,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 +34,7 @@ type SimpleSQLResultStream struct { type ChannelSQLResultStream struct { res chan ISQLResult nextResultCached ISQLResult + provider ColumnProvider } func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { @@ -35,10 +43,25 @@ func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream { } } -func NewChannelSQLResultStream() ISQLResultStream { +func NewChannelSQLResultStream(provider ColumnProvider) ISQLResultStream { return &ChannelSQLResultStream{ - res: make(chan ISQLResult, 1), + res: make(chan ISQLResult, 1), + provider: provider, + } +} + +func (srs *SimpleSQLResultStream) GetColumns() []ISQLColumn { + if srs.res == nil { + return nil + } + return srs.res.GetColumns() +} + +func (srs *ChannelSQLResultStream) GetColumns() []ISQLColumn { + 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 new file mode 100644 index 0000000..e23a8b1 --- /dev/null +++ b/pkg/sqldata/sqldata_test.go @@ -0,0 +1,114 @@ +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 _, 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() }() + 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(want, 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 provider 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) + } + }) + } +} + +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/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/schema_test.go b/schema_test.go new file mode 100644 index 0000000..6c78bd5 --- /dev/null +++ b/schema_test.go @@ -0,0 +1,295 @@ +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(sqldata.NewSQLResult(columns, 0, 0, nil)) + _ = stream.Close() + return stream + default: + stream := sqldata.NewChannelSQLResultStream(sqldata.NewSQLResult(columns, 0, 0, nil)) + 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" { + 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"} { + 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 _, 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"} { + name := kind + "/" + description + + "/" + 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, + } + if description == "describe no data" { + backend.columns = nil + } + 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 description != "execute" { + sendDescribePortal(t, client, "") + if description == "describe no data" { + expectMsg(t, client, types.ServerNoData) + } else { + expectSchema(t, client, format) + } + } + sendExecute(t, client, "", 0) + if description != "describe execute" { + 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 + 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("second Complete = %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) + } +} 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..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,40 +70,14 @@ 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 } - if writer.written == 0 && writer.columns != nil { - err := writer.Empty() - if err != nil { - return err - } - } - defer writer.close() if notices != "" { noticesComplete(writer.client, notices)