use std::collections::HashMap;
use std::{collections::HashSet, convert::Infallible, fmt::Write as _, fs, path::Path};
use axum::body::Body;
use futures::FutureExt;
use http::{HeaderValue, Method, Request};
use jsonrpsee::server::ws::is_upgrade_request;
pub use jsonrpsee::server::ServerHandle;
use jsonrpsee::RpcModule;
use tower::service_fn;
use tower::Service;
use tower::ServiceBuilder;
use crate::builder::*;
use crate::RequestKind;
#[derive(Clone)]
pub struct Router<Ctx> {
nested_routers: Vec<(&'static str, Router<Ctx>)>,
handlers: Vec<HandlerCallbacks<Ctx>>,
}
impl<Ctx> Router<Ctx>
where
Ctx: Clone + Send + Sync + 'static,
{
pub fn new() -> Self {
Self::default()
}
pub fn handler<H: Handler<Ctx>>(mut self, handler: H) -> Self {
self.handlers.push(HandlerCallbacks::from_handler(handler));
self
}
pub fn nest(mut self, namespace: &'static str, router: Router<Ctx>) -> Self {
self.nested_routers.push((namespace, router));
self
}
pub fn write_bindings_to_dir(&self, out_dir: impl AsRef<Path>) {
let out_dir = out_dir.as_ref();
fs::create_dir_all(out_dir).unwrap();
fs::remove_dir_all(out_dir).unwrap();
fs::create_dir_all(out_dir).unwrap();
let header = String::from(include_str!("../header.txt"));
let (imports, exports, _types) = self
.get_handlers()
.into_iter()
.flat_map(|handler| {
(handler.export_all_dependencies_to)(out_dir)
.unwrap()
.into_iter()
.map(|dep| {
(
format!("./{}", dep.output_path.to_str().unwrap()),
dep.ts_name,
)
})
.chain((handler.qubit_types)().into_iter().map(|ty| ty.to_ts()))
})
.fold(
(String::new(), String::new(), HashSet::new()),
|(mut imports, mut exports, mut types), ty| {
if types.contains(&ty) {
return (imports, exports, types);
}
let (package, ty_name) = ty;
writeln!(
&mut imports,
r#"import type {{ {ty_name} }} from "{package}";"#,
)
.unwrap();
writeln!(
&mut exports,
r#"export type {{ {ty_name} }} from "{package}";"#,
)
.unwrap();
types.insert((package, ty_name));
(imports, exports, types)
},
);
let server_type = format!("export type QubitServer = {};", self.get_type());
fs::write(
out_dir.join("index.ts"),
[header, imports, exports, server_type]
.into_iter()
.filter(|part| !part.is_empty())
.collect::<Vec<_>>()
.join("\n"),
)
.unwrap();
}
pub fn to_service(
self,
ctx: Ctx,
) -> (
impl Service<
hyper::Request<axum::body::Body>,
Response = jsonrpsee::server::HttpResponse,
Error = Infallible,
Future = impl Send,
> + Clone,
ServerHandle,
) {
let (stop_handle, server_handle) = jsonrpsee::server::stop_channel();
let mut service = jsonrpsee::server::Server::builder()
.set_http_middleware(ServiceBuilder::new().map_request(|mut req: Request<_>| {
let request_type = if matches!(req.method(), &Method::GET)
&& !is_upgrade_request(&req)
{
*req.method_mut() = Method::POST;
let headers = req.headers_mut();
headers.insert(
hyper::header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
headers.insert(
hyper::header::ACCEPT,
HeaderValue::from_static("application/json"),
);
if let Some(body) = req
.uri()
.query()
.and_then(|query| serde_qs::from_str::<HashMap<String, String>>(query).ok())
.and_then(|mut query| query.remove("input"))
.map(|input| urlencoding::decode(&input).unwrap_or_default().to_string())
{
*req.body_mut() = Body::from(body);
}
RequestKind::Query
} else {
RequestKind::Any
};
req.extensions_mut().insert(request_type);
req
}))
.to_service_builder()
.build(self.build_rpc_module(ctx, None), stop_handle);
(
service_fn(move |req: hyper::Request<axum::body::Body>| {
let call = service.call(req);
async move {
match call.await {
Ok(response) => Ok::<_, Infallible>(response),
Err(_) => unreachable!(),
}
}
.boxed()
}),
server_handle,
)
}
fn get_type(&self) -> String {
let handlers = self
.handlers
.iter()
.map(|handler| {
let handler_type = (handler.get_type)();
format!("{}: {}", handler_type.name, handler_type.signature)
})
.chain(
self.nested_routers.iter().map(|(namespace, router)| {
let router_type = router.get_type();
format!("{namespace}: {router_type}")
}),
)
.collect::<Vec<_>>();
format!("{{ {} }}", handlers.join(", "))
}
fn build_rpc_module(self, ctx: Ctx, namespace: Option<&'static str>) -> RpcModule<Ctx> {
let rpc_module = self
.handlers
.into_iter()
.fold(
RpcBuilder::with_namespace(ctx.clone(), namespace),
|rpc_builder, handler| (handler.register)(rpc_builder),
)
.build();
let parent_namespace = namespace;
self.nested_routers
.into_iter()
.fold(rpc_module, |mut rpc_module, (namespace, router)| {
let namespace = if let Some(parent_namespace) = parent_namespace {
format!("{parent_namespace}.{namespace}").leak()
} else {
namespace
};
rpc_module
.merge(router.build_rpc_module(ctx.clone(), Some(namespace)))
.unwrap();
rpc_module
})
}
fn get_handlers(&self) -> Vec<HandlerCallbacks<Ctx>> {
self.handlers
.iter()
.cloned()
.chain(
self.nested_routers
.iter()
.flat_map(|(_, router)| router.get_handlers()),
)
.collect()
}
}
impl<Ctx> Default for Router<Ctx> {
fn default() -> Self {
Self {
nested_routers: Default::default(),
handlers: Default::default(),
}
}
}