#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
#[cfg_attr(feature = "agent", derive(bincode::Encode, bincode::Decode))]
pub enum AccessLevel {
#[default]
Market,
View,
Trading,
}
impl AccessLevel {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Market => "market",
Self::View => "view",
Self::Trading => "trading",
}
}
#[must_use]
pub fn permits(self, required: Self) -> bool {
self >= required
}
}
impl std::str::FromStr for AccessLevel {
type Err = crate::error::Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.trim().to_ascii_lowercase().as_str() {
"market" => Ok(Self::Market),
"view" => Ok(Self::View),
"trading" => Ok(Self::Trading),
other => Err(crate::error::Error::invalid_request(format!(
"invalid access tier '{other}' (expected one of: {ACCESS_LADDER})"
))),
}
}
}
pub const ACCESS_LADDER: &str = "market < view < trading";
#[must_use]
pub fn required_access_for(method: &str, path: &str) -> AccessLevel {
if !method.eq_ignore_ascii_case("GET") {
return AccessLevel::Trading;
}
let p = path.trim_start_matches('/');
let is_public_market = p == "tickers"
|| p.starts_with("order-book/")
|| p.starts_with("public/")
|| p.starts_with("candles/")
|| p.starts_with("trades/all/")
|| p.starts_with("configuration/");
if is_public_market {
AccessLevel::Market
} else {
AccessLevel::View
}
}
#[must_use]
pub fn access_denied(required: AccessLevel, current: AccessLevel) -> String {
format!(
"access denied: this operation needs `--access {req}`, but the session is running with \
`--access {cur}`. (Re)start with `--access {req}` (or higher; the tiers are {ACCESS_LADDER}) \
to allow it.",
req = required.as_str(),
cur = current.as_str(),
)
}
#[cfg(test)]
mod tests {
use super::{AccessLevel, access_denied, required_access_for};
#[test]
fn tiers_are_ordered_and_cumulative() {
assert!(AccessLevel::Market < AccessLevel::View);
assert!(AccessLevel::View < AccessLevel::Trading);
assert!(AccessLevel::Trading.permits(AccessLevel::Market));
assert!(AccessLevel::View.permits(AccessLevel::Market));
assert!(AccessLevel::View.permits(AccessLevel::View));
assert!(!AccessLevel::Market.permits(AccessLevel::View));
assert!(!AccessLevel::View.permits(AccessLevel::Trading));
assert_eq!(AccessLevel::default(), AccessLevel::Market);
}
#[test]
fn classifies_known_endpoints() {
let m = |p: &str| required_access_for("GET", p);
assert_eq!(m("/tickers"), AccessLevel::Market);
assert_eq!(m("/order-book/BTC-USD"), AccessLevel::Market);
assert_eq!(m("/public/last-trades"), AccessLevel::Market);
assert_eq!(m("/candles/BTC-USD"), AccessLevel::Market);
assert_eq!(m("/trades/all/BTC-USD"), AccessLevel::Market);
assert_eq!(m("/configuration/pairs"), AccessLevel::Market);
assert_eq!(m("/balances"), AccessLevel::View);
assert_eq!(m("/orders/active"), AccessLevel::View);
assert_eq!(m("/orders/fills/abc"), AccessLevel::View);
assert_eq!(m("/trades/private/BTC-USD"), AccessLevel::View);
assert_eq!(required_access_for("POST", "/orders"), AccessLevel::Trading);
assert_eq!(
required_access_for("DELETE", "/orders/abc"),
AccessLevel::Trading
);
}
#[test]
fn unknown_get_endpoint_fails_closed_to_view() {
assert_eq!(
required_access_for("GET", "/some/new/thing"),
AccessLevel::View
);
}
#[test]
fn denial_names_the_required_option() {
let msg = access_denied(AccessLevel::View, AccessLevel::Market);
assert!(msg.contains("--access view"));
assert!(msg.contains("--access market"));
}
}