Files
rvbox/internal/server/session/agent_server_test.go
T

324 lines
11 KiB
Go

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")
}