use crate::common::MessageParser;
use crate::common::error::{FlareError, Result};
use crate::common::protocol::Frame;
use crate::transport::events::ConnectionEvent;
use async_trait::async_trait;
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Clone)]
pub struct MessageContext {
pub frame: Frame,
pub connection_id: Option<String>,
pub parser: MessageParser,
pub metadata: Arc<RwLock<std::collections::HashMap<String, Vec<u8>>>>,
}
impl MessageContext {
pub fn new(frame: Frame, connection_id: Option<String>, parser: MessageParser) -> Self {
Self {
frame,
connection_id,
parser,
metadata: Arc::new(RwLock::new(std::collections::HashMap::new())),
}
}
pub async fn set_metadata(&self, key: String, value: Vec<u8>) {
let mut meta = self.metadata.write().await;
meta.insert(key, value);
}
pub async fn get_metadata(&self, key: &str) -> Option<Vec<u8>> {
let meta = self.metadata.read().await;
meta.get(key).cloned()
}
}
#[async_trait]
pub trait MessageMiddleware: Send + Sync {
async fn before(&self, ctx: &MessageContext) -> Result<Option<Frame>> {
let _ = ctx;
Ok(None)
}
async fn after(&self, ctx: &MessageContext, response: Option<Frame>) -> Result<Option<Frame>> {
let _ = (ctx, response);
Ok(None)
}
fn name(&self) -> &str {
"UnknownMiddleware"
}
fn priority(&self) -> u32 {
100
}
}
pub type ArcMessageMiddleware = Arc<dyn MessageMiddleware>;
#[async_trait]
pub trait MessageProcessor: Send + Sync {
async fn process(&self, ctx: &MessageContext) -> Result<Option<Frame>>;
fn name(&self) -> &str {
"UnknownProcessor"
}
}
pub type ArcMessageProcessor = Arc<dyn MessageProcessor>;
#[derive(Clone)]
pub struct MessagePipeline {
middlewares: Arc<RwLock<Arc<Vec<ArcMessageMiddleware>>>>,
processors: Arc<RwLock<Arc<Vec<ArcMessageProcessor>>>>,
parser: Arc<tokio::sync::Mutex<MessageParser>>,
}
impl MessagePipeline {
pub fn new(parser: MessageParser) -> Self {
Self {
middlewares: Arc::new(RwLock::new(Arc::new(Vec::new()))),
processors: Arc::new(RwLock::new(Arc::new(Vec::new()))),
parser: Arc::new(tokio::sync::Mutex::new(parser)),
}
}
pub async fn update_parser(&self, parser: MessageParser) {
let mut p = self.parser.lock().await;
*p = parser;
}
pub async fn add_middleware(&self, middleware: ArcMessageMiddleware) {
let mut middlewares = self.middlewares.write().await;
let mut next = (**middlewares).clone();
next.push(middleware);
next.sort_by_key(|m| m.priority());
*middlewares = Arc::new(next);
}
pub async fn remove_middleware(&self, middleware: &ArcMessageMiddleware) {
let mut middlewares = self.middlewares.write().await;
let mut next = (**middlewares).clone();
next.retain(|m| !Arc::ptr_eq(m, middleware));
*middlewares = Arc::new(next);
}
pub async fn add_processor(&self, processor: ArcMessageProcessor) {
let mut processors = self.processors.write().await;
let mut next = (**processors).clone();
next.push(processor);
*processors = Arc::new(next);
}
pub async fn remove_processor(&self, processor: &ArcMessageProcessor) {
let mut processors = self.processors.write().await;
let mut next = (**processors).clone();
next.retain(|p| !Arc::ptr_eq(p, processor));
*processors = Arc::new(next);
}
pub async fn process_raw(
&self,
data: &[u8],
connection_id: Option<&str>,
) -> Result<Option<Vec<u8>>> {
let parser = self.parser.lock().await;
let frame = parser.parse(data).map_err(|e| {
FlareError::deserialization_error(format!("Failed to parse message: {}", e))
})?;
let parser_snapshot = parser.clone();
drop(parser);
let response = self
.process_frame_with_parser(&frame, connection_id, parser_snapshot.clone())
.await?;
if let Some(response_frame) = response {
let response_data = parser_snapshot.serialize(&response_frame).map_err(|e| {
FlareError::encoding_error(format!("Failed to serialize response: {}", e))
})?;
Ok(Some(response_data))
} else {
Ok(None)
}
}
async fn middleware_snapshot(&self) -> Arc<Vec<ArcMessageMiddleware>> {
let middlewares = self.middlewares.read().await;
Arc::clone(&middlewares)
}
async fn processor_snapshot(&self) -> Arc<Vec<ArcMessageProcessor>> {
let processors = self.processors.read().await;
Arc::clone(&processors)
}
pub async fn process_frame(
&self,
frame: &Frame,
connection_id: Option<&str>,
) -> Result<Option<Frame>> {
let parser = self.parser.lock().await;
let parser_snapshot = parser.clone();
drop(parser);
self.process_frame_with_parser(frame, connection_id, parser_snapshot)
.await
}
async fn process_frame_with_parser(
&self,
frame: &Frame,
connection_id: Option<&str>,
parser: MessageParser,
) -> Result<Option<Frame>> {
let ctx = MessageContext::new(frame.clone(), connection_id.map(|s| s.to_string()), parser);
let middlewares = self.middleware_snapshot().await;
for middleware in middlewares.iter() {
if let Some(response) = middleware.before(&ctx).await? {
return Ok(Some(response));
}
}
let processors = self.processor_snapshot().await;
let mut response = None;
for processor in processors.iter() {
if let Some(resp) = processor.process(&ctx).await? {
response = Some(resp);
break; }
}
let middlewares = self.middleware_snapshot().await;
for middleware in middlewares.iter() {
if let Some(modified_response) = middleware.after(&ctx, response.clone()).await? {
response = Some(modified_response);
}
}
Ok(response)
}
pub async fn handle_connection_event(
&self,
_event: &ConnectionEvent,
_connection_id: Option<&str>,
) -> Result<()> {
let middlewares = self.middleware_snapshot().await;
for _middleware in middlewares.iter() {
}
Ok(())
}
}
impl Default for MessagePipeline {
fn default() -> Self {
Self::new(MessageParser::protobuf())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::common::protocol::{
Command, FrameBuilder, Reliability, SerializationFormat, frame_with_system_command, ping,
pong,
};
use tokio::sync::{Mutex, oneshot};
struct BlockingMiddleware {
entered_tx: Mutex<Option<oneshot::Sender<()>>>,
release_rx: Mutex<Option<oneshot::Receiver<()>>>,
}
#[async_trait]
impl MessageMiddleware for BlockingMiddleware {
async fn before(&self, _ctx: &MessageContext) -> Result<Option<Frame>> {
if let Some(tx) = self.entered_tx.lock().await.take() {
let _ = tx.send(());
}
if let Some(rx) = self.release_rx.lock().await.take() {
let _ = rx.await;
}
Ok(None)
}
}
struct NoopMiddleware;
#[async_trait]
impl MessageMiddleware for NoopMiddleware {}
struct ParserUpdatingProcessor {
pipeline: MessagePipeline,
}
#[async_trait]
impl MessageProcessor for ParserUpdatingProcessor {
async fn process(&self, _ctx: &MessageContext) -> Result<Option<Frame>> {
self.pipeline.update_parser(MessageParser::protobuf()).await;
Ok(Some(frame_with_system_command(
pong(),
Reliability::BestEffort,
)))
}
}
#[tokio::test]
async fn middleware_updates_do_not_wait_for_inflight_middleware_to_finish() {
let pipeline = MessagePipeline::new(MessageParser::json());
let (entered_tx, entered_rx) = oneshot::channel();
let (release_tx, release_rx) = oneshot::channel();
pipeline
.add_middleware(Arc::new(BlockingMiddleware {
entered_tx: Mutex::new(Some(entered_tx)),
release_rx: Mutex::new(Some(release_rx)),
}))
.await;
let frame = FrameBuilder::new()
.with_command(Command {
r#type: Some(
crate::common::protocol::flare::core::commands::command::Type::System(ping()),
),
})
.with_reliability(Reliability::BestEffort)
.build();
let processing = {
let pipeline = pipeline.clone();
tokio::spawn(async move { pipeline.process_frame(&frame, Some("conn-1")).await })
};
entered_rx.await.expect("blocking middleware should start");
let add_result = tokio::time::timeout(
std::time::Duration::from_millis(50),
pipeline.add_middleware(Arc::new(NoopMiddleware)),
)
.await;
let _ = release_tx.send(());
processing
.await
.expect("pipeline task should not panic")
.expect("pipeline should succeed");
assert!(
add_result.is_ok(),
"middleware updates should use a copy-on-write snapshot and avoid waiting for in-flight middleware"
);
}
#[tokio::test]
async fn process_raw_serializes_response_with_request_parser_snapshot() {
let json_parser = MessageParser::json();
let pipeline = MessagePipeline::new(json_parser.clone());
pipeline
.add_processor(Arc::new(ParserUpdatingProcessor {
pipeline: pipeline.clone(),
}))
.await;
let request = frame_with_system_command(ping(), Reliability::BestEffort);
let request_data = json_parser
.serialize(&request)
.expect("json request should serialize");
let response_data = pipeline
.process_raw(&request_data, Some("conn-1"))
.await
.expect("pipeline should process request")
.expect("processor should produce response");
let response = json_parser.parse_with_format(&response_data, SerializationFormat::Json).expect(
"response should use the same parser snapshot as request even if parser is updated mid-pipeline",
);
let format = response
.command
.and_then(|command| command.r#type)
.and_then(|kind| match kind {
crate::common::protocol::flare::core::commands::command::Type::System(system) => {
SerializationFormat::try_from(system.format).ok()
}
_ => None,
})
.expect("response should be a system command");
assert_eq!(format, SerializationFormat::Protobuf);
}
}