use std::net::SocketAddr;
use std::sync::Arc;
use dig_tls::{BindingPolicy, NodeCert};
use tokio_rustls::TlsAcceptor;
use crate::error::MethodError;
use crate::method::TraversalKind;
use crate::mux::PeerSession;
use crate::peer::PeerConnection;
use crate::relay::RelayTunnel;
use crate::tunnel::RelayTunnelStream;
fn unspecified_addr() -> SocketAddr {
SocketAddr::from(([0u8; 16], 0))
}
#[derive(Clone)]
pub struct RelayAcceptor {
node: Arc<NodeCert>,
binding_policy: BindingPolicy,
relay_endpoint: SocketAddr,
}
impl RelayAcceptor {
pub fn new(node: Arc<NodeCert>) -> Self {
RelayAcceptor {
node,
binding_policy: BindingPolicy::default(),
relay_endpoint: unspecified_addr(),
}
}
pub fn with_binding_policy(mut self, policy: BindingPolicy) -> Self {
self.binding_policy = policy;
self
}
pub fn with_relay_endpoint(mut self, endpoint: SocketAddr) -> Self {
self.relay_endpoint = endpoint;
self
}
pub async fn accept(&self, tunnel: RelayTunnel) -> Result<PeerConnection, MethodError> {
let kind = TraversalKind::Relayed;
let server_tls = dig_tls::server_config_spki_pinned(&self.node, self.binding_policy)
.map_err(|e| MethodError::failed(kind, format!("server cert config: {e}")))?;
let captured = server_tls.captured_peer_id;
let captured_bls = server_tls.captured_bls;
let acceptor = TlsAcceptor::from(server_tls.config);
let stream = RelayTunnelStream::new(tunnel);
let tls = acceptor
.accept(stream)
.await
.map_err(|e| MethodError::failed(kind, format!("mtls accept: {e}")))?;
let verified = captured
.get()
.ok_or_else(|| MethodError::failed(kind, "peer presented no certificate"))?;
let session = PeerSession::server(tls);
Ok(PeerConnection {
peer_id: verified,
method: kind,
remote_addr: self.relay_endpoint,
peer_bls_pub: captured_bls.get(),
session,
})
}
}
impl std::fmt::Debug for RelayAcceptor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RelayAcceptor")
.field("binding_policy", &self.binding_policy)
.field("relay_endpoint", &self.relay_endpoint)
.finish_non_exhaustive()
}
}