1use rand::RngCore;
9use rusqlite::{params, Connection, OptionalExtension};
10
11pub const MAX_UPLOAD_BYTES: usize = 25 * 1024 * 1024;
13
14fn env_num(name: &str, default: u64) -> u64 {
15 std::env::var(name).ok().and_then(|v| v.trim().parse().ok()).unwrap_or(default)
16}
17
18pub fn max_bytes() -> usize {
20 env_num("M4A_MEDIA_MAX_BYTES", MAX_UPLOAD_BYTES as u64) as usize
21}
22
23pub fn federation_enabled() -> bool {
26 cfg!(feature = "media-federation") && std::env::var("M4A_MEDIA_FEDERATION").map(|v| v != "off").unwrap_or(true)
27}
28
29pub fn remote_ttl_ms() -> i64 {
31 env_num("M4A_MEDIA_REMOTE_TTL_SECS", 7 * 24 * 3600) as i64 * 1000
32}
33
34pub fn remote_cache_cap() -> i64 {
36 env_num("M4A_MEDIA_REMOTE_CACHE_BYTES", 512 * 1024 * 1024) as i64
37}
38
39pub fn create_media_schema(conn: &Connection) -> rusqlite::Result<()> {
41 conn.execute_batch(
42 "CREATE TABLE IF NOT EXISTS media (
43 media_id TEXT PRIMARY KEY,
44 owner_user_id INTEGER NOT NULL,
45 content_type TEXT NOT NULL,
46 filename TEXT,
47 created_ms INTEGER NOT NULL,
48 data BLOB NOT NULL
49 );
50 CREATE INDEX IF NOT EXISTS idx_media_created ON media(created_ms);
51 CREATE TABLE IF NOT EXISTS media_pending (
52 media_id TEXT PRIMARY KEY,
53 owner_user_id INTEGER NOT NULL,
54 expires_ms INTEGER NOT NULL
55 );
56 CREATE TABLE IF NOT EXISTS remote_media (
57 server TEXT NOT NULL,
58 media_id TEXT NOT NULL,
59 content_type TEXT NOT NULL,
60 filename TEXT,
61 fetched_ms INTEGER NOT NULL,
62 size INTEGER NOT NULL,
63 data BLOB NOT NULL,
64 PRIMARY KEY (server, media_id)
65 );",
66 )
67}
68
69pub struct Blob {
71 pub content_type: String,
73 pub filename: Option<String>,
75 pub data: Vec<u8>,
77}
78
79pub fn put(conn: &Connection, owner: i64, content_type: &str, filename: Option<&str>, data: &[u8], now_ms: i64) -> rusqlite::Result<String> {
81 let mut raw = [0u8; 18];
82 rand::thread_rng().fill_bytes(&mut raw);
83 let id = base64::Engine::encode(&base64::engine::general_purpose::URL_SAFE_NO_PAD, raw);
84 conn.execute(
85 "INSERT INTO media (media_id, owner_user_id, content_type, filename, created_ms, data) VALUES (?1,?2,?3,?4,?5,?6)",
86 params![id, owner, content_type, filename, now_ms, data],
87 )?;
88 Ok(id)
89}
90
91pub fn get(conn: &Connection, media_id: &str) -> rusqlite::Result<Option<Blob>> {
93 conn.query_row("SELECT content_type, filename, data FROM media WHERE media_id = ?1", params![media_id], |r| {
94 Ok(Blob { content_type: r.get(0)?, filename: r.get(1)?, data: r.get(2)? })
95 })
96 .optional()
97}
98
99pub fn purge_expired(conn: &Connection, now_ms: i64, ttl_ms: i64) -> rusqlite::Result<usize> {
101 conn.execute("DELETE FROM media WHERE created_ms < ?1", params![now_ms - ttl_ms])
102}
103
104pub fn get_remote(conn: &Connection, server: &str, media_id: &str, now_ms: i64) -> rusqlite::Result<Option<Blob>> {
106 conn.query_row(
107 "SELECT content_type, filename, data FROM remote_media WHERE server = ?1 AND media_id = ?2 AND fetched_ms > ?3",
108 params![server, media_id, now_ms - remote_ttl_ms()],
109 |r| Ok(Blob { content_type: r.get(0)?, filename: r.get(1)?, data: r.get(2)? }),
110 )
111 .optional()
112}
113
114pub fn put_remote(conn: &Connection, server: &str, media_id: &str, blob: &Blob, now_ms: i64) -> rusqlite::Result<()> {
116 conn.execute("DELETE FROM remote_media WHERE fetched_ms <= ?1", [now_ms - remote_ttl_ms()])?;
117 conn.execute(
118 "INSERT OR REPLACE INTO remote_media (server, media_id, content_type, filename, fetched_ms, size, data) VALUES (?1,?2,?3,?4,?5,?6,?7)",
119 params![server, media_id, blob.content_type, blob.filename, now_ms, blob.data.len() as i64, blob.data],
120 )?;
121 loop {
122 let total: i64 = conn.query_row("SELECT COALESCE(SUM(size), 0) FROM remote_media", [], |r| r.get(0))?;
123 if total <= remote_cache_cap() {
124 return Ok(());
125 }
126 let gone = conn.execute(
127 "DELETE FROM remote_media WHERE (server, media_id) = (SELECT server, media_id FROM remote_media WHERE NOT (server = ?1 AND media_id = ?2) ORDER BY fetched_ms LIMIT 1)",
128 params![server, media_id],
129 )?;
130 if gone == 0 {
131 return Ok(());
132 }
133 }
134}
135
136pub fn create_pending(conn: &Connection, owner: i64, now_ms: i64) -> rusqlite::Result<Option<(String, i64)>> {
138 conn.execute("DELETE FROM media_pending WHERE expires_ms < ?1", [now_ms])?;
139 let open: i64 = conn.query_row("SELECT COUNT(*) FROM media_pending WHERE owner_user_id = ?1", [owner], |r| r.get(0))?;
140 if open >= 10 {
141 return Ok(None);
142 }
143 let mut raw = [0u8; 18];
144 rand::thread_rng().fill_bytes(&mut raw);
145 let id = base64::Engine::encode(&base64::engine::general_purpose::URL_SAFE_NO_PAD, raw);
146 let expires = now_ms + 24 * 3600 * 1000;
147 conn.execute("INSERT INTO media_pending (media_id, owner_user_id, expires_ms) VALUES (?1,?2,?3)", params![id, owner, expires])?;
148 Ok(Some((id, expires)))
149}
150
151pub fn put_reserved(conn: &Connection, owner: i64, media_id: &str, content_type: &str, filename: Option<&str>, data: &[u8], now_ms: i64) -> rusqlite::Result<bool> {
153 let n = conn.execute("DELETE FROM media_pending WHERE media_id = ?1 AND owner_user_id = ?2 AND expires_ms >= ?3", params![media_id, owner, now_ms])?;
154 if n == 0 {
155 return Ok(false);
156 }
157 conn.execute(
158 "INSERT INTO media (media_id, owner_user_id, content_type, filename, created_ms, data) VALUES (?1,?2,?3,?4,?5,?6)",
159 params![media_id, owner, content_type, filename, now_ms, data],
160 )?;
161 Ok(true)
162}
163
164pub fn is_pending(conn: &Connection, media_id: &str, now_ms: i64) -> bool {
166 conn.query_row("SELECT 1 FROM media_pending WHERE media_id = ?1 AND expires_ms >= ?2", params![media_id, now_ms], |_| Ok(())).is_ok()
167}
168
169pub fn multipart_body(blob: &Blob) -> (String, Vec<u8>) {
171 let mut raw = [0u8; 12];
172 rand::thread_rng().fill_bytes(&mut raw);
173 let boundary: String = raw.iter().map(|b| format!("{b:02x}")).collect();
174 let mut body = Vec::new();
175 body.extend_from_slice(format!("--{boundary}\r\nContent-Type: application/json\r\n\r\n{{}}\r\n--{boundary}\r\nContent-Type: {}\r\n", blob.content_type).as_bytes());
176 if let Some(f) = &blob.filename {
177 body.extend_from_slice(format!("Content-Disposition: attachment; filename=\"{}\"\r\n", f.replace(['"', '\\', '\r', '\n'], "_")).as_bytes());
178 }
179 body.extend_from_slice(b"\r\n");
180 body.extend_from_slice(&blob.data);
181 body.extend_from_slice(format!("\r\n--{boundary}--\r\n").as_bytes());
182 (format!("multipart/mixed; boundary={boundary}"), body)
183}
184
185pub enum Part {
187 Data(Blob),
189 Redirect(String),
191}
192
193fn find(hay: &[u8], needle: &[u8], from: usize) -> Option<usize> {
194 if needle.is_empty() || hay.len() < needle.len() || from > hay.len() - needle.len() {
195 return None;
196 }
197 (from..=hay.len() - needle.len()).find(|&i| &hay[i..i + needle.len()] == needle)
198}
199
200pub fn parse_multipart(content_type: &str, body: &[u8]) -> Option<Part> {
202 let b = content_type.split(';').filter_map(|p| p.trim().strip_prefix("boundary=")).next()?.trim_matches('"');
203 let delim = format!("--{b}").into_bytes();
204 let first = find(body, &delim, 0)? + delim.len();
205 let h1_end = find(body, b"\r\n\r\n", first)? + 4;
206 let next = find(body, &[b"\r\n".as_slice(), &delim].concat(), h1_end)? + 2 + delim.len();
207 let h2_end = find(body, b"\r\n\r\n", next)? + 4;
208 let headers = std::str::from_utf8(&body[next..h2_end]).ok()?;
209 let end = find(body, &[b"\r\n".as_slice(), &delim].concat(), h2_end)?;
210 let mut content_type = "application/octet-stream".to_string();
211 let (mut filename, mut location) = (None, None);
212 for line in headers.split("\r\n") {
213 let Some((k, v)) = line.split_once(':') else { continue };
214 let v = v.trim();
215 match k.trim().to_ascii_lowercase().as_str() {
216 "content-type" if !v.is_empty() && v.len() <= 128 => content_type = v.to_string(),
217 "location" => location = Some(v.to_string()),
218 "content-disposition" => {
219 filename = v.split(';').filter_map(|p| p.trim().strip_prefix("filename=")).next().map(|f| f.trim_matches('"').to_string());
220 }
221 _ => {}
222 }
223 }
224 if let Some(l) = location {
225 return Some(Part::Redirect(l));
226 }
227 Some(Part::Data(Blob { content_type, filename, data: body[h2_end..end].to_vec() }))
228}
229
230#[cfg(feature = "media-thumbnails")]
233pub fn thumbnail(data: &[u8], width: u32, height: u32, crop: bool) -> Option<(Vec<u8>, &'static str)> {
234 use image::{ImageFormat, ImageReader};
235 let reader = ImageReader::new(std::io::Cursor::new(data)).with_guessed_format().ok()?;
236 let fmt = reader.format()?;
237 let (w, h) = reader.into_dimensions().ok()?;
238 if u64::from(w) * u64::from(h) > 40_000_000 || (w <= width && h <= height) {
239 return None;
240 }
241 let img = image::load_from_memory_with_format(data, fmt).ok()?;
242 let out = if crop { img.resize_to_fill(width, height, image::imageops::FilterType::Triangle) } else { img.resize(width, height, image::imageops::FilterType::Triangle) };
243 let (format, mime) = if fmt == ImageFormat::Jpeg { (ImageFormat::Jpeg, "image/jpeg") } else { (ImageFormat::Png, "image/png") };
244 let mut buf = std::io::Cursor::new(Vec::new());
245 if format == ImageFormat::Jpeg { out.to_rgb8().write_to(&mut buf, format).ok()?; } else { out.write_to(&mut buf, format).ok()?; }
246 Some((buf.into_inner(), mime))
247}
248
249#[cfg(not(feature = "media-thumbnails"))]
250pub fn thumbnail(_data: &[u8], _width: u32, _height: u32, _crop: bool) -> Option<(Vec<u8>, &'static str)> {
251 None
252}
253
254#[cfg(test)]
255mod tests {
256 use super::*;
257
258 #[test]
259 fn multipart_round_trips_and_a_location_part_is_a_redirect() {
260 let blob = Blob { content_type: "text/plain".into(), filename: Some("a.txt".into()), data: b"hello\r\n--x not a boundary".to_vec() };
261 let (ct, body) = multipart_body(&blob);
262 let Some(Part::Data(got)) = parse_multipart(&ct, &body) else { panic!("data part") };
263 assert_eq!((got.content_type.as_str(), got.filename.as_deref(), got.data), ("text/plain", Some("a.txt"), blob.data));
264 let redirect = b"--B\r\nContent-Type: application/json\r\n\r\n{}\r\n--B\r\nLocation: https://cdn.example/x\r\n\r\n\r\n--B--\r\n";
265 assert!(matches!(parse_multipart("multipart/mixed; boundary=B", redirect), Some(Part::Redirect(u)) if u == "https://cdn.example/x"));
266 assert!(parse_multipart("multipart/mixed; boundary=B", b"junk").is_none());
267 }
268
269 #[test]
270 fn remote_cache_expires_and_evicts_the_oldest_over_the_cap() {
271 let c = Connection::open_in_memory().unwrap();
272 create_media_schema(&c).unwrap();
273 let blob = |n: usize| Blob { content_type: "a/b".into(), filename: None, data: vec![0; n] };
274 std::env::set_var("M4A_MEDIA_REMOTE_CACHE_BYTES", "100");
275 put_remote(&c, "s", "one", &blob(60), 1_000).unwrap();
276 put_remote(&c, "s", "two", &blob(60), 2_000).unwrap();
277 std::env::remove_var("M4A_MEDIA_REMOTE_CACHE_BYTES");
278 assert!(get_remote(&c, "s", "one", 3_000).unwrap().is_none(), "oldest evicted");
279 assert!(get_remote(&c, "s", "two", 3_000).unwrap().is_some());
280 assert!(get_remote(&c, "s", "two", 2_000 + remote_ttl_ms() + 1).unwrap().is_none(), "expired");
281 }
282
283 #[cfg(feature = "media-thumbnails")]
284 #[test]
285 fn images_are_scaled_or_cropped_and_other_bytes_are_left_alone() {
286 let img = image::RgbImage::from_fn(100, 60, |x, y| image::Rgb([x as u8, y as u8, 7]));
287 let mut png = std::io::Cursor::new(Vec::new());
288 img.write_to(&mut png, image::ImageFormat::Png).unwrap();
289 let png = png.into_inner();
290 let (bytes, mime) = thumbnail(&png, 32, 32, false).unwrap();
291 let t = image::load_from_memory(&bytes).unwrap();
292 assert_eq!((mime, t.width(), t.height()), ("image/png", 32, 19));
293 let (bytes, _) = thumbnail(&png, 32, 32, true).unwrap();
294 let t = image::load_from_memory(&bytes).unwrap();
295 assert_eq!((t.width(), t.height()), (32, 32));
296 assert!(thumbnail(&png, 200, 200, false).is_none(), "never enlarged");
297 assert!(thumbnail(b"ciphertext, not an image", 32, 32, false).is_none());
298 }
299
300 #[test]
301 fn put_get_and_ttl() {
302 let c = Connection::open_in_memory().unwrap();
303 create_media_schema(&c).unwrap();
304 let id = put(&c, 1, "application/octet-stream", Some("a.bin"), &[1, 2, 3], 1_000).unwrap();
305 assert_eq!(get(&c, &id).unwrap().unwrap().data, vec![1, 2, 3]);
306 assert_eq!(purge_expired(&c, 1_500, 1_000).unwrap(), 0);
307 assert_eq!(purge_expired(&c, 5_000, 1_000).unwrap(), 1);
308 assert!(get(&c, &id).unwrap().is_none());
309 }
310}