815 lines
30 KiB
Go
815 lines
30 KiB
Go
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
|
|
}
|
|
connection, err := websocket.Accept(response, request, &websocket.AcceptOptions{
|
|
CompressionMode: websocket.CompressionDisabled,
|
|
})
|
|
if err != nil {
|
|
return
|
|
}
|
|
defer connection.CloseNow()
|
|
server.serveConnection(request.Context(), connection)
|
|
}
|
|
|
|
func (server *AgentServer) serveConnection(parent context.Context, connection *websocket.Conn) {
|
|
started := time.Now()
|
|
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()
|
|
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(), 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.ReconcileTargets(parent, hello.GetClientId())
|
|
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
|
|
}
|
|
|
|
heartbeat := newSynchronizedHeartbeat(server.heartbeatIdle(), server.livenessTimeout(), 0)
|
|
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.ReconcileClientSnapshotForSession(sessionContext, hello.GetClientId(), registration.Generation, snapshot)
|
|
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.RecordCommandAcceptance(sessionContext, issue, hello.GetClientId(), registration.Generation, acknowledgement.GetCommandRevision(), acknowledgement.GetAccepted(), 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()
|
|
}
|
|
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)) (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}) {
|
|
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{}, sentMu *sync.Mutex) (bool, error) {
|
|
intents, err := server.Store.PendingStdin(ctx, clientID)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
for _, intent := range intents {
|
|
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{}, sentMu *sync.Mutex) (bool, error) {
|
|
intents, err := server.Store.PendingSignals(ctx, clientID, generation)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
for _, intent := range intents {
|
|
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{}, 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) {
|
|
select {
|
|
case <-reconciled:
|
|
case <-ctx.Done():
|
|
return
|
|
}
|
|
for {
|
|
for {
|
|
stdinQueued, stdinErr := server.enqueueNextStdin(ctx, queue, clientID, sessionID, generation, stdinSent, capacityMu)
|
|
if stdinErr != nil {
|
|
server.closeForDispatchFailure(cancel)
|
|
return
|
|
}
|
|
if stdinQueued {
|
|
continue
|
|
}
|
|
signalQueued, signalErr := server.enqueueNextSignal(ctx, queue, clientID, sessionID, generation, signalSent, 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
|
|
if candidate.ScriptPresent {
|
|
scriptTransfers[candidate.IssueUUID.String()] = &scriptTransfer{Body: append([]byte(nil), candidate.ScriptContent...), Digest: sha256.Sum256(candidate.ScriptContent)}
|
|
}
|
|
inserted = true
|
|
capacityMu.Unlock()
|
|
})
|
|
if !sent {
|
|
capacityMu.Lock()
|
|
if inserted {
|
|
delete(reservations, 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)
|
|
}
|
|
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
|
|
}
|
|
heartbeat.Observe(time.Since(started))
|
|
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)
|
|
}
|