Skip to main content

matrix_sdk/
message_search.rs

1// Copyright 2026 The Matrix.org Foundation C.I.C.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Messages search facilities and high-level helpers to perform searches across
16//! one or multiple rooms.
17//!
18//! These helpers expose the results as [`Stream`]s of pages, lazily fetching
19//! the next page from the underlying index as the stream is polled. Use the
20//! [`StreamExt`] and [`TryStreamExt`] combinators (`next`, `try_concat`,
21//! `take`, …) to consume them.
22//!
23//! [`StreamExt`]: futures_util::StreamExt
24//! [`TryStreamExt`]: futures_util::TryStreamExt
25//!
26//! # Examples
27//!
28//! ## Searching within a single room
29//!
30//! Use [`Room::search_messages`] to get a stream of pages of `(score,
31//! event_id)` pairs, or [`Room::search_messages_events`] to load the full
32//! [`TimelineEvent`]s.
33//!
34//! ```no_run
35//! # use matrix_sdk::{Room, message_search::SearchResult};
36//! # use futures_util::StreamExt as _;
37//! # async fn example(room: Room) -> anyhow::Result<()> {
38//! let mut stream = Box::pin(room.search_messages("hello world".to_owned()));
39//!
40//! while let Some(page) = stream.next().await {
41//!     let SearchResult { total_count, events } = page?;
42//!
43//!     for (score, event_id) in events {
44//!         println!("Found event {event_id} (score: {score}, total_count: {total_count})");
45//!     }
46//! }
47//! # Ok(())
48//! # }
49//! ```
50//!
51//! ## Searching across all joined rooms
52//!
53//! Use [`Client::search_messages`] to create a [`GlobalSearchBuilder`].
54//! Optionally restrict the working set to DM rooms (or non-DM rooms) before
55//! calling [`GlobalSearchBuilder::build`] to get a stream of pages of results,
56//! sorted by relevance score across all rooms. Use
57//! [`GlobalSearchBuilder::build_events`] to load full [`TimelineEvent`]s
58//! instead of plain event IDs.
59//!
60//! ```no_run
61//! # use matrix_sdk::{Client};
62//! # use futures_util::StreamExt as _;
63//! # async fn example(client: Client) -> anyhow::Result<()> {
64//! // Search only in DM rooms.
65//! let mut stream = Box::pin(
66//!     client
67//!         .search_messages("hello world".to_owned())
68//!         .only_dm_rooms()
69//!         .await?
70//!         .build_events(),
71//! );
72//!
73//! while let Some(page) = stream.next().await {
74//!     for (room_id, event) in page? {
75//!         println!(
76//!             "Found event in room {room_id} with timestamp: {:?}",
77//!             event.timestamp
78//!         );
79//!     }
80//! }
81//! # Ok(())
82//! # }
83//! ```
84
85use 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
98/// Number of results pulled from the index in one go while paginating through a
99/// search stream.
100const SEARCH_RESULTS_PAGE_SIZE: usize = 100;
101
102/// A boxed, score-descending stream of `(score, event_id)` results for a single
103/// room.
104type RoomResultStream = Pin<Box<dyn Stream<Item = Result<(f32, OwnedEventId), IndexError>> + Send>>;
105
106/// A cursor over one room's score-descending search results, used while merging
107/// results across rooms.
108struct RoomStreamCursor {
109    /// The room these results come from.
110    room_id: OwnedRoomId,
111
112    /// The room's score-descending result stream.
113    stream: RoomResultStream,
114
115    /// The next result this room would contribute to the merge: a one-item
116    /// lookahead buffered from `stream`, so we can compare every room's best
117    /// remaining result without consuming it. `None` once the stream is
118    /// exhausted.
119    next_result: Option<(f32, OwnedEventId)>,
120}
121
122impl Room {
123    /// Search this room's [`RoomIndex`] for query and return at most
124    /// max_number_of_results results.
125    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/// An error that can occur while searching messages, using the high-level
137/// search helpers provided by this module.
138#[derive(thiserror::Error, Debug)]
139pub enum SearchError {
140    /// An error occurred while searching through the index for matching events.
141    #[error(transparent)]
142    IndexError(#[from] IndexError),
143    /// An error occurred while loading the event content for a search result.
144    #[error(transparent)]
145    EventLoadError(#[from] crate::Error),
146}
147
148impl Room {
149    /// Search for messages in this room matching the given query, returning a
150    /// stream of pages of `(score, event_id)` results sorted by descending
151    /// relevance score.
152    pub fn search_messages(
153        &self,
154        query: String,
155    ) -> impl Stream<Item = Result<SearchResult, IndexError>> + use<> {
156        let room = self.clone();
157
158        // TODO: use the client/server API search endpoint for public rooms, as
159        // those may require lots of time for indexing all events.
160        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    /// Same as [`Room::search_messages`], but yields pages of full
174    /// [`TimelineEvent`]s instead of event IDs, by loading them from the store
175    /// or from the network.
176    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/// A builder for a global search [`Stream`] that allows configuring the initial
200/// working set of rooms to search in.
201#[derive(Debug)]
202pub struct GlobalSearchBuilder {
203    client: Client,
204
205    /// The search query, directly forwarded to the search API.
206    query: String,
207
208    /// The working set of rooms to search in.
209    room_set: Vec<Room>,
210}
211
212impl GlobalSearchBuilder {
213    /// Create a new global search on all the joined rooms.
214    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    /// Keep only the DM rooms from the initial working set.
220    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    /// Keep only non-DM rooms (groups) from the initial working set.
232    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    /// Build a stream over the search results across all the rooms in the
244    /// working set, yielding pages of `(room_id, score, event_id)` tuples
245    /// sorted by descending relevance score.
246    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            // One score-descending result stream per room, each primed with its
254            // next result so we can merge across rooms by score.
255            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                // Pick the room whose next result has the highest relevance score.
274                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                    // Every room is exhausted.
284                    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                // Refill this room's lookahead for the next iteration.
293                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    /// Same as [`Self::build`], but yields pages of full [`TimelineEvent`]s
311    /// instead of event IDs, by loading them from the store or from the
312    /// network.
313    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    /// Flatten a stream of result pages into a stream of individual results, so
340    /// the cross-room merge can compare results one at a time.
341    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    /// Search across all rooms for events with the given query, returning a
358    /// builder for a stream over the results.
359    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        // Let the search indexer process the new event.
399        sleep(Duration::from_millis(200)).await;
400
401        // Searching for a missing keyword should succeed and yield nothing.
402        {
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        // Search for an existing keyword, by event id.
409        {
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        // Search for an existing keyword, by events.
417        {
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        // Let the search indexer process the new event.
459        sleep(Duration::from_millis(200)).await;
460
461        // Searching for a missing keyword should succeed and yield nothing.
462        {
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        // Search for an existing keyword, by event id.
473        {
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            // Search results order is not guaranteed, so we check that both
478            // expected results are present.
479            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        // Search for an existing keyword, by event.
488        {
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            // Search results order is not guaranteed, so we check that both
497            // expected results are present.
498            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        // Both rooms get two documents of identical length (padded with filler
521        // so document-length normalization and the per-corpus IDF of "world"
522        // match across rooms). The score then depends only on how many times
523        // "world" appears.
524        //
525        // Term frequencies are 4, 3, 2, 1, split so the rooms alternate by
526        // rank: room1 holds the 4x and 2x events, room2 the 3x and 1x events. A
527        // correct cross-room sort therefore interleaves the rooms: r1, r2, r1,
528        // r2.
529        let r1_rank1 = event_id!("$r1_rank1:localhost"); // room1, "world" x4
530        let r2_rank2 = event_id!("$r2_rank2:localhost"); // room2, "world" x3
531        let r1_rank3 = event_id!("$r1_rank3:localhost"); // room1, "world" x2
532        let r2_rank4 = event_id!("$r2_rank4:localhost"); // room2, "world" x1
533
534        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        // Results are interleaved across the two rooms strictly by score.
574        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        // This time, room_id1 is a DM room,
589        let room_id1 = room_id!("!r1:localhost");
590        // While room_id2 isn't.
591        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                    // Note: adding a DM room for room_id1 here.
613                    .add_global_account_data(
614                        f.direct().add_user((*BOB).to_owned().into(), room_id1),
615                    );
616            })
617            .await;
618
619        // Let the search indexer process the new event.
620        sleep(Duration::from_millis(200)).await;
621
622        // Search for an existing keyword, by event id, only in DMs.
623        {
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        // Search for an existing keyword, by event, only in groups.
642        {
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}