use proc_macro2::TokenStream;
use quote::{format_ident, quote};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum BroadcastMode {
#[default]
Parallel,
Sequential,
}
impl BroadcastMode {
#[must_use]
pub fn from_annotation(value: &str) -> Option<Self> {
match value.to_lowercase().as_str() {
"parallel" | "concurrent" => Some(Self::Parallel),
"sequential" | "ordered" | "seq" => Some(Self::Sequential),
_ => None,
}
}
#[must_use]
pub fn annotation_name(&self) -> &'static str {
match self {
Self::Parallel => "parallel",
Self::Sequential => "sequential",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CollectionMode {
#[default]
Ordered,
Unordered,
}
impl CollectionMode {
#[must_use]
pub fn from_annotation(value: &str) -> Option<Self> {
match value.to_lowercase().as_str() {
"ordered" | "sequential" | "ordered_collection" => Some(Self::Ordered),
"unordered" | "parallel" | "any_order" => Some(Self::Unordered),
_ => None,
}
}
#[must_use]
pub fn annotation_name(&self) -> &'static str {
match self {
Self::Ordered => "ordered",
Self::Unordered => "unordered",
}
}
}
#[derive(Debug, Clone)]
pub struct BatchConfig {
pub messages: Vec<String>,
pub auto_unwrap: bool,
}
impl Default for BatchConfig {
fn default() -> Self {
Self {
messages: Vec::new(),
auto_unwrap: true,
}
}
}
pub fn generate_parallel_broadcast(
adapter_name: &str,
recipients_name: &str,
payload_name: &str,
) -> TokenStream {
let adapter = format_ident!("{}", adapter_name);
let recipients = format_ident!("{}", recipients_name);
let payload = format_ident!("{}", payload_name);
quote! {
futures::future::try_join_all(
#recipients.iter().map(|r| #adapter.send(*r, #payload.clone()))
).await?;
}
}
pub fn generate_sequential_broadcast(
adapter_name: &str,
recipients_name: &str,
payload_name: &str,
) -> TokenStream {
let adapter = format_ident!("{}", adapter_name);
let recipients = format_ident!("{}", recipients_name);
let payload = format_ident!("{}", payload_name);
quote! {
for recipient in &#recipients {
#adapter.send(*recipient, #payload.clone()).await?;
}
}
}
pub fn generate_broadcast(
mode: BroadcastMode,
adapter_name: &str,
recipients_name: &str,
payload_name: &str,
) -> TokenStream {
match mode {
BroadcastMode::Parallel => {
generate_parallel_broadcast(adapter_name, recipients_name, payload_name)
}
BroadcastMode::Sequential => {
generate_sequential_broadcast(adapter_name, recipients_name, payload_name)
}
}
}
pub fn generate_ordered_collection(
adapter_name: &str,
senders_name: &str,
response_type: &str,
result_name: &str,
) -> TokenStream {
let adapter = format_ident!("{}", adapter_name);
let senders = format_ident!("{}", senders_name);
let response_ty = format_ident!("{}", response_type);
let result = format_ident!("{}", result_name);
quote! {
let mut #result = Vec::with_capacity(#senders.len());
for sender in &#senders {
let msg = #adapter.recv::<#response_ty>(*sender).await?;
#result.push(msg);
}
}
}
pub fn generate_unordered_collection(
adapter_name: &str,
senders_name: &str,
response_type: &str,
result_name: &str,
) -> TokenStream {
let adapter = format_ident!("{}", adapter_name);
let senders = format_ident!("{}", senders_name);
let response_ty = format_ident!("{}", response_type);
let result = format_ident!("{}", result_name);
quote! {
use futures::future::select_all;
let mut pending: Vec<_> = #senders.iter()
.map(|s| Box::pin(#adapter.recv::<#response_ty>(*s)))
.collect();
let mut #result = Vec::with_capacity(pending.len());
while !pending.is_empty() {
let (recv_result, _idx, remaining) = select_all(pending).await;
#result.push(recv_result?);
pending = remaining;
}
}
}
pub fn generate_collection(
mode: CollectionMode,
adapter_name: &str,
senders_name: &str,
response_type: &str,
result_name: &str,
) -> TokenStream {
match mode {
CollectionMode::Ordered => {
generate_ordered_collection(adapter_name, senders_name, response_type, result_name)
}
CollectionMode::Unordered => {
generate_unordered_collection(adapter_name, senders_name, response_type, result_name)
}
}
}
pub fn generate_batch_send(
adapter_name: &str,
recipient_name: &str,
message_names: &[(&str, &str)], ) -> TokenStream {
let adapter = format_ident!("{}", adapter_name);
let recipient = format_ident!("{}", recipient_name);
let field_assignments: Vec<TokenStream> = message_names
.iter()
.map(|(field, value)| {
let field_ident = format_ident!("{}", field);
let value_ident = format_ident!("{}", value);
quote! { #field_ident: #value_ident }
})
.collect();
quote! {
let batch_payload = BatchPayload {
#(#field_assignments),*
};
#adapter.send(#recipient, batch_payload).await?;
}
}
pub fn generate_batch_recv(
adapter_name: &str,
sender_name: &str,
batch_type: &str,
field_names: &[&str],
) -> TokenStream {
let adapter = format_ident!("{}", adapter_name);
let sender = format_ident!("{}", sender_name);
let batch_ty = format_ident!("{}", batch_type);
let field_extractions: Vec<TokenStream> = field_names
.iter()
.map(|field| {
let field_ident = format_ident!("{}", field);
quote! { let #field_ident = batch.#field_ident; }
})
.collect();
quote! {
let batch: #batch_ty = #adapter.recv(#sender).await?;
#(#field_extractions)*
}
}
#[derive(Debug, Clone, Default)]
pub struct ProtocolConcurrencyConfig {
pub default_broadcast: BroadcastMode,
pub default_collection: CollectionMode,
}
impl ProtocolConcurrencyConfig {
pub fn from_annotations(annotations: &std::collections::HashMap<String, String>) -> Self {
let mut config = Self::default();
if let Some(broadcast) = annotations.get("default_broadcast") {
if let Some(mode) = BroadcastMode::from_annotation(broadcast) {
config.default_broadcast = mode;
}
}
if let Some(collection) = annotations.get("default_collection") {
if let Some(mode) = CollectionMode::from_annotation(collection) {
config.default_collection = mode;
}
}
config
}
}
#[derive(Debug, Clone, Default)]
pub struct StatementConcurrencyConfig {
pub broadcast_mode: Option<BroadcastMode>,
pub collection_mode: Option<CollectionMode>,
pub is_batch: bool,
}
impl StatementConcurrencyConfig {
pub fn from_annotations(annotations: &std::collections::HashMap<String, String>) -> Self {
let mut config = Self::default();
if annotations.contains_key("sequential") {
config.broadcast_mode = Some(BroadcastMode::Sequential);
} else if annotations.contains_key("parallel") {
config.broadcast_mode = Some(BroadcastMode::Parallel);
}
if annotations.contains_key("unordered") {
config.collection_mode = Some(CollectionMode::Unordered);
} else if annotations.contains_key("ordered") {
config.collection_mode = Some(CollectionMode::Ordered);
}
if annotations.contains_key("batch") {
config.is_batch = true;
}
config
}
#[must_use]
pub fn effective_broadcast_mode(
&self,
protocol_config: &ProtocolConcurrencyConfig,
) -> BroadcastMode {
self.broadcast_mode
.unwrap_or(protocol_config.default_broadcast)
}
#[must_use]
pub fn effective_collection_mode(
&self,
protocol_config: &ProtocolConcurrencyConfig,
) -> CollectionMode {
self.collection_mode
.unwrap_or(protocol_config.default_collection)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_broadcast_mode_parsing() {
assert_eq!(
BroadcastMode::from_annotation("parallel"),
Some(BroadcastMode::Parallel)
);
assert_eq!(
BroadcastMode::from_annotation("sequential"),
Some(BroadcastMode::Sequential)
);
assert_eq!(
BroadcastMode::from_annotation("concurrent"),
Some(BroadcastMode::Parallel)
);
assert_eq!(
BroadcastMode::from_annotation("seq"),
Some(BroadcastMode::Sequential)
);
assert_eq!(BroadcastMode::from_annotation("invalid"), None);
}
#[test]
fn test_collection_mode_parsing() {
assert_eq!(
CollectionMode::from_annotation("ordered"),
Some(CollectionMode::Ordered)
);
assert_eq!(
CollectionMode::from_annotation("unordered"),
Some(CollectionMode::Unordered)
);
assert_eq!(
CollectionMode::from_annotation("parallel"),
Some(CollectionMode::Unordered)
);
assert_eq!(CollectionMode::from_annotation("invalid"), None);
}
#[test]
fn test_generate_parallel_broadcast() {
let code = generate_parallel_broadcast("adapter", "witnesses", "msg");
let code_str = code.to_string();
assert!(code_str.contains("try_join_all"));
assert!(code_str.contains("witnesses"));
assert!(code_str.contains("adapter"));
}
#[test]
fn test_generate_sequential_broadcast() {
let code = generate_sequential_broadcast("adapter", "witnesses", "msg");
let code_str = code.to_string();
assert!(code_str.contains("for recipient"));
assert!(code_str.contains("witnesses"));
}
#[test]
fn test_generate_ordered_collection() {
let code = generate_ordered_collection("adapter", "senders", "Response", "responses");
let code_str = code.to_string();
assert!(code_str.contains("for sender"));
assert!(code_str.contains("recv"));
assert!(code_str.contains("Response"));
}
#[test]
fn test_generate_unordered_collection() {
let code = generate_unordered_collection("adapter", "senders", "Response", "responses");
let code_str = code.to_string();
assert!(code_str.contains("select_all"));
assert!(code_str.contains("while"));
}
#[test]
fn test_protocol_config_from_annotations() {
let mut annotations = std::collections::HashMap::new();
annotations.insert("default_broadcast".to_string(), "sequential".to_string());
annotations.insert("default_collection".to_string(), "unordered".to_string());
let config = ProtocolConcurrencyConfig::from_annotations(&annotations);
assert_eq!(config.default_broadcast, BroadcastMode::Sequential);
assert_eq!(config.default_collection, CollectionMode::Unordered);
}
#[test]
fn test_statement_config_from_annotations() {
let mut annotations = std::collections::HashMap::new();
annotations.insert("sequential".to_string(), "true".to_string());
annotations.insert("batch".to_string(), "true".to_string());
let config = StatementConcurrencyConfig::from_annotations(&annotations);
assert_eq!(config.broadcast_mode, Some(BroadcastMode::Sequential));
assert!(config.is_batch);
}
#[test]
fn test_effective_modes() {
let protocol_config = ProtocolConcurrencyConfig {
default_broadcast: BroadcastMode::Sequential,
default_collection: CollectionMode::Unordered,
};
let stmt_config = StatementConcurrencyConfig::default();
assert_eq!(
stmt_config.effective_broadcast_mode(&protocol_config),
BroadcastMode::Sequential
);
assert_eq!(
stmt_config.effective_collection_mode(&protocol_config),
CollectionMode::Unordered
);
let override_config = StatementConcurrencyConfig {
broadcast_mode: Some(BroadcastMode::Parallel),
..Default::default()
};
assert_eq!(
override_config.effective_broadcast_mode(&protocol_config),
BroadcastMode::Parallel
);
}
#[test]
fn test_batch_send_generation() {
let code = generate_batch_send(
"adapter",
"recipient",
&[("msg_a", "a_value"), ("msg_b", "b_value")],
);
let code_str = code.to_string();
assert!(code_str.contains("BatchPayload"));
assert!(code_str.contains("msg_a"));
assert!(code_str.contains("msg_b"));
}
#[test]
fn test_batch_recv_generation() {
let code = generate_batch_recv("adapter", "sender", "BatchPayload", &["msg_a", "msg_b"]);
let code_str = code.to_string();
assert!(
code_str.contains("BatchPayload"),
"Should contain BatchPayload: {}",
code_str
);
assert!(
code_str.contains("batch") && code_str.contains("msg_a"),
"Should reference batch.msg_a: {}",
code_str
);
assert!(
code_str.contains("batch") && code_str.contains("msg_b"),
"Should reference batch.msg_b: {}",
code_str
);
}
}