use std::{
cell::RefCell,
marker::PhantomData,
sync::{Arc, atomic::AtomicI32},
};
use motore::{
layer::{Identity, Layer, Stack},
service::{BoxCloneService, Service},
};
use pilota::thrift::TMessageType;
use tokio::time::Duration;
#[cfg(feature = "shmipc")]
use volo::net::shmipc_fallback::ShmipcMakeTransportWithFallback;
use volo::{
FastStr,
client::WithOptService,
context::{Context, Endpoint, Role, RpcInfo},
discovery::{Discover, DummyDiscover},
loadbalance::{LbConfig, MkLbLayer, random::WeightedRandomBalance},
net::{
Address,
dial::{DefaultMakeTransport, MakeTransport},
},
};
use crate::{
ClientError, EntryMessage, ThriftMessage,
codec::{
DefaultMakeCodec, MakeCodec,
default::{framed::MakeFramedCodec, thrift::MakeThriftCodec, ttheader::MakeTTHeaderCodec},
},
context::{CLIENT_CONTEXT_CACHE, ClientContext, Config},
transport::{pingpong, pool},
};
mod callopt;
pub use callopt::CallOpt;
use self::layer::timeout::TimeoutLayer;
pub mod layer;
pub struct ClientBuilder<IL, OL, MkClient, Req, Resp, MkT, MkC, LB> {
config: Config,
pool: Option<pool::Config>,
callee_name: FastStr,
caller_name: FastStr,
address: Option<Address>, inner_layer: IL,
outer_layer: OL,
make_transport: MkT,
make_codec: MkC,
mk_client: MkClient,
mk_lb: LB,
_marker: PhantomData<(*const Req, *const Resp)>,
disable_timeout_layer: bool,
enable_biz_error: bool,
#[cfg(feature = "multiplex")]
multiplex: bool,
}
impl<C, Req, Resp>
ClientBuilder<
Identity,
Identity,
C,
Req,
Resp,
DefaultMakeTransport,
DefaultMakeCodec<MakeTTHeaderCodec<MakeFramedCodec<MakeThriftCodec>>>,
LbConfig<WeightedRandomBalance<<DummyDiscover as Discover>::Key>, DummyDiscover>,
>
{
pub fn new(service_name: impl AsRef<str>, service_client: C) -> Self {
ClientBuilder {
config: Default::default(),
pool: None,
caller_name: "".into(),
callee_name: FastStr::new(service_name),
address: None,
inner_layer: Identity::new(),
outer_layer: Identity::new(),
mk_client: service_client,
make_transport: DefaultMakeTransport::default(),
make_codec: DefaultMakeCodec::default(),
mk_lb: LbConfig::new(WeightedRandomBalance::new(), DummyDiscover {}),
_marker: PhantomData,
disable_timeout_layer: false,
enable_biz_error: true,
#[cfg(feature = "multiplex")]
multiplex: false,
}
}
}
impl<IL, OL, C, Req, Resp, MkT, MkC, LB, DISC>
ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, LbConfig<LB, DISC>>
{
pub fn load_balance<NLB>(
self,
load_balance: NLB,
) -> ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, LbConfig<NLB, DISC>> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb.load_balance(load_balance),
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
pub fn discover<NDISC>(
self,
discover: NDISC,
) -> ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, LbConfig<LB, NDISC>> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb.discover(discover),
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
pub fn retry_count(mut self, count: usize) -> Self {
self.mk_lb = self.mk_lb.retry_count(count);
self
}
}
impl<IL, OL, C, Req, Resp, MkT, MkC, LB> ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, LB> {
pub fn rpc_timeout(mut self, timeout: Option<Duration>) -> Self {
self.config.set_rpc_timeout(timeout);
self
}
pub fn pool_config(mut self, config: pool::Config) -> Self {
self.pool = Some(config);
self
}
pub fn connect_timeout(mut self, timeout: Option<Duration>) -> Self {
self.config.set_connect_timeout(timeout);
self
}
pub fn read_write_timeout(mut self, timeout: Option<Duration>) -> Self {
self.config.set_read_write_timeout(timeout);
self
}
pub fn caller_name(mut self, name: impl AsRef<str>) -> Self {
self.caller_name = FastStr::new(name);
self
}
#[doc(hidden)]
pub fn disable_timeout_layer(mut self) -> Self {
self.disable_timeout_layer = true;
self
}
#[doc(hidden)]
pub fn enable_biz_error(mut self, enable_biz_error: bool) -> Self {
self.enable_biz_error = enable_biz_error;
self
}
pub fn mk_load_balance<NLB>(
self,
mk_load_balance: NLB,
) -> ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, NLB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: mk_load_balance,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
#[doc(hidden)]
pub fn make_codec<MakeCodec>(
self,
make_codec: MakeCodec,
) -> ClientBuilder<IL, OL, C, Req, Resp, MkT, MakeCodec, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
#[doc(hidden)]
pub fn make_transport<MakeTransport>(
self,
make_transport: MakeTransport,
) -> ClientBuilder<IL, OL, C, Req, Resp, MakeTransport, MkC, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
pub fn address<A: Into<Address>>(mut self, target: A) -> Self {
self.address = Some(target.into());
self
}
pub fn layer_inner<Inner>(
self,
layer: Inner,
) -> ClientBuilder<Stack<Inner, IL>, OL, C, Req, Resp, MkT, MkC, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: Stack::new(layer, self.inner_layer),
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
pub fn layer_inner_front<Inner>(
self,
layer: Inner,
) -> ClientBuilder<Stack<IL, Inner>, OL, C, Req, Resp, MkT, MkC, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: Stack::new(self.inner_layer, layer),
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
pub fn layer_outer<Outer>(
self,
layer: Outer,
) -> ClientBuilder<IL, Stack<Outer, OL>, C, Req, Resp, MkT, MkC, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: Stack::new(layer, self.outer_layer),
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
pub fn layer_outer_front<Outer>(
self,
layer: Outer,
) -> ClientBuilder<IL, Stack<OL, Outer>, C, Req, Resp, MkT, MkC, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: Stack::new(self.outer_layer, layer),
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
#[cfg(feature = "multiplex")]
#[doc(hidden)]
pub fn multiplex(self, multiplex: bool) -> ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, LB> {
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: self.address,
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: self.make_transport,
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
multiplex,
}
}
#[cfg(feature = "shmipc")]
pub fn shmipc_fallback_address<A: Into<Address>>(
self,
fallback_addr: A,
) -> ClientBuilder<IL, OL, C, Req, Resp, ShmipcMakeTransportWithFallback, MkC, LB> {
let shmipc_addr = self
.address
.expect("Must call .address() before .with_fallback_address()");
ClientBuilder {
config: self.config,
pool: self.pool,
caller_name: self.caller_name,
callee_name: self.callee_name,
address: Some(shmipc_addr),
inner_layer: self.inner_layer,
outer_layer: self.outer_layer,
mk_client: self.mk_client,
_marker: PhantomData,
make_transport: ShmipcMakeTransportWithFallback::new(
DefaultMakeTransport::default(),
DefaultMakeTransport::default(),
fallback_addr.into(),
),
make_codec: self.make_codec,
mk_lb: self.mk_lb,
disable_timeout_layer: self.disable_timeout_layer,
enable_biz_error: self.enable_biz_error,
#[cfg(feature = "multiplex")]
multiplex: self.multiplex,
}
}
#[doc(hidden)]
pub fn get_callee_name(&self) -> &FastStr {
&self.callee_name
}
}
#[derive(Clone)]
pub struct MessageService<Resp, MkT, MkC>
where
Resp: EntryMessage + Send + 'static,
MkT: MakeTransport,
MkC: MakeCodec<MkT::ReadHalf, MkT::WriteHalf> + Sync,
{
#[cfg(not(feature = "multiplex"))]
inner: pingpong::Client<Resp, MkT, MkC>,
#[cfg(feature = "multiplex")]
inner: motore::utils::Either<
pingpong::Client<Resp, MkT, MkC>,
crate::transport::multiplex::Client<Resp, MkT, MkC>,
>,
read_biz_error: bool,
}
impl<Req, Resp, MkT, MkC> Service<ClientContext, Req> for MessageService<Resp, MkT, MkC>
where
Req: EntryMessage + 'static + Send,
Resp: Send + 'static + EntryMessage + Sync,
MkT: MakeTransport,
MkC: MakeCodec<MkT::ReadHalf, MkT::WriteHalf> + Sync,
{
type Response = Option<Resp>;
type Error = ClientError;
async fn call(&self, cx: &mut ClientContext, req: Req) -> Result<Self::Response, Self::Error> {
let msg = ThriftMessage::mk_client_msg(cx, req);
let resp = self.inner.call(cx, msg).await;
if self.read_biz_error {
if let Some(biz_err) = cx.common_stats.biz_error() {
return Err(biz_err.clone().into());
}
}
match resp {
Ok(Some(ThriftMessage { data: Ok(data), .. })) => Ok(Some(data)),
Ok(Some(ThriftMessage { data: Err(e), .. })) => Err(e.into()),
Err(e) => Err(e),
Ok(None) => Ok(None),
}
}
}
impl<IL, OL, C, Req, Resp, MkT, MkC, LB> ClientBuilder<IL, OL, C, Req, Resp, MkT, MkC, LB>
where
C: volo::client::MkClient<
Client<
BoxCloneService<
ClientContext,
Req,
Option<Resp>,
<OL::Service as Service<ClientContext, Req>>::Error,
>,
>,
>,
LB: MkLbLayer,
LB::Layer: Layer<IL::Service>,
<LB::Layer as Layer<IL::Service>>::Service: Service<ClientContext, Req, Response = Option<Resp>, Error = ClientError>
+ 'static
+ Send
+ Clone
+ Sync,
Req: EntryMessage + Send + 'static + Sync + Clone,
Resp: EntryMessage + Send + 'static,
IL: Layer<MessageService<Resp, MkT, MkC>>,
IL::Service:
Service<ClientContext, Req, Response = Option<Resp>> + Sync + Clone + Send + 'static,
<IL::Service as Service<ClientContext, Req>>::Error: Send + Into<ClientError>,
MkT: MakeTransport,
MkC: MakeCodec<MkT::ReadHalf, MkT::WriteHalf> + Sync,
OL: Layer<BoxCloneService<ClientContext, Req, Option<Resp>, ClientError>>,
OL::Service:
Service<ClientContext, Req, Response = Option<Resp>> + 'static + Send + Clone + Sync,
<OL::Service as Service<ClientContext, Req>>::Error: Send + Sync + Into<ClientError>,
{
pub fn build(mut self) -> C::Target {
if let Some(timeout) = self.config.connect_timeout() {
self.make_transport.set_connect_timeout(Some(timeout));
}
if let Some(timeout) = self.config.read_write_timeout() {
self.make_transport.set_read_timeout(Some(timeout));
}
if let Some(timeout) = self.config.read_write_timeout() {
self.make_transport.set_write_timeout(Some(timeout));
}
let msg_svc = MessageService {
#[cfg(not(feature = "multiplex"))]
inner: pingpong::Client::new(self.make_transport, self.pool, self.make_codec),
#[cfg(feature = "multiplex")]
inner: if !self.multiplex {
motore::utils::Either::A(pingpong::Client::new(
self.make_transport,
self.pool,
self.make_codec,
))
} else {
motore::utils::Either::B(crate::transport::multiplex::Client::new(
self.make_transport,
self.pool,
self.make_codec,
))
},
read_biz_error: self.enable_biz_error,
};
let transport = if !self.disable_timeout_layer {
BoxCloneService::new(self.outer_layer.layer(BoxCloneService::new(
TimeoutLayer::new().layer(self.mk_lb.make().layer(self.inner_layer.layer(msg_svc))),
)))
} else {
BoxCloneService::new(self.outer_layer.layer(BoxCloneService::new(
self.mk_lb.make().layer(self.inner_layer.layer(msg_svc)),
)))
};
self.mk_client.mk_client(Client {
inner: Arc::new(ClientInner {
callee_name: self.callee_name,
config: self.config,
address: self.address,
caller_name: self.caller_name,
seq_id: AtomicI32::new(0),
}),
transport,
})
}
}
#[derive(Clone)]
pub struct Client<S> {
transport: S,
inner: Arc<ClientInner>,
}
struct ClientInner {
callee_name: FastStr,
caller_name: FastStr,
config: Config,
address: Option<Address>,
seq_id: AtomicI32,
}
impl<S> Client<S> {
pub fn make_cx(&self, method: &str, oneway: bool) -> ClientContext {
CLIENT_CONTEXT_CACHE.with(|cache| {
let mut cache = cache.borrow_mut();
cache
.pop()
.map(|mut cx| {
cx.reset(
self.inner
.seq_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
if oneway {
TMessageType::OneWay
} else {
TMessageType::Call
},
);
cx.rpc_info_mut()
.caller_mut()
.set_service_name(self.inner.caller_name.clone());
cx.rpc_info_mut()
.callee_mut()
.set_service_name(self.inner.callee_name.clone());
if let Some(target) = &self.inner.address {
cx.rpc_info_mut().callee_mut().set_address(target.clone());
}
cx.rpc_info_mut().set_config(self.inner.config);
cx.rpc_info_mut().set_method(FastStr::new(method));
cx
})
.unwrap_or_else(|| {
ClientContext::new(
self.inner
.seq_id
.fetch_add(1, std::sync::atomic::Ordering::Relaxed),
self.make_rpc_info(method),
if oneway {
TMessageType::OneWay
} else {
TMessageType::Call
},
)
})
})
}
fn make_rpc_info(&self, method: &str) -> RpcInfo<Config> {
let caller = Endpoint::new(self.inner.caller_name.clone());
let mut callee = Endpoint::new(self.inner.callee_name.clone());
if let Some(target) = &self.inner.address {
callee.set_address(target.clone());
}
let config = self.inner.config;
RpcInfo::new(Role::Client, FastStr::new(method), caller, callee, config)
}
pub fn with_opt<Opt>(self, opt: Opt) -> Client<WithOptService<S, Opt>> {
Client {
transport: WithOptService::new(self.transport, opt),
inner: self.inner,
}
}
}
macro_rules! impl_client {
(($self: ident, &mut $cx:ident, $req: ident) => async move $e: tt ) => {
impl<S, Req: Send + 'static, Res: 'static>
volo::service::Service<crate::context::ClientContext, Req> for Client<S>
where
S: volo::service::Service<
crate::context::ClientContext,
Req,
Response = Option<Res>,
Error = crate::ClientError,
> + Sync
+ Send
+ 'static,
{
type Response = S::Response;
type Error = S::Error;
async fn call(
&$self,
$cx: &mut crate::context::ClientContext,
$req: Req,
) -> Result<Self::Response, Self::Error> {
$e
}
}
impl<S, Req: Send + 'static, Res: 'static>
volo::client::OneShotService<crate::context::ClientContext, Req> for Client<S>
where
S: volo::client::OneShotService<
crate::context::ClientContext,
Req,
Response = Option<Res>,
Error = crate::ClientError,
> + Sync
+ Send
+ 'static,
{
type Response = S::Response;
type Error = S::Error;
async fn call(
$self,
$cx: &mut crate::context::ClientContext,
$req: Req,
) -> Result<Self::Response, Self::Error> {
$e
}
}
};
}
impl_client!((self, &mut cx, req) => async move {
let has_metainfo = metainfo::METAINFO.try_with(|_| {}).is_ok();
let mk_call = async { self.transport.call(cx, req).await };
if has_metainfo {
mk_call.await
} else {
metainfo::METAINFO
.scope(RefCell::new(metainfo::MetaInfo::default()), mk_call)
.await
}
});