210 lines
11 KiB
Go
210 lines
11 KiB
Go
package agent
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/sha256"
|
|
"errors"
|
|
"path/filepath"
|
|
"testing"
|
|
"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"
|
|
)
|
|
|
|
func TestHandshakeHelloWelcome_HP_SES_07(t *testing.T) {
|
|
t.Parallel()
|
|
welcome := &rvboxv1.AgentEnvelope{SessionId: "issued-session", SessionGeneration: 7, Payload: &rvboxv1.AgentEnvelope_ServerWelcome{ServerWelcome: &rvboxv1.ServerWelcome{SelectedProtocol: &rvboxv1.ProtocolVersion{Major: 1, Minor: 0}, ServerTime: timestamppb.New(time.Date(2026, time.September, 6, 0, 0, 0, 0, time.UTC))}}}
|
|
encoded, err := proto.Marshal(welcome)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
transport := &fakeTransport{read: encoded}
|
|
session, err := Handshake(context.Background(), transport, validHello(), agentproto.DefaultLimits())
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if session.ID != "issued-session" || session.Generation != 7 || session.Protocol.GetMajor() != 1 {
|
|
t.Fatalf("session = %#v", session)
|
|
}
|
|
var sent rvboxv1.AgentEnvelope
|
|
if err := proto.Unmarshal(transport.written, &sent); err != nil || sent.GetClientHello() == nil || sent.GetSessionId() != "" || sent.GetSessionGeneration() != 0 {
|
|
t.Fatalf("sent ClientHello valid=%t session=%q generation=%d err=%v", sent.GetClientHello() != nil, sent.GetSessionId(), sent.GetSessionGeneration(), err)
|
|
}
|
|
}
|
|
|
|
func TestHandshakeRejectsNonWelcome_BH_SES_07(t *testing.T) {
|
|
t.Parallel()
|
|
bad := &rvboxv1.AgentEnvelope{SessionId: "issued-session", SessionGeneration: 7, Payload: &rvboxv1.AgentEnvelope_ClientCapacity{ClientCapacity: &rvboxv1.ClientCapacity{MaxRunningCommands: 1, MaxQueuedCommands: 1}}}
|
|
encoded, err := proto.Marshal(bad)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
_, err = Handshake(context.Background(), &fakeTransport{read: encoded}, validHello(), agentproto.DefaultLimits())
|
|
if !errors.Is(err, ErrProtocolHandshake) {
|
|
t.Fatalf("non-welcome handshake error = %v, want ErrProtocolHandshake", err)
|
|
}
|
|
}
|
|
|
|
func TestReconcileSnapshot_HP_SES_08(t *testing.T) {
|
|
t.Parallel()
|
|
result := &rvboxv1.AgentEnvelope{SessionId: "issued-session", SessionGeneration: 7, Payload: &rvboxv1.AgentEnvelope_ReconcileResult{ReconcileResult: &rvboxv1.ReconcileResult{}}}
|
|
encoded, err := proto.Marshal(result)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
transport := &fakeTransport{read: encoded}
|
|
reconciled, err := Reconcile(context.Background(), transport, Session{ID: "issued-session", Generation: 7}, &rvboxv1.ReconcileSnapshot{}, agentproto.DefaultLimits())
|
|
if err != nil || reconciled == nil {
|
|
t.Fatalf("Reconcile = %#v, %v", reconciled, err)
|
|
}
|
|
var sent rvboxv1.AgentEnvelope
|
|
if err := proto.Unmarshal(transport.written, &sent); err != nil || sent.GetReconcileSnapshot() == nil || sent.GetSessionId() != "issued-session" || sent.GetSessionGeneration() != 7 {
|
|
t.Fatalf("sent snapshot valid=%t session=%q generation=%d err=%v", sent.GetReconcileSnapshot() != nil, sent.GetSessionId(), sent.GetSessionGeneration(), err)
|
|
}
|
|
}
|
|
|
|
func TestApplyReconcileResult_HP_SES_12(t *testing.T) {
|
|
t.Parallel()
|
|
active := "019c46f1-1d02-7000-8000-000000000062"
|
|
terminal := "019c46f1-1d02-7000-8000-000000000063"
|
|
store := &reconcileStore{}
|
|
terminated, err := ApplyReconcileResult(context.Background(), store, &rvboxv1.ReconcileResult{
|
|
TerminateLocalIssueUuids: []string{active},
|
|
DiscardLocalTerminalIssueUuids: []string{terminal},
|
|
}, time.Date(2026, time.September, 6, 0, 0, 0, 0, time.UTC))
|
|
if err != nil || len(terminated) != 1 || terminated[0].String() != active || len(store.discarded) != 1 || store.discarded[0].String() != terminal {
|
|
t.Fatalf("ApplyReconcileResult = terminated=%v discarded=%v err=%v", terminated, store.discarded, err)
|
|
}
|
|
if _, err := ApplyReconcileResult(context.Background(), store, &rvboxv1.ReconcileResult{TerminateLocalIssueUuids: []string{active}, DiscardLocalTerminalIssueUuids: []string{active}}, time.Now()); !errors.Is(err, ErrInvalidReconcileResult) {
|
|
t.Fatalf("overlapping result error = %v", err)
|
|
}
|
|
}
|
|
|
|
func TestPersistDispatchUsesImmutableHash_HP_DISPATCH_04(t *testing.T) {
|
|
store, err := spool.Open(context.Background(), spool.Options{DataDir: filepath.Join(t.TempDir(), "spool"), BusyTimeout: time.Second})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer store.Close()
|
|
dispatch := &rvboxv1.CommandDispatch{IssueUuid: "019c46f1-1d02-7000-8000-000000000064", CommandRevision: 1, TargetSessionGeneration: 7, IssueTime: timestamppb.Now(), ImmutableRequestSha256: []byte("12345678901234567890123456789012"), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "Write-Output ok"}}}
|
|
accepted, err := PersistDispatch(context.Background(), store, Session{Generation: 7}, dispatch, time.Now(), agentproto.DefaultLimits())
|
|
encodedSpec, marshalErr := proto.MarshalOptions{Deterministic: true}.Marshal(dispatch.GetSpec())
|
|
if err != nil || marshalErr != nil || accepted.Duplicate || !bytes.Equal(accepted.Command.ExecutionSpec, encodedSpec) {
|
|
t.Fatalf("PersistDispatch = %#v, %v", accepted, err)
|
|
}
|
|
duplicate, err := PersistDispatch(context.Background(), store, Session{Generation: 7}, dispatch, time.Now(), agentproto.DefaultLimits())
|
|
if err != nil || !duplicate.Duplicate {
|
|
t.Fatalf("duplicate PersistDispatch = %#v, %v", duplicate, err)
|
|
}
|
|
}
|
|
|
|
func TestScriptFramesApplyDurablyAndReplaySafely_HP_SCRIPT_04(t *testing.T) {
|
|
store, err := spool.Open(context.Background(), spool.Options{DataDir: filepath.Join(t.TempDir(), "script-spool"), BusyTimeout: time.Second})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer store.Close()
|
|
body := []byte("Write-Output 'hello'\r\n")
|
|
digest := sha256.Sum256(body)
|
|
dispatch := &rvboxv1.CommandDispatch{IssueUuid: "019c46f1-1d02-7000-8000-000000000067", CommandRevision: 1, TargetSessionGeneration: 9, IssueTime: timestamppb.Now(), ImmutableRequestSha256: []byte("12345678901234567890123456789012"), Spec: &rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_POWERSHELL, Source: &rvboxv1.ExecutionSpec_Script{Script: &rvboxv1.ScriptDescriptor{Filename: "hello.ps1", SizeBytes: uint64(len(body)), Sha256: digest[:]}}}}
|
|
accepted, err := PersistDispatch(context.Background(), store, Session{Generation: 9}, dispatch, time.Now().UTC(), agentproto.DefaultLimits())
|
|
if err != nil || accepted.Duplicate {
|
|
t.Fatalf("script PersistDispatch = %#v, %v", accepted, err)
|
|
}
|
|
chunkDigest := sha256.Sum256(body[:8])
|
|
status, err := ApplyScriptChunk(context.Background(), store, Session{Generation: 9}, &rvboxv1.ScriptChunk{IssueUuid: dispatch.GetIssueUuid(), Offset: 0, Data: body[:8], Sha256: chunkDigest[:]}, agentproto.DefaultLimits())
|
|
if err != nil || status.ReceivedBytes != 8 {
|
|
t.Fatalf("first script chunk = %#v, %v", status, err)
|
|
}
|
|
duplicate, err := ApplyScriptChunk(context.Background(), store, Session{Generation: 9}, &rvboxv1.ScriptChunk{IssueUuid: dispatch.GetIssueUuid(), Offset: 0, Data: body[:8], Sha256: chunkDigest[:]}, agentproto.DefaultLimits())
|
|
if err != nil || !duplicate.Duplicate || duplicate.ReceivedBytes != 8 {
|
|
t.Fatalf("replayed script chunk = %#v, %v", duplicate, err)
|
|
}
|
|
restDigest := sha256.Sum256(body[8:])
|
|
if _, err := ApplyScriptChunk(context.Background(), store, Session{Generation: 9}, &rvboxv1.ScriptChunk{IssueUuid: dispatch.GetIssueUuid(), Offset: 8, Data: body[8:], Sha256: restDigest[:]}, agentproto.DefaultLimits()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
committed, err := ApplyScriptCommit(context.Background(), store, Session{Generation: 9}, &rvboxv1.ScriptCommit{IssueUuid: dispatch.GetIssueUuid(), SizeBytes: uint64(len(body)), Sha256: digest[:]}, agentproto.DefaultLimits())
|
|
if err != nil || !committed.Committed || committed.ReceivedBytes != uint64(len(body)) {
|
|
t.Fatalf("script commit = %#v, %v", committed, err)
|
|
}
|
|
replayed, err := ApplyScriptChunk(context.Background(), store, Session{Generation: 9}, &rvboxv1.ScriptChunk{IssueUuid: dispatch.GetIssueUuid(), Offset: 8, Data: body[8:], Sha256: restDigest[:]}, agentproto.DefaultLimits())
|
|
if err != nil || !replayed.Duplicate || !replayed.Committed {
|
|
t.Fatalf("post-commit chunk replay = %#v, %v", replayed, err)
|
|
}
|
|
}
|
|
|
|
func TestSendCommandEventCanonicalDigest_HP_EVENT_02(t *testing.T) {
|
|
transport := &fakeTransport{}
|
|
event := &rvboxv1.CommandEvent{IssueUuid: "019c46f1-1d02-7000-8000-000000000065", EventSeq: 1, ObservedAt: timestamppb.Now(), Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, CommandRevision: 1}}}
|
|
if err := SendCommandEvent(context.Background(), transport, Session{ID: "issued-session", Generation: 7}, event, agentproto.DefaultLimits()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var envelope rvboxv1.AgentEnvelope
|
|
if err := proto.Unmarshal(transport.written, &envelope); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sent := envelope.GetCommandEvent()
|
|
digest, err := agentproto.CommandEventDigest(sent)
|
|
if err != nil || sent == nil || string(digest[:]) != string(sent.GetImmutableEventSha256()) {
|
|
t.Fatalf("sent event = %#v, %v", sent, err)
|
|
}
|
|
}
|
|
|
|
func TestApplyEventAckValidatesAndReleasesPrefix_HP_EVENT_03(t *testing.T) {
|
|
store := &ackStore{}
|
|
ack := &rvboxv1.EventAck{IssueUuid: "019c46f1-1d02-7000-8000-000000000066", ThroughEventSeq: 4}
|
|
if err := ApplyEventAck(context.Background(), store, ack); err != nil || store.issue.String() != ack.IssueUuid || store.sequence != 4 {
|
|
t.Fatalf("ApplyEventAck = issue=%s sequence=%d err=%v", store.issue, store.sequence, err)
|
|
}
|
|
if err := ApplyEventAck(context.Background(), store, &rvboxv1.EventAck{IssueUuid: "bad", ThroughEventSeq: 1}); !errors.Is(err, ErrInvalidEventAck) {
|
|
t.Fatalf("invalid ack error = %v", err)
|
|
}
|
|
}
|
|
|
|
type fakeTransport struct {
|
|
written []byte
|
|
read []byte
|
|
}
|
|
|
|
type reconcileStore struct{ discarded []domain.UUID }
|
|
|
|
func (store *reconcileStore) DiscardTerminal(_ context.Context, issue domain.UUID, _ time.Time) error {
|
|
store.discarded = append(store.discarded, issue)
|
|
return nil
|
|
}
|
|
|
|
type ackStore struct {
|
|
issue domain.UUID
|
|
sequence uint64
|
|
}
|
|
|
|
func (store *ackStore) Ack(_ context.Context, issue domain.UUID, sequence uint64) error {
|
|
store.issue, store.sequence = issue, sequence
|
|
return nil
|
|
}
|
|
|
|
func (transport *fakeTransport) Write(_ context.Context, value []byte) error {
|
|
transport.written = append([]byte(nil), value...)
|
|
return nil
|
|
}
|
|
|
|
func (transport *fakeTransport) Read(context.Context) ([]byte, error) {
|
|
if transport.read == nil {
|
|
return nil, ErrSessionClosed
|
|
}
|
|
return append([]byte(nil), transport.read...), nil
|
|
}
|
|
|
|
func (transport *fakeTransport) Close() error { return nil }
|
|
|
|
func validHello() *rvboxv1.ClientHello {
|
|
return &rvboxv1.ClientHello{ClientId: "win-client", SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: "019c46f1-1d02-7000-8000-000000000061", MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now()}
|
|
}
|