From 6d4698fc1de8a758ba4ac941b65e1c053e2656ab Mon Sep 17 00:00:00 2001 From: Pavel Sviderski Date: Mon, 30 Sep 2024 19:00:51 +1000 Subject: [PATCH] implement Query in corrosion API client --- internal/corrosion/client.go | 267 +++++++++++++++++++++++++++++++---- 1 file changed, 243 insertions(+), 24 deletions(-) diff --git a/internal/corrosion/client.go b/internal/corrosion/client.go index d5bd99cf..f9b12cf8 100644 --- a/internal/corrosion/client.go +++ b/internal/corrosion/client.go @@ -23,28 +23,12 @@ const ( http2Timeout = 30 * time.Second ) +// APIClient is a client for the Corrosion API. type APIClient struct { baseURL *url.URL client *http.Client } -type ExecResponse struct { - Results []ExecResult `json:"results"` - Time float64 `json:"time"` - Version *uint64 `json:"version"` -} - -type ExecResult struct { - RowsAffected *uint `json:"rows_affected"` - Time *float64 `json:"time"` - Error *string `json:"error"` -} - -type Statement struct { - Query string `json:"query"` - Params []any `json:"params"` -} - func NewAPIClient(addr netip.AddrPort) (*APIClient, error) { baseURL, err := url.Parse(fmt.Sprintf("http://%s", addr)) if err != nil { @@ -67,16 +51,34 @@ func NewAPIClient(addr netip.AddrPort) (*APIClient, error) { }, nil } -// Exec writes changes to the Corrosion database for propagation through the cluster. Corrosion does not sync schema -// changes made using this method. Use Corrosion's schema_files to create and update the cluster's database schema. -func (c *APIClient) Exec(ctx context.Context, query string, args ...any) (*ExecResult, error) { +type Statement struct { + Query string `json:"query"` + Params []any `json:"params"` +} + +type ExecResponse struct { + Results []ExecResult `json:"results"` + Time float64 `json:"time"` + Version *uint `json:"version"` +} + +type ExecResult struct { + RowsAffected uint `json:"rows_affected"` + Time float64 `json:"time"` + Error *string `json:"error"` +} + +// ExecContext writes changes to the Corrosion database for propagation through the cluster. The args are for any +// placeholder parameters in the query. Corrosion does not sync schema changes made using this method. Use Corrosion's +// schema_files to create and update the cluster's database schema. +func (c *APIClient) ExecContext(ctx context.Context, query string, args ...any) (*ExecResult, error) { statements := []Statement{ { Query: query, Params: args, }, } - resp, err := c.ExecMulti(ctx, statements...) + resp, err := c.ExecMultiContext(ctx, statements...) if err != nil { return nil, err } @@ -87,9 +89,9 @@ func (c *APIClient) Exec(ctx context.Context, query string, args ...any) (*ExecR return &resp.Results[0], nil } -// ExecMulti writes changes to the Corrosion database for propagation through the cluster. -// Unlike Exec, this method allows multiple statements to be executed in a single transaction. -func (c *APIClient) ExecMulti(ctx context.Context, statements ...Statement) (*ExecResponse, error) { +// ExecMultiContext writes changes to the Corrosion database for propagation through the cluster. +// Unlike ExecContext, this method allows multiple statements to be executed in a single transaction. +func (c *APIClient) ExecMultiContext(ctx context.Context, statements ...Statement) (*ExecResponse, error) { body, err := json.Marshal(statements) if err != nil { return nil, fmt.Errorf("marshal queries: %w", err) @@ -143,3 +145,220 @@ func (c *APIClient) ExecMulti(ctx context.Context, statements ...Statement) (*Ex } return nil, fmt.Errorf("unexpected status code %d: %s", resp.StatusCode, respBody) } + +type QueryEvent struct { + Columns []string `json:"columns"` + Row *RowEvent `json:"row"` + EOQ *EndOfQuery `json:"eoq"` + // TODO: implement event type Change to support subscriptions. + //Change []any `json:"change"` + // Error is a server-side error that occurred during query execution. It's considered fatal for the client + // as it cannot be recovered from server-side. + Error *string `json:"error"` +} + +type EndOfQuery struct { + Time float64 `json:"time"` + ChangeID *uint64 `json:"change_id"` +} + +type RowEvent struct { + RowID uint64 + Values []json.RawMessage +} + +func (re *RowEvent) UnmarshalJSON(data []byte) error { + var raw []json.RawMessage + if err := json.Unmarshal(data, &raw); err != nil { + return fmt.Errorf("invalid row event: %w", err) + } + if len(raw) != 2 { + return fmt.Errorf("invalid row event: expected an array of 2 elements") + } + if err := json.Unmarshal(raw[0], &re.RowID); err != nil { + return fmt.Errorf("invalid row event: %w", err) + } + if err := json.Unmarshal(raw[1], &re.Values); err != nil { + return fmt.Errorf("invalid row event: %w", err) + } + return nil +} + +func (re *RowEvent) MarshalJSON() ([]byte, error) { + return json.Marshal([]any{re.RowID, re.Values}) +} + +// QueryContext executes a query that returns rows, typically a SELECT. +// The args are for any placeholder parameters in the query. +func (c *APIClient) QueryContext(ctx context.Context, query string, args ...any) (*Rows, error) { + statement := Statement{ + Query: query, + Params: args, + } + body, err := json.Marshal(statement) + if err != nil { + return nil, fmt.Errorf("marshal query: %w", err) + } + + queriesURL := c.baseURL.JoinPath("/v1/queries").String() + req, err := http.NewRequestWithContext(ctx, "POST", queriesURL, bytes.NewReader(body)) + if err != nil { + return nil, fmt.Errorf("create request: %w", err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Accept", "application/json") + + resp, err := c.client.Do(req) + if err != nil { + return nil, fmt.Errorf("send request: %w", err) + } + + if resp.StatusCode != http.StatusOK { + respBody, err := io.ReadAll(resp.Body) + resp.Body.Close() + if err != nil { + return nil, fmt.Errorf("read response body: %w", err) + } + return nil, fmt.Errorf("unexpected status code %d: %s", resp.StatusCode, respBody) + } + + rows, err := newRows(ctx, resp.Body) + if err != nil { + resp.Body.Close() + return nil, fmt.Errorf("parse query response: %w", err) + } + return rows, nil +} + +// Rows is the result of a query. Its cursor starts before the first row of the result set. +// Use [Rows.Next] to advance from row to row. +type Rows struct { + ctx context.Context + body io.ReadCloser + decoder *json.Decoder + + columns []string + row RowEvent + time float64 + err error +} + +func newRows(ctx context.Context, body io.ReadCloser) (*Rows, error) { + select { + case <-ctx.Done(): + return nil, ctx.Err() + default: + } + + decoder := json.NewDecoder(body) + var e QueryEvent + if err := decoder.Decode(&e); err != nil { + return nil, fmt.Errorf("decode query event: %w", err) + } + if e.Columns == nil { + return nil, fmt.Errorf("expected columns event, got: %+v", e) + } + + return &Rows{ + ctx: ctx, + body: body, + decoder: decoder, + columns: e.Columns, + }, nil +} + +// Columns returns the column names. +func (rs *Rows) Columns() []string { + return rs.columns +} + +// Next prepares the next result row for reading with the [Rows.Scan] method. It returns true on success, or false +// if there is no next result row or an error happened while preparing it. [Rows.Err] should be consulted to distinguish +// between the two cases. +// +// Every call to [Rows.Scan], even the first one, must be preceded by a call to [Rows.Next]. +func (rs *Rows) Next() bool { + select { + case <-rs.ctx.Done(): + rs.err = rs.ctx.Err() + _ = rs.Close() + return false + default: + } + + var e QueryEvent + if err := rs.decoder.Decode(&e); err != nil { + rs.err = fmt.Errorf("decode query event: %w", err) + _ = rs.Close() + return false + } + // Server-side query error. + if e.Error != nil { + rs.err = fmt.Errorf("query error: %s", *e.Error) + _ = rs.Close() + return false + } + + if e.Row != nil { + if len(e.Row.Values) != len(rs.columns) { + rs.err = fmt.Errorf("expected %d column values, got %d", len(rs.columns), len(e.Row.Values)) + _ = rs.Close() + return false + } + rs.row = *e.Row + return true + } + if e.EOQ != nil { + rs.time = e.EOQ.Time + _ = rs.Close() + return false + } + + rs.err = fmt.Errorf("expected row or eof event, got: %+v", e) + _ = rs.Close() + return false +} + +// Err returns the error, if any, that was encountered during iteration. +// Err may be called after an explicit or implicit [Rows.Close]. +func (rs *Rows) Err() error { + return rs.err +} + +// Scan copies the columns in the current row into the values pointed at by dest. +// The number of values in dest must be the same as the number of columns in [Rows]. +// Scan converts JSON-encoded column values to the provided Go types using [json.Unmarshal]. +func (rs *Rows) Scan(dest ...any) error { + if rs.err != nil { + return rs.err + } + if len(dest) != len(rs.columns) { + return fmt.Errorf("expected %d values, got %d", len(rs.columns), len(dest)) + } + + for i, v := range rs.row.Values { + if err := json.Unmarshal(v, dest[i]); err != nil { + return fmt.Errorf("unmarshal column value #%d: %w", i, err) + } + } + return nil +} + +// Time returns the time taken to execute the query in seconds. It's only available after all rows have been consumed. +// It doesn't include the time to send the query, receive the response, or iterate over the rows. +func (rs *Rows) Time() (float64, error) { + if rs.time == 0 { + if rs.Err() != nil { + return 0, fmt.Errorf("time is not available: %w", rs.Err()) + } + return 0, errors.New("time is not available until all rows are consumed") + } + return rs.time, nil +} + +// Close closes the [Rows], preventing further enumeration. If [Rows.Next] is called and returns false, +// the [Rows] are closed automatically and it will suffice to check the result of [Rows.Err]. +// Close is idempotent and does not affect the result of [Rows.Err]. +func (rs *Rows) Close() error { + return rs.body.Close() +}