turso
This commit is contained in:
11
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/driver.go
generated
vendored
Normal file
11
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/driver.go
generated
vendored
Normal file
@@ -0,0 +1,11 @@
|
||||
package http
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
|
||||
"github.com/tursodatabase/libsql-client-go/libsql/internal/http/hranaV2"
|
||||
)
|
||||
|
||||
func Connect(url, jwt, host string) driver.Conn {
|
||||
return hranaV2.Connect(url, jwt, host)
|
||||
}
|
||||
476
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/hranaV2/hranaV2.go
generated
vendored
Normal file
476
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/hranaV2/hranaV2.go
generated
vendored
Normal file
@@ -0,0 +1,476 @@
|
||||
package hranaV2
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
net_url "net/url"
|
||||
"runtime/debug"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/tursodatabase/libsql-client-go/libsql/internal/hrana"
|
||||
"github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared"
|
||||
)
|
||||
|
||||
var commitHash string
|
||||
|
||||
func init() {
|
||||
if info, ok := debug.ReadBuildInfo(); ok {
|
||||
for _, module := range info.Deps {
|
||||
if module.Path == "github.com/tursodatabase/libsql-client-go" {
|
||||
parts := strings.Split(module.Version, "-")
|
||||
if len(parts) == 3 {
|
||||
commitHash = parts[2][:6]
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
commitHash = "unknown"
|
||||
}
|
||||
|
||||
func Connect(url, jwt, host string) driver.Conn {
|
||||
return &hranaV2Conn{url, jwt, host, "", false, 0}
|
||||
}
|
||||
|
||||
type hranaV2Stmt struct {
|
||||
conn *hranaV2Conn
|
||||
numInput int
|
||||
sql string
|
||||
}
|
||||
|
||||
func (s *hranaV2Stmt) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *hranaV2Stmt) NumInput() int {
|
||||
return s.numInput
|
||||
}
|
||||
|
||||
func convertToNamed(args []driver.Value) []driver.NamedValue {
|
||||
if len(args) == 0 {
|
||||
return nil
|
||||
}
|
||||
var result []driver.NamedValue
|
||||
for idx := range args {
|
||||
result = append(result, driver.NamedValue{Ordinal: idx, Value: args[idx]})
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *hranaV2Stmt) Exec(args []driver.Value) (driver.Result, error) {
|
||||
return s.ExecContext(context.Background(), convertToNamed(args))
|
||||
}
|
||||
|
||||
func (s *hranaV2Stmt) Query(args []driver.Value) (driver.Rows, error) {
|
||||
return s.QueryContext(context.Background(), convertToNamed(args))
|
||||
}
|
||||
|
||||
func (s *hranaV2Stmt) ExecContext(ctx context.Context, args []driver.NamedValue) (driver.Result, error) {
|
||||
return s.conn.ExecContext(ctx, s.sql, args)
|
||||
}
|
||||
|
||||
func (s *hranaV2Stmt) QueryContext(ctx context.Context, args []driver.NamedValue) (driver.Rows, error) {
|
||||
return s.conn.QueryContext(ctx, s.sql, args)
|
||||
}
|
||||
|
||||
type hranaV2Conn struct {
|
||||
url string
|
||||
jwt string
|
||||
host string
|
||||
baton string
|
||||
streamClosed bool
|
||||
replicationIndex uint64
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) Ping() error {
|
||||
return h.PingContext(context.Background())
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) PingContext(ctx context.Context) error {
|
||||
_, err := h.executeStmt(ctx, "SELECT 1", nil, false)
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) Prepare(query string) (driver.Stmt, error) {
|
||||
return h.PrepareContext(context.Background(), query)
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) PrepareContext(ctx context.Context, query string) (driver.Stmt, error) {
|
||||
stmts, paramInfos, err := shared.ParseStatement(query)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(stmts) != 1 {
|
||||
return nil, fmt.Errorf("only one statement is supported got %d", len(stmts))
|
||||
}
|
||||
numInput := -1
|
||||
if len(paramInfos[0].NamedParameters) == 0 {
|
||||
numInput = paramInfos[0].PositionalParametersCount
|
||||
}
|
||||
return &hranaV2Stmt{h, numInput, query}, nil
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) Close() error {
|
||||
if h.baton != "" {
|
||||
go func(baton, url, jwt, host string) {
|
||||
msg := hrana.PipelineRequest{Baton: baton}
|
||||
msg.Add(hrana.CloseStream())
|
||||
_, _, _ = sendPipelineRequest(context.Background(), &msg, url, jwt, host)
|
||||
}(h.baton, h.url, h.jwt, h.host)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) Begin() (driver.Tx, error) {
|
||||
return h.BeginTx(context.Background(), driver.TxOptions{})
|
||||
}
|
||||
|
||||
type hranaV2Tx struct {
|
||||
conn *hranaV2Conn
|
||||
}
|
||||
|
||||
func (h hranaV2Tx) Commit() error {
|
||||
_, err := h.conn.ExecContext(context.Background(), "COMMIT", nil)
|
||||
return err
|
||||
}
|
||||
|
||||
func (h hranaV2Tx) Rollback() error {
|
||||
_, err := h.conn.ExecContext(context.Background(), "ROLLBACK", nil)
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) BeginTx(ctx context.Context, opts driver.TxOptions) (driver.Tx, error) {
|
||||
if opts.ReadOnly {
|
||||
return nil, fmt.Errorf("read only transactions are not supported")
|
||||
}
|
||||
if opts.Isolation != driver.IsolationLevel(sql.LevelDefault) {
|
||||
return nil, fmt.Errorf("isolation level %d is not supported", opts.Isolation)
|
||||
}
|
||||
_, err := h.ExecContext(ctx, "BEGIN", nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &hranaV2Tx{h}, nil
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) sendPipelineRequest(ctx context.Context, msg *hrana.PipelineRequest, streamClose bool) (*hrana.PipelineResponse, error) {
|
||||
if h.streamClosed {
|
||||
// If the stream is closed, we can't send any more requests using this connection.
|
||||
return nil, fmt.Errorf("stream is closed: %w", driver.ErrBadConn)
|
||||
}
|
||||
if h.baton != "" {
|
||||
msg.Baton = h.baton
|
||||
}
|
||||
if h.replicationIndex > 0 {
|
||||
addReplicationIndex(msg, h.replicationIndex)
|
||||
}
|
||||
result, streamClosed, err := sendPipelineRequest(ctx, msg, h.url, h.jwt, h.host)
|
||||
if streamClosed {
|
||||
h.streamClosed = true
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
h.baton = result.Baton
|
||||
if result.Baton == "" && !streamClose {
|
||||
// We need to remember that the stream is closed so we don't try to send any more requests using this connection.
|
||||
h.streamClosed = true
|
||||
}
|
||||
if result.BaseUrl != "" {
|
||||
h.url = result.BaseUrl
|
||||
}
|
||||
if idx := getReplicationIndex(&result); idx > h.replicationIndex {
|
||||
h.replicationIndex = idx
|
||||
}
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func addReplicationIndex(msg *hrana.PipelineRequest, replicationIndex uint64) {
|
||||
for i := range msg.Requests {
|
||||
if msg.Requests[i].Stmt != nil && msg.Requests[i].Stmt.ReplicationIndex == nil {
|
||||
msg.Requests[i].Stmt.ReplicationIndex = &replicationIndex
|
||||
} else if msg.Requests[i].Batch != nil && msg.Requests[i].Batch.ReplicationIndex == nil {
|
||||
msg.Requests[i].Batch.ReplicationIndex = &replicationIndex
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func getReplicationIndex(response *hrana.PipelineResponse) uint64 {
|
||||
if response == nil || len(response.Results) == 0 {
|
||||
return 0
|
||||
}
|
||||
var replicationIndex uint64
|
||||
for _, result := range response.Results {
|
||||
if result.Response == nil {
|
||||
continue
|
||||
}
|
||||
if result.Response.Type == "execute" {
|
||||
if res, err := result.Response.ExecuteResult(); err == nil && res.ReplicationIndex != nil {
|
||||
if *res.ReplicationIndex > replicationIndex {
|
||||
replicationIndex = *res.ReplicationIndex
|
||||
}
|
||||
}
|
||||
} else if result.Response.Type == "batch" {
|
||||
if res, err := result.Response.BatchResult(); err == nil && res.ReplicationIndex != nil {
|
||||
if *res.ReplicationIndex > replicationIndex {
|
||||
replicationIndex = *res.ReplicationIndex
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return replicationIndex
|
||||
}
|
||||
|
||||
func sendPipelineRequest(ctx context.Context, msg *hrana.PipelineRequest, url string, jwt string, host string) (result hrana.PipelineResponse, streamClosed bool, err error) {
|
||||
reqBody, err := json.Marshal(msg)
|
||||
if err != nil {
|
||||
return hrana.PipelineResponse{}, false, err
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(ctx, 60*time.Second)
|
||||
defer cancel()
|
||||
pipelineURL, err := net_url.JoinPath(url, "/v2/pipeline")
|
||||
if err != nil {
|
||||
return hrana.PipelineResponse{}, false, err
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", pipelineURL, bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return hrana.PipelineResponse{}, false, err
|
||||
}
|
||||
if len(jwt) > 0 {
|
||||
req.Header.Set("Authorization", "Bearer "+jwt)
|
||||
}
|
||||
req.Header.Set("x-libsql-client-version", "libsql-remote-go-"+commitHash)
|
||||
req.Host = host
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return hrana.PipelineResponse{}, false, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return hrana.PipelineResponse{}, false, err
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
// We need to remember that the stream is closed so we don't try to send any more requests using this connection.
|
||||
var serverError struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &serverError); err == nil {
|
||||
return hrana.PipelineResponse{}, true, fmt.Errorf("error code %d: %s", resp.StatusCode, serverError.Error)
|
||||
}
|
||||
var errResponse hrana.Error
|
||||
if err := json.Unmarshal(body, &errResponse); err == nil {
|
||||
if errResponse.Code != nil {
|
||||
if *errResponse.Code == "STREAM_EXPIRED" {
|
||||
return hrana.PipelineResponse{}, true, fmt.Errorf("error code %s: %s\n%w", *errResponse.Code, errResponse.Message, driver.ErrBadConn)
|
||||
} else {
|
||||
return hrana.PipelineResponse{}, true, fmt.Errorf("error code %s: %s", *errResponse.Code, errResponse.Message)
|
||||
}
|
||||
}
|
||||
return hrana.PipelineResponse{}, true, errors.New(errResponse.Message)
|
||||
}
|
||||
return hrana.PipelineResponse{}, true, fmt.Errorf("error code %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
if err = json.Unmarshal(body, &result); err != nil {
|
||||
return hrana.PipelineResponse{}, false, err
|
||||
}
|
||||
return result, false, nil
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) executeStmt(ctx context.Context, query string, args []driver.NamedValue, wantRows bool) (*hrana.PipelineResponse, error) {
|
||||
stmts, params, err := shared.ParseStatementAndArgs(query, args)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%w", query, err)
|
||||
}
|
||||
msg := &hrana.PipelineRequest{}
|
||||
if len(stmts) == 1 {
|
||||
executeStream, err := hrana.ExecuteStream(stmts[0], params[0], wantRows)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%w", query, err)
|
||||
}
|
||||
msg.Add(*executeStream)
|
||||
} else {
|
||||
batchStream, err := hrana.BatchStream(stmts, params, wantRows)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%w", query, err)
|
||||
}
|
||||
msg.Add(*batchStream)
|
||||
}
|
||||
|
||||
result, err := h.sendPipelineRequest(ctx, msg, false)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%w", query, err)
|
||||
}
|
||||
|
||||
if result.Results[0].Error != nil {
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%s", query, result.Results[0].Error.Message)
|
||||
}
|
||||
if result.Results[0].Response == nil {
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%s", query, "no response received")
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) ExecContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Result, error) {
|
||||
result, err := h.executeStmt(ctx, query, args, false)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch result.Results[0].Response.Type {
|
||||
case "execute":
|
||||
res, err := result.Results[0].Response.ExecuteResult()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return shared.NewResult(res.GetLastInsertRowId(), int64(res.AffectedRowCount)), nil
|
||||
case "batch":
|
||||
res, err := result.Results[0].Response.BatchResult()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
lastInsertRowId := int64(0)
|
||||
affectedRowCount := int64(0)
|
||||
for _, r := range res.StepResults {
|
||||
rowId := r.GetLastInsertRowId()
|
||||
if rowId > 0 {
|
||||
lastInsertRowId = rowId
|
||||
}
|
||||
affectedRowCount += int64(r.AffectedRowCount)
|
||||
}
|
||||
return shared.NewResult(lastInsertRowId, affectedRowCount), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%s", query, "unknown response type")
|
||||
}
|
||||
}
|
||||
|
||||
type StmtResultRowsProvider struct {
|
||||
r *hrana.StmtResult
|
||||
}
|
||||
|
||||
func (p *StmtResultRowsProvider) SetsCount() int {
|
||||
return 1
|
||||
}
|
||||
|
||||
func (p *StmtResultRowsProvider) RowsCount(setIdx int) int {
|
||||
if setIdx != 0 {
|
||||
return 0
|
||||
}
|
||||
return len(p.r.Rows)
|
||||
}
|
||||
|
||||
func (p *StmtResultRowsProvider) Columns(setIdx int) []string {
|
||||
if setIdx != 0 {
|
||||
return nil
|
||||
}
|
||||
res := make([]string, len(p.r.Cols))
|
||||
for i, c := range p.r.Cols {
|
||||
if c.Name != nil {
|
||||
res[i] = *c.Name
|
||||
}
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func (p *StmtResultRowsProvider) FieldValue(setIdx, rowIdx, colIdx int) driver.Value {
|
||||
if setIdx != 0 {
|
||||
return nil
|
||||
}
|
||||
return p.r.Rows[rowIdx][colIdx].ToValue(p.r.Cols[colIdx].Type)
|
||||
}
|
||||
|
||||
func (p *StmtResultRowsProvider) Error(setIdx int) string {
|
||||
return ""
|
||||
}
|
||||
|
||||
func (p *StmtResultRowsProvider) HasResult(setIdx int) bool {
|
||||
return setIdx == 0
|
||||
}
|
||||
|
||||
type BatchResultRowsProvider struct {
|
||||
r *hrana.BatchResult
|
||||
}
|
||||
|
||||
func (p *BatchResultRowsProvider) SetsCount() int {
|
||||
return len(p.r.StepResults)
|
||||
}
|
||||
|
||||
func (p *BatchResultRowsProvider) RowsCount(setIdx int) int {
|
||||
if setIdx >= len(p.r.StepResults) || p.r.StepResults[setIdx] == nil {
|
||||
return 0
|
||||
}
|
||||
return len(p.r.StepResults[setIdx].Rows)
|
||||
}
|
||||
|
||||
func (p *BatchResultRowsProvider) Columns(setIdx int) []string {
|
||||
if setIdx >= len(p.r.StepResults) || p.r.StepResults[setIdx] == nil {
|
||||
return nil
|
||||
}
|
||||
res := make([]string, len(p.r.StepResults[setIdx].Cols))
|
||||
for i, c := range p.r.StepResults[setIdx].Cols {
|
||||
if c.Name != nil {
|
||||
res[i] = *c.Name
|
||||
}
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func (p *BatchResultRowsProvider) FieldValue(setIdx, rowIdx, colIdx int) driver.Value {
|
||||
if setIdx >= len(p.r.StepResults) || p.r.StepResults[setIdx] == nil {
|
||||
return nil
|
||||
}
|
||||
return p.r.StepResults[setIdx].Rows[rowIdx][colIdx].ToValue(p.r.StepResults[setIdx].Cols[colIdx].Type)
|
||||
}
|
||||
|
||||
func (p *BatchResultRowsProvider) Error(setIdx int) string {
|
||||
if setIdx >= len(p.r.StepErrors) || p.r.StepErrors[setIdx] == nil {
|
||||
return ""
|
||||
}
|
||||
return p.r.StepErrors[setIdx].Message
|
||||
}
|
||||
|
||||
func (p *BatchResultRowsProvider) HasResult(setIdx int) bool {
|
||||
return setIdx < len(p.r.StepResults) && p.r.StepResults[setIdx] != nil
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) QueryContext(ctx context.Context, query string, args []driver.NamedValue) (driver.Rows, error) {
|
||||
result, err := h.executeStmt(ctx, query, args, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch result.Results[0].Response.Type {
|
||||
case "execute":
|
||||
res, err := result.Results[0].Response.ExecuteResult()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return shared.NewRows(&StmtResultRowsProvider{res}), nil
|
||||
case "batch":
|
||||
res, err := result.Results[0].Response.BatchResult()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return shared.NewRows(&BatchResultRowsProvider{res}), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("failed to execute SQL: %s\n%s", query, "unknown response type")
|
||||
}
|
||||
}
|
||||
|
||||
func (h *hranaV2Conn) ResetSession(ctx context.Context) error {
|
||||
if h.baton != "" {
|
||||
go func(baton, url, jwt, host string) {
|
||||
msg := hrana.PipelineRequest{Baton: baton}
|
||||
msg.Add(hrana.CloseStream())
|
||||
_, _, _ = sendPipelineRequest(context.Background(), &msg, url, jwt, host)
|
||||
}(h.baton, h.url, h.jwt, h.host)
|
||||
h.baton = ""
|
||||
}
|
||||
return nil
|
||||
}
|
||||
18
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared/result.go
generated
vendored
Normal file
18
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared/result.go
generated
vendored
Normal file
@@ -0,0 +1,18 @@
|
||||
package shared
|
||||
|
||||
type result struct {
|
||||
id int64
|
||||
changes int64
|
||||
}
|
||||
|
||||
func NewResult(id, changes int64) *result {
|
||||
return &result{id: id, changes: changes}
|
||||
}
|
||||
|
||||
func (r *result) LastInsertId() (int64, error) {
|
||||
return r.id, nil
|
||||
}
|
||||
|
||||
func (r *result) RowsAffected() (int64, error) {
|
||||
return r.changes, nil
|
||||
}
|
||||
69
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared/rows.go
generated
vendored
Normal file
69
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared/rows.go
generated
vendored
Normal file
@@ -0,0 +1,69 @@
|
||||
package shared
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"io"
|
||||
)
|
||||
|
||||
type rowsProvider interface {
|
||||
SetsCount() int
|
||||
RowsCount(setIdx int) int
|
||||
Columns(setIdx int) []string
|
||||
FieldValue(setIdx, rowIdx int, columnIdx int) driver.Value
|
||||
Error(setIdx int) string
|
||||
HasResult(setIdx int) bool
|
||||
}
|
||||
|
||||
func NewRows(result rowsProvider) driver.Rows {
|
||||
return &rows{result: result}
|
||||
}
|
||||
|
||||
type rows struct {
|
||||
result rowsProvider
|
||||
currentResultSetIndex int
|
||||
currentRowIdx int
|
||||
}
|
||||
|
||||
func (r *rows) Columns() []string {
|
||||
return r.result.Columns(r.currentResultSetIndex)
|
||||
}
|
||||
|
||||
func (r *rows) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *rows) Next(dest []driver.Value) error {
|
||||
if r.currentRowIdx == r.result.RowsCount(r.currentResultSetIndex) {
|
||||
return io.EOF
|
||||
}
|
||||
count := len(r.result.Columns(r.currentResultSetIndex))
|
||||
for idx := 0; idx < count; idx++ {
|
||||
dest[idx] = r.result.FieldValue(r.currentResultSetIndex, r.currentRowIdx, idx)
|
||||
}
|
||||
r.currentRowIdx++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *rows) HasNextResultSet() bool {
|
||||
return r.currentResultSetIndex < r.result.SetsCount()-1
|
||||
}
|
||||
|
||||
func (r *rows) NextResultSet() error {
|
||||
if !r.HasNextResultSet() {
|
||||
return io.EOF
|
||||
}
|
||||
|
||||
r.currentResultSetIndex++
|
||||
r.currentRowIdx = 0
|
||||
|
||||
errStr := r.result.Error(r.currentResultSetIndex)
|
||||
if errStr != "" {
|
||||
return fmt.Errorf("failed to execute statement\n%s", errStr)
|
||||
}
|
||||
if !r.result.HasResult(r.currentResultSetIndex) {
|
||||
return fmt.Errorf("no results for statement")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
240
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared/statement.go
generated
vendored
Normal file
240
vendor/github.com/tursodatabase/libsql-client-go/libsql/internal/http/shared/statement.go
generated
vendored
Normal file
@@ -0,0 +1,240 @@
|
||||
package shared
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
"sort"
|
||||
|
||||
"github.com/antlr/antlr4/runtime/Go/antlr/v4"
|
||||
"github.com/libsql/sqlite-antlr4-parser/sqliteparser"
|
||||
"github.com/libsql/sqlite-antlr4-parser/sqliteparserutils"
|
||||
)
|
||||
|
||||
type ParamsInfo struct {
|
||||
NamedParameters []string
|
||||
PositionalParametersCount int
|
||||
}
|
||||
|
||||
func ParseStatement(sql string) ([]string, []ParamsInfo, error) {
|
||||
stmts, _ := sqliteparserutils.SplitStatement(sql)
|
||||
|
||||
stmtsParams := make([]ParamsInfo, len(stmts))
|
||||
for idx, stmt := range stmts {
|
||||
nameParams, positionalParamsCount, err := extractParameters(stmt)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
stmtsParams[idx] = ParamsInfo{nameParams, positionalParamsCount}
|
||||
}
|
||||
return stmts, stmtsParams, nil
|
||||
}
|
||||
|
||||
func ParseStatementAndArgs(sql string, args []driver.NamedValue) ([]string, []Params, error) {
|
||||
parameters, err := ConvertArgs(args)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
stmts, _ := sqliteparserutils.SplitStatement(sql)
|
||||
|
||||
stmtsParams := make([]Params, len(stmts))
|
||||
totalParametersAlreadyUsed := 0
|
||||
for idx, stmt := range stmts {
|
||||
stmtParams, err := generateStatementParameters(stmt, parameters, totalParametersAlreadyUsed)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("fail to generate statement parameter. statement: %s. error: %v", stmt, err)
|
||||
}
|
||||
stmtsParams[idx] = stmtParams
|
||||
totalParametersAlreadyUsed += stmtParams.Len()
|
||||
}
|
||||
return stmts, stmtsParams, nil
|
||||
}
|
||||
|
||||
type paramsType int
|
||||
|
||||
const (
|
||||
namedParameters paramsType = iota
|
||||
positionalParameters
|
||||
)
|
||||
|
||||
type Params struct {
|
||||
positional []any
|
||||
named map[string]any
|
||||
}
|
||||
|
||||
func (p *Params) MarshalJSON() ([]byte, error) {
|
||||
if len(p.named) > 0 {
|
||||
return json.Marshal(p.named)
|
||||
}
|
||||
if len(p.positional) > 0 {
|
||||
return json.Marshal(p.positional)
|
||||
}
|
||||
return json.Marshal(make([]any, 0))
|
||||
}
|
||||
|
||||
func (p *Params) Named() map[string]any {
|
||||
return p.named
|
||||
}
|
||||
|
||||
func (p *Params) Positional() []any {
|
||||
return p.positional
|
||||
}
|
||||
|
||||
func (p *Params) Len() int {
|
||||
if p.named != nil {
|
||||
return len(p.named)
|
||||
}
|
||||
|
||||
return len(p.positional)
|
||||
}
|
||||
|
||||
func (p *Params) Type() paramsType {
|
||||
if p.named != nil {
|
||||
return namedParameters
|
||||
}
|
||||
|
||||
return positionalParameters
|
||||
}
|
||||
|
||||
func NewParams(t paramsType) Params {
|
||||
p := Params{}
|
||||
switch t {
|
||||
case namedParameters:
|
||||
p.named = make(map[string]any)
|
||||
case positionalParameters:
|
||||
p.positional = make([]any, 0)
|
||||
}
|
||||
|
||||
return p
|
||||
}
|
||||
|
||||
func getParamType(arg *driver.NamedValue) paramsType {
|
||||
if arg.Name == "" {
|
||||
return positionalParameters
|
||||
}
|
||||
return namedParameters
|
||||
}
|
||||
|
||||
func ConvertArgs(args []driver.NamedValue) (Params, error) {
|
||||
if len(args) == 0 {
|
||||
return NewParams(positionalParameters), nil
|
||||
}
|
||||
|
||||
var sortedArgs []*driver.NamedValue
|
||||
for idx := range args {
|
||||
sortedArgs = append(sortedArgs, &args[idx])
|
||||
}
|
||||
sort.Slice(sortedArgs, func(i, j int) bool {
|
||||
return sortedArgs[i].Ordinal < sortedArgs[j].Ordinal
|
||||
})
|
||||
|
||||
parametersType := getParamType(sortedArgs[0])
|
||||
parameters := NewParams(parametersType)
|
||||
for _, arg := range sortedArgs {
|
||||
if parametersType != getParamType(arg) {
|
||||
return Params{}, fmt.Errorf("driver does not accept positional and named parameters at the same time")
|
||||
}
|
||||
|
||||
switch parametersType {
|
||||
case positionalParameters:
|
||||
parameters.positional = append(parameters.positional, arg.Value)
|
||||
case namedParameters:
|
||||
parameters.named[arg.Name] = arg.Value
|
||||
}
|
||||
}
|
||||
return parameters, nil
|
||||
}
|
||||
|
||||
func generateStatementParameters(stmt string, queryParams Params, positionalParametersOffset int) (Params, error) {
|
||||
nameParams, positionalParamsCount, err := extractParameters(stmt)
|
||||
if err != nil {
|
||||
return Params{}, err
|
||||
}
|
||||
|
||||
stmtParams := NewParams(queryParams.Type())
|
||||
|
||||
switch queryParams.Type() {
|
||||
case positionalParameters:
|
||||
if positionalParametersOffset+positionalParamsCount > len(queryParams.positional) {
|
||||
return Params{}, fmt.Errorf("missing positional parameters")
|
||||
}
|
||||
stmtParams.positional = queryParams.positional[positionalParametersOffset : positionalParametersOffset+positionalParamsCount]
|
||||
case namedParameters:
|
||||
stmtParametersNeeded := make(map[string]bool)
|
||||
for _, stmtParametersName := range nameParams {
|
||||
stmtParametersNeeded[stmtParametersName] = true
|
||||
}
|
||||
for queryParamsName, queryParamsValue := range queryParams.named {
|
||||
if stmtParametersNeeded[queryParamsName] {
|
||||
stmtParams.named[queryParamsName] = queryParamsValue
|
||||
delete(stmtParametersNeeded, queryParamsName)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return stmtParams, nil
|
||||
}
|
||||
|
||||
func extractParameters(stmt string) (nameParams []string, positionalParamsCount int, err error) {
|
||||
statementStream := antlr.NewInputStream(stmt)
|
||||
sqliteparser.NewSQLiteLexer(statementStream)
|
||||
lexer := sqliteparser.NewSQLiteLexer(statementStream)
|
||||
|
||||
allTokens := lexer.GetAllTokens()
|
||||
|
||||
nameParamsSet := make(map[string]bool)
|
||||
|
||||
for _, token := range allTokens {
|
||||
tokenType := token.GetTokenType()
|
||||
if tokenType == sqliteparser.SQLiteLexerBIND_PARAMETER {
|
||||
parameter := token.GetText()
|
||||
|
||||
isPositionalParameter, err := isPositionalParameter(parameter)
|
||||
if err != nil {
|
||||
return []string{}, 0, err
|
||||
}
|
||||
|
||||
if isPositionalParameter {
|
||||
positionalParamsCount++
|
||||
} else {
|
||||
paramWithoutPrefix, err := removeParamPrefix(parameter)
|
||||
if err != nil {
|
||||
return []string{}, 0, err
|
||||
} else {
|
||||
nameParamsSet[paramWithoutPrefix] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
nameParams = make([]string, 0, len(nameParamsSet))
|
||||
for k := range nameParamsSet {
|
||||
nameParams = append(nameParams, k)
|
||||
}
|
||||
|
||||
return nameParams, positionalParamsCount, nil
|
||||
}
|
||||
|
||||
func isPositionalParameter(param string) (ok bool, err error) {
|
||||
re := regexp.MustCompile(`\?([0-9]*).*`)
|
||||
match := re.FindSubmatch([]byte(param))
|
||||
if match == nil {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
posS := string(match[1])
|
||||
if posS == "" {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
return true, fmt.Errorf("unsuppoted positional parameter. This driver does not accept positional parameters with indexes (like ?<number>)")
|
||||
}
|
||||
|
||||
func removeParamPrefix(paramName string) (string, error) {
|
||||
if paramName[0] == ':' || paramName[0] == '@' || paramName[0] == '$' {
|
||||
return paramName[1:], nil
|
||||
}
|
||||
return "", fmt.Errorf("all named parameters must start with ':', or '@' or '$'")
|
||||
}
|
||||
Reference in New Issue
Block a user