package session import ( "context" "crypto/rand" "crypto/sha256" "encoding/base64" "errors" "fmt" "net/http" "sort" "sync" "time" "github.com/coder/websocket" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/server/store" "google.golang.org/protobuf/proto" "google.golang.org/protobuf/types/known/timestamppb" ) const ( defaultAgentPath = "/v1/agent" defaultWriteWait = 10 * time.Second ) var ( ErrUnexpectedOrigin = errors.New("agent connections must not send an Origin header") ErrUnexpectedMessage = errors.New("agent message is not a binary protobuf envelope") ErrStaleSession = errors.New("agent message is for an unknown or fenced session") ErrDispatchDataFull = errors.New("session dispatch data lane is full") ) // AgentServer is the transport edge for agent WebSocket sessions. Command // processing is intentionally injected above this admission/fencing boundary. type AgentServer struct { Store *store.Store Registry *Registry Path string Limits agentproto.Limits SupportedProtocol *rvboxv1.ProtocolRange WriteDeadline time.Duration HeartbeatIdle time.Duration LivenessTimeout time.Duration Now func() time.Time } func (server *AgentServer) ServeHTTP(response http.ResponseWriter, request *http.Request) { if request.URL.Path != server.agentPath() { http.NotFound(response, request) return } if request.Header.Get("Origin") != "" { http.Error(response, ErrUnexpectedOrigin.Error(), http.StatusForbidden) return } if server.Store == nil { http.Error(response, "agent server storage is unavailable", http.StatusServiceUnavailable) return } if server.Registry == nil { http.Error(response, "agent session registry is unavailable", http.StatusServiceUnavailable) return } started := time.Now() heartbeat := newSynchronizedHeartbeat(server.heartbeatIdle(), server.livenessTimeout(), 0) connection, err := websocket.Accept(response, request, &websocket.AcceptOptions{ CompressionMode: websocket.CompressionDisabled, // coder/websocket consumes control frames inside Read. Count ping and // pong callbacks as inbound activity so a healthy, otherwise idle agent // is not mistaken for a dead peer; application Read results are observed // by serveConnection as usual. OnPingReceived: func(context.Context, []byte) bool { heartbeat.Observe(time.Since(started)) return true }, OnPongReceived: func(context.Context, []byte) { heartbeat.Observe(time.Since(started)) }, }) if err != nil { return } defer connection.CloseNow() server.serveConnection(request.Context(), connection, heartbeat, started) } func (server *AgentServer) serveConnection(parent context.Context, connection *websocket.Conn, heartbeat *synchronizedHeartbeat, started time.Time) { messageType, payload, err := connection.Read(parent) if err != nil { return } if messageType != websocket.MessageBinary { server.close(connection, websocket.StatusUnsupportedData, ErrUnexpectedMessage.Error()) return } helloEnvelope, err := agentproto.DecodeEnvelope(payload, server.limits(), rvboxv1.Platform_PLATFORM_UNSPECIFIED) if err != nil || helloEnvelope.GetClientHello() == nil { server.close(connection, websocket.StatusPolicyViolation, "invalid ClientHello") return } hello := helloEnvelope.GetClientHello() // This cutoff precedes registration, which is when control callers can // first observe the live client. Reconciliation must compare only the // state included in its request; later control commands are fresh work. reconcileBoundary := server.now() instanceID, err := domain.ParseUUIDv7(hello.GetClientInstanceId()) if err != nil { server.close(connection, websocket.StatusPolicyViolation, "invalid client instance ID") return } selected, err := domain.SelectProtocol(server.protocolRange(), hello.GetSupportedProtocol()) if err != nil { server.close(connection, websocket.StatusPolicyViolation, "unsupported protocol") return } sessionID, err := randomSessionID() if err != nil { server.close(connection, websocket.StatusInternalError, "could not allocate session") return } shells, err := proto.Marshal(&rvboxv1.ClientHello{SupportedShells: hello.GetSupportedShells()}) if err != nil { server.close(connection, websocket.StatusInternalError, "could not persist client capabilities") return } registration, err := server.Store.RegisterClientSession(parent, store.ClientRegistration{ ClientID: hello.GetClientId(), Platform: uint32(hello.GetPlatform()), Architecture: hello.GetArchitecture(), DaemonVersion: hello.GetDaemonVersion(), DaemonCWD: hello.GetDaemonCwd(), SupportedShells: shells, ClientInstanceID: [16]byte(instanceID), SessionID: sessionID, ConnectedAt: server.now(), }) if err != nil { server.close(connection, websocket.StatusPolicyViolation, sessionCloseReason(err)) return } registry := server.Registry handle, err := registry.Install(hello.GetClientId(), sessionID, registration.Generation) if err != nil { _ = server.Store.CloseLiveSession(context.Background(), sessionID, registration.Generation, "registry installation failed", server.now()) server.close(connection, websocket.StatusInternalError, "could not install session") return } defer func() { registry.Remove(handle) _ = server.Store.CloseLiveSession(context.Background(), sessionID, registration.Generation, "connection closed", server.now()) }() sessionContext, cancel := context.WithCancel(parent) stopFence := context.AfterFunc(handle.Context, cancel) defer func() { stopFence() cancel() }() queue := NewWriterQueue(16, 64, 8) defer queue.Close() capacity := &CapacityShadow{} var capacityMu sync.Mutex reservations := make(map[string]DispatchLane) stdinSent := make(map[string]struct{}) signalSent := make(map[string]struct{}) scriptTransfers := make(map[string]*scriptTransfer) // A dispatch may have been durably claimed immediately before a server or // client reconnect. Reconstruct every still-active script from the command // store before enabling the dispatcher so an uncertain upload is replayed // from its immutable body instead of being silently lost. pendingScripts, err := server.Store.PendingScriptDispatches(parent, hello.GetClientId()) if err != nil { server.close(connection, websocket.StatusInternalError, "could not restore script transfers") return } for _, pending := range pendingScripts { scriptTransfers[pending.IssueUUID.String()] = &scriptTransfer{Body: append([]byte(nil), pending.Body...), Digest: pending.Digest} } 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(), handle.SignalDispatch, reconciled, &capacityMu, capacity, reservations, stdinSent, signalSent, scriptTransfers, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel) encodedSessionID := encodeSessionID(sessionID) welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, Payload: &rvboxv1.AgentEnvelope_ServerWelcome{ServerWelcome: &rvboxv1.ServerWelcome{ SelectedProtocol: selected, ServerTime: timestamppb.New(server.now()), }}, }) if err != nil { server.close(connection, websocket.StatusInternalError, "could not encode welcome") return } if err := queue.EnqueueControl(Frame{Kind: FrameControl, Payload: welcome}); err != nil { server.close(connection, websocket.StatusInternalError, "could not queue welcome") return } targets, err := server.Store.ReconcileTargetsAt(parent, hello.GetClientId(), reconcileBoundary) if err != nil { server.close(connection, websocket.StatusInternalError, "could not build reconciliation request") return } reconcileRequest, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, Payload: &rvboxv1.AgentEnvelope_ReconcileRequest{ReconcileRequest: &rvboxv1.ReconcileRequest{Targets: targets}}, }) if err != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: reconcileRequest}) != nil { server.close(connection, websocket.StatusInternalError, "could not queue reconciliation request") return } writerDone := make(chan struct{}) go func() { defer close(writerDone) server.writeLoop(sessionContext, connection, queue, heartbeat, started) cancel() }() defer func() { <-writerDone }() for { messageType, payload, err = connection.Read(sessionContext) if err != nil { return } // Elapsed time is monotonic; peer-provided wall timestamps are never used // to determine liveness. heartbeat.Observe(time.Since(started)) if messageType != websocket.MessageBinary { server.close(connection, websocket.StatusUnsupportedData, ErrUnexpectedMessage.Error()) return } envelope, err := agentproto.DecodeEnvelope(payload, server.limits(), hello.GetPlatform()) if err != nil || envelope.GetClientHello() != nil || envelope.GetSessionId() != encodedSessionID || envelope.GetSessionGeneration() != registration.Generation { server.close(connection, websocket.StatusPolicyViolation, ErrStaleSession.Error()) return } clientID, err := server.Store.ValidateLiveSession(sessionContext, sessionID, registration.Generation) if err != nil || clientID != hello.GetClientId() { server.close(connection, websocket.StatusPolicyViolation, ErrStaleSession.Error()) return } if snapshot := envelope.GetReconcileSnapshot(); snapshot != nil { result, reconcileErr := server.Store.ReconcileClientSnapshotForSessionAt(sessionContext, hello.GetClientId(), registration.Generation, snapshot, reconcileBoundary) if reconcileErr != nil { server.close(connection, websocket.StatusPolicyViolation, "reconciliation failed") return } encoded, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, Payload: &rvboxv1.AgentEnvelope_ReconcileResult{ReconcileResult: result}, }) resultWritten := make(chan struct{}) if err != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded, Written: resultWritten}) != nil { server.close(connection, websocket.StatusInternalError, "could not queue reconciliation result") return } select { case <-resultWritten: case <-sessionContext.Done(): return } reconcileOnce.Do(func() { close(reconciled) }) handle.SignalDispatch() continue } if advertised := envelope.GetClientCapacity(); advertised != nil { 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 { issue, parseErr := domain.ParseUUIDv7(acknowledgement.GetIssueUuid()) if parseErr != nil { server.close(connection, websocket.StatusPolicyViolation, "invalid command acceptance") return } if _, acceptErr := server.Store.RecordCommandAcceptanceWithRejection(sessionContext, issue, hello.GetClientId(), registration.Generation, acknowledgement.GetCommandRevision(), acknowledgement.GetAccepted(), acknowledgement.GetRejection(), server.now()); acceptErr != nil { 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 { var scriptIssue domain.UUID if scriptStatus := event.GetScriptStatus(); scriptStatus != nil { var parseErr error scriptIssue, parseErr = domain.ParseUUIDv7(event.GetIssueUuid()) if parseErr != nil { server.close(connection, websocket.StatusPolicyViolation, "invalid script status") return } capacityMu.Lock() transfer := scriptTransfers[scriptIssue.String()] validStatus := transfer != nil && scriptStatus.GetReceivedBytes() <= uint64(len(transfer.Body)) && scriptStatus.GetReceivedBytes() <= transfer.NextOffset if validStatus && scriptStatus.GetComplete() { validStatus = scriptStatus.GetReceivedBytes() == uint64(len(transfer.Body)) && transfer.CommitQueued } capacityMu.Unlock() if !validStatus { server.close(connection, websocket.StatusPolicyViolation, "invalid script progress") return } } appendEvent, eventErr := eventAppendFromWire(event, hello.GetClientId(), registration.Generation, server.now()) if eventErr != nil { server.close(connection, websocket.StatusPolicyViolation, "invalid command event") return } appended, eventErr := server.Store.AppendCommandEvent(sessionContext, appendEvent) if eventErr != nil { server.close(connection, websocket.StatusPolicyViolation, "command event was not accepted") return } ack, eventErr := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: encodedSessionID, SessionGeneration: registration.Generation, Payload: &rvboxv1.AgentEnvelope_EventAck{EventAck: &rvboxv1.EventAck{IssueUuid: event.GetIssueUuid(), ThroughEventSeq: appended.ThroughEventSeq}}}) if eventErr != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: ack}) != nil { server.close(connection, websocket.StatusInternalError, "could not acknowledge command event") return } if stdinAck := event.GetStdinAck(); stdinAck != nil { issue, _ := domain.ParseUUIDv7(event.GetIssueUuid()) capacityMu.Lock() delete(stdinSent, stdinIntentKey(issue, stdinAck.GetWriteSeq())) capacityMu.Unlock() handle.SignalDispatch() } if signalResult := event.GetSignalResult(); signalResult != nil { issue, _ := domain.ParseUUIDv7(event.GetIssueUuid()) capacityMu.Lock() delete(signalSent, signalIntentKey(issue, signalResult.GetCommandRevision(), signalResult.GetSignal())) capacityMu.Unlock() handle.SignalDispatch() } if scriptStatus := event.GetScriptStatus(); scriptStatus != nil { capacityMu.Lock() transfer := scriptTransfers[scriptIssue.String()] if scriptStatus.GetReceivedBytes() > transfer.AcknowledgedOffset { transfer.AcknowledgedOffset = scriptStatus.GetReceivedBytes() } if scriptStatus.GetComplete() { transfer.Completed = true } capacityMu.Unlock() handle.SignalDispatch() } } } } func eventAppendFromWire(event *rvboxv1.CommandEvent, clientID string, generation uint64, receipt time.Time) (store.EventAppend, error) { issue, err := domain.ParseUUIDv7(event.GetIssueUuid()) if err != nil { return store.EventAppend{}, err } payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(event) if err != nil { return store.EventAppend{}, err } var immutable [32]byte copy(immutable[:], event.GetImmutableEventSha256()) result := store.EventAppend{IssueUUID: [16]byte(issue), ClientID: clientID, SessionGeneration: generation, EventSeq: event.GetEventSeq(), ObservedUnixNano: event.GetObservedAt().AsTime().UnixNano(), ReceiptUnixNano: receipt.UnixNano(), EventType: eventType(event), Compression: 1, RawLength: uint64(len(payload)), Payload: payload, ImmutableSHA256: immutable} if lifecycle := event.GetLifecycle(); lifecycle != nil { value := lifecycle.GetLifecycle() result.Lifecycle, result.LifecycleRevision = &value, lifecycle.GetCommandRevision() result.WindowsIdentity = lifecycle.GetWindowsExecutionIdentity() } if output := event.GetOutput(); output != nil { result.Stream = uint16(output.GetStream()) result.Output = true } if result.EventType == 0 { return store.EventAppend{}, errors.New("unsupported command event") } return result, nil } func eventType(event *rvboxv1.CommandEvent) uint16 { switch event.Payload.(type) { case *rvboxv1.CommandEvent_Lifecycle: return 4 case *rvboxv1.CommandEvent_Output: return 5 case *rvboxv1.CommandEvent_Resource: return 6 case *rvboxv1.CommandEvent_StdinAck: return 7 case *rvboxv1.CommandEvent_SignalResult: return 8 case *rvboxv1.CommandEvent_ScriptStatus: return 9 case *rvboxv1.CommandEvent_OutputTruncation: return 10 case *rvboxv1.CommandEvent_OutputIncomplete: return 11 default: return 0 } } // 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, beforeEnqueue func(*store.DispatchCandidate), onWritten func(domain.UUID)) (domain.UUID, bool, error) { candidate, err := server.Store.ClaimNextDispatch(ctx, clientID, generation, server.now()) if err != nil || candidate == nil { return domain.UUID{}, false, err } requeue := func(cause error) (bool, error) { _, rollbackErr := server.Store.RequeueDispatch(ctx, candidate.IssueUUID, clientID, generation) if rollbackErr != nil { return false, rollbackErr } return false, cause } spec := &rvboxv1.ExecutionSpec{} if err := proto.Unmarshal(candidate.ExecutionSpec, spec); err != nil { 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 { 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, IssueTime: timestamppb.New(candidate.IssueTime), Spec: spec, ImmutableRequestSha256: candidate.ImmutableSHA256[:], } if candidate.QueueExpiryTime != nil { dispatch.QueueExpiryTime = timestamppb.New(*candidate.QueueExpiryTime) } encoded, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_CommandDispatch{CommandDispatch: dispatch}, }) if err != nil { sent, requeueErr := requeue(err) return candidate.IssueUUID, sent, requeueErr } if beforeEnqueue != nil { beforeEnqueue(candidate) } if !queue.EnqueueData(Frame{Kind: FrameData, Payload: encoded, OnWritten: func() { if onWritten != nil { onWritten(candidate.IssueUUID) } }}) { sent, requeueErr := requeue(ErrDispatchDataFull) return candidate.IssueUUID, sent, requeueErr } return candidate.IssueUUID, true, nil } const scriptSendWindowBytes uint64 = 1 << 20 type scriptTransfer struct { Body []byte Digest [sha256.Size]byte NextOffset uint64 AcknowledgedOffset uint64 CommitQueued bool Completed bool } // enqueueNextScript sends one bounded script frame. The dispatcher never // queues more than a 1 MiB unacknowledged window, and the data lane remains // bounded; a reconnect simply starts the exact payload from offset zero. func (server *AgentServer) enqueueNextScript(queue *WriterQueue, sessionID string, generation uint64, transfers map[string]*scriptTransfer, sentMu *sync.Mutex, chunkLimit uint64) (bool, error) { if chunkLimit == 0 { return false, errors.New("script chunk limit is zero") } sentMu.Lock() keys := make([]string, 0, len(transfers)) for key := range transfers { keys = append(keys, key) } sort.Strings(keys) for _, key := range keys { transfer := transfers[key] if transfer == nil || transfer.Completed || transfer.NextOffset-transfer.AcknowledgedOffset >= scriptSendWindowBytes { continue } var envelope *rvboxv1.AgentEnvelope if transfer.NextOffset < uint64(len(transfer.Body)) { end := transfer.NextOffset + chunkLimit if end > uint64(len(transfer.Body)) { end = uint64(len(transfer.Body)) } chunk := transfer.Body[transfer.NextOffset:end] digest := sha256.Sum256(chunk) envelope = &rvboxv1.AgentEnvelope{SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_ScriptChunk{ScriptChunk: &rvboxv1.ScriptChunk{IssueUuid: key, Offset: transfer.NextOffset, Data: append([]byte(nil), chunk...), Sha256: digest[:]}}} } else if !transfer.CommitQueued { transfer.CommitQueued = true envelope = &rvboxv1.AgentEnvelope{SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_ScriptCommit{ScriptCommit: &rvboxv1.ScriptCommit{IssueUuid: key, SizeBytes: uint64(len(transfer.Body)), Sha256: transfer.Digest[:]}}} } else { continue } encoded, err := proto.Marshal(envelope) if err != nil { if transfer.CommitQueued && transfer.NextOffset == uint64(len(transfer.Body)) { transfer.CommitQueued = false } sentMu.Unlock() return false, err } if !queue.EnqueueData(Frame{Kind: FrameData, Payload: encoded}) { if transfer.CommitQueued && transfer.NextOffset == uint64(len(transfer.Body)) { transfer.CommitQueued = false } sentMu.Unlock() return false, ErrDispatchDataFull } if transfer.NextOffset < uint64(len(transfer.Body)) { transfer.NextOffset += uint64(len(envelope.GetScriptChunk().GetData())) } sentMu.Unlock() return true, nil } sentMu.Unlock() return false, nil } // enqueueNextStdin exposes one durable input intent on the essential control // lane. Session-local sent tracking suppresses duplicate frames while a live // connection remains usable; reconnecting naturally replays unacknowledged // writes from storage. func (server *AgentServer) enqueueNextStdin(ctx context.Context, queue *WriterQueue, clientID, sessionID string, generation uint64, sent map[string]struct{}, dispatchWritten map[string]bool, sentMu *sync.Mutex) (bool, error) { intents, err := server.Store.PendingStdin(ctx, clientID) if err != nil { return false, err } for _, intent := range intents { sentMu.Lock() written, waitingForDispatch := dispatchWritten[intent.IssueUUID.String()] sentMu.Unlock() if waitingForDispatch && !written { continue } key := stdinIntentKey(intent.IssueUUID, intent.WriteSeq) sentMu.Lock() _, alreadySent := sent[key] if !alreadySent { sent[key] = struct{}{} } sentMu.Unlock() if alreadySent { continue } envelope := &rvboxv1.AgentEnvelope{SessionId: sessionID, SessionGeneration: generation} if intent.Close { envelope.Payload = &rvboxv1.AgentEnvelope_CloseStdin{CloseStdin: &rvboxv1.CloseStdin{IssueUuid: intent.IssueUUID.String(), WriteSeq: intent.WriteSeq}} } else { envelope.Payload = &rvboxv1.AgentEnvelope_StdinWrite{StdinWrite: &rvboxv1.StdinWrite{IssueUuid: intent.IssueUUID.String(), WriteSeq: intent.WriteSeq, Data: append([]byte(nil), intent.Data...), AppendNewline: intent.AppendNewline}} } encoded, marshalErr := proto.Marshal(envelope) if marshalErr != nil { sentMu.Lock() delete(sent, key) sentMu.Unlock() return false, marshalErr } if enqueueErr := queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded}); enqueueErr != nil { sentMu.Lock() delete(sent, key) sentMu.Unlock() return false, enqueueErr } return true, nil } return false, nil } func stdinIntentKey(issue domain.UUID, writeSeq uint64) string { return issue.String() + ":" + fmt.Sprint(writeSeq) } func signalIntentKey(issue domain.UUID, revision uint64, signal rvboxv1.SignalKind) string { return issue.String() + ":" + fmt.Sprint(revision) + ":" + fmt.Sprint(signal) } func (server *AgentServer) enqueueNextSignal(ctx context.Context, queue *WriterQueue, clientID, sessionID string, generation uint64, sent map[string]struct{}, dispatchWritten map[string]bool, sentMu *sync.Mutex) (bool, error) { intents, err := server.Store.PendingSignals(ctx, clientID, generation) if err != nil { return false, err } for _, intent := range intents { sentMu.Lock() written, waitingForDispatch := dispatchWritten[intent.IssueUUID.String()] sentMu.Unlock() if waitingForDispatch && !written { continue } key := signalIntentKey(intent.IssueUUID, intent.CommandRevision, intent.Signal) sentMu.Lock() _, alreadySent := sent[key] if !alreadySent { sent[key] = struct{}{} } sentMu.Unlock() if alreadySent { continue } envelope := &rvboxv1.AgentEnvelope{SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_SignalCommand{SignalCommand: &rvboxv1.SignalCommand{IssueUuid: intent.IssueUUID.String(), CommandRevision: intent.CommandRevision, Signal: intent.Signal}}} encoded, marshalErr := proto.Marshal(envelope) if marshalErr != nil { sentMu.Lock() delete(sent, key) sentMu.Unlock() return false, marshalErr } if enqueueErr := queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded}); enqueueErr != nil { sentMu.Lock() delete(sent, key) sentMu.Unlock() return false, enqueueErr } return true, nil } return false, 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{}, signalWake func(), reconciled <-chan struct{}, capacityMu *sync.Mutex, capacity *CapacityShadow, reservations map[string]DispatchLane, stdinSent, signalSent map[string]struct{}, scriptTransfers map[string]*scriptTransfer, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) { // A control frame may use the essential queue and therefore overtake data. // Keep its issue blocked only until the dispatch frame has actually crossed // the socket writer; after that, WebSocket ordering preserves the dependency. dispatchWritten := make(map[string]bool) select { case <-reconciled: case <-ctx.Done(): return } for { for { stdinQueued, stdinErr := server.enqueueNextStdin(ctx, queue, clientID, sessionID, generation, stdinSent, dispatchWritten, capacityMu) if stdinErr != nil { server.closeForDispatchFailure(cancel) return } if stdinQueued { continue } signalQueued, signalErr := server.enqueueNextSignal(ctx, queue, clientID, sessionID, generation, signalSent, dispatchWritten, capacityMu) if signalErr != nil { server.closeForDispatchFailure(cancel) return } if signalQueued { continue } scriptQueued, scriptErr := server.enqueueNextScript(queue, sessionID, generation, scriptTransfers, capacityMu, server.limits().MaxRawChunkBytes) if scriptErr != nil && !errors.Is(scriptErr, ErrDispatchDataFull) { server.closeForDispatchFailure(cancel) return } if scriptQueued { continue } 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(candidate *store.DispatchCandidate) { capacityMu.Lock() reservations[candidate.IssueUUID.String()] = lane dispatchWritten[candidate.IssueUUID.String()] = false if candidate.ScriptPresent { scriptTransfers[candidate.IssueUUID.String()] = &scriptTransfer{Body: append([]byte(nil), candidate.ScriptContent...), Digest: sha256.Sum256(candidate.ScriptContent)} } inserted = true capacityMu.Unlock() }, func(issue domain.UUID) { capacityMu.Lock() dispatchWritten[issue.String()] = true capacityMu.Unlock() if signalWake != nil { signalWake() } }) if !sent { capacityMu.Lock() if inserted { delete(reservations, issue.String()) delete(dispatchWritten, issue.String()) delete(scriptTransfers, 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() } } func (server *AgentServer) writeLoop(ctx context.Context, connection *websocket.Conn, queue *WriterQueue, heartbeat *synchronizedHeartbeat, started time.Time) { for { frameContext, cancel := context.WithTimeout(ctx, server.heartbeatPollInterval()) frame, err := queue.Next(frameContext) cancel() if err == nil { writeContext, writeCancel := context.WithTimeout(ctx, server.writeDeadline()) err = connection.Write(writeContext, websocket.MessageBinary, frame.Payload) writeCancel() if err != nil { return } if frame.Written != nil { close(frame.Written) } if frame.OnWritten != nil { frame.OnWritten() } continue } if !errors.Is(err, context.DeadlineExceeded) { return } switch heartbeat.Check(time.Since(started)) { case HeartbeatPing: pingContext, pingCancel := context.WithTimeout(ctx, server.writeDeadline()) err = connection.Ping(pingContext) pingCancel() if err != nil { return } case HeartbeatClose: return } } } func (server *AgentServer) close(connection *websocket.Conn, status websocket.StatusCode, reason string) { _ = connection.Close(status, reason) } func (server *AgentServer) agentPath() string { if server.Path == "" { return defaultAgentPath } return server.Path } func (server *AgentServer) limits() agentproto.Limits { if server.Limits.MaxEnvelopeBytes == 0 { return agentproto.DefaultLimits() } return server.Limits } func (server *AgentServer) protocolRange() *rvboxv1.ProtocolRange { if server.SupportedProtocol == nil { return &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0} } return server.SupportedProtocol } func (server *AgentServer) writeDeadline() time.Duration { if server.WriteDeadline <= 0 { return defaultWriteWait } return server.WriteDeadline } func (server *AgentServer) heartbeatIdle() time.Duration { if server.HeartbeatIdle <= 0 { return 10 * time.Second } return server.HeartbeatIdle } func (server *AgentServer) livenessTimeout() time.Duration { if server.LivenessTimeout <= server.heartbeatIdle() { return 30 * time.Second } return server.LivenessTimeout } func (server *AgentServer) heartbeatPollInterval() time.Duration { interval := server.heartbeatIdle() / 2 if interval <= 0 || interval > server.writeDeadline() { return server.writeDeadline() } return interval } func (server *AgentServer) now() time.Time { if server.Now != nil { return server.Now().UTC() } return time.Now().UTC() } func randomSessionID() ([16]byte, error) { var result [16]byte if _, err := rand.Read(result[:]); err != nil { return result, fmt.Errorf("read session entropy: %w", err) } return result, nil } func encodeSessionID(value [16]byte) string { return base64.RawURLEncoding.EncodeToString(value[:]) } func sessionCloseReason(err error) string { if errors.Is(err, store.ErrTakeoverRequired) { return "client takeover authorization required" } return "client registration rejected" } type synchronizedHeartbeat struct { mu sync.Mutex heartbeat *Heartbeat } func newSynchronizedHeartbeat(idle, timeout, now time.Duration) *synchronizedHeartbeat { return &synchronizedHeartbeat{heartbeat: NewHeartbeat(idle, timeout, now)} } func (heartbeat *synchronizedHeartbeat) Observe(now time.Duration) { heartbeat.mu.Lock() defer heartbeat.mu.Unlock() heartbeat.heartbeat.ObserveInbound(now) } func (heartbeat *synchronizedHeartbeat) Check(now time.Duration) HeartbeatAction { heartbeat.mu.Lock() defer heartbeat.mu.Unlock() return heartbeat.heartbeat.Check(now) }