Skip to main content

active_call/media/
cache.rs

1use anyhow::{Result, anyhow};
2use bytes::BytesMut;
3use once_cell::sync::Lazy;
4use sha2::{Digest, Sha256};
5use std::sync::RwLock;
6use std::{
7    io::{IoSlice, SeekFrom},
8    path::PathBuf,
9};
10use tokio::io::{AsyncReadExt, AsyncSeekExt};
11use tokio::{fs::create_dir_all, io::AsyncWriteExt};
12use tracing::{debug, info};
13
14// Default cache directory
15static DEFAULT_CACHE_DIR: &str = "/tmp/mediacache";
16
17// Global cache configuration
18static CACHE_CONFIG: Lazy<RwLock<CacheConfig>> = Lazy::new(|| {
19    RwLock::new(CacheConfig {
20        cache_dir: PathBuf::from(DEFAULT_CACHE_DIR),
21    })
22});
23
24#[derive(Debug, Clone)]
25pub struct CacheConfig {
26    pub cache_dir: PathBuf,
27}
28
29/// Set the cache directory for the media cache
30pub fn set_cache_dir(path: &str) -> Result<()> {
31    let path = PathBuf::from(path);
32    let mut config = CACHE_CONFIG
33        .write()
34        .map_err(|_| anyhow!("Failed to acquire write lock"))?;
35    config.cache_dir = path;
36    Ok(())
37}
38
39/// Get the current cache directory
40pub fn get_cache_dir() -> Result<PathBuf> {
41    let config = CACHE_CONFIG
42        .read()
43        .map_err(|_| anyhow!("Failed to acquire read lock"))?;
44    Ok(config.cache_dir.clone())
45}
46
47/// Ensure the cache directory exists
48pub async fn ensure_cache_dir() -> Result<()> {
49    let cache_dir = get_cache_dir()?;
50
51    if !cache_dir.exists() {
52        debug!("Creating cache directory: {:?}", cache_dir);
53        create_dir_all(&cache_dir).await?;
54    }
55
56    Ok(())
57}
58
59/// Generate a cache key from text or URL
60pub fn generate_cache_key(
61    input: &str,
62    sample_rate: u32,
63    speaker: Option<&String>,
64    speed: Option<f32>,
65) -> String {
66    let mut hasher = Sha256::new();
67    hasher.update(input.as_bytes());
68    let result = hasher.finalize();
69    match speaker {
70        Some(speaker) => format!(
71            "{}_{}_{}_{}",
72            hex::encode(result),
73            sample_rate,
74            speaker,
75            speed.unwrap_or(1.0)
76        ),
77        None => format!(
78            "{}_{}_{}",
79            hex::encode(result),
80            sample_rate,
81            speed.unwrap_or(1.0)
82        ),
83    }
84}
85
86/// Get the full path for a cached file
87pub fn get_cache_path(key: &str) -> Result<PathBuf> {
88    let cache_dir = get_cache_dir()?;
89    Ok(cache_dir.join(key).with_extension("pcm"))
90}
91
92/// Check if a file exists in the cache
93pub async fn is_cached(key: &str) -> Result<bool> {
94    let path = get_cache_path(key)?;
95    Ok(tokio::fs::try_exists(&path).await?)
96}
97
98/// Store data in the cache
99pub async fn store_in_cache(key: &str, data: &Vec<u8>) -> Result<()> {
100    ensure_cache_dir().await?;
101    let path = get_cache_path(key)?;
102    tokio::fs::write(&path.with_extension(".tmp"), data).await?;
103    tokio::fs::rename(&path.with_extension(".tmp"), &path).await?;
104    info!("cache: Stored {} -> {} bytes", key, data.len());
105    Ok(())
106}
107
108// Store datas in the cache
109pub async fn store_in_cache_vectored(key: &str, data: &[impl AsRef<[u8]>]) -> Result<()> {
110    ensure_cache_dir().await?;
111    let path = get_cache_path(key)?;
112    let tmp_path = path.with_extension(".tmp");
113    let mut file = tokio::fs::File::create(tmp_path.clone()).await?;
114    let io_slices = data
115        .iter()
116        .map(|d| IoSlice::new(d.as_ref()))
117        .collect::<Vec<_>>();
118    let n = file.write_vectored(&io_slices).await?;
119    tokio::fs::rename(&tmp_path, &path).await?;
120    info!("cache: Stored {} -> {} bytes", key, n);
121    Ok(())
122}
123
124/// Store decoded PCM samples (i16) in the cache as raw little-endian bytes
125pub async fn store_pcm_in_cache(key: &str, samples: &[i16]) -> Result<()> {
126    let mut bytes = Vec::with_capacity(samples.len() * 2);
127    for sample in samples {
128        bytes.extend_from_slice(&sample.to_le_bytes());
129    }
130    store_in_cache(key, &bytes).await
131}
132
133/// Retrieve decoded PCM samples (i16) from the cache stored as raw little-endian bytes
134pub async fn retrieve_pcm_from_cache(key: &str) -> Result<Vec<i16>> {
135    retrieve_pcm_from_cache_at(key, 0).await
136}
137
138/// Retrieve decoded PCM samples (i16) from the cache starting at `sample_offset`
139/// samples in, seeking past the skipped bytes instead of reading them.
140pub async fn retrieve_pcm_from_cache_at(key: &str, sample_offset: usize) -> Result<Vec<i16>> {
141    let path = get_cache_path(key)?;
142    let mut file = tokio::fs::File::open(&path).await?;
143    let file_size = file.metadata().await?.len();
144    let byte_offset = ((sample_offset as u64) * 2).min(file_size);
145    if byte_offset > 0 {
146        file.seek(SeekFrom::Start(byte_offset)).await?;
147    }
148
149    let remaining = (file_size - byte_offset) as usize;
150    let mut bytes = Vec::with_capacity(remaining);
151    file.read_to_end(&mut bytes).await?;
152
153    if bytes.len() % 2 != 0 {
154        return Err(anyhow!(
155            "cache: pcm data length {} is not aligned to i16 for key: {}",
156            bytes.len(),
157            key
158        ));
159    }
160    let samples = bytes
161        .chunks_exact(2)
162        .map(|b| i16::from_le_bytes([b[0], b[1]]))
163        .collect();
164    Ok(samples)
165}
166
167/// Retrieve data from the cache
168pub async fn retrieve_from_cache(key: &str) -> Result<Vec<u8>> {
169    let path = get_cache_path(key)?;
170
171    if !tokio::fs::try_exists(&path).await? {
172        return Err(anyhow!("Cache file not found for key: {}", key));
173    }
174
175    let data = tokio::fs::read(&path).await?;
176    debug!(key, size = data.len(), "retrieved file from cache");
177    Ok(data)
178}
179
180// Retrieve data from the cache with a buffer
181pub async fn retrieve_from_cache_with_buffer(key: &str, buffer: &mut BytesMut) -> Result<()> {
182    let path = get_cache_path(key)?;
183    let mut file = tokio::fs::File::open(path).await?;
184    let metadata = file.metadata().await?;
185    let file_size = metadata.len() as usize;
186    buffer.reserve(file_size);
187
188    while file.read_buf(buffer).await? > 0 {}
189    Ok(())
190}
191
192/// Delete a specific file from the cache
193pub async fn delete_from_cache(key: &str) -> Result<()> {
194    let path = get_cache_path(key)?;
195
196    if tokio::fs::try_exists(&path).await? {
197        tokio::fs::remove_file(path).await?;
198        debug!("Deleted file from cache with key: {}", key);
199    }
200
201    Ok(())
202}
203
204#[cfg(test)]
205mod tests {
206    use super::*;
207    #[tokio::test]
208    async fn test_cache_operations() -> Result<()> {
209        ensure_cache_dir().await?;
210
211        // Generate a cache key
212        let key = generate_cache_key("test_data", 8000, None, None);
213
214        // Test storing data in cache
215        let test_data = b"TEST DATA".to_vec();
216        store_in_cache(&key, &test_data).await?;
217
218        // Test if data is cached
219        assert!(is_cached(&key).await?);
220
221        // Test retrieving data from cache
222        let retrieved_data = retrieve_from_cache(&key).await?;
223        assert_eq!(retrieved_data, test_data);
224
225        // Test deleting data from cache
226        delete_from_cache(&key).await?;
227        assert!(!is_cached(&key).await?);
228
229        // Test clean cache
230        let key2 = generate_cache_key("test_data2", 16000, None, None);
231        store_in_cache(&key2, &test_data).await?;
232        Ok(())
233    }
234
235    #[tokio::test]
236    async fn test_pcm_cache_roundtrip_and_offset() -> Result<()> {
237        let key = generate_cache_key("test_pcm_offset", 16000, None, None);
238        delete_from_cache(&key).await.ok();
239
240        let samples: Vec<i16> = (0..1000i16).collect();
241        store_pcm_in_cache(&key, &samples).await?;
242
243        // Full retrieval matches what we stored.
244        let full = retrieve_pcm_from_cache(&key).await?;
245        assert_eq!(full, samples);
246
247        // Offset retrieval matches slicing the full buffer.
248        let offset = 250;
249        let tail = retrieve_pcm_from_cache_at(&key, offset).await?;
250        assert_eq!(tail, samples[offset..]);
251
252        // Offset past the end yields an empty buffer rather than erroring.
253        let empty = retrieve_pcm_from_cache_at(&key, samples.len() + 100).await?;
254        assert!(empty.is_empty());
255
256        delete_from_cache(&key).await?;
257        Ok(())
258    }
259
260    #[test]
261    fn test_generate_cache_key() {
262        let key1 = generate_cache_key("hello", 16000, None, None);
263        let key2 = generate_cache_key("hello", 8000, None, None);
264        let key3 = generate_cache_key("world", 16000, None, None);
265
266        // Same input with different sample rates should produce different keys
267        assert_ne!(key1, key2);
268
269        // Different inputs with same sample rate should produce different keys
270        assert_ne!(key1, key3);
271    }
272
273    #[tokio::test]
274    async fn test_store_in_cache_v() -> Result<()> {
275        // Test data as multiple slices
276        let data1 = b"Hello, ";
277        let data2 = b"world!";
278        let data3 = b" This is a test.";
279        let data_slices = [data1.as_slice(), data2.as_slice(), data3.as_slice()];
280
281        let key = generate_cache_key("test_vectored_store", 16000, None, None);
282
283        // Ensure the key doesn't exist initially
284        delete_from_cache(&key).await.ok();
285        assert!(!is_cached(&key).await?);
286
287        // Store using vectored write
288        store_in_cache_vectored(&key, &data_slices).await?;
289
290        // Verify it was stored
291        assert!(is_cached(&key).await?);
292
293        // Retrieve and verify content
294        let retrieved = retrieve_from_cache(&key).await?;
295        let expected = [data1.as_slice(), data2.as_slice(), data3.as_slice()].concat();
296        assert_eq!(retrieved, expected);
297
298        // Clean up
299        delete_from_cache(&key).await?;
300        Ok(())
301    }
302
303    #[tokio::test]
304    async fn test_store_in_cache_v_empty_slices() -> Result<()> {
305        let empty_data: &[&[u8]] = &[];
306        let key = generate_cache_key("test_empty_vectored", 16000, None, None);
307
308        // Clean up first
309        delete_from_cache(&key).await.ok();
310
311        // Store empty data
312        store_in_cache_vectored(&key, empty_data).await?;
313
314        // Verify it was stored as empty file
315        assert!(is_cached(&key).await?);
316        let retrieved = retrieve_from_cache(&key).await?;
317        assert_eq!(retrieved.len(), 0);
318
319        // Clean up
320        delete_from_cache(&key).await?;
321        Ok(())
322    }
323
324    #[tokio::test]
325    async fn test_retrieve_from_cache_with_buffer() -> Result<()> {
326        // Test data
327        let test_data = b"This is test data for buffer retrieval testing.";
328        let key = generate_cache_key("test_buffer_retrieve", 16000, None, None);
329
330        // Clean up first
331        delete_from_cache(&key).await.ok();
332
333        // Store test data using regular store
334        store_in_cache(&key, &test_data.to_vec()).await?;
335
336        // Retrieve using buffer method
337        let mut buffer = BytesMut::new();
338        retrieve_from_cache_with_buffer(&key, &mut buffer).await?;
339
340        // Verify content
341        assert_eq!(buffer.as_ref(), test_data);
342        assert_eq!(buffer.len(), test_data.len());
343
344        // Clean up
345        delete_from_cache(&key).await?;
346        Ok(())
347    }
348
349    #[tokio::test]
350    async fn test_retrieve_from_cache_with_buffer_large_file() -> Result<()> {
351        // Create larger test data (1MB)
352        let large_data: Vec<u8> = vec![7; 1024 * 1024];
353        let key = generate_cache_key("test_large_buffer", 16000, None, None);
354
355        // Clean up first
356        delete_from_cache(&key).await.ok();
357
358        // Store large data
359        store_in_cache(&key, &large_data).await?;
360
361        // Retrieve using buffer method
362        let mut buffer = BytesMut::new();
363        retrieve_from_cache_with_buffer(&key, &mut buffer).await?;
364
365        // Verify content
366        assert_eq!(buffer.len(), large_data.len());
367        assert_eq!(buffer.as_ref(), large_data.as_slice());
368
369        // Clean up
370        delete_from_cache(&key).await?;
371        Ok(())
372    }
373
374    #[tokio::test]
375    async fn test_retrieve_from_cache_with_buffer_nonexistent() -> Result<()> {
376        let nonexistent_key = generate_cache_key("nonexistent_file", 16000, None, None);
377        let mut buffer = BytesMut::new();
378
379        // Should fail for nonexistent file
380        let result = retrieve_from_cache_with_buffer(&nonexistent_key, &mut buffer).await;
381        assert!(result.is_err());
382
383        Ok(())
384    }
385
386    #[tokio::test]
387    async fn test_store_v_and_retrieve_buffer_integration() -> Result<()> {
388        // Test integration between vectored store and buffer retrieve
389        let data_parts = [
390            b"Part 1: Hello".as_slice(),
391            b", Part 2: World".as_slice(),
392            b", Part 3: Integration Test!".as_slice(),
393        ];
394        let key = generate_cache_key("test_integration", 16000, None, None);
395
396        // Clean up first
397        delete_from_cache(&key).await.ok();
398
399        // Store using vectored write
400        store_in_cache_vectored(&key, &data_parts).await?;
401
402        // Retrieve using buffer method
403        let mut buffer = BytesMut::new();
404        retrieve_from_cache_with_buffer(&key, &mut buffer).await?;
405
406        // Verify the data was correctly concatenated
407        let expected = data_parts.concat();
408        assert_eq!(buffer.as_ref(), expected.as_slice());
409        assert_eq!(buffer.len(), expected.len());
410
411        // Also verify with regular retrieve for double-check
412        let regular_retrieve = retrieve_from_cache(&key).await?;
413        assert_eq!(regular_retrieve, expected);
414        assert_eq!(buffer.as_ref(), regular_retrieve.as_slice());
415
416        // Clean up
417        delete_from_cache(&key).await?;
418        Ok(())
419    }
420}