feat: admit fenced agent websocket sessions

This commit is contained in:
2026-09-02 06:08:23 +00:00
parent cfb2693930
commit 083cca1ce7
6 changed files with 521 additions and 0 deletions
@@ -0,0 +1,183 @@
package session
import (
"context"
"errors"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/coder/websocket"
rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1"
"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 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 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 newTestAgentServer(t *testing.T) (*httptest.Server, 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, 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)
}
return connection, &welcome
}
func agentHTTPURL(serverURL string) string { return serverURL + "/v1/agent" }
func websocketURL(httpURL string) string {
return "ws" + strings.TrimPrefix(agentHTTPURL(httpURL), "http")
}