use async_trait::async_trait;
use std::collections::BTreeMap;
use std::sync::{Arc, Mutex};
use super::adapter::{ChoiceLabel, ChoreographicAdapter, Message};
use crate::effects::{ChoreographyError, RoleId};
type MessageQueues = BTreeMap<(String, String), Vec<Vec<u8>>>;
pub struct TestAdapter<R: RoleId> {
role: R,
families: BTreeMap<String, Vec<R>>,
sent: Arc<Mutex<MessageQueues>>,
received: Arc<Mutex<MessageQueues>>,
}
impl<R: RoleId> TestAdapter<R> {
pub fn new(role: R) -> Self {
Self {
role,
families: BTreeMap::new(),
sent: Arc::new(Mutex::new(BTreeMap::new())),
received: Arc::new(Mutex::new(BTreeMap::new())),
}
}
pub fn with_family(mut self, name: &str, instances: Vec<R>) -> Self {
self.families.insert(name.to_string(), instances);
self
}
pub fn add_family(&mut self, name: &str, instances: Vec<R>) {
self.families.insert(name.to_string(), instances);
}
pub fn role(&self) -> R {
self.role
}
pub fn sent_messages(&self) -> MessageQueues {
self.sent.lock().unwrap().clone()
}
pub fn queue_message(&self, from: R, to: R, message: Vec<u8>) {
let key = (from.role_name().to_string(), to.role_name().to_string());
let mut received = self.received.lock().unwrap();
received.entry(key).or_default().push(message);
}
pub fn queue_typed_message<M: Message>(&self, from: R, to: R, message: M) {
let bytes = bincode::serialize(&message).expect("serialization should succeed");
self.queue_message(from, to, bytes);
}
pub fn linked_pair(role_a: R, role_b: R) -> (Self, Self) {
let shared_a_to_b = Arc::new(Mutex::new(BTreeMap::new()));
let shared_b_to_a = Arc::new(Mutex::new(BTreeMap::new()));
let adapter_a = Self {
role: role_a,
families: BTreeMap::new(),
sent: shared_a_to_b.clone(),
received: shared_b_to_a.clone(),
};
let adapter_b = Self {
role: role_b,
families: BTreeMap::new(),
sent: shared_b_to_a,
received: shared_a_to_b,
};
(adapter_a, adapter_b)
}
}
#[derive(Debug, thiserror::Error)]
pub enum TestAdapterError {
#[error("role family '{0}' not configured")]
FamilyNotConfigured(String),
#[error("invalid role range for '{family}': [{start}, {end})")]
InvalidRange {
family: String,
start: u32,
end: u32,
},
#[error("no message available from {from} to {to}")]
NoMessageAvailable { from: String, to: String },
#[error("serialization error: {0}")]
Serialization(String),
#[error("role family '{0}' resolved to empty set")]
EmptyFamily(String),
}
impl From<TestAdapterError> for ChoreographyError {
fn from(err: TestAdapterError) -> Self {
match err {
TestAdapterError::FamilyNotConfigured(name) => {
ChoreographyError::RoleFamilyNotFound(name)
}
TestAdapterError::InvalidRange { family, start, end } => {
ChoreographyError::InvalidRoleRange { family, start, end }
}
TestAdapterError::NoMessageAvailable { from, to } => {
ChoreographyError::Transport(format!("no message from {} to {}", from, to))
}
TestAdapterError::Serialization(msg) => ChoreographyError::Serialization(msg),
TestAdapterError::EmptyFamily(name) => ChoreographyError::EmptyRoleFamily(name),
}
}
}
#[async_trait]
impl<R: RoleId + 'static> ChoreographicAdapter for TestAdapter<R> {
type Error = ChoreographyError;
type Role = R;
async fn send<M: Message>(&mut self, to: Self::Role, msg: M) -> Result<(), Self::Error> {
let bytes = bincode::serialize(&msg)
.map_err(|e| ChoreographyError::Serialization(e.to_string()))?;
let key = (
self.role.role_name().to_string(),
to.role_name().to_string(),
);
let mut sent = self.sent.lock().unwrap();
sent.entry(key).or_default().push(bytes);
Ok(())
}
async fn recv<M: Message>(&mut self, from: Self::Role) -> Result<M, Self::Error> {
let key = (
from.role_name().to_string(),
self.role.role_name().to_string(),
);
let bytes = {
let mut received = self.received.lock().unwrap();
let queue = received.entry(key.clone()).or_default();
if queue.is_empty() {
return Err(TestAdapterError::NoMessageAvailable {
from: key.0,
to: key.1,
}
.into());
}
queue.remove(0)
};
bincode::deserialize(&bytes).map_err(|e| ChoreographyError::Serialization(e.to_string()))
}
async fn choose(
&mut self,
to: Self::Role,
label: <Self::Role as RoleId>::Label,
) -> Result<(), Self::Error> {
self.send(to, ChoiceLabel(label)).await
}
async fn offer(
&mut self,
from: Self::Role,
) -> Result<<Self::Role as RoleId>::Label, Self::Error> {
let choice: ChoiceLabel<<Self::Role as RoleId>::Label> = self.recv(from).await?;
Ok(choice.0)
}
fn resolve_family(&self, family: &str) -> Result<Vec<Self::Role>, Self::Error> {
self.families
.get(family)
.cloned()
.ok_or_else(|| ChoreographyError::RoleFamilyNotFound(family.to_string()))
}
fn resolve_range(
&self,
family: &str,
start: u32,
end: u32,
) -> Result<Vec<Self::Role>, Self::Error> {
let instances = self.resolve_family(family)?;
if start >= end {
return Err(ChoreographyError::InvalidRoleRange {
family: family.to_string(),
start,
end,
});
}
let start_idx = start as usize;
let end_idx = (end as usize).min(instances.len());
if start_idx >= instances.len() {
return Err(ChoreographyError::InvalidRoleRange {
family: family.to_string(),
start,
end,
});
}
Ok(instances[start_idx..end_idx].to_vec())
}
}
#[cfg(all(test, not(target_arch = "wasm32")))]
mod tests {
use super::*;
use crate::effects::LabelId;
use crate::identifiers::RoleName;
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
enum TestRole {
Coordinator,
Witness(u32),
}
#[derive(Copy, Clone, Debug, PartialEq, Eq, Hash)]
enum TestLabel {
Commit,
Abort,
}
impl LabelId for TestLabel {
fn as_str(&self) -> &'static str {
match self {
TestLabel::Commit => "Commit",
TestLabel::Abort => "Abort",
}
}
fn from_str(label: &str) -> Option<Self> {
match label {
"Commit" => Some(TestLabel::Commit),
"Abort" => Some(TestLabel::Abort),
_ => None,
}
}
}
impl RoleId for TestRole {
type Label = TestLabel;
fn role_name(&self) -> RoleName {
match self {
TestRole::Coordinator => RoleName::from_static("Coordinator"),
TestRole::Witness(_) => RoleName::from_static("Witness"),
}
}
fn role_index(&self) -> Option<u32> {
match self {
TestRole::Witness(index) => Some(*index),
_ => None,
}
}
}
#[test]
fn test_resolve_family() {
let adapter = TestAdapter::new(TestRole::Coordinator).with_family(
"Witness",
vec![
TestRole::Witness(0),
TestRole::Witness(1),
TestRole::Witness(2),
],
);
let witnesses = adapter.resolve_family("Witness").unwrap();
assert_eq!(witnesses.len(), 3);
assert_eq!(witnesses[0], TestRole::Witness(0));
assert_eq!(witnesses[1], TestRole::Witness(1));
assert_eq!(witnesses[2], TestRole::Witness(2));
}
#[test]
fn test_resolve_family_not_found() {
let adapter = TestAdapter::new(TestRole::Coordinator);
let result = adapter.resolve_family("Unknown");
assert!(result.is_err());
}
#[test]
fn test_resolve_range() {
let adapter = TestAdapter::new(TestRole::Coordinator).with_family(
"Witness",
vec![
TestRole::Witness(0),
TestRole::Witness(1),
TestRole::Witness(2),
TestRole::Witness(3),
TestRole::Witness(4),
],
);
let witnesses = adapter.resolve_range("Witness", 1, 4).unwrap();
assert_eq!(witnesses.len(), 3);
assert_eq!(witnesses[0], TestRole::Witness(1));
assert_eq!(witnesses[1], TestRole::Witness(2));
assert_eq!(witnesses[2], TestRole::Witness(3));
}
#[test]
fn test_resolve_range_clamps_to_bounds() {
let adapter = TestAdapter::new(TestRole::Coordinator).with_family(
"Witness",
vec![
TestRole::Witness(0),
TestRole::Witness(1),
TestRole::Witness(2),
],
);
let witnesses = adapter.resolve_range("Witness", 1, 10).unwrap();
assert_eq!(witnesses.len(), 2);
assert_eq!(witnesses[0], TestRole::Witness(1));
assert_eq!(witnesses[1], TestRole::Witness(2));
}
#[test]
fn test_resolve_range_invalid() {
let adapter = TestAdapter::new(TestRole::Coordinator)
.with_family("Witness", vec![TestRole::Witness(0)]);
let result = adapter.resolve_range("Witness", 5, 3);
assert!(result.is_err());
let result = adapter.resolve_range("Witness", 10, 15);
assert!(result.is_err());
}
#[test]
fn test_family_size() {
let adapter = TestAdapter::new(TestRole::Coordinator).with_family(
"Witness",
vec![
TestRole::Witness(0),
TestRole::Witness(1),
TestRole::Witness(2),
],
);
assert_eq!(adapter.family_size("Witness").unwrap(), 3);
}
#[tokio::test]
async fn test_send_recv() {
let (mut coordinator, mut witness) =
TestAdapter::linked_pair(TestRole::Coordinator, TestRole::Witness(0));
coordinator
.send(TestRole::Witness(0), "hello".to_string())
.await
.unwrap();
let msg: String = witness.recv(TestRole::Coordinator).await.unwrap();
assert_eq!(msg, "hello");
}
#[tokio::test]
async fn test_broadcast() {
let mut adapter = TestAdapter::new(TestRole::Coordinator).with_family(
"Witness",
vec![
TestRole::Witness(0),
TestRole::Witness(1),
TestRole::Witness(2),
],
);
let witnesses = adapter.resolve_family("Witness").unwrap();
adapter
.broadcast(&witnesses, "broadcast message".to_string())
.await
.unwrap();
let sent = adapter.sent_messages();
let key = ("Coordinator".to_string(), "Witness".to_string());
let messages = sent.get(&key).unwrap();
assert_eq!(messages.len(), 3);
}
}