use ahash::AHashMap;
use motore::{BoxCloneService, service::Service};
use pilota::thrift::{ApplicationException, ApplicationExceptionKind};
use volo::FastStr;
use crate::{
Bytes, ServerError,
context::{ServerContext, ThriftContext},
};
pub trait NamedService {
const NAME: &'static str;
}
type BoxedService = BoxCloneService<ServerContext, Bytes, Bytes, ServerError>;
pub struct Router {
services: AHashMap<FastStr, BoxedService>,
default_service: Option<BoxedService>,
}
impl Default for Router {
#[inline]
fn default() -> Self {
Self::new()
}
}
impl Clone for Router {
#[inline]
fn clone(&self) -> Self {
Self {
services: self.services.clone(),
default_service: self.default_service.clone(),
}
}
}
impl Router {
pub fn new() -> Self {
Self {
services: AHashMap::new(),
default_service: None,
}
}
pub fn with_default_service<S>(mut self, service: S) -> Self
where
S: Service<ServerContext, Bytes, Response = Bytes, Error = ServerError>
+ NamedService
+ Clone
+ Send
+ Sync
+ 'static,
{
let name = FastStr::from_static_str(S::NAME);
let boxed = BoxCloneService::new(service);
self.default_service = Some(boxed.clone());
self.services.insert(name, boxed);
self
}
pub fn add_service<S>(mut self, service: S) -> Self
where
S: Service<ServerContext, Bytes, Response = Bytes, Error = ServerError>
+ NamedService
+ Clone
+ Send
+ Sync
+ 'static,
{
let name = FastStr::from_static_str(S::NAME);
self.services.insert(name, BoxCloneService::new(service));
self
}
pub fn service_count(&self) -> usize {
self.services.len()
}
pub fn has_default_service(&self) -> bool {
self.default_service.is_some()
}
}
impl Service<ServerContext, Bytes> for Router {
type Response = Bytes;
type Error = ServerError;
#[inline]
async fn call(
&self,
cx: &mut ServerContext,
payload: Bytes,
) -> Result<Self::Response, Self::Error> {
let service_name = cx.idl_service_name();
let service = match service_name {
Some(name) => self.services.get(name).or(self.default_service.as_ref()),
None => self.default_service.as_ref(),
};
match service {
Some(svc) => svc.call(cx, payload).await,
None => Err(ServerError::Application(ApplicationException::new(
ApplicationExceptionKind::UNKNOWN_METHOD,
format!(
"service not found: {:?}",
service_name.map(|s: &FastStr| s.as_str())
),
))),
}
}
}
impl std::fmt::Debug for Router {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Router")
.field("services", &self.services.keys().collect::<Vec<_>>())
.field("has_default_service", &self.default_service.is_some())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Clone)]
struct MockService {
name: &'static str,
}
impl NamedService for MockService {
const NAME: &'static str = "MockService";
}
impl Service<ServerContext, Bytes> for MockService {
type Response = Bytes;
type Error = ServerError;
async fn call(
&self,
_cx: &mut ServerContext,
_payload: Bytes,
) -> Result<Self::Response, Self::Error> {
Ok(Bytes::from(self.name))
}
}
#[derive(Clone)]
struct AnotherMockService;
impl NamedService for AnotherMockService {
const NAME: &'static str = "AnotherService";
}
impl Service<ServerContext, Bytes> for AnotherMockService {
type Response = Bytes;
type Error = ServerError;
async fn call(
&self,
_cx: &mut ServerContext,
_payload: Bytes,
) -> Result<Self::Response, Self::Error> {
Ok(Bytes::from("another"))
}
}
#[test]
fn test_router_new() {
let router = Router::new();
assert_eq!(router.service_count(), 0);
assert!(!router.has_default_service());
}
#[test]
fn test_router_with_default_service() {
let router = Router::new().with_default_service(MockService { name: "default" });
assert_eq!(router.service_count(), 1);
assert!(router.has_default_service());
}
#[test]
fn test_router_add_service() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
assert_eq!(router.service_count(), 2);
assert!(router.has_default_service());
}
#[tokio::test]
async fn test_router_routes_by_isn() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
let mut cx = ServerContext::default();
cx.set_idl_service_name(FastStr::from_static_str("AnotherService"));
let result = router.call(&mut cx, Bytes::new()).await.unwrap();
assert_eq!(result, Bytes::from("another"));
}
#[tokio::test]
async fn test_router_routes_to_default_without_isn() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
let mut cx = ServerContext::default();
let result = router.call(&mut cx, Bytes::new()).await.unwrap();
assert_eq!(result, Bytes::from("default"));
}
#[tokio::test]
async fn test_router_routes_to_default_with_unknown_isn() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
let mut cx = ServerContext::default();
cx.set_idl_service_name(FastStr::from_static_str("UnknownService"));
let result = router.call(&mut cx, Bytes::new()).await.unwrap();
assert_eq!(result, Bytes::from("default"));
}
#[tokio::test]
async fn test_router_error_no_service_found() {
let router = Router::new().add_service(AnotherMockService);
let mut cx = ServerContext::default();
cx.set_idl_service_name(FastStr::from_static_str("UnknownService"));
let result = router.call(&mut cx, Bytes::new()).await;
assert!(result.is_err());
match result {
Err(ServerError::Application(e)) => {
assert_eq!(e.kind(), ApplicationExceptionKind::UNKNOWN_METHOD);
assert!(e.message().contains("UnknownService"));
}
_ => panic!("Expected ApplicationException"),
}
}
#[tokio::test]
async fn test_router_error_no_default_no_isn() {
let router = Router::new().add_service(AnotherMockService);
let mut cx = ServerContext::default();
let result = router.call(&mut cx, Bytes::new()).await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_router_routes_to_named_service_by_isn() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
let mut cx = ServerContext::default();
cx.set_idl_service_name(FastStr::from_static_str("MockService"));
let result = router.call(&mut cx, Bytes::new()).await.unwrap();
assert_eq!(result, Bytes::from("default"));
}
#[test]
fn test_router_clone() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
let cloned = router.clone();
assert_eq!(cloned.service_count(), 2);
assert!(cloned.has_default_service());
}
#[test]
fn test_router_debug() {
let router = Router::new()
.with_default_service(MockService { name: "default" })
.add_service(AnotherMockService);
let debug_str = format!("{:?}", router);
assert!(debug_str.contains("Router"));
assert!(debug_str.contains("has_default_service: true"));
}
}