Files
rvbox/internal/server/session/agent_server.go
T

485 lines
17 KiB
Go

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")
ErrDispatchDataFull = errors.New("session dispatch data lane is full")
)
// 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()
capacity := &CapacityShadow{}
if !capacity.UpdateAdvertised(0, 0, hello.GetMaxRunningCommands(), hello.GetMaxQueuedCommands()) {
server.close(connection, websocket.StatusPolicyViolation, "invalid initial client capacity")
return
}
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},
})
resultWritten := make(chan struct{})
if err != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: encoded, Written: resultWritten}) != nil {
server.close(connection, websocket.StatusInternalError, "could not queue reconciliation result")
return
}
select {
case <-resultWritten:
case <-sessionContext.Done():
return
}
lane := capacity.Reserve()
dispatched, dispatchErr := false, error(nil)
if lane != DispatchNone {
dispatched, dispatchErr = server.enqueueNextDispatch(sessionContext, queue, hello.GetClientId(), hello.GetPlatform(), encodedSessionID, registration.Generation)
if !dispatched {
capacity.Release(lane)
}
}
if dispatchErr != nil && !errors.Is(dispatchErr, ErrDispatchDataFull) {
server.close(connection, websocket.StatusInternalError, "could not queue command dispatch")
return
}
continue
}
if advertised := envelope.GetClientCapacity(); advertised != nil {
if !capacity.UpdateAdvertised(advertised.GetRunningCommands(), advertised.GetQueuedCommands(), advertised.GetMaxRunningCommands(), advertised.GetMaxQueuedCommands()) {
server.close(connection, websocket.StatusPolicyViolation, "invalid client capacity")
return
}
continue
}
if acknowledgement := envelope.GetCommandAccepted(); acknowledgement != nil {
issue, parseErr := domain.ParseUUIDv7(acknowledgement.GetIssueUuid())
if parseErr != nil {
server.close(connection, websocket.StatusPolicyViolation, "invalid command acceptance")
return
}
if _, acceptErr := server.Store.RecordCommandAcceptance(sessionContext, issue, hello.GetClientId(), registration.Generation, acknowledgement.GetCommandRevision(), acknowledgement.GetAccepted(), server.now()); acceptErr != nil {
server.close(connection, websocket.StatusPolicyViolation, "invalid command acceptance")
return
}
continue
}
if event := envelope.GetCommandEvent(); event != nil {
appendEvent, eventErr := eventAppendFromWire(event, hello.GetClientId(), registration.Generation, server.now())
if eventErr != nil {
server.close(connection, websocket.StatusPolicyViolation, "invalid command event")
return
}
appended, eventErr := server.Store.AppendCommandEvent(sessionContext, appendEvent)
if eventErr != nil {
server.close(connection, websocket.StatusPolicyViolation, "command event was not accepted")
return
}
ack, eventErr := proto.Marshal(&rvboxv1.AgentEnvelope{SessionId: encodedSessionID, SessionGeneration: registration.Generation, Payload: &rvboxv1.AgentEnvelope_EventAck{EventAck: &rvboxv1.EventAck{IssueUuid: event.GetIssueUuid(), ThroughEventSeq: appended.ThroughEventSeq}}})
if eventErr != nil || queue.EnqueueControl(Frame{Kind: FrameControl, Payload: ack}) != nil {
server.close(connection, websocket.StatusInternalError, "could not acknowledge command event")
return
}
}
}
}
func eventAppendFromWire(event *rvboxv1.CommandEvent, clientID string, generation uint64, receipt time.Time) (store.EventAppend, error) {
issue, err := domain.ParseUUIDv7(event.GetIssueUuid())
if err != nil {
return store.EventAppend{}, err
}
payload, err := proto.MarshalOptions{Deterministic: true}.Marshal(event)
if err != nil {
return store.EventAppend{}, err
}
var immutable [32]byte
copy(immutable[:], event.GetImmutableEventSha256())
result := store.EventAppend{IssueUUID: [16]byte(issue), ClientID: clientID, SessionGeneration: generation, EventSeq: event.GetEventSeq(), ObservedUnixNano: event.GetObservedAt().AsTime().UnixNano(), ReceiptUnixNano: receipt.UnixNano(), EventType: eventType(event), Compression: 1, RawLength: uint64(len(payload)), Payload: payload, ImmutableSHA256: immutable}
if output := event.GetOutput(); output != nil {
result.Stream = uint16(output.GetStream())
result.Output = true
}
if result.EventType == 0 {
return store.EventAppend{}, errors.New("unsupported command event")
}
return result, nil
}
func eventType(event *rvboxv1.CommandEvent) uint16 {
switch event.Payload.(type) {
case *rvboxv1.CommandEvent_Lifecycle:
return 4
case *rvboxv1.CommandEvent_Output:
return 5
case *rvboxv1.CommandEvent_Resource:
return 6
case *rvboxv1.CommandEvent_StdinAck:
return 7
case *rvboxv1.CommandEvent_SignalResult:
return 8
case *rvboxv1.CommandEvent_ScriptStatus:
return 9
case *rvboxv1.CommandEvent_OutputTruncation:
return 10
case *rvboxv1.CommandEvent_OutputIncomplete:
return 11
default:
return 0
}
}
// enqueueNextDispatch records the queued-to-dispatched transition before
// exposing work to the network. A full data lane is a pre-write failure, so
// only the owning generation can put the command back into the queue.
func (server *AgentServer) enqueueNextDispatch(ctx context.Context, queue *WriterQueue, clientID string, platform rvboxv1.Platform, sessionID string, generation uint64) (bool, error) {
candidate, err := server.Store.ClaimNextDispatch(ctx, clientID, generation, server.now())
if err != nil || candidate == nil {
return false, err
}
requeue := func(cause error) (bool, error) {
_, rollbackErr := server.Store.RequeueDispatch(ctx, candidate.IssueUUID, clientID, generation)
if rollbackErr != nil {
return false, rollbackErr
}
return false, cause
}
spec := &rvboxv1.ExecutionSpec{}
if err := proto.Unmarshal(candidate.ExecutionSpec, spec); err != nil {
return requeue(fmt.Errorf("decode persisted execution spec: %w", err))
}
if err := agentproto.ValidateExecutionSpec(spec, server.limits(), platform); err != nil {
return requeue(fmt.Errorf("validate persisted execution spec: %w", err))
}
dispatch := &rvboxv1.CommandDispatch{
IssueUuid: candidate.IssueUUID.String(), CommandRevision: candidate.Revision, TargetSessionGeneration: generation,
IssueTime: timestamppb.New(candidate.IssueTime), Spec: spec, ImmutableRequestSha256: candidate.ImmutableSHA256[:],
}
if candidate.QueueExpiryTime != nil {
dispatch.QueueExpiryTime = timestamppb.New(*candidate.QueueExpiryTime)
}
encoded, err := proto.Marshal(&rvboxv1.AgentEnvelope{
SessionId: sessionID, SessionGeneration: generation, Payload: &rvboxv1.AgentEnvelope_CommandDispatch{CommandDispatch: dispatch},
})
if err != nil {
return requeue(err)
}
if !queue.EnqueueData(Frame{Kind: FrameData, Payload: encoded}) {
return requeue(ErrDispatchDataFull)
}
return true, nil
}
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
}
if frame.Written != nil {
close(frame.Written)
}
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)
}