feat: persist and transfer script command payloads
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user