tina-core 0.0.2

Tina platform
Documentation
//! service 扩展

use std::{
    convert::Infallible,
    fmt::Debug,
    sync::{Arc, Mutex, MutexGuard},
    task::Poll,
};

use futures::{future::BoxFuture, FutureExt};
use http::{Request, Response, Version};
use http_body::Body;
use once_cell::sync::Lazy;
use prost::Message;
use tonic::{
    body::BoxBody,
    transport::{server::Router, NamedService, Server},
    Code, Status,
};
use tower::Service;

use crate::tina::{
    grpc::{ApiGrpcService, IntoGrpcResponse},
    server::grpc::service::GrpcServiceServer,
};

use super::request_ext::RequestExt;

static GLOBAL_GRPC_SERVICE_REGISTER: Lazy<Arc<Mutex<GrpcServiceRegister>>> = Lazy::new(|| {
    let server = Server::builder();
    Arc::new(Mutex::new(GrpcServiceRegister {
        inner: Some(Inner::Server(server)),
        service_names: vec![],
    }))
});

/// gRPC服务注册
pub struct GrpcServiceRegister {
    inner: Option<Inner>,
    service_names: Vec<&'static str>,
}

enum Inner {
    Server(Server),
    Router(Router),
}

impl GrpcServiceRegister {
    /// 添加服务
    pub fn add_service<S, M>(service: S)
    where
        S: ApiGrpcService,
        M: prost::Message + Default + Debug + Send + Sync + 'static,
        S::Rejection: IntoGrpcResponse<Response = tonic::Response<M>>,
    {
        let mut register: MutexGuard<GrpcServiceRegister> =
            GLOBAL_GRPC_SERVICE_REGISTER.lock().expect("lock GLOBAL_GRPC_SERVICE_REGISTER failed");
        match register.inner.take() {
            Some(inner) => match inner {
                Inner::Server(mut server) => {
                    let router = server.add_service(GrpcServiceServer::new(service));
                    register.inner = Some(Inner::Router(router));
                }
                Inner::Router(router) => {
                    let router = router.add_service(GrpcServiceServer::new(service));
                    register.inner = Some(Inner::Router(router));
                }
            },
            None => {
                let mut server = Server::builder();
                register.inner = Some(Inner::Router(server.add_service(GrpcServiceServer::new(service))));
            }
        }
        register.service_names.push(S::NAME);
    }
    /// 添加tonic服务
    pub fn add_tonic_service<S>(svc: S)
    where
        S: Service<Request<hyper::Body>, Response = Response<BoxBody>, Error = Infallible> + NamedService + Clone + Send + 'static,
        S::Future: Send + 'static,
    {
        let mut register: MutexGuard<GrpcServiceRegister> =
            GLOBAL_GRPC_SERVICE_REGISTER.lock().expect("lock GLOBAL_GRPC_SERVICE_REGISTER failed");
        match register.inner.take() {
            Some(inner) => match inner {
                Inner::Server(mut server) => {
                    let router = server.add_service(svc);
                    register.inner = Some(Inner::Router(router));
                }
                Inner::Router(router) => {
                    let router = router.add_service(svc);
                    register.inner = Some(Inner::Router(router));
                }
            },
            None => {
                let mut server = Server::builder();
                register.inner = Some(Inner::Router(server.add_service(svc)));
            }
        }
        register.service_names.push(S::NAME);
    }
    /// 获取Router
    pub fn get_router() -> Router {
        let mut register: MutexGuard<GrpcServiceRegister> =
            GLOBAL_GRPC_SERVICE_REGISTER.lock().expect("lock GLOBAL_GRPC_SERVICE_REGISTER failed");
        match register.inner.take() {
            Some(inner) => match inner {
                Inner::Server(mut server) => server.add_optional_service(None as Option<NoService>),
                Inner::Router(router) => router,
            },
            None => {
                let mut server = Server::builder();
                server.add_optional_service(None as Option<NoService>)
            }
        }
    }
    /// 获取注册的服务名列表
    pub fn get_registed_service_names() -> Vec<&'static str> {
        let register: MutexGuard<GrpcServiceRegister> =
            GLOBAL_GRPC_SERVICE_REGISTER.lock().expect("lock GLOBAL_GRPC_SERVICE_REGISTER failed");
        register.service_names.clone()
    }
}

impl<S, M> Service<Request<hyper::Body>> for GrpcServiceServer<S>
where
    S: ApiGrpcService,
    M: Message + Debug + Default + Send + Sync + 'static,
    S::Rejection: IntoGrpcResponse<Response = tonic::Response<M>>,
{
    type Response = Response<BoxBody>;

    type Error = Infallible;

    type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;

    fn poll_ready(&mut self, _cx: &mut std::task::Context<'_>) -> std::task::Poll<Result<(), Self::Error>> {
        std::task::Poll::Ready(Ok(()))
    }

    fn call(&mut self, req: Request<hyper::Body>) -> Self::Future {
        let req: crate::tina::grpc::request::Request = crate::tina::grpc::request::Request {
            inner: req,
        };
        let service = self.service.clone();
        let locale = req.get_locale();
        async move {
            let res = service.call(req).await;
            match res {
                Ok(res) => {
                    let (metadata, body, extensions) = res.inner.into_parts();
                    let body = body
                        .map_err(move |err| {
                            let message = err.get_return_msg().get_string(locale.as_str());
                            let code = tonic::Code::from(&err);
                            Status::new(code, message)
                        })
                        .boxed_unsync();
                    let headers = metadata.into_headers();
                    let extensions = extensions.into_http();
                    let mut res = http::Response::new(body);
                    *res.version_mut() = Version::HTTP_2;
                    *res.headers_mut() = headers;
                    *res.extensions_mut() = extensions;
                    tracing::trace!("GrpcServiceServer send ok response: {res:?}");
                    Ok(res)
                }
                Err(err) => {
                    let (metadata, body, extensions) = err.into_grpc_response().await.into_parts();
                    let res = crate::tina::grpc::response::Response::from_parts(metadata, body, extensions);
                    let (metadata, body, extensions) = res.inner.into_parts();
                    let body = body
                        .map_err(move |err1| {
                            let message = err1.get_return_msg().get_string(&locale);
                            Status::new(Code::from(&err1), message)
                        })
                        .boxed_unsync();
                    let headers = metadata.into_headers();
                    let extensions = extensions.into_http();
                    let mut res = http::Response::new(body);
                    *res.version_mut() = Version::HTTP_2;
                    *res.headers_mut() = headers;
                    *res.extensions_mut() = extensions;
                    tracing::trace!("GrpcServiceServer send err response: {res:?}");
                    Ok(res)
                }
            }
        }
        .boxed()
    }
}

impl<S> NamedService for GrpcServiceServer<S>
where
    S: ApiGrpcService,
{
    const NAME: &'static str = S::NAME;
}

#[derive(Clone)]
pub(crate) struct NoService;

impl Service<Request<hyper::Body>> for NoService {
    type Response = Response<BoxBody>;

    type Error = Infallible;

    type Future = BoxFuture<'static, Result<Self::Response, Self::Error>>;

    fn poll_ready(&mut self, _cx: &mut std::task::Context<'_>) -> std::task::Poll<Result<(), Self::Error>> {
        Poll::Ready(Ok(()))
    }

    fn call(&mut self, _req: Request<hyper::Body>) -> Self::Future {
        let res = Response::new(BoxBody::default());
        async move { Ok(res) }.boxed()
    }
}

impl NamedService for NoService {
    const NAME: &'static str = "NoService";
}