use std::time::Duration;
use base64::engine::general_purpose;
use base64::Engine;
use bitreq::{post, Client as BitreqClient, Error as BitreqError, RequestExt};
use serde_json::Value;
use thiserror::Error;
use tokio::time::sleep;
use tracing::warn;
#[derive(Debug, Error, serde::Serialize, serde::Deserialize)]
pub enum TransportError {
#[error("HTTP error: {0}")]
Http(String),
#[error("JSON error: {0}")]
Json(String),
#[error("RPC error: {0}")]
Rpc(String),
#[error("Connection error: {0}")]
ConnectionError(String),
#[error("HttpRedirect: {0}")]
HttpRedirect(String),
#[error("Malformed Response: {0}")]
MalformedResponse(String),
#[error("Error parsing rpc response: {0}")]
Parse(String),
#[error("Max retries {0} exceeded")]
MaxRetriesExceeded(u8),
}
impl From<BitreqError> for TransportError {
fn from(value: BitreqError) -> Self {
match value {
BitreqError::AddressNotFound
| BitreqError::IoError(_)
| BitreqError::RustlsCreateConnection(_) =>
TransportError::ConnectionError(value.to_string()),
BitreqError::RedirectLocationMissing
| BitreqError::InfiniteRedirectionLoop
| BitreqError::TooManyRedirections => TransportError::HttpRedirect(value.to_string()),
BitreqError::HeadersOverflow
| BitreqError::StatusLineOverflow
| BitreqError::BodyOverflow
| BitreqError::MalformedChunkLength
| BitreqError::MalformedChunkEnd
| BitreqError::MalformedContentLength
| BitreqError::InvalidUtf8InResponse
| BitreqError::InvalidUtf8InBody(_) =>
TransportError::MalformedResponse(value.to_string()),
_ => TransportError::Http(value.to_string()),
}
}
}
impl From<serde_json::Error> for TransportError {
fn from(err: serde_json::Error) -> Self { TransportError::Json(err.to_string()) }
}
impl From<std::io::Error> for TransportError {
fn from(err: std::io::Error) -> Self { TransportError::Rpc(err.to_string()) }
}
pub trait TransportTrait: Send + Sync {
fn send_request<'a>(
&'a self,
method: &'a str,
params: &'a [Value],
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Value, TransportError>> + Send + 'a>,
>;
fn send_batch<'a>(
&'a self,
bodies: &'a [Value],
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Vec<Value>, TransportError>> + Send + 'a>,
>;
fn url(&self) -> &str;
}
pub trait TransportExt {
fn call<'a, T: serde::de::DeserializeOwned>(
&'a self,
method: &'a str,
params: &'a [Value],
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<T, TransportError>> + Send + 'a>>;
}
impl<T: TransportTrait> TransportExt for T {
fn call<'a, T2: serde::de::DeserializeOwned>(
&'a self,
method: &'a str,
params: &'a [Value],
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<T2, TransportError>> + Send + 'a>>
{
Box::pin(async move {
let result = self.send_request(method, params).await?;
Ok(serde_json::from_value(result)?)
})
}
}
#[derive(Clone)]
pub struct DefaultTransport {
client: BitreqClient,
url: String,
authorization: Option<String>,
timeout_secs: u64,
max_retries: u8,
retry_interval: u64,
wallet_name: Option<String>,
}
const DEFAULT_HTTP_CLIENT_CAPACITY: usize = 10;
const DEFAULT_TIMEOUT_SECONDS: u64 = 30;
const DEFAULT_MAX_RETRIES: u8 = 3;
const DEFAULT_RETRY_INTERVAL_MS: u64 = 1_000;
impl DefaultTransport {
pub fn new(url: impl Into<String>, auth: Option<(String, String)>) -> Self {
let authorization = auth.as_ref().map(|(u, p)| {
format!("Basic {}", general_purpose::STANDARD.encode(format!("{}:{}", u, p)))
});
Self {
client: BitreqClient::new(DEFAULT_HTTP_CLIENT_CAPACITY),
url: url.into(),
authorization,
timeout_secs: DEFAULT_TIMEOUT_SECONDS,
max_retries: DEFAULT_MAX_RETRIES,
retry_interval: DEFAULT_RETRY_INTERVAL_MS,
wallet_name: None,
}
}
pub fn with_wallet(mut self, wallet_name: impl Into<String>) -> Self {
self.wallet_name = Some(wallet_name.into());
self
}
fn is_bitreq_error_recoverable(err: &BitreqError) -> bool {
match err {
BitreqError::AddressNotFound
| BitreqError::IoError(_)
| BitreqError::RustlsCreateConnection(_) => {
warn!(err = %err, "connection error, retrying...");
true
}
BitreqError::RedirectLocationMissing => false,
BitreqError::InfiniteRedirectionLoop => false,
BitreqError::TooManyRedirections => false,
BitreqError::HeadersOverflow => false,
BitreqError::StatusLineOverflow => false,
BitreqError::BodyOverflow => false,
BitreqError::MalformedChunkLength
| BitreqError::MalformedChunkEnd
| BitreqError::MalformedContentLength
| BitreqError::InvalidUtf8InResponse => {
warn!(err = %err, "malformed response, retrying...");
true
}
BitreqError::InvalidUtf8InBody(_) => false,
BitreqError::HttpsFeatureNotEnabled => false,
BitreqError::Other(_) => false,
_ => false,
}
}
}
impl std::fmt::Debug for DefaultTransport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DefaultTransport")
.field("url", &self.url)
.field("timeout_secs", &self.timeout_secs)
.field("max_retries", &self.max_retries)
.field("retry_interval", &self.retry_interval)
.field("wallet_name", &self.wallet_name)
.finish_non_exhaustive()
}
}
enum DoRequestError {
Network(BitreqError),
Transport(TransportError),
}
impl TransportTrait for DefaultTransport {
fn send_request<'a>(
&'a self,
method: &'a str,
params: &'a [Value],
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Value, TransportError>> + Send + 'a>,
> {
let client = self.client.clone();
let url = self.url.clone();
let authorization = self.authorization.clone();
let wallet_name = self.wallet_name.clone();
let timeout_secs = self.timeout_secs;
let max_retries = self.max_retries;
let retry_interval = self.retry_interval;
async fn do_request(
client: &BitreqClient,
url: &str,
authorization: &Option<String>,
request: &serde_json::Value,
timeout_secs: u64,
) -> Result<Value, DoRequestError> {
let body = serde_json::to_vec(request)
.map_err(|e| DoRequestError::Transport(TransportError::Json(e.to_string())))?;
let mut req = post(url)
.with_header("Content-Type", "application/json")
.with_body(body)
.with_timeout(timeout_secs);
if let Some(ref h) = authorization {
req = req.with_header("Authorization", h);
}
let response =
req.send_async_with_client(client).await.map_err(DoRequestError::Network)?;
let status_code = response.status_code;
if !(200..300).contains(&status_code) {
return Err(DoRequestError::Transport(TransportError::Http(format!(
"{} {}",
status_code, response.reason_phrase
))));
}
let raw = response.as_str().map_err(|e: BitreqError| {
DoRequestError::Transport(TransportError::Parse(e.to_string()))
})?;
let json: Value = serde_json::from_str(raw)
.map_err(|e| DoRequestError::Transport(TransportError::Parse(e.to_string())))?;
if let Some(error) = json.get("error") {
if !error.is_null() {
return Err(DoRequestError::Transport(TransportError::Rpc(error.to_string())));
}
}
json.get("result").cloned().ok_or_else(|| {
DoRequestError::Transport(TransportError::Rpc("No result field".to_string()))
})
}
Box::pin(async move {
let request = serde_json::json!({
"jsonrpc": "2.0", "id": "1", "method": method, "params": params
});
let mut retries = 0u8;
loop {
let target_url = if let Some(ref wallet) = wallet_name {
format!("{}/wallet/{}", url.trim_end_matches('/'), wallet)
} else {
url.clone()
};
match do_request(&client, &target_url, &authorization, &request, timeout_secs).await
{
Ok(v) => return Ok(v),
Err(DoRequestError::Transport(TransportError::Rpc(ref msg)))
if wallet_name.is_some() && msg.contains("\"code\":-32601") =>
match do_request(&client, &url, &authorization, &request, timeout_secs)
.await
{
Ok(v) => return Ok(v),
Err(DoRequestError::Network(e)) => return Err(TransportError::from(e)),
Err(DoRequestError::Transport(e)) => return Err(e),
},
Err(DoRequestError::Network(bitreq_err)) => {
if !Self::is_bitreq_error_recoverable(&bitreq_err) {
return Err(TransportError::from(bitreq_err));
}
}
Err(DoRequestError::Transport(err)) => {
return Err(err);
}
}
retries += 1;
if retries >= max_retries {
return Err(TransportError::MaxRetriesExceeded(max_retries));
}
sleep(Duration::from_millis(retry_interval)).await;
}
})
}
fn send_batch<'a>(
&'a self,
bodies: &'a [Value],
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Vec<Value>, TransportError>> + Send + 'a>,
> {
let client = self.client.clone();
let url = self.url.clone();
let authorization = self.authorization.clone();
let timeout_secs = self.timeout_secs;
Box::pin(async move {
let bodies_vec: Vec<Value> = bodies.to_vec();
let body =
serde_json::to_vec(&bodies_vec).map_err(|e| TransportError::Json(e.to_string()))?;
let mut req = post(&url)
.with_header("Content-Type", "application/json")
.with_body(body)
.with_timeout(timeout_secs);
if let Some(ref h) = authorization {
req = req.with_header("Authorization", h);
}
let response = req.send_async_with_client(&client).await?;
let status_code = response.status_code;
if !(200..300).contains(&status_code) {
return Err(TransportError::Http(format!("HTTP {}", status_code)));
}
let raw =
response.as_str().map_err(|e: BitreqError| TransportError::Parse(e.to_string()))?;
let v: Vec<Value> =
serde_json::from_str(raw).map_err(|e| TransportError::Parse(e.to_string()))?;
Ok(v)
})
}
fn url(&self) -> &str { &self.url }
}