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