use std::sync::Mutex;
use async_trait::async_trait;
use gluonscan_core::{Chain, ChainProvider, Clock, Error, Http, Timestamp};
use serde::Deserialize;
#[derive(Debug, Clone)]
pub enum Match {
Any,
PrimaryContains(String),
PrimaryIs(String),
BodyContains(String),
JsonEq {
pointer: String,
value: serde_json::Value,
},
All(Vec<Match>),
AnyOf(Vec<Match>),
}
impl Match {
pub fn primary_contains(s: impl Into<String>) -> Match {
Match::PrimaryContains(s.into())
}
pub fn method(s: impl Into<String>) -> Match {
Match::PrimaryIs(s.into())
}
pub fn body_contains(s: impl Into<String>) -> Match {
Match::BodyContains(s.into())
}
pub fn json_eq(pointer: impl Into<String>, value: serde_json::Value) -> Match {
Match::JsonEq {
pointer: pointer.into(),
value,
}
}
pub fn all(m: impl IntoIterator<Item = Match>) -> Match {
Match::All(m.into_iter().collect())
}
pub fn any_of(m: impl IntoIterator<Item = Match>) -> Match {
Match::AnyOf(m.into_iter().collect())
}
pub fn matches(&self, primary: &str, body: &str) -> bool {
match self {
Match::Any => true,
Match::PrimaryContains(s) => primary.contains(s.as_str()),
Match::PrimaryIs(s) => primary == s,
Match::BodyContains(s) => body.contains(s.as_str()),
Match::JsonEq { pointer, value } => serde_json::from_str::<serde_json::Value>(body)
.ok()
.and_then(|v| v.pointer(pointer).cloned())
.is_some_and(|found| &found == value),
Match::All(ms) => ms.iter().all(|m| m.matches(primary, body)),
Match::AnyOf(ms) => ms.iter().any(|m| m.matches(primary, body)),
}
}
}
#[derive(Debug, Clone)]
pub struct Contract {
pub when: Match,
pub reply: String,
}
#[derive(Debug, Clone, Deserialize)]
pub struct ContractSpec {
#[serde(default)]
pub primary_contains: Option<String>,
#[serde(default)]
pub method: Option<String>,
#[serde(default)]
pub body_contains: Option<String>,
#[serde(default)]
pub reply: Option<String>,
#[serde(default)]
pub reply_path: Option<String>,
}
impl ContractSpec {
pub fn into_contract(self, base_dir: &std::path::Path) -> std::io::Result<Contract> {
let mut ms = Vec::new();
if let Some(s) = self.primary_contains {
ms.push(Match::PrimaryContains(s));
}
if let Some(s) = self.method {
ms.push(Match::PrimaryIs(s));
}
if let Some(s) = self.body_contains {
ms.push(Match::BodyContains(s));
}
let when = if ms.is_empty() {
Match::Any
} else {
Match::All(ms)
};
let reply = match (self.reply, self.reply_path) {
(Some(r), _) => r,
(None, Some(p)) => std::fs::read_to_string(base_dir.join(p))?,
(None, None) => String::new(),
};
Ok(Contract { when, reply })
}
}
pub fn load_contracts(dir: impl AsRef<std::path::Path>) -> std::io::Result<Vec<Contract>> {
let dir = dir.as_ref();
let mut out = Vec::new();
for entry in std::fs::read_dir(dir)? {
let path = entry?.path();
if path.extension().and_then(|e| e.to_str()) == Some("json") {
let spec: ContractSpec = serde_json::from_str(&std::fs::read_to_string(&path)?)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?;
out.push(spec.into_contract(dir)?);
}
}
Ok(out)
}
fn no_match(primary: &str, body: &str) -> Error {
let preview: String = body.chars().take(200).collect();
Error::Permanent {
message: format!("no contract matched request to `{primary}` with body: {preview}"),
}
}
type ErrorContract = (Match, Box<dyn Fn() -> Error + Send + Sync>);
#[derive(Default)]
pub struct MockHttp {
contracts: Vec<Contract>,
errors: Vec<ErrorContract>,
calls: Mutex<Vec<(String, String)>>,
}
impl MockHttp {
pub fn new() -> Self {
MockHttp::default()
}
pub fn on(mut self, when: Match, reply: impl Into<String>) -> Self {
self.contracts.push(Contract {
when,
reply: reply.into(),
});
self
}
pub fn on_err(
mut self,
when: Match,
error: impl Fn() -> Error + Send + Sync + 'static,
) -> Self {
self.errors.push((when, Box::new(error)));
self
}
pub fn on_transient(self, when: Match) -> Self {
self.on_err(when, || Error::Transient {
message: "mock transient (HTTP 429)".into(),
retry_after: None,
})
}
pub fn from_contracts(contracts: Vec<Contract>) -> Self {
MockHttp {
contracts,
errors: Vec::new(),
calls: Mutex::new(Vec::new()),
}
}
pub fn calls(&self) -> Vec<(String, String)> {
self.calls.lock().unwrap().clone()
}
fn reply(&self, primary: &str, body: &str) -> Result<String, Error> {
if let Some((_, factory)) = self.errors.iter().find(|(w, _)| w.matches(primary, body)) {
return Err(factory());
}
self.contracts
.iter()
.find(|c| c.when.matches(primary, body))
.map(|c| c.reply.clone())
.ok_or_else(|| no_match(primary, body))
}
}
#[async_trait]
impl Http for MockHttp {
async fn post(
&self,
url: &str,
body: String,
_headers: &[(&str, &str)],
) -> Result<String, Error> {
self.calls
.lock()
.unwrap()
.push((url.to_string(), body.clone()));
self.reply(url, &body)
}
async fn get(&self, url: &str, _headers: &[(&str, &str)]) -> Result<String, Error> {
self.calls
.lock()
.unwrap()
.push((url.to_string(), String::new()));
self.reply(url, "")
}
}
#[derive(Default)]
pub struct MockChainProvider {
contracts: Vec<Contract>,
calls: Mutex<Vec<(String, String)>>,
}
impl MockChainProvider {
pub fn new() -> Self {
MockChainProvider::default()
}
pub fn on(mut self, when: Match, reply: impl Into<String>) -> Self {
self.contracts.push(Contract {
when,
reply: reply.into(),
});
self
}
pub fn calls(&self) -> Vec<(String, String)> {
self.calls.lock().unwrap().clone()
}
}
#[async_trait]
impl ChainProvider for MockChainProvider {
async fn call(&self, _chain: Chain, method: &str, params: String) -> Result<String, Error> {
self.calls
.lock()
.unwrap()
.push((method.to_string(), params.clone()));
self.contracts
.iter()
.find(|c| c.when.matches(method, ¶ms))
.map(|c| c.reply.clone())
.ok_or_else(|| no_match(method, ¶ms))
}
}
#[derive(Debug, Clone, Copy)]
pub struct MockClock(pub i64);
impl Clock for MockClock {
fn now(&self) -> Timestamp {
Timestamp(self.0)
}
}