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 // OnSessionError observes one failed dial, handshake, protocol, or active // transport session before normal reconnect backoff. It must not block; the // durable spool and retry policy remain owned by Run. OnSessionError func(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() sessionErr := runOnce(ctx, options) if ctx.Err() != nil { return nil } if sessionErr != nil && options.OnSessionError != nil { options.OnSessionError(sessionErr) } 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 fmt.Errorf("handle command dispatch %s: %w", envelope.GetCommandDispatch().GetIssueUuid(), 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 fmt.Errorf("flush dispatched command %s events: %w", issue, err) } } case envelope.GetEventAck() != nil: if err := ApplyEventAck(ctx, options.Store, envelope.GetEventAck()); err != nil { return fmt.Errorf("apply event acknowledgement for %s: %w", envelope.GetEventAck().GetIssueUuid(), 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()) }