feat: harden trading SDK examples

This commit is contained in:
0xfnzero
2026-07-12 14:35:58 +08:00
parent 47cef59d15
commit b09902c527
70 changed files with 1447 additions and 471 deletions
+50
View File
@@ -1348,6 +1348,12 @@ impl TradingClient {
(bool, Vec<Signature>, Option<TradeError>, Vec<(crate::swqos::SwqosType, i64)>),
anyhow::Error,
> {
validate_trade_safety(
"buy",
params.input_token_amount,
params.fixed_output_token_amount,
params.slippage_basis_points,
)?;
if params.recent_blockhash.is_none() && params.durable_nonce.is_none() {
return Err(anyhow::anyhow!(
"Must provide either recent_blockhash or durable_nonce for buy (required for transaction validity)"
@@ -1476,6 +1482,12 @@ impl TradingClient {
(bool, Vec<Signature>, Option<TradeError>, Vec<(crate::swqos::SwqosType, i64)>),
anyhow::Error,
> {
validate_trade_safety(
"sell",
params.input_token_amount,
params.fixed_output_token_amount,
params.slippage_basis_points,
)?;
#[cfg(feature = "perf-trace")]
if sdk_log::sdk_log_enabled() && params.slippage_basis_points.is_none() {
debug!(
@@ -1820,6 +1832,30 @@ impl TradingClient {
}
}
fn validate_trade_safety(
side: &str,
input_amount: u64,
fixed_output_amount: Option<u64>,
slippage_basis_points: Option<u64>,
) -> Result<(), anyhow::Error> {
if input_amount == 0 {
return Err(anyhow::anyhow!("{} input amount must be greater than zero", side));
}
if fixed_output_amount == Some(0) {
return Err(anyhow::anyhow!("{} fixed output amount must be greater than zero", side));
}
if let Some(bps) = slippage_basis_points {
if bps >= 10_000 {
return Err(anyhow::anyhow!(
"{} slippage_basis_points must be below 10000, got {}",
side,
bps
));
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
@@ -1841,6 +1877,20 @@ mod tests {
})
}
#[test]
fn trade_safety_rejects_zero_amounts_and_unbounded_slippage() {
assert!(validate_trade_safety("buy", 0, None, Some(100)).is_err());
assert!(validate_trade_safety("buy", 1, Some(0), Some(100)).is_err());
assert!(validate_trade_safety("sell", 1, None, Some(10_000)).is_err());
assert!(validate_trade_safety("sell", 1, None, Some(u64::MAX)).is_err());
}
#[test]
fn trade_safety_accepts_bounded_values() {
assert!(validate_trade_safety("buy", 1, None, None).is_ok());
assert!(validate_trade_safety("buy", 1, Some(1), Some(9_999)).is_ok());
}
#[test]
fn normalize_swqos_configs_adds_default_rpc_route() {
let configs = vec![SwqosConfig::Jito("uuid".to_string(), SwqosRegion::Frankfurt, None)];
+50
View File
@@ -0,0 +1,50 @@
use anyhow::{Context, Result};
use solana_sdk::signature::Keypair;
/// Load a Solana keypair from a base58 string or a 64-byte JSON array in an environment variable.
pub fn load_keypair_from_env(name: &str) -> Result<Keypair> {
let value = std::env::var(name).with_context(|| format!("{} is required", name))?;
load_keypair_from_string(&value).with_context(|| format!("invalid {}", name))
}
/// Parse a Solana keypair without the panic behavior of `Keypair::from_base58_string`.
pub fn load_keypair_from_string(value: &str) -> Result<Keypair> {
let value = value.trim();
if value.is_empty() {
anyhow::bail!("keypair is empty");
}
let bytes = if value.starts_with('[') {
serde_json::from_str::<Vec<u8>>(value).context("keypair JSON must be a byte array")?
} else {
bs58::decode(value)
.into_vec()
.context("keypair must be valid base58 or a 64-byte JSON array")?
};
if bytes.len() != 64 {
anyhow::bail!("keypair must contain 64 bytes, got {}", bytes.len());
}
Keypair::try_from(bytes.as_slice()).context("invalid 64-byte Solana keypair")
}
#[cfg(test)]
mod tests {
use super::*;
use solana_sdk::signature::Signer;
#[test]
fn parses_base58_and_json_keypairs() {
let keypair = Keypair::new();
let base58 = bs58::encode(keypair.to_bytes()).into_string();
let json = serde_json::to_string(&keypair.to_bytes().to_vec()).unwrap();
assert_eq!(load_keypair_from_string(&base58).unwrap().pubkey(), keypair.pubkey());
assert_eq!(load_keypair_from_string(&json).unwrap().pubkey(), keypair.pubkey());
}
#[test]
fn rejects_placeholders_and_wrong_lengths() {
assert!(load_keypair_from_string("use_your_payer_keypair_here").is_err());
assert!(load_keypair_from_string("[1,2,3]").is_err());
}
}
+1
View File
@@ -5,6 +5,7 @@ pub mod fast_fn;
pub mod fast_timing;
pub mod gas_fee_strategy;
pub mod global;
pub mod keypair;
pub mod nonce_cache;
pub mod sdk_log;
pub mod seed;