From ef43e9592caf48fcf67976fcaec48f4e8a67a6c9 Mon Sep 17 00:00:00 2001 From: cabbage Date: Sun, 6 Sep 2026 10:03:11 +0000 Subject: [PATCH] feat: wake reconciled sessions for queued dispatch --- cmd/rvbox-server/main.go | 6 +- internal/server/control/service.go | 7 +- internal/server/session/agent_server.go | 106 ++++++++++++++---- internal/server/session/session.go | 37 +++++- internal/server/session/session_test.go | 11 ++ test/coverage.toml | 6 + .../clientagent_integration_test.go | 61 ++++++++++ 7 files changed, 210 insertions(+), 24 deletions(-) diff --git a/cmd/rvbox-server/main.go b/cmd/rvbox-server/main.go index 3a52604..c5184bc 100644 --- a/cmd/rvbox-server/main.go +++ b/cmd/rvbox-server/main.go @@ -69,9 +69,11 @@ func run(configPath string) error { return err } defer func() { _ = cleanupControl() }() + registry := session.NewRegistry() controlService, err := control.NewService(control.Options{ Store: persistence, DefaultQueueTTL: configured.Queue.DefaultTTL, DefaultQueueTTLSet: true, - Limits: agentproto.Limits{MaxEnvelopeBytes: configured.Protocol.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: configured.Protocol.MaxExecutionSpecBytes, MaxRawChunkBytes: configured.Protocol.MaxRawChunkBytes, MaxScriptBytes: configured.Protocol.MaxScriptBytes, MaxDetailBytes: configured.Storage.ProtocolDetailMaxBytes}, + WakeClient: registry.Wake, + Limits: agentproto.Limits{MaxEnvelopeBytes: configured.Protocol.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: configured.Protocol.MaxExecutionSpecBytes, MaxRawChunkBytes: configured.Protocol.MaxRawChunkBytes, MaxScriptBytes: configured.Protocol.MaxScriptBytes, MaxDetailBytes: configured.Storage.ProtocolDetailMaxBytes}, }) if err != nil { return err @@ -79,7 +81,7 @@ func run(configPath string) error { grpcServer := grpc.NewServer(grpc.MaxRecvMsgSize(int(configured.Protocol.MaxControlRequestBytes)), grpc.MaxSendMsgSize(int(configured.Protocol.MaxControlRequestBytes))) rvboxv1.RegisterControlServer(grpcServer, controlService) agent := &session.AgentServer{ - Store: persistence, Registry: session.NewRegistry(), Path: configured.Server.AgentPath, + Store: persistence, Registry: registry, Path: configured.Server.AgentPath, Limits: agentproto.Limits{ MaxEnvelopeBytes: configured.Protocol.MaxAgentEnvelopeBytes, MaxExecutionSpecBytes: configured.Protocol.MaxExecutionSpecBytes, diff --git a/internal/server/control/service.go b/internal/server/control/service.go index 288da0a..8b90b77 100644 --- a/internal/server/control/service.go +++ b/internal/server/control/service.go @@ -44,6 +44,7 @@ type Options struct { Limits agentproto.Limits Now func() time.Time CursorKey []byte + WakeClient func(string) bool } type Service struct { @@ -53,6 +54,7 @@ type Service struct { limits agentproto.Limits now func() time.Time cursors *domain.CursorCodec + wakeClient func(string) bool } func NewService(options Options) (*Service, error) { @@ -82,7 +84,7 @@ func NewService(options Options) (*Service, error) { if err != nil { return nil, err } - return &Service{store: options.Store, defaultQueueTTL: options.DefaultQueueTTL, limits: options.Limits, now: options.Now, cursors: codec}, nil + return &Service{store: options.Store, defaultQueueTTL: options.DefaultQueueTTL, limits: options.Limits, now: options.Now, cursors: codec, wakeClient: options.WakeClient}, nil } func (service *Service) ListClients(ctx context.Context, request *rvboxv1.ListClientsRequest) (*rvboxv1.ListClientsResponse, error) { @@ -264,6 +266,9 @@ func (service *Service) RunCommand(ctx context.Context, request *rvboxv1.RunComm if err != nil { return nil, mapStoreError(err) } + if service.wakeClient != nil { + _ = service.wakeClient(request.GetTargetClientId()) + } return &rvboxv1.RunCommandResponse{IssueUuid: issue.String(), Lifecycle: rvboxv1.CommandLifecycle(queued.Lifecycle)}, nil } diff --git a/internal/server/session/agent_server.go b/internal/server/session/agent_server.go index 16f539c..6620fdb 100644 --- a/internal/server/session/agent_server.go +++ b/internal/server/session/agent_server.go @@ -139,10 +139,15 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w queue := NewWriterQueue(16, 64, 8) defer queue.Close() capacity := &CapacityShadow{} + var capacityMu sync.Mutex + reservations := make(map[string]DispatchLane) + reconciled := make(chan struct{}) + var reconcileOnce sync.Once if !capacity.UpdateAdvertised(0, 0, hello.GetMaxRunningCommands(), hello.GetMaxQueuedCommands()) { server.close(connection, websocket.StatusPolicyViolation, "invalid initial client capacity") return } + go server.dispatchLoop(sessionContext, queue, handle.DispatchWake(), reconciled, &capacityMu, capacity, reservations, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel) encodedSessionID := encodeSessionID(sessionID) welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, @@ -210,25 +215,19 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w case <-sessionContext.Done(): return } - lane := capacity.Reserve() - dispatched, dispatchErr := false, error(nil) - if lane != DispatchNone { - dispatched, dispatchErr = server.enqueueNextDispatch(sessionContext, queue, hello.GetClientId(), hello.GetPlatform(), encodedSessionID, registration.Generation) - if !dispatched { - capacity.Release(lane) - } - } - if dispatchErr != nil && !errors.Is(dispatchErr, ErrDispatchDataFull) { - server.close(connection, websocket.StatusInternalError, "could not queue command dispatch") - return - } + reconcileOnce.Do(func() { close(reconciled) }) + handle.SignalDispatch() continue } if advertised := envelope.GetClientCapacity(); advertised != nil { - if !capacity.UpdateAdvertised(advertised.GetRunningCommands(), advertised.GetQueuedCommands(), advertised.GetMaxRunningCommands(), advertised.GetMaxQueuedCommands()) { + capacityMu.Lock() + validCapacity := capacity.UpdateAdvertised(advertised.GetRunningCommands(), advertised.GetQueuedCommands(), advertised.GetMaxRunningCommands(), advertised.GetMaxQueuedCommands()) + capacityMu.Unlock() + if !validCapacity { server.close(connection, websocket.StatusPolicyViolation, "invalid client capacity") return } + handle.SignalDispatch() continue } if acknowledgement := envelope.GetCommandAccepted(); acknowledgement != nil { @@ -241,6 +240,13 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w server.close(connection, websocket.StatusPolicyViolation, "invalid command acceptance") return } + capacityMu.Lock() + if lane, reserved := reservations[issue.String()]; reserved { + capacity.Release(lane) + delete(reservations, issue.String()) + } + capacityMu.Unlock() + handle.SignalDispatch() continue } if event := envelope.GetCommandEvent(); event != nil { @@ -315,10 +321,10 @@ func eventType(event *rvboxv1.CommandEvent) uint16 { // enqueueNextDispatch records the queued-to-dispatched transition before // exposing work to the network. A full data lane is a pre-write failure, so // only the owning generation can put the command back into the queue. -func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *WriterQueue, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64) (bool, error) { +func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *WriterQueue, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, beforeEnqueue func(domain.UUID)) (domain.UUID, bool, error) { candidate, err := server.Store.ClaimNextDispatch(ctx, clientID, generation, server.now()) if err != nil || candidate == nil { - return false, err + return domain.UUID{}, false, err } requeue := func(cause error) (bool, error) { _, rollbackErr := server.Store.RequeueDispatch(ctx, candidate.IssueUUID, clientID, generation) @@ -329,10 +335,12 @@ func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *Write } spec := &rvboxv1.ExecutionSpec{} if err := proto.Unmarshal(candidate.ExecutionSpec, spec); err != nil { - return requeue(fmt.Errorf("decode persisted execution spec: %w", err)) + sent, requeueErr := requeue(fmt.Errorf("decode persisted execution spec: %w", err)) + return candidate.IssueUUID, sent, requeueErr } if err := agentproto.ValidateExecutionSpec(spec, server.limits(), platform); err != nil { - return requeue(fmt.Errorf("validate persisted execution spec: %w", err)) + sent, requeueErr := requeue(fmt.Errorf("validate persisted execution spec: %w", err)) + return candidate.IssueUUID, sent, requeueErr } dispatch := &rvboxv1.CommandDispatch{ IssueUuid: candidate.IssueUUID.String(), CommandRevision: candidate.Revision, TargetSessionGeneration: generation, @@ -345,12 +353,70 @@ func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *Write SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_CommandDispatch{CommandDispatch: dispatch}, }) if err != nil { - return requeue(err) + sent, requeueErr := requeue(err) + return candidate.IssueUUID, sent, requeueErr + } + if beforeEnqueue != nil { + beforeEnqueue(candidate.IssueUUID) } if !queue.EnqueueData(Frame{Kind: FrameData, Payload: encoded}) { - return requeue(ErrDispatchDataFull) + sent, requeueErr := requeue(ErrDispatchDataFull) + return candidate.IssueUUID, sent, requeueErr + } + return candidate.IssueUUID, true, nil +} + +// dispatchLoop is the per-session serialized dispatcher. It waits for a +// complete reconciliation result before consuming queued work, then coalesces +// wakeups from local control RPCs, capacity advertisements, and acceptances. +func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue, wake <-chan struct{}, reconciled <-chan struct{}, capacityMu *sync.Mutex, capacity *CapacityShadow, reservations map[string]DispatchLane, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) { + select { + case <-reconciled: + case <-ctx.Done(): + return + } + for { + for { + capacityMu.Lock() + lane := capacity.Reserve() + capacityMu.Unlock() + if lane == DispatchNone { + break + } + inserted := false + issue, sent, err := server.enqueueNextDispatch(ctx, queue, clientID, platform, sessionID, generation, func(issue domain.UUID) { + capacityMu.Lock() + reservations[issue.String()] = lane + inserted = true + capacityMu.Unlock() + }) + if !sent { + capacityMu.Lock() + if inserted { + delete(reservations, issue.String()) + } + capacity.Release(lane) + capacityMu.Unlock() + if err != nil && !errors.Is(err, ErrDispatchDataFull) { + server.closeForDispatchFailure(cancel) + return + } + break + } + } + select { + case <-wake: + continue + case <-ctx.Done(): + return + } + } +} + +func (server *AgentServer) closeForDispatchFailure(cancel context.CancelFunc) { + if cancel != nil { + cancel() } - return true, nil } func (server *AgentServer) writeLoop(ctx context.Context, connection *websocket.Conn, queue *WriterQueue, heartbeat *synchronizedHeartbeat, started time.Time) { diff --git a/internal/server/session/session.go b/internal/server/session/session.go index 0390218..5b35acb 100644 --- a/internal/server/session/session.go +++ b/internal/server/session/session.go @@ -231,6 +231,7 @@ type Handle struct { SessionID [16]byte Generation uint64 Context context.Context + wake chan struct{} cancel context.CancelFunc } @@ -253,11 +254,45 @@ func (registry *Registry) Install(clientID string, sessionID [16]byte, generatio old.cancel() } ctx, cancel := context.WithCancel(context.Background()) - handle := &Handle{ClientID: clientID, SessionID: sessionID, Generation: generation, Context: ctx, cancel: cancel} + handle := &Handle{ClientID: clientID, SessionID: sessionID, Generation: generation, Context: ctx, wake: make(chan struct{}, 1), cancel: cancel} registry.clients[clientID] = handle return handle, nil } +// Wake requests a best-effort dispatch pass for the current generation. The +// signal is edge-triggered and coalesced; durable queue state is the source of +// truth, so dropping a redundant wake cannot lose work. +func (registry *Registry) Wake(clientID string) bool { + if registry == nil || clientID == "" { + return false + } + registry.mu.Lock() + handle := registry.clients[clientID] + registry.mu.Unlock() + if handle == nil { + return false + } + handle.SignalDispatch() + return true +} + +func (handle *Handle) DispatchWake() <-chan struct{} { + if handle == nil { + return nil + } + return handle.wake +} + +func (handle *Handle) SignalDispatch() { + if handle == nil || handle.wake == nil { + return + } + select { + case handle.wake <- struct{}{}: + default: + } +} + func (registry *Registry) Remove(handle *Handle) bool { if handle == nil { return false diff --git a/internal/server/session/session_test.go b/internal/server/session/session_test.go index 4bbfaa4..5f8a359 100644 --- a/internal/server/session/session_test.go +++ b/internal/server/session/session_test.go @@ -100,6 +100,17 @@ func TestCapacityShadowAndRegistry_BH_SES_03(t *testing.T) { if registry.Remove(first) || registry.Get("client-a") != second { t.Fatal("stale handle removed current session") } + if !registry.Wake("client-a") { + t.Fatal("wake did not find current session") + } + select { + case <-second.DispatchWake(): + default: + t.Fatal("wake was not delivered to current session") + } + if registry.Wake("missing") { + t.Fatal("wake reported a missing session") + } if !registry.Remove(second) { t.Fatal("current handle removal failed") } diff --git a/test/coverage.toml b/test/coverage.toml index 65e1498..2bc9176 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -254,6 +254,12 @@ layer = "unit" status = "implemented" tests = ["internal/server/store/command_test.go:TestRecordCommandAcceptanceFencesGeneration_HP_DISPATCH_05"] +[[requirements]] +id = "HP-DISPATCH-06" +layer = "integration" +status = "implemented" +tests = ["test/integration/clientagent/clientagent_integration_test.go:TestControlQueueWakesReconciledSession_HP_DISPATCH_06"] + [[requirements]] id = "HP-EVENT-01" layer = "unit" diff --git a/test/integration/clientagent/clientagent_integration_test.go b/test/integration/clientagent/clientagent_integration_test.go index e5e9af3..b066580 100644 --- a/test/integration/clientagent/clientagent_integration_test.go +++ b/test/integration/clientagent/clientagent_integration_test.go @@ -13,8 +13,11 @@ import ( "github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/client/agent" "github.com/rvbox/rvbox/internal/domain" + "github.com/rvbox/rvbox/internal/server/control" "github.com/rvbox/rvbox/internal/server/session" "github.com/rvbox/rvbox/internal/server/store" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "google.golang.org/protobuf/proto" "google.golang.org/protobuf/types/known/timestamppb" ) @@ -137,3 +140,61 @@ func TestWebSocketDispatchAfterReconciliation_HP_DISPATCH_03(t *testing.T) { t.Fatalf("command lifecycle = %d, %v", lifecycle, err) } } + +func TestControlQueueWakesReconciledSession_HP_DISPATCH_06(t *testing.T) { + ctx := context.Background() + persistence, err := store.Open(ctx, store.Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer persistence.Close() + registry := session.NewRegistry() + server := httptest.NewServer(&session.AgentServer{Store: persistence, Registry: registry, Path: "/v1/agent", Limits: agentproto.DefaultLimits()}) + defer server.Close() + controlService, err := control.NewService(control.Options{Store: persistence, WakeClient: registry.Wake, CursorKey: []byte("0123456789abcdef0123456789abcdef")}) + if err != nil { + t.Fatal(err) + } + address := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/agent" + transport, err := agent.DialWebSocket(ctx, address, nil) + if err != nil { + t.Fatal(err) + } + defer transport.Close() + hello := &rvboxv1.ClientHello{ClientId: "win-wake-client", SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: "019c46f1-1d02-7000-8000-000000000076", MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now()} + accepted, err := agent.Handshake(ctx, transport, hello, agentproto.DefaultLimits()) + if err != nil { + t.Fatal(err) + } + if _, err := agent.Reconcile(ctx, transport, accepted, &rvboxv1.ReconcileSnapshot{}, agentproto.DefaultLimits()); err != nil { + t.Fatal(err) + } + issue, err := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000077") + if err != nil { + t.Fatal(err) + } + response, err := controlService.RunCommand(ctx, &rvboxv1.RunCommandRequest{TargetClientId: hello.GetClientId(), RequestId: issue.String(), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "Write-Output wake"}}}) + if err != nil || response.GetIssueUuid() != issue.String() { + t.Fatalf("control queue response = %#v, %v", response, err) + } + readContext, cancel := context.WithTimeout(ctx, time.Second) + defer cancel() + encoded, err := transport.Read(readContext) + if err != nil { + t.Fatal(err) + } + envelope, err := agentproto.DecodeEnvelope(encoded, agentproto.DefaultLimits(), rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil || envelope.GetCommandDispatch() == nil || envelope.GetCommandDispatch().GetIssueUuid() != issue.String() { + t.Fatalf("woken dispatch = %#v, %v", envelope, err) + } + _, missingErr := controlService.RunCommand(ctx, &rvboxv1.RunCommandRequest{TargetClientId: "missing", RequestId: fixedIssueID(0x78), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "x"}}}) + if status.Code(missingErr) != codes.NotFound { + t.Fatal("missing target did not return not found") + } +} + +func fixedIssueID(last byte) string { + issue, _ := domain.ParseUUIDv7("019c46f1-1d02-7000-8000-000000000001") + issue[15] = last + return issue.String() +}