use std::future::Future;
use std::pin::Pin;
use futures::{future, FutureExt};
use http_body_util::{BodyExt, Limited};
use hyper::header;
use jsonrpsee::{
core::BoxError,
server::{HttpBody, HttpRequest, HttpResponse},
};
use jsonrpsee_types::ErrorObject;
use serde::{Deserialize, Serialize};
use tower::Service;
use super::cookie::Cookie;
use base64::{engine::general_purpose::STANDARD, Engine as _};
#[derive(Clone, Debug)]
pub struct HttpRequestMiddleware<S> {
service: S,
cookie: Option<Cookie>,
max_request_body_size: usize,
}
impl<S> HttpRequestMiddleware<S> {
pub fn new(service: S, cookie: Option<Cookie>, max_request_body_size: usize) -> Self {
Self {
service,
cookie,
max_request_body_size,
}
}
pub fn check_credentials(&self, headers: &header::HeaderMap) -> bool {
self.cookie.as_ref().is_none_or(|internal_cookie| {
headers
.get(header::AUTHORIZATION)
.and_then(|auth_header| auth_header.to_str().ok())
.and_then(|auth_header| auth_header.split_whitespace().nth(1))
.and_then(|encoded| STANDARD.decode(encoded).ok())
.and_then(|decoded| String::from_utf8(decoded).ok())
.and_then(|request_cookie| request_cookie.split(':').nth(1).map(String::from))
.is_some_and(|passwd| internal_cookie.authenticate(passwd))
})
}
pub fn insert_or_replace_content_type_header(headers: &mut header::HeaderMap) {
if !headers.contains_key(header::CONTENT_TYPE)
|| headers
.get(header::CONTENT_TYPE)
.filter(|value| {
value
.to_str()
.ok()
.unwrap_or_default()
.starts_with("text/plain")
})
.is_some()
{
headers.insert(
header::CONTENT_TYPE,
header::HeaderValue::from_static("application/json"),
);
}
}
async fn request_to_json_rpc_2<B>(
request: HttpRequest<B>,
max_request_body_size: usize,
) -> Result<(JsonRpcVersion, HttpRequest<HttpBody>), BoxError>
where
B: hyper::body::Body<Data = hyper::body::Bytes> + Send + 'static,
B::Error: Into<BoxError>,
{
let (parts, body) = request.into_parts();
let bytes = Limited::new(body, max_request_body_size)
.collect()
.await?
.to_bytes();
let (version, bytes) =
if let Ok(request) = serde_json::from_slice::<'_, JsonRpcRequest>(bytes.as_ref()) {
let version = request.version();
if matches!(version, JsonRpcVersion::Unknown) {
(version, bytes)
} else {
(
version,
serde_json::to_vec(&request.into_2()).expect("valid").into(),
)
}
} else {
(JsonRpcVersion::Unknown, bytes)
};
Ok((
version,
HttpRequest::from_parts(parts, HttpBody::from(bytes.as_ref().to_vec())),
))
}
async fn response_from_json_rpc_2(
version: JsonRpcVersion,
response: HttpResponse<HttpBody>,
) -> Result<HttpResponse<HttpBody>, BoxError> {
let (parts, body) = response.into_parts();
let bytes = body.collect().await?.to_bytes();
let bytes =
if let Ok(response) = serde_json::from_slice::<'_, JsonRpcResponse>(bytes.as_ref()) {
serde_json::to_vec(&response.into_version(version))
.expect("valid")
.into()
} else {
bytes
};
Ok(HttpResponse::from_parts(
parts,
HttpBody::from(bytes.as_ref().to_vec()),
))
}
}
#[derive(Clone)]
pub struct HttpRequestMiddlewareLayer {
cookie: Option<Cookie>,
max_request_body_size: usize,
}
impl HttpRequestMiddlewareLayer {
pub fn new(cookie: Option<Cookie>, max_request_body_size: usize) -> Self {
Self {
cookie,
max_request_body_size,
}
}
}
impl<S> tower::Layer<S> for HttpRequestMiddlewareLayer {
type Service = HttpRequestMiddleware<S>;
fn layer(&self, service: S) -> Self::Service {
HttpRequestMiddleware::new(service, self.cookie.clone(), self.max_request_body_size)
}
}
impl<S, B> Service<HttpRequest<B>> for HttpRequestMiddleware<S>
where
S: Service<HttpRequest, Response = HttpResponse> + std::clone::Clone + Send + 'static,
S::Error: Into<BoxError> + 'static,
S::Future: Send + 'static,
B: hyper::body::Body<Data = hyper::body::Bytes> + Send + 'static,
B::Error: Into<BoxError>,
{
type Response = S::Response;
type Error = BoxError;
type Future =
Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send + 'static>>;
fn poll_ready(
&mut self,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Result<(), Self::Error>> {
self.service.poll_ready(cx).map_err(Into::into)
}
fn call(&mut self, mut request: HttpRequest<B>) -> Self::Future {
if !self.check_credentials(request.headers_mut()) {
let error = ErrorObject::borrowed(401, "unauthenticated method", None);
return future::err(BoxError::from(error)).boxed();
}
Self::insert_or_replace_content_type_header(request.headers_mut());
let mut service = self.service.clone();
let max_request_body_size = self.max_request_body_size;
async move {
let (version, request) =
Self::request_to_json_rpc_2(request, max_request_body_size).await?;
let response = service.call(request).await.map_err(Into::into)?;
Self::response_from_json_rpc_2(version, response).await
}
.boxed()
}
}
#[derive(Clone, Copy, Debug)]
enum JsonRpcVersion {
Bitcoind,
Lightwalletd,
TwoPointZero,
Unknown,
}
#[derive(Debug, Deserialize, Serialize)]
struct JsonRpcRequest {
#[serde(skip_serializing_if = "Option::is_none")]
jsonrpc: Option<String>,
method: String,
#[serde(skip_serializing_if = "Option::is_none")]
params: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
id: Option<serde_json::Value>,
}
impl JsonRpcRequest {
fn version(&self) -> JsonRpcVersion {
match (self.jsonrpc.as_deref(), &self.params, &self.id) {
(
Some("2.0"),
_,
None
| Some(
serde_json::Value::Null
| serde_json::Value::String(_)
| serde_json::Value::Number(_),
),
) => JsonRpcVersion::TwoPointZero,
(Some("1.0"), Some(_), Some(_)) => JsonRpcVersion::Lightwalletd,
(None, Some(_), Some(_)) => JsonRpcVersion::Bitcoind,
_ => JsonRpcVersion::Unknown,
}
}
fn into_2(mut self) -> Self {
self.jsonrpc = Some("2.0".into());
self
}
}
#[derive(Debug, Deserialize, Serialize)]
struct JsonRpcResponse {
#[serde(skip_serializing_if = "Option::is_none")]
jsonrpc: Option<String>,
id: serde_json::Value,
#[serde(skip_serializing_if = "Option::is_none")]
result: Option<Box<serde_json::value::RawValue>>,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<serde_json::Value>,
}
impl JsonRpcResponse {
fn into_version(mut self, version: JsonRpcVersion) -> Self {
match version {
JsonRpcVersion::Bitcoind => {
self.jsonrpc = None;
self.result = self
.result
.or_else(|| serde_json::value::to_raw_value(&()).ok());
self.error = self.error.or(Some(serde_json::Value::Null));
}
JsonRpcVersion::Lightwalletd => {
self.jsonrpc = Some("1.0".into());
self.result = self
.result
.or_else(|| serde_json::value::to_raw_value(&()).ok());
self.error = self.error.or(Some(serde_json::Value::Null));
}
JsonRpcVersion::TwoPointZero => {
assert_eq!(self.jsonrpc.as_deref(), Some("2.0"));
if self.error.is_none() {
self.result = self
.result
.or_else(|| serde_json::value::to_raw_value(&()).ok());
} else {
assert!(self.result.is_none());
}
}
JsonRpcVersion::Unknown => (),
}
self
}
}