use std::{
marker::PhantomData,
sync::{Arc, atomic::Ordering},
time::Duration,
};
use futures::future::BoxFuture;
use motore::{
BoxError,
layer::{Identity, Layer, Stack},
service::Service,
};
use scopeguard::defer;
use tokio::{
io::{AsyncRead, AsyncWrite},
sync::Notify,
};
use tracing::{info, trace};
use volo::{
net::{
Address,
conn::{OwnedReadHalf, OwnedWriteHalf},
incoming::Incoming,
},
service::BoxService,
};
use crate::{
EntryMessage,
codec::{
DefaultMakeCodec, MakeCodec,
default::{framed::MakeFramedCodec, thrift::MakeThriftCodec, ttheader::MakeTTHeaderCodec},
},
context::ServerContext,
server::layer::biz_error::BizErrorLayer,
tracing::{DefaultProvider, SpanProvider},
};
mod layer;
pub mod panic_handler;
#[doc(hidden)]
pub type TraceFn = fn(&ServerContext);
pub struct Server<S, L, Req, MkC, SP> {
service: S,
layer: L,
make_codec: MkC,
stat_tracer: Vec<TraceFn>,
#[cfg(feature = "multiplex")]
multiplex: bool,
span_provider: SP,
shutdown_hooks: Vec<Box<dyn FnOnce() -> BoxFuture<'static, ()> + Send>>,
_marker: PhantomData<Req>,
}
impl<S, Req>
Server<
S,
Identity,
Req,
DefaultMakeCodec<MakeTTHeaderCodec<MakeFramedCodec<MakeThriftCodec>>>,
DefaultProvider,
>
{
pub fn new(service: S) -> Self
where
S: Service<ServerContext, Req>,
{
Self {
make_codec: DefaultMakeCodec::default(),
service,
layer: Identity::new(),
stat_tracer: Vec::new(),
#[cfg(feature = "multiplex")]
multiplex: false,
span_provider: DefaultProvider {},
shutdown_hooks: Vec::new(),
_marker: PhantomData,
}
}
}
impl<S, L, Req, MkC, SP> Server<S, L, Req, MkC, SP> {
pub fn register_shutdown_hook(
mut self,
hook: impl FnOnce() -> BoxFuture<'static, ()> + 'static + Send,
) -> Self {
self.shutdown_hooks.push(Box::new(hook));
self
}
pub fn layer<Inner>(self, layer: Inner) -> Server<S, Stack<Inner, L>, Req, MkC, SP> {
Server {
layer: Stack::new(layer, self.layer),
service: self.service,
make_codec: self.make_codec,
stat_tracer: self.stat_tracer,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
span_provider: self.span_provider,
shutdown_hooks: self.shutdown_hooks,
_marker: PhantomData,
}
}
pub fn layer_front<Front>(self, layer: Front) -> Server<S, Stack<L, Front>, Req, MkC, SP> {
Server {
layer: Stack::new(self.layer, layer),
service: self.service,
make_codec: self.make_codec,
stat_tracer: self.stat_tracer,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
span_provider: self.span_provider,
shutdown_hooks: self.shutdown_hooks,
_marker: PhantomData,
}
}
#[doc(hidden)]
pub fn stat_tracer(mut self, trace_fn: TraceFn) -> Self {
self.stat_tracer.push(trace_fn);
self
}
#[doc(hidden)]
pub fn make_codec<MakeCodec>(self, make_codec: MakeCodec) -> Server<S, L, Req, MakeCodec, SP> {
Server {
layer: self.layer,
service: self.service,
make_codec,
stat_tracer: self.stat_tracer,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
span_provider: self.span_provider,
shutdown_hooks: self.shutdown_hooks,
_marker: PhantomData,
}
}
pub async fn run<MI: volo::net::incoming::MakeIncoming>(
self,
make_incoming: MI,
) -> Result<(), BoxError>
where
L: Layer<BoxService<ServerContext, Req, S::Response, crate::ServerError>>,
MkC: MakeCodec<OwnedReadHalf, OwnedWriteHalf>,
L::Service: Service<ServerContext, Req, Response = S::Response, Error = crate::ServerError>
+ Send
+ 'static
+ Sync,
S: Service<ServerContext, Req, Error = crate::ServerError> + Send + 'static + Sync,
S::Response: EntryMessage + Send + 'static + Sync,
Req: EntryMessage + Send + 'static,
SP: SpanProvider,
{
let service = Arc::new(
self.layer
.layer(BoxService::new(BizErrorLayer::new().layer(self.service))),
);
let stat_tracer: Arc<[TraceFn]> = Arc::from(self.stat_tracer);
let mut incoming = make_incoming.make_incoming().await?;
info!("[VOLO] server start at: {:?}", incoming);
let conn_cnt = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let gconn_cnt = conn_cnt.clone();
let (exit_notify, exit_flag, exit_mark) = (
Arc::new(Notify::const_new()),
Arc::new(parking_lot::RwLock::new(false)),
Arc::new(std::sync::atomic::AtomicBool::default()),
);
let (exit_notify_inner, exit_flag_inner, exit_mark_inner) =
(exit_notify.clone(), exit_flag.clone(), exit_mark.clone());
let handler = tokio::spawn(async move {
let exit_flag = exit_flag_inner.clone();
loop {
if *exit_flag.read() {
break Ok(());
}
match incoming.accept().await {
Ok(Some(conn)) => {
let peer_addr = conn.info.peer_addr;
trace!("[VOLO] accept connection from: {:?}", peer_addr);
let (rh, wh) = conn.stream.into_split();
#[cfg(feature = "multiplex")]
if self.multiplex {
#[cfg(feature = "shmipc")]
if peer_addr.as_ref().is_some_and(Address::is_shmipc) {
tracing::error!("multiplex is not supported when using shmipc");
let _ = rh.shmipc_helper().close().await;
continue;
}
tokio::spawn(handle_conn_multiplex(
rh,
wh,
service.clone(),
self.make_codec.clone(),
stat_tracer.clone(),
exit_notify_inner.clone(),
exit_mark_inner.clone(),
conn_cnt.clone(),
peer_addr,
));
} else {
tokio::spawn(handle_conn(
rh,
wh,
service.clone(),
self.make_codec.clone(),
stat_tracer.clone(),
exit_notify_inner.clone(),
exit_mark_inner.clone(),
conn_cnt.clone(),
peer_addr,
self.span_provider.clone(),
));
}
#[cfg(not(feature = "multiplex"))]
tokio::spawn(handle_conn(
rh,
wh,
service.clone(),
self.make_codec.clone(),
stat_tracer.clone(),
exit_notify_inner.clone(),
exit_mark_inner.clone(),
conn_cnt.clone(),
peer_addr,
self.span_provider.clone(),
));
}
Ok(None) => break Ok(()),
Err(e) => break Err(e),
}
}
});
#[cfg(target_family = "unix")]
{
let mut sigint =
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::interrupt())?;
let mut sighup =
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::hangup())?;
let mut sigterm =
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())?;
tokio::select! {
_ = sigint.recv() => {}
_ = sighup.recv() => {}
_ = sigterm.recv() => {}
res = handler => {
match res {
Ok(res) => {
match res {
Ok(()) => {}
Err(e) => return Err(Box::new(e))
};
}
Err(e) => return Err(Box::new(e)),
}
}
}
}
#[cfg(target_family = "windows")]
tokio::select! {
_ = tokio::signal::ctrl_c() => {}
res = handler => {
match res {
Ok(res) => {
match res {
Ok(()) => {}
Err(e) => return Err(Box::new(e))
};
}
Err(e) => return Err(Box::new(e)),
}
}
}
if !self.shutdown_hooks.is_empty() {
info!("[VOLO] call shutdown hooks");
for hook in self.shutdown_hooks {
(hook)().await;
}
}
info!("[VOLO] received signal, gracefully exiting now");
*exit_flag.write() = true;
exit_mark.store(true, Ordering::Relaxed);
if gconn_cnt.load(Ordering::Relaxed) != 0 {
tokio::time::sleep(Duration::from_secs(2)).await;
}
exit_notify.notify_waiters();
for _ in 0..28 {
if gconn_cnt.load(Ordering::Relaxed) == 0 {
break;
}
trace!(
"[VOLO] gracefully exiting, remaining connection count: {}",
gconn_cnt.load(Ordering::Relaxed)
);
tokio::time::sleep(Duration::from_secs(1)).await;
}
Ok(())
}
#[cfg(feature = "multiplex")]
#[doc(hidden)]
pub fn multiplex(self, multiplex: bool) -> Server<S, L, Req, MkC, SP> {
Server {
layer: self.layer,
service: self.service,
make_codec: self.make_codec,
stat_tracer: self.stat_tracer,
multiplex,
span_provider: self.span_provider,
shutdown_hooks: self.shutdown_hooks,
_marker: PhantomData,
}
}
pub fn span_provider<P: SpanProvider>(self, provider: P) -> Server<S, L, Req, MkC, P> {
Server {
layer: self.layer,
service: self.service,
make_codec: self.make_codec,
stat_tracer: self.stat_tracer,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
span_provider: provider,
shutdown_hooks: self.shutdown_hooks,
_marker: PhantomData,
}
}
}
#[allow(clippy::too_many_arguments)]
async fn handle_conn<R, W, Req, Svc, Resp, MkC, SP>(
rh: R,
wh: W,
service: Svc,
make_codec: MkC,
stat_tracer: Arc<[TraceFn]>,
exit_notify: Arc<Notify>,
exit_mark: Arc<std::sync::atomic::AtomicBool>,
conn_cnt: Arc<std::sync::atomic::AtomicUsize>,
peer_addr: Option<Address>,
span_provider: SP,
) where
R: AsyncRead + Unpin + Send + Sync + 'static,
W: AsyncWrite + Unpin + Send + Sync + 'static,
Svc: Service<ServerContext, Req, Response = Resp> + Clone + Send + 'static,
Svc::Error: Send,
Svc::Error: Into<crate::ServerError>,
Req: EntryMessage + Send + 'static,
Resp: EntryMessage + Send + 'static,
MkC: MakeCodec<R, W>,
SP: SpanProvider,
{
conn_cnt.fetch_add(1, Ordering::Relaxed);
defer! {
conn_cnt.fetch_sub(1, Ordering::Relaxed);
}
let (encoder, decoder) = make_codec.make_codec(rh, wh);
#[cfg(feature = "shmipc")]
let _guard = crate::codec::Decoder::shmipc_helper(&decoder).close_guard();
tracing::trace!(
"[VOLO] handle conn by ping-pong, peer_addr: {:?}",
peer_addr
);
crate::transport::pingpong::serve(
encoder,
decoder,
exit_notify.notified(),
exit_mark,
&service,
stat_tracer,
peer_addr,
span_provider,
)
.await;
}
#[cfg(feature = "multiplex")]
#[allow(clippy::too_many_arguments)]
async fn handle_conn_multiplex<R, W, Req, Svc, Resp, MkC>(
rh: R,
wh: W,
service: Svc,
make_codec: MkC,
stat_tracer: Arc<[TraceFn]>,
exit_notify: Arc<Notify>,
exit_mark: Arc<std::sync::atomic::AtomicBool>,
conn_cnt: Arc<std::sync::atomic::AtomicUsize>,
peer_addr: Option<Address>,
) where
R: AsyncRead + Unpin + Send + Sync + 'static,
W: AsyncWrite + Unpin + Send + Sync + 'static,
Svc: Service<ServerContext, Req, Response = Resp> + Clone + Send + 'static + Sync,
Svc::Error: Into<crate::ServerError> + Send,
Req: EntryMessage + Send + 'static,
Resp: EntryMessage + Send + 'static,
MkC: MakeCodec<R, W>,
{
conn_cnt.fetch_add(1, Ordering::Relaxed);
defer! {
conn_cnt.fetch_sub(1, Ordering::Relaxed);
}
let (encoder, decoder) = make_codec.make_codec(rh, wh);
info!(
"[VOLO] handle conn by multiplex, peer_addr: {:?}",
peer_addr
);
crate::transport::multiplex::serve(
encoder,
decoder,
exit_notify.notified(),
exit_mark,
service,
stat_tracer,
peer_addr,
)
.await;
}