use crate::error::Result;
use crate::ids::{ConversationId, MessageId, ParticipantId};
use crate::message::{ContentType, Message};
use chrono::{DateTime, Utc};
use dashmap::DashMap;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
#[derive(Clone, Debug, Default)]
pub struct MessageFilter {
pub from_participant: Option<ParticipantId>,
pub content_types: Option<Vec<ContentType>>,
pub since: Option<DateTime<Utc>>,
pub until: Option<DateTime<Utc>>,
pub page_size: Option<usize>,
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub struct PageCursor {
pub offset: usize,
}
pub struct MessagePage {
pub messages: Vec<Message>,
pub next: Option<PageCursor>,
}
#[async_trait::async_trait]
pub trait MessageStore: Send + Sync {
async fn put(&self, message: Message) -> Result<()>;
async fn get(&self, id: &MessageId) -> Result<Option<Message>>;
async fn list(
&self,
conversation_id: &ConversationId,
filter: MessageFilter,
cursor: Option<PageCursor>,
) -> Result<MessagePage>;
async fn mark_read(
&self,
id: &MessageId,
by: &ParticipantId,
) -> Result<()>;
}
#[derive(Clone, Debug, Default)]
pub struct MemoryMessageStore {
log: Arc<DashMap<ConversationId, Vec<Message>>>,
read_by: Arc<DashMap<MessageId, Vec<ParticipantId>>>,
}
impl MemoryMessageStore {
pub fn new() -> Self {
Self::default()
}
pub fn read_receipts(&self, id: &MessageId) -> Vec<ParticipantId> {
self.read_by
.get(id)
.map(|e| e.value().clone())
.unwrap_or_default()
}
}
#[async_trait::async_trait]
impl MessageStore for MemoryMessageStore {
async fn put(&self, message: Message) -> Result<()> {
let cid = message.conversation_id.clone();
self.log.entry(cid).or_default().push(message);
Ok(())
}
async fn get(&self, id: &MessageId) -> Result<Option<Message>> {
for entry in self.log.iter() {
if let Some(m) = entry.value().iter().find(|m| &m.id == id) {
return Ok(Some(m.clone()));
}
}
Ok(None)
}
async fn list(
&self,
conversation_id: &ConversationId,
filter: MessageFilter,
cursor: Option<PageCursor>,
) -> Result<MessagePage> {
let log = self
.log
.get(conversation_id)
.map(|e| e.value().clone())
.unwrap_or_default();
let start = cursor.map(|c| c.offset).unwrap_or(0);
let limit = filter.page_size.unwrap_or(50);
let filtered: Vec<Message> = log
.into_iter()
.skip(start)
.filter(|m| {
filter
.from_participant
.as_ref()
.map_or(true, |p| &m.from_participant == p)
&& filter
.content_types
.as_ref()
.map_or(true, |cts| cts.contains(&m.content_type))
&& filter.since.map_or(true, |t| m.timestamp >= t)
&& filter.until.map_or(true, |t| m.timestamp <= t)
})
.take(limit)
.collect();
let next = if filtered.len() == limit {
Some(PageCursor {
offset: start + limit,
})
} else {
None
};
Ok(MessagePage {
messages: filtered,
next,
})
}
async fn mark_read(
&self,
id: &MessageId,
by: &ParticipantId,
) -> Result<()> {
let mut e = self.read_by.entry(id.clone()).or_default();
if !e.contains(by) {
e.push(by.clone());
}
Ok(())
}
}