feat: admit fenced agent websocket sessions
This commit is contained in:
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user