From 844e548c9522d2e222c3af4c7e6dc7278235a9a1 Mon Sep 17 00:00:00 2001
From: sechmachine <97589681+sechmachine727@users.noreply.github.com>
Date: Wed, 29 Jul 2026 11:27:31 +0700
Subject: [PATCH] feat(gateway): fetch sealed Apollo session work
---
cmd/verse-gateway/main.go | 16 ++------
gateway/apollo_native.go | 77 ++++++++++++++++++++++++++---------
gateway/apollo_native_test.go | 73 +++++++++++++++++++++++++++++++++
gateway/control_plane.go | 12 ++++++
gateway/gateway_test.go | 16 +++++++-
gateway/provider.go | 38 +++++++++++++----
gateway/transport.go | 25 +++++++++++-
go.mod | 2 +-
go.sum | 2 +
9 files changed, 217 insertions(+), 44 deletions(-)
create mode 100644 gateway/apollo_native_test.go
diff --git a/cmd/verse-gateway/main.go b/cmd/verse-gateway/main.go
index 0c3580b..0f56834 100644
--- a/cmd/verse-gateway/main.go
+++ b/cmd/verse-gateway/main.go
@@ -30,7 +30,6 @@ func main() {
func run() error {
var listen, advertiseAddress, controlPlane, certFile, keyFile, clientCAFile string
var gatewayID, instanceIdentity, certificateIdentity, publicIdentity string
- var providerManagement, providerRTSPAddress, providerRTSPURL, providerIdentity string
flag.StringVar(&listen, "listen", "0.0.0.0:443", "gateway QUIC listen address")
flag.StringVar(&advertiseAddress, "advertise-address", "", "client-visible gateway address host:port")
flag.StringVar(&controlPlane, "control-plane", "", "Connection Server HTTPS base URL")
@@ -41,12 +40,8 @@ func run() error {
flag.StringVar(&instanceIdentity, "instance-identity", "", "gateway instance identity")
flag.StringVar(&certificateIdentity, "certificate-identity", "", "gateway certificate identity")
flag.StringVar(&publicIdentity, "public-identity", "gateway", "gateway public identity")
- flag.StringVar(&providerManagement, "provider-management", "", "internal Apollo management URL")
- flag.StringVar(&providerRTSPAddress, "provider-rtsp-address", "", "internal Apollo RTSP address")
- flag.StringVar(&providerRTSPURL, "provider-rtsp-url", "", "internal Apollo RTSP URL")
- flag.StringVar(&providerIdentity, "provider-identity", "", "enrolled Apollo identity unique-id#fingerprint")
flag.Parse()
- for name, value := range map[string]string{"control-plane": controlPlane, "advertise-address": advertiseAddress, "cert": certFile, "key": keyFile, "client-ca": clientCAFile, "gateway-id": gatewayID, "instance-identity": instanceIdentity, "certificate-identity": certificateIdentity, "provider-management": providerManagement, "provider-rtsp-address": providerRTSPAddress, "provider-rtsp-url": providerRTSPURL, "provider-identity": providerIdentity} {
+ for name, value := range map[string]string{"control-plane": controlPlane, "advertise-address": advertiseAddress, "cert": certFile, "key": keyFile, "client-ca": clientCAFile, "gateway-id": gatewayID, "instance-identity": instanceIdentity, "certificate-identity": certificateIdentity} {
if value == "" {
return fmt.Errorf("-%s is required", name)
}
@@ -54,24 +49,19 @@ func run() error {
if err := validateAdvertisedAddress(advertiseAddress); err != nil {
return err
}
- expectedIdentity, err := parseProviderIdentity(providerIdentity)
- if err != nil {
- return err
- }
serverTLS, clientTLS, err := loadTLS(certFile, keyFile, clientCAFile)
if err != nil {
return err
}
transport := &http.Transport{TLSClientConfig: clientTLS}
controlPlaneClient := gateway.NewControlPlaneClient(controlPlane, &http.Client{Transport: transport, Timeout: 5 * time.Second})
- providerBackend := gateway.NewNativeApolloBackend(providerManagement, providerRTSPAddress, providerRTSPURL, &http.Client{Timeout: 5 * time.Second})
- provider := gateway.NewApolloAdapter(providerBackend, expectedIdentity)
+ provider := gateway.NewApolloAdapter(gateway.NewNativeApolloBackend(), gateway.ProviderIdentity{})
capabilities := gateway.DefaultCapabilities()
server, err := gateway.NewServer(gateway.ServerConfig{ListenAddress: listen, TLSConfig: serverTLS, GatewayID: gatewayID, Capabilities: capabilities, ProviderCapabilities: capabilities, Admission: controlPlaneClient, ProviderStateReporter: controlPlaneClient, Provider: provider, PacerKbps: 100000})
if err != nil {
return err
}
- registration := protocol.GatewayRegistration{Version: "1", GatewayID: gatewayID, InstanceIdentity: instanceIdentity, CertificateIdentity: certificateIdentity, PublicIdentity: publicIdentity, Address: advertiseAddress, ProviderIdentity: providerIdentity, ProtocolMinVersion: 1, ProtocolMaxVersion: 1, ConnectionCapacity: 8, BandwidthCapacityKbps: 100000, Features: []string{"quic-tls13", "datagram.media", "apollo"}, Capabilities: capabilities}
+ registration := protocol.GatewayRegistration{Version: "1", GatewayID: gatewayID, InstanceIdentity: instanceIdentity, CertificateIdentity: certificateIdentity, PublicIdentity: publicIdentity, Address: advertiseAddress, ProviderIdentity: "server-derived", ProtocolMinVersion: 1, ProtocolMaxVersion: 1, ConnectionCapacity: 8, BandwidthCapacityKbps: 100000, Features: []string{"quic-tls13", "datagram.media", "apollo"}, Capabilities: capabilities}
if _, err := controlPlaneClient.Register(context.Background(), registration); err != nil {
_ = server.Close()
return err
diff --git a/gateway/apollo_native.go b/gateway/apollo_native.go
index ec3376b..57e9c45 100644
--- a/gateway/apollo_native.go
+++ b/gateway/apollo_native.go
@@ -2,12 +2,17 @@ package gateway
import (
"bufio"
+ "bytes"
"context"
+ "crypto/sha256"
"crypto/tls"
+ "crypto/x509"
+ "encoding/hex"
"fmt"
"io"
"net"
"net/http"
+ "strconv"
"strings"
"sync"
"time"
@@ -16,36 +21,35 @@ import (
)
// NativeApolloBackend keeps provider sockets inside the gateway process. The
-// RTSP endpoint is configuration owned by the gateway and is never serialized
-// into a client manifest or authority.
+// session-scoped Server work is the sole source of provider endpoint and mTLS
+// material; it is never serialized into a client manifest or authority.
type NativeApolloBackend struct {
- ManagementURL string
- RTSPAddress string
- RTSPURL string
- HTTPClient *http.Client
- Dialer *net.Dialer
- TLSConfig *tls.Config
+ Dialer *net.Dialer
mu sync.Mutex
pending map[string]net.Conn
}
-func NewNativeApolloBackend(managementURL, rtspAddress, rtspURL string, client *http.Client) *NativeApolloBackend {
- if client == nil {
- client = &http.Client{Timeout: 5 * time.Second}
- }
- return &NativeApolloBackend{ManagementURL: managementURL, RTSPAddress: rtspAddress, RTSPURL: rtspURL, HTTPClient: client, Dialer: &net.Dialer{Timeout: 5 * time.Second}, pending: make(map[string]net.Conn)}
+func NewNativeApolloBackend() *NativeApolloBackend {
+ return &NativeApolloBackend{Dialer: &net.Dialer{Timeout: 5 * time.Second}, pending: make(map[string]net.Conn)}
}
-func (b *NativeApolloBackend) Management(ctx context.Context) ([]byte, error) {
- if b.ManagementURL == "" {
+func (b *NativeApolloBackend) Management(ctx context.Context, request LaunchRequest) ([]byte, error) {
+ work := request.ProviderWork
+ if err := work.Validate(); err != nil || work.ProviderProfile != ProviderProfileApollo {
return nil, ErrProviderMalformed
}
- request, err := http.NewRequestWithContext(ctx, http.MethodGet, b.ManagementURL, nil)
+ tlsConfig, err := pinnedApolloTLSConfig(work)
+ if err != nil {
+ return nil, err
+ }
+ client := &http.Client{Transport: &http.Transport{TLSClientConfig: tlsConfig}, Timeout: 5 * time.Second}
+ managementURL := "https://" + net.JoinHostPort(work.ManagementHost, strconv.FormatInt(work.ManagementPort, 10)) + "/serverinfo"
+ httpRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, managementURL, nil)
if err != nil {
return nil, ErrProviderMalformed
}
- response, err := b.HTTPClient.Do(request)
+ response, err := client.Do(httpRequest)
if err != nil {
return nil, err
}
@@ -56,18 +60,51 @@ func (b *NativeApolloBackend) Management(ctx context.Context) ([]byte, error) {
return readBounded(response.Body, 64*1024)
}
+func pinnedApolloTLSConfig(work protocol.ProviderSessionWork) (*tls.Config, error) {
+ identity, ok := providerIdentityFromKey(work.ProviderIdentity)
+ if !ok || !strings.HasPrefix(identity.Fingerprint, "sha256:") {
+ return nil, ErrProviderIdentity
+ }
+ pinned, err := hex.DecodeString(strings.TrimPrefix(identity.Fingerprint, "sha256:"))
+ if err != nil || len(pinned) != sha256.Size {
+ return nil, ErrProviderIdentity
+ }
+ certificate, err := tls.X509KeyPair([]byte(work.ClientCertificatePem), []byte(work.ClientPrivateKeyPem))
+ if err != nil {
+ return nil, ErrProviderIdentity
+ }
+ trust := x509.NewCertPool()
+ if !trust.AppendCertsFromPEM([]byte(work.ServerCertificatePem)) {
+ return nil, ErrProviderIdentity
+ }
+ return &tls.Config{
+ MinVersion: tls.VersionTLS13, Certificates: []tls.Certificate{certificate}, RootCAs: trust,
+ VerifyPeerCertificate: func(rawCertificates [][]byte, _ [][]*x509.Certificate) error {
+ if len(rawCertificates) == 0 {
+ return ErrProviderIdentity
+ }
+ digest := sha256.Sum256(rawCertificates[0])
+ if !bytes.Equal(digest[:], pinned) {
+ return ErrProviderIdentity
+ }
+ return nil
+ },
+ }, nil
+}
+
func (b *NativeApolloBackend) Setup(ctx context.Context, request LaunchRequest) ([]byte, error) {
- if b.RTSPAddress == "" || b.RTSPURL == "" || request.SessionID == "" {
+ work := request.ProviderWork
+ if err := work.Validate(); err != nil || request.SessionID == "" {
return nil, ErrProviderMalformed
}
- conn, err := b.Dialer.DialContext(ctx, "tcp", b.RTSPAddress)
+ conn, err := b.Dialer.DialContext(ctx, "tcp", net.JoinHostPort(work.StreamHost, strconv.FormatInt(work.StreamPort, 10)))
if err != nil {
return nil, err
}
if deadline, ok := ctx.Deadline(); ok {
_ = conn.SetDeadline(deadline)
}
- requestText := "SETUP " + b.RTSPURL + " RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\nSession: " + request.SessionID + "\r\n\r\n"
+ requestText := "SETUP rtsp://" + work.StreamHost + "/streamid=video/0/0 RTSP/1.0\r\nCSeq: 1\r\nTransport: RTP/AVP/TCP;interleaved=0-1\r\nSession: " + request.SessionID + "\r\n\r\n"
if _, err := io.WriteString(conn, requestText); err != nil {
_ = conn.Close()
return nil, err
diff --git a/gateway/apollo_native_test.go b/gateway/apollo_native_test.go
new file mode 100644
index 0000000..63a4543
--- /dev/null
+++ b/gateway/apollo_native_test.go
@@ -0,0 +1,73 @@
+package gateway
+
+import (
+ "context"
+ "crypto/sha256"
+ "crypto/tls"
+ "crypto/x509"
+ "encoding/hex"
+ "encoding/pem"
+ "net"
+ "net/http"
+ "net/http/httptest"
+ "strconv"
+ "testing"
+
+ protocol "git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol/gen/go/protocol"
+)
+
+func TestNativeApolloManagementUsesSessionScopedMTLS(t *testing.T) {
+ serverTLS, clientTLS := testTLS(t)
+ server := httptest.NewUnstartedServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
+ if request.URL.Path != "/serverinfo" || request.TLS == nil || len(request.TLS.PeerCertificates) != 1 {
+ http.Error(response, "mTLS required", http.StatusUnauthorized)
+ return
+ }
+ _, _ = response.Write([]byte("apollo-server"))
+ }))
+ server.TLS = serverTLS
+ server.StartTLS()
+ defer server.Close()
+ host, portText, err := net.SplitHostPort(server.Listener.Addr().String())
+ if err != nil {
+ t.Fatal(err)
+ }
+ port, err := strconv.ParseInt(portText, 10, 64)
+ if err != nil {
+ t.Fatal(err)
+ }
+ pinned := sha256.Sum256(serverTLS.Certificates[0].Certificate[0])
+ work := protocol.ProviderSessionWork{
+ Version: "1", SessionID: "session-1", GatewayID: "gateway-1", ReconnectSequence: 0,
+ ExpiresAt: "2099-01-01T00:00:00Z", ProviderProfile: ProviderProfileApollo,
+ ProviderIdentity: "apollo-server#sha256:" + hex.EncodeToString(pinned[:]), PolicyVersionID: "policy-1", ApplicationID: "1",
+ ManagementHost: host, ManagementPort: port, StreamHost: host, StreamPort: 47984,
+ ClientCertificatePem: certificatePEM(t, clientTLS.Certificates[0]),
+ ClientPrivateKeyPem: privateKeyPEM(t, clientTLS.Certificates[0]),
+ ServerCertificatePem: certificatePEM(t, tls.Certificate{Certificate: [][]byte{serverTLS.Certificates[0].Certificate[1]}}),
+ }
+ data, err := NewNativeApolloBackend().Management(context.Background(), LaunchRequest{SessionID: "session-1", ProviderProfile: ProviderProfileApollo, ProviderWork: work})
+ if err != nil {
+ t.Fatalf("Management() error = %v", err)
+ }
+ if info, err := ParseManagementXML(data); err != nil || info.Identity.UniqueID != "apollo-server" {
+ t.Fatalf("ParseManagementXML() = %+v, %v", info, err)
+ }
+}
+
+func certificatePEM(t *testing.T, certificate tls.Certificate) string {
+ t.Helper()
+ if len(certificate.Certificate) == 0 {
+ t.Fatal("expected certificate")
+ }
+ return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certificate.Certificate[0]}))
+}
+
+func privateKeyPEM(t *testing.T, certificate tls.Certificate) string {
+ t.Helper()
+ encoded, err := x509.MarshalPKCS8PrivateKey(certificate.PrivateKey)
+ if err != nil {
+ t.Fatal(err)
+ }
+ return string(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: encoded}))
+}
diff --git a/gateway/control_plane.go b/gateway/control_plane.go
index 25c8b19..dc048d2 100644
--- a/gateway/control_plane.go
+++ b/gateway/control_plane.go
@@ -76,6 +76,18 @@ func (c *ControlPlaneClient) Release(ctx context.Context, authority protocol.Ses
return err
}
+func (c *ControlPlaneClient) ProviderWork(ctx context.Context, authority protocol.SessionAuthority) (protocol.ProviderSessionWork, error) {
+ payload, err := protocol.EncodeSessionAuthority(authority)
+ if err != nil {
+ return protocol.ProviderSessionWork{}, err
+ }
+ response, err := c.post(ctx, "/api/v1/gateway/provider-work", payload)
+ if err != nil {
+ return protocol.ProviderSessionWork{}, err
+ }
+ return protocol.DecodeProviderSessionWork(response)
+}
+
func (c *ControlPlaneClient) ReportProviderState(ctx context.Context, state protocol.ProviderState) error {
payload, err := protocol.EncodeProviderState(state)
if err != nil {
diff --git a/gateway/gateway_test.go b/gateway/gateway_test.go
index e5733e2..561db78 100644
--- a/gateway/gateway_test.go
+++ b/gateway/gateway_test.go
@@ -307,7 +307,7 @@ func testTLS(t *testing.T) (*tls.Config, *tls.Config) {
if keyErr != nil {
t.Fatal(keyErr)
}
- template := &x509.Certificate{SerialNumber: big.NewInt(serial), Subject: pkix.Name{CommonName: dns}, DNSNames: []string{dns}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), ExtKeyUsage: []x509.ExtKeyUsage{usage}, KeyUsage: x509.KeyUsageDigitalSignature}
+ template := &x509.Certificate{SerialNumber: big.NewInt(serial), Subject: pkix.Name{CommonName: dns}, DNSNames: []string{dns}, IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, NotBefore: time.Now().Add(-time.Hour), NotAfter: time.Now().Add(time.Hour), ExtKeyUsage: []x509.ExtKeyUsage{usage}, KeyUsage: x509.KeyUsageDigitalSignature}
der, createErr := x509.CreateCertificate(rand.Reader, template, caCert, &key.PublicKey, caKey)
if createErr != nil {
t.Fatal(createErr)
@@ -353,6 +353,20 @@ func (a *oneTimeAdmission) Admit(context.Context, protocol.TunnelAdmissionReques
return a.authority, nil
}
+func (a *oneTimeAdmission) ProviderWork(_ context.Context, authority protocol.SessionAuthority) (protocol.ProviderSessionWork, error) {
+ if authority != a.authority {
+ return protocol.ProviderSessionWork{}, ErrAdmissionRejected
+ }
+ return protocol.ProviderSessionWork{
+ Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID,
+ ReconnectSequence: authority.ReconnectSequence, ExpiresAt: authority.ExpiresAt,
+ ProviderProfile: ProviderProfileApollo, ProviderIdentity: authority.ProviderIdentity,
+ PolicyVersionID: "policy-1", ApplicationID: "1", ManagementHost: "apollo.test", ManagementPort: 47990,
+ StreamHost: "apollo.test", StreamPort: 47984, ClientCertificatePem: "certificate",
+ ClientPrivateKeyPem: "private-key", ServerCertificatePem: "server-certificate",
+ }, nil
+}
+
func (a *oneTimeAdmission) Release(context.Context, protocol.SessionAuthority) error {
if a.releases.Add(1) == 1 {
close(a.released)
diff --git a/gateway/provider.go b/gateway/provider.go
index c7fe839..28eda6b 100644
--- a/gateway/provider.go
+++ b/gateway/provider.go
@@ -44,8 +44,17 @@ 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 == "" || i.Fingerprint == "" || i.UniqueID != expected.UniqueID || i.Fingerprint != expected.Fingerprint {
+ 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) {
@@ -68,7 +77,8 @@ func ParseManagementXML(data []byte) (ManagementInfo, error) {
}
var document struct {
XMLName xml.Name `xml:"root"`
- UniqueID string `xml:"unique_id"`
+ UniqueID string `xml:"uniqueid"`
+ LegacyID string `xml:"unique_id"`
Fingerprint string `xml:"fingerprint"`
NotBefore string `xml:"not_before"`
NotAfter string `xml:"not_after"`
@@ -79,6 +89,9 @@ func ParseManagementXML(data []byte) (ManagementInfo, error) {
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 != "" {
@@ -93,7 +106,7 @@ func ParseManagementXML(data []byte) (ManagementInfo, error) {
return ManagementInfo{}, ErrProviderMalformed
}
}
- if identity.UniqueID == "" || len(identity.UniqueID) > 128 || identity.Fingerprint == "" || len(identity.Fingerprint) > 256 {
+ if identity.UniqueID == "" || len(identity.UniqueID) > 128 || len(identity.Fingerprint) > 256 {
return ManagementInfo{}, ErrProviderMalformed
}
return ManagementInfo{Identity: identity, Name: document.Name}, nil
@@ -181,6 +194,7 @@ type LaunchRequest struct {
Capabilities protocol.CapabilityProfile
ProviderProfile string
ProviderIdentity string
+ ProviderWork protocol.ProviderSessionWork
}
type InputEvent struct {
@@ -213,7 +227,7 @@ type ProviderSession interface {
}
type ApolloBackend interface {
- Management(context.Context) ([]byte, error)
+ Management(context.Context, LaunchRequest) ([]byte, error)
Setup(context.Context, LaunchRequest) ([]byte, error)
Open(context.Context, LaunchRequest, RTSPResponse) (ProviderSession, error)
}
@@ -233,7 +247,7 @@ func (a *ApolloAdapter) Start(ctx context.Context, request LaunchRequest) (Provi
if a == nil || a.backend == nil || request.ProviderProfile != ProviderProfileApollo {
return nil, ErrProviderIdentity
}
- management, err := a.backend.Management(ctx)
+ management, err := a.backend.Management(ctx, request)
if err != nil {
return nil, err
}
@@ -241,10 +255,18 @@ func (a *ApolloAdapter) Start(ctx context.Context, request LaunchRequest) (Provi
if err != nil {
return nil, err
}
- if err := info.Identity.Validate(a.now(), a.expected); err != nil {
+ 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 != "" && request.ProviderIdentity != info.Identity.Key() {
+ if request.ProviderIdentity != "" && info.Identity.UniqueID != expected.UniqueID {
return nil, ErrProviderIdentity
}
rawRTSP, err := a.backend.Setup(ctx, request)
@@ -320,7 +342,7 @@ func NewFakeApollo(config FakeApolloConfig) *FakeApollo {
return &FakeApollo{config: config}
}
-func (f *FakeApollo) Management(context.Context) ([]byte, error) {
+func (f *FakeApollo) Management(context.Context, LaunchRequest) ([]byte, error) {
if f.config.Failure == FakeFailureMalformed {
return []byte(""), nil
}
diff --git a/gateway/transport.go b/gateway/transport.go
index 39c9a6d..4c11175 100644
--- a/gateway/transport.go
+++ b/gateway/transport.go
@@ -33,6 +33,7 @@ var (
type Admission interface {
Admit(context.Context, protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error)
+ ProviderWork(context.Context, protocol.SessionAuthority) (protocol.ProviderSessionWork, error)
Release(context.Context, protocol.SessionAuthority) error
}
@@ -44,6 +45,10 @@ func (f AdmissionFunc) Admit(ctx context.Context, request protocol.TunnelAdmissi
func (AdmissionFunc) Release(context.Context, protocol.SessionAuthority) error { return nil }
+func (AdmissionFunc) ProviderWork(context.Context, protocol.SessionAuthority) (protocol.ProviderSessionWork, error) {
+ return protocol.ProviderSessionWork{}, ErrAdmissionRejected
+}
+
type ProviderStateReporter interface {
ReportProviderState(context.Context, protocol.ProviderState) error
}
@@ -210,6 +215,12 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
_ = writeStableError(stream, "invalid_authority", err, false)
return
}
+ work, err := s.config.Admission.ProviderWork(ctx, authority)
+ if err != nil || s.validateProviderWork(work, authority) != nil {
+ _ = s.config.Admission.Release(context.Background(), authority)
+ _ = writeStableError(stream, "provider_work_unavailable", ErrAdmissionRejected, err != nil)
+ return
+ }
selected, err := IntersectCapabilities(s.config.Capabilities, s.config.ProviderCapabilities, request.Capabilities, authority.Capabilities)
if err != nil {
_ = s.config.Admission.Release(context.Background(), authority)
@@ -222,7 +233,7 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
_ = writeStableError(stream, "provider_state_unavailable", err, true)
return
}
- providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: authority.ProviderIdentity})
+ providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: work.ProviderIdentity, ProviderWork: work})
if err != nil {
s.metrics.ProviderErrors.Add(1)
_ = s.reportProviderState(context.Background(), protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateFailed, CleanupPending: false, Channels: []string{"video", "audio", "input", "feedback"}})
@@ -279,6 +290,18 @@ func (s *Server) validateAuthority(authority protocol.SessionAuthority, request
return nil
}
+func (s *Server) validateProviderWork(work protocol.ProviderSessionWork, authority protocol.SessionAuthority) error {
+ if err := work.Validate(); err != nil {
+ return err
+ }
+ if work.SessionID != authority.SessionID || work.GatewayID != authority.GatewayID ||
+ work.ReconnectSequence != authority.ReconnectSequence || work.ExpiresAt != authority.ExpiresAt ||
+ work.ProviderProfile != authority.ProviderProfile {
+ return ErrAdmissionRejected
+ }
+ return nil
+}
+
func (s *Server) addSession(session *gatewaySession) {
s.mu.Lock()
s.sessions[session] = struct{}{}
diff --git a/go.mod b/go.mod
index 6c0f698..03c9996 100644
--- a/go.mod
+++ b/go.mod
@@ -3,7 +3,7 @@ module git.sechmachine.io.vn/sechmachine/VerseVDI-Data-Plane
go 1.26.5
require (
- git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol v1.0.0-phase3c-gateway-rc.3
+ git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol v1.0.0-phase3c-gateway-rc.4
github.com/quic-go/quic-go v0.61.0
)
diff --git a/go.sum b/go.sum
index 68d8533..689d616 100644
--- a/go.sum
+++ b/go.sum
@@ -1,5 +1,7 @@
git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol v1.0.0-phase3c-gateway-rc.3 h1:ZoXbg9CRwlypVbDO0EaXwHVOKTGlIfZDC7s/4JuOISE=
git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol v1.0.0-phase3c-gateway-rc.3/go.mod h1:7PhFIDhjtr20btWoEb2GqB+7dBpzJt43olrnHVutWoc=
+git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol v1.0.0-phase3c-gateway-rc.4 h1:uwNoKtzRlpdbvq9kHblncKr+K+dck0ydtozzPKItQzA=
+git.sechmachine.io.vn/sechmachine/VerseVDI-Protocol v1.0.0-phase3c-gateway-rc.4/go.mod h1:7PhFIDhjtr20btWoEb2GqB+7dBpzJt43olrnHVutWoc=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=