Skip to content
Open
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
11 changes: 0 additions & 11 deletions client.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,17 +40,6 @@ type ClientOpts struct {
RpcInterceptors []ClientRPCInterceptor
MultiRPCInterceptors []ClientMultiRPCInterceptor
StreamInterceptors []StreamInterceptor
SkipClaim func() bool
}

// WithClientSkipClaim lets a queue rpc bypass the claim handshake while enabled,
// which is only sound if the bus delivers a queue subscription to exactly one
// subscriber. Consulted per request so it can be revoked at runtime without a
// redeploy, and disabled when unset.
func WithClientSkipClaim(enabled func() bool) ClientOption {
return func(o *ClientOpts) {
o.SkipClaim = enabled
}
}

func WithClientID(id string) ClientOption {
Expand Down
156 changes: 78 additions & 78 deletions internal/internal.pb.go

Large diffs are not rendered by default.

6 changes: 4 additions & 2 deletions internal/internal.proto
Original file line number Diff line number Diff line change
Expand Up @@ -38,8 +38,10 @@ message Request {
google.protobuf.Any request = 6;
map<string, string> metadata = 7;
bytes raw_request = 8;
// Advisory: the caller still answers a claim, so older servers are unaffected.
bool skip_claim = 9;
// 9 was the caller-elected skip, which older servers honor unconditionally.
reserved 9;
// Advertises that an announcement may replace the claim; the server decides.
bool skip_claim = 10;
}

message Response {
Expand Down
99 changes: 90 additions & 9 deletions internal/test/skipclaim_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -45,10 +45,9 @@ func TestSkipClaim(t *testing.T) {
b := newBus(t)

s := server.NewRPCServer(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithServerObserver(obs))
psrpc.WithServerObserver(obs), psrpc.WithServerSkipClaim(enabled))
t.Cleanup(func() { s.Close(true) })
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithClientSkipClaim(enabled))
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b)
require.NoError(t, err)
t.Cleanup(func() { c.Close() })

Expand Down Expand Up @@ -162,10 +161,10 @@ func TestSkipClaimSlowHandler(t *testing.T) {
}
}))

s := server.NewRPCServer(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b)
s := server.NewRPCServer(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithServerSkipClaim(enabled))
t.Cleanup(func() { s.Close(true) })
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithClientSkipClaim(enabled))
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b)
require.NoError(t, err)
t.Cleanup(func() { c.Close() })

Expand Down Expand Up @@ -218,10 +217,9 @@ func TestSkipClaimRevokedAtRuntime(t *testing.T) {
on.Store(true)

s := server.NewRPCServer(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithServerObserver(obs))
psrpc.WithServerObserver(obs), psrpc.WithServerSkipClaim(on.Load))
t.Cleanup(func() { s.Close(true) })
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithClientSkipClaim(on.Load))
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b)
require.NoError(t, err)
t.Cleanup(func() { c.Close() })

Expand All @@ -248,3 +246,86 @@ func TestSkipClaimRevokedAtRuntime(t *testing.T) {
psrpc.ClaimSkipped, psrpc.ClaimGranted, psrpc.ClaimSkipped,
}, claims)
}

// A caller that does not advertise must be negotiated with, even by a server
// that elected to skip.
func TestSkipClaimCallerDoesNotAdvertise(t *testing.T) {
obs := &recordingObserver{}
b := testutils.NewTestBus(bus.NewLocalMessageBus(),
testutils.WithPublishInterceptor(func(next testutils.PublishHandler) testutils.PublishHandler {
return func(ctx context.Context, channel testutils.Channel, msg proto.Message) error {
if req, ok := msg.(*internal.Request); ok {
// As a client predating the field would leave it.
req.SkipClaim = false
}
return next(ctx, channel, msg)
}
}))

s := server.NewRPCServer(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithServerObserver(obs), psrpc.WithServerSkipClaim(enabled))
t.Cleanup(func() { s.Close(true) })
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b)
require.NoError(t, err)
t.Cleanup(func() { c.Close() })

s.RegisterMethod("queued", false, false, true, true)
c.RegisterMethod("queued", false, false, true, true)
require.NoError(t, server.RegisterHandler(s, "queued", nil,
func(context.Context, *internal.Request) (*internal.Response, error) {
return &internal.Response{}, nil
}, nil))

_, err = client.RequestSingle[*internal.Response](context.Background(), c, "queued", nil, &internal.Request{})
require.NoError(t, err)

_, claims := obs.snapshot()
require.Equal(t, []psrpc.ClaimOutcome{psrpc.ClaimGranted}, claims,
"a caller that did not advertise must be negotiated with")
}

