use anyhow::{anyhow, bail, Context, Result};
use async_trait::async_trait;
use quinn::{Endpoint, ReadError, RecvStream, SendStream};
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
use rustls::pki_types::{CertificateDer, PrivatePkcs8KeyDer, ServerName, UnixTime};
use rustls::SignatureScheme;
use serde_json::Value;
use sha2::{Digest, Sha256};
use std::net::{IpAddr, SocketAddr, UdpSocket};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
use crate::net::Transport;
pub fn direct_enabled() -> bool {
std::env::var("FILAMENT_DIRECT").map(|v| v != "0").unwrap_or(true)
|| std::env::var("FILAMENT_L2").map(|v| v == "1").unwrap_or(false)
}
fn test_block() -> bool {
std::env::var("FILAMENT_DIRECT_TEST_BLOCK").map(|v| v == "1").unwrap_or(false)
}
#[cfg(feature = "test-hooks")]
fn freeze_after_bytes() -> Option<u64> {
std::env::var("FILAMENT_TEST_FREEZE_AFTER_BYTES")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.filter(|n| *n > 0)
}
#[cfg(feature = "test-hooks")]
fn freeze_persist() -> bool {
std::env::var("FILAMENT_TEST_FREEZE_PERSIST").map(|v| v == "1").unwrap_or(false)
}
#[cfg(feature = "test-hooks")]
fn direct_unblock_after_ms() -> Option<u64> {
std::env::var("FILAMENT_TEST_DIRECT_UNBLOCK_MS")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.filter(|n| *n > 0)
}
#[cfg(feature = "test-hooks")]
fn direct_flaky_upgrade() -> bool {
std::env::var("FILAMENT_TEST_DIRECT_FLAKY").map(|v| v == "1").unwrap_or(false)
}
#[cfg(feature = "test-hooks")]
const FLAKY_REFREEZE_BYTES: u64 = 4_096;
#[cfg(feature = "test-hooks")]
static FROZE_ONCE: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub const DIRECT_BUDGET: std::time::Duration = std::time::Duration::from_secs(5);
const ALPN: &[u8] = b"filament-direct/1";
pub const MAX_DIRECT_PAYLOAD: usize = 1024 * 1024;
fn hmac_sha256_raw(key: &[u8], msg: &[u8]) -> [u8; 32] {
let mut k = [0u8; 64];
if key.len() > 64 {
let mut h = Sha256::new();
h.update(key);
k[..32].copy_from_slice(&h.finalize());
} else {
k[..key.len()].copy_from_slice(key);
}
let ipad: Vec<u8> = k.iter().map(|b| b ^ 0x36).collect();
let opad: Vec<u8> = k.iter().map(|b| b ^ 0x5c).collect();
let mut inner = Sha256::new();
inner.update(&ipad);
inner.update(msg);
let inner = inner.finalize();
let mut outer = Sha256::new();
outer.update(&opad);
outer.update(inner);
let mut out = [0u8; 32];
out.copy_from_slice(&outer.finalize());
out
}
pub fn transport_key(secret: &str) -> [u8; 32] {
let prk = hmac_sha256_raw(&[0u8; 32], secret.as_bytes());
let mut info = b"filament-direct-transport-v1".to_vec();
info.push(0x01);
hmac_sha256_raw(&prk, &info)
}
fn ct_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut diff = 0u8;
for (x, y) in a.iter().zip(b.iter()) {
diff |= x ^ y;
}
diff == 0
}
struct CidrRule {
ip: IpAddr,
prefix_len: u8,
}
fn parse_cidr(s: &str) -> Option<CidrRule> {
let s = s.trim();
let (ip_str, len_str) = s.split_once('/')?;
let ip: IpAddr = ip_str.parse().ok()?;
let prefix_len: u8 = len_str.parse().ok()?;
let max = if ip.is_ipv4() { 32 } else { 128 };
if prefix_len > max {
return None;
}
Some(CidrRule { ip, prefix_len })
}
fn ip_matches_cidr(ip: &IpAddr, rule: &CidrRule) -> bool {
match (ip, &rule.ip) {
(IpAddr::V4(a), IpAddr::V4(net)) => {
let mask = u32::MAX.checked_shl(32 - rule.prefix_len as u32).unwrap_or(0);
u32::from(*a) & mask == u32::from(*net) & mask
}
(IpAddr::V6(a), IpAddr::V6(net)) => {
let mask = u128::MAX.checked_shl(128 - rule.prefix_len as u32).unwrap_or(0);
u128::from(*a) & mask == u128::from(*net) & mask
}
_ => false,
}
}
fn expand_group(item: &str) -> Vec<CidrRule> {
match item.to_ascii_lowercase().as_str() {
"tailscale" => parse_cidr("100.64.0.0/10").into_iter().collect(),
"docker" => parse_cidr("172.17.0.0/16").into_iter().collect(),
_ => vec![],
}
}
fn parse_rule_item(item: &str) -> Vec<CidrRule> {
let item = item.trim();
if item.is_empty() {
return vec![];
}
if let Some(r) = parse_cidr(item) {
return vec![r];
}
let groups = expand_group(item);
if !groups.is_empty() {
return groups;
}
vec![]
}
fn parse_rule_list(csv: &str) -> Vec<CidrRule> {
csv.split(',')
.flat_map(parse_rule_item)
.collect()
}
fn read_iface_setting(key: &str, peer: Option<&str>) -> String {
if key == "interfaces" {
return crate::settings::raw_membership(peer);
}
crate::settings::get_str(key, peer).unwrap_or_default()
}
fn filter_local_ips(ips: Vec<IpAddr>, peer: Option<&str>) -> Vec<IpAddr> {
let membership_raw = read_iface_setting("interfaces", peer);
let (membership_mode, membership_items) = parse_membership(&membership_raw);
let avoid_rules = if membership_mode == "avoid" {
parse_rule_list(&membership_items)
} else {
vec![]
};
let only_rules = if membership_mode == "only" {
parse_rule_list(&membership_items)
} else {
vec![]
};
let prefer = read_iface_setting("prefer", peer);
let prefer_rules = parse_rule_list(&prefer);
if membership_mode.is_empty() && prefer.is_empty() {
return ips;
}
let mut kept: Vec<(IpAddr, u8)> = ips
.into_iter()
.filter(|ip| {
if ip.is_loopback() {
return true; }
for r in &avoid_rules {
if ip_matches_cidr(ip, r) {
return false;
}
}
if !only_rules.is_empty() {
return only_rules.iter().any(|r| ip_matches_cidr(ip, r));
}
true
})
.map(|ip| {
let weight: u8 = if prefer_rules.iter().any(|r| ip_matches_cidr(&ip, r)) {
0
} else {
1
};
(ip, weight)
})
.collect();
if kept.is_empty() {
crate::ui::debug("membership filter removed ALL host candidates");
return vec![];
}
kept.sort_by_key(|(_, w)| *w);
kept.into_iter().map(|(ip, _)| ip).collect()
}
fn parse_membership(raw: &str) -> (String, String) {
if raw.is_empty() {
return (String::new(), String::new());
}
if let Ok(v) = serde_json::from_str::<serde_json::Value>(raw) {
let mode = v["m"].as_str().unwrap_or("").to_string();
let items = v["i"].as_str().unwrap_or("").to_string();
return (mode, items);
}
(String::new(), String::new())
}
fn local_ips() -> Vec<IpAddr> {
let mut out = Vec::new();
if let Ok(s) = UdpSocket::bind("0.0.0.0:0") {
if s.connect(("8.8.8.8", 9)).is_ok() {
if let Ok(la) = s.local_addr() {
let ip = la.ip();
if !ip.is_loopback() && !out.contains(&ip) {
out.push(ip);
}
}
}
}
if let Ok(s) = UdpSocket::bind("[::]:0") {
if s.connect(("2001:4860:4860::8888", 9)).is_ok() {
if let Ok(la) = s.local_addr() {
let ip = la.ip();
if !ip.is_loopback() && !out.contains(&ip) {
out.push(ip);
}
}
}
}
let lo: IpAddr = "127.0.0.1".parse().unwrap();
out.push(lo);
out
}
pub fn local_ip_snapshot() -> Vec<String> {
let mut v: Vec<String> = local_ips()
.into_iter()
.filter(|ip| !ip.is_loopback())
.map(|ip| ip.to_string())
.collect();
v.sort();
v.dedup();
v
}
const PUBIP_TTL: std::time::Duration = std::time::Duration::from_secs(300);
struct PubIpEntry {
server: String,
ip: IpAddr,
at: std::time::Instant,
}
static PUBIP_CACHE: std::sync::OnceLock<Mutex<Option<PubIpEntry>>> = std::sync::OnceLock::new();
fn pubip_fresh(cached_server: &str, age: std::time::Duration, want_server: &str, ttl: std::time::Duration) -> bool {
cached_server == want_server && age < ttl
}
async fn public_ip(server: &str) -> Option<IpAddr> {
if let Ok(v) = std::env::var("FILAMENT_PUBLIC_IP") {
if let Ok(ip) = v.trim().parse::<IpAddr>() {
return Some(ip);
}
}
let cache = PUBIP_CACHE.get_or_init(|| Mutex::new(None));
{
let guard = cache.lock().await;
if let Some(e) = guard.as_ref() {
if pubip_fresh(&e.server, e.at.elapsed(), server, PUBIP_TTL) {
return Some(e.ip);
}
}
}
let url = format!("{server}/api/whoami");
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(3))
.build()
.ok()?;
let resp = client.get(&url).send().await.ok()?;
if !resp.status().is_success() {
return None;
}
let v: Value = resp.json().await.ok()?;
let ip = v["ip"].as_str()?.trim().parse::<IpAddr>().ok()?;
let mut guard = cache.lock().await;
*guard = Some(PubIpEntry { server: server.to_string(), ip, at: std::time::Instant::now() });
Some(ip)
}
pub async fn warm_public_ip(server: &str) {
let _ = public_ip(server).await;
}
fn cand_str(ip: IpAddr, port: u16) -> String {
SocketAddr::new(ip, port).to_string()
}
fn suppress_public() -> bool {
std::env::var("FILAMENT_DIRECT_NO_PUBLIC").map(|v| v == "1").unwrap_or(false)
}
#[cfg(feature = "test-hooks")]
fn loopback_only() -> bool {
std::env::var("FILAMENT_DIRECT_LOOPBACK_ONLY").map(|v| v == "1").unwrap_or(false)
}
#[cfg(not(feature = "test-hooks"))]
#[inline]
fn loopback_only() -> bool {
false
}
pub async fn gather_candidates(server: &str, port: u16) -> Vec<String> {
let peer: Option<&str> = None; if loopback_only() {
let lo: IpAddr = "127.0.0.1".parse().unwrap();
return vec![cand_str(lo, port)];
}
let ips = filter_local_ips(local_ips(), peer);
let mut cands: Vec<String> = ips.into_iter().map(|ip| cand_str(ip, port)).collect();
if cands.is_empty() && !local_ips().is_empty() {
crate::ui::debug("iface filter removed ALL host candidates; connectivity may fail");
}
if !suppress_public() {
if let Some(pip) = public_ip(server).await {
let s = cand_str(pip, port);
if !cands.contains(&s) {
cands.push(s);
}
}
}
cands
}
#[derive(Debug)]
struct AcceptAnyCert(Arc<rustls::crypto::CryptoProvider>);
impl ServerCertVerifier for AcceptAnyCert {
fn verify_server_cert(
&self,
_end_entity: &CertificateDer<'_>,
_intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
_now: UnixTime,
) -> Result<ServerCertVerified, rustls::Error> {
Ok(ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&self.0.signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&self.0.signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
self.0.signature_verification_algorithms.supported_schemes()
}
}
fn provider() -> Arc<rustls::crypto::CryptoProvider> {
Arc::new(rustls::crypto::ring::default_provider())
}
fn direct_transport_config() -> Arc<quinn::TransportConfig> {
let mut tc = quinn::TransportConfig::default();
tc.keep_alive_interval(Some(std::time::Duration::from_secs(7)));
if let Ok(idle) = quinn::IdleTimeout::try_from(std::time::Duration::from_secs(21)) {
tc.max_idle_timeout(Some(idle));
}
tc.stream_receive_window(quinn::VarInt::from_u32(16_777_216));
tc.receive_window(quinn::VarInt::from_u32(33_554_432));
tc.send_window(16_777_216_u64);
tc.congestion_controller_factory(Arc::new(quinn_proto::congestion::BbrConfig::default()));
tc.datagram_receive_buffer_size(Some(1024 * 1024));
tc.datagram_send_buffer_size(1024 * 1024);
Arc::new(tc)
}
pub(crate) fn server_config() -> Result<quinn::ServerConfig> {
let ck = rcgen::generate_simple_self_signed(vec!["filament-direct".to_string()])
.context("self-signed cert")?;
let cert_der = CertificateDer::from(ck.cert.der().clone());
let key_der = PrivatePkcs8KeyDer::from(ck.key_pair.serialize_der());
let mut crypto = rustls::ServerConfig::builder_with_provider(provider())
.with_safe_default_protocol_versions()
.context("server tls versions")?
.with_no_client_auth()
.with_single_cert(vec![cert_der], key_der.into())
.context("server single cert")?;
crypto.alpn_protocols = vec![ALPN.to_vec()];
let qsc = quinn::crypto::rustls::QuicServerConfig::try_from(crypto)
.context("quic server config")?;
let mut sc = quinn::ServerConfig::with_crypto(Arc::new(qsc));
sc.transport_config(direct_transport_config());
Ok(sc)
}
pub(crate) fn client_config() -> Result<quinn::ClientConfig> {
let mut crypto = rustls::ClientConfig::builder_with_provider(provider())
.with_safe_default_protocol_versions()
.context("client tls versions")?
.dangerous()
.with_custom_certificate_verifier(Arc::new(AcceptAnyCert(provider())))
.with_no_client_auth();
crypto.alpn_protocols = vec![ALPN.to_vec()];
let qcc = quinn::crypto::rustls::QuicClientConfig::try_from(crypto)
.context("quic client config")?;
let mut cc = quinn::ClientConfig::new(Arc::new(qcc));
cc.transport_config(direct_transport_config());
Ok(cc)
}
pub fn bind_endpoint() -> Result<(Endpoint, u16)> {
use socket2::{Domain, Protocol, Socket, Type};
let sock = Socket::new(Domain::IPV4, Type::DGRAM, Some(Protocol::UDP))
.context("socket2 UDP socket")?;
let buf_size: usize = 8 * 1024 * 1024; sock.set_recv_buffer_size(buf_size)
.or_else(|_| sock.set_recv_buffer_size(4 * 1024 * 1024))
.or_else(|_| sock.set_recv_buffer_size(1024 * 1024))
.context("set SO_RCVBUF")?;
sock.set_send_buffer_size(buf_size)
.or_else(|_| sock.set_send_buffer_size(4 * 1024 * 1024))
.or_else(|_| sock.set_send_buffer_size(1024 * 1024))
.context("set SO_SNDBUF")?;
sock.bind(&"0.0.0.0:0".parse::<std::net::SocketAddr>().unwrap().into())
.context("bind socket")?;
let port = sock.local_addr().context("sock local addr")?.as_socket().unwrap().port();
let socket: std::net::UdpSocket = sock.into();
socket.set_nonblocking(true).context("set nonblocking")?;
let runtime = quinn::default_runtime()
.ok_or_else(|| anyhow::anyhow!("no async runtime found"))?;
let mut ep = Endpoint::new_with_abstract_socket(
quinn::EndpointConfig::default(),
Some(server_config()?),
runtime.wrap_udp_socket(socket)?,
runtime,
)?;
ep.set_default_client_config(client_config()?);
Ok((ep, port))
}
fn keying_material(conn: &quinn::Connection) -> Result<[u8; 32]> {
let mut out = [0u8; 32];
conn.export_keying_material(&mut out, b"filament-direct-auth", b"")
.map_err(|e| anyhow!("export_keying_material failed: {e:?}"))?;
Ok(out)
}
fn auth_tag(tkey: &[u8; 32], km: &[u8; 32], who: &str) -> [u8; 32] {
let mut msg = Vec::with_capacity(32 + 16);
msg.extend_from_slice(km);
msg.push(b'|');
msg.extend_from_slice(who.as_bytes());
hmac_sha256_raw(tkey, &msg)
}
async fn authenticate(
conn: &quinn::Connection,
tkey: &[u8; 32],
is_dialer: bool,
) -> Result<(SendStream, RecvStream)> {
let km = keying_material(conn)?;
let (my_who, their_who) = if is_dialer { ("dialer", "acceptor") } else { ("acceptor", "dialer") };
let my_tag = auth_tag(tkey, &km, my_who);
let their_expected = auth_tag(tkey, &km, their_who);
let (mut send, mut recv) = if is_dialer {
conn.open_bi().await.context("open auth stream")?
} else {
conn.accept_bi().await.context("accept auth stream")?
};
send.write_all(&my_tag).await.context("write auth tag")?;
let mut peer_tag = [0u8; 32];
recv.read_exact(&mut peer_tag).await.context("read auth tag")?;
if !ct_eq(&peer_tag, &their_expected) {
bail!("DIRECT-AUTH-FAIL: pair-secret MAC mismatch, rejecting peer");
}
Ok((send, recv))
}
pub async fn serve_tun_listen(bind: SocketAddr, secret: &[u8; 32]) -> Result<quinn::Connection> {
let ep = Endpoint::server(server_config()?, bind).context("bind serve-tun listener")?;
let secret = *secret;
let (tx, mut rx) = tokio::sync::mpsc::channel::<quinn::Connection>(1);
loop {
tokio::select! {
Some(conn) = rx.recv() => {
keep_endpoint_alive(ep, &conn);
return Ok(conn);
}
incoming = ep.accept() => {
let Some(incoming) = incoming else {
return Err(anyhow!("listener closed before a peer authenticated"));
};
let tx = tx.clone();
tokio::spawn(async move {
let ok = tokio::time::timeout(std::time::Duration::from_secs(10), async {
let conn = incoming.await?;
authenticate(&conn, &secret, false).await?;
Ok::<_, anyhow::Error>(conn)
})
.await;
if let Ok(Ok(conn)) = ok {
let _ = tx.send(conn).await;
}
});
}
}
}
}
pub async fn serve_tun_connect(peer: SocketAddr, secret: &[u8; 32]) -> Result<quinn::Connection> {
let mut ep = Endpoint::client("0.0.0.0:0".parse().unwrap()).context("bind serve-tun client")?;
ep.set_default_client_config(client_config()?);
let mut last = anyhow!("serve-tun connect: no attempt made");
for _ in 0..8 {
let attempt = async {
let conn = ep.connect(peer, "filament-direct")?.await?;
authenticate(&conn, secret, true).await?;
Ok::<_, anyhow::Error>(conn)
};
match attempt.await {
Ok(conn) => {
keep_endpoint_alive(ep, &conn);
return Ok(conn);
}
Err(e) => {
last = e;
tokio::time::sleep(std::time::Duration::from_millis(800)).await;
}
}
}
Err(last.context("serve-tun connect failed after retries"))
}
fn keep_endpoint_alive(ep: Endpoint, conn: &quinn::Connection) {
let c = conn.clone();
let ep_addr = ep.local_addr().ok();
let conn_id = conn.stable_id();
tokio::spawn(async move {
c.closed().await;
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[KEEP-EP-DROP] conn_id={conn_id} ep={ep_addr:?} — connection closed, dropping endpoint");
drop(ep);
});
}
pub struct DirectTransport {
conn: quinn::Connection,
send: Arc<Mutex<SendStream>>,
last_activity: Arc<std::sync::atomic::AtomicU64>,
dead: Arc<std::sync::atomic::AtomicBool>,
answerer: bool,
ep: Option<quinn::Endpoint>,
#[cfg(feature = "test-hooks")]
sent_data: std::sync::atomic::AtomicU64,
#[cfg(feature = "test-hooks")]
frozen: std::sync::atomic::AtomicBool,
#[cfg(feature = "test-hooks")]
born_ms: u64,
}
impl Drop for DirectTransport {
fn drop(&mut self) {
let ep_info = self.ep.as_ref().map(|e| format!("ep={:?}", e.local_addr())).unwrap_or_else(|| "ep=None".to_string());
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[DROP] DirectTransport stable_id={} answerer={} {ep_info} close_reason={:?}",
self.conn.stable_id(), self.answerer, self.conn.close_reason());
}
}
const KIND_CONTROL: u8 = 0;
const KIND_DATA: u8 = 1;
fn now_ms() -> u64 {
use std::sync::OnceLock;
static EPOCH: OnceLock<std::time::Instant> = OnceLock::new();
EPOCH.get_or_init(std::time::Instant::now).elapsed().as_millis() as u64
}
impl DirectTransport {
async fn write_framed(&self, kind: u8, payload: &[u8]) -> Result<()> {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return Err(anyhow!("direct connection closed"));
}
let mut hdr = [0u8; 5];
hdr[0] = kind;
hdr[1..5].copy_from_slice(&(payload.len() as u32).to_be_bytes());
let mut s = self.send.lock().await;
if let Err(e) = s.write_all(&hdr).await {
self.dead.store(true, std::sync::atomic::Ordering::Relaxed);
return Err(anyhow!("direct write hdr: {e}"));
}
if let Err(e) = s.write_all(payload).await {
self.dead.store(true, std::sync::atomic::Ordering::Relaxed);
return Err(anyhow!("direct write body: {e}"));
}
Ok(())
}
}
#[async_trait]
impl Transport for DirectTransport {
async fn send_control(&self, msg: &Value) -> Result<()> {
let text = msg.to_string();
self.write_framed(KIND_CONTROL, text.as_bytes()).await
}
async fn send_frame(&self, sid: u32, offset: u64, payload: &[u8]) -> Result<()> {
#[cfg(feature = "test-hooks")]
let unblocked = match direct_unblock_after_ms() {
Some(after) => self.born_ms >= after,
None => false,
};
#[cfg(feature = "test-hooks")]
if unblocked && !self.frozen.load(std::sync::atomic::Ordering::Relaxed) {
if !direct_flaky_upgrade() {
let mut framed = Vec::with_capacity(4 + 8 + payload.len());
framed.extend_from_slice(&sid.to_be_bytes());
framed.extend_from_slice(&offset.to_be_bytes());
framed.extend_from_slice(payload);
self.write_framed(KIND_DATA, &framed).await?;
self.last_activity.store(now_ms(), std::sync::atomic::Ordering::Relaxed);
return Ok(());
}
let prior = self
.sent_data
.fetch_add(payload.len() as u64, std::sync::atomic::Ordering::Relaxed);
if prior + (payload.len() as u64) >= FLAKY_REFREEZE_BYTES {
self.frozen.store(true, std::sync::atomic::Ordering::Relaxed);
eprintln!("[test] FLAKY direct standby re-froze at {prior} bytes, verify must discard it");
loop {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return Err(anyhow!("direct connection closed (flaky standby discarded)"));
}
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
}
}
let mut framed = Vec::with_capacity(4 + 8 + payload.len());
framed.extend_from_slice(&sid.to_be_bytes());
framed.extend_from_slice(&offset.to_be_bytes());
framed.extend_from_slice(payload);
self.write_framed(KIND_DATA, &framed).await?;
self.last_activity.store(now_ms(), std::sync::atomic::Ordering::Relaxed);
return Ok(());
}
#[cfg(feature = "test-hooks")]
if let Some(limit) = freeze_after_bytes() {
if !self.frozen.load(std::sync::atomic::Ordering::Relaxed) {
let prior = self
.sent_data
.fetch_add(payload.len() as u64, std::sync::atomic::Ordering::Relaxed);
if freeze_persist() {
let already_froze = FROZE_ONCE.load(std::sync::atomic::Ordering::SeqCst);
let cross = prior + (payload.len() as u64) >= limit;
if already_froze || cross {
FROZE_ONCE.store(true, std::sync::atomic::Ordering::SeqCst);
self.frozen.store(true, std::sync::atomic::Ordering::Relaxed);
eprintln!("[test] data-path FREEZE engaged at {} bytes, black-holing this transport", prior);
}
} else if prior + (payload.len() as u64) >= limit
&& !FROZE_ONCE.swap(true, std::sync::atomic::Ordering::SeqCst)
{
self.frozen.store(true, std::sync::atomic::Ordering::Relaxed);
eprintln!("[test] data-path FREEZE engaged at {} bytes, black-holing this transport", prior);
}
}
if self.frozen.load(std::sync::atomic::Ordering::Relaxed) {
loop {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return Err(anyhow!("direct connection closed (frozen path repaired)"));
}
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
}
}
}
let mut framed = Vec::with_capacity(4 + 8 + payload.len());
framed.extend_from_slice(&sid.to_be_bytes());
framed.extend_from_slice(&offset.to_be_bytes());
framed.extend_from_slice(payload);
self.write_framed(KIND_DATA, &framed).await?;
self.last_activity.store(now_ms(), std::sync::atomic::Ordering::Relaxed);
Ok(())
}
async fn flush(&self) -> Result<()> {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return Err(anyhow!("direct connection closed while flushing"));
}
Ok(())
}
async fn drain_finish(&self) -> Result<()> {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return Ok(()); }
let stopped = {
let mut s = self.send.lock().await;
let _ = s.finish(); s.stopped()
};
match tokio::time::timeout(std::time::Duration::from_secs(180), stopped).await {
Ok(Ok(_)) => Ok(()),
Ok(Err(e)) => Err(anyhow!("direct drain: peer dropped before full ack: {e}")),
Err(_) => Err(anyhow!("direct drain: timed out after 180s awaiting peer ack")),
}
}
fn max_payload(&self) -> usize {
MAX_DIRECT_PAYLOAD
}
fn supports_datagrams(&self) -> bool {
true
}
fn is_dead(&self) -> bool {
self.dead.load(std::sync::atomic::Ordering::Relaxed)
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn force_close(&self) {
self.conn.close(0u32.into(), b"dropped");
}
async fn open_stream(&self) -> Option<(quinn::SendStream, quinn::RecvStream, quinn::Connection)> {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return None;
}
match self.conn.open_bi().await {
Ok(s) => Some((s.0, s.1, self.conn.clone())),
Err(_) => None,
}
}
fn send_datagram(&self, packet: &[u8]) -> Result<()> {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return Err(anyhow!("direct connection closed"));
}
self.conn
.send_datagram(bytes::Bytes::copy_from_slice(packet))
.map_err(|e| anyhow!("send_datagram: {e}"))
}
async fn recv_datagram(&self) -> Result<bytes::Bytes> {
self.conn
.read_datagram()
.await
.map_err(|e| anyhow!("read_datagram: {e}"))
}
fn max_datagram_size(&self) -> Option<usize> {
self.conn.max_datagram_size()
}
fn channel_binding(&self) -> Option<Vec<u8>> {
keying_material(&self.conn).ok().map(|km| km.to_vec())
}
fn sid_answerer(&self) -> bool {
self.answerer
}
fn idle_ms(&self) -> u64 {
if self.dead.load(std::sync::atomic::Ordering::Relaxed) {
return u64::MAX;
}
let _ = &self.conn; now_ms().saturating_sub(self.last_activity.load(std::sync::atomic::Ordering::Relaxed))
}
fn is_alive(&self) -> bool {
!self.dead.load(std::sync::atomic::Ordering::Relaxed)
&& self.conn.close_reason().is_none()
}
fn remote_addr(&self) -> Option<std::net::SocketAddr> {
Some(self.conn.remote_address())
}
fn rtt_ms(&self) -> Option<u64> {
Some(self.conn.rtt().as_millis() as u64)
}
fn local_ip(&self) -> Option<std::net::IpAddr> {
self.conn.local_ip()
}
}
fn spawn_reader(
peer_id: String,
mut recv: RecvStream,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
last_activity: Arc<std::sync::atomic::AtomicU64>,
dead: Arc<std::sync::atomic::AtomicBool>,
answerer: bool,
) {
tokio::spawn(async move {
loop {
let mut hdr = [0u8; 5];
if let Err(e) = recv.read_exact(&mut hdr).await {
if matches!(&e, quinn::ReadExactError::FinishedEarly(_)) {
break;
}
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[DEAD] spawn_reader: peer={} answerer={answerer} read hdr error: {e:?}", peer_id);
dead.store(true, std::sync::atomic::Ordering::Relaxed);
break;
}
let kind = hdr[0];
let len = u32::from_be_bytes([hdr[1], hdr[2], hdr[3], hdr[4]]) as usize;
if len > MAX_DIRECT_PAYLOAD + 64 {
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[DEAD] spawn_reader: peer={} absurd len={}", peer_id, len);
dead.store(true, std::sync::atomic::Ordering::Relaxed);
break;
}
let mut body = vec![0u8; len];
if let Err(e) = recv.read_exact(&mut body).await {
if matches!(&e, quinn::ReadExactError::FinishedEarly(_)) {
break;
}
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[DEAD] spawn_reader: peer={} read body error kind={} err={e:?}", peer_id, kind);
dead.store(true, std::sync::atomic::Ordering::Relaxed);
break;
}
match kind {
KIND_CONTROL => {
if let Ok(v) = serde_json::from_slice::<Value>(&body) {
let _ = tx.send(crate::net::Ev::Control(peer_id.clone(), v));
}
}
KIND_DATA => {
if body.len() >= 12 {
last_activity.store(now_ms(), std::sync::atomic::Ordering::Relaxed);
let sid = u32::from_be_bytes([body[0], body[1], body[2], body[3]]);
let offset = u64::from_be_bytes([
body[4], body[5], body[6], body[7],
body[8], body[9], body[10], body[11],
]);
let bytes = bytes::Bytes::from(body);
let _ = tx.send(crate::net::Ev::Chunk(
peer_id.clone(),
sid,
Some(offset),
bytes.slice(12..),
));
}
}
_ => {}
}
}
});
}
pub fn make_transport(
peer_id: String,
conn: quinn::Connection,
send: SendStream,
recv: RecvStream,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
answerer: bool,
ep: Option<quinn::Endpoint>,
) -> Arc<dyn Transport> {
let last_activity = Arc::new(std::sync::atomic::AtomicU64::new(now_ms()));
let dead = Arc::new(std::sync::atomic::AtomicBool::new(false));
spawn_reader(peer_id, recv, tx, last_activity.clone(), dead.clone(), answerer);
Arc::new(DirectTransport {
conn,
send: Arc::new(Mutex::new(send)),
last_activity,
dead,
answerer,
ep,
#[cfg(feature = "test-hooks")]
sent_data: std::sync::atomic::AtomicU64::new(0),
#[cfg(feature = "test-hooks")]
frozen: std::sync::atomic::AtomicBool::new(false),
#[cfg(feature = "test-hooks")]
born_ms: now_ms(),
})
}
pub fn spawn_mesh_accept(
pid: String,
t: &std::sync::Arc<dyn crate::net::Transport>,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
) {
if let Some(dt) = t.as_any().downcast_ref::<DirectTransport>() {
let conn = dt.conn.clone();
tokio::spawn(async move {
loop {
match conn.accept_bi().await {
Ok((send, recv)) => {
let conn2 = conn.clone();
let worker = make_transport(pid.clone(), conn2, send, recv, tx.clone(), true, None);
let _ = tx.send(crate::net::Ev::DirectWorkersReady(pid.clone(), vec![worker]));
}
Err(_) => return,
}
}
});
}
}
pub async fn dial_workers(
ports: Vec<u16>,
peer_ip: IpAddr,
tkey: [u8; 32],
peer_id: String,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
count: usize,
) -> Vec<Arc<dyn Transport>> {
const WORKER_BUDGET: Duration = Duration::from_secs(5);
if count == 0 || ports.is_empty() {
return vec![];
}
let mut workers = Vec::with_capacity(count.min(ports.len()));
let deadline = tokio::time::sleep(WORKER_BUDGET);
tokio::pin!(deadline);
use futures_util::stream::StreamExt;
let mut futs: futures_util::stream::FuturesUnordered<_> = ports
.into_iter()
.map(|port| {
let tkey = tkey;
Box::pin(async move {
let (ep, _) = bind_endpoint()?;
let peer = SocketAddr::new(peer_ip, port);
let connecting =
ep.connect(peer, "filament-direct").context("worker connect")?;
let conn = connecting.await.context("worker dial handshake")?;
let (s, r) = authenticate(&conn, &tkey, true).await?;
Ok::<_, anyhow::Error>((conn, s, r, ep))
})
})
.collect();
loop {
tokio::select! {
result = futs.next() => {
match result {
Some(Ok((conn, send, recv, ep))) => {
workers.push(make_transport(
peer_id.clone(), conn, send, recv, tx.clone(), false, Some(ep),
));
if workers.len() >= count { break; }
}
Some(Err(e)) => {
crate::ui::debug(&format!("worker dial: {e}"));
}
None => break,
}
}
_ = &mut deadline => break,
}
}
workers
}
pub async fn accept_workers(
endpoints: Vec<Endpoint>,
tkey: [u8; 32],
peer_id: String,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
count: usize,
) -> Vec<Arc<dyn Transport>> {
const WORKER_BUDGET: Duration = Duration::from_secs(5);
if count == 0 || endpoints.is_empty() {
return vec![];
}
let mut workers = Vec::with_capacity(count.min(endpoints.len()));
let deadline = tokio::time::sleep(WORKER_BUDGET);
tokio::pin!(deadline);
use futures_util::stream::StreamExt;
let mut futs: futures_util::stream::FuturesUnordered<_> = endpoints
.into_iter()
.map(|ep| {
let tkey = tkey;
Box::pin(async move {
match ep.accept().await {
Some(incoming) => {
let conn = incoming.await.context("worker accept handshake")?;
let (s, r) = authenticate(&conn, &tkey, false).await?;
Ok::<_, anyhow::Error>((conn, s, r, ep))
}
None => bail!("worker endpoint closed before connection arrived"),
}
})
})
.collect();
loop {
tokio::select! {
result = futs.next() => {
match result {
Some(Ok((conn, send, recv, ep))) => {
workers.push(make_transport(
peer_id.clone(), conn, send, recv, tx.clone(), true, Some(ep),
));
if workers.len() >= count { break; }
}
Some(Err(e)) => {
crate::ui::debug(&format!("worker accept: {e}"));
}
None => break,
}
}
_ = &mut deadline => break,
}
}
workers
}
pub async fn race_connect(
endpoint: Endpoint,
peer_cands: Vec<String>,
secret: &str,
peer_id: String,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
answerer: bool,
) -> Option<Arc<dyn Transport>> {
race_connect_labeled(endpoint, peer_cands, secret, peer_id, tx, "direct-quic", answerer).await
}
pub async fn race_connect_labeled(
endpoint: Endpoint,
peer_cands: Vec<String>,
secret: &str,
peer_id: String,
tx: tokio::sync::mpsc::UnboundedSender<crate::net::Ev>,
route: &str,
answerer: bool,
) -> Option<Arc<dyn Transport>> {
if route == "direct-quic" && test_block() {
eprintln!("filament: DIRECT-BLOCKED (test), forcing WebRTC fallback");
tokio::time::sleep(DIRECT_BUDGET).await;
endpoint.close(0u32.into(), b"test-block");
return None;
}
let tkey = transport_key(secret);
async fn auth_conn(
conn: quinn::Connection,
tkey: [u8; 32],
is_dialer: bool,
) -> Result<(quinn::Connection, SendStream, RecvStream)> {
let (s, r) = authenticate(&conn, &tkey, is_dialer).await?;
Ok((conn, s, r))
}
let mut futs: Vec<std::pin::Pin<Box<dyn std::future::Future<Output = Result<(quinn::Connection, SendStream, RecvStream)>> + Send>>> = Vec::new();
{
let ep = endpoint.clone();
let tkey = tkey;
futs.push(Box::pin(async move {
let incoming = ep.accept().await.ok_or_else(|| anyhow!("endpoint closed"))?;
let conn = incoming.await.context("accept handshake")?;
auth_conn(conn, tkey, false).await
}));
}
for cand in peer_cands {
let Ok(addr) = cand.parse::<SocketAddr>() else { continue };
let ep = endpoint.clone();
let tkey = tkey;
futs.push(Box::pin(async move {
let connecting = ep
.connect(addr, "filament-direct")
.context("connect")?;
let conn = connecting.await.context("dial handshake")?;
auth_conn(conn, tkey, true).await
}));
}
let race = async {
use futures_util::stream::{FuturesUnordered, StreamExt};
let mut set: FuturesUnordered<_> = futs.into_iter().collect();
while let Some(res) = set.next().await {
match res {
Ok((conn, send, recv)) => return Some((conn, send, recv)),
Err(e) => {
let s = e.to_string();
if s.contains("DIRECT-AUTH-FAIL") {
crate::ui::trace(&format!("filament: {s}"));
}
continue;
}
}
}
None
};
let winner = match tokio::time::timeout(DIRECT_BUDGET, race).await {
Ok(Some(w)) => w,
_ => {
endpoint.close(0u32.into(), b"direct-timeout");
return None;
}
};
let (conn, send, recv) = winner;
crate::ui::debug(&format!(
"filament: DIRECT-CONNECT ok (route: {}) peer={} remote={}",
route,
peer_id,
conn.remote_address()
));
let ep_for_transport = endpoint.clone(); {
let conn2 = conn.clone();
tokio::spawn(async move {
conn2.closed().await;
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[KEEP-EP-DROP] endpoint for primary conn_stable={} ep={:?}",
conn2.stable_id(), endpoint.local_addr());
drop(endpoint);
});
}
let conn_id = conn.stable_id();
#[cfg(debug_assertions)]
#[cfg(debug_assertions)]
eprintln!("[PRIMARY-MADE] conn_stable={conn_id} answerer={answerer}");
Some(make_transport(peer_id, conn, send, recv, tx, answerer, Some(ep_for_transport)))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hkdf_is_deterministic_and_independent() {
let a = transport_key("secret-one");
let b = transport_key("secret-one");
let c = transport_key("secret-two");
assert_eq!(a, b, "same secret -> same key");
assert_ne!(a, c, "different secret -> different key");
assert_ne!(&a[..], b"secret-one\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0\0");
}
#[test]
fn pubip_cache_freshness() {
use std::time::Duration;
let ttl = Duration::from_secs(300);
let server = "https://api.filament.autumated.com";
assert!(pubip_fresh(server, Duration::from_secs(10), server, ttl));
assert!(!pubip_fresh(server, ttl, server, ttl));
assert!(!pubip_fresh(server, Duration::from_secs(600), server, ttl));
assert!(!pubip_fresh("https://other.example.com", Duration::from_secs(1), server, ttl));
}
#[test]
fn ct_eq_basic() {
assert!(ct_eq(b"abc", b"abc"));
assert!(!ct_eq(b"abc", b"abd"));
assert!(!ct_eq(b"abc", b"ab"));
}
#[test]
fn auth_tags_directional_and_secret_bound() {
let k1 = transport_key("right");
let k2 = transport_key("wrong");
let km = [7u8; 32];
let dialer = auth_tag(&k1, &km, "dialer");
let acceptor = auth_tag(&k1, &km, "acceptor");
assert_ne!(dialer, acceptor, "direction-tagged tags differ");
assert_ne!(auth_tag(&k2, &km, "dialer"), dialer);
}
}