fix(transport): harden QUIC admission boundaries
This commit is contained in:
+56
-5
@@ -1,6 +1,6 @@
|
||||
use std::fmt;
|
||||
use std::io::Cursor;
|
||||
use std::sync::Arc;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use rustls::client::ResolvesClientCert;
|
||||
use rustls::pki_types::CertificateDer;
|
||||
@@ -13,7 +13,7 @@ use crate::wire::NativeTunnelCredential;
|
||||
pub(crate) type SignCallback = dyn Fn(&[u8]) -> Result<[u8; 64]> + Send + Sync;
|
||||
|
||||
pub(crate) struct CallbackSigningKey {
|
||||
callback: Arc<SignCallback>,
|
||||
callback: Arc<Mutex<Option<Arc<SignCallback>>>>,
|
||||
}
|
||||
|
||||
impl fmt::Debug for CallbackSigningKey {
|
||||
@@ -24,7 +24,9 @@ impl fmt::Debug for CallbackSigningKey {
|
||||
|
||||
impl CallbackSigningKey {
|
||||
pub(crate) fn new(callback: Arc<SignCallback>) -> Self {
|
||||
Self { callback }
|
||||
Self {
|
||||
callback: Arc::new(Mutex::new(Some(callback))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -40,7 +42,7 @@ impl SigningKey for CallbackSigningKey {
|
||||
}
|
||||
}
|
||||
|
||||
struct CallbackSigner(Arc<SignCallback>);
|
||||
struct CallbackSigner(Arc<Mutex<Option<Arc<SignCallback>>>>);
|
||||
|
||||
impl fmt::Debug for CallbackSigner {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
@@ -50,7 +52,13 @@ impl fmt::Debug for CallbackSigner {
|
||||
|
||||
impl Signer for CallbackSigner {
|
||||
fn sign(&self, message: &[u8]) -> std::result::Result<Vec<u8>, rustls::Error> {
|
||||
(self.0)(message)
|
||||
let callback = self
|
||||
.0
|
||||
.lock()
|
||||
.map_err(|_| rustls::Error::General("client signing state failed".to_owned()))?
|
||||
.take()
|
||||
.ok_or_else(|| rustls::Error::General("client signer already used".to_owned()))?;
|
||||
callback(message)
|
||||
.map(|signature| signature.to_vec())
|
||||
.map_err(|_| rustls::Error::General("client signing failed".to_owned()))
|
||||
}
|
||||
@@ -112,3 +120,46 @@ pub(crate) fn client_config(
|
||||
config.enable_early_data = false;
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
use super::*;
|
||||
|
||||
struct DropMarker(Arc<AtomicBool>);
|
||||
|
||||
impl Drop for DropMarker {
|
||||
fn drop(&mut self) {
|
||||
self.0.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tls_callback_is_consumed_and_released_after_one_signature() {
|
||||
let dropped = Arc::new(AtomicBool::new(false));
|
||||
let marker = DropMarker(Arc::clone(&dropped));
|
||||
let callback: Arc<SignCallback> = Arc::new(move |_| {
|
||||
let _ = ▮
|
||||
Ok([7; 64])
|
||||
});
|
||||
let key = CallbackSigningKey::new(callback);
|
||||
let signer = key
|
||||
.choose_scheme(&[SignatureScheme::ED25519])
|
||||
.expect("ED25519 signer");
|
||||
let second_signer = key
|
||||
.choose_scheme(&[SignatureScheme::ED25519])
|
||||
.expect("second ED25519 signer");
|
||||
drop(key);
|
||||
|
||||
assert_eq!(
|
||||
signer.sign(b"handshake").expect("first signature"),
|
||||
vec![7; 64]
|
||||
);
|
||||
assert!(dropped.load(Ordering::SeqCst), "callback remained retained");
|
||||
assert!(
|
||||
second_signer.sign(b"second request").is_err(),
|
||||
"signer was reusable"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+87
-17
@@ -1,4 +1,5 @@
|
||||
use std::fmt;
|
||||
use std::io;
|
||||
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
@@ -157,15 +158,6 @@ pub async fn connect_with_cancellation(
|
||||
}
|
||||
Ok(signature)
|
||||
});
|
||||
let rustls = client_config(credential, checked_tls_signer)?;
|
||||
let crypto = QuicClientConfig::try_from(rustls).map_err(|_| CoreError::Tls)?;
|
||||
let mut quinn_config = quinn::ClientConfig::new(Arc::new(crypto));
|
||||
let mut transport = TransportConfig::default();
|
||||
transport.max_concurrent_bidi_streams(VarInt::from_u32(1));
|
||||
transport.max_concurrent_uni_streams(VarInt::from_u32(0));
|
||||
transport.datagram_receive_buffer_size(Some(64 * 1024));
|
||||
quinn_config.transport_config(Arc::new(transport));
|
||||
|
||||
let remaining = deadline
|
||||
.checked_sub(started.elapsed())
|
||||
.ok_or(CoreError::Cancelled)?;
|
||||
@@ -173,16 +165,25 @@ pub async fn connect_with_cancellation(
|
||||
let address_count = manifest.addresses().len();
|
||||
for (address_index, address) in manifest.addresses().iter().enumerate() {
|
||||
cancellation.check()?;
|
||||
let Ok(resolved) = tokio::net::lookup_host(address).await else {
|
||||
continue;
|
||||
};
|
||||
let resolved = resolved.take(8).collect::<Vec<_>>();
|
||||
let address_budget = deadline.saturating_sub(started.elapsed())
|
||||
/ u32::try_from(address_count.saturating_sub(address_index).max(1)).unwrap_or(1);
|
||||
let address_deadline = Instant::now() + address_budget;
|
||||
let resolved = match bounded_lookup(
|
||||
address_budget,
|
||||
cancellation,
|
||||
tokio::net::lookup_host(address),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(resolved) => resolved,
|
||||
Err(CoreError::Transport) => continue,
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
let remote_count = resolved.len();
|
||||
for (remote_index, remote) in resolved.into_iter().enumerate() {
|
||||
cancellation.check()?;
|
||||
let quinn_config =
|
||||
client_transport_config(credential, Arc::clone(&checked_tls_signer))?;
|
||||
let remaining_remotes = remote_count.saturating_sub(remote_index).max(1);
|
||||
let attempt_budget = address_deadline.saturating_duration_since(Instant::now())
|
||||
/ u32::try_from(remaining_remotes).unwrap_or(1);
|
||||
@@ -222,6 +223,21 @@ pub async fn connect_with_cancellation(
|
||||
}
|
||||
}
|
||||
|
||||
fn client_transport_config(
|
||||
credential: &NativeTunnelCredential,
|
||||
callback: Arc<SignCallback>,
|
||||
) -> Result<quinn::ClientConfig> {
|
||||
let rustls = client_config(credential, callback)?;
|
||||
let crypto = QuicClientConfig::try_from(rustls).map_err(|_| CoreError::Tls)?;
|
||||
let mut config = quinn::ClientConfig::new(Arc::new(crypto));
|
||||
let mut transport = TransportConfig::default();
|
||||
transport.max_concurrent_bidi_streams(VarInt::from_u32(1));
|
||||
transport.max_concurrent_uni_streams(VarInt::from_u32(0));
|
||||
transport.datagram_receive_buffer_size(Some(64 * 1024));
|
||||
config.transport_config(Arc::new(transport));
|
||||
Ok(config)
|
||||
}
|
||||
|
||||
struct AttemptContext<'a> {
|
||||
manifest: &'a ConnectionManifest,
|
||||
offered: &'a CapabilityProfile,
|
||||
@@ -253,6 +269,20 @@ async fn cancellable_timeout<T>(
|
||||
}
|
||||
}
|
||||
|
||||
async fn bounded_lookup<T>(
|
||||
duration: Duration,
|
||||
cancellation: &Cancellation,
|
||||
future: impl std::future::Future<Output = io::Result<T>>,
|
||||
) -> Result<Vec<SocketAddr>>
|
||||
where
|
||||
T: Iterator<Item = SocketAddr>,
|
||||
{
|
||||
let resolved = cancellable_timeout(duration, cancellation, future)
|
||||
.await?
|
||||
.map_err(|_| CoreError::Transport)?;
|
||||
Ok(resolved.take(8).collect())
|
||||
}
|
||||
|
||||
async fn dial(
|
||||
remote: SocketAddr,
|
||||
config: quinn::ClientConfig,
|
||||
@@ -270,7 +300,7 @@ async fn dial(
|
||||
let negotiated = connecting
|
||||
.handshake_data()
|
||||
.await
|
||||
.map_err(|_| AttemptError::Terminal(CoreError::Tls))?
|
||||
.map_err(|_| AttemptError::Terminal(tls_error(context)))?
|
||||
.downcast::<quinn::crypto::rustls::HandshakeData>()
|
||||
.ok()
|
||||
.and_then(|data| data.protocol.clone());
|
||||
@@ -279,7 +309,7 @@ async fn dial(
|
||||
}
|
||||
let connection = connecting
|
||||
.await
|
||||
.map_err(|_| AttemptError::Terminal(CoreError::Tls))?;
|
||||
.map_err(|_| AttemptError::Terminal(tls_error(context)))?;
|
||||
let payload = admission_payload(
|
||||
context.manifest,
|
||||
context.offered,
|
||||
@@ -302,7 +332,7 @@ async fn dial(
|
||||
.map_err(|_| AttemptError::Terminal(connection_error(&connection)))?;
|
||||
let Ok(authority) = ClientSessionAuthority::decode(&response) else {
|
||||
let stable = decode_stable_error(&response)?;
|
||||
return if stable.retryable {
|
||||
return if stable.retryable && stable.code == "gateway_draining" {
|
||||
Err(AttemptError::Retry)
|
||||
} else {
|
||||
Err(AttemptError::Terminal(stable.error))
|
||||
@@ -318,6 +348,14 @@ async fn dial(
|
||||
})
|
||||
}
|
||||
|
||||
fn tls_error(context: &AttemptContext<'_>) -> CoreError {
|
||||
if context.cancellation.check().is_err() || context.started.elapsed() >= context.deadline {
|
||||
CoreError::Cancelled
|
||||
} else {
|
||||
CoreError::Tls
|
||||
}
|
||||
}
|
||||
|
||||
fn admission_payload(
|
||||
manifest: &ConnectionManifest,
|
||||
offered: &CapabilityProfile,
|
||||
@@ -410,7 +448,11 @@ fn hello_length(length: usize) -> Result<u32> {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{hello_length, HELLO_LIMIT};
|
||||
use std::io;
|
||||
use std::net::SocketAddr;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use super::{bounded_lookup, hello_length, Cancellation, HELLO_LIMIT};
|
||||
use crate::error::CoreError;
|
||||
|
||||
#[test]
|
||||
@@ -419,4 +461,32 @@ mod tests {
|
||||
assert_eq!(hello_length(0), Err(CoreError::Protocol));
|
||||
assert_eq!(hello_length(HELLO_LIMIT + 1), Err(CoreError::Protocol));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dns_lookup_is_bounded_by_address_share() {
|
||||
let started = Instant::now();
|
||||
let result = tokio::runtime::Runtime::new()
|
||||
.expect("runtime")
|
||||
.block_on(bounded_lookup(
|
||||
Duration::from_millis(20),
|
||||
&Cancellation::new(),
|
||||
std::future::pending::<io::Result<std::vec::IntoIter<SocketAddr>>>(),
|
||||
));
|
||||
assert_eq!(result.err(), Some(CoreError::Transport));
|
||||
assert!(started.elapsed() < Duration::from_secs(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cancellation_interrupts_dns_lookup() {
|
||||
let cancellation = Cancellation::new();
|
||||
cancellation.cancel();
|
||||
let result = tokio::runtime::Runtime::new()
|
||||
.expect("runtime")
|
||||
.block_on(bounded_lookup(
|
||||
Duration::from_secs(1),
|
||||
&cancellation,
|
||||
std::future::pending::<io::Result<std::vec::IntoIter<SocketAddr>>>(),
|
||||
));
|
||||
assert_eq!(result.err(), Some(CoreError::Cancelled));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -388,6 +388,7 @@ impl ConnectionManifest {
|
||||
.iter()
|
||||
.any(|address| !bounded(address, 1, 256))
|
||||
|| !valid_dns_name(&self.gateway.public_identity)
|
||||
|| self.gateway.public_identity == self.gateway.id
|
||||
|| !(1..=4).contains(&self.tunnel.versions.len())
|
||||
|| self
|
||||
.tunnel
|
||||
@@ -621,6 +622,7 @@ struct StableError {
|
||||
|
||||
pub(crate) struct DecodedStableError {
|
||||
pub(crate) error: CoreError,
|
||||
pub(crate) code: String,
|
||||
pub(crate) retryable: bool,
|
||||
}
|
||||
|
||||
@@ -703,6 +705,7 @@ pub(crate) fn decode_stable_error(bytes: &[u8]) -> Result<DecodedStableError> {
|
||||
};
|
||||
Ok(DecodedStableError {
|
||||
error,
|
||||
code: stable.code,
|
||||
retryable: stable.retryable,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -46,6 +46,7 @@ type admission struct {
|
||||
lastTranscript []byte
|
||||
reusable bool
|
||||
failWorkOnce bool
|
||||
releaseFails bool
|
||||
authority protocol.SessionAuthority
|
||||
work protocol.ProviderSessionWork
|
||||
public ed25519.PublicKey
|
||||
@@ -67,11 +68,15 @@ func (a *admission) ProviderWork(context.Context, protocol.SessionAuthority) (pr
|
||||
if a.failWorkOnce { a.failWorkOnce = false; return protocol.ProviderSessionWork{}, context.DeadlineExceeded }
|
||||
return a.work, nil
|
||||
}
|
||||
func (a *admission) Release(context.Context, protocol.SessionAuthority) error { return nil }
|
||||
func (a *admission) Release(context.Context, protocol.SessionAuthority) error {
|
||||
if a.releaseFails { return errors.New("release failed") }
|
||||
return nil
|
||||
}
|
||||
|
||||
type provider struct{}
|
||||
type provider struct { failStart bool }
|
||||
type session struct { state protocol.ProviderState; video chan gateway.ProviderMedia; audio chan gateway.ProviderMedia; events chan gateway.ProviderEvent }
|
||||
func (provider) Start(_ context.Context, request gateway.LaunchRequest) (gateway.ProviderSession, error) {
|
||||
func (p provider) Start(_ context.Context, request gateway.LaunchRequest) (gateway.ProviderSession, error) {
|
||||
if p.failStart { return nil, context.DeadlineExceeded }
|
||||
return &session{state: protocol.ProviderState{Version:"1", SessionID:request.SessionID, State:gateway.ProviderStateReady, Channels:[]string{"video","audio","input","feedback"}}, video:make(chan gateway.ProviderMedia), audio:make(chan gateway.ProviderMedia), events:make(chan gateway.ProviderEvent)}, nil
|
||||
}
|
||||
func (s *session) Ready(context.Context) error { return nil }
|
||||
@@ -125,9 +130,9 @@ func main() {
|
||||
_, firstReplayErr := replayGuard.Admit(context.Background(),replayRequest)
|
||||
_, secondReplayErr := replayGuard.Admit(context.Background(),replayRequest)
|
||||
replayRejected := firstReplayErr == nil && errors.Is(secondReplayErr,gateway.ErrAdmissionRejected)
|
||||
admissionService := &admission{authority:authority,work:work,public:admissionKey.Public().(ed25519.PublicKey),reusable:mode == "reusable",failWorkOnce:mode == "post-retryable"}
|
||||
admissionService := &admission{authority:authority,work:work,public:admissionKey.Public().(ed25519.PublicKey),reusable:mode == "reusable",failWorkOnce:mode == "post-retryable",releaseFails:mode == "cleanup-release-failure"}
|
||||
newService := func() *gateway.Server {
|
||||
service, err := gateway.NewServer(gateway.ServerConfig{ListenAddress:"127.0.0.1:0",TLSConfig:server,GatewayID:authority.GatewayID,Admission:admissionService,Provider:provider{}}); if err != nil { panic(err) }
|
||||
service, err := gateway.NewServer(gateway.ServerConfig{ListenAddress:"127.0.0.1:0",TLSConfig:server,GatewayID:authority.GatewayID,Admission:admissionService,Provider:provider{failStart:mode == "provider-start-lost-response" || mode == "cleanup-release-failure"}}); if err != nil { panic(err) }
|
||||
return service
|
||||
}
|
||||
service := newService()
|
||||
@@ -139,7 +144,7 @@ func main() {
|
||||
services = []*gateway.Server{draining, service}
|
||||
addresses = []string{draining.Addr().String(), service.Addr().String()}
|
||||
}
|
||||
if mode == "post-retryable" {
|
||||
if mode == "post-retryable" || mode == "provider-start-lost-response" || mode == "cleanup-release-failure" {
|
||||
second := newService()
|
||||
services = []*gateway.Server{service, second}
|
||||
addresses = []string{service.Addr().String(), second.Addr().String()}
|
||||
@@ -303,7 +308,7 @@ fn callback_ed25519_signer_completes_tls13_quic_admission_without_private_key_in
|
||||
|
||||
assert_eq!(session.authority().session_id(), "session-1");
|
||||
assert_eq!(admission_calls.load(Ordering::SeqCst), 1);
|
||||
assert!(tls_calls.load(Ordering::SeqCst) >= 1);
|
||||
assert_eq!(tls_calls.load(Ordering::SeqCst), 1);
|
||||
assert!(admission_inputs
|
||||
.lock()
|
||||
.expect("admission transcript")
|
||||
@@ -598,6 +603,41 @@ fn overall_deadline_includes_synchronous_admission_signing() {
|
||||
assert!(started.elapsed() < Duration::from_secs(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn tls_signer_finite_deadline_overrun_is_cancelled_not_tls() {
|
||||
let oracle = Oracle::start("");
|
||||
let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest");
|
||||
let credential =
|
||||
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
|
||||
let admission_key = test_key(&oracle.ready.admission_key);
|
||||
let tls_key = test_key(&oracle.ready.client_key);
|
||||
let cancellation = Cancellation::new();
|
||||
let trigger = cancellation.clone();
|
||||
let canceller = thread::spawn(move || {
|
||||
thread::sleep(Duration::from_millis(25));
|
||||
trigger.cancel();
|
||||
});
|
||||
let result =
|
||||
tokio::runtime::Runtime::new()
|
||||
.expect("runtime")
|
||||
.block_on(connect_with_cancellation(
|
||||
&manifest,
|
||||
&credential,
|
||||
Signers::new(
|
||||
AdmissionSigner::new(signer(admission_key)),
|
||||
TlsEd25519Signer::new(move |message| {
|
||||
thread::sleep(Duration::from_millis(100));
|
||||
signer(Arc::clone(&tls_key))(message)
|
||||
}),
|
||||
),
|
||||
"2026-08-12T00:00:00Z",
|
||||
Duration::from_secs(5),
|
||||
&cancellation,
|
||||
));
|
||||
canceller.join().expect("canceller");
|
||||
assert_eq!(result.err(), Some(CoreError::Cancelled));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_cancellation_interrupts_network_wait() {
|
||||
let oracle = Oracle::start("");
|
||||
@@ -733,29 +773,6 @@ fn invalid_first_manifest_address_does_not_block_live_second_address() {
|
||||
fn retryable_stable_error_advances_to_live_manifest_address() {
|
||||
let oracle = Oracle::start("retryable");
|
||||
let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest");
|
||||
let credential =
|
||||
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
|
||||
let runtime = tokio::runtime::Runtime::new().expect("runtime");
|
||||
let session = runtime
|
||||
.block_on(connect(
|
||||
&manifest,
|
||||
&credential,
|
||||
recording_signers(
|
||||
test_key(&oracle.ready.admission_key),
|
||||
test_key(&oracle.ready.client_key),
|
||||
Arc::new(Mutex::new(Vec::new())),
|
||||
),
|
||||
"2026-08-12T00:00:00Z",
|
||||
Duration::from_secs(5),
|
||||
))
|
||||
.expect("retryable draining response advanced to live address");
|
||||
runtime.block_on(session.close());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn post_admission_retry_uses_a_fresh_signed_nonce() {
|
||||
let oracle = Oracle::start("post-retryable");
|
||||
let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest");
|
||||
let credential =
|
||||
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
|
||||
let transcripts = Arc::new(Mutex::new(Vec::new()));
|
||||
@@ -772,13 +789,48 @@ fn post_admission_retry_uses_a_fresh_signed_nonce() {
|
||||
"2026-08-12T00:00:00Z",
|
||||
Duration::from_secs(5),
|
||||
))
|
||||
.expect("post-admission retry used a fresh request");
|
||||
.expect("retryable draining response advanced to live address");
|
||||
runtime.block_on(session.close());
|
||||
let transcripts = transcripts.lock().expect("transcripts");
|
||||
assert_eq!(transcripts.len(), 2);
|
||||
assert_ne!(transcripts[0], transcripts[1]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn post_admission_failure_is_terminal_and_does_not_mutate_twice() {
|
||||
for mode in [
|
||||
"post-retryable",
|
||||
"provider-start-lost-response",
|
||||
"cleanup-release-failure",
|
||||
] {
|
||||
let oracle = Oracle::start(mode);
|
||||
let manifest =
|
||||
ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest");
|
||||
let credential =
|
||||
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
|
||||
let transcripts = Arc::new(Mutex::new(Vec::new()));
|
||||
let runtime = tokio::runtime::Runtime::new().expect("runtime");
|
||||
let result = runtime.block_on(connect(
|
||||
&manifest,
|
||||
&credential,
|
||||
recording_signers(
|
||||
test_key(&oracle.ready.admission_key),
|
||||
test_key(&oracle.ready.client_key),
|
||||
Arc::clone(&transcripts),
|
||||
),
|
||||
"2026-08-12T00:00:00Z",
|
||||
Duration::from_secs(5),
|
||||
));
|
||||
assert!(result.is_err(), "{mode} was retried to success");
|
||||
let transcripts = transcripts.lock().expect("transcripts");
|
||||
assert_eq!(
|
||||
transcripts.len(),
|
||||
1,
|
||||
"{mode} attempted admission mutation twice"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repeated_production_gateway_admissions_remain_bounded() {
|
||||
let runtime = tokio::runtime::Runtime::new().expect("runtime");
|
||||
|
||||
@@ -145,6 +145,14 @@ fn manifest_public_identity_requires_dns_sni_not_ip_or_uuid() {
|
||||
"non-DNS SNI accepted: {invalid_identity}"
|
||||
);
|
||||
}
|
||||
|
||||
let identity_equals_dns_shaped_gateway_id = String::from_utf8(valid_manifest().to_vec())
|
||||
.expect("fixture is UTF-8")
|
||||
.replace("\"id\":\"gateway\"", "\"id\":\"gateway.test\"");
|
||||
assert!(
|
||||
ConnectionManifest::decode(identity_equals_dns_shaped_gateway_id.as_bytes()).is_err(),
|
||||
"public SNI identity matched the logical gateway id"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
+212
-1
@@ -536,6 +536,164 @@ func TestStableErrorFramePrecedesConnectionTeardown(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestStableErrorDrainOutlivesExpiredHandlerContext(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()}
|
||||
admission := AdmissionFunc(func(ctx context.Context, _ protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error) {
|
||||
<-ctx.Done()
|
||||
return protocol.SessionAuthority{}, ctx.Err()
|
||||
})
|
||||
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: fake})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server.helloTimeout = 20 * time.Millisecond
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
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)
|
||||
}
|
||||
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.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 with expired handler context: %v", err)
|
||||
}
|
||||
if stable, err := protocol.DecodeStableError(response); err != nil || stable.Code != "admission_rejected" {
|
||||
t.Fatalf("stable response = %#v, %v", stable, err)
|
||||
}
|
||||
_ = connection.CloseWithError(applicationError, "done")
|
||||
cancel()
|
||||
_ = server.Close()
|
||||
}
|
||||
|
||||
func TestServerCloseInterruptsPartialHelloClient(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)
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
serveDone := make(chan error, 1)
|
||||
go func() { serveDone <- 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)
|
||||
}
|
||||
if _, err := stream.Write([]byte{0, 0}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
closed := make(chan error, 1)
|
||||
go func() { closed <- server.Close() }()
|
||||
select {
|
||||
case err := <-closed:
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("Server.Close blocked on partial hello")
|
||||
}
|
||||
select {
|
||||
case <-connection.Context().Done():
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("partial hello connection remained open")
|
||||
}
|
||||
cancel()
|
||||
if err := <-serveDone; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHelloUsesOneAbsoluteDeadlineAgainstTrickle(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)
|
||||
}
|
||||
server.helloTimeout = 40 * time.Millisecond
|
||||
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)
|
||||
}
|
||||
if _, err := stream.Write([]byte{0}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
started := time.Now()
|
||||
time.Sleep(25 * time.Millisecond)
|
||||
_, _ = stream.Write([]byte{0})
|
||||
response, err := readWire(stream, defaultHelloLimit)
|
||||
if err != nil {
|
||||
t.Fatalf("absolute hello deadline did not produce a stable error: %v", err)
|
||||
}
|
||||
stable, err := protocol.DecodeStableError(response)
|
||||
if err != nil || stable.Code != "invalid_hello" {
|
||||
t.Fatalf("stable response = %#v, %v", stable, err)
|
||||
}
|
||||
if time.Since(started) > 250*time.Millisecond {
|
||||
t.Fatal("trickling hello reset or escaped its absolute deadline")
|
||||
}
|
||||
_ = connection.CloseWithError(applicationError, "done")
|
||||
_ = server.Close()
|
||||
}
|
||||
|
||||
func TestHelloDeadlineIsClearedAfterAdmission(t *testing.T) {
|
||||
serverTLS, clientTLS := testTLS(t)
|
||||
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()}
|
||||
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true}
|
||||
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: fake})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server.helloTimeout = 40 * time.Millisecond
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
go func() { _ = server.Serve(ctx) }()
|
||||
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
|
||||
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
time.Sleep(80 * time.Millisecond)
|
||||
select {
|
||||
case <-client.connection.Context().Done():
|
||||
t.Fatal("hello deadline leaked into the admitted session")
|
||||
default:
|
||||
}
|
||||
_ = client.Close()
|
||||
_ = 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 {
|
||||
@@ -554,6 +712,58 @@ func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostAdmissionProviderFailureIsTerminalEvenWhenReleaseFails(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()}
|
||||
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true, releaseErr: errors.New("release failed")}
|
||||
provider := providerStartFunc(func(context.Context, LaunchRequest) (ProviderSession, error) {
|
||||
return nil, context.DeadlineExceeded
|
||||
})
|
||||
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: provider})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
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)
|
||||
}
|
||||
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.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.Fatal(err)
|
||||
}
|
||||
stable, err := protocol.DecodeStableError(response)
|
||||
if err != nil || stable.Code != "provider_timeout" || stable.Retryable {
|
||||
t.Fatalf("post-admission stable response = %#v, %v", stable, err)
|
||||
}
|
||||
select {
|
||||
case <-admission.released:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("post-admission failure did not attempt release")
|
||||
}
|
||||
if got := admission.releases.Load(); got != 1 {
|
||||
t.Fatalf("release attempts = %d, want 1", got)
|
||||
}
|
||||
_ = connection.CloseWithError(applicationError, "done")
|
||||
_ = server.Close()
|
||||
}
|
||||
|
||||
func TestGatewayTelemetrySeparatesQueueProcessingAndPacing(t *testing.T) {
|
||||
serverTLS, clientTLS := testTLS(t)
|
||||
session := &fakeSession{
|
||||
@@ -1668,6 +1878,7 @@ type oneTimeAdmission struct {
|
||||
streamPolicy protocol.ProviderStreamPolicy
|
||||
providerWork *protocol.ProviderSessionWork
|
||||
disableClipboard bool
|
||||
releaseErr error
|
||||
}
|
||||
|
||||
type recordingProviderStateReporter struct {
|
||||
@@ -1742,7 +1953,7 @@ func (a *oneTimeAdmission) Release(_ context.Context, authority protocol.Session
|
||||
a.releaseAuthority = authority
|
||||
close(a.released)
|
||||
}
|
||||
return nil
|
||||
return a.releaseErr
|
||||
}
|
||||
|
||||
func mustRead(t *testing.T, path string) []byte {
|
||||
|
||||
+59
-24
@@ -21,6 +21,7 @@ import (
|
||||
|
||||
const (
|
||||
defaultHelloLimit = 16 * 1024
|
||||
defaultHelloTimeout = 10 * time.Second
|
||||
defaultControlLimit = 128 * 1024
|
||||
clientControlBacklog = 64
|
||||
terminalAckTimeout = 2 * time.Second
|
||||
@@ -88,16 +89,18 @@ type mediaTimingObservation struct {
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
listener *quic.Listener
|
||||
config ServerConfig
|
||||
metrics *Metrics
|
||||
pacer *fairPacer
|
||||
mu sync.Mutex
|
||||
sessions map[*gatewaySession]struct{}
|
||||
draining atomic.Bool
|
||||
closed atomic.Bool
|
||||
closeOnce sync.Once
|
||||
workers sync.WaitGroup
|
||||
listener *quic.Listener
|
||||
config ServerConfig
|
||||
metrics *Metrics
|
||||
pacer *fairPacer
|
||||
mu sync.Mutex
|
||||
sessions map[*gatewaySession]struct{}
|
||||
connections map[*quic.Conn]struct{}
|
||||
helloTimeout time.Duration
|
||||
draining atomic.Bool
|
||||
closed atomic.Bool
|
||||
closeOnce sync.Once
|
||||
workers sync.WaitGroup
|
||||
}
|
||||
|
||||
func NewServer(config ServerConfig) (*Server, error) {
|
||||
@@ -138,7 +141,7 @@ func NewServer(config ServerConfig) (*Server, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Server{listener: listener, config: config, metrics: &Metrics{}, pacer: newFairPacer(config.PacerKbps), sessions: make(map[*gatewaySession]struct{})}, nil
|
||||
return &Server{listener: listener, config: config, metrics: &Metrics{}, pacer: newFairPacer(config.PacerKbps), sessions: make(map[*gatewaySession]struct{}), connections: make(map[*quic.Conn]struct{}), helloTimeout: defaultHelloTimeout}, nil
|
||||
}
|
||||
|
||||
func validateServerTLS(config *tls.Config) error {
|
||||
@@ -174,9 +177,22 @@ func (s *Server) Serve(ctx context.Context) error {
|
||||
}
|
||||
return err
|
||||
}
|
||||
s.mu.Lock()
|
||||
if s.closed.Load() {
|
||||
s.mu.Unlock()
|
||||
_ = connection.CloseWithError(applicationError, "server closed")
|
||||
continue
|
||||
}
|
||||
s.connections[connection] = struct{}{}
|
||||
s.workers.Add(1)
|
||||
s.mu.Unlock()
|
||||
go func() {
|
||||
defer s.workers.Done()
|
||||
defer func() {
|
||||
s.mu.Lock()
|
||||
delete(s.connections, connection)
|
||||
s.mu.Unlock()
|
||||
}()
|
||||
s.handleConnection(ctx, connection)
|
||||
}()
|
||||
}
|
||||
@@ -186,13 +202,24 @@ func (s *Server) Close() error {
|
||||
var err error
|
||||
s.closeOnce.Do(func() {
|
||||
s.BeginDrain()
|
||||
s.closed.Store(true)
|
||||
err = s.listener.Close()
|
||||
s.mu.Lock()
|
||||
s.closed.Store(true)
|
||||
sessions := make([]*gatewaySession, 0, len(s.sessions))
|
||||
for session := range s.sessions {
|
||||
session.cancel()
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
connections := make([]*quic.Conn, 0, len(s.connections))
|
||||
for connection := range s.connections {
|
||||
connections = append(connections, connection)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
err = s.listener.Close()
|
||||
for _, session := range sessions {
|
||||
session.cancel()
|
||||
}
|
||||
for _, connection := range connections {
|
||||
_ = connection.CloseWithError(applicationError, "server closed")
|
||||
}
|
||||
})
|
||||
s.workers.Wait()
|
||||
return err
|
||||
@@ -200,15 +227,23 @@ func (s *Server) Close() error {
|
||||
|
||||
func (s *Server) handleConnection(parent context.Context, connection *quic.Conn) {
|
||||
defer connection.CloseWithError(applicationError, "connection closed")
|
||||
ctx, cancel := context.WithTimeout(parent, 10*time.Second)
|
||||
helloDeadline := time.Now().Add(s.helloTimeout)
|
||||
ctx, cancel := context.WithDeadline(parent, helloDeadline)
|
||||
defer cancel()
|
||||
stream, err := connection.AcceptStream(ctx)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if err := stream.SetDeadline(helloDeadline); err != nil {
|
||||
return
|
||||
}
|
||||
writeError := func(code string, err error, retryable bool) {
|
||||
responseDeadline := time.Now().Add(time.Second)
|
||||
if stream.SetDeadline(responseDeadline) != nil {
|
||||
return
|
||||
}
|
||||
if writeStableError(stream, code, err, retryable) == nil && stream.Close() == nil {
|
||||
responseCtx, responseCancel := context.WithTimeout(ctx, time.Second)
|
||||
responseCtx, responseCancel := context.WithDeadline(context.Background(), responseDeadline)
|
||||
defer responseCancel()
|
||||
select {
|
||||
case <-connection.Context().Done():
|
||||
@@ -237,12 +272,12 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
authority, err := s.config.Admission.Admit(ctx, request)
|
||||
if err != nil {
|
||||
s.metrics.AdmissionRejects.Add(1)
|
||||
writeError(stableAdmissionCode(err), err, errors.Is(err, context.DeadlineExceeded))
|
||||
writeError(stableAdmissionCode(err), err, false)
|
||||
return
|
||||
}
|
||||
if s.Draining() {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("gateway_draining", ErrGatewayDraining, true)
|
||||
writeError("gateway_draining", ErrGatewayDraining, false)
|
||||
return
|
||||
}
|
||||
if err := s.validateAuthority(authority, request); err != nil {
|
||||
@@ -253,7 +288,7 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
work, err := s.config.Admission.ProviderWork(ctx, authority)
|
||||
if err != nil || s.validateProviderWork(work, authority) != nil {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("provider_work_unavailable", ErrAdmissionRejected, err != nil)
|
||||
writeError("provider_work_unavailable", ErrAdmissionRejected, false)
|
||||
return
|
||||
}
|
||||
selected, err := IntersectCapabilities(s.config.Capabilities, s.config.ProviderCapabilities, request.Capabilities, authority.Capabilities)
|
||||
@@ -278,12 +313,12 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
}
|
||||
if (work.ClipboardPolicy.ClientToProviderEnabled || work.ClipboardPolicy.ProviderToClientEnabled) && s.config.ClipboardAuditReporter == nil {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, true)
|
||||
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, false)
|
||||
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)
|
||||
writeError("provider_state_unavailable", err, true)
|
||||
writeError("provider_state_unavailable", err, false)
|
||||
return
|
||||
}
|
||||
providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: work.ProviderIdentity, ProviderWork: work})
|
||||
@@ -291,14 +326,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)
|
||||
writeError(stableProviderCode(err), err, errors.Is(err, context.DeadlineExceeded))
|
||||
writeError(stableProviderCode(err), err, false)
|
||||
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)
|
||||
writeError("provider_state_unavailable", err, true)
|
||||
writeError("provider_state_unavailable", err, false)
|
||||
return
|
||||
}
|
||||
clientAuthority := protocol.ClientSessionAuthority{
|
||||
@@ -306,7 +341,7 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
ReconnectSequence: authority.ReconnectSequence, ExpiresAt: authority.ExpiresAt, Capabilities: selected,
|
||||
}
|
||||
authorityBytes, err := protocol.EncodeClientSessionAuthority(clientAuthority)
|
||||
if err != nil || writeWire(stream, authorityBytes, defaultHelloLimit) != nil {
|
||||
if err != nil || writeWire(stream, authorityBytes, defaultHelloLimit) != nil || stream.SetDeadline(time.Time{}) != nil {
|
||||
_ = providerSession.ReleaseAll(context.Background())
|
||||
_ = providerSession.Terminate(context.Background())
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
|
||||
Reference in New Issue
Block a user