1use std::{collections::HashSet, pin::Pin};
86
87use async_stream::try_stream;
88use futures_util::{Stream, StreamExt as _};
89use matrix_sdk_base::{RoomStateFilter, deserialized_responses::TimelineEvent};
90use matrix_sdk_search::error::IndexError;
91#[cfg(doc)]
92use matrix_sdk_search::index::RoomIndex;
93pub use matrix_sdk_search::index::SearchResult;
94use ruma::{OwnedEventId, OwnedRoomId};
95
96use crate::{Client, Room};
97
98const SEARCH_RESULTS_PAGE_SIZE: usize = 100;
101
102type RoomResultStream = Pin<Box<dyn Stream<Item = Result<(f32, OwnedEventId), IndexError>> + Send>>;
105
106struct RoomStreamCursor {
109 room_id: OwnedRoomId,
111
112 stream: RoomResultStream,
114
115 next_result: Option<(f32, OwnedEventId)>,
120}
121
122impl Room {
123 pub async fn search(
126 &self,
127 query: &str,
128 max_number_of_results: usize,
129 pagination_offset: Option<usize>,
130 ) -> Result<SearchResult, IndexError> {
131 let mut search_index_guard = self.client.search_index().lock().await;
132 search_index_guard.search(query, max_number_of_results, pagination_offset, self.room_id())
133 }
134}
135
136#[derive(thiserror::Error, Debug)]
139pub enum SearchError {
140 #[error(transparent)]
142 IndexError(#[from] IndexError),
143 #[error(transparent)]
145 EventLoadError(#[from] crate::Error),
146}
147
148impl Room {
149 pub fn search_messages(
153 &self,
154 query: String,
155 ) -> impl Stream<Item = Result<SearchResult, IndexError>> + use<> {
156 let room = self.clone();
157
158 try_stream! {
161 let mut offset = 0;
162 loop {
163 let page = room.search(&query, SEARCH_RESULTS_PAGE_SIZE, Some(offset)).await?;
164 if page.events.is_empty() {
165 break;
166 }
167 offset += page.events.len();
168 yield page;
169 }
170 }
171 }
172
173 pub fn search_messages_events(
177 &self,
178 query: String,
179 ) -> impl Stream<Item = Result<Vec<TimelineEvent>, SearchError>> + use<> {
180 let room = self.clone();
181
182 try_stream! {
183 let mut pages = Box::pin(room.search_messages(query));
184
185 while let Some(page) = pages.next().await {
186 let page = page?;
187 let mut events = Vec::with_capacity(page.events.len());
188
189 for (_score, event_id) in page.events {
190 events.push(room.load_or_fetch_event(&event_id, None).await?);
191 }
192
193 yield events;
194 }
195 }
196 }
197}
198
199#[derive(Debug)]
202pub struct GlobalSearchBuilder {
203 client: Client,
204
205 query: String,
207
208 room_set: Vec<Room>,
210}
211
212impl GlobalSearchBuilder {
213 fn new(client: Client, query: String) -> Self {
215 let room_set = client.rooms_filtered(RoomStateFilter::JOINED);
216 Self { client, query, room_set }
217 }
218
219 pub async fn only_dm_rooms(mut self) -> Result<Self, crate::Error> {
221 let mut to_remove = HashSet::new();
222 for room in &self.room_set {
223 if !room.compute_is_dm().await? {
224 to_remove.insert(room.room_id().to_owned());
225 }
226 }
227 self.room_set.retain(|room| !to_remove.contains(room.room_id()));
228 Ok(self)
229 }
230
231 pub async fn no_dms(mut self) -> Result<Self, crate::Error> {
233 let mut to_remove = HashSet::new();
234 for room in &self.room_set {
235 if room.compute_is_dm().await? {
236 to_remove.insert(room.room_id().to_owned());
237 }
238 }
239 self.room_set.retain(|room| !to_remove.contains(room.room_id()));
240 Ok(self)
241 }
242
243 pub fn build(
247 self,
248 ) -> impl Stream<Item = Result<Vec<(OwnedRoomId, f32, OwnedEventId)>, IndexError>> {
249 let query = self.query;
250 let rooms = self.room_set;
251
252 try_stream! {
253 let mut cursors: Vec<RoomStreamCursor> = Vec::with_capacity(rooms.len());
256 for room in rooms {
257 let room_id = room.room_id().to_owned();
258
259 let stream = Box::pin(Self::flatten_pages(room.search_messages(query.clone())));
260 cursors.push(RoomStreamCursor { room_id, stream, next_result: None });
261 }
262
263 for cursor in &mut cursors {
264 cursor.next_result = match cursor.stream.next().await {
265 Some(result) => Some(result?),
266 None => None,
267 };
268 }
269
270 let mut page = Vec::with_capacity(SEARCH_RESULTS_PAGE_SIZE);
271
272 loop {
273 let best = cursors
275 .iter()
276 .enumerate()
277 .filter_map(|(index, cursor)| {
278 cursor.next_result.as_ref().map(|(score, _)| (index, *score))
279 })
280 .max_by(|(_, a), (_, b)| a.total_cmp(b));
281
282 let Some((index, _)) = best else {
283 break;
285 };
286
287 let cursor = &mut cursors[index];
288 let (score, event_id) =
289 cursor.next_result.take().expect("the chosen room must have a next result");
290 let room_id = cursor.room_id.clone();
291
292 cursor.next_result = match cursor.stream.next().await {
294 Some(result) => Some(result?),
295 None => None,
296 };
297
298 page.push((room_id, score, event_id));
299 if page.len() == SEARCH_RESULTS_PAGE_SIZE {
300 yield std::mem::take(&mut page);
301 }
302 }
303
304 if !page.is_empty() {
305 yield page;
306 }
307 }
308 }
309
310 pub fn build_events(
314 self,
315 ) -> impl Stream<Item = Result<Vec<(OwnedRoomId, TimelineEvent)>, SearchError>> {
316 let client = self.client.clone();
317 let pages = self.build();
318
319 try_stream! {
320 let mut pages = Box::pin(pages);
321
322 while let Some(page) = pages.next().await {
323 let page = page?;
324 let mut events = Vec::with_capacity(page.len());
325
326 for (room_id, _score, event_id) in page {
327 let Some(room) = client.get_room(&room_id) else {
328 continue;
329 };
330
331 events.push((room_id, room.load_or_fetch_event(&event_id, None).await?));
332 }
333
334 yield events;
335 }
336 }
337 }
338
339 fn flatten_pages(
342 pages: impl Stream<Item = Result<SearchResult, IndexError>>,
343 ) -> impl Stream<Item = Result<(f32, OwnedEventId), IndexError>> {
344 try_stream! {
345 let mut pages = Box::pin(pages);
346
347 while let Some(page) = pages.next().await {
348 for result in page?.events {
349 yield result;
350 }
351 }
352 }
353 }
354}
355
356impl Client {
357 pub fn search_messages(&self, query: String) -> GlobalSearchBuilder {
360 GlobalSearchBuilder::new(self.clone(), query)
361 }
362}
363
364#[cfg(test)]
365mod tests {
366 use std::time::Duration;
367
368 use futures_util::TryStreamExt as _;
369 use matrix_sdk_search::index::SearchResult;
370 use matrix_sdk_test::{BOB, JoinedRoomBuilder, async_test, event_factory::EventFactory};
371 use ruma::{OwnedEventId, OwnedRoomId, event_id, room_id, user_id};
372
373 use crate::{sleep::sleep, test_utils::mocks::MatrixMockServer};
374
375 #[async_test]
376 async fn test_room_message_search() {
377 let server = MatrixMockServer::new().await;
378 let client = server.client_builder().build().await;
379
380 let event_cache = client.event_cache();
381 event_cache.subscribe().unwrap();
382
383 let room_id = room_id!("!room_id:localhost");
384 let room = server.sync_joined_room(&client, room_id).await;
385
386 let f = EventFactory::new().room(room_id).sender(user_id!("@user_id:localhost"));
387
388 let event_id = event_id!("$event_id:localhost");
389
390 server
391 .sync_room(
392 &client,
393 JoinedRoomBuilder::new(room_id)
394 .add_timeline_event(f.text_msg("hello world").event_id(event_id)),
395 )
396 .await;
397
398 sleep(Duration::from_millis(200)).await;
400
401 {
403 let results: Vec<SearchResult> =
404 room.search_messages("search query".to_owned()).try_collect().await.unwrap();
405 assert!(results.is_empty());
406 }
407
408 {
410 let results: Vec<SearchResult> =
411 room.search_messages("world".to_owned()).try_collect().await.unwrap();
412 assert_eq!(results[0].events.len(), 1);
413 assert_eq!(results[0].events[0].1, event_id);
414 }
415
416 {
418 let results: Vec<_> =
419 room.search_messages_events("world".to_owned()).try_concat().await.unwrap();
420 assert_eq!(results.len(), 1);
421 assert_eq!(results[0].event_id().unwrap(), event_id);
422 }
423 }
424
425 #[async_test]
426 async fn test_global_message_search() {
427 let server = MatrixMockServer::new().await;
428 let client = server.client_builder().build().await;
429
430 let event_cache = client.event_cache();
431 event_cache.subscribe().unwrap();
432
433 let room_id1 = room_id!("!r1:localhost");
434 let room_id2 = room_id!("!r2:localhost");
435
436 let f = EventFactory::new().sender(user_id!("@user_id:localhost"));
437
438 let result_event_id1 = event_id!("$result1:localhost");
439 let result_event_id2 = event_id!("$result2:localhost");
440
441 server
442 .mock_sync()
443 .ok_and_run(&client, |sync_builder| {
444 sync_builder
445 .add_joined_room(
446 JoinedRoomBuilder::new(room_id1)
447 .add_timeline_event(
448 f.text_msg("hello world").room(room_id1).event_id(result_event_id1),
449 )
450 .add_timeline_event(f.text_msg("hello back").room(room_id1)),
451 )
452 .add_joined_room(JoinedRoomBuilder::new(room_id2).add_timeline_event(
453 f.text_msg("it's a mad world").room(room_id2).event_id(result_event_id2),
454 ));
455 })
456 .await;
457
458 sleep(Duration::from_millis(200)).await;
460
461 {
463 let results: Vec<(OwnedRoomId, f32, OwnedEventId)> = client
464 .search_messages("search query".to_owned())
465 .build()
466 .try_concat()
467 .await
468 .unwrap();
469 assert!(results.is_empty());
470 }
471
472 {
474 let results: Vec<(OwnedRoomId, f32, OwnedEventId)> =
475 client.search_messages("world".to_owned()).build().try_concat().await.unwrap();
476 assert_eq!(results.len(), 2);
477 assert!(results.iter().any(|(room_id, _, event_id)| {
480 room_id == room_id1 && event_id == result_event_id1
481 }));
482 assert!(results.iter().any(|(room_id, _, event_id)| {
483 room_id == room_id2 && event_id == result_event_id2
484 }));
485 }
486
487 {
489 let results: Vec<_> = client
490 .search_messages("world".to_owned())
491 .build_events()
492 .try_concat()
493 .await
494 .unwrap();
495 assert_eq!(results.len(), 2);
496 assert!(results.iter().any(|(room_id, event)| {
499 room_id == room_id1 && event.event_id() == Some(result_event_id1)
500 }));
501 assert!(results.iter().any(|(room_id, event)| {
502 room_id == room_id2 && event.event_id() == Some(result_event_id2)
503 }));
504 }
505 }
506
507 #[async_test]
508 async fn test_global_message_search_score_ordering() {
509 let server = MatrixMockServer::new().await;
510 let client = server.client_builder().build().await;
511
512 let event_cache = client.event_cache();
513 event_cache.subscribe().unwrap();
514
515 let room_id1 = room_id!("!r1:localhost");
516 let room_id2 = room_id!("!r2:localhost");
517
518 let f = EventFactory::new().sender(user_id!("@user_id:localhost"));
519
520 let r1_rank1 = event_id!("$r1_rank1:localhost"); let r2_rank2 = event_id!("$r2_rank2:localhost"); let r1_rank3 = event_id!("$r1_rank3:localhost"); let r2_rank4 = event_id!("$r2_rank4:localhost"); server
535 .mock_sync()
536 .ok_and_run(&client, |sync_builder| {
537 sync_builder
538 .add_joined_room(
539 JoinedRoomBuilder::new(room_id1)
540 .add_timeline_event(
541 f.text_msg("world world world world filler filler filler filler filler filler")
542 .room(room_id1)
543 .event_id(r1_rank1),
544 )
545 .add_timeline_event(
546 f.text_msg("world world filler filler filler filler filler filler filler filler")
547 .room(room_id1)
548 .event_id(r1_rank3),
549 ),
550 )
551 .add_joined_room(
552 JoinedRoomBuilder::new(room_id2)
553 .add_timeline_event(
554 f.text_msg("world world world filler filler filler filler filler filler filler")
555 .room(room_id2)
556 .event_id(r2_rank2),
557 )
558 .add_timeline_event(
559 f.text_msg("world filler filler filler filler filler filler filler filler filler")
560 .room(room_id2)
561 .event_id(r2_rank4),
562 ),
563 );
564 })
565 .await;
566
567 sleep(Duration::from_millis(200)).await;
568
569 let results: Vec<(OwnedRoomId, f32, OwnedEventId)> =
570 client.search_messages("world".to_owned()).build().try_concat().await.unwrap();
571 assert_eq!(results.len(), 4);
572
573 assert_eq!((&results[0].0, &results[0].2), (&room_id1.to_owned(), &r1_rank1.to_owned()));
575 assert_eq!((&results[1].0, &results[1].2), (&room_id2.to_owned(), &r2_rank2.to_owned()));
576 assert_eq!((&results[2].0, &results[2].2), (&room_id1.to_owned(), &r1_rank3.to_owned()));
577 assert_eq!((&results[3].0, &results[3].2), (&room_id2.to_owned(), &r2_rank4.to_owned()));
578 }
579
580 #[async_test]
581 async fn test_global_message_search_dm_or_groups() {
582 let server = MatrixMockServer::new().await;
583 let client = server.client_builder().build().await;
584
585 let event_cache = client.event_cache();
586 event_cache.subscribe().unwrap();
587
588 let room_id1 = room_id!("!r1:localhost");
590 let room_id2 = room_id!("!r2:localhost");
592
593 let f = EventFactory::new().sender(user_id!("@user_id:localhost"));
594
595 let result_event_id1 = event_id!("$result1:localhost");
596 let result_event_id2 = event_id!("$result2:localhost");
597
598 server
599 .mock_sync()
600 .ok_and_run(&client, |sync_builder| {
601 sync_builder
602 .add_joined_room(
603 JoinedRoomBuilder::new(room_id1)
604 .add_timeline_event(
605 f.text_msg("hello world").room(room_id1).event_id(result_event_id1),
606 )
607 .add_timeline_event(f.text_msg("hello back").room(room_id1)),
608 )
609 .add_joined_room(JoinedRoomBuilder::new(room_id2).add_timeline_event(
610 f.text_msg("it's a mad world").room(room_id2).event_id(result_event_id2),
611 ))
612 .add_global_account_data(
614 f.direct().add_user((*BOB).to_owned().into(), room_id1),
615 );
616 })
617 .await;
618
619 sleep(Duration::from_millis(200)).await;
621
622 {
624 let results: Vec<(OwnedRoomId, f32, OwnedEventId)> = client
625 .search_messages("world".to_owned())
626 .only_dm_rooms()
627 .await
628 .unwrap()
629 .build()
630 .try_concat()
631 .await
632 .unwrap();
633
634 assert_eq!(results.len(), 1);
635 assert_eq!(
636 (&results[0].0, &results[0].2),
637 (&room_id1.to_owned(), &result_event_id1.to_owned())
638 );
639 }
640
641 {
643 let results: Vec<_> = client
644 .search_messages("world".to_owned())
645 .no_dms()
646 .await
647 .unwrap()
648 .build_events()
649 .try_concat()
650 .await
651 .unwrap();
652
653 assert_eq!(results.len(), 1);
654 assert_eq!(results[0].0, room_id2);
655 assert_eq!(results[0].1.event_id().unwrap(), result_event_id2);
656 }
657 }
658}