#![deny(missing_docs)]
use std::collections::HashMap;
use std::io::{Error, Result};
use std::sync::{Arc, Mutex, Weak};
pub use sbd_client::{PubKey, SbdClientConfig};
pub mod protocol;
mod sodoken_crypto;
pub use sodoken_crypto::*;
pub struct Config {
pub client_config: SbdClientConfig,
pub listener: bool,
pub max_connections: usize,
pub max_idle: std::time::Duration,
}
impl Default for Config {
fn default() -> Self {
Self {
client_config: Default::default(),
listener: false,
max_connections: 4096,
max_idle: std::time::Duration::from_secs(10),
}
}
}
type ClientSync = tokio::sync::Mutex<sbd_client::SbdClient>;
pub struct MsgRecv {
inner: Arc<Mutex<Inner>>,
recv: sbd_client::MsgRecv,
client: Weak<ClientSync>,
}
impl MsgRecv {
pub async fn recv(&mut self) -> Option<(PubKey, bytes::Bytes)> {
while let Some(msg) = self.recv.recv().await {
let pk = msg.pub_key();
let dec = self.inner.lock().unwrap().dec(msg);
match dec {
DecRes::Ok(msg) => return Some((pk, msg)),
DecRes::Ignore => (),
DecRes::ReqNewStream => {
if let Some(client) = self.client.upgrade() {
let msg =
protocol::Protocol::request_new_stream(&*pk.0);
if let Err(err) =
client.lock().await.send(&pk, msg.base_msg()).await
{
tracing::debug!(?err, "failure sending request_new_stream in message receive handler");
}
} else {
return None;
}
}
}
}
None
}
}
pub struct SbdClientCrypto {
pub_key: PubKey,
inner: Arc<Mutex<Inner>>,
client: Arc<ClientSync>,
}
impl SbdClientCrypto {
pub async fn new(
url: &str,
config: Arc<Config>,
) -> Result<(Self, MsgRecv)> {
let crypto = SodokenCrypto::new()?;
use sbd_client::Crypto;
let pub_key = PubKey(Arc::new(*crypto.pub_key()));
let (client, recv) = sbd_client::SbdClient::connect_config(
url,
&crypto,
config.client_config.clone(),
)
.await?;
let client = Arc::new(tokio::sync::Mutex::new(client));
let inner = Arc::new(Mutex::new(Inner::new(config, crypto)));
let recv = MsgRecv {
inner: inner.clone(),
recv,
client: Arc::downgrade(&client),
};
let this = Self {
pub_key,
inner,
client,
};
Ok((this, recv))
}
pub fn pub_key(&self) -> &PubKey {
&self.pub_key
}
pub fn active_peers(&self) -> Vec<PubKey> {
let mut inner = self.inner.lock().unwrap();
let max_idle = inner.config.max_idle;
Inner::prune(&mut inner.c_map, max_idle);
inner.c_map.keys().cloned().collect()
}
pub async fn assert(&self, pk: &PubKey) -> Result<()> {
let enc = self.inner.lock().unwrap().enc(pk, None)?;
{
let client = self.client.lock().await;
for enc in enc {
client.send(pk, &enc).await?;
}
}
Ok(())
}
pub async fn send(&self, pk: &PubKey, msg: &[u8]) -> Result<()> {
const SBD_MAX: usize = 20_000;
const SBD_HDR: usize = 32;
const SBD_SS_HDR: usize = 1;
const SS_ABYTES: usize = sodoken::secretstream::ABYTES;
const MAX_MSG: usize = SBD_MAX - SBD_HDR - SBD_SS_HDR - SS_ABYTES;
if msg.len() > MAX_MSG {
return Err(Error::other("message too long"));
}
let enc = self.inner.lock().unwrap().enc(pk, Some(msg))?;
{
let client = self.client.lock().await;
for enc in enc {
client.send(pk, &enc).await?;
}
}
Ok(())
}
pub async fn close_peer(&self, pk: &PubKey) {
self.inner.lock().unwrap().close(pk);
}
pub async fn close(&self) {
self.client.lock().await.close().await;
}
}
enum DecRes {
Ok(bytes::Bytes),
Ignore,
ReqNewStream,
}
struct InnerRec {
enc: Option<Encryptor>,
dec: Option<Decryptor>,
last_active: std::time::Instant,
}
impl InnerRec {
pub fn new() -> Self {
Self {
enc: None,
dec: None,
last_active: std::time::Instant::now(),
}
}
}
struct Inner {
config: Arc<Config>,
crypto: SodokenCrypto,
c_map: HashMap<PubKey, InnerRec>,
}
impl Inner {
fn new(config: Arc<Config>, crypto: SodokenCrypto) -> Self {
Self {
config,
crypto,
c_map: HashMap::new(),
}
}
fn close(&mut self, pk: &PubKey) {
self.c_map.remove(pk);
}
fn prune(
c_map: &mut HashMap<PubKey, InnerRec>,
max_idle: std::time::Duration,
) {
let now = std::time::Instant::now();
c_map.retain(|_pk, r| now - r.last_active < max_idle);
}
fn loc_assert<'a>(
config: &'a Config,
c_map: &'a mut HashMap<PubKey, InnerRec>,
pk: PubKey,
do_create: bool,
) -> Result<&'a mut InnerRec> {
use std::collections::hash_map::Entry;
let tot = c_map.len();
Self::prune(c_map, config.max_idle);
match c_map.entry(pk.clone()) {
Entry::Vacant(e) => {
if do_create {
if tot >= config.max_connections {
return Err(Error::other("too many connections"));
}
Ok(e.insert(InnerRec::new()))
} else {
Err(Error::other("ignore unsolicited"))
}
}
Entry::Occupied(e) => Ok(e.into_mut()),
}
}
fn enc(
&mut self,
pk: &PubKey,
msg: Option<&[u8]>,
) -> Result<Vec<bytes::Bytes>> {
let Self {
config,
crypto,
c_map,
} = self;
let mut out = Vec::new();
let rec = Self::loc_assert(config, c_map, pk.clone(), true)?;
rec.last_active = std::time::Instant::now();
if rec.enc.is_none() {
let (enc, hdr) = crypto.new_enc(pk)?;
rec.enc = Some(enc);
let msg = protocol::Protocol::new_stream(&**pk, &hdr);
out.push(msg.base_msg().clone());
}
if let Some(msg) = msg {
let msg = rec.enc.as_mut().unwrap().encrypt(&*pk.0, msg)?;
out.push(msg.base_msg().clone());
}
Ok(out)
}
fn dec(&mut self, msg: sbd_client::Msg) -> DecRes {
let Self {
config,
crypto,
c_map,
} = self;
let rec = match Self::loc_assert(
config,
c_map,
msg.pub_key(),
config.listener,
) {
Ok(rec) => rec,
Err(_) => {
return DecRes::Ignore;
}
};
rec.last_active = std::time::Instant::now();
let dec = match protocol::Protocol::from_full(
bytes::Bytes::copy_from_slice(&msg.0),
) {
Some(dec) => dec,
None => {
rec.dec = None;
return DecRes::ReqNewStream;
}
};
match dec {
protocol::Protocol::NewStream { header, .. } => {
let dec =
match crypto.new_dec(msg.pub_key_ref(), header.as_ref()) {
Ok(dec) => dec,
Err(_) => return DecRes::ReqNewStream,
};
rec.dec = Some(dec);
DecRes::Ignore
}
protocol::Protocol::Message { message, .. } => {
match rec.dec.as_mut() {
Some(dec) => match dec.decrypt(message.as_ref()) {
Ok(message) => DecRes::Ok(message),
Err(_) => DecRes::ReqNewStream,
},
None => {
DecRes::Ignore
}
}
}
protocol::Protocol::RequestNewStream { .. } => {
rec.enc = None;
DecRes::Ignore
}
}
}
}
#[cfg(test)]
mod test;