diff --git a/docs/dependency-decisions.md b/docs/dependency-decisions.md index 11ff156..aa494b5 100644 --- a/docs/dependency-decisions.md +++ b/docs/dependency-decisions.md @@ -5,6 +5,7 @@ All versions are exact in `go.mod`, generated code, or the toolchain image. | Dependency | Purpose | Decision | | --- | --- | --- | +| coder/websocket 1.8.15 | Context-aware WebSocket server/client transport | Maintained Go WebSocket implementation with bounded message reads and explicit Ping support; RVBox keeps protobuf validation, session fencing, origin policy, and backpressure in its own code. | | Buf 1.72.0 | Protobuf formatting, linting, and deterministic generation | Current stable release when the v1 implementation began; installed only in the toolchain image. | | protoc 35.0 | Protobuf compiler and well-known includes | Pinned archive with a checked SHA-256; included even though normal generation is driven by Buf. | | protobuf-go 1.36.12 | Go protobuf runtime and generator | Official maintained Go protobuf implementation. | diff --git a/go.mod b/go.mod index 8e78ff5..e3cf41f 100644 --- a/go.mod +++ b/go.mod @@ -3,6 +3,7 @@ module github.com/rvbox/rvbox go 1.27.0 require ( + github.com/coder/websocket v1.8.15 github.com/google/uuid v1.6.0 github.com/klauspost/compress v1.19.0 github.com/pelletier/go-toml/v2 v2.3.1 diff --git a/go.sum b/go.sum index 78b969e..9926ea7 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,7 @@ github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/coder/websocket v1.8.15 h1:6B2JPeOGlpff2Uz6vOEH1Vzpi0iUz20A+lPVhPHtNUA= +github.com/coder/websocket v1.8.15/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI= diff --git a/internal/server/session/agent_server.go b/internal/server/session/agent_server.go new file mode 100644 index 0000000..d8af37a --- /dev/null +++ b/internal/server/session/agent_server.go @@ -0,0 +1,322 @@ +package session + +import ( + "context" + "crypto/rand" + "encoding/base64" + "errors" + "fmt" + "net/http" + "sync" + "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" +) + +const ( + defaultAgentPath = "/v1/agent" + defaultWriteWait = 10 * time.Second +) + +var ( + ErrUnexpectedOrigin = errors.New("agent connections must not send an Origin header") + ErrUnexpectedMessage = errors.New("agent message is not a binary protobuf envelope") + ErrStaleSession = errors.New("agent message is for an unknown or fenced session") +) + +// AgentServer is the transport edge for agent WebSocket sessions. Command +// processing is intentionally injected above this admission/fencing boundary. +type AgentServer struct { + Store *store.Store + Registry *Registry + Path string + Limits agentproto.Limits + SupportedProtocol *rvboxv1.ProtocolRange + WriteDeadline time.Duration + HeartbeatIdle time.Duration + LivenessTimeout time.Duration + Now func() time.Time +} + +func (server *AgentServer) ServeHTTP(response http.ResponseWriter, request *http.Request) { + if request.URL.Path != server.agentPath() { + http.NotFound(response, request) + return + } + if request.Header.Get("Origin") != "" { + http.Error(response, ErrUnexpectedOrigin.Error(), http.StatusForbidden) + return + } + if server.Store == nil { + http.Error(response, "agent server storage is unavailable", http.StatusServiceUnavailable) + return + } + if server.Registry == nil { + http.Error(response, "agent session registry is unavailable", http.StatusServiceUnavailable) + return + } + connection, err := websocket.Accept(response, request, &websocket.AcceptOptions{ + CompressionMode: websocket.CompressionDisabled, + }) + if err != nil { + return + } + defer connection.CloseNow() + server.serveConnection(request.Context(), connection) +} + +func (server *AgentServer) serveConnection(parent context.Context, connection *websocket.Conn) { + started := time.Now() + messageType, payload, err := connection.Read(parent) + if err != nil { + return + } + if messageType != websocket.MessageBinary { + server.close(connection, websocket.StatusUnsupportedData, ErrUnexpectedMessage.Error()) + return + } + helloEnvelope, err := agentproto.DecodeEnvelope(payload, server.limits(), rvboxv1.Platform_PLATFORM_UNSPECIFIED) + if err != nil || helloEnvelope.GetClientHello() == nil { + server.close(connection, websocket.StatusPolicyViolation, "invalid ClientHello") + return + } + hello := helloEnvelope.GetClientHello() + instanceID, err := domain.ParseUUIDv7(hello.GetClientInstanceId()) + if err != nil { + server.close(connection, websocket.StatusPolicyViolation, "invalid client instance ID") + return + } + selected, err := domain.SelectProtocol(server.protocolRange(), hello.GetSupportedProtocol()) + if err != nil { + server.close(connection, websocket.StatusPolicyViolation, "unsupported protocol") + return + } + sessionID, err := randomSessionID() + if err != nil { + server.close(connection, websocket.StatusInternalError, "could not allocate session") + return + } + shells, err := proto.Marshal(&rvboxv1.ClientHello{SupportedShells: hello.GetSupportedShells()}) + if err != nil { + server.close(connection, websocket.StatusInternalError, "could not persist client capabilities") + return + } + registration, err := server.Store.RegisterClientSession(parent, store.ClientRegistration{ + ClientID: hello.GetClientId(), Platform: uint32(hello.GetPlatform()), Architecture: hello.GetArchitecture(), + DaemonVersion: hello.GetDaemonVersion(), DaemonCWD: hello.GetDaemonCwd(), SupportedShells: shells, + ClientInstanceID: [16]byte(instanceID), SessionID: sessionID, ConnectedAt: server.now(), + }) + if err != nil { + server.close(connection, websocket.StatusPolicyViolation, sessionCloseReason(err)) + return + } + registry := server.Registry + handle, err := registry.Install(hello.GetClientId(), sessionID, registration.Generation) + if err != nil { + _ = server.Store.CloseLiveSession(context.Background(), sessionID, registration.Generation, "registry installation failed", server.now()) + server.close(connection, websocket.StatusInternalError, "could not install session") + return + } + defer func() { + registry.Remove(handle) + _ = server.Store.CloseLiveSession(context.Background(), sessionID, registration.Generation, "connection closed", server.now()) + }() + + sessionContext, cancel := context.WithCancel(parent) + stopFence := context.AfterFunc(handle.Context, cancel) + defer func() { + stopFence() + cancel() + }() + + queue := NewWriterQueue(16, 64, 8) + defer queue.Close() + encodedSessionID := encodeSessionID(sessionID) + welcome, err := proto.Marshal(&rvboxv1.AgentEnvelope{ + SessionId: encodedSessionID, SessionGeneration: registration.Generation, + Payload: &rvboxv1.AgentEnvelope_ServerWelcome{ServerWelcome: &rvboxv1.ServerWelcome{ + SelectedProtocol: selected, ServerTime: timestamppb.New(server.now()), + }}, + }) + if err != nil { + server.close(connection, websocket.StatusInternalError, "could not encode welcome") + return + } + if err := queue.EnqueueControl(Frame{Kind: FrameControl, Payload: welcome}); err != nil { + server.close(connection, websocket.StatusInternalError, "could not queue welcome") + return + } + + heartbeat := newSynchronizedHeartbeat(server.heartbeatIdle(), server.livenessTimeout(), 0) + writerDone := make(chan struct{}) + go func() { + defer close(writerDone) + server.writeLoop(sessionContext, connection, queue, heartbeat, started) + cancel() + }() + defer func() { <-writerDone }() + + for { + messageType, payload, err = connection.Read(sessionContext) + if err != nil { + return + } + // Elapsed time is monotonic; peer-provided wall timestamps are never used + // to determine liveness. + heartbeat.Observe(time.Since(started)) + if messageType != websocket.MessageBinary { + server.close(connection, websocket.StatusUnsupportedData, ErrUnexpectedMessage.Error()) + return + } + envelope, err := agentproto.DecodeEnvelope(payload, server.limits(), hello.GetPlatform()) + if err != nil || envelope.GetClientHello() != nil || envelope.GetSessionId() != encodedSessionID || envelope.GetSessionGeneration() != registration.Generation { + server.close(connection, websocket.StatusPolicyViolation, ErrStaleSession.Error()) + return + } + clientID, err := server.Store.ValidateLiveSession(sessionContext, sessionID, registration.Generation) + if err != nil || clientID != hello.GetClientId() { + server.close(connection, websocket.StatusPolicyViolation, ErrStaleSession.Error()) + return + } + } +} + +func (server *AgentServer) writeLoop(ctx context.Context, connection *websocket.Conn, queue *WriterQueue, heartbeat *synchronizedHeartbeat, started time.Time) { + for { + frameContext, cancel := context.WithTimeout(ctx, server.heartbeatPollInterval()) + frame, err := queue.Next(frameContext) + cancel() + if err == nil { + writeContext, writeCancel := context.WithTimeout(ctx, server.writeDeadline()) + err = connection.Write(writeContext, websocket.MessageBinary, frame.Payload) + writeCancel() + if err != nil { + return + } + continue + } + if !errors.Is(err, context.DeadlineExceeded) { + return + } + switch heartbeat.Check(time.Since(started)) { + case HeartbeatPing: + pingContext, pingCancel := context.WithTimeout(ctx, server.writeDeadline()) + err = connection.Ping(pingContext) + pingCancel() + if err != nil { + return + } + heartbeat.Observe(time.Since(started)) + case HeartbeatClose: + return + } + } +} + +func (server *AgentServer) close(connection *websocket.Conn, status websocket.StatusCode, reason string) { + ctx, cancel := context.WithTimeout(context.Background(), server.writeDeadline()) + defer cancel() + _ = connection.Close(status, reason) + _ = ctx +} + +func (server *AgentServer) agentPath() string { + if server.Path == "" { + return defaultAgentPath + } + return server.Path +} + +func (server *AgentServer) limits() agentproto.Limits { + if server.Limits.MaxEnvelopeBytes == 0 { + return agentproto.DefaultLimits() + } + return server.Limits +} + +func (server *AgentServer) protocolRange() *rvboxv1.ProtocolRange { + if server.SupportedProtocol == nil { + return &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0} + } + return server.SupportedProtocol +} + +func (server *AgentServer) writeDeadline() time.Duration { + if server.WriteDeadline <= 0 { + return defaultWriteWait + } + return server.WriteDeadline +} + +func (server *AgentServer) heartbeatIdle() time.Duration { + if server.HeartbeatIdle <= 0 { + return 10 * time.Second + } + return server.HeartbeatIdle +} + +func (server *AgentServer) livenessTimeout() time.Duration { + if server.LivenessTimeout <= server.heartbeatIdle() { + return 30 * time.Second + } + return server.LivenessTimeout +} + +func (server *AgentServer) heartbeatPollInterval() time.Duration { + interval := server.heartbeatIdle() / 2 + if interval <= 0 || interval > server.writeDeadline() { + return server.writeDeadline() + } + return interval +} + +func (server *AgentServer) now() time.Time { + if server.Now != nil { + return server.Now().UTC() + } + return time.Now().UTC() +} + +func randomSessionID() ([16]byte, error) { + var result [16]byte + if _, err := rand.Read(result[:]); err != nil { + return result, fmt.Errorf("read session entropy: %w", err) + } + return result, nil +} + +func encodeSessionID(value [16]byte) string { return base64.RawURLEncoding.EncodeToString(value[:]) } + +func sessionCloseReason(err error) string { + if errors.Is(err, store.ErrTakeoverRequired) { + return "client takeover authorization required" + } + return "client registration rejected" +} + +type synchronizedHeartbeat struct { + mu sync.Mutex + heartbeat *Heartbeat +} + +func newSynchronizedHeartbeat(idle, timeout, now time.Duration) *synchronizedHeartbeat { + return &synchronizedHeartbeat{heartbeat: NewHeartbeat(idle, timeout, now)} +} + +func (heartbeat *synchronizedHeartbeat) Observe(now time.Duration) { + heartbeat.mu.Lock() + defer heartbeat.mu.Unlock() + heartbeat.heartbeat.ObserveInbound(now) +} + +func (heartbeat *synchronizedHeartbeat) Check(now time.Duration) HeartbeatAction { + heartbeat.mu.Lock() + defer heartbeat.mu.Unlock() + return heartbeat.heartbeat.Check(now) +} diff --git a/internal/server/session/agent_server_test.go b/internal/server/session/agent_server_test.go new file mode 100644 index 0000000..1269d67 --- /dev/null +++ b/internal/server/session/agent_server_test.go @@ -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") +} diff --git a/test/coverage.toml b/test/coverage.toml index c3cde0f..c0e071b 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -121,6 +121,18 @@ layer = "unit" status = "implemented" tests = ["internal/domain/protocol_test.go:TestSelectHighestCompatibleProtocol_HP_SES_01"] +[[requirements]] +id = "HP-SES-05" +layer = "integration" +status = "implemented" +tests = ["internal/server/session/agent_server_test.go:TestAgentServerRegistrationAndReplacement_HP_SES_05"] + +[[requirements]] +id = "BH-SES-04" +layer = "integration" +status = "implemented" +tests = ["internal/server/session/agent_server_test.go:TestAgentServerRejectsHostileWireInputs_BH_SES_04"] + [[requirements]] id = "HP-SES-11" layer = "unit"