1#[cfg(feature = "e2e-encryption")]
16use std::io::Read;
17use std::{
18 fmt::{Debug, Formatter},
19 sync::Arc,
20};
21
22use api::{
23 download::{
24 encrypted::DownloadAndScanEncryptedMediaRequest, unencrypted::DownloadAndScanMediaRequest,
25 },
26 public_server_key::PublicServerKeyRequest,
27};
28use matrix_sdk::{
29 BoxFuture, Client, Error, IdParseError,
30 encryption::vodozemac::pk_encryption::Message,
31 locks::Mutex,
32 media::{MediaFetcher, MediaRequestParameters},
33 ruma::events::room::MediaSource,
34};
35#[cfg(feature = "e2e-encryption")]
36use matrix_sdk_crypto::AttachmentDecryptor;
37use matrix_sdk_crypto::olm::Curve25519PublicKey;
38use ruma::{
39 events::room::EncryptedFile,
40 serde::{Base64, base64::Standard},
41};
42use serde::{Deserialize, Serialize};
43use tracing::trace;
44
45#[cfg(feature = "uniffi")]
46uniffi::setup_scaffolding!();
47
48pub use crate::api::scan::MediaScanResponse;
49use crate::api::{
50 DownloadAndScanMediaResponse,
51 scan::{encrypted::EncryptedMediaScanRequest, unencrypted::MediaScanRequest},
52};
53
54mod api;
55
56#[derive(Debug)]
58pub struct ContentScanner {
59 scanner_url: String,
60 public_server_key: Arc<Mutex<Option<String>>>,
61}
62
63impl ContentScanner {
64 pub fn new(scanner_url: impl Into<String>) -> Self {
66 Self { scanner_url: scanner_url.into(), public_server_key: Arc::new(Mutex::new(None)) }
67 }
68
69 pub(crate) async fn fetch_public_server_key(&self, client: &Client) -> Result<String, Error> {
70 let response = client.send(PublicServerKeyRequest::new(self.scanner_url.clone())).await?;
71 Ok(response.public_key)
72 }
73
74 async fn get_or_fetch_public_server_key(&self, client: &Client) -> Option<Curve25519PublicKey> {
75 let public_server_key =
76 if let Some(public_server_key) = (*self.public_server_key.lock()).clone() {
77 trace!("Using cached public server key");
78 Some(public_server_key)
79 } else {
80 trace!("Using cached public server key");
81 let ret = self.fetch_public_server_key(client).await.ok();
82
83 if let Some(public_server_key) = &ret {
84 trace!("Saved new public server key");
85 let mut guard = self.public_server_key.lock();
86 let _ = guard.insert(public_server_key.clone());
87 }
88
89 ret
90 };
91
92 public_server_key.and_then(|key| Curve25519PublicKey::from_base64(&key).ok())
93 }
94
95 pub(crate) async fn get_media(
96 &self,
97 client: &Client,
98 media_source: &MediaSource,
99 ) -> Result<DownloadAndScanMediaResponse, Error> {
100 match &media_source {
101 MediaSource::Encrypted(encrypted) => {
102 let public_server_key = self.get_or_fetch_public_server_key(client).await;
104
105 Ok(client
106 .send(DownloadAndScanEncryptedMediaRequest::new(
107 self.scanner_url.clone(),
108 public_server_key,
109 *encrypted.clone(),
110 ))
111 .await?)
112 }
113 MediaSource::Plain(mxc) => {
114 let (server_name, media_id) =
115 mxc.parts().map_err(|e| Error::Identifier(IdParseError::InvalidMxcUri(e)))?;
116 Ok(client
117 .send(DownloadAndScanMediaRequest::new(
118 &self.scanner_url,
119 server_name.as_str(),
120 media_id,
121 ))
122 .await?)
123 }
124 }
125 }
126
127 pub async fn scan(
130 &self,
131 client: &Client,
132 media_source: &MediaSource,
133 ) -> Result<MediaScanResponse, Error> {
134 match &media_source {
135 MediaSource::Encrypted(encrypted) => {
136 let public_server_key = self.get_or_fetch_public_server_key(client).await;
138
139 Ok(client
140 .send(EncryptedMediaScanRequest::new(
141 self.scanner_url.clone(),
142 public_server_key,
143 *encrypted.clone(),
144 ))
145 .await?)
146 }
147 MediaSource::Plain(mxc) => {
148 let (server_name, media_id) =
149 mxc.parts().map_err(|e| Error::Identifier(IdParseError::InvalidMxcUri(e)))?;
150 Ok(client
151 .send(MediaScanRequest::new(
152 self.scanner_url.clone(),
153 server_name.to_string(),
154 media_id.to_owned(),
155 ))
156 .await?)
157 }
158 }
159 }
160}
161
162#[derive(Debug, Clone, Serialize)]
163struct EncryptedBody {
164 ciphertext: String,
165 mac: String,
166 ephemeral: String,
167}
168
169impl From<Message> for EncryptedBody {
170 fn from(value: Message) -> Self {
171 Self {
172 ciphertext: Base64::<Standard>::new(value.ciphertext).to_string(),
173 mac: Base64::<Standard>::new(value.mac).to_string(),
174 ephemeral: value.ephemeral_key.to_base64(),
175 }
176 }
177}
178
179#[derive(Debug, Clone, Serialize)]
180pub(crate) struct EncryptedFileRequest {
181 #[serde(skip_serializing_if = "Option::is_none")]
182 pub file: Option<EncryptedFile>,
183 #[serde(skip_serializing_if = "Option::is_none")]
184 pub encrypted_body: Option<EncryptedBody>,
185}
186
187impl EncryptedFileRequest {
188 pub(crate) fn from_file_info(file_info: EncryptedFile) -> Self {
189 Self { file: Some(file_info), encrypted_body: None }
190 }
191
192 pub(crate) fn from_encrypted_body(encrypted_body: EncryptedBody) -> Self {
193 Self { file: None, encrypted_body: Some(encrypted_body) }
194 }
195}
196
197pub struct ContentScannerMediaFetcher {
199 pub content_scanner: Arc<ContentScanner>,
200}
201
202impl ContentScannerMediaFetcher {
203 pub fn new(scanner_url: impl Into<String>) -> Self {
205 Self { content_scanner: Arc::new(ContentScanner::new(scanner_url.into())) }
206 }
207
208 pub fn with_content_scanner(content_scanner: Arc<ContentScanner>) -> Self {
209 Self { content_scanner }
210 }
211}
212
213impl MediaFetcher for ContentScannerMediaFetcher {
214 fn fetch_media_content<'a>(
215 &'a self,
216 client: &'a Client,
217 request: &'a MediaRequestParameters,
218 ) -> BoxFuture<'a, matrix_sdk::Result<Vec<u8>, Error>> {
219 Box::pin(async move {
220 let content = self.content_scanner.get_media(client, &request.source).await?.content;
221 #[cfg(feature = "e2e-encryption")]
222 let content = {
223 match &request.source {
224 MediaSource::Encrypted(file) => {
225 let content_len = content.len();
226 let mut cursor = std::io::Cursor::new(content);
227 let mut reader =
228 AttachmentDecryptor::new(&mut cursor, file.as_ref().clone().into())?;
229
230 let mut decrypted = Vec::with_capacity(content_len);
233
234 reader.read_to_end(&mut decrypted)?;
235
236 decrypted
237 }
238 MediaSource::Plain(_) => content,
239 }
240 };
241 Ok(content)
242 })
243 }
244}
245
246impl Debug for ContentScannerMediaFetcher {
247 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
248 f.write_str("ContentScannerMediaFetcher")
249 }
250}
251
252#[derive(Debug, Deserialize)]
254pub struct ContentScannerError {
255 pub info: String,
256 pub reason: ErrorReason,
257}
258
259#[allow(non_camel_case_types)]
261#[cfg_attr(feature = "uniffi", derive(uniffi::Enum))]
262#[derive(Clone, Debug, Deserialize)]
263pub enum ErrorReason {
264 MCS_MALFORMED_JSON,
266 MCS_MEDIA_FAILED_TO_DECRYPT,
268 M_MISSING_TOKEN,
270 M_UNKNOWN_TOKEN,
272 M_NOT_FOUND,
274 MCS_MEDIA_NOT_CLEAN,
276 MCS_MIME_TYPE_FORBIDDEN,
278 MCS_BAD_DECRYPTION,
280 M_UNKNOWN,
282 MCS_MEDIA_REQUEST_FAILED,
284}
285
286#[cfg(test)]
287mod tests {
288 use std::ops::Not;
289
290 use assert_matches2::assert_matches;
291 use matrix_sdk::{HttpError, RumaApiError, test_utils::mocks::MatrixMockServer};
292 use matrix_sdk_test::async_test;
293 use ruma::{
294 api::{
295 MatrixVersion,
296 error::{ErrorBody, FromHttpResponseError},
297 },
298 events::room::{
299 EncryptedFile, EncryptedFileHash, EncryptedFileHashes, EncryptedFileInfo, MediaSource,
300 V2EncryptedFileInfo,
301 },
302 exports::{http::StatusCode, serde_json::json},
303 owned_mxc_uri,
304 serde::Base64,
305 };
306 use serde::Deserialize;
307 use wiremock::{
308 Mock, MockServer, ResponseTemplate,
309 matchers::{header_exists, method, path, path_regex},
310 };
311
312 use crate::{ContentScanner, ContentScannerError, ErrorReason};
313
314 #[async_test]
315 async fn test_fetch_public_key() {
316 let server = MatrixMockServer::new().await;
317 let client =
318 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
319
320 let content_scanner_server = MockServer::start().await;
321 Mock::given(method("GET"))
322 .and(path("/_matrix/media_proxy/unstable/public_key"))
323 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
324 "public_key": "1234567890"
325 })))
326 .mount(&content_scanner_server)
327 .await;
328
329 let content_scanner = ContentScanner::new(content_scanner_server.uri());
330 content_scanner.fetch_public_server_key(&client).await.expect("Load public key");
331 }
332
333 #[async_test]
334 async fn test_get_media() {
335 let server = MatrixMockServer::new().await;
336 let client =
337 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
338
339 let content_scanner_server = MockServer::start().await;
340 Mock::given(method("GET"))
341 .and(path_regex(r"/_matrix/media_proxy/unstable/download/.+/.+"))
342 .and(header_exists("Authorization"))
343 .respond_with(
344 ResponseTemplate::new(200).set_body_bytes(vec![1, 2, 3, 4, 5, 6, 7, 8, 9]),
345 )
346 .mount(&content_scanner_server)
347 .await;
348
349 let content_scanner = ContentScanner::new(content_scanner_server.uri());
350 let media_source =
351 MediaSource::Plain(owned_mxc_uri!("mxc://matrix.org/RhfpOXOzAwzkuqcmbgMwQUrJ"));
352 content_scanner.get_media(&client, &media_source).await.expect("Get media");
353 }
354
355 #[async_test]
356 async fn test_get_media_unsupported() {
357 let server = MatrixMockServer::new().await;
358 let client =
359 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
360
361 let content_scanner_server = MockServer::start().await;
362 Mock::given(method("GET"))
363 .and(path_regex(r"/_matrix/media_proxy/unstable/download/.+/.+"))
364 .and(header_exists("Authorization"))
365 .respond_with(ResponseTemplate::new(403).set_body_json(json!({
366 "reason": "MCS_MIME_TYPE_FORBIDDEN",
367 "info": "File type: application/octet-stream not allowed",
368 })))
369 .mount(&content_scanner_server)
370 .await;
371
372 let content_scanner = ContentScanner::new(content_scanner_server.uri());
373 let media_source =
374 MediaSource::Plain(owned_mxc_uri!("mxc://matrix.org/ckTaStcNnFXLzKApkBmgRDoC"));
375 let err =
376 content_scanner.get_media(&client, &media_source).await.expect_err("Get media error");
377 let client_error = err.as_client_api_error().expect("Get client error");
378 assert_eq!(client_error.status_code, StatusCode::FORBIDDEN);
379 assert_eq!(
380 client_error.to_string(),
381 "[403] {\"info\":\"File type: application/octet-stream not allowed\",\"reason\":\"MCS_MIME_TYPE_FORBIDDEN\"}"
382 );
383 }
384
385 #[async_test]
386 async fn test_get_encrypted_media() {
387 let server = MatrixMockServer::new().await;
388 let client =
389 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
390
391 let content_scanner_server = MockServer::start().await;
392 Mock::given(method("POST"))
393 .and(path("/_matrix/media_proxy/unstable/download_encrypted"))
394 .and(header_exists("Authorization"))
395 .respond_with(
396 ResponseTemplate::new(200).set_body_bytes(vec![1, 2, 3, 4, 5, 6, 7, 8, 9]),
397 )
398 .mount(&content_scanner_server)
399 .await;
400
401 let content_scanner = ContentScanner::new(content_scanner_server.uri());
402 let file_info = EncryptedFileInfo::V2(V2EncryptedFileInfo::new(
403 Base64::parse("9lpOscZyMOZRCF3v867nPPo3WPNMZt9JXMsuYiWRszc".as_bytes()).expect("k"),
404 Base64::parse("czvdfKSjfLEAAAAAAAAAAA".as_bytes()).expect("iv"),
405 ));
406 let mut hashes = EncryptedFileHashes::new();
407 hashes.insert(EncryptedFileHash::Sha256(
408 Base64::parse("SBbJ3hINT2LgwXK8ev82enjnhubUy5UuKGDF3SezAhs".as_bytes()).expect("hash"),
409 ));
410 let media_source = MediaSource::Encrypted(Box::new(EncryptedFile::new(
411 owned_mxc_uri!(
412 "mxc://element.io/b50f38aa8ae820c75992370e4e944a045481e3932057062074730676224"
413 ),
414 file_info,
415 hashes,
416 )));
417 content_scanner.get_media(&client, &media_source).await.expect("Get media");
418 }
419
420 #[async_test]
421 async fn test_get_encrypted_media_unsupported() {
422 let server = MatrixMockServer::new().await;
423 let client =
424 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
425
426 let content_scanner_server = MockServer::start().await;
427
428 Mock::given(method("POST"))
429 .and(path("/_matrix/media_proxy/unstable/download_encrypted"))
430 .and(header_exists("Authorization"))
431 .respond_with(ResponseTemplate::new(403).set_body_json(json!({
432 "reason": "MCS_MIME_TYPE_FORBIDDEN",
433 "info": "File type: application/octet-stream not allowed",
434 })))
435 .mount(&content_scanner_server)
436 .await;
437
438 let content_scanner = ContentScanner::new(content_scanner_server.uri());
439 let file_info = EncryptedFileInfo::V2(V2EncryptedFileInfo::new(
440 Base64::parse("tdHdCI5mc-g29IYfhYx2wkA5o-bILP9-nXY6Np1uSnM".as_bytes()).expect("k"),
441 Base64::parse("IBFdH65KqhoAAAAAAAAAAA".as_bytes()).expect("iv"),
442 ));
443 let mut hashes = EncryptedFileHashes::new();
444 hashes.insert(EncryptedFileHash::Sha256(
445 Base64::parse("HSkkamvMSvF3Q30HInorh0ccPrxjgu+wp1vyUOmov/8".as_bytes()).expect("hash"),
446 ));
447 let media_source = MediaSource::Encrypted(Box::new(EncryptedFile::new(
448 owned_mxc_uri!("mxc://matrix.org/WlfuejQQdpvWiWVpAGwfIKJL"),
449 file_info,
450 hashes,
451 )));
452 let err = content_scanner
453 .get_media(&client, &media_source)
454 .await
455 .expect_err("Invalid type error");
456 let client_error = err.as_client_api_error().expect("Invalid error");
457 assert_eq!(client_error.status_code, StatusCode::FORBIDDEN);
458 assert_eq!(
459 client_error.to_string(),
460 "[403] {\"info\":\"File type: application/octet-stream not allowed\",\"reason\":\"MCS_MIME_TYPE_FORBIDDEN\"}"
461 );
462 }
463
464 #[async_test]
465 async fn test_scan_media() {
466 let server = MatrixMockServer::new().await;
467 let client =
468 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
469
470 let content_scanner_server = MockServer::start().await;
471 Mock::given(method("GET"))
472 .and(path_regex(r"/_matrix/media_proxy/unstable/scan/.+/.+"))
473 .and(header_exists("Authorization"))
474 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
475 "clean": true,
476 "info": "All clear!"
477 })))
478 .mount(&content_scanner_server)
479 .await;
480
481 let content_scanner = ContentScanner::new(content_scanner_server.uri());
482 let media_source =
483 MediaSource::Plain(owned_mxc_uri!("mxc://matrix.org/RhfpOXOzAwzkuqcmbgMwQUrJ"));
484 let response = content_scanner.scan(&client, &media_source).await.expect("Get media");
485 assert!(response.clean);
486 }
487
488 #[async_test]
489 async fn test_scan_encrypted_media() {
490 let server = MatrixMockServer::new().await;
491 let client =
492 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
493
494 let content_scanner_server = MockServer::start().await;
495 Mock::given(method("POST"))
496 .and(path("/_matrix/media_proxy/unstable/scan_encrypted"))
497 .and(header_exists("Authorization"))
498 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
499 "clean": true,
500 "info": "All clear!"
501 })))
502 .mount(&content_scanner_server)
503 .await;
504
505 let content_scanner = ContentScanner::new(content_scanner_server.uri());
506 let file_info = EncryptedFileInfo::V2(V2EncryptedFileInfo::new(
507 Base64::parse("9lpOscZyMOZRCF3v867nPPo3WPNMZt9JXMsuYiWRszc".as_bytes()).expect("k"),
508 Base64::parse("czvdfKSjfLEAAAAAAAAAAA".as_bytes()).expect("iv"),
509 ));
510 let mut hashes = EncryptedFileHashes::new();
511 hashes.insert(EncryptedFileHash::Sha256(
512 Base64::parse("SBbJ3hINT2LgwXK8ev82enjnhubUy5UuKGDF3SezAhs".as_bytes()).expect("hash"),
513 ));
514 let media_source = MediaSource::Encrypted(Box::new(EncryptedFile::new(
515 owned_mxc_uri!(
516 "mxc://element.io/b50f38aa8ae820c75992370e4e944a045481e3932057062074730676224"
517 ),
518 file_info,
519 hashes,
520 )));
521 let response = content_scanner.scan(&client, &media_source).await.expect("Get media");
522 assert!(response.clean);
523 }
524
525 #[async_test]
526 async fn test_scan_media_unsupported() {
527 let server = MatrixMockServer::new().await;
528 let client =
529 server.client_builder().server_versions(vec![MatrixVersion::V1_11]).build().await;
530
531 let content_scanner_server = MockServer::start().await;
532 Mock::given(method("GET"))
533 .and(path_regex(r"/_matrix/media_proxy/unstable/scan/.+/.+"))
534 .and(header_exists("Authorization"))
535 .respond_with(
536 ResponseTemplate::new(200).set_body_json(json!({
538 "clean": false,
539 "info": "***VIRUS DETECTED***"
540 })),
541 )
542 .mount(&content_scanner_server)
543 .await;
544
545 let content_scanner = ContentScanner::new(content_scanner_server.uri());
546 let media_source =
547 MediaSource::Plain(owned_mxc_uri!("mxc://matrix.org/RhfpOXOzAwzkuqcmbgMwQUrJ"));
548 let response = content_scanner.scan(&client, &media_source).await.expect("Get media");
549 assert!(response.clean.not());
550 }
551
552 #[test]
553 fn test_error_mapping() {
554 let error = HttpError::Api(Box::new(FromHttpResponseError::Server(
555 RumaApiError::MatrixError(ruma::api::error::Error::new(
556 StatusCode::FORBIDDEN,
557 ErrorBody::Json(json!({
558 "info": "***VIRUS DETECTED***",
559 "reason": "MCS_MEDIA_NOT_CLEAN"
560 })),
561 )),
562 )));
563 let api_error = error.as_client_api_error().expect("error as api error");
564 assert_eq!(
565 api_error.to_string(),
566 "[403] {\"info\":\"***VIRUS DETECTED***\",\"reason\":\"MCS_MEDIA_NOT_CLEAN\"}"
567 );
568 assert_matches!(&api_error.body, ErrorBody::Json(json_body));
569 let content_scanner_error =
570 ContentScannerError::deserialize(json_body).expect("deserialize");
571 assert_eq!(content_scanner_error.info, "***VIRUS DETECTED***");
572 assert_matches!(content_scanner_error.reason, ErrorReason::MCS_MEDIA_NOT_CLEAN);
573 }
574}