50 lines
2.1 KiB
Go
50 lines
2.1 KiB
Go
package agent
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"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"
|
|
)
|
|
|
|
// PersistDispatch admits a server dispatch to the durable spool before any
|
|
// CommandAccepted reply. UUID/hash equality makes a retransmission harmless;
|
|
// a different immutable request for the same UUID is rejected by the spool.
|
|
func PersistDispatch(ctx context.Context, store *spool.Store, session Session, dispatch *rvboxv1.CommandDispatch, now time.Time, limits agentproto.Limits) (spool.Acceptance, error) {
|
|
if store == nil || dispatch == nil || session.Generation == 0 || dispatch.GetTargetSessionGeneration() != session.Generation || now.IsZero() {
|
|
return spool.Acceptance{}, ErrProtocolHandshake
|
|
}
|
|
if err := agentproto.ValidateExecutionSpec(dispatch.GetSpec(), limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil {
|
|
return spool.Acceptance{}, err
|
|
}
|
|
issue, err := domain.ParseUUIDv7(dispatch.GetIssueUuid())
|
|
if err != nil || dispatch.GetCommandRevision() == 0 || len(dispatch.GetImmutableRequestSha256()) != 32 {
|
|
return spool.Acceptance{}, ErrProtocolHandshake
|
|
}
|
|
var immutable [32]byte
|
|
copy(immutable[:], dispatch.GetImmutableRequestSha256())
|
|
if immutable == [32]byte{} {
|
|
return spool.Acceptance{}, errors.New("empty dispatch immutable hash")
|
|
}
|
|
spec, err := proto.MarshalOptions{Deterministic: true}.Marshal(dispatch.GetSpec())
|
|
if err != nil {
|
|
return spool.Acceptance{}, err
|
|
}
|
|
acceptance, err := store.AcceptCommand(ctx, spool.Command{IssueUUID: issue, ImmutableSHA256: immutable, Revision: dispatch.GetCommandRevision(), Phase: uint32(rvboxv1.CommandLifecycle_COMMAND_ACCEPTED), ExecutionSpec: spec}, now)
|
|
if err != nil {
|
|
return spool.Acceptance{}, err
|
|
}
|
|
if descriptor := dispatch.GetSpec().GetScript(); descriptor != nil {
|
|
_, err = store.BeginScript(ctx, issue, spool.ScriptDescriptor{SizeBytes: descriptor.GetSizeBytes(), SHA256: bytesToDigest(descriptor.GetSha256())})
|
|
if err != nil {
|
|
return spool.Acceptance{}, err
|
|
}
|
|
}
|
|
return acceptance, nil
|
|
}
|