use std::collections::HashMap;
use std::io::{Cursor, Read, Seek, SeekFrom, Write};
use std::sync::{Arc, Mutex};
use parking_lot::RwLock;
use crate::error::{LaurusError, Result};
use crate::storage::{
LockManager, Storage, StorageError, StorageInput, StorageLock, StorageOutput,
};
#[derive(Debug, Clone)]
pub struct MemoryStorageConfig {
pub initial_capacity: usize,
}
impl Default for MemoryStorageConfig {
fn default() -> Self {
MemoryStorageConfig {
initial_capacity: 16,
}
}
}
#[derive(Debug)]
struct MemFileInner {
data: Vec<u8>,
committed: usize,
}
#[derive(Debug)]
struct MemFile {
inner: RwLock<MemFileInner>,
}
impl MemFile {
fn new(data: Vec<u8>, committed: usize) -> Self {
MemFile {
inner: RwLock::new(MemFileInner { data, committed }),
}
}
}
#[derive(Debug)]
pub struct MemoryStorage {
files: Arc<RwLock<HashMap<String, Arc<MemFile>>>>,
lock_manager: Arc<MemoryLockManager>,
#[allow(dead_code)]
config: MemoryStorageConfig,
closed: bool,
}
impl Default for MemoryStorage {
fn default() -> Self {
Self::new(MemoryStorageConfig::default())
}
}
impl MemoryStorage {
pub fn new(config: MemoryStorageConfig) -> Self {
let initial_capacity = config.initial_capacity;
MemoryStorage {
files: Arc::new(RwLock::new(HashMap::with_capacity(initial_capacity))),
lock_manager: Arc::new(MemoryLockManager::new()),
config,
closed: false,
}
}
fn check_closed(&self) -> Result<()> {
if self.closed {
Err(StorageError::StorageClosed.into())
} else {
Ok(())
}
}
#[inline]
pub fn file_count(&self) -> usize {
self.files.read().len()
}
pub fn total_size(&self) -> u64 {
let files = self.files.read();
files
.values()
.map(|file| file.inner.read().committed as u64)
.sum()
}
pub fn clear(&self) -> Result<()> {
self.check_closed()?;
let mut files = self.files.write();
files.clear();
Ok(())
}
}
impl Storage for MemoryStorage {
#[inline]
fn open_input(&self, name: &str) -> Result<Box<dyn StorageInput>> {
self.check_closed()?;
let file = {
let files = self.files.read();
files
.get(name)
.cloned()
.ok_or_else(|| StorageError::FileNotFound(name.to_string()))?
};
let snapshot = {
let inner = file.inner.read();
inner.data[..inner.committed].to_vec()
};
Ok(Box::new(MemoryInput::new(snapshot)))
}
fn create_output(&self, name: &str) -> Result<Box<dyn StorageOutput>> {
self.check_closed()?;
Ok(Box::new(MemoryOutput::new(
name.to_string(),
Arc::clone(&self.files),
)))
}
fn create_output_append(&self, name: &str) -> Result<Box<dyn StorageOutput>> {
self.check_closed()?;
Ok(Box::new(MemoryOutput::new_append(
name.to_string(),
Arc::clone(&self.files),
)))
}
fn file_exists(&self, name: &str) -> bool {
if self.closed {
return false;
}
let files = self.files.read();
files.contains_key(name)
}
fn delete_file(&self, name: &str) -> Result<()> {
self.check_closed()?;
let mut files = self.files.write();
files.remove(name);
Ok(())
}
fn list_files(&self) -> Result<Vec<String>> {
self.check_closed()?;
let files = self.files.read();
let mut file_names: Vec<String> = files.keys().cloned().collect();
file_names.sort();
Ok(file_names)
}
fn file_size(&self, name: &str) -> Result<u64> {
self.check_closed()?;
let files = self.files.read();
let file = files
.get(name)
.ok_or_else(|| StorageError::FileNotFound(name.to_string()))?;
Ok(file.inner.read().committed as u64)
}
fn metadata(&self, name: &str) -> Result<crate::storage::FileMetadata> {
self.check_closed()?;
let files = self.files.read();
if let Some(file) = files.get(name) {
let now = crate::util::time::now_secs();
Ok(crate::storage::FileMetadata {
size: file.inner.read().committed as u64,
modified: now,
created: now,
readonly: false,
})
} else {
Err(LaurusError::storage(format!("File not found: {name}")))
}
}
fn rename_file(&self, old_name: &str, new_name: &str) -> Result<()> {
self.check_closed()?;
let mut files = self.files.write();
let data = files
.remove(old_name)
.ok_or_else(|| StorageError::FileNotFound(old_name.to_string()))?;
files.insert(new_name.to_string(), data);
Ok(())
}
fn create_temp_output(&self, prefix: &str) -> Result<(String, Box<dyn StorageOutput>)> {
self.check_closed()?;
let mut counter = 0;
let mut temp_name;
loop {
temp_name = format!("{prefix}_{counter}.tmp");
if !self.file_exists(&temp_name) {
break;
}
counter += 1;
if counter > 10000 {
return Err(
StorageError::IoError("Could not create temporary file".to_string()).into(),
);
}
}
let output = self.create_output(&temp_name)?;
Ok((temp_name, output))
}
fn sync(&self) -> Result<()> {
self.check_closed()?;
Ok(())
}
fn close(&mut self) -> Result<()> {
self.closed = true;
self.lock_manager.release_all()?;
Ok(())
}
}
#[derive(Debug)]
pub struct MemoryInput {
cursor: Cursor<Vec<u8>>,
size: u64,
}
impl MemoryInput {
fn new(data: Vec<u8>) -> Self {
let size = data.len() as u64;
let cursor = Cursor::new(data);
MemoryInput { cursor, size }
}
}
impl Read for MemoryInput {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
self.cursor.read(buf)
}
}
impl Seek for MemoryInput {
fn seek(&mut self, pos: SeekFrom) -> std::io::Result<u64> {
self.cursor.seek(pos)
}
}
impl StorageInput for MemoryInput {
fn size(&self) -> Result<u64> {
Ok(self.size)
}
fn clone_input(&self) -> Result<Box<dyn StorageInput>> {
Ok(Box::new(MemoryInput::new(self.cursor.get_ref().clone())))
}
fn close(&mut self) -> Result<()> {
Ok(())
}
fn as_slice(&self) -> Option<&[u8]> {
let pos = self.cursor.position() as usize;
let data = self.cursor.get_ref();
Some(&data[pos.min(data.len())..])
}
}
#[derive(Debug)]
pub struct MemoryOutput {
name: String,
memfile: Arc<MemFile>,
files: Arc<RwLock<HashMap<String, Arc<MemFile>>>>,
position: u64,
closed: bool,
}
impl MemoryOutput {
fn new(name: String, files: Arc<RwLock<HashMap<String, Arc<MemFile>>>>) -> Self {
MemoryOutput {
name,
memfile: Arc::new(MemFile::new(Vec::new(), 0)),
files,
position: 0,
closed: false,
}
}
fn new_append(name: String, files: Arc<RwLock<HashMap<String, Arc<MemFile>>>>) -> Self {
let existing_data = {
let files_guard = files.read();
files_guard
.get(&name)
.map(|file| {
let inner = file.inner.read();
inner.data[..inner.committed].to_vec()
})
.unwrap_or_default()
};
let position = existing_data.len() as u64;
let committed = existing_data.len();
MemoryOutput {
name,
memfile: Arc::new(MemFile::new(existing_data, committed)),
files,
position,
closed: false,
}
}
}
impl Write for MemoryOutput {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
if self.closed {
return Err(std::io::Error::other("Output is closed"));
}
let end_pos = (self.position as usize)
.checked_add(buf.len())
.ok_or_else(|| std::io::Error::other("File too large"))?;
let mut inner = self.memfile.inner.write();
if end_pos > inner.data.len() {
inner.data.resize(end_pos, 0);
}
inner.data[self.position as usize..end_pos].copy_from_slice(buf);
drop(inner);
self.position += buf.len() as u64;
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl Seek for MemoryOutput {
fn seek(&mut self, pos: SeekFrom) -> std::io::Result<u64> {
if self.closed {
return Err(std::io::Error::other("Output is closed"));
}
let new_pos = match pos {
SeekFrom::Start(offset) => offset,
SeekFrom::End(offset) => {
let len = self.memfile.inner.read().data.len() as u64;
if offset < 0 {
let abs_offset = (-offset) as u64;
if abs_offset > len {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Invalid seek position",
));
}
len - abs_offset
} else {
len + offset as u64
}
}
SeekFrom::Current(offset) => {
if offset < 0 {
let abs_offset = (-offset) as u64;
if abs_offset > self.position {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Invalid seek position",
));
}
self.position - abs_offset
} else {
self.position + offset as u64
}
}
};
self.position = new_pos;
Ok(new_pos)
}
}
impl StorageOutput for MemoryOutput {
fn flush_and_sync(&mut self) -> Result<()> {
self.publish();
Ok(())
}
fn position(&self) -> Result<u64> {
Ok(self.position)
}
fn close(&mut self) -> Result<()> {
if !self.closed {
self.publish();
self.closed = true;
}
Ok(())
}
}
impl MemoryOutput {
fn publish(&mut self) {
{
let mut inner = self.memfile.inner.write();
inner.committed = inner.data.len();
}
self.files
.write()
.insert(self.name.clone(), Arc::clone(&self.memfile));
}
}
impl Drop for MemoryOutput {
fn drop(&mut self) {
let _ = self.close();
}
}
#[derive(Debug)]
pub struct MemoryLockManager {
locks: Arc<Mutex<HashMap<String, Arc<Mutex<MemoryLock>>>>>,
}
impl MemoryLockManager {
fn new() -> Self {
MemoryLockManager {
locks: Arc::new(Mutex::new(HashMap::new())),
}
}
}
impl LockManager for MemoryLockManager {
fn acquire_lock(&self, name: &str) -> Result<Box<dyn StorageLock>> {
let mut locks = self.locks.lock().unwrap();
if locks.contains_key(name) {
return Err(StorageError::LockFailed(name.to_string()).into());
}
let lock = Arc::new(Mutex::new(MemoryLock::new(name.to_string())));
locks.insert(name.to_string(), lock.clone());
Ok(Box::new(MemoryLockWrapper { lock }))
}
fn try_acquire_lock(&self, name: &str) -> Result<Option<Box<dyn StorageLock>>> {
match self.acquire_lock(name) {
Ok(lock) => Ok(Some(lock)),
Err(e) => {
if let LaurusError::Storage(ref msg) = e
&& msg.contains("Failed to acquire lock")
{
return Ok(None);
}
Err(e)
}
}
}
fn lock_exists(&self, name: &str) -> bool {
let locks = self.locks.lock().unwrap();
locks.contains_key(name)
}
fn release_all(&self) -> Result<()> {
let mut locks = self.locks.lock().unwrap();
locks.clear();
Ok(())
}
}
#[derive(Debug)]
struct MemoryLock {
#[allow(dead_code)]
name: String,
released: bool,
}
impl MemoryLock {
fn new(name: String) -> Self {
MemoryLock {
name,
released: false,
}
}
}
#[derive(Debug)]
struct MemoryLockWrapper {
lock: Arc<Mutex<MemoryLock>>,
}
impl StorageLock for MemoryLockWrapper {
fn name(&self) -> &str {
"memory_lock"
}
fn release(&mut self) -> Result<()> {
let mut lock = self.lock.lock().unwrap();
lock.released = true;
Ok(())
}
fn is_valid(&self) -> bool {
let lock = self.lock.lock().unwrap();
!lock.released
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
#[test]
fn test_memory_storage_creation() {
let storage = MemoryStorage::default();
assert_eq!(storage.file_count(), 0);
assert_eq!(storage.total_size(), 0);
}
#[test]
fn test_create_and_read_file() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("test.txt").unwrap();
output.write_all(b"Hello, Memory!").unwrap();
output.close().unwrap();
let mut input = storage.open_input("test.txt").unwrap();
let mut buffer = Vec::new();
input.read_to_end(&mut buffer).unwrap();
assert_eq!(buffer, b"Hello, Memory!");
assert_eq!(input.size().unwrap(), 14);
assert_eq!(storage.file_count(), 1);
assert_eq!(storage.total_size(), 14);
}
#[test]
fn as_slice_returns_remaining_bytes() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("data.bin").unwrap();
output.write_all(b"abcdefghij").unwrap();
output.close().unwrap();
let mut input = storage.open_input("data.bin").unwrap();
assert_eq!(input.as_slice(), Some(&b"abcdefghij"[..]));
let mut head = [0u8; 3];
input.read_exact(&mut head).unwrap();
assert_eq!(&head, b"abc");
assert_eq!(input.as_slice(), Some(&b"defghij"[..]));
input.seek(SeekFrom::End(0)).unwrap();
assert_eq!(input.as_slice(), Some(&[][..]));
}
#[test]
fn test_file_operations() {
let storage = MemoryStorage::default();
assert!(!storage.file_exists("nonexistent.txt"));
let mut output = storage.create_output("test.txt").unwrap();
output.write_all(b"Test content").unwrap();
output.close().unwrap();
assert!(storage.file_exists("test.txt"));
assert_eq!(storage.file_size("test.txt").unwrap(), 12);
let files = storage.list_files().unwrap();
assert_eq!(files, vec!["test.txt"]);
storage.rename_file("test.txt", "renamed.txt").unwrap();
assert!(!storage.file_exists("test.txt"));
assert!(storage.file_exists("renamed.txt"));
storage.delete_file("renamed.txt").unwrap();
assert!(!storage.file_exists("renamed.txt"));
assert_eq!(storage.file_count(), 0);
}
#[test]
fn test_multiple_files() {
let storage = MemoryStorage::default();
for i in 0..5 {
let mut output = storage.create_output(&format!("file_{i}.txt")).unwrap();
output.write_all(format!("Content {i}").as_bytes()).unwrap();
output.close().unwrap();
}
assert_eq!(storage.file_count(), 5);
let files = storage.list_files().unwrap();
assert_eq!(files.len(), 5);
for (i, file) in files.iter().enumerate().take(5) {
assert_eq!(file, &format!("file_{i}.txt"));
}
}
#[test]
fn test_temp_file_creation() {
let storage = MemoryStorage::default();
let (temp_name, mut output) = storage.create_temp_output("test").unwrap();
assert!(temp_name.starts_with("test_"));
assert!(temp_name.ends_with(".tmp"));
output.write_all(b"Temporary content").unwrap();
output.close().unwrap();
assert!(storage.file_exists(&temp_name));
assert_eq!(storage.file_size(&temp_name).unwrap(), 17);
}
#[test]
fn test_input_clone() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("test.txt").unwrap();
output.write_all(b"Hello, Clone!").unwrap();
output.close().unwrap();
let mut input1 = storage.open_input("test.txt").unwrap();
let mut input2 = input1.clone_input().unwrap();
let mut buffer1 = Vec::new();
let mut buffer2 = Vec::new();
input1.read_to_end(&mut buffer1).unwrap();
input2.read_to_end(&mut buffer2).unwrap();
assert_eq!(buffer1, b"Hello, Clone!");
assert_eq!(buffer2, b"Hello, Clone!");
assert_eq!(buffer1, buffer2);
}
#[test]
fn test_seek_operations() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("test.txt").unwrap();
output.write_all(b"0123456789").unwrap();
output.close().unwrap();
let mut input = storage.open_input("test.txt").unwrap();
input.seek(SeekFrom::Start(5)).unwrap();
let mut buffer = [0u8; 3];
input.read_exact(&mut buffer).unwrap();
assert_eq!(&buffer, b"567");
input.seek(SeekFrom::End(-2)).unwrap();
let mut buffer = [0u8; 2];
input.read_exact(&mut buffer).unwrap();
assert_eq!(&buffer, b"89");
}
#[test]
fn test_file_not_found() {
let storage = MemoryStorage::default();
let result = storage.open_input("nonexistent.txt");
assert!(result.is_err());
let result = storage.file_size("nonexistent.txt");
assert!(result.is_err());
}
#[test]
fn test_storage_close() {
let mut storage = MemoryStorage::default();
storage.close().unwrap();
assert!(storage.closed);
let result = storage.create_output("test.txt");
assert!(result.is_err());
}
#[test]
fn test_clear_storage() {
let storage = MemoryStorage::default();
for i in 0..3 {
let mut output = storage.create_output(&format!("file_{i}.txt")).unwrap();
output.write_all(b"content").unwrap();
output.close().unwrap();
}
assert_eq!(storage.file_count(), 3);
storage.clear().unwrap();
assert_eq!(storage.file_count(), 0);
assert_eq!(storage.total_size(), 0);
}
fn read_file(storage: &MemoryStorage, name: &str) -> Vec<u8> {
let mut input = storage.open_input(name).unwrap();
let mut buf = Vec::new();
input.read_to_end(&mut buf).unwrap();
buf
}
#[test]
fn flush_then_open_input_sees_only_committed() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("wal").unwrap();
output.write_all(b"AAAA").unwrap();
output.flush_and_sync().unwrap();
assert_eq!(read_file(&storage, "wal"), b"AAAA");
output.write_all(b"BBBB").unwrap();
assert_eq!(read_file(&storage, "wal"), b"AAAA");
assert_eq!(storage.file_size("wal").unwrap(), 4);
output.flush_and_sync().unwrap();
assert_eq!(read_file(&storage, "wal"), b"AAAABBBB");
assert_eq!(storage.file_size("wal").unwrap(), 8);
}
#[test]
fn repeated_flush_is_amortized_and_correct() {
let storage = MemoryStorage::default();
let mut output = storage.create_output_append("wal").unwrap();
let mut expected = Vec::new();
for i in 0u32..64 {
let record = i.to_le_bytes();
output.write_all(&record).unwrap();
output.flush_and_sync().unwrap();
expected.extend_from_slice(&record);
assert_eq!(storage.file_size("wal").unwrap(), expected.len() as u64);
}
assert_eq!(read_file(&storage, "wal"), expected);
}
#[test]
fn truncate_keeps_old_content_until_republish() {
let storage = MemoryStorage::default();
let mut old = storage.create_output("f").unwrap();
old.write_all(b"oldcontent").unwrap();
old.close().unwrap();
assert_eq!(read_file(&storage, "f"), b"oldcontent");
let mut fresh = storage.create_output("f").unwrap();
assert_eq!(read_file(&storage, "f"), b"oldcontent");
fresh.write_all(b"new").unwrap();
fresh.close().unwrap();
assert_eq!(read_file(&storage, "f"), b"new");
assert_eq!(storage.file_size("f").unwrap(), 3);
}
#[test]
fn create_output_close_without_write_publishes_empty_file() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("empty.log").unwrap();
output.close().unwrap();
assert!(storage.file_exists("empty.log"));
assert_eq!(storage.file_size("empty.log").unwrap(), 0);
assert_eq!(read_file(&storage, "empty.log"), b"");
}
#[test]
fn delete_then_writer_flush_resurrects_file() {
let storage = MemoryStorage::default();
let mut output = storage.create_output_append("wal").unwrap();
output.write_all(b"AAAA").unwrap();
output.flush_and_sync().unwrap();
assert!(storage.file_exists("wal"));
storage.delete_file("wal").unwrap();
assert!(!storage.file_exists("wal"));
output.write_all(b"BBBB").unwrap();
output.flush_and_sync().unwrap();
assert!(storage.file_exists("wal"));
assert_eq!(read_file(&storage, "wal"), b"AAAABBBB");
}
#[test]
fn seek_back_overwrite_header_then_commit() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("bkd").unwrap();
output.write_all(&[0u8; 4]).unwrap();
output.write_all(b"PAYLOAD").unwrap();
output.seek(SeekFrom::Start(0)).unwrap();
output.write_all(b"HEAD").unwrap();
output.seek(SeekFrom::End(0)).unwrap();
output.close().unwrap();
assert_eq!(read_file(&storage, "bkd"), b"HEADPAYLOAD");
assert_eq!(storage.file_size("bkd").unwrap(), 11);
}
#[test]
fn size_surfaces_report_committed_not_written_tail() {
let storage = MemoryStorage::default();
let mut output = storage.create_output("f").unwrap();
output.write_all(b"AAAA").unwrap();
output.flush_and_sync().unwrap();
output.write_all(b"BBBBBB").unwrap();
assert_eq!(storage.file_size("f").unwrap(), 4);
assert_eq!(storage.total_size(), 4);
assert_eq!(storage.metadata("f").unwrap().size, 4);
assert_eq!(storage.open_input("f").unwrap().size().unwrap(), 4);
}
}