256 lines
8.0 KiB
Rust
256 lines
8.0 KiB
Rust
use crate::domain::dto::reaction::ReactionGroupResponse;
|
|
use crate::domain::events::message::{MessageReactionAddedEvent, MessageReactionRemovedEvent};
|
|
use crate::http::error::HTTPError;
|
|
use crate::models::{emoji, message_reaction};
|
|
use crate::services::ServicesContext;
|
|
use event_bus::Scope;
|
|
use sea_orm::{ColumnTrait, EntityTrait, QueryFilter, Set};
|
|
use std::collections::HashMap;
|
|
use std::sync::Arc;
|
|
use std::time::Instant;
|
|
use uuid::Uuid;
|
|
|
|
#[derive(Debug, Clone)]
|
|
pub struct MessageReactionService {
|
|
context: Arc<ServicesContext>,
|
|
}
|
|
|
|
impl MessageReactionService {
|
|
pub fn new(context: Arc<ServicesContext>) -> Self {
|
|
Self { context }
|
|
}
|
|
|
|
fn normalize_skin_tone(skin_tone: Option<u8>) -> Result<i32, HTTPError> {
|
|
match skin_tone {
|
|
None => Ok(0),
|
|
Some(tone @ 1..=5) => Ok(tone as i32),
|
|
Some(_) => Err(HTTPError::BadRequest(
|
|
"skin_tone must be between 1 and 5".into(),
|
|
)),
|
|
}
|
|
}
|
|
|
|
async fn validate(
|
|
&self,
|
|
message_id: Uuid,
|
|
emoji_id: Uuid,
|
|
skin_tone: Option<u8>,
|
|
) -> Result<
|
|
(
|
|
crate::models::message::Model,
|
|
emoji::Model,
|
|
i32,
|
|
Option<Uuid>,
|
|
),
|
|
HTTPError,
|
|
> {
|
|
let message = self
|
|
.context
|
|
.repositories
|
|
.message
|
|
.get_by_id(message_id)
|
|
.await?
|
|
.ok_or(HTTPError::NotFound)?;
|
|
let channel = self
|
|
.context
|
|
.repositories
|
|
.channel
|
|
.get_by_id(message.channel_id)
|
|
.await?
|
|
.ok_or(HTTPError::NotFound)?;
|
|
let emoji = self
|
|
.context
|
|
.repositories
|
|
.emoji
|
|
.get_by_id(emoji_id)
|
|
.await?
|
|
.ok_or_else(|| HTTPError::BadRequest("Emoji not found".into()))?;
|
|
|
|
if emoji.server_id.is_some() && emoji.server_id != channel.server_id {
|
|
return Err(HTTPError::BadRequest(
|
|
"Emoji is not available in this message scope".into(),
|
|
));
|
|
}
|
|
|
|
let normalized_tone = Self::normalize_skin_tone(skin_tone)?;
|
|
if normalized_tone != 0 && (emoji.emoji_type != "unicode" || !emoji.supports_skin_tone) {
|
|
return Err(HTTPError::BadRequest(
|
|
"This emoji does not support skin tones".into(),
|
|
));
|
|
}
|
|
|
|
Ok((message, emoji, normalized_tone, channel.server_id))
|
|
}
|
|
|
|
pub async fn add(
|
|
&self,
|
|
message_id: Uuid,
|
|
user_id: Uuid,
|
|
emoji_id: Uuid,
|
|
skin_tone: Option<u8>,
|
|
) -> Result<(message_reaction::Model, bool), HTTPError> {
|
|
let started_at = Instant::now();
|
|
let (message, _, skin_tone, server_id) =
|
|
self.validate(message_id, emoji_id, skin_tone).await?;
|
|
|
|
if let Some(existing) = self
|
|
.context
|
|
.repositories
|
|
.message_reaction
|
|
.find(message_id, user_id, emoji_id, skin_tone)
|
|
.await?
|
|
{
|
|
return Ok((existing, false));
|
|
}
|
|
|
|
let reaction = self
|
|
.context
|
|
.repositories
|
|
.message_reaction
|
|
.create(message_reaction::ActiveModel {
|
|
id: Set(Uuid::now_v7()),
|
|
message_id: Set(message_id),
|
|
user_id: Set(user_id),
|
|
emoji_id: Set(emoji_id),
|
|
skin_tone: Set(skin_tone),
|
|
..Default::default()
|
|
})
|
|
.await?;
|
|
|
|
let mut scopes = vec![Scope::uuid("channel", message.channel_id)];
|
|
if let Some(server_id) = server_id {
|
|
scopes.push(Scope::uuid("server", server_id));
|
|
}
|
|
self.context.event_bus.emit_scoped(
|
|
"message_reaction_added",
|
|
scopes,
|
|
MessageReactionAddedEvent {
|
|
server_id,
|
|
channel_id: message.channel_id,
|
|
reaction: reaction.clone(),
|
|
},
|
|
);
|
|
|
|
Ok((reaction, true))
|
|
}
|
|
|
|
pub async fn remove(
|
|
&self,
|
|
message_id: Uuid,
|
|
user_id: Uuid,
|
|
emoji_id: Uuid,
|
|
skin_tone: Option<u8>,
|
|
) -> Result<message_reaction::Model, HTTPError> {
|
|
let (message, _, skin_tone, server_id) =
|
|
self.validate(message_id, emoji_id, skin_tone).await?;
|
|
let reaction = self
|
|
.context
|
|
.repositories
|
|
.message_reaction
|
|
.delete(message_id, user_id, emoji_id, skin_tone)
|
|
.await?
|
|
.ok_or(HTTPError::NotFound)?;
|
|
|
|
let mut scopes = vec![Scope::uuid("channel", message.channel_id)];
|
|
if let Some(server_id) = server_id {
|
|
scopes.push(Scope::uuid("server", server_id));
|
|
}
|
|
self.context.event_bus.emit_scoped(
|
|
"message_reaction_removed",
|
|
scopes,
|
|
MessageReactionRemovedEvent {
|
|
server_id,
|
|
channel_id: message.channel_id,
|
|
reaction: reaction.clone(),
|
|
},
|
|
);
|
|
|
|
Ok(reaction)
|
|
}
|
|
|
|
pub async fn grouped_for_messages(
|
|
&self,
|
|
message_ids: &[Uuid],
|
|
) -> Result<HashMap<Uuid, Vec<ReactionGroupResponse>>, HTTPError> {
|
|
let reactions = self
|
|
.context
|
|
.repositories
|
|
.message_reaction
|
|
.list_by_message_ids(message_ids)
|
|
.await?;
|
|
if reactions.is_empty() {
|
|
return Ok(HashMap::new());
|
|
}
|
|
let emoji_ids: Vec<_> = reactions.iter().map(|reaction| reaction.emoji_id).collect();
|
|
let emojis: HashMap<_, _> = emoji::Entity::find()
|
|
.filter(emoji::Column::Id.is_in(emoji_ids))
|
|
.all(&self.context.repositories.emoji.context.db)
|
|
.await?
|
|
.into_iter()
|
|
.map(|emoji| (emoji.id, emoji))
|
|
.collect();
|
|
|
|
let mut grouped: HashMap<Uuid, HashMap<(Uuid, i32), ReactionGroupResponse>> =
|
|
HashMap::new();
|
|
for reaction in reactions {
|
|
let Some(emoji) = emojis.get(&reaction.emoji_id) else {
|
|
continue;
|
|
};
|
|
let groups = grouped.entry(reaction.message_id).or_default();
|
|
let group = groups
|
|
.entry((reaction.emoji_id, reaction.skin_tone))
|
|
.or_insert_with(|| ReactionGroupResponse {
|
|
emoji_id: emoji.id,
|
|
name: emoji.name.clone(),
|
|
unicode_sequence: emoji.unicode_sequence.clone(),
|
|
asset_url: emoji
|
|
.file_path
|
|
.as_ref()
|
|
.map(|_| format!("/api/emojis/{}/asset", emoji.id)),
|
|
skin_tone: (reaction.skin_tone != 0).then_some(reaction.skin_tone as u8),
|
|
count: 0,
|
|
user_ids: Vec::new(),
|
|
});
|
|
group.count += 1;
|
|
group.user_ids.push(reaction.user_id);
|
|
}
|
|
|
|
Ok(grouped
|
|
.into_iter()
|
|
.map(|(message_id, groups)| {
|
|
let mut groups: Vec<_> = groups.into_values().collect();
|
|
for group in &mut groups {
|
|
group.user_ids.sort_unstable();
|
|
}
|
|
groups.sort_by(|left, right| {
|
|
left.name
|
|
.cmp(&right.name)
|
|
.then(left.skin_tone.cmp(&right.skin_tone))
|
|
});
|
|
(message_id, groups)
|
|
})
|
|
.collect())
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::MessageReactionService;
|
|
|
|
#[test]
|
|
fn skin_tone_validation_accepts_only_the_supported_range() {
|
|
assert_eq!(
|
|
MessageReactionService::normalize_skin_tone(None).unwrap(),
|
|
0
|
|
);
|
|
for tone in 1..=5 {
|
|
assert_eq!(
|
|
MessageReactionService::normalize_skin_tone(Some(tone)).unwrap(),
|
|
tone as i32
|
|
);
|
|
}
|
|
assert!(MessageReactionService::normalize_skin_tone(Some(0)).is_err());
|
|
assert!(MessageReactionService::normalize_skin_tone(Some(6)).is_err());
|
|
}
|
|
}
|