Skip to main content

wasm_sandbox/communication/
memory.rs

1//! Shared memory abstractions for host-guest communication
2
3use std::sync::{Arc, Mutex};
4use std::collections::HashMap;
5
6use crate::error::{Error, Result};
7
8/// Shared memory region
9pub struct SharedMemoryRegion {
10    /// Region name
11    name: String,
12    
13    /// Memory buffer
14    buffer: Mutex<Vec<u8>>,
15    
16    /// Region size
17    size: usize,
18}
19
20impl SharedMemoryRegion {
21    /// Create a new shared memory region
22    pub fn new(name: &str, size: usize) -> Self {
23        Self {
24            name: name.to_string(),
25            buffer: Mutex::new(vec![0; size]),
26            size,
27        }
28    }
29    
30    /// Read from the shared memory region
31    pub fn read(&self, offset: usize, buf: &mut [u8]) -> Result<usize> {
32        // Get the buffer
33        let buffer = self.buffer.lock().unwrap();
34        
35        // Check if the offset is valid
36        if offset >= self.size {
37            return Err(Error::Communication {
38                channel: "shared_memory".to_string(),
39                reason: format!("Invalid offset: {}", offset),
40                instance_id: None,
41            });
42        }
43        
44        // Calculate the number of bytes to read
45        let n = std::cmp::min(buf.len(), self.size - offset);
46        
47        // Copy the data
48        buf[..n].copy_from_slice(&buffer[offset..offset+n]);
49        
50        Ok(n)
51    }
52    
53    /// Write to the shared memory region
54    pub fn write(&self, offset: usize, data: &[u8]) -> Result<usize> {
55        // Get the buffer
56        let mut buffer = self.buffer.lock().unwrap();
57        
58        // Check if the offset is valid
59        if offset >= self.size {
60            return Err(Error::Communication {
61                channel: "shared_memory".to_string(),
62                reason: format!("Invalid offset: {}", offset),
63                instance_id: None,
64            });
65        }
66        
67        // Calculate the number of bytes to write
68        let n = std::cmp::min(data.len(), self.size - offset);
69        
70        // Copy the data
71        buffer[offset..offset+n].copy_from_slice(&data[..n]);
72        
73        Ok(n)
74    }
75    
76    /// Get the region size
77    pub fn size(&self) -> usize {
78        self.size
79    }
80    
81    /// Get the region name
82    pub fn name(&self) -> &str {
83        &self.name
84    }
85}
86
87/// Shared memory manager
88pub struct SharedMemoryManager {
89    /// Shared memory regions
90    regions: Mutex<HashMap<String, Arc<SharedMemoryRegion>>>,
91}
92
93impl SharedMemoryManager {
94    /// Create a new shared memory manager
95    pub fn new() -> Self {
96        Self {
97            regions: Mutex::new(HashMap::new()),
98        }
99    }
100    
101    /// Create a new shared memory region
102    pub fn create_region(&self, name: &str, size: usize) -> Result<Arc<SharedMemoryRegion>> {
103        // Check if the region already exists
104        let mut regions = self.regions.lock().unwrap();
105        if regions.contains_key(name) {
106            return Err(Error::Communication {
107                channel: "shared_memory_manager".to_string(),
108                reason: format!("Region already exists: {}", name),
109                instance_id: None,
110            });
111        }
112        
113        // Create the region
114        let region = Arc::new(SharedMemoryRegion::new(name, size));
115        
116        // Register the region
117        regions.insert(name.to_string(), region.clone());
118        
119        Ok(region)
120    }
121    
122    /// Get a shared memory region
123    pub fn get_region(&self, name: &str) -> Option<Arc<SharedMemoryRegion>> {
124        let regions = self.regions.lock().unwrap();
125        regions.get(name).cloned()
126    }
127    
128    /// Delete a shared memory region
129    pub fn delete_region(&self, name: &str) -> Result<()> {
130        let mut regions = self.regions.lock().unwrap();
131        if regions.remove(name).is_none() {
132            return Err(Error::Communication {
133                channel: "shared_memory_manager".to_string(),
134                reason: format!("Region not found: {}", name),
135                instance_id: None,
136            });
137        }
138        
139        Ok(())
140    }
141    
142    /// List all shared memory regions
143    pub fn list_regions(&self) -> Vec<String> {
144        let regions = self.regions.lock().unwrap();
145        regions.keys().cloned().collect()
146    }
147}