use crate::{Channel, Pusher, PusherError, Result};
use base64::{Engine as _, engine::general_purpose::STANDARD as BASE64};
use serde::{Deserialize, Serialize};
use sonic_rs::{Value, json};
use std::collections::HashMap;
use std::fmt;
#[cfg(all(feature = "encryption", feature = "sodiumoxide"))]
use std::sync::Once;
#[cfg(all(feature = "encryption", feature = "sodiumoxide"))]
static SODIUM_INIT: Once = Once::new();
#[cfg(all(feature = "encryption", feature = "sodiumoxide"))]
fn init_sodium() -> Result<()> {
SODIUM_INIT.call_once(|| {
sodiumoxide::init().expect("Failed to initialize sodiumoxide");
});
Ok(())
}
#[derive(Debug, Clone, PartialEq)]
pub enum EventData {
String(String),
Json(Value),
}
impl EventData {
pub fn from_string(s: impl Into<String>) -> Self {
EventData::String(s.into())
}
pub fn from_json(value: Value) -> Self {
EventData::Json(value)
}
pub fn to_string(&self) -> String {
match self {
EventData::String(s) => s.clone(),
EventData::Json(v) => sonic_rs::to_string(v).unwrap_or_default(),
}
}
pub fn as_json(&self) -> Result<Value> {
match self {
EventData::String(s) => sonic_rs::from_str(s).map_err(|e| PusherError::Json(e)),
EventData::Json(v) => Ok(v.clone()),
}
}
}
impl fmt::Display for EventData {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.to_string())
}
}
impl From<String> for EventData {
fn from(s: String) -> Self {
EventData::String(s)
}
}
impl From<&str> for EventData {
fn from(s: &str) -> Self {
EventData::String(s.to_string())
}
}
impl From<Value> for EventData {
fn from(v: Value) -> Self {
EventData::Json(v)
}
}
#[derive(Debug, Serialize)]
pub struct Event {
pub name: String,
pub data: String,
pub channels: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub socket_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub info: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tags: Option<HashMap<String, String>>,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct BatchEvent {
pub name: String,
pub channel: String,
pub data: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub socket_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub info: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tags: Option<HashMap<String, String>>,
}
impl BatchEvent {
pub fn new(
name: impl Into<String>,
channel: impl Into<String>,
data: impl Into<EventData>,
) -> Self {
Self {
name: name.into(),
channel: channel.into(),
data: data.into().to_string(),
socket_id: None,
info: None,
tags: None,
}
}
pub fn with_socket_id(mut self, socket_id: impl Into<String>) -> Self {
self.socket_id = Some(socket_id.into());
self
}
pub fn with_info(mut self, info: impl Into<String>) -> Self {
self.info = Some(info.into());
self
}
pub fn with_tags(mut self, tags: HashMap<String, String>) -> Self {
self.tags = Some(tags);
self
}
}
#[derive(Debug, Clone, Default)]
pub struct TriggerParams {
pub socket_id: Option<String>,
pub info: Option<String>,
pub tags: Option<HashMap<String, String>>,
}
impl TriggerParams {
pub fn builder() -> TriggerParamsBuilder {
TriggerParamsBuilder::default()
}
}
#[derive(Debug, Default)]
pub struct TriggerParamsBuilder {
socket_id: Option<String>,
info: Option<String>,
tags: Option<HashMap<String, String>>,
}
impl TriggerParamsBuilder {
pub fn socket_id(mut self, socket_id: impl Into<String>) -> Self {
self.socket_id = Some(socket_id.into());
self
}
pub fn info(mut self, info: impl Into<String>) -> Self {
self.info = Some(info.into());
self
}
pub fn tags(mut self, tags: HashMap<String, String>) -> Self {
self.tags = Some(tags);
self
}
pub fn build(self) -> TriggerParams {
TriggerParams {
socket_id: self.socket_id,
info: self.info,
tags: self.tags,
}
}
}
#[cfg(feature = "encryption")]
fn encrypt(pusher: &Pusher, channel: &str, data: &EventData) -> Result<String> {
#[cfg(feature = "sodiumoxide")]
{
encrypt_sodiumoxide(pusher, channel, data)
}
#[cfg(not(feature = "sodiumoxide"))]
{
encrypt_pure_rust(pusher, channel, data)
}
}
#[cfg(all(feature = "encryption", feature = "sodiumoxide"))]
fn encrypt_sodiumoxide(pusher: &Pusher, channel: &str, data: &EventData) -> Result<String> {
init_sodium()?;
let _master_key =
pusher
.config()
.encryption_master_key()
.ok_or_else(|| PusherError::Encryption {
message: "Set encryptionMasterKey before triggering events on encrypted channels"
.to_string(),
})?;
let nonce_bytes =
sodiumoxide::randombytes::randombytes(sodiumoxide::crypto::secretbox::NONCEBYTES);
let nonce =
sodiumoxide::crypto::secretbox::Nonce::from_slice(&nonce_bytes).ok_or_else(|| {
PusherError::Encryption {
message: "Failed to create nonce from random bytes".to_string(),
}
})?;
let shared_secret_bytes = pusher.channel_shared_secret(channel)?;
let key =
sodiumoxide::crypto::secretbox::Key::from_slice(&shared_secret_bytes).ok_or_else(|| {
PusherError::Encryption {
message: format!(
"Channel shared secret must be {} bytes long, but was {} bytes.",
sodiumoxide::crypto::secretbox::KEYBYTES,
shared_secret_bytes.len()
),
}
})?;
let data_string = data.to_string();
let data_bytes = data_string.as_bytes();
let ciphertext = sodiumoxide::crypto::secretbox::seal(data_bytes, &nonce, &key);
let encrypted_payload = json!({
"nonce": BASE64.encode(nonce.as_ref()),
"ciphertext": BASE64.encode(&ciphertext),
});
Ok(sonic_rs::to_string(&encrypted_payload)?)
}
#[cfg(all(feature = "encryption", not(feature = "sodiumoxide")))]
fn encrypt_pure_rust(pusher: &Pusher, channel: &str, data: &EventData) -> Result<String> {
use chacha20poly1305::{
ChaCha20Poly1305, Nonce,
aead::{Aead, AeadCore, KeyInit, OsRng},
};
let _master_key =
pusher
.config()
.encryption_master_key()
.ok_or_else(|| PusherError::Encryption {
message: "Set encryptionMasterKey before triggering events on encrypted channels"
.to_string(),
})?;
let shared_secret_bytes = pusher.channel_shared_secret(channel)?;
let cipher = ChaCha20Poly1305::new_from_slice(&shared_secret_bytes).map_err(|_| {
PusherError::Encryption {
message: "Failed to create cipher from shared secret".to_string(),
}
})?;
let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
let data_string = data.to_string();
let ciphertext = cipher
.encrypt(&nonce, data_string.as_bytes())
.map_err(|_| PusherError::Encryption {
message: "Encryption failed".to_string(),
})?;
let encrypted_payload = json!({
"nonce": BASE64.encode(&nonce),
"ciphertext": BASE64.encode(&ciphertext),
});
Ok(sonic_rs::to_string(&encrypted_payload)?)
}
#[cfg(not(feature = "encryption"))]
fn encrypt(_pusher: &Pusher, _channel: &str, _data: &EventData) -> Result<String> {
Err(PusherError::Encryption {
message: "Encryption support is not enabled. Enable the 'encryption' feature to use encrypted channels.".to_string(),
})
}
pub async fn trigger<D: Into<EventData>>(
pusher: &Pusher,
channels: &[Channel],
event_name: impl AsRef<str>,
data: D,
params: Option<&TriggerParams>,
) -> Result<reqwest::Response> {
let data = data.into();
let event_name = event_name.as_ref();
if event_name.len() > 200 {
return Err(PusherError::Validation {
message: format!("Event name too long: '{}' (max 200 characters)", event_name),
});
}
let channel_strings: Vec<String> = channels.iter().map(|c| c.full_name()).collect();
if channels.len() == 1 && channels[0].is_encrypted() {
#[cfg(feature = "encryption")]
{
let encrypted_data = encrypt(pusher, &channel_strings[0], &data)?;
let mut event = Event {
name: event_name.to_string(),
data: encrypted_data,
channels: channel_strings,
socket_id: None,
info: None,
tags: None,
};
if let Some(params) = params {
event.socket_id = params.socket_id.clone();
event.info = params.info.clone();
event.tags = params.tags.clone();
}
let event_json = sonic_rs::to_value(&event)?;
pusher.post("/events", &event_json).await
}
#[cfg(not(feature = "encryption"))]
{
Err(PusherError::Encryption {
message: "Encryption support is not enabled. Enable the 'encryption' feature to use encrypted channels.".to_string(),
})
}
} else {
for channel in channels {
if channel.is_encrypted() {
return Err(PusherError::Validation {
message:
"You cannot trigger to multiple channels when using encrypted channels"
.to_string(),
});
}
}
let mut event = Event {
name: event_name.to_string(),
data: data.to_string(),
channels: channel_strings,
socket_id: None,
info: None,
tags: None,
};
if let Some(params) = params {
event.socket_id = params.socket_id.clone();
event.info = params.info.clone();
event.tags = params.tags.clone();
}
let event_json = sonic_rs::to_value(&event)?;
pusher.post("/events", &event_json).await
}
}
pub async fn trigger_on_channels<D: Into<EventData>>(
pusher: &Pusher,
channels: &[String],
event_name: impl AsRef<str>,
data: D,
params: Option<&TriggerParams>,
) -> Result<reqwest::Response> {
let channels: Result<Vec<Channel>> = channels.iter().map(|c| Channel::from_string(c)).collect();
let channels = channels?;
trigger(pusher, &channels, event_name, data, params).await
}
pub async fn trigger_batch(
pusher: &Pusher,
mut batch: Vec<BatchEvent>,
) -> Result<reqwest::Response> {
if batch.is_empty() {
return Err(PusherError::Validation {
message: "Batch cannot be empty".to_string(),
});
}
if batch.len() > 10 {
return Err(PusherError::Validation {
message: format!("Batch too large: {} events (max 10)", batch.len()),
});
}
for event in &mut batch {
let channel = Channel::from_string(&event.channel)?;
if channel.is_encrypted() {
#[cfg(feature = "encryption")]
{
let data = EventData::String(event.data.clone());
event.data = encrypt(pusher, &event.channel, &data)?;
}
#[cfg(not(feature = "encryption"))]
{
return Err(PusherError::Encryption {
message: "Encryption support is not enabled. Enable the 'encryption' feature to use encrypted channels.".to_string(),
});
}
}
}
let batch_payload = json!({ "batch": batch });
pusher.post("/batch_events", &batch_payload).await
}
#[cfg(test)]
mod tests {
use super::*;
use sonic_rs::json;
#[test]
fn test_event_data_conversions() {
let data = EventData::from_string("hello");
assert_eq!(data.to_string(), "hello");
let json_data = json!({"key": "value"});
let data = EventData::from_json(json_data.clone());
assert_eq!(data.as_json().unwrap(), json_data);
let data: EventData = "test".into();
assert!(matches!(data, EventData::String(_)));
let data: EventData = json!({"test": 123}).into();
assert!(matches!(data, EventData::Json(_)));
}
#[test]
fn test_batch_event_builder() {
let event = BatchEvent::new("test-event", "test-channel", "test-data")
.with_socket_id("123.456")
.with_info("test-info");
assert_eq!(event.name, "test-event");
assert_eq!(event.channel, "test-channel");
assert_eq!(event.data, "test-data");
assert_eq!(event.socket_id, Some("123.456".to_string()));
assert_eq!(event.info, Some("test-info".to_string()));
}
#[test]
fn test_batch_event_with_tags() {
let mut tags = HashMap::new();
tags.insert("symbol".to_string(), "BONK".to_string());
tags.insert("price_usd".to_string(), "0.00001".to_string());
let event =
BatchEvent::new("test-event", "test-channel", "test-data").with_tags(tags.clone());
assert_eq!(event.tags, Some(tags));
}
#[test]
fn test_trigger_params_builder() {
let params = TriggerParams::builder()
.socket_id("123.456")
.info("test-info")
.build();
assert_eq!(params.socket_id, Some("123.456".to_string()));
assert_eq!(params.info, Some("test-info".to_string()));
}
#[test]
fn test_trigger_params_builder_with_tags() {
let mut tags = HashMap::new();
tags.insert("event_type".to_string(), "goal".to_string());
let params = TriggerParams::builder().tags(tags.clone()).build();
assert_eq!(params.tags, Some(tags));
}
}