use serde::{Deserialize, Serialize};
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClientChannelFunding {
pub params_json: String,
pub funding_proofs_json: String,
pub channel_secret_hex: String,
pub keyset_info_json: String,
pub sender_pubkey_hex: String,
pub capacity: u64,
pub funding_token_amount: u64,
pub mint_url: String,
pub created_at: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ClientPaymentState {
pub balance: u64,
pub signature: String,
pub payment_count: u64,
pub last_payment_at: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
pub enum ClientChannelState {
#[default]
Open,
Closed,
}
pub trait ClientStorage {
fn save_funding(&mut self, channel_id: &str, funding: ClientChannelFunding);
fn get_funding(&self, channel_id: &str) -> Option<&ClientChannelFunding>;
fn get_payment_state(&self, channel_id: &str) -> Option<&ClientPaymentState>;
fn save_payment_state(&mut self, channel_id: &str, state: ClientPaymentState);
fn get_state(&self, channel_id: &str) -> ClientChannelState;
fn set_closed(&mut self, channel_id: &str);
fn list_channel_ids(&self) -> Vec<String>;
fn delete(&mut self, channel_id: &str);
}
#[derive(Debug, Default)]
pub struct MemoryClientStorage {
funding: HashMap<String, ClientChannelFunding>,
payments: HashMap<String, ClientPaymentState>,
closed: HashSet<String>,
}
impl MemoryClientStorage {
pub fn new() -> Self {
Self::default()
}
pub fn channel_count(&self) -> usize {
self.funding.len()
}
}
impl ClientStorage for MemoryClientStorage {
fn save_funding(&mut self, channel_id: &str, funding: ClientChannelFunding) {
self.funding.insert(channel_id.to_string(), funding);
}
fn get_funding(&self, channel_id: &str) -> Option<&ClientChannelFunding> {
self.funding.get(channel_id)
}
fn get_payment_state(&self, channel_id: &str) -> Option<&ClientPaymentState> {
self.payments.get(channel_id)
}
fn save_payment_state(&mut self, channel_id: &str, state: ClientPaymentState) {
self.payments.insert(channel_id.to_string(), state);
}
fn get_state(&self, channel_id: &str) -> ClientChannelState {
if self.closed.contains(channel_id) {
ClientChannelState::Closed
} else if self.funding.contains_key(channel_id) {
ClientChannelState::Open
} else {
ClientChannelState::Closed
}
}
fn set_closed(&mut self, channel_id: &str) {
self.closed.insert(channel_id.to_string());
}
fn list_channel_ids(&self) -> Vec<String> {
self.funding.keys().cloned().collect()
}
fn delete(&mut self, channel_id: &str) {
self.funding.remove(channel_id);
self.payments.remove(channel_id);
self.closed.remove(channel_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_test_funding() -> ClientChannelFunding {
ClientChannelFunding {
params_json: r#"{"test": true}"#.to_string(),
funding_proofs_json: "[]".to_string(),
channel_secret_hex: "aa".repeat(32),
keyset_info_json: "{}".to_string(),
sender_pubkey_hex: "02".to_string() + &"bb".repeat(32),
capacity: 1000,
funding_token_amount: 1100,
mint_url: "https://mint.example.com".to_string(),
created_at: 1234567890,
}
}
fn make_test_payment_state(balance: u64) -> ClientPaymentState {
ClientPaymentState {
balance,
signature: "sig".to_string(),
payment_count: 1,
last_payment_at: 1234567890,
}
}
#[test]
fn test_memory_storage_funding() {
let mut storage = MemoryClientStorage::new();
let channel_id = "test_channel_1";
assert!(storage.get_funding(channel_id).is_none());
assert_eq!(storage.channel_count(), 0);
storage.save_funding(channel_id, make_test_funding());
let funding = storage.get_funding(channel_id).unwrap();
assert_eq!(funding.capacity, 1000);
assert_eq!(storage.channel_count(), 1);
assert_eq!(storage.get_state(channel_id), ClientChannelState::Open);
}
#[test]
fn test_memory_storage_payments() {
let mut storage = MemoryClientStorage::new();
let channel_id = "test_channel_1";
storage.save_funding(channel_id, make_test_funding());
assert!(storage.get_payment_state(channel_id).is_none());
storage.save_payment_state(channel_id, make_test_payment_state(100));
let state = storage.get_payment_state(channel_id).unwrap();
assert_eq!(state.balance, 100);
assert_eq!(state.payment_count, 1);
storage.save_payment_state(channel_id, make_test_payment_state(200));
let state = storage.get_payment_state(channel_id).unwrap();
assert_eq!(state.balance, 200);
}
#[test]
fn test_memory_storage_lifecycle() {
let mut storage = MemoryClientStorage::new();
let channel_id = "test_channel_1";
assert_eq!(storage.get_state(channel_id), ClientChannelState::Closed);
storage.save_funding(channel_id, make_test_funding());
assert_eq!(storage.get_state(channel_id), ClientChannelState::Open);
storage.set_closed(channel_id);
assert_eq!(storage.get_state(channel_id), ClientChannelState::Closed);
}
#[test]
fn test_memory_storage_delete() {
let mut storage = MemoryClientStorage::new();
let channel_id = "test_channel_1";
storage.save_funding(channel_id, make_test_funding());
storage.save_payment_state(channel_id, make_test_payment_state(100));
storage.set_closed(channel_id);
assert_eq!(storage.channel_count(), 1);
storage.delete(channel_id);
assert_eq!(storage.channel_count(), 0);
assert!(storage.get_funding(channel_id).is_none());
assert!(storage.get_payment_state(channel_id).is_none());
assert_eq!(storage.get_state(channel_id), ClientChannelState::Closed);
}
#[test]
fn test_memory_storage_list() {
let mut storage = MemoryClientStorage::new();
storage.save_funding("channel_1", make_test_funding());
storage.save_funding("channel_2", make_test_funding());
storage.save_funding("channel_3", make_test_funding());
let mut ids = storage.list_channel_ids();
ids.sort();
assert_eq!(ids, vec!["channel_1", "channel_2", "channel_3"]);
}
}