Files
rvbox/internal/client/agent/runner.go
T

558 lines
19 KiB
Go

package agent
import (
"context"
"crypto/rand"
"errors"
"fmt"
"math/big"
"time"
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
"github.com/rvbox/rvbox/internal/agentproto"
"github.com/rvbox/rvbox/internal/client/spool"
"github.com/rvbox/rvbox/internal/domain"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/timestamppb"
)
// RunnerOptions contains the replaceable network edge and the durable client
// state used by one daemon. The supervisor is deliberately a callback here:
// process execution can outlive a network session and is owned by the caller.
type RunnerOptions struct {
Store *spool.Store
Dial func(context.Context) (Transport, error)
Hello *rvboxv1.ClientHello
Limits agentproto.Limits
Backoff BackoffOptions
Jitter Jitter
Now func() time.Time
OnDispatch func(context.Context, Session, *rvboxv1.CommandDispatch) error
OnStdin func(context.Context, Session, *rvboxv1.StdinWrite) error
OnCloseStdin func(context.Context, Session, *rvboxv1.CloseStdin) error
OnSignal func(context.Context, Session, *rvboxv1.SignalCommand) error
OnScriptReady func(context.Context, Session, domain.UUID) error
OnTerminate func(context.Context, domain.UUID) error
// EventReady wakes the active session after a supervisor worker appends a
// durable event. The network loop remains the sole writer; a reconnect can
// safely ignore a stale notification because replay reads the spool again.
EventReady <-chan domain.UUID
}
// BackoffOptions and Jitter are kept at the agent boundary so callers do not
// need to depend on the internal session-machine implementation.
type BackoffOptions struct {
Initial time.Duration
Maximum time.Duration
StableReset time.Duration
}
func (options BackoffOptions) Validate() error {
if options.Initial <= 0 || options.Maximum < options.Initial || options.StableReset <= 0 {
return errors.New("invalid client reconnect backoff options")
}
return nil
}
type Jitter func(time.Duration) time.Duration
func (options RunnerOptions) validate() error {
if options.Store == nil || options.Dial == nil || options.Hello == nil {
return errors.New("client runner requires store, dialer, and hello")
}
if options.Limits.MaxEnvelopeBytes == 0 {
options.Limits = agentproto.DefaultLimits()
}
if options.Backoff.Initial == 0 {
options.Backoff = BackoffOptions{Initial: time.Second, Maximum: time.Minute, StableReset: time.Minute}
}
if err := options.Backoff.Validate(); err != nil {
return err
}
if options.Jitter == nil {
return errors.New("client runner jitter is required")
}
if options.Now == nil {
return errors.New("client runner clock is required")
}
return nil
}
// Run reconnects until ctx is cancelled. A failed handshake or transport is
// treated as a normal network failure; durable commands and their spool remain
// untouched. The first failure uses the configured initial backoff and a
// successful session resets the exponential history only after StableReset.
func Run(ctx context.Context, options RunnerOptions) error {
if err := options.validate(); err != nil {
return err
}
delay := time.Duration(0)
failures := uint32(0)
for {
if err := waitUntil(ctx, delay); err != nil {
return nil
}
sessionStarted := options.Now()
_ = runOnce(ctx, options)
if ctx.Err() != nil {
return nil
}
if options.Now().Sub(sessionStarted) >= options.Backoff.StableReset {
failures = 0
}
if failures < ^uint32(0) {
failures++
}
cap := options.Backoff.Initial
for index := uint32(1); index < failures && cap < options.Backoff.Maximum; index++ {
if cap > options.Backoff.Maximum/2 {
cap = options.Backoff.Maximum
break
}
cap *= 2
}
if cap > options.Backoff.Maximum {
cap = options.Backoff.Maximum
}
delay = options.Jitter(cap)
if delay < 0 {
delay = 0
}
if delay > cap {
delay = cap
}
}
}
func waitUntil(ctx context.Context, delay time.Duration) error {
if delay <= 0 {
return nil
}
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}
// RunOnce performs one complete session and returns when the transport is
// lost or a protocol/storage error makes that session unusable.
func RunOnce(ctx context.Context, options RunnerOptions) error {
if err := options.validate(); err != nil {
return err
}
return runOnce(ctx, options)
}
func runOnce(ctx context.Context, options RunnerOptions) (resultErr error) {
transport, err := options.Dial(ctx)
if err != nil {
return err
}
defer func() {
if closeErr := transport.Close(); resultErr == nil && closeErr != nil {
resultErr = closeErr
}
}()
limits := options.Limits
if limits.MaxEnvelopeBytes == 0 {
limits = agentproto.DefaultLimits()
}
session, err := Handshake(ctx, transport, options.Hello, limits)
if err != nil {
return err
}
snapshot, err := options.Store.ReconcileSnapshot(ctx)
if err != nil {
return fmt.Errorf("build client reconciliation snapshot: %w", err)
}
result, err := Reconcile(ctx, transport, session, snapshot, limits)
if err != nil {
return err
}
terminated, err := ApplyReconcileResult(ctx, options.Store, result, options.Now())
if err != nil {
return fmt.Errorf("apply server reconciliation: %w", err)
}
for _, issue := range terminated {
if options.OnTerminate != nil {
if err := options.OnTerminate(ctx, issue); err != nil {
return err
}
}
}
sentEvents := make(map[domain.UUID]uint64)
if err := replayEvents(ctx, transport, options.Store, session, snapshot, limits, sentEvents); err != nil {
return err
}
if err := sendCapacity(ctx, transport, session, options.Hello, snapshot, limits); err != nil {
return err
}
return serveActive(ctx, transport, options, session, sentEvents)
}
func replayEvents(ctx context.Context, transport Transport, store *spool.Store, session Session, snapshot *rvboxv1.ReconcileSnapshot, limits agentproto.Limits, sent map[domain.UUID]uint64) error {
for _, retained := range snapshot.GetRetainedCommands() {
if retained == nil || retained.GetTombstoned() {
continue
}
issue, err := domain.ParseUUIDv7(retained.GetIssueUuid())
if err != nil {
return err
}
if _, err := store.AssignSendWindow(ctx, issue, 128, 1<<20); err != nil {
return err
}
events, err := store.PendingEvents(ctx, issue)
if err != nil {
return err
}
for _, event := range events {
if err := SendStoredEvent(ctx, transport, session, event, limits); err != nil {
return err
}
if event.EventSeq > sent[issue] {
sent[issue] = event.EventSeq
}
}
}
return nil
}
// SendStoredEvent converts the compact local representation back to a
// canonical CommandEvent. Output rows store only their typed OutputChunk to
// avoid duplicating the envelope; metadata rows may store a full event.
func SendStoredEvent(ctx context.Context, transport Transport, session Session, event spool.Event, limits agentproto.Limits) error {
commandEvent, err := commandEventFromStored(event)
if err != nil {
return err
}
commandEvent.IssueUuid = event.IssueUUID.String()
commandEvent.EventSeq = event.EventSeq
commandEvent.ObservedAt = timestamppb.New(event.CreatedAt)
return SendCommandEvent(ctx, transport, session, commandEvent, limits)
}
func commandEventFromStored(event spool.Event) (*rvboxv1.CommandEvent, error) {
if event.EventSeq == 0 || event.IssueUUID == (domain.UUID{}) || event.CreatedAt.IsZero() {
return nil, errors.New("stored event lacks assigned sequence or owner")
}
if event.Kind == spool.EventKindOutput {
var output rvboxv1.OutputChunk
if err := proto.Unmarshal(event.Payload, &output); err != nil {
return nil, fmt.Errorf("decode stored output event: %w", err)
}
return &rvboxv1.CommandEvent{Payload: &rvboxv1.CommandEvent_Output{Output: &output}}, nil
}
if event.Kind == spool.EventKindOutputTruncation {
var marker rvboxv1.OutputTruncation
if err := proto.Unmarshal(event.Payload, &marker); err != nil {
return nil, fmt.Errorf("decode stored truncation event: %w", err)
}
return &rvboxv1.CommandEvent{Payload: &rvboxv1.CommandEvent_OutputTruncation{OutputTruncation: &marker}}, nil
}
var stored rvboxv1.CommandEvent
if err := proto.Unmarshal(event.Payload, &stored); err != nil || stored.Payload == nil {
return nil, fmt.Errorf("decode stored command event: %w", err)
}
return &stored, nil
}
func sendCapacity(ctx context.Context, transport Transport, session Session, hello *rvboxv1.ClientHello, snapshot *rvboxv1.ReconcileSnapshot, limits agentproto.Limits) error {
var running, queued uint32
for _, retained := range snapshot.GetRetainedCommands() {
if retained == nil || retained.GetTombstoned() {
continue
}
switch retained.GetLifecycle() {
case rvboxv1.CommandLifecycle_COMMAND_ACCEPTED, rvboxv1.CommandLifecycle_COMMAND_RUNNING:
running++
case rvboxv1.CommandLifecycle_COMMAND_QUEUED, rvboxv1.CommandLifecycle_COMMAND_DISPATCHED:
queued++
}
}
if running > hello.GetMaxRunningCommands() || queued > hello.GetMaxQueuedCommands() {
return fmt.Errorf("durable client command count exceeds advertised capacity")
}
envelope := &rvboxv1.AgentEnvelope{SessionId: session.ID, SessionGeneration: session.Generation, Payload: &rvboxv1.AgentEnvelope_ClientCapacity{ClientCapacity: &rvboxv1.ClientCapacity{RunningCommands: running, QueuedCommands: queued, MaxRunningCommands: hello.GetMaxRunningCommands(), MaxQueuedCommands: hello.GetMaxQueuedCommands()}}}
if err := agentproto.ValidateEnvelope(envelope, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil {
return err
}
encoded, err := proto.Marshal(envelope)
if err != nil {
return err
}
return transport.Write(ctx, encoded)
}
func serveActive(ctx context.Context, transport Transport, options RunnerOptions, session Session, sent map[domain.UUID]uint64) error {
limits := options.Limits
if limits.MaxEnvelopeBytes == 0 {
limits = agentproto.DefaultLimits()
}
readContext, cancelRead := context.WithCancel(ctx)
defer cancelRead()
frames := make(chan []byte, 1)
readErrors := make(chan error, 1)
go func() {
for {
encoded, err := transport.Read(readContext)
if err != nil {
select {
case readErrors <- err:
case <-readContext.Done():
}
return
}
select {
case frames <- encoded:
case <-readContext.Done():
return
}
}
}()
for {
var encoded []byte
select {
case <-ctx.Done():
return nil
case err := <-readErrors:
return err
case issue := <-options.EventReady:
if issue == (domain.UUID{}) {
continue
}
if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil {
return err
}
continue
case encoded = <-frames:
}
envelope, err := agentproto.DecodeEnvelope(encoded, limits, rvboxv1.Platform_PLATFORM_WINDOWS)
if err != nil {
return err
}
if envelope.GetSessionId() != session.ID || envelope.GetSessionGeneration() != session.Generation {
return ErrProtocolHandshake
}
switch {
case envelope.GetCommandDispatch() != nil:
if err := handleDispatch(ctx, transport, options, session, envelope.GetCommandDispatch(), limits); err != nil {
return err
}
issue, parseErr := domain.ParseUUIDv7(envelope.GetCommandDispatch().GetIssueUuid())
if parseErr == nil {
if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil {
return err
}
}
case envelope.GetEventAck() != nil:
if err := ApplyEventAck(ctx, options.Store, envelope.GetEventAck()); err != nil {
return err
}
case envelope.GetScriptChunk() != nil:
if err := handleScriptChunk(ctx, transport, options.Store, session, envelope.GetScriptChunk(), limits, options.Now, sent); err != nil {
return err
}
case envelope.GetScriptCommit() != nil:
if err := handleScriptCommit(ctx, transport, options.Store, session, envelope.GetScriptCommit(), limits, options.Now, sent); err != nil {
return err
}
if options.OnScriptReady != nil {
issue, err := domain.ParseUUIDv7(envelope.GetScriptCommit().GetIssueUuid())
if err != nil {
return err
}
if err := options.OnScriptReady(ctx, session, issue); err != nil {
return err
}
}
case envelope.GetStdinWrite() != nil:
if options.OnStdin != nil {
if err := options.OnStdin(ctx, session, envelope.GetStdinWrite()); err != nil {
return err
}
}
if issue, parseErr := domain.ParseUUIDv7(envelope.GetStdinWrite().GetIssueUuid()); parseErr == nil {
if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil {
return err
}
}
case envelope.GetCloseStdin() != nil:
if options.OnCloseStdin != nil {
if err := options.OnCloseStdin(ctx, session, envelope.GetCloseStdin()); err != nil {
return err
}
}
if issue, parseErr := domain.ParseUUIDv7(envelope.GetCloseStdin().GetIssueUuid()); parseErr == nil {
if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil {
return err
}
}
case envelope.GetSignalCommand() != nil:
if options.OnSignal != nil {
if err := options.OnSignal(ctx, session, envelope.GetSignalCommand()); err != nil {
return err
}
}
if issue, parseErr := domain.ParseUUIDv7(envelope.GetSignalCommand().GetIssueUuid()); parseErr == nil {
if err := flushEvents(ctx, transport, options.Store, session, issue, limits, sent); err != nil {
return err
}
}
case envelope.GetError() != nil:
if envelope.GetError().GetCloseSession() {
return fmt.Errorf("server closed agent session: %s", envelope.GetError().GetError().GetMessage())
}
case envelope.GetReconcileRequest() != nil:
// A request is advisory after the snapshot/result barrier. The next
// reconnect repeats the complete snapshot; never apply partial targets.
default:
return ErrUnexpectedMessage
}
}
}
func flushEvents(ctx context.Context, transport Transport, store *spool.Store, session Session, issue domain.UUID, limits agentproto.Limits, sent map[domain.UUID]uint64) error {
if issue == (domain.UUID{}) {
return nil
}
if _, err := store.AssignSendWindow(ctx, issue, 128, 1<<20); err != nil {
return err
}
events, err := store.PendingEvents(ctx, issue)
if err != nil {
return err
}
for _, event := range events {
if event.EventSeq == 0 || event.EventSeq <= sent[issue] {
continue
}
if err := SendStoredEvent(ctx, transport, session, event, limits); err != nil {
return err
}
sent[issue] = event.EventSeq
}
return nil
}
func handleDispatch(ctx context.Context, transport Transport, options RunnerOptions, session Session, dispatch *rvboxv1.CommandDispatch, limits agentproto.Limits) error {
acceptance, err := PersistDispatch(ctx, options.Store, session, dispatch, options.Now(), limits)
accepted := err == nil
ack := &rvboxv1.CommandAccepted{IssueUuid: dispatch.GetIssueUuid(), CommandRevision: dispatch.GetCommandRevision(), Accepted: accepted}
if err != nil {
ack.Rejection = rejectionForError(err, dispatch.GetIssueUuid())
}
envelope := &rvboxv1.AgentEnvelope{SessionId: session.ID, SessionGeneration: session.Generation, Payload: &rvboxv1.AgentEnvelope_CommandAccepted{CommandAccepted: ack}}
if validateErr := agentproto.ValidateEnvelope(envelope, limits, rvboxv1.Platform_PLATFORM_WINDOWS); validateErr != nil {
return validateErr
}
encoded, marshalErr := proto.Marshal(envelope)
if marshalErr != nil {
return marshalErr
}
if err := transport.Write(ctx, encoded); err != nil {
return err
}
if accepted && options.OnDispatch != nil {
return options.OnDispatch(ctx, session, dispatch)
}
_ = acceptance
return nil
}
func rejectionForError(err error, issue string) *rvboxv1.ControlError {
code := rvboxv1.ControlError_TRANSIENT
if errors.Is(err, spool.ErrCommandConflict) || errors.Is(err, spool.ErrAlreadyExecuted) {
code = rvboxv1.ControlError_CONFLICT
} else if errors.Is(err, agentproto.ErrInvalidExecutionSpec) || errors.Is(err, ErrProtocolHandshake) {
code = rvboxv1.ControlError_INVALID_ARGUMENT
}
return &rvboxv1.ControlError{Code: code, Message: boundedError(err), Retryable: code == rvboxv1.ControlError_TRANSIENT, IssueUuid: issue}
}
func boundedError(err error) string {
if err == nil {
return ""
}
message := err.Error()
if len(message) > 1024 {
message = message[:1024]
}
return message
}
func handleScriptChunk(ctx context.Context, transport Transport, store *spool.Store, session Session, chunk *rvboxv1.ScriptChunk, limits agentproto.Limits, now func() time.Time, sent map[domain.UUID]uint64) error {
status, err := ApplyScriptChunk(ctx, store, session, chunk, limits)
if err != nil {
return err
}
return appendAndSendScriptStatus(ctx, transport, store, session, chunk.GetIssueUuid(), status, limits, now, sent)
}
func handleScriptCommit(ctx context.Context, transport Transport, store *spool.Store, session Session, commit *rvboxv1.ScriptCommit, limits agentproto.Limits, now func() time.Time, sent map[domain.UUID]uint64) error {
status, err := ApplyScriptCommit(ctx, store, session, commit, limits)
if err != nil {
return err
}
return appendAndSendScriptStatus(ctx, transport, store, session, commit.GetIssueUuid(), status, limits, now, sent)
}
func appendAndSendScriptStatus(ctx context.Context, transport Transport, store *spool.Store, session Session, issueText string, status spool.ScriptStatus, limits agentproto.Limits, now func() time.Time, sent map[domain.UUID]uint64) error {
issue, err := domain.ParseUUIDv7(issueText)
if err != nil {
return err
}
if sent == nil {
sent = make(map[domain.UUID]uint64)
}
if now == nil {
now = func() time.Time { return time.Now().UTC() }
}
observedAt := now()
if observedAt.IsZero() {
return errors.New("script status clock returned zero")
}
event := &rvboxv1.CommandEvent{IssueUuid: issueText, ObservedAt: timestamppb.New(observedAt), Payload: &rvboxv1.CommandEvent_ScriptStatus{ScriptStatus: &rvboxv1.ScriptUploadStatus{ReceivedBytes: status.ReceivedBytes, Complete: status.Committed}}}
payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(event)
if err != nil {
return err
}
if _, err := store.AppendEvent(ctx, issue, spool.EventInput{Kind: 9, Compression: 1, RawBytes: uint64(len(payload)), Payload: payload, CreatedAt: event.GetObservedAt().AsTime()}); err != nil {
return err
}
assigned, err := store.AssignSendWindow(ctx, issue, 1, 1<<20)
if err != nil {
return err
}
for _, item := range assigned {
if item.EventSeq == 0 || item.EventSeq <= sent[issue] {
continue
}
if err := SendStoredEvent(ctx, transport, session, item, limits); err != nil {
return err
}
sent[issue] = item.EventSeq
}
return nil
}
// CryptoJitter returns a full-jitter delay without relying on math/rand's
// process-global state. Tests inject a deterministic jitter instead.
func CryptoJitter(capacity time.Duration) time.Duration {
if capacity <= 0 {
return 0
}
value, err := rand.Int(rand.Reader, big.NewInt(int64(capacity)+1))
if err != nil {
return capacity
}
return time.Duration(value.Int64())
}