package session import ( "bytes" "context" "crypto/sha256" "errors" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "sync" "testing" "time" "github.com/coder/websocket" rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" "github.com/rvbox/rvbox/internal/agentproto" "github.com/rvbox/rvbox/internal/domain" "github.com/rvbox/rvbox/internal/server/store" "google.golang.org/protobuf/proto" "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() instance, err := domain.NewUUIDv7() if err != nil { t.Fatal(err) } first, firstWelcome := dialAndHello(t, server, "client-a", instance.String()) defer first.CloseNow() if firstWelcome.GetSessionId() == "" || firstWelcome.GetSessionGeneration() != 1 { t.Fatalf("first welcome = %+v", firstWelcome) } second, secondWelcome := dialAndHello(t, server, "client-a", instance.String()) defer second.CloseNow() if secondWelcome.GetSessionId() == firstWelcome.GetSessionId() || secondWelcome.GetSessionGeneration() != 2 { t.Fatalf("replacement welcome = %+v, first = %+v", secondWelcome, firstWelcome) } readContext, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, _, err = first.Read(readContext) if err == nil { t.Fatal("fenced connection remained readable") } } func TestAgentServerDispatchesCommandAdmittedDuringReconcile_HP_SES_11(t *testing.T) { server, agent, cleanup := newTestAgentServerWithStore(t) defer cleanup() instance, err := domain.NewUUIDv7() if err != nil { t.Fatal(err) } connection, welcome := dialAndHello(t, server, "client-race", instance.String()) defer connection.CloseNow() issue, err := domain.NewUUIDv7() if err != nil { t.Fatal(err) } spec, err := proto.Marshal(&rvboxv1.ExecutionSpec{ShellType: rvboxv1.ShellType_SHELL_CMD, Source: &rvboxv1.ExecutionSpec_CommandText{CommandText: "echo dispatch"}}) if err != nil { t.Fatal(err) } now := time.Now().UTC() if _, err := agent.Store.QueueCommand(context.Background(), store.QueueCommandInput{IssueUUID: issue, ClientID: "client-race", IssueTime: now, ReceiptTime: now, ImmutableSHA256: sha256.Sum256([]byte("during-reconcile")), ExecutionSpec: spec}); err != nil { t.Fatal(err) } snapshot, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: welcome.GetSessionId(), SessionGeneration: welcome.GetSessionGeneration(), Payload: &rvboxv1.AgentEnvelope_ReconcileSnapshot{ReconcileSnapshot: &rvboxv1.ReconcileSnapshot{}}, }) if err != nil { t.Fatal(err) } if err := connection.Write(context.Background(), websocket.MessageBinary, snapshot); err != nil { t.Fatal(err) } readContext, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() for index := 0; index < 2; index++ { _, payload, err := connection.Read(readContext) if err != nil { t.Fatal(err) } var envelope rvboxv1.AgentEnvelope if err := proto.Unmarshal(payload, &envelope); err != nil { t.Fatal(err) } if index == 0 && envelope.GetReconcileResult() == nil { t.Fatalf("first post-snapshot envelope = %T, want reconcile result", envelope.Payload) } if index == 1 && (envelope.GetCommandDispatch() == nil || envelope.GetCommandDispatch().GetIssueUuid() != issue.String()) { t.Fatalf("second post-snapshot envelope = %+v, want dispatch %s", envelope.Payload, issue) } } } func TestAgentServerRejectsHostileWireInputs_BH_SES_04(t *testing.T) { server, cleanup := newTestAgentServer(t) defer cleanup() url := agentHTTPURL(server.URL) request, err := http.NewRequest(http.MethodGet, url, nil) if err != nil { t.Fatal(err) } request.Header.Set("Origin", "https://untrusted.example") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } if response.StatusCode != http.StatusForbidden { t.Fatalf("origin response status = %d", response.StatusCode) } _ = response.Body.Close() connection, _, err := websocket.Dial(context.Background(), websocketURL(server.URL), nil) if err != nil { t.Fatal(err) } defer connection.CloseNow() if err := connection.Write(context.Background(), websocket.MessageText, []byte("not protobuf")); err != nil { t.Fatal(err) } readContext, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() _, _, err = connection.Read(readContext) if err == nil { t.Fatal("text first message was not rejected") } instance, err := domain.NewUUIDv7() if err != nil { t.Fatal(err) } valid, welcome := dialAndHello(t, server, "client-b", instance.String()) defer valid.CloseNow() stale, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: "not-the-issued-session", SessionGeneration: welcome.GetSessionGeneration(), Payload: &rvboxv1.AgentEnvelope_ClientCapacity{ClientCapacity: &rvboxv1.ClientCapacity{ MaxRunningCommands: 1, MaxQueuedCommands: 1, }}, }) if err != nil { t.Fatal(err) } if err := valid.Write(context.Background(), websocket.MessageBinary, stale); err != nil { t.Fatal(err) } _, _, err = valid.Read(readContext) if err == nil { t.Fatal("stale session envelope was not rejected") } var closeError websocket.CloseError if !errors.As(err, &closeError) || closeError.Code != websocket.StatusPolicyViolation { t.Fatalf("stale envelope close = %v, want policy violation", err) } } func TestWireEventAppendCarriesClientBinding_HP_EVENT_01(t *testing.T) { event := &rvboxv1.CommandEvent{IssueUuid: "019c46f1-1d02-7000-8000-000000000096", EventSeq: 1, ObservedAt: timestamppb.Now(), Payload: &rvboxv1.CommandEvent_Lifecycle{Lifecycle: &rvboxv1.LifecycleChange{Lifecycle: rvboxv1.CommandLifecycle_COMMAND_RUNNING, CommandRevision: 1}}} digest, err := agentproto.CommandEventDigest(event) if err != nil { t.Fatal(err) } event.ImmutableEventSha256 = digest[:] appendEvent, err := eventAppendFromWire(event, "client-a", 7, time.Now()) if err != nil || appendEvent.ClientID != "client-a" || appendEvent.SessionGeneration != 7 || appendEvent.EventSeq != 1 || appendEvent.EventType != 4 || appendEvent.ImmutableSHA256 != digest { t.Fatalf("wire event append = %#v, %v", appendEvent, err) } } func newTestAgentServer(t *testing.T) (*httptest.Server, func()) { server, _, cleanup := newTestAgentServerWithStore(t) return server, cleanup } func newTestAgentServerWithStore(t *testing.T) (*httptest.Server, *AgentServer, func()) { t.Helper() dataDirectory := filepath.Join(t.TempDir(), "store") if err := os.Mkdir(dataDirectory, 0o700); err != nil { t.Fatal(err) } persistence, err := store.Open(context.Background(), store.Options{ DataDir: dataDirectory, BusyTimeout: time.Second, }) if err != nil { t.Fatal(err) } agent := &AgentServer{ Store: persistence, Registry: NewRegistry(), WriteDeadline: 200 * time.Millisecond, HeartbeatIdle: time.Second, LivenessTimeout: 2 * time.Second, } httpServer := httptest.NewServer(agent) return httpServer, agent, func() { httpServer.Close() if err := persistence.Close(); err != nil { t.Errorf("close persistence: %v", err) } } } func dialAndHello(t *testing.T, server *httptest.Server, clientID, instanceID string) (*websocket.Conn, *rvboxv1.AgentEnvelope) { t.Helper() connection, _, err := websocket.Dial(context.Background(), websocketURL(server.URL), nil) if err != nil { t.Fatal(err) } hello, err := proto.Marshal(&rvboxv1.AgentEnvelope{Payload: &rvboxv1.AgentEnvelope_ClientHello{ClientHello: &rvboxv1.ClientHello{ ClientId: clientID, SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\ProgramData\RVBox\work`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: instanceID, MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now(), }}}) if err != nil { connection.CloseNow() t.Fatal(err) } if err := connection.Write(context.Background(), websocket.MessageBinary, hello); err != nil { connection.CloseNow() t.Fatal(err) } readContext, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() messageType, welcomeBytes, err := connection.Read(readContext) if err != nil { connection.CloseNow() t.Fatal(err) } if messageType != websocket.MessageBinary { connection.CloseNow() t.Fatalf("welcome message type = %v", messageType) } var welcome rvboxv1.AgentEnvelope if err := proto.Unmarshal(welcomeBytes, &welcome); err != nil { connection.CloseNow() t.Fatal(err) } if welcome.GetServerWelcome() == nil { connection.CloseNow() t.Fatalf("welcome payload = %T", welcome.Payload) } // Reconnecting agents receive an advisory target list immediately after // welcome. Consume it here so callers that exercise fencing or malformed // frames observe the next read from the actual session boundary. readContext, cancel = context.WithTimeout(context.Background(), time.Second) defer cancel() messageType, advisoryBytes, err := connection.Read(readContext) if err != nil { connection.CloseNow() t.Fatal(err) } if messageType != websocket.MessageBinary { connection.CloseNow() t.Fatalf("reconciliation message type = %v", messageType) } var advisory rvboxv1.AgentEnvelope if err := proto.Unmarshal(advisoryBytes, &advisory); err != nil || advisory.GetReconcileRequest() == nil { connection.CloseNow() t.Fatalf("reconciliation advisory = %+v, %v", advisory.Payload, err) } return connection, &welcome } func agentHTTPURL(serverURL string) string { return serverURL + "/v1/agent" } func websocketURL(httpURL string) string { return "ws" + strings.TrimPrefix(agentHTTPURL(httpURL), "http") }