use std::collections::BTreeSet;
use std::fmt;
use serde::{Deserialize, Deserializer, Serialize, Serializer, de};
use crate::ToolsetFormatError;
pub(crate) const CAPABILITY_KEY: &str = "stellar-agent-capabilities";
pub(crate) const RESERVED_PREFIX: &str = "stellar-agent-";
pub(crate) const SIGN_TRANSACTION_TOKEN: &str = "sign-transaction";
pub(crate) const SIGN_PAYMENT_TOKEN: &str = "sign-payment";
pub(crate) const SIGN_RULE_CREATE_TOKEN: &str = "sign-rule-create";
#[non_exhaustive]
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum Capability {
ReadBalance,
ProposeTransaction,
SuggestDestination,
ObserveEvent,
SignPayment,
ReadRules,
SignRuleCreate,
}
impl Capability {
#[must_use]
pub fn is_key_touching(self) -> bool {
match self {
Self::ReadBalance => false,
Self::ProposeTransaction => false,
Self::SuggestDestination => false,
Self::ObserveEvent => false,
Self::SignPayment => true,
Self::ReadRules => false,
Self::SignRuleCreate => true,
}
}
}
impl fmt::Display for Capability {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::ReadBalance => f.write_str("read-balance"),
Self::ProposeTransaction => f.write_str("propose-transaction"),
Self::SuggestDestination => f.write_str("suggest-destination"),
Self::ObserveEvent => f.write_str("observe-event"),
Self::SignPayment => f.write_str("sign-payment"),
Self::ReadRules => f.write_str("read-rules"),
Self::SignRuleCreate => f.write_str("sign-rule-create"),
}
}
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct CapabilitySet(BTreeSet<Capability>);
impl CapabilitySet {
#[must_use]
pub fn empty() -> Self {
Self(BTreeSet::new())
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
#[must_use]
pub fn len(&self) -> usize {
self.0.len()
}
#[must_use]
pub fn contains(&self, cap: Capability) -> bool {
self.0.contains(&cap)
}
pub fn iter(&self) -> impl Iterator<Item = Capability> + '_ {
self.0.iter().copied()
}
}
impl IntoIterator for CapabilitySet {
type Item = Capability;
type IntoIter = std::collections::btree_set::IntoIter<Capability>;
fn into_iter(self) -> Self::IntoIter {
self.0.into_iter()
}
}
impl Serialize for CapabilitySet {
fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
use serde::ser::SerializeSeq;
let mut seq = s.serialize_seq(Some(self.0.len()))?;
for cap in &self.0 {
seq.serialize_element(&cap.to_string())?;
}
seq.end()
}
}
impl<'de> Deserialize<'de> for CapabilitySet {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
struct CapSetVisitor;
impl<'de> de::Visitor<'de> for CapSetVisitor {
type Value = CapabilitySet;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("an array of capability token strings")
}
fn visit_seq<A: de::SeqAccess<'de>>(self, mut seq: A) -> Result<Self::Value, A::Error> {
let mut set = BTreeSet::new();
while let Some(token) = seq.next_element::<String>()? {
if !token.chars().all(is_valid_token_char) {
return Err(de::Error::custom(format!(
"capability token '{token}' contains characters outside [a-z0-9-]"
)));
}
match match_capability_token(&token) {
Ok(cap) => {
set.insert(cap);
}
Err(ToolsetFormatError::UnknownCapability { .. }) => {
}
Err(ToolsetFormatError::BareSignTransactionForbidden) => {
return Err(de::Error::custom(
"capability token 'sign-transaction' is forbidden in a \
stored record; the record is structurally malformed",
));
}
Err(other) => {
return Err(de::Error::custom(format!(
"invalid capability token '{token}': {other}"
)));
}
}
}
Ok(CapabilitySet(set))
}
}
d.deserialize_seq(CapSetVisitor)
}
}
pub(crate) fn parse_capability_value(value: &str) -> Result<CapabilitySet, ToolsetFormatError> {
let mut set = BTreeSet::new();
for token in value.split_ascii_whitespace() {
if !token.chars().all(is_valid_token_char) {
return Err(ToolsetFormatError::CapabilityTokenInvalidChar {
token: token.to_owned(),
});
}
let cap = match_capability_token(token)?;
set.insert(cap);
}
Ok(CapabilitySet(set))
}
#[inline]
pub(crate) fn is_valid_token_char(ch: char) -> bool {
ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '-'
}
#[cfg(any(test, feature = "test-helpers"))]
pub fn parse_capability_value_pub(value: &str) -> Result<CapabilitySet, ToolsetFormatError> {
parse_capability_value(value)
}
fn match_capability_token(token: &str) -> Result<Capability, ToolsetFormatError> {
match token {
SIGN_TRANSACTION_TOKEN => Err(ToolsetFormatError::BareSignTransactionForbidden),
"read-balance" => Ok(Capability::ReadBalance),
"propose-transaction" => Ok(Capability::ProposeTransaction),
"suggest-destination" => Ok(Capability::SuggestDestination),
"observe-event" => Ok(Capability::ObserveEvent),
SIGN_PAYMENT_TOKEN => Ok(Capability::SignPayment),
"read-rules" => Ok(Capability::ReadRules),
SIGN_RULE_CREATE_TOKEN => Ok(Capability::SignRuleCreate),
other => Err(ToolsetFormatError::UnknownCapability {
token: other.to_owned(),
}),
}
}
#[cfg(test)]
mod tests {
#![allow(
clippy::unwrap_used,
clippy::expect_used,
reason = "test-only; panics acceptable in unit tests"
)]
use super::*;
#[test]
fn empty_value_yields_empty_set() {
assert!(parse_capability_value("").unwrap().is_empty());
}
#[test]
fn whitespace_only_yields_empty_set() {
assert!(parse_capability_value(" \t ").unwrap().is_empty());
}
#[test]
fn all_taxonomy_tokens_parse() {
let set = parse_capability_value(
"read-balance propose-transaction suggest-destination observe-event sign-payment",
)
.unwrap();
assert!(set.contains(Capability::ReadBalance));
assert!(set.contains(Capability::ProposeTransaction));
assert!(set.contains(Capability::SuggestDestination));
assert!(set.contains(Capability::ObserveEvent));
assert!(set.contains(Capability::SignPayment));
assert_eq!(set.len(), 5);
}
#[test]
fn sign_payment_parses_to_sign_payment_capability() {
let set = parse_capability_value("sign-payment").unwrap();
assert!(set.contains(Capability::SignPayment));
assert_eq!(set.len(), 1);
}
#[test]
fn sign_rule_create_parses_to_sign_rule_create_capability() {
let set = parse_capability_value("sign-rule-create").unwrap();
assert!(set.contains(Capability::SignRuleCreate));
assert_eq!(set.len(), 1);
}
#[test]
fn sign_rule_create_display_roundtrip() {
assert_eq!(Capability::SignRuleCreate.to_string(), "sign-rule-create");
}
#[test]
fn sign_payment_display_roundtrip() {
assert_eq!(Capability::SignPayment.to_string(), "sign-payment");
}
#[test]
fn duplicate_tokens_deduplicate() {
let set = parse_capability_value("read-balance read-balance read-balance").unwrap();
assert_eq!(set.len(), 1);
}
#[test]
fn bare_sign_transaction_forbidden() {
let err = parse_capability_value("sign-transaction").unwrap_err();
assert!(
matches!(err, ToolsetFormatError::BareSignTransactionForbidden),
"expected BareSignTransactionForbidden, got {err:?}"
);
}
#[test]
fn sign_transaction_uppercase_refused_at_charset_gate() {
let err = parse_capability_value("Sign-Transaction").unwrap_err();
assert!(
matches!(err, ToolsetFormatError::CapabilityTokenInvalidChar { .. }),
"expected CapabilityTokenInvalidChar, got {err:?}"
);
}
#[test]
fn sign_transaction_all_caps_refused_at_charset_gate() {
let err = parse_capability_value("SIGN-TRANSACTION").unwrap_err();
assert!(
matches!(err, ToolsetFormatError::CapabilityTokenInvalidChar { .. }),
"expected CapabilityTokenInvalidChar, got {err:?}"
);
}
#[test]
fn sign_transaction_underscore_refused_at_charset_gate() {
let err = parse_capability_value("sign_transaction").unwrap_err();
assert!(
matches!(err, ToolsetFormatError::CapabilityTokenInvalidChar { .. }),
"expected CapabilityTokenInvalidChar, got {err:?}"
);
}
#[test]
fn sign_transaction_whitespace_padded_refused() {
let err = parse_capability_value(" sign-transaction ").unwrap_err();
assert!(
matches!(err, ToolsetFormatError::BareSignTransactionForbidden),
"expected BareSignTransactionForbidden, got {err:?}"
);
}
#[test]
fn sign_transaction_unicode_homoglyph_refused_at_charset_gate() {
let homoglyph = "ѕign-transaction"; let err = parse_capability_value(homoglyph).unwrap_err();
assert!(
matches!(err, ToolsetFormatError::CapabilityTokenInvalidChar { .. }),
"expected CapabilityTokenInvalidChar, got {err:?}"
);
}
#[test]
fn read_balance_homoglyph_refused_at_charset_gate() {
let homoglyph = "re\u{0430}d-balance"; let err = parse_capability_value(homoglyph).unwrap_err();
assert!(
matches!(err, ToolsetFormatError::CapabilityTokenInvalidChar { .. }),
"expected CapabilityTokenInvalidChar, got {err:?}"
);
}
#[test]
fn unknown_token_refused() {
let err = parse_capability_value("send-xdr").unwrap_err();
assert!(
matches!(err, ToolsetFormatError::UnknownCapability { .. }),
"expected UnknownCapability, got {err:?}"
);
}
#[test]
fn tab_in_token_refused_at_charset_gate() {
let s = "read\tbalance";
let err = parse_capability_value(s).unwrap_err();
assert!(
matches!(err, ToolsetFormatError::UnknownCapability { .. }),
"got {err:?}"
);
}
#[test]
fn deserialize_sign_transaction_is_serde_error() {
let json = r#"["sign-transaction"]"#;
let result: Result<CapabilitySet, _> = serde_json::from_str(json);
assert!(
result.is_err(),
"deserialising sign-transaction must be a serde error, not a silent drop"
);
let msg = result.unwrap_err().to_string();
assert!(
msg.contains("sign-transaction"),
"error message must mention the forbidden token: {msg}"
);
}
#[test]
fn deserialize_known_tokens_roundtrip() {
let json = r#"["read-balance","propose-transaction","suggest-destination","observe-event","sign-payment"]"#;
let set: CapabilitySet = serde_json::from_str(json).unwrap();
assert!(set.contains(Capability::ReadBalance));
assert!(set.contains(Capability::ProposeTransaction));
assert!(set.contains(Capability::SuggestDestination));
assert!(set.contains(Capability::ObserveEvent));
assert!(set.contains(Capability::SignPayment));
assert_eq!(set.len(), 5);
}
#[test]
fn deserialize_sign_payment_is_ok() {
let json = r#"["sign-payment"]"#;
let set: CapabilitySet = serde_json::from_str(json).unwrap();
assert!(
set.contains(Capability::SignPayment),
"sign-payment must deserialise to SignPayment variant"
);
}
#[test]
fn deserialize_sign_rule_create_is_ok() {
let json = r#"["sign-rule-create"]"#;
let set: CapabilitySet = serde_json::from_str(json).unwrap();
assert!(
set.contains(Capability::SignRuleCreate),
"sign-rule-create must deserialise to SignRuleCreate variant"
);
}
#[test]
fn deserialize_unknown_token_silently_skipped() {
let json = r#"["read-balance","future-unknown-cap"]"#;
let set: CapabilitySet = serde_json::from_str(json).unwrap();
assert!(set.contains(Capability::ReadBalance));
assert_eq!(set.len(), 1);
}
#[test]
fn deserialize_invalid_charset_token_is_serde_error() {
let json = r#"["read-balance","UPPERCASE"]"#;
let result: Result<CapabilitySet, _> = serde_json::from_str(json);
assert!(
result.is_err(),
"deserialising a token with invalid charset must be a serde error"
);
}
#[test]
fn serialize_deserialize_roundtrip() {
let set = parse_capability_value("read-balance propose-transaction").unwrap();
let json = serde_json::to_string(&set).unwrap();
let restored: CapabilitySet = serde_json::from_str(&json).unwrap();
assert_eq!(set, restored);
}
#[test]
fn capability_set_iter_is_sorted() {
let set = parse_capability_value("observe-event read-balance propose-transaction").unwrap();
let v: Vec<Capability> = set.iter().collect();
assert_eq!(v[0], Capability::ReadBalance);
assert_eq!(v[1], Capability::ProposeTransaction);
assert_eq!(v[2], Capability::ObserveEvent);
}
#[test]
fn into_iterator_consumes_set_by_value() {
let set =
parse_capability_value("read-balance propose-transaction suggest-destination").unwrap();
assert_eq!(set.len(), 3);
let mut seen = Vec::new();
for cap in set {
seen.push(cap);
}
assert_eq!(seen.len(), 3);
assert_eq!(seen[0], Capability::ReadBalance);
assert_eq!(seen[1], Capability::ProposeTransaction);
assert_eq!(seen[2], Capability::SuggestDestination);
}
#[test]
fn capability_display() {
assert_eq!(Capability::ReadBalance.to_string(), "read-balance");
assert_eq!(
Capability::ProposeTransaction.to_string(),
"propose-transaction"
);
assert_eq!(
Capability::SuggestDestination.to_string(),
"suggest-destination"
);
assert_eq!(Capability::ObserveEvent.to_string(), "observe-event");
assert_eq!(Capability::SignPayment.to_string(), "sign-payment");
assert_eq!(Capability::ReadRules.to_string(), "read-rules");
assert_eq!(Capability::SignRuleCreate.to_string(), "sign-rule-create");
}
#[test]
fn sign_payment_is_key_touching() {
assert!(
Capability::SignPayment.is_key_touching(),
"SignPayment must be key-touching (accesses signing key)"
);
}
#[test]
fn sign_rule_create_is_key_touching() {
assert!(
Capability::SignRuleCreate.is_key_touching(),
"SignRuleCreate must be key-touching (installs a rule via the signing key)"
);
}
#[test]
fn non_signing_capabilities_are_not_key_touching() {
let non_key_touching = [
Capability::ReadBalance,
Capability::ProposeTransaction,
Capability::SuggestDestination,
Capability::ObserveEvent,
];
for cap in non_key_touching {
assert!(
!cap.is_key_touching(),
"{cap} must NOT be key-touching (does not access signing key)"
);
}
}
#[test]
fn is_key_touching_exhaustive_table() {
let table: &[(Capability, bool)] = &[
(Capability::ReadBalance, false),
(Capability::ProposeTransaction, false),
(Capability::SuggestDestination, false),
(Capability::ObserveEvent, false),
(Capability::SignPayment, true),
(Capability::ReadRules, false),
(Capability::SignRuleCreate, true),
];
for (cap, expected) in table {
assert_eq!(
cap.is_key_touching(),
*expected,
"is_key_touching({cap}) should be {expected}"
);
}
}
}