use crate::api::{Ciphertext, GroupId};
use anyhow::Result;
use bytes::Bytes;
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
pub enum MlsFrameType {
ApplicationData = 0x01,
Handshake = 0x02,
Welcome = 0x03,
Commit = 0x04,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MlsFrame {
pub frame_type: MlsFrameType,
pub group_id: GroupId,
pub payload: Bytes,
}
#[derive(Debug)]
pub struct QuicStreamManager {
connections: Arc<RwLock<HashMap<GroupId, Vec<u8>>>>,
stream_ids: Arc<RwLock<HashMap<GroupId, Vec<u64>>>>,
}
impl Default for QuicStreamManager {
fn default() -> Self {
Self::new()
}
}
impl QuicStreamManager {
pub fn new() -> Self {
Self {
connections: Arc::new(RwLock::new(HashMap::new())),
stream_ids: Arc::new(RwLock::new(HashMap::new())),
}
}
pub fn register_connection(&self, group_id: GroupId, _connection_data: Vec<u8>) {
let mut connections = self.connections.write();
connections.insert(group_id, _connection_data);
}
pub async fn send_frame(&self, frame: &MlsFrame) -> Result<()> {
let connections = self.connections.read();
let _connection_data = connections
.get(&frame.group_id)
.ok_or_else(|| anyhow::anyhow!("No connection for group"))?;
let _data = postcard::to_stdvec(frame)?;
let mut stream_ids = self.stream_ids.write();
stream_ids
.entry(frame.group_id.clone())
.or_default()
.push(Self::stream_for_frame_type(frame.frame_type));
Ok(())
}
pub async fn send_application_data(
&self,
group_id: &GroupId,
ciphertext: &Ciphertext,
) -> Result<()> {
let frame = MlsFrame {
frame_type: MlsFrameType::ApplicationData,
group_id: group_id.clone(),
payload: ciphertext.data.clone(),
};
self.send_frame(&frame).await
}
pub async fn receive_frame(&self, _group_id: &GroupId) -> Result<MlsFrame> {
Ok(MlsFrame {
frame_type: MlsFrameType::ApplicationData,
group_id: GroupId::generate(),
payload: Bytes::new(),
})
}
pub fn stream_for_frame_type(frame_type: MlsFrameType) -> u64 {
match frame_type {
MlsFrameType::ApplicationData => 0,
MlsFrameType::Handshake => 1,
MlsFrameType::Welcome => 2,
MlsFrameType::Commit => 3,
}
}
pub async fn close_group(&self, group_id: &GroupId) -> Result<()> {
let mut connections = self.connections.write();
connections.remove(group_id);
let mut stream_ids = self.stream_ids.write();
stream_ids.remove(group_id);
Ok(())
}
}
pub fn encode_frame(frame: &MlsFrame) -> Result<Vec<u8>> {
Ok(postcard::to_stdvec(frame)?)
}
pub fn decode_frame(data: &[u8]) -> Result<MlsFrame> {
Ok(postcard::from_bytes(data)?)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_frame_encoding() {
let frame = MlsFrame {
frame_type: MlsFrameType::ApplicationData,
group_id: GroupId::generate(),
payload: Bytes::from(b"test payload".to_vec()),
};
let encoded = encode_frame(&frame).unwrap();
let decoded = decode_frame(&encoded).unwrap();
assert_eq!(frame.frame_type as u8, decoded.frame_type as u8);
assert_eq!(frame.group_id, decoded.group_id);
assert_eq!(frame.payload, decoded.payload);
}
#[test]
fn test_stream_mapping() {
assert_eq!(
QuicStreamManager::stream_for_frame_type(MlsFrameType::ApplicationData),
0
);
assert_eq!(
QuicStreamManager::stream_for_frame_type(MlsFrameType::Handshake),
1
);
assert_eq!(
QuicStreamManager::stream_for_frame_type(MlsFrameType::Welcome),
2
);
assert_eq!(
QuicStreamManager::stream_for_frame_type(MlsFrameType::Commit),
3
);
}
}