use std::marker::PhantomData;
use crate::{
derec_message::DeRecMessageBuilderError,
protocol_version::ProtocolVersion,
types::ChannelId,
};
use derec_proto::{DeRecMessage, MessageBody};
use prost_types::Timestamp;
#[derive(Debug)]
pub struct NotEncrypted;
#[derive(Debug)]
pub struct Encrypted;
#[derive(Debug)]
pub struct PairingMode;
#[derive(Debug)]
pub struct ChannelMode;
#[derive(Debug)]
pub struct DeRecMessageBuilder<State, Mode> {
pub(crate) sequence: Option<u32>,
pub(crate) channel_id: Option<ChannelId>,
pub(crate) timestamp: Option<Timestamp>,
pub(crate) message: Option<MessageBody>,
pub(crate) trace_id: Option<u64>,
encrypted: Vec<u8>,
_state: PhantomData<State>,
_mode: PhantomData<Mode>,
}
impl<State, Mode> DeRecMessageBuilder<State, Mode> {
pub fn channel_id(mut self, channel_id: ChannelId) -> Self {
self.channel_id = Some(channel_id);
self
}
pub fn sequence(mut self, sequence: u32) -> Self {
self.sequence = Some(sequence);
self
}
pub fn timestamp(mut self, timestamp: Timestamp) -> Self {
self.timestamp = Some(timestamp);
self
}
pub fn trace_id(mut self, trace_id: u64) -> Self {
self.trace_id = Some(trace_id);
self
}
pub fn auto_trace_id(mut self) -> Self {
use rand::Rng as _;
self.trace_id = Some(rand::rng().next_u64());
self
}
pub fn message_body(mut self, body: MessageBody) -> Self {
self.message = Some(body);
self
}
}
impl DeRecMessageBuilder<NotEncrypted, PairingMode> {
pub fn pairing() -> Self {
Self {
sequence: None,
channel_id: None,
timestamp: None,
message: None,
trace_id: None,
encrypted: Vec::new(),
_state: PhantomData,
_mode: PhantomData,
}
}
pub fn encrypt_pairing(
self,
public_key: impl AsRef<[u8]>,
) -> Result<DeRecMessageBuilder<Encrypted, PairingMode>, DeRecMessageBuilderError> {
if self.message.is_none() {
return Err(DeRecMessageBuilderError::MissingMessage);
}
let encoded = self.message.unwrap().encode_to_vec();
let encrypted =
derec_cryptography::pairing::envelope::encrypt(&encoded, public_key.as_ref())?;
Ok(DeRecMessageBuilder {
message: None,
encrypted,
timestamp: self.timestamp,
sequence: self.sequence,
channel_id: self.channel_id,
trace_id: self.trace_id,
_state: PhantomData,
_mode: PhantomData,
})
}
}
impl DeRecMessageBuilder<NotEncrypted, ChannelMode> {
pub fn channel() -> Self {
Self {
sequence: None,
channel_id: None,
timestamp: None,
message: None,
trace_id: None,
encrypted: Vec::new(),
_state: PhantomData,
_mode: PhantomData,
}
}
pub fn encrypt(
self,
shared_key: &[u8; 32],
) -> Result<DeRecMessageBuilder<Encrypted, ChannelMode>, DeRecMessageBuilderError> {
if self.message.is_none() {
return Err(DeRecMessageBuilderError::MissingMessage);
}
let channel_id = self
.channel_id
.ok_or(DeRecMessageBuilderError::MissingChannelId)?;
let mut nonce = [0u8; 32];
nonce[24..].copy_from_slice(&u64::from(channel_id).to_be_bytes());
let encoded = self.message.unwrap().encode_to_vec();
let encrypted = derec_cryptography::channel::encrypt_message(&encoded, shared_key, &nonce)?;
Ok(DeRecMessageBuilder {
message: None,
encrypted,
timestamp: self.timestamp,
sequence: self.sequence,
channel_id: Some(channel_id),
trace_id: self.trace_id,
_state: PhantomData,
_mode: PhantomData,
})
}
}
impl<Mode> DeRecMessageBuilder<Encrypted, Mode> {
pub fn build(self) -> Result<DeRecMessage, DeRecMessageBuilderError> {
let protocol_version = ProtocolVersion::current();
if self.timestamp.is_none() {
return Err(DeRecMessageBuilderError::MissingTimestamp);
}
if self.encrypted.is_empty() {
return Err(DeRecMessageBuilderError::MissingMessage);
}
Ok(DeRecMessage {
protocol_version_major: protocol_version.major,
protocol_version_minor: protocol_version.minor,
sequence: self.sequence.unwrap_or_default(),
channel_id: self
.channel_id
.ok_or(DeRecMessageBuilderError::MissingChannelId)?
.into(),
timestamp: Some(
self.timestamp
.ok_or(DeRecMessageBuilderError::MissingTimestamp)?,
),
message: self.encrypted,
trace_id: self.trace_id.unwrap_or_default(),
})
}
}
#[cfg(not(target_arch = "wasm32"))]
pub fn current_timestamp() -> Timestamp {
use std::time::{SystemTime, UNIX_EPOCH};
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("time went backwards");
Timestamp {
seconds: now.as_secs() as i64,
nanos: now.subsec_nanos() as i32,
}
}
#[cfg(target_arch = "wasm32")]
pub fn current_timestamp() -> Timestamp {
let millis = js_sys::Date::now() as i64;
Timestamp {
seconds: millis / 1000,
nanos: ((millis % 1000) * 1_000_000) as i32,
}
}