use std::sync::Arc;
use dashmap::DashMap;
use serde::de::DeserializeOwned;
use serde::Serialize;
use super::types::IpcError;
use crate::traits::ActonMessage;
type DeserializerFn = Arc<
dyn Fn(&[u8]) -> Result<Box<dyn ActonMessage + Send + Sync>, String> + Send + Sync,
>;
type SerializerFn = Arc<
dyn Fn(&dyn ActonMessage) -> Result<serde_json::Value, String> + Send + Sync,
>;
#[derive(Default)]
pub struct IpcTypeRegistry {
deserializers: DashMap<String, DeserializerFn>,
type_id_to_name: DashMap<std::any::TypeId, String>,
serializers: DashMap<std::any::TypeId, SerializerFn>,
}
impl std::fmt::Debug for IpcTypeRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("IpcTypeRegistry")
.field("registered_types", &self.deserializers.len())
.field("type_id_mappings", &self.type_id_to_name.len())
.field("serializers", &self.serializers.len())
.finish()
}
}
impl IpcTypeRegistry {
#[must_use]
pub fn new() -> Self {
Self {
deserializers: DashMap::new(),
type_id_to_name: DashMap::new(),
serializers: DashMap::new(),
}
}
pub fn register<M>(&self, name: &str)
where
M: ActonMessage + Serialize + DeserializeOwned + 'static,
{
let deserializer: DeserializerFn = Arc::new(|bytes: &[u8]| {
let msg: M = serde_json::from_slice(bytes).map_err(|e| e.to_string())?;
Ok(Box::new(msg))
});
self.deserializers.insert(name.to_string(), deserializer);
let type_id = std::any::TypeId::of::<M>();
self.type_id_to_name.insert(type_id, name.to_string());
let serializer: SerializerFn = Arc::new(|msg: &dyn ActonMessage| {
let concrete = msg
.as_any()
.downcast_ref::<M>()
.ok_or_else(|| "Type mismatch during serialization".to_string())?;
serde_json::to_value(concrete).map_err(|e| e.to_string())
});
self.serializers.insert(type_id, serializer);
}
pub fn register_with_type_name<M>(&self)
where
M: ActonMessage + Serialize + DeserializeOwned + 'static,
{
let type_name = std::any::type_name::<M>();
self.register::<M>(type_name);
}
pub fn deserialize(
&self,
type_name: &str,
bytes: &[u8],
) -> Result<Box<dyn ActonMessage + Send + Sync>, IpcError> {
let deserializer = self
.deserializers
.get(type_name)
.ok_or_else(|| IpcError::UnknownMessageType(type_name.to_string()))?;
deserializer(bytes).map_err(IpcError::SerializationError)
}
pub fn deserialize_value(
&self,
type_name: &str,
value: &serde_json::Value,
) -> Result<Box<dyn ActonMessage + Send + Sync>, IpcError> {
let bytes = serde_json::to_vec(value)?;
self.deserialize(type_name, &bytes)
}
#[must_use]
pub fn is_registered(&self, type_name: &str) -> bool {
self.deserializers.contains_key(type_name)
}
#[must_use]
pub fn len(&self) -> usize {
self.deserializers.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.deserializers.is_empty()
}
pub fn type_names(&self) -> impl Iterator<Item = String> + '_ {
self.deserializers.iter().map(|entry| entry.key().clone())
}
#[must_use]
pub fn get_type_name_by_id(&self, type_id: &std::any::TypeId) -> Option<String> {
self.type_id_to_name.get(type_id).map(|r| r.clone())
}
pub fn serialize_by_type_id(
&self,
type_id: &std::any::TypeId,
message: &dyn ActonMessage,
) -> Result<serde_json::Value, String> {
let serializer = self
.serializers
.get(type_id)
.ok_or_else(|| "Type not registered for IPC serialization".to_string())?;
serializer(message)
}
}
impl Clone for IpcTypeRegistry {
fn clone(&self) -> Self {
let new_deserializers = DashMap::new();
for entry in &self.deserializers {
new_deserializers.insert(entry.key().clone(), entry.value().clone());
}
let new_type_id_map = DashMap::new();
for entry in &self.type_id_to_name {
new_type_id_map.insert(*entry.key(), entry.value().clone());
}
let new_serializers = DashMap::new();
for entry in &self.serializers {
new_serializers.insert(*entry.key(), entry.value().clone());
}
Self {
deserializers: new_deserializers,
type_id_to_name: new_type_id_map,
serializers: new_serializers,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
struct TestMessage {
value: i32,
text: String,
}
#[test]
fn test_register_and_deserialize() {
let registry = IpcTypeRegistry::new();
registry.register::<TestMessage>("TestMessage");
assert!(registry.is_registered("TestMessage"));
assert!(!registry.is_registered("Unknown"));
let msg = TestMessage {
value: 42,
text: "hello".to_string(),
};
let bytes = serde_json::to_vec(&msg).unwrap();
let result = registry.deserialize("TestMessage", &bytes);
assert!(result.is_ok());
let boxed = result.unwrap();
let downcast = (*boxed).as_any().downcast_ref::<TestMessage>();
assert!(downcast.is_some());
assert_eq!(downcast.unwrap(), &msg);
}
#[test]
fn test_register_with_type_name() {
let registry = IpcTypeRegistry::new();
registry.register_with_type_name::<TestMessage>();
let type_name = std::any::type_name::<TestMessage>();
assert!(registry.is_registered(type_name));
}
#[test]
fn test_deserialize_unknown_type() {
let registry = IpcTypeRegistry::new();
let result = registry.deserialize("Unknown", b"{}");
assert!(matches!(result, Err(IpcError::UnknownMessageType(_))));
}
#[test]
fn test_deserialize_invalid_json() {
let registry = IpcTypeRegistry::new();
registry.register::<TestMessage>("TestMessage");
let result = registry.deserialize("TestMessage", b"not valid json");
assert!(matches!(result, Err(IpcError::SerializationError(_))));
}
#[test]
fn test_deserialize_value() {
let registry = IpcTypeRegistry::new();
registry.register::<TestMessage>("TestMessage");
let value = serde_json::json!({
"value": 42,
"text": "hello"
});
let result = registry.deserialize_value("TestMessage", &value);
assert!(result.is_ok());
}
#[test]
fn test_registry_len_and_empty() {
let registry = IpcTypeRegistry::new();
assert!(registry.is_empty());
assert_eq!(registry.len(), 0);
registry.register::<TestMessage>("TestMessage");
assert!(!registry.is_empty());
assert_eq!(registry.len(), 1);
}
#[test]
fn test_type_names_iterator() {
let registry = IpcTypeRegistry::new();
registry.register::<TestMessage>("Type1");
registry.register::<TestMessage>("Type2");
let names: Vec<String> = registry.type_names().collect();
assert_eq!(names.len(), 2);
assert!(names.contains(&"Type1".to_string()));
assert!(names.contains(&"Type2".to_string()));
}
#[test]
fn test_registry_clone() {
let registry = IpcTypeRegistry::new();
registry.register::<TestMessage>("TestMessage");
let cloned = registry.clone();
assert!(cloned.is_registered("TestMessage"));
let msg = TestMessage {
value: 99,
text: "cloned".to_string(),
};
let bytes = serde_json::to_vec(&msg).unwrap();
let original_result = registry.deserialize("TestMessage", &bytes);
assert!(original_result.is_ok());
let clone_result = cloned.deserialize("TestMessage", &bytes);
assert!(clone_result.is_ok());
let boxed = clone_result.unwrap();
let downcast = (*boxed).as_any().downcast_ref::<TestMessage>();
assert!(downcast.is_some());
assert_eq!(downcast.unwrap(), &msg);
}
}