Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 59 additions & 49 deletions command.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@
"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"
Expand Down Expand Up @@ -191,10 +191,9 @@
return nil
}


func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) error {
if srv.SimpleQuery == nil && srv.SQLBackendFactory == nil {
ErrorCode(cn, NewErrUnimplementedMessageType(types.ClientSimpleQuery))

Check failure on line 196 in command.go

View workflow job for this annotation

GitHub Actions / lint

Error return value is not checked (errcheck)
return readyForQuery(cn, types.ServerIdle)
}

Expand All @@ -205,6 +204,13 @@

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 {
Expand All @@ -214,52 +220,32 @@
if q == "" {
if i == len(qArr)-1 {
// trailing semicolon, ignore
commandComplete(cn, "OK")

Check failure on line 223 in command.go

View workflow job for this annotation

GitHub Actions / lint

Error return value is not checked (errcheck)
return readyForQuery(cn, types.ServerIdle)
}
continue
}
rdr, err := cn.HandleSimpleQuery(ctx, q)
if err != nil {
ErrorCode(cn, err)

Check failure on line 230 in command.go

View workflow job for this annotation

GitHub Actions / lint

Error return value is not checked (errcheck)
return readyForQuery(cn, types.ServerIdle)
}
dw := &dataWriter{
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)
}
}

Expand All @@ -269,34 +255,58 @@
})

if err != nil {
ErrorCode(cn, err)

Check failure on line 258 in command.go

View workflow job for this annotation

GitHub Actions / lint

Error return value is not checked (errcheck)
return readyForQuery(cn, types.ServerIdle)
}

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 {
Expand Down
22 changes: 21 additions & 1 deletion docs/postgres_emulation.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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.

53 changes: 22 additions & 31 deletions extended_query.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@ package wire
import (
"context"
"errors"
"io"

"github.com/lib/pq/oid"
"github.com/stackql/psql-wire/internal/buffer"
Expand Down Expand Up @@ -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)
}

Expand Down Expand Up @@ -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")
Expand All @@ -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.
Expand Down Expand Up @@ -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(),
Expand All @@ -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
Expand Down
27 changes: 25 additions & 2 deletions pkg/sqldata/sqldata.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -27,6 +34,7 @@ type SimpleSQLResultStream struct {
type ChannelSQLResultStream struct {
res chan ISQLResult
nextResultCached ISQLResult
provider ColumnProvider
}

func NewSimpleSQLResultStream(res ISQLResult) ISQLResultStream {
Expand All @@ -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 {
Comment thread
general-kroll-4-life marked this conversation as resolved.
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) {
Expand Down
Loading
Loading