package gateway import ( "context" "encoding/xml" "errors" "fmt" "strings" "sync" "time" protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol" ) const ( ProviderProfileApollo = "apollo" ProviderStateStarting = "starting" ProviderStateReady = "ready" ProviderStateDisconnected = "disconnected" ProviderStateTerminating = "terminating" ProviderStateTerminated = "terminated" ProviderStateCleanup = "cleanup_pending" ProviderStateFailed = "failed" ) var ( ErrProviderIdentity = errors.New("provider identity rejected") ErrProviderMalformed = errors.New("provider response malformed") ErrProviderTimeout = errors.New("provider operation timed out") ErrProviderDisconnected = errors.New("provider disconnected") ErrProviderCleanup = errors.New("provider cleanup pending") ErrProviderTerminated = errors.New("provider session terminated") ) type ProviderIdentity struct { UniqueID string Fingerprint string NotBefore time.Time NotAfter time.Time } func (i ProviderIdentity) Key() string { return i.UniqueID + "#" + i.Fingerprint } func providerIdentityFromKey(value string) (ProviderIdentity, bool) { uniqueID, fingerprint, ok := strings.Cut(strings.TrimSpace(value), "#") if !ok || uniqueID == "" || fingerprint == "" || strings.Contains(fingerprint, "#") || len(uniqueID) > 128 || len(fingerprint) > 256 { return ProviderIdentity{}, false } return ProviderIdentity{UniqueID: uniqueID, Fingerprint: fingerprint}, true } func (i ProviderIdentity) Validate(now time.Time, expected ProviderIdentity) error { if i.UniqueID == "" || expected.UniqueID == "" || i.UniqueID != expected.UniqueID || (i.Fingerprint != "" && i.Fingerprint != expected.Fingerprint) { return ErrProviderIdentity } if !i.NotBefore.IsZero() && now.Before(i.NotBefore) { return ErrProviderIdentity } if !i.NotAfter.IsZero() && !now.Before(i.NotAfter) { return ErrProviderIdentity } return nil } type ManagementInfo struct { Identity ProviderIdentity Name string } func ParseManagementXML(data []byte) (ManagementInfo, error) { if len(data) == 0 || len(data) > 64*1024 { return ManagementInfo{}, ErrProviderMalformed } var document struct { XMLName xml.Name `xml:"root"` UniqueID string `xml:"uniqueid"` LegacyID string `xml:"unique_id"` Fingerprint string `xml:"fingerprint"` NotBefore string `xml:"not_before"` NotAfter string `xml:"not_after"` Name string `xml:"name"` } decoder := xml.NewDecoder(strings.NewReader(string(data))) decoder.Strict = true if err := decoder.Decode(&document); err != nil { return ManagementInfo{}, fmt.Errorf("%w: %v", ErrProviderMalformed, err) } if document.UniqueID == "" { document.UniqueID = document.LegacyID } identity := ProviderIdentity{UniqueID: document.UniqueID, Fingerprint: document.Fingerprint} var err error if document.NotBefore != "" { identity.NotBefore, err = time.Parse(time.RFC3339Nano, document.NotBefore) if err != nil { return ManagementInfo{}, ErrProviderMalformed } } if document.NotAfter != "" { identity.NotAfter, err = time.Parse(time.RFC3339Nano, document.NotAfter) if err != nil { return ManagementInfo{}, ErrProviderMalformed } } if identity.UniqueID == "" || len(identity.UniqueID) > 128 || len(identity.Fingerprint) > 256 { return ManagementInfo{}, ErrProviderMalformed } return ManagementInfo{Identity: identity, Name: document.Name}, nil } type RTSPResponse struct { StatusCode int Session string Transport string } func ParseRTSPResponse(data []byte) (RTSPResponse, error) { if len(data) == 0 || len(data) > 16*1024 { return RTSPResponse{}, ErrProviderMalformed } text := string(data) if !strings.Contains(text, "\r\n") { return RTSPResponse{}, ErrProviderMalformed } lines := strings.Split(text, "\r\n") if len(lines) < 2 { return RTSPResponse{}, ErrProviderMalformed } statusParts := strings.SplitN(lines[0], " ", 3) if len(statusParts) < 2 || statusParts[0] != "RTSP/1.0" { return RTSPResponse{}, ErrProviderMalformed } var response RTSPResponse if _, err := fmt.Sscanf(statusParts[1], "%d", &response.StatusCode); err != nil || response.StatusCode != 200 { return RTSPResponse{}, ErrProviderMalformed } for _, line := range lines[1:] { if line == "" { break } key, value, ok := strings.Cut(line, ":") if !ok { return RTSPResponse{}, ErrProviderMalformed } switch strings.ToLower(strings.TrimSpace(key)) { case "session": response.Session = strings.TrimSpace(value) case "transport": response.Transport = strings.TrimSpace(value) } } if response.Session == "" || response.Transport == "" || len(response.Session) > 256 || len(response.Transport) > 1024 { return RTSPResponse{}, ErrProviderMalformed } return response, nil } type LaunchRequest struct { SessionID string Capabilities protocol.CapabilityProfile ProviderProfile string ProviderIdentity string ProviderWork protocol.ProviderSessionWork } type InputEvent struct { Sequence uint32 Device string Code int32 Pressed bool Payload []byte } type Feedback struct { Sequence uint32 Kind FeedbackKind Payload []byte } type ProviderEventKind uint8 const ( ProviderEventTerminated ProviderEventKind = iota + 1 ProviderEventRumble ProviderEventHDR ) type ProviderEvent struct { Kind ProviderEventKind Payload []byte } // ProviderTelemetry holds measured provider-channel state only; it never // contains provider routes, credentials, or payload bytes. type ProviderTelemetry struct { State string ControlRTT time.Duration ControlJitter time.Duration ReliableSent uint64 ReliableRetransmits uint64 PendingReliable uint64 MediaDrops uint64 } type Provider interface { Start(context.Context, LaunchRequest) (ProviderSession, error) } type ProviderSession interface { Ready(context.Context) error Video() <-chan []byte Audio() <-chan []byte Events() <-chan ProviderEvent Input(context.Context, InputEvent) error Feedback(context.Context, Feedback) error ReadClipboard(context.Context) (string, error) WriteClipboard(context.Context, string) error Telemetry() ProviderTelemetry ReleaseAll(context.Context) error Terminate(context.Context) error State() protocol.ProviderState } type ApolloBackend interface { Management(context.Context, LaunchRequest) ([]byte, error) Setup(context.Context, LaunchRequest) ([]byte, error) Open(context.Context, LaunchRequest, RTSPResponse) (ProviderSession, error) } type ApolloAdapter struct { backend ApolloBackend expected ProviderIdentity now func() time.Time readyTTL time.Duration } func NewApolloAdapter(backend ApolloBackend, expected ProviderIdentity) *ApolloAdapter { return &ApolloAdapter{backend: backend, expected: expected, now: time.Now, readyTTL: 2 * time.Second} } func (a *ApolloAdapter) Start(ctx context.Context, request LaunchRequest) (ProviderSession, error) { if a == nil || a.backend == nil || request.ProviderProfile != ProviderProfileApollo { return nil, ErrProviderIdentity } management, err := a.backend.Management(ctx, request) if err != nil { return nil, err } info, err := ParseManagementXML(management) if err != nil { return nil, err } expected := a.expected if request.ProviderWork.ProviderIdentity != "" { parsed, ok := providerIdentityFromKey(request.ProviderWork.ProviderIdentity) if !ok { return nil, ErrProviderIdentity } expected = parsed } if err := info.Identity.Validate(a.now(), expected); err != nil { return nil, err } if request.ProviderIdentity != "" && info.Identity.UniqueID != expected.UniqueID { return nil, ErrProviderIdentity } rawRTSP, err := a.backend.Setup(ctx, request) if err != nil { return nil, err } rtsp, err := ParseRTSPResponse(rawRTSP) if err != nil { return nil, err } session, err := a.backend.Open(ctx, request, rtsp) if err != nil { return nil, err } readyCtx, cancel := context.WithTimeout(ctx, a.readyTTL) defer cancel() if err := session.Ready(readyCtx); err != nil { cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), time.Second) cleanupErr := session.Terminate(cleanupCtx) cleanupCancel() if cleanupErr != nil { return nil, fmt.Errorf("%w: %v", ErrProviderCleanup, err) } if errors.Is(err, context.DeadlineExceeded) { return nil, ErrProviderTimeout } return nil, err } return session, nil } type FakeFailure string const ( FakeFailureNone FakeFailure = "" FakeFailureIdentity FakeFailure = "identity" FakeFailureMalformed FakeFailure = "malformed" FakeFailureReadinessTimeout FakeFailure = "readiness-timeout" FakeFailureProviderDisconnect FakeFailure = "provider-disconnect" FakeFailureTerminationTimeout FakeFailure = "termination-timeout" ) type FakeApolloConfig struct { Identity ProviderIdentity Failure FakeFailure Video [][]byte Audio [][]byte Now time.Time } type FakeApollo struct { config FakeApolloConfig mu sync.Mutex last *fakeSession } func NewFakeApollo(config FakeApolloConfig) *FakeApollo { if config.Identity.UniqueID == "" { config.Identity.UniqueID = "apollo-fixture-1" } if config.Identity.Fingerprint == "" { config.Identity.Fingerprint = "sha256:fixture-apollo-1" } if config.Now.IsZero() { config.Now = time.Date(2026, time.January, 1, 0, 0, 0, 0, time.UTC) } if len(config.Video) == 0 { config.Video = [][]byte{[]byte{0x00, 0x00, 0x01, 0x65, 0x01, 0x02}} } if len(config.Audio) == 0 { config.Audio = [][]byte{[]byte{0x4f, 0x70, 0x75, 0x73, 0x01}} } return &FakeApollo{config: config} } func (f *FakeApollo) Management(context.Context, LaunchRequest) ([]byte, error) { if f.config.Failure == FakeFailureMalformed { return []byte(""), nil } identity := f.config.Identity if f.config.Failure == FakeFailureIdentity { identity.Fingerprint = "sha256:changed-fixture" } return []byte(fmt.Sprintf("%s%s%s%sfixture-apollo", identity.UniqueID, identity.Fingerprint, f.config.Now.Add(-time.Hour).Format(time.RFC3339), f.config.Now.Add(time.Hour).Format(time.RFC3339))), nil } func (f *FakeApollo) Setup(context.Context, LaunchRequest) ([]byte, error) { if f.config.Failure == FakeFailureMalformed { return []byte("RTSP/1.0 200 OK\r\n\r\n"), nil } return []byte("RTSP/1.0 200 OK\r\nSession: fixture-session\r\nTransport: unicast;server_port=43000\r\n\r\n"), nil } func (f *FakeApollo) Open(_ context.Context, request LaunchRequest, _ RTSPResponse) (ProviderSession, error) { session := &fakeSession{ failure: f.config.Failure, video: make(chan []byte, 16), audio: make(chan []byte, 16), events: make(chan ProviderEvent, 16), clipboardWrites: make(chan string, 1), state: protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, Channels: []string{"video", "audio", "input", "feedback"}}, pressed: make(map[string]struct{}), } for _, payload := range f.config.Video { session.EmitVideo(payload) } for _, payload := range f.config.Audio { session.EmitAudio(payload) } f.mu.Lock() f.last = session f.mu.Unlock() return session, nil } func (f *FakeApollo) Start(ctx context.Context, request LaunchRequest) (ProviderSession, error) { adapter := NewApolloAdapter(f, f.config.Identity) adapter.now = func() time.Time { return f.config.Now } return adapter.Start(ctx, request) } func (f *FakeApollo) LastSession() ProviderSession { f.mu.Lock() defer f.mu.Unlock() return f.last } func (f *FakeApollo) DisconnectProvider() { f.mu.Lock() session := f.last f.mu.Unlock() if session != nil { session.Disconnect() } } type fakeSession struct { mu sync.Mutex failure FakeFailure video chan []byte audio chan []byte events chan ProviderEvent state protocol.ProviderState pressed map[string]struct{} inputs []InputEvent feedback []Feedback clipboard string clipboardWrites chan string releaseAll int closeOnce sync.Once } func (s *fakeSession) Ready(ctx context.Context) error { if s.failure == FakeFailureReadinessTimeout { <-ctx.Done() return ctx.Err() } s.mu.Lock() defer s.mu.Unlock() if s.state.State == ProviderStateDisconnected { return ErrProviderDisconnected } s.state.State = ProviderStateReady return nil } func (s *fakeSession) Video() <-chan []byte { return s.video } func (s *fakeSession) Audio() <-chan []byte { return s.audio } func (s *fakeSession) Events() <-chan ProviderEvent { return s.events } func (s *fakeSession) EmitEvent(event ProviderEvent) { select { case s.events <- ProviderEvent{Kind: event.Kind, Payload: append([]byte(nil), event.Payload...)}: default: } } func (s *fakeSession) EmitVideo(payload []byte) { s.mu.Lock() defer s.mu.Unlock() if s.state.State == ProviderStateTerminating || s.state.State == ProviderStateTerminated || s.state.State == ProviderStateDisconnected { return } select { case s.video <- append([]byte(nil), payload...): default: select { case <-s.video: default: } select { case s.video <- append([]byte(nil), payload...): default: } } } func (s *fakeSession) EmitAudio(payload []byte) { s.mu.Lock() defer s.mu.Unlock() if s.state.State == ProviderStateTerminating || s.state.State == ProviderStateTerminated || s.state.State == ProviderStateDisconnected { return } select { case s.audio <- append([]byte(nil), payload...): default: select { case <-s.audio: default: } select { case s.audio <- append([]byte(nil), payload...): default: } } } func (s *fakeSession) Input(_ context.Context, event InputEvent) error { s.mu.Lock() defer s.mu.Unlock() if s.state.State != ProviderStateReady { return ErrProviderDisconnected } s.inputs = append(s.inputs, event) key := fmt.Sprintf("%s:%d", event.Device, event.Code) if event.Pressed { s.pressed[key] = struct{}{} } else { delete(s.pressed, key) } return nil } func (s *fakeSession) Feedback(_ context.Context, feedback Feedback) error { s.mu.Lock() defer s.mu.Unlock() if s.state.State != ProviderStateReady { return ErrProviderDisconnected } s.feedback = append(s.feedback, feedback) return nil } func (s *fakeSession) ReadClipboard(ctx context.Context) (string, error) { if err := ctx.Err(); err != nil { return "", err } s.mu.Lock() defer s.mu.Unlock() if s.state.State != ProviderStateReady { return "", ErrProviderDisconnected } return s.clipboard, nil } func (s *fakeSession) WriteClipboard(ctx context.Context, value string) error { if err := ctx.Err(); err != nil { return err } s.mu.Lock() if s.state.State != ProviderStateReady { s.mu.Unlock() return ErrProviderDisconnected } s.clipboard = value s.mu.Unlock() select { case s.clipboardWrites <- value: default: } return nil } func (s *fakeSession) ReleaseAll(_ context.Context) error { s.mu.Lock() defer s.mu.Unlock() if s.failure == FakeFailureProviderDisconnect { return ErrProviderDisconnected } s.pressed = make(map[string]struct{}) s.releaseAll++ return nil } func (s *fakeSession) Terminate(ctx context.Context) error { if s.failure == FakeFailureTerminationTimeout { <-ctx.Done() s.mu.Lock() s.state.State = ProviderStateCleanup s.state.CleanupPending = true s.mu.Unlock() return ctx.Err() } s.mu.Lock() if s.state.State == ProviderStateTerminated { s.mu.Unlock() return nil } s.state.State = ProviderStateTerminating s.closeOnce.Do(func() { close(s.video) close(s.audio) }) s.state.State = ProviderStateTerminated s.mu.Unlock() return nil } func (s *fakeSession) State() protocol.ProviderState { s.mu.Lock() defer s.mu.Unlock() return s.state } func (s *fakeSession) Telemetry() ProviderTelemetry { return ProviderTelemetry{State: s.State().State} } func (s *fakeSession) Disconnect() { s.mu.Lock() s.state.State = ProviderStateDisconnected s.mu.Unlock() } func (s *fakeSession) ReleaseCount() int { s.mu.Lock() defer s.mu.Unlock() return s.releaseAll }