use std::collections::HashMap;
use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
use quinn::VarInt;
use tokio::sync::Mutex;
use weida_core::{EndpointAddr, Error, Fingerprint};
use weida_protocol::codes;
use crate::config::{ClientTls, Discovery, RuntimeConfig};
use crate::conn::{ConnCtx, ConnHandle, conn_error};
use crate::listener::Namespace;
use crate::runtime::{Exec, Shared};
use crate::tls;
use crate::transport::Link;
fn dial_port(written: Option<u16>, discovery: Discovery, host: &str) -> Result<u16, Error> {
match (written, discovery) {
(Some(port), _) => Ok(port),
(None, Discovery::Aware) => Ok(weida_core::DEFAULT_PORT),
(None, Discovery::Single) => Err(Error::InvalidAddress(format!(
"{host} names no port, which means a set of nodes, and this runtime is \
configured `Discovery::Single`: write the port, or allow discovery"
))),
}
}
pub(crate) struct ClientPool {
state: Mutex<PoolState>,
shared: Arc<Shared>,
}
type PeerKey = (String, u16, Arc<ClientTls>, Option<Fingerprint>);
type ConnKey = (PeerKey, String);
#[derive(Default)]
struct PoolState {
endpoint: Option<quinn::Endpoint>,
connections: HashMap<ConnKey, ConnHandle>,
}
impl ClientPool {
pub(crate) fn new(shared: Arc<Shared>) -> ClientPool {
ClientPool {
state: Mutex::new(PoolState::default()),
shared,
}
}
pub(crate) async fn take_endpoint(&self) -> Option<quinn::Endpoint> {
let mut state = self.state.lock().await;
state.connections.clear();
state.endpoint.take()
}
pub(crate) async fn connect(
&self,
config: &RuntimeConfig,
exec: &Exec,
addr: &EndpointAddr,
tls: &Arc<ClientTls>,
) -> Result<ConnHandle, Error> {
let EndpointAddr {
host,
port,
path,
peer: expected,
} = addr;
let (host, expected) = (host.as_str(), *expected);
let written_port = *port;
let port = dial_port(written_port, config.discovery, host)?;
let peer: PeerKey = (host.to_owned(), port, Arc::clone(tls), expected);
let key: ConnKey = (peer.clone(), path.clone());
let (endpoint, known) = {
let mut state = self.state.lock().await;
if let Some(existing) = state.connections.get(&key) {
if existing.conn.close_reason().is_none() {
return Ok(ConnHandle::clone(existing));
}
state.connections.remove(&key);
}
state
.connections
.retain(|_, handle| handle.conn.close_reason().is_none());
let endpoint = match &state.endpoint {
Some(endpoint) => endpoint.clone(),
None => {
let endpoint = bind_client_endpoint(exec)?;
state.endpoint = Some(endpoint.clone());
endpoint
}
};
(endpoint, state.peer_identity(&peer))
};
let (client_config, refused) = tls::client_config(tls, expected, &config.limits)?;
let addrs = config
.resolver
.resolve(exec, host, written_port, config.max_resolved_addresses)
.await?;
let mut last: Option<Error> = None;
let mut established = None;
for (index, addr) in addrs.iter().enumerate() {
let is_last = index + 1 == addrs.len();
tracing::debug!(%addr, host, is_last, "dialling");
let connecting = {
let _guard = exec.enter();
match endpoint.connect_with(client_config.clone(), *addr, host) {
Ok(connecting) => connecting,
Err(e) => {
last = Some(Error::Transport(format!("connect to {addr} failed: {e}")));
continue;
}
}
};
let attempt = exec.spawn(connecting);
let joined = if is_last {
Some(attempt.await)
} else {
exec.within(config.connect_attempt_timeout, attempt).await
};
let handshake = match joined {
Some(joined) => {
joined.map_err(|e| Error::Runtime(format!("dial task failed: {e}")))?
}
None => {
last = Some(Error::Transport(format!(
"connect to {addr} did not answer within {:?}",
config.connect_attempt_timeout
)));
continue;
}
};
match handshake {
Ok(conn) => {
established = Some(conn);
break;
}
Err(e) => {
let refused = refused.lock().expect("refusal record poisoned").take();
if let Some(presented) = refused {
return Err(Error::Untrusted(presented));
}
last = Some(conn_error(e));
}
}
}
let conn = match established {
Some(conn) => conn,
None => {
return Err(last.unwrap_or_else(|| {
Error::InvalidAddress(format!("{host}:{port} resolved to no addresses"))
}));
}
};
let presented = crate::tls::peer_fingerprint(&conn);
if let Some(known) = known
&& known != presented
{
conn.close(
VarInt::from_u32(codes::SHUTDOWN as u32),
b"peer identity differs from this peer's other connections",
);
return Err(match presented {
Some(fp) => Error::Untrusted(fp),
None => Error::Tls(
"the peer that answered proved no identity, so it cannot be the peer this \
runtime is already connected to"
.into(),
),
});
}
let handle = ConnCtx::spawn(
Link::Quic(conn),
config.limits,
Arc::new(Namespace::new()),
None,
exec.clone(),
config.guarantees,
Arc::clone(&self.shared),
);
handle.negotiated().await?;
let mut state = self.state.lock().await;
let winner = state
.connections
.get(&key)
.filter(|held| held.conn.close_reason().is_none())
.map(ConnHandle::clone);
if let Some(winner) = winner {
handle
.conn
.close(codes::SHUTDOWN, "a concurrent dial won this pool key");
return Ok(winner);
}
state.connections.insert(key, ConnHandle::clone(&handle));
Ok(handle)
}
}
impl PoolState {
fn peer_identity(&self, peer: &PeerKey) -> Option<Option<Fingerprint>> {
self.connections
.iter()
.find(|((p, _), handle)| p == peer && handle.conn.close_reason().is_none())
.map(|(_, handle)| handle.peer.as_ref().and_then(|id| id.key()))
}
}
fn bind_client_endpoint(exec: &Exec) -> Result<quinn::Endpoint, Error> {
let _guard = exec.enter();
let v6 = SocketAddr::from((Ipv6Addr::UNSPECIFIED, 0));
match quinn::Endpoint::client(v6) {
Ok(endpoint) => Ok(endpoint),
Err(e) => {
tracing::debug!(error = %e, "IPv6 client socket unavailable; falling back to IPv4");
quinn::Endpoint::client(SocketAddr::from((Ipv4Addr::UNSPECIFIED, 0))).map_err(Error::Io)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_dial_port_follows_the_url_then_the_discovery_mode() {
assert_eq!(
dial_port(Some(9000), Discovery::Aware, "h").expect("written"),
9000
);
assert_eq!(
dial_port(Some(9000), Discovery::Single, "h").expect("written"),
9000,
"a written port is never overridden by the mode"
);
assert_eq!(
dial_port(None, Discovery::Aware, "jobs.example").expect("discovered"),
weida_core::DEFAULT_PORT
);
let refused = dial_port(None, Discovery::Single, "jobs.example")
.expect_err("a set is not one endpoint");
let message = refused.to_string();
assert!(
message.contains("jobs.example") && message.contains("Discovery::Single"),
"the refusal names the host and the reason: {message}"
);
}
#[tokio::test]
async fn the_client_socket_binds() {
let exec = Exec::current().expect("ambient runtime");
let endpoint = bind_client_endpoint(&exec).unwrap();
assert_ne!(endpoint.local_addr().unwrap().port(), 0);
}
}