feat: persist and transfer script command payloads

This commit is contained in:
2026-09-06 11:11:36 +00:00
parent 26ac4c1cac
commit 3f84d3b2f1
18 changed files with 1155 additions and 31 deletions
+134 -6
View File
@@ -3,10 +3,12 @@ package session
import (
"context"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"errors"
"fmt"
"net/http"
"sort"
"sync"
"time"
@@ -143,13 +145,26 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w
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, hello.GetClientId(), hello.GetPlatform(), encodeSessionID(sessionID), registration.Generation, cancel)
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,
@@ -252,6 +267,26 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w
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")
@@ -281,6 +316,18 @@ func (server *AgentServer) serveConnection(parent context.Context, connection *w
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()
}
}
}
}
@@ -337,7 +384,7 @@ 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, beforeEnqueue func(domain.UUID)) (domain.UUID, bool, error) {
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
@@ -373,7 +420,7 @@ func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *Write
return candidate.IssueUUID, sent, requeueErr
}
if beforeEnqueue != nil {
beforeEnqueue(candidate.IssueUUID)
beforeEnqueue(candidate)
}
if !queue.EnqueueData(Frame{Kind: FrameData, Payload: encoded}) {
sent, requeueErr := requeue(ErrDispatchDataFull)
@@ -382,6 +429,75 @@ func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *Write
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
@@ -472,7 +588,7 @@ func (server *AgentServer) enqueueNextSignal(ctx context.Context, queue *WriterQ
// 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{}, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64, cancel context.CancelFunc) {
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():
@@ -496,6 +612,14 @@ func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue,
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()
@@ -503,9 +627,12 @@ func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue,
break
}
inserted := false
issue, sent, err := server.enqueueNextDispatch(ctx, queue, clientID, platform, sessionID, generation, func(issue domain.UUID) {
issue, sent, err := server.enqueueNextDispatch(ctx, queue, clientID, platform, sessionID, generation, func(candidate *store.DispatchCandidate) {
capacityMu.Lock()
reservations[issue.String()] = lane
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()
})
@@ -513,6 +640,7 @@ func (server *AgentServer) dispatchLoop(ctx context.Context, queue *WriterQueue,
capacityMu.Lock()
if inserted {
delete(reservations, issue.String())
delete(scriptTransfers, issue.String())
}
capacity.Release(lane)
capacityMu.Unlock()
@@ -1,13 +1,16 @@
package session
import (
"bytes"
"context"
"crypto/sha256"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
@@ -20,6 +23,54 @@ import (
"google.golang.org/protobuf/types/known/timestamppb"
)
func TestScriptTransferUsesChecksummedChunksAndCommit_HP_SCRIPT_03(t *testing.T) {
t.Parallel()
body := bytes.Repeat([]byte("abcd"), 8)
digest := sha256.Sum256(body)
transfer := map[string]*scriptTransfer{"019c46f1-1d02-7000-8000-0000000000b1": {Body: body, Digest: digest}}
queue := NewWriterQueue(4, 8, 2)
defer queue.Close()
for offset := uint64(0); offset < uint64(len(body)); {
queued, err := (&AgentServer{}).enqueueNextScript(queue, "session", 7, transfer, &sync.Mutex{}, 5)
if err != nil || !queued {
t.Fatalf("chunk at offset %d: queued=%t err=%v", offset, queued, err)
}
frame, err := queue.Next(context.Background())
if err != nil {
t.Fatal(err)
}
var envelope rvboxv1.AgentEnvelope
if err := proto.Unmarshal(frame.Payload, &envelope); err != nil {
t.Fatal(err)
}
chunk := envelope.GetScriptChunk()
if chunk == nil || chunk.GetOffset() != offset || !bytes.Equal(chunk.GetData(), body[offset:offset+uint64(len(chunk.GetData()))]) {
t.Fatalf("chunk = %+v at offset %d", chunk, offset)
}
chunkDigest := sha256.Sum256(chunk.GetData())
if !bytes.Equal(chunk.GetSha256(), chunkDigest[:]) {
t.Fatal("chunk digest mismatch")
}
offset += uint64(len(chunk.GetData()))
}
queued, err := (&AgentServer{}).enqueueNextScript(queue, "session", 7, transfer, &sync.Mutex{}, 5)
if err != nil || !queued {
t.Fatalf("commit queued=%t err=%v", queued, err)
}
frame, err := queue.Next(context.Background())
if err != nil {
t.Fatal(err)
}
var envelope rvboxv1.AgentEnvelope
if err := proto.Unmarshal(frame.Payload, &envelope); err != nil {
t.Fatal(err)
}
commit := envelope.GetScriptCommit()
if commit == nil || commit.GetSizeBytes() != uint64(len(body)) || !bytes.Equal(commit.GetSha256(), digest[:]) {
t.Fatalf("commit = %+v", commit)
}
}
func TestAgentServerRegistrationAndReplacement_HP_SES_05(t *testing.T) {
server, cleanup := newTestAgentServer(t)
defer cleanup()