use serde::Deserialize;
use solana_address::Address;
use std::str::FromStr;
use crate::error::{Error, Result};
use crate::check::{Check, CheckKind, Scenario};
use crate::replay::{
AccountAssert, CmpOp, Expect, FeatureToggle, Mutation, StateCheck, TimeTravel,
};
#[derive(Deserialize, Clone)]
pub struct FeatureInput {
pub id: String,
#[serde(default)]
pub active: bool,
}
impl FeatureInput {
pub fn into_toggle(self) -> Result<FeatureToggle> {
let id = Address::from_str(self.id.trim()).map_err(|_| {
Error::InvalidSpec(format!("bad feature id (not a pubkey): {}", self.id))
})?;
Ok(FeatureToggle {
id,
active: self.active,
})
}
}
pub fn feature_toggles(features: Vec<FeatureInput>) -> Result<Vec<FeatureToggle>> {
features
.into_iter()
.map(FeatureInput::into_toggle)
.collect()
}
#[derive(Deserialize)]
#[serde(tag = "kind", rename_all = "lowercase")]
pub enum MutationInput {
Lamports {
address: String,
lamports: u64,
},
Data {
address: String,
offset: usize,
bytes_hex: String,
},
SkipIx {
index: usize,
},
IxData {
index: usize,
bytes_hex: String,
},
MoveIx {
from: usize,
to: usize,
},
IxArg {
index: usize,
arg: String,
value: serde_json::Value,
},
Field {
address: String,
field: String,
value: serde_json::Value,
},
}
pub(crate) fn hex_decode(s: &str) -> Result<Vec<u8>> {
let s = s.trim().trim_start_matches("0x").replace([' ', '_'], "");
if s.is_empty() || !s.len().is_multiple_of(2) || !s.is_ascii() {
return Err(Error::InvalidSpec(
"hex bytes must be a non-empty, even-length hex string".into(),
));
}
let bytes = s.as_bytes();
(0..bytes.len())
.step_by(2)
.map(|i| {
let pair = std::str::from_utf8(&bytes[i..i + 2]).expect("ascii checked above");
u8::from_str_radix(pair, 16)
.map_err(|_| Error::InvalidSpec(format!("invalid hex: {s}")))
})
.collect()
}
impl MutationInput {
pub fn into_mutation(self) -> Result<Mutation> {
Ok(match self {
MutationInput::Lamports { address, lamports } => Mutation::Lamports {
address,
value: lamports,
},
MutationInput::Data {
address,
offset,
bytes_hex,
} => Mutation::DataPatch {
address,
offset,
bytes: hex_decode(&bytes_hex)?,
},
MutationInput::IxArg { index, arg, value } => Mutation::IxArg { index, arg, value },
MutationInput::SkipIx { index } => Mutation::SkipIx { index },
MutationInput::IxData { index, bytes_hex } => Mutation::IxDataReplace {
index,
bytes: hex_decode(&bytes_hex)?,
},
MutationInput::MoveIx { from, to } => Mutation::MoveIx { from, to },
MutationInput::Field {
address,
field,
value,
} => Mutation::Field {
address,
field,
value: value
.as_i64()
.map(i128::from)
.or_else(|| value.as_u64().map(i128::from))
.or_else(|| value.as_str().and_then(|s| s.trim().parse::<i128>().ok()))
.ok_or_else(|| {
crate::Error::InvalidSpec(format!(
"field value must be an integer, got {value}"
))
})?,
},
})
}
}
#[derive(Deserialize)]
pub struct AssertInput {
pub address: String,
#[serde(default = "default_kind")]
pub kind: String,
#[serde(default)]
pub offset: usize,
#[serde(default)]
pub field: Option<String>,
#[serde(default = "default_op")]
pub op: String,
#[serde(default)]
pub value: Option<i64>,
}
fn default_kind() -> String {
"u64".into()
}
fn default_op() -> String {
"==".into()
}
impl AssertInput {
fn into_assert(self) -> Result<AccountAssert> {
let op = match self.op.as_str() {
"==" | "eq" => CmpOp::Eq,
"!=" | "ne" => CmpOp::Ne,
"<" | "lt" => CmpOp::Lt,
"<=" | "le" => CmpOp::Le,
">" | "gt" => CmpOp::Gt,
">=" | "ge" => CmpOp::Ge,
other => return Err(Error::InvalidSpec(format!("unknown assert op: {other}"))),
};
let value = || -> Result<i64> {
self.value.ok_or_else(|| {
Error::InvalidSpec(format!("assert kind \"{}\" needs a \"value\"", self.kind))
})
};
let unsigned = || -> Result<u64> {
u64::try_from(value()?)
.map_err(|_| Error::InvalidSpec(format!("{} value must be ≥ 0", self.kind)))
};
let field = || -> Result<String> {
match &self.field {
Some(f) if !f.trim().is_empty() => Ok(f.trim().to_string()),
_ => Err(Error::InvalidSpec(format!(
"assert kind \"{}\" needs a \"field\" name",
self.kind
))),
}
};
let check = match self.kind.as_str() {
"lamports" => StateCheck::Lamports {
op,
value: unsigned()? as i128,
},
"u64" => StateCheck::U64At {
offset: self.offset,
op,
value: unsigned()? as i128,
},
"token_amount" => StateCheck::U64At {
offset: 64,
op,
value: unsigned()? as i128,
},
"lamports_delta" => StateCheck::LamportsDelta {
op,
value: value()? as i128,
},
"token_delta" => StateCheck::TokenDelta {
op,
value: value()? as i128,
},
"field" => StateCheck::Field {
name: field()?,
op,
value: value()? as i128,
},
"field_delta" => StateCheck::FieldDelta {
name: field()?,
op,
value: value()? as i128,
},
"field_unchanged" => StateCheck::FieldUnchanged { name: field()? },
other => return Err(Error::InvalidSpec(format!("unknown assert kind: {other}"))),
};
Ok(AccountAssert {
address: self.address,
check,
})
}
}
#[derive(Deserialize)]
pub struct ScenarioInput {
pub name: String,
#[serde(default = "default_expect")]
pub expect: String,
#[serde(default)]
pub contains: Option<String>,
#[serde(default)]
pub mutations: Vec<MutationInput>,
#[serde(default)]
pub asserts: Vec<AssertInput>,
}
fn default_expect() -> String {
"any".into()
}
impl ScenarioInput {
pub fn into_scenario(self) -> Result<Scenario> {
let expect = match self.expect.trim() {
"success" | "pass" => Expect::Success,
"revert" | "fail" => match self.contains {
Some(s) if !s.is_empty() => Expect::RevertContains(s),
_ => Expect::Revert,
},
"any" => Expect::Any,
other => {
return Err(Error::InvalidSpec(format!(
"unknown expect \"{other}\" (use success, revert, or any)"
)))
}
};
let mutations = self
.mutations
.into_iter()
.map(MutationInput::into_mutation)
.collect::<Result<_>>()?;
let mut checks = vec![Check(CheckKind::Outcome(expect))];
for a in self.asserts {
checks.push(Check(CheckKind::Account(vec![a.into_assert()?])));
}
Ok(Scenario {
name: self.name,
mutations,
checks,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hex_decode_tolerates_prefix_spaces_underscores() {
assert_eq!(
hex_decode("0xDEAD_beef").unwrap(),
vec![0xde, 0xad, 0xbe, 0xef]
);
assert_eq!(hex_decode("00 ff").unwrap(), vec![0x00, 0xff]);
}
#[test]
fn hex_decode_rejects_bad_input() {
assert!(hex_decode("").is_err());
assert!(hex_decode("abc").is_err()); assert!(hex_decode("zz").is_err());
}
#[test]
fn hex_decode_rejects_multibyte_without_panicking() {
assert!(hex_decode("aΩb").is_err());
assert!(hex_decode("0x€€").is_err());
}
#[test]
fn unknown_expect_string_is_an_error_not_a_vacuous_pass() {
let s: ScenarioInput =
serde_json::from_str(r#"{"name":"typo","expect":"sucess"}"#).unwrap();
assert!(matches!(s.into_scenario(), Err(Error::InvalidSpec(_))));
}
#[test]
fn field_mutation_parses() {
let m: MutationInput =
serde_json::from_str(r#"{"kind":"field","address":"X","field":"count","value":-5}"#)
.unwrap();
match m.into_mutation().unwrap() {
Mutation::Field {
address,
field,
value,
} => {
assert_eq!(
(address.as_str(), field.as_str(), value),
("X", "count", -5)
);
}
other => panic!("expected Field, got {other:?}"),
}
}
#[test]
fn data_mutation_becomes_patch() {
let m: MutationInput = serde_json::from_str(
r#"{"kind":"data","address":"X","offset":64,"bytes_hex":"0000000000000000"}"#,
)
.unwrap();
match m.into_mutation().unwrap() {
Mutation::DataPatch {
address,
offset,
bytes,
} => {
assert_eq!(address, "X");
assert_eq!(offset, 64);
assert_eq!(bytes, vec![0u8; 8]);
}
_ => panic!("expected DataPatch"),
}
}
#[test]
fn assert_kinds_and_ops_resolve() {
let a: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"token_amount","op":">=","value":5}"#)
.unwrap();
match a.into_assert().unwrap().check {
StateCheck::U64At {
offset: 64,
op: CmpOp::Ge,
value: 5,
} => {}
_ => panic!("expected U64At @64 >= 5"),
}
let d: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"token_delta","op":"<","value":-3}"#)
.unwrap();
match d.into_assert().unwrap().check {
StateCheck::TokenDelta {
op: CmpOp::Lt,
value: -3,
} => {}
_ => panic!("expected TokenDelta < -3"),
}
}
#[test]
fn field_asserts_resolve_and_allow_negatives() {
let a: AssertInput = serde_json::from_str(
r#"{"address":"X","kind":"field","field":"pool.reserveA","op":">=","value":1000}"#,
)
.unwrap();
match a.into_assert().unwrap().check {
StateCheck::Field {
name,
op: CmpOp::Ge,
value: 1000,
} => assert_eq!(name, "pool.reserveA"),
_ => panic!("expected Field >= 1000"),
}
let d: AssertInput = serde_json::from_str(
r#"{"address":"X","kind":"field_delta","field":"reserveA","value":-500}"#,
)
.unwrap();
match d.into_assert().unwrap().check {
StateCheck::FieldDelta {
name,
op: CmpOp::Eq,
value: -500,
} => assert_eq!(name, "reserveA"),
_ => panic!("expected FieldDelta == -500"),
}
}
#[test]
fn field_unchanged_needs_no_value_and_other_kinds_do() {
let a: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"field_unchanged","field":"authority"}"#)
.unwrap();
match a.into_assert().unwrap().check {
StateCheck::FieldUnchanged { name } => assert_eq!(name, "authority"),
other => panic!("expected FieldUnchanged, got {other:?}"),
}
let missing: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"lamports"}"#).unwrap();
let err = missing.into_assert().unwrap_err().to_string();
assert!(err.contains("needs a \"value\""), "{err}");
}
#[test]
fn field_assert_requires_a_field_name() {
let a: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"field","value":1}"#).unwrap();
assert!(a.into_assert().unwrap_err().to_string().contains("field"));
}
#[test]
fn non_delta_assert_rejects_negative_value() {
let a: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"lamports","value":-1}"#).unwrap();
assert!(a.into_assert().is_err());
}
#[test]
fn unknown_op_and_kind_error() {
let a: AssertInput =
serde_json::from_str(r#"{"address":"X","op":"~=","value":1}"#).unwrap();
assert!(a.into_assert().is_err());
let k: AssertInput =
serde_json::from_str(r#"{"address":"X","kind":"balancez","value":1}"#).unwrap();
assert!(k.into_assert().is_err());
}
#[test]
fn scenario_expect_revert_with_contains() {
let s: ScenarioInput = serde_json::from_str(
r#"{"name":"drain","expect":"revert","contains":"Slippage","mutations":[],"asserts":[]}"#,
)
.unwrap();
let scenario = s.into_scenario().unwrap();
match &scenario.checks[0].0 {
CheckKind::Outcome(Expect::RevertContains(t)) => assert_eq!(t, "Slippage"),
other => panic!("expected RevertContains, got {other:?}"),
}
}
}
#[derive(Deserialize)]
pub struct SuiteRequest {
#[serde(default)]
pub signature: Option<String>,
#[serde(default)]
pub fixture: Option<String>,
#[serde(default)]
pub cluster: Option<String>,
#[serde(default)]
pub rpc: Option<String>,
#[serde(default)]
pub time_travel: TimeTravel,
#[serde(default)]
pub features: Vec<FeatureInput>,
pub scenarios: Vec<ScenarioInput>,
}