// CS-1992: an error response that beat the announcement was stashed while the
// caller waited out the timeout. The bus delays announcements to force that order.
func TestSkipClaimFastFailingHandler(t *testing.T) {
b := testutils.NewTestBus(bus.NewLocalMessageBus(),
testutils.WithPublishInterceptor(func(next testutils.PublishHandler) testutils.PublishHandler {
return func(ctx context.Context, channel testutils.Channel, msg proto.Message) error {
if _, ok := msg.(*internal.ClaimRequest); ok {
go func() {
time.Sleep(50 * time.Millisecond)
_ = next(ctx, channel, msg)
}()
return nil
}
return next(ctx, channel, msg)
}
}))

s := server.NewRPCServer(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b,
psrpc.WithServerSkipClaim(enabled))
t.Cleanup(func() { s.Close(true) })
c, err := client.NewRPCClient(&info.ServiceDefinition{Name: "test", ID: rand.NewString()}, b)
require.NoError(t, err)
t.Cleanup(func() { c.Close() })

s.RegisterMethod("queued", false, false, true, true)
c.RegisterMethod("queued", false, false, true, true)
require.NoError(t, server.RegisterHandler(s, "queued", nil,
func(context.Context, *internal.Request) (*internal.Response, error) {
return nil, psrpc.NewErrorf(psrpc.NotFound, "requested room does not exist")
}, nil))

const timeout = time.Second

start := time.Now()
_, err = client.RequestSingle[*internal.Response](context.Background(), c, "queued", nil,
&internal.Request{}, psrpc.WithRequestTimeout(timeout))

require.Error(t, err)
require.NotErrorIs(t, err, psrpc.ErrRequestTimedOut,
"the handler answered, so the caller must not time out")
code, ok := psrpc.GetErrorCode(err)
require.True(t, ok)
require.Equal(t, psrpc.NotFound, code, "the handler's error must reach the caller")
require.Less(t, time.Since(start), timeout/2, "the answer must not wait out the timeout")
}
37 changes: 33 additions & 4 deletions pkg/client/client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ func testAffinity(t *testing.T, opts psrpc.SelectionOpts, expectedID string) {
Affinity: 0.9,
}
}()
sel, err := selectServer(context.Background(), c, nil, opts)
sel, err := selectServer(context.Background(), c, nil, opts, false)
require.NoError(t, err)
require.Equal(t, expectedID, sel.serverID)
}
Expand All @@ -99,7 +99,7 @@ func TestSelectServerGrantsABid(t *testing.T) {
claims <- &internal.ClaimRequest{RequestId: "1", ServerId: "2", Affinity: 1}

sel, err := selectServer(context.Background(), claims, make(chan *internal.Response, 1),
psrpc.SelectionOpts{AcceptFirstAvailable: true})
psrpc.SelectionOpts{AcceptFirstAvailable: true}, true)
require.NoError(t, err)
require.Equal(t, "2", sel.serverID)
require.False(t, sel.handling, "a bid still needs granting")
Expand All @@ -113,7 +113,7 @@ func TestSelectServerHonorsAnnouncement(t *testing.T) {
claims <- &internal.ClaimRequest{RequestId: "1", ServerId: "2", Affinity: 1, Handling: true}

sel, err := selectServer(context.Background(), claims, make(chan *internal.Response, 1),
psrpc.SelectionOpts{MinimumAffinity: 2, AffinityTimeout: time.Second})
psrpc.SelectionOpts{MinimumAffinity: 2, AffinityTimeout: time.Second}, true)
require.NoError(t, err)
require.Equal(t, "2", sel.serverID)
require.True(t, sel.handling)
Expand All @@ -127,9 +127,38 @@ func TestSelectServerReturnsEarlyResponse(t *testing.T) {
responses <- &internal.Response{RequestId: "1", ServerId: "2"}

sel, err := selectServer(context.Background(), make(chan *internal.ClaimRequest, 1), responses,
psrpc.SelectionOpts{AcceptFirstAvailable: true})
psrpc.SelectionOpts{AcceptFirstAvailable: true}, true)
require.NoError(t, err)
require.NotNil(t, sel.res)
require.Equal(t, "2", sel.res.ServerId)
require.Empty(t, sel.serverID)
}

// CS-1992: on queue an error response is the answer, not a fallback.
func TestSelectServerReturnsQueueError(t *testing.T) {
responses := make(chan *internal.Response, 1)
responses <- &internal.Response{RequestId: "1", ServerId: "2", Error: "not found", Code: "not_found"}

sel, err := selectServer(context.Background(), make(chan *internal.ClaimRequest, 1), responses,
psrpc.SelectionOpts{AcceptFirstAvailable: true}, true)
require.NoError(t, err)
require.NotNil(t, sel.res, "the error response is the answer, not a fallback")
}

