fix(transport): harden QUIC admission boundaries
This commit is contained in:
+56
-5
@@ -1,6 +1,6 @@
|
|||||||
use std::fmt;
|
use std::fmt;
|
||||||
use std::io::Cursor;
|
use std::io::Cursor;
|
||||||
use std::sync::Arc;
|
use std::sync::{Arc, Mutex};
|
||||||
|
|
||||||
use rustls::client::ResolvesClientCert;
|
use rustls::client::ResolvesClientCert;
|
||||||
use rustls::pki_types::CertificateDer;
|
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) type SignCallback = dyn Fn(&[u8]) -> Result<[u8; 64]> + Send + Sync;
|
||||||
|
|
||||||
pub(crate) struct CallbackSigningKey {
|
pub(crate) struct CallbackSigningKey {
|
||||||
callback: Arc<SignCallback>,
|
callback: Arc<Mutex<Option<Arc<SignCallback>>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl fmt::Debug for CallbackSigningKey {
|
impl fmt::Debug for CallbackSigningKey {
|
||||||
@@ -24,7 +24,9 @@ impl fmt::Debug for CallbackSigningKey {
|
|||||||
|
|
||||||
impl CallbackSigningKey {
|
impl CallbackSigningKey {
|
||||||
pub(crate) fn new(callback: Arc<SignCallback>) -> Self {
|
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 {
|
impl fmt::Debug for CallbackSigner {
|
||||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
@@ -50,7 +52,13 @@ impl fmt::Debug for CallbackSigner {
|
|||||||
|
|
||||||
impl Signer for CallbackSigner {
|
impl Signer for CallbackSigner {
|
||||||
fn sign(&self, message: &[u8]) -> std::result::Result<Vec<u8>, rustls::Error> {
|
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(|signature| signature.to_vec())
|
||||||
.map_err(|_| rustls::Error::General("client signing failed".to_owned()))
|
.map_err(|_| rustls::Error::General("client signing failed".to_owned()))
|
||||||
}
|
}
|
||||||
@@ -112,3 +120,46 @@ pub(crate) fn client_config(
|
|||||||
config.enable_early_data = false;
|
config.enable_early_data = false;
|
||||||
Ok(config)
|
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::fmt;
|
||||||
|
use std::io;
|
||||||
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
|
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
|
||||||
use std::sync::atomic::{AtomicBool, Ordering};
|
use std::sync::atomic::{AtomicBool, Ordering};
|
||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
@@ -157,15 +158,6 @@ pub async fn connect_with_cancellation(
|
|||||||
}
|
}
|
||||||
Ok(signature)
|
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
|
let remaining = deadline
|
||||||
.checked_sub(started.elapsed())
|
.checked_sub(started.elapsed())
|
||||||
.ok_or(CoreError::Cancelled)?;
|
.ok_or(CoreError::Cancelled)?;
|
||||||
@@ -173,16 +165,25 @@ pub async fn connect_with_cancellation(
|
|||||||
let address_count = manifest.addresses().len();
|
let address_count = manifest.addresses().len();
|
||||||
for (address_index, address) in manifest.addresses().iter().enumerate() {
|
for (address_index, address) in manifest.addresses().iter().enumerate() {
|
||||||
cancellation.check()?;
|
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())
|
let address_budget = deadline.saturating_sub(started.elapsed())
|
||||||
/ u32::try_from(address_count.saturating_sub(address_index).max(1)).unwrap_or(1);
|
/ u32::try_from(address_count.saturating_sub(address_index).max(1)).unwrap_or(1);
|
||||||
let address_deadline = Instant::now() + address_budget;
|
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();
|
let remote_count = resolved.len();
|
||||||
for (remote_index, remote) in resolved.into_iter().enumerate() {
|
for (remote_index, remote) in resolved.into_iter().enumerate() {
|
||||||
cancellation.check()?;
|
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 remaining_remotes = remote_count.saturating_sub(remote_index).max(1);
|
||||||
let attempt_budget = address_deadline.saturating_duration_since(Instant::now())
|
let attempt_budget = address_deadline.saturating_duration_since(Instant::now())
|
||||||
/ u32::try_from(remaining_remotes).unwrap_or(1);
|
/ 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> {
|
struct AttemptContext<'a> {
|
||||||
manifest: &'a ConnectionManifest,
|
manifest: &'a ConnectionManifest,
|
||||||
offered: &'a CapabilityProfile,
|
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(
|
async fn dial(
|
||||||
remote: SocketAddr,
|
remote: SocketAddr,
|
||||||
config: quinn::ClientConfig,
|
config: quinn::ClientConfig,
|
||||||
@@ -270,7 +300,7 @@ async fn dial(
|
|||||||
let negotiated = connecting
|
let negotiated = connecting
|
||||||
.handshake_data()
|
.handshake_data()
|
||||||
.await
|
.await
|
||||||
.map_err(|_| AttemptError::Terminal(CoreError::Tls))?
|
.map_err(|_| AttemptError::Terminal(tls_error(context)))?
|
||||||
.downcast::<quinn::crypto::rustls::HandshakeData>()
|
.downcast::<quinn::crypto::rustls::HandshakeData>()
|
||||||
.ok()
|
.ok()
|
||||||
.and_then(|data| data.protocol.clone());
|
.and_then(|data| data.protocol.clone());
|
||||||
@@ -279,7 +309,7 @@ async fn dial(
|
|||||||
}
|
}
|
||||||
let connection = connecting
|
let connection = connecting
|
||||||
.await
|
.await
|
||||||
.map_err(|_| AttemptError::Terminal(CoreError::Tls))?;
|
.map_err(|_| AttemptError::Terminal(tls_error(context)))?;
|
||||||
let payload = admission_payload(
|
let payload = admission_payload(
|
||||||
context.manifest,
|
context.manifest,
|
||||||
context.offered,
|
context.offered,
|
||||||
@@ -302,7 +332,7 @@ async fn dial(
|
|||||||
.map_err(|_| AttemptError::Terminal(connection_error(&connection)))?;
|
.map_err(|_| AttemptError::Terminal(connection_error(&connection)))?;
|
||||||
let Ok(authority) = ClientSessionAuthority::decode(&response) else {
|
let Ok(authority) = ClientSessionAuthority::decode(&response) else {
|
||||||
let stable = decode_stable_error(&response)?;
|
let stable = decode_stable_error(&response)?;
|
||||||
return if stable.retryable {
|
return if stable.retryable && stable.code == "gateway_draining" {
|
||||||
Err(AttemptError::Retry)
|
Err(AttemptError::Retry)
|
||||||
} else {
|
} else {
|
||||||
Err(AttemptError::Terminal(stable.error))
|
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(
|
fn admission_payload(
|
||||||
manifest: &ConnectionManifest,
|
manifest: &ConnectionManifest,
|
||||||
offered: &CapabilityProfile,
|
offered: &CapabilityProfile,
|
||||||
@@ -410,7 +448,11 @@ fn hello_length(length: usize) -> Result<u32> {
|
|||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
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;
|
use crate::error::CoreError;
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
@@ -419,4 +461,32 @@ mod tests {
|
|||||||
assert_eq!(hello_length(0), Err(CoreError::Protocol));
|
assert_eq!(hello_length(0), Err(CoreError::Protocol));
|
||||||
assert_eq!(hello_length(HELLO_LIMIT + 1), 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()
|
.iter()
|
||||||
.any(|address| !bounded(address, 1, 256))
|
.any(|address| !bounded(address, 1, 256))
|
||||||
|| !valid_dns_name(&self.gateway.public_identity)
|
|| !valid_dns_name(&self.gateway.public_identity)
|
||||||
|
|| self.gateway.public_identity == self.gateway.id
|
||||||
|| !(1..=4).contains(&self.tunnel.versions.len())
|
|| !(1..=4).contains(&self.tunnel.versions.len())
|
||||||
|| self
|
|| self
|
||||||
.tunnel
|
.tunnel
|
||||||
@@ -621,6 +622,7 @@ struct StableError {
|
|||||||
|
|
||||||
pub(crate) struct DecodedStableError {
|
pub(crate) struct DecodedStableError {
|
||||||
pub(crate) error: CoreError,
|
pub(crate) error: CoreError,
|
||||||
|
pub(crate) code: String,
|
||||||
pub(crate) retryable: bool,
|
pub(crate) retryable: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -703,6 +705,7 @@ pub(crate) fn decode_stable_error(bytes: &[u8]) -> Result<DecodedStableError> {
|
|||||||
};
|
};
|
||||||
Ok(DecodedStableError {
|
Ok(DecodedStableError {
|
||||||
error,
|
error,
|
||||||
|
code: stable.code,
|
||||||
retryable: stable.retryable,
|
retryable: stable.retryable,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ type admission struct {
|
|||||||
lastTranscript []byte
|
lastTranscript []byte
|
||||||
reusable bool
|
reusable bool
|
||||||
failWorkOnce bool
|
failWorkOnce bool
|
||||||
|
releaseFails bool
|
||||||
authority protocol.SessionAuthority
|
authority protocol.SessionAuthority
|
||||||
work protocol.ProviderSessionWork
|
work protocol.ProviderSessionWork
|
||||||
public ed25519.PublicKey
|
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 }
|
if a.failWorkOnce { a.failWorkOnce = false; return protocol.ProviderSessionWork{}, context.DeadlineExceeded }
|
||||||
return a.work, nil
|
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 }
|
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
|
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 }
|
func (s *session) Ready(context.Context) error { return nil }
|
||||||
@@ -125,9 +130,9 @@ func main() {
|
|||||||
_, firstReplayErr := replayGuard.Admit(context.Background(),replayRequest)
|
_, firstReplayErr := replayGuard.Admit(context.Background(),replayRequest)
|
||||||
_, secondReplayErr := replayGuard.Admit(context.Background(),replayRequest)
|
_, secondReplayErr := replayGuard.Admit(context.Background(),replayRequest)
|
||||||
replayRejected := firstReplayErr == nil && errors.Is(secondReplayErr,gateway.ErrAdmissionRejected)
|
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 {
|
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
|
return service
|
||||||
}
|
}
|
||||||
service := newService()
|
service := newService()
|
||||||
@@ -139,7 +144,7 @@ func main() {
|
|||||||
services = []*gateway.Server{draining, service}
|
services = []*gateway.Server{draining, service}
|
||||||
addresses = []string{draining.Addr().String(), service.Addr().String()}
|
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()
|
second := newService()
|
||||||
services = []*gateway.Server{service, second}
|
services = []*gateway.Server{service, second}
|
||||||
addresses = []string{service.Addr().String(), second.Addr().String()}
|
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!(session.authority().session_id(), "session-1");
|
||||||
assert_eq!(admission_calls.load(Ordering::SeqCst), 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
|
assert!(admission_inputs
|
||||||
.lock()
|
.lock()
|
||||||
.expect("admission transcript")
|
.expect("admission transcript")
|
||||||
@@ -598,6 +603,41 @@ fn overall_deadline_includes_synchronous_admission_signing() {
|
|||||||
assert!(started.elapsed() < Duration::from_secs(2));
|
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]
|
#[test]
|
||||||
fn explicit_cancellation_interrupts_network_wait() {
|
fn explicit_cancellation_interrupts_network_wait() {
|
||||||
let oracle = Oracle::start("");
|
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() {
|
fn retryable_stable_error_advances_to_live_manifest_address() {
|
||||||
let oracle = Oracle::start("retryable");
|
let oracle = Oracle::start("retryable");
|
||||||
let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest");
|
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 =
|
let credential =
|
||||||
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
|
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
|
||||||
let transcripts = Arc::new(Mutex::new(Vec::new()));
|
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",
|
"2026-08-12T00:00:00Z",
|
||||||
Duration::from_secs(5),
|
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());
|
runtime.block_on(session.close());
|
||||||
let transcripts = transcripts.lock().expect("transcripts");
|
let transcripts = transcripts.lock().expect("transcripts");
|
||||||
assert_eq!(transcripts.len(), 2);
|
assert_eq!(transcripts.len(), 2);
|
||||||
assert_ne!(transcripts[0], transcripts[1]);
|
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]
|
#[test]
|
||||||
fn repeated_production_gateway_admissions_remain_bounded() {
|
fn repeated_production_gateway_admissions_remain_bounded() {
|
||||||
let runtime = tokio::runtime::Runtime::new().expect("runtime");
|
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}"
|
"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]
|
#[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) {
|
func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
|
||||||
var wire bytes.Buffer
|
var wire bytes.Buffer
|
||||||
if err := writeStableError(&wire, "provider_unavailable", errors.New("https://provider.invalid/launch?rikey=secret-sentinel"), true); err != nil {
|
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) {
|
func TestGatewayTelemetrySeparatesQueueProcessingAndPacing(t *testing.T) {
|
||||||
serverTLS, clientTLS := testTLS(t)
|
serverTLS, clientTLS := testTLS(t)
|
||||||
session := &fakeSession{
|
session := &fakeSession{
|
||||||
@@ -1668,6 +1878,7 @@ type oneTimeAdmission struct {
|
|||||||
streamPolicy protocol.ProviderStreamPolicy
|
streamPolicy protocol.ProviderStreamPolicy
|
||||||
providerWork *protocol.ProviderSessionWork
|
providerWork *protocol.ProviderSessionWork
|
||||||
disableClipboard bool
|
disableClipboard bool
|
||||||
|
releaseErr error
|
||||||
}
|
}
|
||||||
|
|
||||||
type recordingProviderStateReporter struct {
|
type recordingProviderStateReporter struct {
|
||||||
@@ -1742,7 +1953,7 @@ func (a *oneTimeAdmission) Release(_ context.Context, authority protocol.Session
|
|||||||
a.releaseAuthority = authority
|
a.releaseAuthority = authority
|
||||||
close(a.released)
|
close(a.released)
|
||||||
}
|
}
|
||||||
return nil
|
return a.releaseErr
|
||||||
}
|
}
|
||||||
|
|
||||||
func mustRead(t *testing.T, path string) []byte {
|
func mustRead(t *testing.T, path string) []byte {
|
||||||
|
|||||||
+59
-24
@@ -21,6 +21,7 @@ import (
|
|||||||
|
|
||||||
const (
|
const (
|
||||||
defaultHelloLimit = 16 * 1024
|
defaultHelloLimit = 16 * 1024
|
||||||
|
defaultHelloTimeout = 10 * time.Second
|
||||||
defaultControlLimit = 128 * 1024
|
defaultControlLimit = 128 * 1024
|
||||||
clientControlBacklog = 64
|
clientControlBacklog = 64
|
||||||
terminalAckTimeout = 2 * time.Second
|
terminalAckTimeout = 2 * time.Second
|
||||||
@@ -88,16 +89,18 @@ type mediaTimingObservation struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Server struct {
|
type Server struct {
|
||||||
listener *quic.Listener
|
listener *quic.Listener
|
||||||
config ServerConfig
|
config ServerConfig
|
||||||
metrics *Metrics
|
metrics *Metrics
|
||||||
pacer *fairPacer
|
pacer *fairPacer
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
sessions map[*gatewaySession]struct{}
|
sessions map[*gatewaySession]struct{}
|
||||||
draining atomic.Bool
|
connections map[*quic.Conn]struct{}
|
||||||
closed atomic.Bool
|
helloTimeout time.Duration
|
||||||
closeOnce sync.Once
|
draining atomic.Bool
|
||||||
workers sync.WaitGroup
|
closed atomic.Bool
|
||||||
|
closeOnce sync.Once
|
||||||
|
workers sync.WaitGroup
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewServer(config ServerConfig) (*Server, error) {
|
func NewServer(config ServerConfig) (*Server, error) {
|
||||||
@@ -138,7 +141,7 @@ func NewServer(config ServerConfig) (*Server, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
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 {
|
func validateServerTLS(config *tls.Config) error {
|
||||||
@@ -174,9 +177,22 @@ func (s *Server) Serve(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
return err
|
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.workers.Add(1)
|
||||||
|
s.mu.Unlock()
|
||||||
go func() {
|
go func() {
|
||||||
defer s.workers.Done()
|
defer s.workers.Done()
|
||||||
|
defer func() {
|
||||||
|
s.mu.Lock()
|
||||||
|
delete(s.connections, connection)
|
||||||
|
s.mu.Unlock()
|
||||||
|
}()
|
||||||
s.handleConnection(ctx, connection)
|
s.handleConnection(ctx, connection)
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
@@ -186,13 +202,24 @@ func (s *Server) Close() error {
|
|||||||
var err error
|
var err error
|
||||||
s.closeOnce.Do(func() {
|
s.closeOnce.Do(func() {
|
||||||
s.BeginDrain()
|
s.BeginDrain()
|
||||||
s.closed.Store(true)
|
|
||||||
err = s.listener.Close()
|
|
||||||
s.mu.Lock()
|
s.mu.Lock()
|
||||||
|
s.closed.Store(true)
|
||||||
|
sessions := make([]*gatewaySession, 0, len(s.sessions))
|
||||||
for session := range 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()
|
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()
|
s.workers.Wait()
|
||||||
return err
|
return err
|
||||||
@@ -200,15 +227,23 @@ func (s *Server) Close() error {
|
|||||||
|
|
||||||
func (s *Server) handleConnection(parent context.Context, connection *quic.Conn) {
|
func (s *Server) handleConnection(parent context.Context, connection *quic.Conn) {
|
||||||
defer connection.CloseWithError(applicationError, "connection closed")
|
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()
|
defer cancel()
|
||||||
stream, err := connection.AcceptStream(ctx)
|
stream, err := connection.AcceptStream(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
if err := stream.SetDeadline(helloDeadline); err != nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
writeError := func(code string, err error, retryable bool) {
|
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 {
|
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()
|
defer responseCancel()
|
||||||
select {
|
select {
|
||||||
case <-connection.Context().Done():
|
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)
|
authority, err := s.config.Admission.Admit(ctx, request)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
s.metrics.AdmissionRejects.Add(1)
|
s.metrics.AdmissionRejects.Add(1)
|
||||||
writeError(stableAdmissionCode(err), err, errors.Is(err, context.DeadlineExceeded))
|
writeError(stableAdmissionCode(err), err, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if s.Draining() {
|
if s.Draining() {
|
||||||
_ = s.config.Admission.Release(context.Background(), authority)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
writeError("gateway_draining", ErrGatewayDraining, true)
|
writeError("gateway_draining", ErrGatewayDraining, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := s.validateAuthority(authority, request); err != nil {
|
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)
|
work, err := s.config.Admission.ProviderWork(ctx, authority)
|
||||||
if err != nil || s.validateProviderWork(work, authority) != nil {
|
if err != nil || s.validateProviderWork(work, authority) != nil {
|
||||||
_ = s.config.Admission.Release(context.Background(), authority)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
writeError("provider_work_unavailable", ErrAdmissionRejected, err != nil)
|
writeError("provider_work_unavailable", ErrAdmissionRejected, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
selected, err := IntersectCapabilities(s.config.Capabilities, s.config.ProviderCapabilities, request.Capabilities, authority.Capabilities)
|
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 {
|
if (work.ClipboardPolicy.ClientToProviderEnabled || work.ClipboardPolicy.ProviderToClientEnabled) && s.config.ClipboardAuditReporter == nil {
|
||||||
_ = s.config.Admission.Release(context.Background(), authority)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, true)
|
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, false)
|
||||||
return
|
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 {
|
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)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
writeError("provider_state_unavailable", err, true)
|
writeError("provider_state_unavailable", err, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: work.ProviderIdentity, ProviderWork: work})
|
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.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.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)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
writeError(stableProviderCode(err), err, errors.Is(err, context.DeadlineExceeded))
|
writeError(stableProviderCode(err), err, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := s.reportProviderState(ctx, providerSession.State()); err != nil {
|
if err := s.reportProviderState(ctx, providerSession.State()); err != nil {
|
||||||
_ = providerSession.ReleaseAll(context.Background())
|
_ = providerSession.ReleaseAll(context.Background())
|
||||||
_ = providerSession.Terminate(context.Background())
|
_ = providerSession.Terminate(context.Background())
|
||||||
_ = s.config.Admission.Release(context.Background(), authority)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
writeError("provider_state_unavailable", err, true)
|
writeError("provider_state_unavailable", err, false)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
clientAuthority := protocol.ClientSessionAuthority{
|
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,
|
ReconnectSequence: authority.ReconnectSequence, ExpiresAt: authority.ExpiresAt, Capabilities: selected,
|
||||||
}
|
}
|
||||||
authorityBytes, err := protocol.EncodeClientSessionAuthority(clientAuthority)
|
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.ReleaseAll(context.Background())
|
||||||
_ = providerSession.Terminate(context.Background())
|
_ = providerSession.Terminate(context.Background())
|
||||||
_ = s.config.Admission.Release(context.Background(), authority)
|
_ = s.config.Admission.Release(context.Background(), authority)
|
||||||
|
|||||||
Reference in New Issue
Block a user