feat: admit fenced agent websocket sessions
This commit is contained in:
@@ -5,6 +5,7 @@ All versions are exact in `go.mod`, generated code, or the toolchain image.
|
|||||||
|
|
||||||
| Dependency | Purpose | Decision |
|
| 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. |
|
| 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. |
|
| 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. |
|
| protobuf-go 1.36.12 | Go protobuf runtime and generator | Official maintained Go protobuf implementation. |
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ module github.com/rvbox/rvbox
|
|||||||
go 1.27.0
|
go 1.27.0
|
||||||
|
|
||||||
require (
|
require (
|
||||||
|
github.com/coder/websocket v1.8.15
|
||||||
github.com/google/uuid v1.6.0
|
github.com/google/uuid v1.6.0
|
||||||
github.com/klauspost/compress v1.19.0
|
github.com/klauspost/compress v1.19.0
|
||||||
github.com/pelletier/go-toml/v2 v2.3.1
|
github.com/pelletier/go-toml/v2 v2.3.1
|
||||||
|
|||||||
@@ -1,5 +1,7 @@
|
|||||||
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
|
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/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 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
|
||||||
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
|
||||||
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
@@ -121,6 +121,18 @@ layer = "unit"
|
|||||||
status = "implemented"
|
status = "implemented"
|
||||||
tests = ["internal/domain/protocol_test.go:TestSelectHighestCompatibleProtocol_HP_SES_01"]
|
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]]
|
[[requirements]]
|
||||||
id = "HP-SES-11"
|
id = "HP-SES-11"
|
||||||
layer = "unit"
|
layer = "unit"
|
||||||
|
|||||||
Reference in New Issue
Block a user