use std::net::{Ipv6Addr, SocketAddr};
use std::os::unix::fs::PermissionsExt;
use std::sync::Arc;
use bytes::Bytes;
use tokio::net::{TcpListener, UnixListener};
use crate::api::common::VERSIONS;
use crate::models::{ConnectionID, Limits, Port, Version};
use crate::protocol::base::{AnyConnection, Connection, Transport};
use crate::protocol::common::Error;
use crate::api::cluster::Cluster;
use crate::api::gate::{Gate, Permit};
use crate::protocol::handler::{Incoming, Negotiation};
use crate::protocol::quic::{self, QUICIncoming};
use crate::tls::Identity;
use crate::helpers::sync;
#[derive(Clone)]
pub struct ServerConfig {
pub versions: Vec<Version>,
pub limits: ServerLimits,
pub identity: Option<Identity>,
pub tls: crate::tls::TLSConfig,
pub ech: Option<crate::tls::ECHKeys>,
pub hsts: Option<crate::hsts::HSTSPolicy>,
pub reuseport: bool,
pub uds_mode: u32,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
versions: VERSIONS.to_vec(),
limits: ServerLimits::default(),
identity: None,
tls: crate::tls::TLSConfig::default(),
ech: None,
hsts: None,
reuseport: true,
uds_mode: 0o666,
}
}
}
#[derive(Clone, Default)]
pub struct Server {
pub config: ServerConfig,
}
impl Server {
pub fn new(config: ServerConfig) -> Self {
Self { config }
}
pub async fn bind(&self, target: Port) -> Result<Listener, Error> {
let socket = self.open(&target)?;
self.attach(&target, socket).await
}
pub fn open(&self, target: &Port) -> Result<RawSocket, Error> {
Ok(match target {
Port::UDS(path) => {
let listener = std::os::unix::net::UnixListener::bind(path)?;
listener.set_nonblocking(true)?;
if self.config.uds_mode != 0 {
std::fs::set_permissions(path, std::fs::Permissions::from_mode(self.config.uds_mode))?;
}
RawSocket::UDS(listener)
}
Port::TCP(port) => RawSocket::TCP(self.socket(*port, socket2::Type::STREAM)?.into()),
Port::QUIC(port) => RawSocket::QUIC(self.socket(*port, socket2::Type::DGRAM)?.into()),
})
}
pub fn socket(&self, port: u16, kind: socket2::Type) -> Result<socket2::Socket, Error> {
let socket = socket2::Socket::new(socket2::Domain::IPV6, kind, None)?;
if self.config.reuseport {
socket.set_reuse_port(true)?;
}
socket.set_nonblocking(true)?;
if kind == socket2::Type::STREAM {
socket.set_reuse_address(true)?;
}
socket.bind(&SocketAddr::from((Ipv6Addr::UNSPECIFIED, port)).into())?;
if kind == socket2::Type::STREAM {
socket.listen(self.config.limits.backlog.min(i32::MAX as u32) as i32)?;
}
Ok(socket)
}
pub async fn attach(&self, target: &Port, socket: RawSocket) -> Result<Listener, Error> {
let versions = target.offers(&self.config.versions);
if versions.is_empty() {
return Err(Error::Version(format!("no configured version runs over {:?}", target.transport())));
}
let socket = match socket {
RawSocket::UDS(listener) => Socket::UDS(UnixListener::from_std(listener)?),
RawSocket::TCP(listener) => Socket::TCP(TcpListener::from_std(listener)?),
RawSocket::QUIC(udp) => {
let identity = self
.config
.identity
.as_ref()
.ok_or_else(|| Error::TLS("a QUIC port needs a certificate and a key".into()))?;
let config = quic::QUICConfig {
versions: versions.clone(),
idle_timeout: self.config.limits.message.read_timeout,
max_streams_bidi: Some(self.config.limits.message.max_concurrent_streams as u64),
enable_dgram: false,
};
let hook = std::sync::Arc::new(quic::QUICServerTLS {
identity: identity.clone(),
ech: self.config.ech.clone(),
tls: self.config.tls.clone(),
});
let (incoming, address) = quic::QUICListener::bind(udp, &config, hook)?;
Socket::QUIC { incoming: tokio::sync::Mutex::new(incoming), address }
}
};
let acceptor = match (&self.config.identity, &socket) {
(Some(identity), Socket::TCP(_)) => Some(Arc::new(self.config.tls.server(identity, &versions, self.config.ech.as_ref())?)),
_ => None,
};
let response_finalizer = crate::finalizer::ResponseFinalizer::new(self.config.hsts);
let negotiation = Negotiation { versions, limits: self.config.limits.message, acceptor, response_finalizer };
let (negotiating, negotiated) = tokio::sync::mpsc::channel(negotiation.limits.max_pending_handshakes.max(1) as usize);
let gate = self.config.limits.gate();
Ok(Listener { socket, gate, negotiation: Arc::new(negotiation), negotiating, negotiated })
}
}
pub enum RawSocket {
UDS(std::os::unix::net::UnixListener),
TCP(std::net::TcpListener),
QUIC(std::net::UdpSocket),
}
impl RawSocket {
pub fn address(&self) -> Result<SocketAddr, Error> {
match self {
Self::TCP(listener) => Ok(listener.local_addr()?),
Self::QUIC(socket) => Ok(socket.local_addr()?),
Self::UDS(_) => Err(Error::IO(std::io::Error::other("a unix socket has no address"))),
}
}
pub fn share(&self) -> Result<Self, Error> {
Ok(match self {
Self::UDS(listener) => Self::UDS(listener.try_clone()?),
Self::TCP(listener) => Self::TCP(listener.try_clone()?),
Self::QUIC(socket) => Self::QUIC(socket.try_clone()?),
})
}
}
pub enum Socket {
UDS(UnixListener),
TCP(TcpListener),
QUIC {
incoming: tokio::sync::Mutex<tokio::sync::mpsc::Receiver<std::io::Result<QUICIncoming>>>,
address: std::net::SocketAddr,
},
}
impl Socket {
pub fn address(&self) -> Result<SocketAddr, Error> {
match self {
Self::TCP(listener) => Ok(listener.local_addr()?),
Self::QUIC { address, .. } => Ok(*address),
Self::UDS(_) => Err(Error::IO(std::io::Error::other("a unix socket has no address"))),
}
}
pub async fn accept(&self) -> Result<Incoming, Error> {
match self {
Self::QUIC { incoming, .. } => {
let incoming = incoming.lock().await.recv().await.ok_or(Error::Closed)?.map_err(Error::IO)?;
Ok(Incoming::QUIC(incoming))
}
Self::TCP(listener) => {
let (transport, address) = listener.accept().await?;
let _ = transport.set_nodelay(true);
let id = ConnectionID(Bytes::from(address.to_string()));
Ok(Incoming::Stream { transport: Box::new(transport), id, client: Some(address) })
}
Self::UDS(listener) => {
let (transport, _) = listener.accept().await?;
let id = ConnectionID(Bytes::from_static(b"unix"));
Ok(Incoming::Stream { transport: Box::new(transport), id, client: None })
}
}
}
}
pub struct Admitted {
pub connection: AnyConnection,
pub permit: Permit,
}
pub struct Listener {
socket: Socket,
gate: Arc<Gate>,
negotiation: Arc<Negotiation>,
negotiating: tokio::sync::mpsc::Sender<Result<Admitted, Error>>,
negotiated: tokio::sync::mpsc::Receiver<Result<Admitted, Error>>,
}
impl Listener {
pub fn with_gate(mut self, gate: Arc<Gate>) -> Self {
self.gate = gate;
self
}
pub fn versions(&self) -> &[Version] {
&self.negotiation.versions
}
pub fn gate(&self) -> &Arc<Gate> {
&self.gate
}
pub fn limits(&self) -> &Limits {
&self.negotiation.limits
}
pub fn socket(&self) -> &Socket {
&self.socket
}
pub fn negotiation(&self) -> &Negotiation {
&self.negotiation
}
pub fn pending(&self) -> usize {
self.negotiating.max_capacity() - self.negotiating.capacity()
}
pub fn address(&self) -> Result<SocketAddr, Error> {
self.socket.address()
}
pub async fn accept(&mut self) -> Result<Admitted, Error> {
loop {
tokio::select! {
biased;
Some(negotiated) = self.negotiated.recv() => {
if let Ok(admitted) = negotiated {
return Ok(admitted);
}
}
incoming = self.socket.accept(), if self.negotiating.capacity() > 0 => {
let incoming = incoming?;
let Some(permit) = self.gate.admit(incoming.client().map(|address| address.ip()), std::time::Instant::now()) else {
self.negotiation.refuse(incoming);
continue;
};
let Ok(slot) = self.negotiating.clone().try_reserve_owned() else {
self.negotiation.refuse(incoming);
continue;
};
let negotiation = Arc::clone(&self.negotiation);
tokio::spawn(async move { slot.send(negotiation.accept(incoming).await.map(|connection| Admitted { connection, permit })) });
}
}
}
}
}
#[derive(Debug, Clone)]
pub struct ServerLimits {
pub message: Limits,
pub backlog: u32,
pub max_connections: u32,
pub max_connections_per_ip: u32,
pub max_connection_rate: Vec<(f64, u32)>,
pub max_connection_history: usize,
pub worker_stack_size: usize,
}
impl ServerLimits {
pub fn gate(&self) -> Arc<Gate> {
Gate::new(self.max_connections, self.max_connections_per_ip, self.max_connection_rate.clone(), self.max_connection_history)
}
}
impl Default for ServerLimits {
fn default() -> Self {
Self {
message: Limits::default(),
backlog: 1024,
max_connections: 0,
max_connections_per_ip: 0,
max_connection_rate: Vec::new(),
max_connection_history: 1024,
worker_stack_size: 8 * 1024 * 1024,
}
}
}
pub trait Handler: Send + Sync + 'static {
fn websocket_limits(&self) -> crate::websocket::WebSocketLimits {
crate::websocket::WebSocketLimits::default()
}
fn on_connection(&self, connection: AnyConnection) -> impl std::future::Future<Output = ()> + Send {
async move {
let mut connection = connection;
loop {
let Ok(request) = connection.receive().await else {
break;
};
if crate::websocket::Handshake::requested(&request) {
match crate::websocket::Handshake::answer(connection, &request, self.websocket_limits()).await {
crate::websocket::Answer::Accepted(socket) => {
self.on_websocket(socket).await;
return;
}
crate::websocket::Answer::Refused(kept) => {
connection = kept;
continue;
}
crate::websocket::Answer::Failed => return,
}
}
let mut response = crate::models::Message::response(200, connection.version());
response.stream_id = request.stream_id;
response.body = Some(crate::models::Body::Data(Bytes::from_static(b"This is the default response from Soyokaze.")));
if connection.send(response).await.is_err() {
break;
}
if !connection.reusable() {
break;
}
}
connection.close().await;
}
}
fn on_websocket(&self, socket: crate::websocket::WebSocketConnection<Box<dyn Transport>>) -> impl std::future::Future<Output = ()> + Send {
async move {
let mut socket = socket;
socket.close(crate::websocket::CloseCode::InternalError, "WebSocket is not configured").await;
}
}
}
pub struct ServerHandle {
pub shutdown: tokio::sync::watch::Sender<bool>,
pub tasks: Arc<tokio::sync::Mutex<tokio::task::JoinSet<()>>>,
pub accept_loops: Vec<tokio::task::JoinHandle<()>>,
pub addresses: Vec<std::net::SocketAddr>,
}
impl ServerHandle {
pub fn address(&self) -> Option<std::net::SocketAddr> {
self.addresses.first().copied()
}
pub fn addresses(&self) -> &[std::net::SocketAddr] {
&self.addresses
}
pub async fn close(self, timeout: Option<f64>) {
let _ = self.shutdown.send(true);
for accept_loop in self.accept_loops {
let _ = accept_loop.await;
}
let mut tasks = self.tasks.lock().await;
let drain = async {
while tasks.join_next().await.is_some() {}
};
match timeout.and_then(sync::Timeout::duration) {
Some(wait) => {
if tokio::time::timeout(wait, drain).await.is_err() {
tasks.abort_all();
while tasks.join_next().await.is_some() {}
}
}
None => drain.await,
}
}
}
impl Server {
pub async fn serve<H: Handler>(&self, handler: H, ports: &[Port]) -> Result<ServerHandle, Error> {
let gate = self.config.limits.gate();
let mut listeners = Vec::with_capacity(ports.len());
for port in ports {
listeners.push(self.bind(port.clone()).await?);
}
Ok(self.launch(Arc::new(handler), listeners, gate))
}
pub fn launch<H: Handler>(&self, handler: Arc<H>, listeners: Vec<Listener>, gate: Arc<Gate>) -> ServerHandle {
let (shutdown, receiver) = tokio::sync::watch::channel(false);
let tasks = Arc::new(tokio::sync::Mutex::new(tokio::task::JoinSet::new()));
let mut accept_loops = Vec::new();
let mut addresses = Vec::new();
for listener in listeners {
if let Ok(address) = listener.address() {
addresses.push(address);
}
let mut listener = listener.with_gate(gate.clone());
let handler = handler.clone();
let tasks = tasks.clone();
let mut receiver = receiver.clone();
accept_loops.push(tokio::spawn(async move {
loop {
tokio::select! {
_ = receiver.changed() => break,
result = listener.accept() => {
let Ok(Admitted { connection, permit }) = result else {
break;
};
let handler = handler.clone();
let mut set = tasks.lock().await;
while set.try_join_next().is_some() {}
set.spawn(async move {
handler.on_connection(connection).await;
drop(permit);
});
}
}
}
}));
}
ServerHandle { shutdown, tasks, accept_loops, addresses }
}
pub fn run<H: Handler>(&self, handler: H, ports: &[Port], workers: usize) -> Result<Cluster, Error> {
let workers = workers.max(1);
if workers > 1 && !self.config.reuseport && ports.iter().any(|port| matches!(port, Port::QUIC(_))) {
let reason = "a QUIC port needs reuseport to run on more than one worker";
return Err(Error::IO(std::io::Error::other(reason)));
}
let handler = Arc::new(handler);
let gate = self.config.limits.gate();
let mut targets = Vec::with_capacity(ports.len());
let mut queues: Vec<Vec<RawSocket>> = Vec::with_capacity(ports.len());
let mut addresses = Vec::new();
for port in ports {
let opened = self.open(port)?;
let address = opened.address().ok();
let target = match (port, address) {
(Port::TCP(_), Some(address)) => Port::TCP(address.port()),
(Port::QUIC(_), Some(address)) => Port::QUIC(address.port()),
_ => port.clone(),
};
addresses.extend(address);
let independent = self.config.reuseport && !matches!(target, Port::UDS(_));
let mut queue = Vec::with_capacity(workers);
queue.push(opened);
while queue.len() < workers {
queue.push(if independent { self.open(&target)? } else { queue[0].share()? });
}
targets.push(target);
queues.push(queue);
}
let (shutdown, receiver) = tokio::sync::watch::channel(None::<f64>);
let (ready, started) = std::sync::mpsc::channel();
let mut threads = Vec::with_capacity(workers);
let mut failure = None;
for index in 0..workers {
let sockets: Vec<RawSocket> = queues.iter_mut().filter_map(Vec::pop).collect();
let targets = targets.clone();
let server = self.clone();
let handler = handler.clone();
let gate = gate.clone();
let mut receiver = receiver.clone();
let ready = ready.clone();
let worker = move || {
let runtime = match tokio::runtime::Builder::new_current_thread().enable_all().build() {
Ok(runtime) => runtime,
Err(error) => {
let _ = ready.send(Err(Error::IO(error)));
return;
}
};
runtime.block_on(async move {
let mut listeners = Vec::with_capacity(targets.len());
for (target, socket) in targets.iter().zip(sockets) {
match server.attach(target, socket).await {
Ok(listener) => listeners.push(listener),
Err(error) => {
let _ = ready.send(Err(error));
return;
}
}
}
let handle = server.launch(handler, listeners, gate);
if ready.send(Ok(())).is_err() {
return;
}
drop(ready);
let _ = receiver.changed().await;
let timeout = *receiver.borrow_and_update();
handle.close(timeout).await;
});
};
match std::thread::Builder::new().name(format!("soyokaze-{index}")).stack_size(self.config.limits.worker_stack_size).spawn(worker) {
Ok(thread) => threads.push(thread),
Err(error) => {
failure = Some(Error::IO(error));
break;
}
}
}
drop(ready);
for _ in 0..threads.len() {
match started.recv() {
Ok(Ok(())) => continue,
Ok(Err(error)) => failure = failure.or(Some(error)),
Err(_) => failure = failure.or(Some(Error::Closed)),
}
}
if let Some(error) = failure {
let _ = shutdown.send(None);
for thread in threads {
let _ = thread.join();
}
return Err(error);
}
Ok(Cluster::new(shutdown, threads, addresses))
}
}