feat(core): add QUIC TLS admission transport

This commit is contained in:
sechmachine
2026-08-12 19:08:32 +07:00
parent 7947ebcc75
commit 111092becb
11 changed files with 2304 additions and 38 deletions
+78
View File
@@ -14,6 +14,7 @@ import (
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"math/big"
"net"
@@ -476,6 +477,83 @@ func TestAdmissionQUICMTLSRelayAndCleanup(t *testing.T) {
}
}
func TestStableErrorFramePrecedesConnectionTeardown(t *testing.T) {
for _, retryable := range []bool{false, true} {
t.Run(fmt.Sprintf("retryable=%t", retryable), func(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
authority := protocol.SessionAuthority{
Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway",
ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(),
ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key(),
}
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: &oneTimeAdmission{authority: authority, released: make(chan struct{})}, Provider: fake})
if err != nil {
t.Fatal(err)
}
if retryable {
server.BeginDrain()
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { _ = server.Serve(ctx) }()
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
if err != nil {
t.Fatal(err)
}
stream, err := connection.OpenStreamSync(context.Background())
if err != nil {
t.Fatal(err)
}
gatewayID := authority.GatewayID
if !retryable {
gatewayID = "wrong-gateway"
}
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: gatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
payload, err := protocol.EncodeTunnelAdmissionRequest(request)
if err != nil {
t.Fatal(err)
}
if err := writeWire(stream, payload, defaultHelloLimit); err != nil {
t.Fatal(err)
}
response, err := readWire(stream, defaultHelloLimit)
if err != nil {
t.Fatalf("stable response lost before connection teardown: %v", err)
}
stable, err := protocol.DecodeStableError(response)
if err != nil || stable.Retryable != retryable {
t.Fatalf("stable response = %#v, %v", stable, err)
}
select {
case <-connection.Context().Done():
case <-time.After(2 * time.Second):
t.Fatal("server retained rejected connection")
}
_ = server.Close()
})
}
}
func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
var wire bytes.Buffer
if err := writeStableError(&wire, "provider_unavailable", errors.New("https://provider.invalid/launch?rikey=secret-sentinel"), true); err != nil {
t.Fatal(err)
}
response, err := readWire(&wire, defaultHelloLimit)
if err != nil {
t.Fatal(err)
}
stable, err := protocol.DecodeStableError(response)
if err != nil {
t.Fatal(err)
}
if strings.Contains(stable.Message, "provider.invalid") || strings.Contains(stable.Message, "secret-sentinel") {
t.Fatalf("stable error leaked internal details: %q", stable.Message)
}
}
func TestGatewayTelemetrySeparatesQueueProcessingAndPacing(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
session := &fakeSession{
+50 -20
View File
@@ -206,74 +206,84 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
if err != nil {
return
}
writeError := func(code string, err error, retryable bool) {
if writeStableError(stream, code, err, retryable) == nil && stream.Close() == nil {
responseCtx, responseCancel := context.WithTimeout(ctx, time.Second)
defer responseCancel()
select {
case <-connection.Context().Done():
case <-responseCtx.Done():
}
}
}
requestBytes, err := readWire(stream, defaultHelloLimit)
if err != nil {
_ = writeStableError(stream, "invalid_hello", err, false)
writeError("invalid_hello", err, false)
return
}
request, err := protocol.DecodeTunnelAdmissionRequest(requestBytes)
if err != nil {
_ = writeStableError(stream, "invalid_hello", err, false)
writeError("invalid_hello", err, false)
return
}
if s.Draining() {
_ = writeStableError(stream, "gateway_draining", ErrGatewayDraining, true)
writeError("gateway_draining", ErrGatewayDraining, true)
return
}
if request.GatewayID != s.config.GatewayID {
_ = writeStableError(stream, "wrong_gateway", ErrAdmissionRejected, false)
writeError("wrong_gateway", ErrAdmissionRejected, false)
return
}
authority, err := s.config.Admission.Admit(ctx, request)
if err != nil {
s.metrics.AdmissionRejects.Add(1)
_ = writeStableError(stream, stableAdmissionCode(err), err, errors.Is(err, context.DeadlineExceeded))
writeError(stableAdmissionCode(err), err, errors.Is(err, context.DeadlineExceeded))
return
}
if s.Draining() {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "gateway_draining", ErrGatewayDraining, true)
writeError("gateway_draining", ErrGatewayDraining, true)
return
}
if err := s.validateAuthority(authority, request); err != nil {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "invalid_authority", err, false)
writeError("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)
writeError("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)
s.metrics.AdmissionRejects.Add(1)
_ = writeStableError(stream, "no_capability_overlap", err, false)
writeError("no_capability_overlap", err, false)
return
}
selected, err = selectApolloPolicyCapabilities(work.StreamPolicy, selected)
if err != nil {
_ = s.config.Admission.Release(context.Background(), authority)
s.metrics.AdmissionRejects.Add(1)
_ = writeStableError(stream, "no_capability_overlap", err, false)
writeError("no_capability_overlap", err, false)
return
}
clipboard, err := newClipboardGate(work.ClipboardPolicy, time.Now)
if err != nil {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "provider_work_unavailable", ErrAdmissionRejected, false)
writeError("provider_work_unavailable", ErrAdmissionRejected, false)
return
}
if (work.ClipboardPolicy.ClientToProviderEnabled || work.ClipboardPolicy.ProviderToClientEnabled) && s.config.ClipboardAuditReporter == nil {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "clipboard_audit_unavailable", ErrAdmissionRejected, true)
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, true)
return
}
if err := s.reportProviderState(ctx, protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, CleanupPending: false, Channels: []string{"video", "audio", "input", "feedback"}}); err != nil {
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "provider_state_unavailable", err, true)
writeError("provider_state_unavailable", err, true)
return
}
providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: work.ProviderIdentity, ProviderWork: work})
@@ -281,14 +291,14 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
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"}})
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, stableProviderCode(err), err, errors.Is(err, context.DeadlineExceeded))
writeError(stableProviderCode(err), err, errors.Is(err, context.DeadlineExceeded))
return
}
if err := s.reportProviderState(ctx, providerSession.State()); err != nil {
_ = providerSession.ReleaseAll(context.Background())
_ = providerSession.Terminate(context.Background())
_ = s.config.Admission.Release(context.Background(), authority)
_ = writeStableError(stream, "provider_state_unavailable", err, true)
writeError("provider_state_unavailable", err, true)
return
}
clientAuthority := protocol.ClientSessionAuthority{
@@ -943,11 +953,8 @@ func (s *gatewaySession) cleanup() {
})
}
func writeStableError(writer io.Writer, code string, err error, retryable bool) error {
message := err.Error()
if len(message) > 256 {
message = message[:256]
}
func writeStableError(writer io.Writer, code string, _ error, retryable bool) error {
message := stableErrorMessage(code)
payload, encodeErr := protocol.EncodeStableError(protocol.StableError{Version: "1", Code: code, Message: message, Retryable: retryable})
if encodeErr != nil {
return encodeErr
@@ -955,6 +962,29 @@ func writeStableError(writer io.Writer, code string, err error, retryable bool)
return writeWire(writer, payload, defaultHelloLimit)
}
func stableErrorMessage(code string) string {
switch code {
case "invalid_hello":
return "invalid client hello"
case "gateway_draining":
return "gateway is draining"
case "wrong_gateway":
return "gateway does not match admission request"
case "admission_rejected", "expired_grant":
return "admission rejected"
case "invalid_authority":
return "invalid session authority"
case "no_capability_overlap":
return "no compatible capability"
case "clipboard_audit_unavailable":
return "clipboard audit unavailable"
case "provider_work_unavailable", "provider_identity_rejected", "provider_malformed", "provider_timeout", "provider_unavailable", "provider_state_unavailable":
return "provider unavailable"
default:
return "request failed"
}
}
func stableAdmissionCode(err error) string {
if errors.Is(err, ErrGatewayDraining) {
return "gateway_draining"