// On broadcast an early error is held back so a healthy bid can win.
func TestSelectServerStashesBroadcastError(t *testing.T) {
responses := make(chan *internal.Response, 1)
responses <- &internal.Response{RequestId: "1", ServerId: "2", Error: "boom", Code: "internal"}
claims := make(chan *internal.ClaimRequest, 1)

go func() {
time.Sleep(50 * time.Millisecond)
claims <- &internal.ClaimRequest{RequestId: "1", ServerId: "3", Affinity: 1}
}()

sel, err := selectServer(context.Background(), claims, responses,
psrpc.SelectionOpts{AcceptFirstAvailable: true, AffinityTimeout: time.Second}, false)
require.NoError(t, err)
require.Equal(t, "3", sel.serverID, "a healthy bid must win over a stashed rejection")
require.Nil(t, sel.res)
}
14 changes: 7 additions & 7 deletions pkg/client/rpc.go
Original file line number Diff line number Diff line change
Expand Up @@ -109,8 +109,8 @@ func newRPC[ResponseType proto.Message](c *RPCClient, i *info.RequestInfo) psrpc
Multi: false,
RawRequest: b,
Metadata: metadata.OutgoingContextMetadata(ctx),
// The queue already chose the server; the claim only ratifies it.
SkipClaim: i.Queue && c.SkipClaim != nil && c.SkipClaim(),
// Advertises that an announcement may replace the claim; making one is the server's call.
SkipClaim: i.Queue,
}

var claimChan chan *internal.ClaimRequest
Expand Down Expand Up @@ -144,7 +144,7 @@ func newRPC[ResponseType proto.Message](c *RPCClient, i *info.RequestInfo) psrpc
var res *internal.Response

if i.RequireClaim {
sel, err := selectServer(ctx, claimChan, resChan, o.SelectionOpts)
sel, err := selectServer(ctx, claimChan, resChan, o.SelectionOpts, i.Queue)
if err != nil {
return nil, err
}
Expand Down Expand Up @@ -208,6 +208,7 @@ func selectServer(
claimChan chan *internal.ClaimRequest,
resChan chan *internal.Response,
opts psrpc.SelectionOpts,
queue bool,
) (selection, error) {

ctx, cancel := context.WithCancel(ctx)
Expand Down Expand Up @@ -268,12 +269,11 @@ func selectServer(
}

case res := <-resChan:
if res.Error == "" {
// Only a server that never waited to be granted answers this early,
// and consuming it here would strand the response.
// On queue the sole responder's answer is final, error or not; on
// broadcast an early error may yet be outbid, so it is held back.
if res.Error == "" || queue {
return selection{res: res}, nil
}
// otherwise a malformed request, which is answered before any claim
resErr = psrpc.NewErrorFromResponse(res.Code, res.Error, res.ErrorDetails...)
}
}
Expand Down
3 changes: 2 additions & 1 deletion pkg/client/stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -95,7 +95,8 @@ func OpenStream[SendType, RecvType proto.Message](
}

if i.RequireClaim {
sel, err := selectServer(ctx, claimChan, nil, o.SelectionOpts)
// nil resChan, so queue-ness is moot
sel, err := selectServer(ctx, claimChan, nil, o.SelectionOpts, false)
if err != nil {
_ = cs.Close(err)
return nil, err
Expand Down
7 changes: 3 additions & 4 deletions pkg/server/rpc.go
Original file line number Diff line number Diff line change
Expand Up @@ -219,10 +219,9 @@ func (h *rpcHandlerImpl[RequestType, ResponseType]) claimRequest(
affinity = 1
}

// A queue subscription already chose this server, so the claim is announced
// rather than negotiated. Queue is re-checked because honoring SkipClaim on a
// broadcast rpc would let every server run the handler.
handling := ir.SkipClaim && h.i.Queue
// The queue re-check keeps a broadcast rpc from running on every server.
serverSkip := s.SkipClaim != nil && s.SkipClaim()
handling := ir.SkipClaim && serverSkip && h.i.Queue

var claimResponseChan chan *internal.ClaimResponse
if !handling {
Expand Down
9 changes: 9 additions & 0 deletions server.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,15 @@ type ServerOpts struct {
StreamInterceptors []StreamInterceptor
ChainedInterceptor ServerRPCInterceptor
RequestObserver RequestObserver
SkipClaim func() bool
}

// WithServerSkipClaim answers advertised queue rpcs with an announcement rather
// than a claim. Read per request, so revocable at runtime; off when unset.
func WithServerSkipClaim(enabled func() bool) ServerOption {
return func(o *ServerOpts) {
o.SkipClaim = enabled
}
}

func WithServerID(id string) ServerOption {
Expand Down
Loading