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 } if snapshot := envelope.GetReconcileSnapshot(); snapshot != nil { result, reconcileErr := server.Store.ReconcileClientSnapshot(sessionContext, hello.GetClientId(), snapshot) if reconcileErr != nil { server.close(connection, websocket.StatusPolicyViolation, "reconciliation failed") return } encoded, err := proto.Marshal(&rvboxv1.AgentEnvelope{ SessionId: encodedSessionID, SessionGeneration: registration.Generation, Payload: &rvboxv1.AgentEnvelope_ReconcileResult{ReconcileResult: result}, }) if err != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded}) != nil { server.close(connection, websocket.StatusInternalError, "could not queue reconciliation result") 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) { _ = connection.Close(status, reason) } 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) }