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
54 changes: 23 additions & 31 deletions database_role_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,29 +2,17 @@ package main

import (
"context"
"net"

"testing"
"time"

sppb "cloud.google.com/go/spanner/apiv1/spannerpb"
"github.com/alecthomas/kong"
"google.golang.org/grpc"
"github.com/apstndb/execspansql/internal/grpctest"
"github.com/apstndb/spanemuboost"
"google.golang.org/api/option"
)

type databaseRoleSpannerServer struct {
sppb.UnimplementedSpannerServer
createSession chan *sppb.CreateSessionRequest
}

func (s *databaseRoleSpannerServer) CreateSession(_ context.Context, req *sppb.CreateSessionRequest) (*sppb.Session, error) {
s.createSession <- req
return &sppb.Session{
Name: req.Database + "/sessions/test",
CreatorRole: req.GetSession().GetCreatorRole(),
Multiplexed: true,
}, nil
}

func TestDatabaseRoleFlag(t *testing.T) {
t.Parallel()

Expand All @@ -46,6 +34,13 @@ func TestDatabaseRoleFlag(t *testing.T) {
}

func TestNewClientSendsDatabaseRoleWhenCreatingSession(t *testing.T) {
env, err := spanemuboost.RunEmulatorWithClients(context.Background())
if err != nil {
t.Fatal(err)
}
defer env.Close() //nolint:errcheck
t.Setenv("SPANNER_EMULATOR_HOST", env.Emulator().URI())

for _, tt := range []struct {
name string
role string
Expand All @@ -54,31 +49,28 @@ func TestNewClientSendsDatabaseRoleWhenCreatingSession(t *testing.T) {
{name: "omitted"},
} {
t.Run(tt.name, func(t *testing.T) {
server := &databaseRoleSpannerServer{createSession: make(chan *sppb.CreateSessionRequest, 1)}
grpcServer := grpc.NewServer()
sppb.RegisterSpannerServer(grpcServer, server)
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
grpcServer.Stop()
_ = listener.Close()
roles := make(chan string, 1)
dialOptions := grpctest.Inspect(func(_ string, req any) error {
if req, ok := req.(*sppb.CreateSessionRequest); ok {
select {
case roles <- req.GetSession().GetCreatorRole():
default:
}
}
return nil
})
go func() { _ = grpcServer.Serve(listener) }()
t.Setenv("SPANNER_EMULATOR_HOST", listener.Addr().String())

ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
client, err := newClient(ctx, "p", "i", "d", tt.role, logGrpcModeOff, false)
client, err := newClient(ctx, env.ProjectID, env.InstanceID, env.DatabaseID, tt.role, logGrpcModeOff, false, option.WithGRPCDialOption(dialOptions[0]), option.WithGRPCDialOption(dialOptions[1]))
if err != nil {
t.Fatal(err)
}
defer client.Close()

select {
case req := <-server.createSession:
if got := req.GetSession().GetCreatorRole(); got != tt.role {
case got := <-roles:
if got != tt.role {
t.Fatalf("session CreatorRole = %q, want %q", got, tt.role)
}
case <-ctx.Done():
Expand Down
58 changes: 58 additions & 0 deletions internal/grpctest/intercept.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
// Package grpctest provides client-side RPC inspection for transport tests.
// It complements emulator tests without implementing Spanner server behavior.
package grpctest

import (
"context"

"google.golang.org/grpc"
)

// Inspect calls inspect before each unary request or stream message is sent.
// Returning a non-nil gRPC status error prevents that message from being sent.
// The callback may run concurrently; it must not mutate or retain the message
// without copying it. Returning nil preserves the actual server response.
func Inspect(inspect func(method string, request any) error) []grpc.DialOption {
return []grpc.DialOption{
grpc.WithChainUnaryInterceptor(func(ctx context.Context, method string, req, reply any, cc *grpc.ClientConn, invoker grpc.UnaryInvoker, opts ...grpc.CallOption) error {
if err := inspect(method, req); err != nil {
return err
}
return invoker(ctx, method, req, reply, cc, opts...)
}),
grpc.WithChainStreamInterceptor(func(ctx context.Context, desc *grpc.StreamDesc, cc *grpc.ClientConn, method string, streamer grpc.Streamer, opts ...grpc.CallOption) (grpc.ClientStream, error) {
ctx, cancel := context.WithCancel(ctx)
stream, err := streamer(ctx, desc, cc, method, opts...)
if err != nil {
cancel()
return nil, err
}
return &inspectingStream{ClientStream: stream, inspect: inspect, method: method, cancel: cancel}, nil
}),
}
}

type inspectingStream struct {
grpc.ClientStream
inspect func(string, any) error
method string
cancel context.CancelFunc
}

func (s *inspectingStream) SendMsg(m any) error {
if err := s.inspect(s.method, m); err != nil {
s.cancel()
return err
}
// Preserve transport errors (including io.EOF); RecvMsg may still need
// to retrieve the server status.
return s.ClientStream.SendMsg(m)
}

func (s *inspectingStream) RecvMsg(m any) error {
err := s.ClientStream.RecvMsg(m)
if err != nil {
s.cancel()
}
return err
}
68 changes: 68 additions & 0 deletions internal/grpctest/intercept_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
package grpctest

import (
"context"
"io"
"testing"

"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)

type transportStream struct {
grpc.ClientStream
sent any
sendErr error
recvErr error
}

func (s *transportStream) SendMsg(m any) error { s.sent = m; return s.sendErr }
func (s *transportStream) RecvMsg(any) error { return s.recvErr }

func TestInspectingStreamRejectsBeforeSending(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
transport := &transportStream{}
rejected := status.Error(codes.InvalidArgument, "captured")
stream := &inspectingStream{ClientStream: transport, method: "/test/Query", cancel: cancel,
inspect: func(method string, req any) error {
if method != "/test/Query" || req != "request" {
t.Fatalf("inspection = %s %v", method, req)
}
return rejected
}}
if err := stream.SendMsg("request"); err != rejected {
t.Fatalf("error = %v", err)
}
if transport.sent != nil {
t.Fatal("rejected request reached transport")
}
if ctx.Err() == nil {
t.Fatal("rejected stream was not cancelled")
}
}

func TestInspectingStreamPreservesServerStatusAfterSendEOF(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
serverErr := status.Error(codes.PermissionDenied, "denied")
transport := &transportStream{sendErr: io.EOF, recvErr: serverErr}
stream := &inspectingStream{ClientStream: transport, cancel: cancel,
inspect: func(string, any) error { return nil }}
if err := stream.SendMsg("request"); err != io.EOF {
t.Fatalf("send error = %v", err)
}
if transport.sent != "request" {
t.Fatal("request was not forwarded")
}
if ctx.Err() != nil {
t.Fatal("cancelled before receiving server status")
}
if err := stream.RecvMsg(nil); err != serverErr {
t.Fatalf("receive error = %v", err)
}
if ctx.Err() == nil {
t.Fatal("finished stream was not cancelled")
}
}
43 changes: 43 additions & 0 deletions jqresult/protojson_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,9 @@
package jqresult

import (
sppb "cloud.google.com/go/spanner/apiv1/spannerpb"
"github.com/apstndb/execspansql/resultset"
"reflect"
"testing"

"github.com/apstndb/spaniter"
Expand Down Expand Up @@ -52,3 +55,43 @@ func TestResultSetMapFromRowIteratorNil(t *testing.T) {
t.Fatal("error = nil, want nil row iterator error")
}
}

// Both eager result materialization and lazy stats conversion must preserve
// arbitrary plan/stat contents, independently of RPC query mode support.
func TestPlanAndStatsAcrossResultPaths(t *testing.T) {
t.Parallel()
for _, withPlan := range []bool{false, true} {
for _, withStats := range []bool{false, true} {
stats := spaniter.Stats{}
want := map[string]any{}
if withPlan {
stats.QueryPlan = &sppb.QueryPlan{PlanNodes: []*sppb.PlanNode{{DisplayName: "Test Scan"}}}
want["queryPlan"] = map[string]any{"planNodes": []any{map[string]any{"displayName": "Test Scan"}}}
}
if withStats {
stats.QueryStats = map[string]any{"elapsed_time": "1 msecs", "nested": map[string]any{"complete": true}}
want["queryStats"] = stats.QueryStats
}
input := spaniter.RowIteratorResult{Stats: stats}
lazy, err := StatsMapFromResult(input)
if err != nil {
t.Fatal(err)
}
rs, err := resultset.FromIteratorResult(nil, input)
if err != nil {
t.Fatal(err)
}
eager, err := ResultSetMap(rs)
if err != nil {
t.Fatal(err)
}
if len(want) == 0 {
if lazy != nil || eager["stats"] != nil {
t.Fatalf("empty stats: lazy=%v eager=%v", lazy, eager)
}
} else if !reflect.DeepEqual(lazy, want) || !reflect.DeepEqual(eager["stats"], want) {
t.Fatalf("plan=%v stats=%v: lazy=%v eager=%v want=%v", withPlan, withStats, lazy, eager["stats"], want)
}
}
}
}
11 changes: 8 additions & 3 deletions main.go
Original file line number Diff line number Diff line change
Expand Up @@ -402,6 +402,11 @@ func runInNewTransaction(ctx context.Context, client *spanner.Client, stmt spann
}

func _main() error {
return runCLI()
}

// runCLI accepts client options so transport tests can inspect outgoing RPCs.
func runCLI(clientOptions ...option.ClientOption) error {
o, err := processFlags()
if err != nil {
os.Exit(1)
Expand Down Expand Up @@ -469,7 +474,7 @@ func _main() error {
}()
}

client, err := newClient(ctx, o.Project, o.Instance, o.Database, o.DatabaseRole, o.LogGrpc, tracingEnabled(o))
client, err := newClient(ctx, o.Project, o.Instance, o.Database, o.DatabaseRole, o.LogGrpc, tracingEnabled(o), clientOptions...)
if err != nil {
return err
}
Expand Down Expand Up @@ -598,7 +603,7 @@ func writeCsvFromResultSet(writer io.Writer, rs *sppb.ResultSet) error {
return csvWriter.Flush()
}

func newClient(ctx context.Context, project, instance, database, databaseRole string, logGrpcMode string, doTrace bool) (*spanner.Client, error) {
func newClient(ctx context.Context, project, instance, database, databaseRole string, logGrpcMode string, doTrace bool, clientOptions ...option.ClientOption) (*spanner.Client, error) {
name, err := databaseResourceName(project, instance, database)
if err != nil {
return nil, err
Expand All @@ -612,7 +617,7 @@ func newClient(ctx context.Context, project, instance, database, databaseRole st
if doTrace {
copts = append(copts, option.WithGRPCDialOption(grpc.WithChainStreamInterceptor(interceptor.StreamInterceptor(interceptor.WithDefaultDecorators()))))
}
return spanner.NewClientWithConfig(ctx, name, spanner.ClientConfig{DatabaseRole: databaseRole}, copts...)
return spanner.NewClientWithConfig(ctx, name, spanner.ClientConfig{DatabaseRole: databaseRole}, append(copts, clientOptions...)...)
}

type encoder interface {
Expand Down
5 changes: 3 additions & 2 deletions pdml_query_mode_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (

"cloud.google.com/go/spanner"
"github.com/apstndb/spanemuboost"
"google.golang.org/api/option"
)

const pdmlQueryModeAuditSQL = "UPDATE PdmlQueryModeAudit SET V=99 WHERE TRUE"
Expand Down Expand Up @@ -73,12 +74,12 @@ func TestPartitionedDMLQueryMode(t *testing.T) {
}
}

func runMain(t *testing.T, args []string) error {
func runMain(t *testing.T, args []string, options ...option.ClientOption) error {
t.Helper()
old := os.Args
os.Args = append([]string{"execspansql"}, args...)
defer func() { os.Args = old }()
return _main()
return runCLI(options...)
}

func readPdmlQueryModeV(t *testing.T, ctx context.Context, client *spanner.Client) int64 {
Expand Down
Loading
Loading