use std::pin::Pin;
use eyeball::{ObservableWriteGuard, SharedObservable, Subscriber};
use eyeball_im::{ObservableVector, Vector, VectorSubscriberBatchedStream};
use futures_util::{Stream, StreamExt as _};
use matrix_sdk::{
Client, deserialized_responses::TimelineEvent, message_search::SearchError, room::Room,
};
use ruma::{MilliSecondsSinceUnixEpoch, OwnedEventId, OwnedRoomId, OwnedUserId};
use tokio::sync::Mutex as AsyncMutex;
use crate::timeline::{Profile, TimelineDetails, TimelineItemContent};
type ResultsStream =
Pin<Box<dyn Stream<Item = Result<Vec<(OwnedRoomId, TimelineEvent)>, SearchError>> + Send>>;
#[cfg_attr(feature = "uniffi", derive(uniffi::Enum), uniffi(name = "SearchServicePaginationState"))]
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum PaginationState {
Idle { end_reached: bool },
Loading,
}
#[derive(Debug, Clone)]
pub enum ResultType {
Message(MessageResult),
}
#[derive(Debug, Clone)]
pub struct MessageResult {
pub room_id: OwnedRoomId,
pub event_id: OwnedEventId,
pub sender: OwnedUserId,
pub sender_profile: TimelineDetails<Profile>,
pub content: TimelineItemContent,
pub timestamp: MilliSecondsSinceUnixEpoch,
}
impl MessageResult {
async fn from_event(room: &Room, event: TimelineEvent) -> Option<Self> {
let sender = event.sender()?;
let event_id = event.event_id()?.to_owned();
let timestamp = event.timestamp().unwrap_or_else(MilliSecondsSinceUnixEpoch::now);
let content = TimelineItemContent::from_event(room, event).await?;
let sender_profile =
TimelineDetails::from_initial_value(Profile::load(room, &sender).await);
Some(Self {
room_id: room.room_id().to_owned(),
event_id,
sender,
sender_profile,
content,
timestamp,
})
}
}
pub struct SearchService {
client: Client,
stream: AsyncMutex<Option<ResultsStream>>,
pagination_state: SharedObservable<PaginationState>,
results: AsyncMutex<ObservableVector<ResultType>>,
}
impl SearchService {
pub fn new(client: Client) -> Self {
Self {
client,
stream: AsyncMutex::new(None),
pagination_state: SharedObservable::new(PaginationState::Idle { end_reached: false }),
results: AsyncMutex::new(ObservableVector::new()),
}
}
pub async fn set_query(&self, query: String) -> Result<(), SearchError> {
let stream = self.client.search_messages(query).build_events();
*self.stream.lock().await = Some(Box::pin(stream));
self.results.lock().await.clear();
self.pagination_state.set(PaginationState::Idle { end_reached: false });
self.paginate().await
}
pub fn pagination_state(&self) -> PaginationState {
self.pagination_state.get()
}
pub fn subscribe_to_pagination_state_updates(&self) -> Subscriber<PaginationState> {
self.pagination_state.subscribe()
}
pub async fn results(&self) -> Vec<ResultType> {
self.results.lock().await.iter().cloned().collect()
}
pub async fn subscribe_to_results(
&self,
) -> (Vector<ResultType>, VectorSubscriberBatchedStream<ResultType>) {
self.results.lock().await.subscribe().into_values_and_batched_stream()
}
pub async fn paginate(&self) -> Result<(), SearchError> {
{
let mut pagination_state = self.pagination_state.write();
match *pagination_state {
PaginationState::Idle { end_reached } if end_reached => return Ok(()),
PaginationState::Loading => return Ok(()),
_ => {}
}
ObservableWriteGuard::set(&mut pagination_state, PaginationState::Loading);
}
let mut stream = self.stream.lock().await;
let Some(stream) = stream.as_mut() else {
self.pagination_state.set(PaginationState::Idle { end_reached: true });
return Ok(());
};
match stream.next().await {
None => {
self.pagination_state.set(PaginationState::Idle { end_reached: true });
}
Some(Err(err)) => {
self.pagination_state.set(PaginationState::Idle { end_reached: false });
return Err(err);
}
Some(Ok(page)) => {
let mut resolved = Vector::new();
for (room_id, event) in page {
let Some(room) = self.client.get_room(&room_id) else {
continue;
};
let Some(result) = MessageResult::from_event(&room, event).await else {
continue;
};
resolved.push_back(ResultType::Message(result));
}
if !resolved.is_empty() {
self.results.lock().await.append(resolved);
}
self.pagination_state.set(PaginationState::Idle { end_reached: false });
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use assert_matches2::assert_let;
use eyeball_im::VectorDiff;
use futures_util::pin_mut;
use matrix_sdk::test_utils::mocks::MatrixMockServer;
use matrix_sdk_test::{JoinedRoomBuilder, async_test, event_factory::EventFactory};
use ruma::{event_id, room_id, user_id};
use stream_assert::{assert_next_matches, assert_pending};
use tokio::time::sleep;
use super::{PaginationState, ResultType, SearchService};
#[async_test]
async fn test_search_pagination() {
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:localhost");
let event_id = event_id!("$event:localhost");
let f = EventFactory::new().sender(user_id!("@user:localhost"));
server
.mock_sync()
.ok_and_run(&client, |builder| {
builder.add_joined_room(JoinedRoomBuilder::new(room_id).add_timeline_event(
f.text_msg("hello world").room(room_id).event_id(event_id),
));
})
.await;
sleep(Duration::from_millis(300)).await;
let search = SearchService::new(client);
assert_eq!(search.pagination_state(), PaginationState::Idle { end_reached: false });
assert!(search.results().await.is_empty());
search.set_query("world".to_owned()).await.unwrap();
assert_eq!(search.pagination_state(), PaginationState::Idle { end_reached: false });
let results = search.results().await;
assert_eq!(results.len(), 1);
assert_let!(ResultType::Message(message) = &results[0]);
assert_eq!(message.event_id, event_id);
let (initial, results_stream) = search.subscribe_to_results().await;
assert_eq!(initial.len(), 1);
pin_mut!(results_stream);
assert_pending!(results_stream);
search.paginate().await.unwrap();
assert_pending!(results_stream);
assert_eq!(search.pagination_state(), PaginationState::Idle { end_reached: true });
assert_eq!(search.results().await.len(), 1);
}
#[async_test]
async fn test_search_resets_on_query_change() {
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:localhost");
let apple_event = event_id!("$apple:localhost");
let banana_event = event_id!("$banana:localhost");
let f = EventFactory::new().sender(user_id!("@user:localhost"));
server
.mock_sync()
.ok_and_run(&client, |builder| {
builder.add_joined_room(
JoinedRoomBuilder::new(room_id)
.add_timeline_event(
f.text_msg("apple pie").room(room_id).event_id(apple_event),
)
.add_timeline_event(
f.text_msg("banana split").room(room_id).event_id(banana_event),
),
);
})
.await;
sleep(Duration::from_millis(300)).await;
let search = SearchService::new(client);
search.set_query("apple".to_owned()).await.unwrap();
let (initial, results_stream) = search.subscribe_to_results().await;
assert_eq!(initial.len(), 1);
assert_let!(ResultType::Message(message) = &initial[0]);
assert_eq!(message.event_id, apple_event);
pin_mut!(results_stream);
assert_pending!(results_stream);
search.set_query("banana".to_owned()).await.unwrap();
assert_next_matches!(results_stream, diffs => {
assert_let!([VectorDiff::Clear, VectorDiff::Append { values }] = diffs.as_slice());
assert_eq!(values.len(), 1);
assert_let!(ResultType::Message(message) = &values[0]);
assert_eq!(message.event_id, banana_event);
});
let results = search.results().await;
assert_eq!(results.len(), 1);
assert_let!(ResultType::Message(message) = &results[0]);
assert_eq!(message.event_id, banana_event);
}
}