use std::{collections::HashSet, pin::Pin};
use async_stream::try_stream;
use futures_util::{Stream, StreamExt as _};
use matrix_sdk_base::{RoomStateFilter, deserialized_responses::TimelineEvent};
use matrix_sdk_search::error::IndexError;
#[cfg(doc)]
use matrix_sdk_search::index::RoomIndex;
use ruma::{OwnedEventId, OwnedRoomId};
use crate::{Client, Room};
const SEARCH_RESULTS_PAGE_SIZE: usize = 100;
type RoomResultStream = Pin<Box<dyn Stream<Item = Result<(f32, OwnedEventId), IndexError>> + Send>>;
struct RoomStreamCursor {
room_id: OwnedRoomId,
stream: RoomResultStream,
next_result: Option<(f32, OwnedEventId)>,
}
impl Room {
pub async fn search(
&self,
query: &str,
max_number_of_results: usize,
pagination_offset: Option<usize>,
) -> Result<Vec<(f32, OwnedEventId)>, IndexError> {
let mut search_index_guard = self.client.search_index().lock().await;
search_index_guard.search(query, max_number_of_results, pagination_offset, self.room_id())
}
}
#[derive(thiserror::Error, Debug)]
pub enum SearchError {
#[error(transparent)]
IndexError(#[from] IndexError),
#[error(transparent)]
EventLoadError(#[from] crate::Error),
}
impl Room {
pub fn search_messages(
&self,
query: String,
) -> impl Stream<Item = Result<Vec<(f32, OwnedEventId)>, IndexError>> + use<> {
let room = self.clone();
try_stream! {
let mut offset = 0;
loop {
let page = room.search(&query, SEARCH_RESULTS_PAGE_SIZE, Some(offset)).await?;
if page.is_empty() {
break;
}
offset += page.len();
yield page;
}
}
}
pub fn search_messages_events(
&self,
query: String,
) -> impl Stream<Item = Result<Vec<TimelineEvent>, SearchError>> + use<> {
let room = self.clone();
try_stream! {
let mut pages = Box::pin(room.search_messages(query));
while let Some(page) = pages.next().await {
let page = page?;
let mut events = Vec::with_capacity(page.len());
for (_score, event_id) in page {
events.push(room.load_or_fetch_event(&event_id, None).await?);
}
yield events;
}
}
}
}
#[derive(Debug)]
pub struct GlobalSearchBuilder {
client: Client,
query: String,
room_set: Vec<Room>,
}
impl GlobalSearchBuilder {
fn new(client: Client, query: String) -> Self {
let room_set = client.rooms_filtered(RoomStateFilter::JOINED);
Self { client, query, room_set }
}
pub async fn only_dm_rooms(mut self) -> Result<Self, crate::Error> {
let mut to_remove = HashSet::new();
for room in &self.room_set {
if !room.compute_is_dm().await? {
to_remove.insert(room.room_id().to_owned());
}
}
self.room_set.retain(|room| !to_remove.contains(room.room_id()));
Ok(self)
}
pub async fn no_dms(mut self) -> Result<Self, crate::Error> {
let mut to_remove = HashSet::new();
for room in &self.room_set {
if room.compute_is_dm().await? {
to_remove.insert(room.room_id().to_owned());
}
}
self.room_set.retain(|room| !to_remove.contains(room.room_id()));
Ok(self)
}
pub fn build(
self,
) -> impl Stream<Item = Result<Vec<(OwnedRoomId, f32, OwnedEventId)>, IndexError>> {
let query = self.query;
let rooms = self.room_set;
try_stream! {
let mut cursors: Vec<RoomStreamCursor> = Vec::with_capacity(rooms.len());
for room in rooms {
let room_id = room.room_id().to_owned();
let stream = Box::pin(Self::flatten_pages(room.search_messages(query.clone())));
cursors.push(RoomStreamCursor { room_id, stream, next_result: None });
}
for cursor in &mut cursors {
cursor.next_result = match cursor.stream.next().await {
Some(result) => Some(result?),
None => None,
};
}
let mut page = Vec::with_capacity(SEARCH_RESULTS_PAGE_SIZE);
loop {
let best = cursors
.iter()
.enumerate()
.filter_map(|(index, cursor)| {
cursor.next_result.as_ref().map(|(score, _)| (index, *score))
})
.max_by(|(_, a), (_, b)| a.total_cmp(b));
let Some((index, _)) = best else {
break;
};
let cursor = &mut cursors[index];
let (score, event_id) =
cursor.next_result.take().expect("the chosen room must have a next result");
let room_id = cursor.room_id.clone();
cursor.next_result = match cursor.stream.next().await {
Some(result) => Some(result?),
None => None,
};
page.push((room_id, score, event_id));
if page.len() == SEARCH_RESULTS_PAGE_SIZE {
yield std::mem::take(&mut page);
}
}
if !page.is_empty() {
yield page;
}
}
}
pub fn build_events(
self,
) -> impl Stream<Item = Result<Vec<(OwnedRoomId, TimelineEvent)>, SearchError>> {
let client = self.client.clone();
let pages = self.build();
try_stream! {
let mut pages = Box::pin(pages);
while let Some(page) = pages.next().await {
let page = page?;
let mut events = Vec::with_capacity(page.len());
for (room_id, _score, event_id) in page {
let Some(room) = client.get_room(&room_id) else {
continue;
};
events.push((room_id, room.load_or_fetch_event(&event_id, None).await?));
}
yield events;
}
}
}
fn flatten_pages(
pages: impl Stream<Item = Result<Vec<(f32, OwnedEventId)>, IndexError>>,
) -> impl Stream<Item = Result<(f32, OwnedEventId), IndexError>> {
try_stream! {
let mut pages = Box::pin(pages);
while let Some(page) = pages.next().await {
for result in page? {
yield result;
}
}
}
}
}
impl Client {
pub fn search_messages(&self, query: String) -> GlobalSearchBuilder {
GlobalSearchBuilder::new(self.clone(), query)
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use futures_util::TryStreamExt as _;
use matrix_sdk_test::{BOB, JoinedRoomBuilder, async_test, event_factory::EventFactory};
use ruma::{OwnedEventId, OwnedRoomId, event_id, room_id, user_id};
use crate::{sleep::sleep, test_utils::mocks::MatrixMockServer};
#[async_test]
async fn test_room_message_search() {
let server = MatrixMockServer::new().await;
let client = server.client_builder().build().await;
let event_cache = client.event_cache();
event_cache.subscribe().unwrap();
let room_id = room_id!("!room_id:localhost");
let room = server.sync_joined_room(&client, room_id).await;
let f = EventFactory::new().room(room_id).sender(user_id!("@user_id:localhost"));
let event_id = event_id!("$event_id:localhost");
server
.sync_room(
&client,
JoinedRoomBuilder::new(room_id)
.add_timeline_event(f.text_msg("hello world").event_id(event_id)),
)
.await;
sleep(Duration::from_millis(200)).await;
{
let results: Vec<(f32, OwnedEventId)> =
room.search_messages("search query".to_owned()).try_concat().await.unwrap();
assert!(results.is_empty());
}
{
let results: Vec<(f32, OwnedEventId)> =
room.search_messages("world".to_owned()).try_concat().await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].1, event_id);
}
{
let events: Vec<_> =
room.search_messages_events("world".to_owned()).try_concat().await.unwrap();
assert_eq!(events.len(), 1);
assert_eq!(events[0].event_id().unwrap(), event_id);
}
}
#[async_test]
async fn test_global_message_search() {
let server = MatrixMockServer::new().await;
let client = server.client_builder().build().await;
let event_cache = client.event_cache();
event_cache.subscribe().unwrap();
let room_id1 = room_id!("!r1:localhost");
let room_id2 = room_id!("!r2:localhost");
let f = EventFactory::new().sender(user_id!("@user_id:localhost"));
let result_event_id1 = event_id!("$result1:localhost");
let result_event_id2 = event_id!("$result2:localhost");
server
.mock_sync()
.ok_and_run(&client, |sync_builder| {
sync_builder
.add_joined_room(
JoinedRoomBuilder::new(room_id1)
.add_timeline_event(
f.text_msg("hello world").room(room_id1).event_id(result_event_id1),
)
.add_timeline_event(f.text_msg("hello back").room(room_id1)),
)
.add_joined_room(JoinedRoomBuilder::new(room_id2).add_timeline_event(
f.text_msg("it's a mad world").room(room_id2).event_id(result_event_id2),
));
})
.await;
sleep(Duration::from_millis(200)).await;
{
let results: Vec<(OwnedRoomId, f32, OwnedEventId)> = client
.search_messages("search query".to_owned())
.build()
.try_concat()
.await
.unwrap();
assert!(results.is_empty());
}
{
let results: Vec<(OwnedRoomId, f32, OwnedEventId)> =
client.search_messages("world".to_owned()).build().try_concat().await.unwrap();
assert_eq!(results.len(), 2);
assert!(results.iter().any(|(room_id, _, event_id)| {
room_id == room_id1 && event_id == result_event_id1
}));
assert!(results.iter().any(|(room_id, _, event_id)| {
room_id == room_id2 && event_id == result_event_id2
}));
}
{
let results: Vec<_> = client
.search_messages("world".to_owned())
.build_events()
.try_concat()
.await
.unwrap();
assert_eq!(results.len(), 2);
assert!(results.iter().any(|(room_id, event)| {
room_id == room_id1 && event.event_id() == Some(result_event_id1)
}));
assert!(results.iter().any(|(room_id, event)| {
room_id == room_id2 && event.event_id() == Some(result_event_id2)
}));
}
}
#[async_test]
async fn test_global_message_search_score_ordering() {
let server = MatrixMockServer::new().await;
let client = server.client_builder().build().await;
let event_cache = client.event_cache();
event_cache.subscribe().unwrap();
let room_id1 = room_id!("!r1:localhost");
let room_id2 = room_id!("!r2:localhost");
let f = EventFactory::new().sender(user_id!("@user_id:localhost"));
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
.mock_sync()
.ok_and_run(&client, |sync_builder| {
sync_builder
.add_joined_room(
JoinedRoomBuilder::new(room_id1)
.add_timeline_event(
f.text_msg("world world world world filler filler filler filler filler filler")
.room(room_id1)
.event_id(r1_rank1),
)
.add_timeline_event(
f.text_msg("world world filler filler filler filler filler filler filler filler")
.room(room_id1)
.event_id(r1_rank3),
),
)
.add_joined_room(
JoinedRoomBuilder::new(room_id2)
.add_timeline_event(
f.text_msg("world world world filler filler filler filler filler filler filler")
.room(room_id2)
.event_id(r2_rank2),
)
.add_timeline_event(
f.text_msg("world filler filler filler filler filler filler filler filler filler")
.room(room_id2)
.event_id(r2_rank4),
),
);
})
.await;
sleep(Duration::from_millis(200)).await;
let results: Vec<(OwnedRoomId, f32, OwnedEventId)> =
client.search_messages("world".to_owned()).build().try_concat().await.unwrap();
assert_eq!(results.len(), 4);
assert_eq!((&results[0].0, &results[0].2), (&room_id1.to_owned(), &r1_rank1.to_owned()));
assert_eq!((&results[1].0, &results[1].2), (&room_id2.to_owned(), &r2_rank2.to_owned()));
assert_eq!((&results[2].0, &results[2].2), (&room_id1.to_owned(), &r1_rank3.to_owned()));
assert_eq!((&results[3].0, &results[3].2), (&room_id2.to_owned(), &r2_rank4.to_owned()));
}
#[async_test]
async fn test_global_message_search_dm_or_groups() {
let server = MatrixMockServer::new().await;
let client = server.client_builder().build().await;
let event_cache = client.event_cache();
event_cache.subscribe().unwrap();
let room_id1 = room_id!("!r1:localhost");
let room_id2 = room_id!("!r2:localhost");
let f = EventFactory::new().sender(user_id!("@user_id:localhost"));
let result_event_id1 = event_id!("$result1:localhost");
let result_event_id2 = event_id!("$result2:localhost");
server
.mock_sync()
.ok_and_run(&client, |sync_builder| {
sync_builder
.add_joined_room(
JoinedRoomBuilder::new(room_id1)
.add_timeline_event(
f.text_msg("hello world").room(room_id1).event_id(result_event_id1),
)
.add_timeline_event(f.text_msg("hello back").room(room_id1)),
)
.add_joined_room(JoinedRoomBuilder::new(room_id2).add_timeline_event(
f.text_msg("it's a mad world").room(room_id2).event_id(result_event_id2),
))
.add_global_account_data(
f.direct().add_user((*BOB).to_owned().into(), room_id1),
);
})
.await;
sleep(Duration::from_millis(200)).await;
{
let results: Vec<(OwnedRoomId, f32, OwnedEventId)> = client
.search_messages("world".to_owned())
.only_dm_rooms()
.await
.unwrap()
.build()
.try_concat()
.await
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(
(&results[0].0, &results[0].2),
(&room_id1.to_owned(), &result_event_id1.to_owned())
);
}
{
let results: Vec<_> = client
.search_messages("world".to_owned())
.no_dms()
.await
.unwrap()
.build_events()
.try_concat()
.await
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].0, room_id2);
assert_eq!(results[0].1.event_id().unwrap(), result_event_id2);
}
}
}