[Fix] Stability
This commit is contained in:
parent
dfe8e6efa7
commit
ad208fd298
12 changed files with 281 additions and 98 deletions
|
|
@ -31,7 +31,7 @@ pub struct AnonymousClientConnection {
|
|||
|
||||
impl AnonymousClientConnection {
|
||||
pub async fn from_general(general: Arc<GeneralConnection>, user_id: u64) -> Arc<Self> {
|
||||
let username: String = generate_username();
|
||||
let username = generate_username(user_id);
|
||||
Arc::new(Self {
|
||||
state: general.state.clone(),
|
||||
user_id: user_id,
|
||||
|
|
@ -47,6 +47,7 @@ impl AnonymousClientConnection {
|
|||
})
|
||||
}
|
||||
pub fn start(self: Arc<Self>) {
|
||||
anonymous_manager::add_anonymous_user(self.clone());
|
||||
let self_clone = self.clone();
|
||||
tokio::spawn(async move {
|
||||
while let Ok(cv) = self_clone.receiver.receive().await {
|
||||
|
|
@ -608,10 +609,13 @@ impl AnonymousClientConnection {
|
|||
}
|
||||
}
|
||||
|
||||
/// Handle connection close
|
||||
pub async fn handle_close(&self) {
|
||||
|
||||
// TODO delete temp user
|
||||
*self.is_open.write().await = false;
|
||||
self.state
|
||||
.call_manager
|
||||
.remove_user_from_calls(self.user_id)
|
||||
.await;
|
||||
anonymous_manager::remove_anonymous_user(self.user_id).await;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
use dashmap::DashMap;
|
||||
use dashmap::mapref::entry::Entry;
|
||||
use once_cell::sync::Lazy;
|
||||
use rand::prelude::{IndexedRandom, RngExt};
|
||||
use std::sync::Arc;
|
||||
|
|
@ -7,15 +8,20 @@ use crate::anonymous_clients::anonymous_client_connection::AnonymousClientConnec
|
|||
|
||||
static ANONYMOUS_USERS: Lazy<DashMap<u64, Arc<AnonymousClientConnection>>> =
|
||||
Lazy::new(|| DashMap::new());
|
||||
static ANONYMOUS_USERNAMES: Lazy<DashMap<String, u64>> = Lazy::new(|| DashMap::new());
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub async fn add_anonymous_user(connection: Arc<AnonymousClientConnection>) {
|
||||
pub fn add_anonymous_user(connection: Arc<AnonymousClientConnection>) {
|
||||
ANONYMOUS_USERS.insert(connection.get_user_id(), connection);
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub async fn remove_anonymous_user(user_id: u64) {
|
||||
ANONYMOUS_USERS.remove(&user_id);
|
||||
if let Some((_, connection)) = ANONYMOUS_USERS.remove(&user_id) {
|
||||
let username = connection.get_user_name().await;
|
||||
ANONYMOUS_USERNAMES.remove_if(&username, |_, reserved_user_id| {
|
||||
*reserved_user_id == user_id
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_anonymous_user(user_id: u64) -> Option<Arc<AnonymousClientConnection>> {
|
||||
|
|
@ -25,30 +31,52 @@ pub async fn get_anonymous_user(user_id: u64) -> Option<Arc<AnonymousClientConne
|
|||
pub async fn get_anonymous_user_by_name(
|
||||
username: String,
|
||||
) -> Option<Arc<AnonymousClientConnection>> {
|
||||
let users: Vec<_> = ANONYMOUS_USERS
|
||||
.iter()
|
||||
.map(|ref_multi| ref_multi.value().clone())
|
||||
.collect();
|
||||
|
||||
for user_conn in users {
|
||||
if user_conn.get_user_name().await == username {
|
||||
return Some(user_conn);
|
||||
}
|
||||
}
|
||||
|
||||
return None;
|
||||
let user_id = ANONYMOUS_USERNAMES
|
||||
.get(&username.to_lowercase())?
|
||||
.value()
|
||||
.to_owned();
|
||||
get_anonymous_user(user_id).await
|
||||
}
|
||||
|
||||
// TODO: implement check if taken
|
||||
pub fn generate_username() -> String {
|
||||
pub fn generate_username(user_id: u64) -> String {
|
||||
let adjectives = ["Swift", "Clever", "Brave", "Sneaky", "Fierce"];
|
||||
let nouns = ["Tiger", "Eagle", "Shark", "Wolf", "Dragon"];
|
||||
|
||||
let mut rng = rand::rng();
|
||||
let adj = adjectives.choose(&mut rng).unwrap();
|
||||
let noun = nouns.choose(&mut rng).unwrap();
|
||||
loop {
|
||||
let Some(adj) = adjectives.choose(&mut rng) else {
|
||||
continue;
|
||||
};
|
||||
let Some(noun) = nouns.choose(&mut rng) else {
|
||||
continue;
|
||||
};
|
||||
let username = format!("{}{}{}", adj, noun, rng.random_range(0..10000));
|
||||
let canonical_username = username.to_lowercase();
|
||||
|
||||
let number: u16 = rng.random_range(0..10000);
|
||||
if reserve_username(canonical_username, user_id) {
|
||||
return username;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
format!("{}{}{}", adj, noun, number)
|
||||
fn reserve_username(username: String, user_id: u64) -> bool {
|
||||
if let Entry::Vacant(entry) = ANONYMOUS_USERNAMES.entry(username) {
|
||||
entry.insert(user_id);
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn username_reservation_is_unique() {
|
||||
let username = "anonymous-manager-reservation-test".to_string();
|
||||
assert!(reserve_username(username.clone(), 1));
|
||||
assert!(!reserve_username(username.clone(), 2));
|
||||
ANONYMOUS_USERNAMES.remove(&username);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -45,6 +45,10 @@ impl AppState {
|
|||
}
|
||||
|
||||
pub fn keyring_for_host(&self) -> Result<Keyring, String> {
|
||||
Keyring::from_bytes(&self.keyring.to_bytes()).map_err(|error| error.to_string())
|
||||
let bytes = self
|
||||
.keyring
|
||||
.try_to_bytes()
|
||||
.map_err(|error| error.to_string())?;
|
||||
Keyring::from_bytes(&bytes).map_err(|error| error.to_string())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -260,6 +260,7 @@ impl CallGroup {
|
|||
.write()
|
||||
.await
|
||||
.retain(|caller| caller.user_id != user_id);
|
||||
self.secrets.write().await.remove(&user_id);
|
||||
}
|
||||
|
||||
pub async fn get_short_link(self: Arc<Self>) -> Option<String> {
|
||||
|
|
|
|||
|
|
@ -78,6 +78,13 @@ impl CallManager {
|
|||
call_groups
|
||||
}
|
||||
|
||||
pub async fn remove_user_from_calls(&self, user_id: u64) {
|
||||
let call_groups = self.get_call_groups(user_id).await;
|
||||
for call_group in call_groups {
|
||||
call_group.remove_caller(user_id).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_call_token(&self, user_id: u64, call_id: Uuid) -> Result<String, CallError> {
|
||||
if let Some(cg) = self.groups.get(&call_id) {
|
||||
let mut members = cg.members.write().await;
|
||||
|
|
@ -221,6 +228,30 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn removing_user_from_calls_removes_membership_and_invitation_secret() {
|
||||
let call_id = Uuid::new_v4();
|
||||
let sender_id = 55;
|
||||
let anonymous_user_id = 66;
|
||||
let group = Arc::new(CallGroup::new(
|
||||
call_id,
|
||||
Arc::new(Caller::new(sender_id, call_id, true)),
|
||||
));
|
||||
let manager = CallManager::default();
|
||||
manager.groups.insert(call_id, group.clone());
|
||||
|
||||
assert!(
|
||||
manager
|
||||
.add_invite(call_id, sender_id, anonymous_user_id, envelope("anonymous"),)
|
||||
.await
|
||||
);
|
||||
|
||||
manager.remove_user_from_calls(anonymous_user_id).await;
|
||||
|
||||
assert!(group.get_caller(anonymous_user_id).await.is_none());
|
||||
assert!(group.get_secret_for_user(anonymous_user_id).await.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn self_invites_are_not_forwarded() {
|
||||
let manager = CallManager::default();
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
use std::{env, time::Duration};
|
||||
use std::{env, net::IpAddr, time::Duration};
|
||||
|
||||
use thiserror::Error;
|
||||
|
||||
const DEFAULT_RHO_PORT: u16 = 443;
|
||||
const DEFAULT_BIND_ADDRESS: &str = "0.0.0.0";
|
||||
const DEFAULT_OMEGA_HOST: &str = "tensamin.net";
|
||||
const DEFAULT_OMEGA_PORT: u16 = 9187;
|
||||
const DEFAULT_OMEGA_SYNC_TIMEOUT_SECONDS: u64 = 20;
|
||||
|
|
@ -18,6 +19,7 @@ pub struct LiveKitConfig {
|
|||
#[derive(Clone, Debug, Eq, PartialEq)]
|
||||
pub struct Config {
|
||||
pub rho_port: u16,
|
||||
pub bind_address: IpAddr,
|
||||
pub omega_host: String,
|
||||
pub omega_port: u16,
|
||||
pub omikron_id: u64,
|
||||
|
|
@ -43,6 +45,7 @@ pub enum ConfigError {
|
|||
impl Config {
|
||||
pub fn from_environment() -> Result<Self, ConfigError> {
|
||||
let rho_port = parse_or_default("RHO_PORT", DEFAULT_RHO_PORT)?;
|
||||
let bind_address = parse_bind_address(env::var("BIND_ADDRESS").ok())?;
|
||||
let omega_port = parse_or_default("OMEGA_PORT", DEFAULT_OMEGA_PORT)?;
|
||||
let omikron_id = parse_or_default("ID", 0_u64)?;
|
||||
// Each synchronization request is bounded by this timeout. After
|
||||
|
|
@ -68,6 +71,7 @@ impl Config {
|
|||
|
||||
Ok(Self {
|
||||
rho_port,
|
||||
bind_address,
|
||||
omega_host,
|
||||
omega_port,
|
||||
omikron_id,
|
||||
|
|
@ -78,6 +82,18 @@ impl Config {
|
|||
}
|
||||
}
|
||||
|
||||
fn parse_bind_address(value: Option<String>) -> Result<IpAddr, ConfigError> {
|
||||
value
|
||||
.as_deref()
|
||||
.unwrap_or(DEFAULT_BIND_ADDRESS)
|
||||
.trim()
|
||||
.parse()
|
||||
.map_err(|_| ConfigError::InvalidValue {
|
||||
name: "BIND_ADDRESS",
|
||||
kind: "IP address",
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_or_default<T>(name: &'static str, default: T) -> Result<T, ConfigError>
|
||||
where
|
||||
T: std::str::FromStr,
|
||||
|
|
@ -148,4 +164,23 @@ mod tests {
|
|||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_configured_bind_address() {
|
||||
assert_eq!(
|
||||
parse_bind_address(Some("10.200.2.0".to_string())),
|
||||
Ok(IpAddr::V4(std::net::Ipv4Addr::new(10, 200, 2, 0)))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_bind_address() {
|
||||
assert_eq!(
|
||||
parse_bind_address(Some("not-an-address".to_string())),
|
||||
Err(ConfigError::InvalidValue {
|
||||
name: "BIND_ADDRESS",
|
||||
kind: "IP address",
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
50
src/main.rs
50
src/main.rs
|
|
@ -8,7 +8,7 @@ mod rho;
|
|||
mod services;
|
||||
mod util;
|
||||
|
||||
use std::path::PathBuf;
|
||||
use std::{fmt::Write, path::PathBuf};
|
||||
|
||||
use dotenv::dotenv;
|
||||
use once_cell::sync::Lazy;
|
||||
|
|
@ -27,7 +27,7 @@ use crate::{
|
|||
omega::omega_connection::{OmegaConnection, start_task_cleanup_loop},
|
||||
rho::rho_manager::RhoManager,
|
||||
rho::server::start,
|
||||
util::logger::startup,
|
||||
util::logger::{PrintType, startup},
|
||||
};
|
||||
|
||||
const KEYRING_PATH: &str = "./omikron.mk";
|
||||
|
|
@ -71,6 +71,18 @@ async fn main() {
|
|||
return;
|
||||
}
|
||||
};
|
||||
let public_key = match omega_database_public_key(&keyring) {
|
||||
Ok(public_key) => public_key,
|
||||
Err(error) => {
|
||||
eprintln!("Unable to serialize public key for Omega: {error}");
|
||||
return;
|
||||
}
|
||||
};
|
||||
log!(
|
||||
0,
|
||||
PrintType::Omikron,
|
||||
"Omikron public_key for Omega enrollment: {public_key}"
|
||||
);
|
||||
|
||||
let omega_keyring = match keyring_for_omega(&keyring) {
|
||||
Ok(keyring) => keyring,
|
||||
|
|
@ -102,5 +114,37 @@ async fn main() {
|
|||
}
|
||||
|
||||
fn keyring_for_omega(keyring: &Keyring) -> Result<Keyring, String> {
|
||||
Keyring::from_bytes(&keyring.to_bytes()).map_err(|error| error.to_string())
|
||||
let bytes = keyring.try_to_bytes().map_err(|error| error.to_string())?;
|
||||
Keyring::from_bytes(&bytes).map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
fn omega_database_public_key(keyring: &Keyring) -> Result<String, String> {
|
||||
let public_key = keyring
|
||||
.public_key_bundle()
|
||||
.try_as_bytes()
|
||||
.map_err(|error| error.to_string())?;
|
||||
omega_database_blob_literal(&public_key)
|
||||
}
|
||||
|
||||
fn omega_database_blob_literal(public_key: &[u8]) -> Result<String, String> {
|
||||
let mut literal = String::with_capacity(3 + public_key.len() * 2);
|
||||
literal.push_str("X'");
|
||||
for byte in public_key {
|
||||
write!(&mut literal, "{byte:02X}").map_err(|error| error.to_string())?;
|
||||
}
|
||||
literal.push('\'');
|
||||
Ok(literal)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn public_key_blob_literal_is_valid_mysql_hex_syntax() {
|
||||
assert_eq!(
|
||||
omega_database_blob_literal(&[0x00, 0x1A, 0xFF]),
|
||||
Ok("X'001AFF'".to_string())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
use std::net::{IpAddr, Ipv4Addr};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
|
|
@ -75,10 +74,23 @@ pub async fn complete_register(
|
|||
}
|
||||
println!("Iota connection request");
|
||||
|
||||
let pub_key_bytes = match pub_key.try_as_bytes() {
|
||||
Ok(bytes) => bytes,
|
||||
Err(error) => {
|
||||
log_err!(
|
||||
0,
|
||||
PrintType::General,
|
||||
"Failed to serialize Iota public key: {}",
|
||||
error
|
||||
);
|
||||
return 0;
|
||||
}
|
||||
};
|
||||
|
||||
let request = CommunicationValue::new(CommunicationType::CompleteRegisterIota)
|
||||
.add_typed_default(
|
||||
DataType::PublicKey,
|
||||
DataValue::Str(BASE64_STD.encode(pub_key.as_bytes())),
|
||||
DataValue::Str(BASE64_STD.encode(pub_key_bytes)),
|
||||
);
|
||||
|
||||
let response = match omega
|
||||
|
|
@ -109,7 +121,7 @@ pub async fn start(state: Arc<AppState>) -> Result<(), Box<dyn std::error::Error
|
|||
let key_pem = load_file_vec("certs", "key.pem").expect("Error loading Keyfile");
|
||||
|
||||
let host_config = HostConfig::new(
|
||||
IpAddr::from(Ipv4Addr::new(0, 0, 0, 0)),
|
||||
state.config.bind_address,
|
||||
state.config.rho_port,
|
||||
cert_pem,
|
||||
key_pem,
|
||||
|
|
@ -154,7 +166,8 @@ pub async fn start(state: Arc<AppState>) -> Result<(), Box<dyn std::error::Error
|
|||
log!(
|
||||
0,
|
||||
PrintType::General,
|
||||
"Server listening on port {}.",
|
||||
"Server listening on {}:{}.",
|
||||
state.config.bind_address,
|
||||
state.config.rho_port
|
||||
);
|
||||
|
||||
|
|
|
|||
|
|
@ -205,7 +205,7 @@ fn format_data_container(data: Vec<(DataTypeId, DataValue)>, version: Version) -
|
|||
let key_str = key.to_string();
|
||||
|
||||
match value {
|
||||
DataValue::Str(s) => format!("{}=\"{}\"", key_str, s),
|
||||
DataValue::Str(s) => format!("{}=\"{}\"", key_str, abbreviate_string(&s)),
|
||||
|
||||
DataValue::Container(inner) => {
|
||||
let inner_formatted = format_data_container(inner, version.clone());
|
||||
|
|
@ -238,7 +238,7 @@ fn format_array(arr: Vec<DataValue>, version: Version) -> String {
|
|||
let parts: Vec<String> = arr
|
||||
.into_iter()
|
||||
.map(|value| match value {
|
||||
DataValue::Str(s) => format!("\"{}\"", s),
|
||||
DataValue::Str(s) => format!("\"{}\"", abbreviate_string(&s)),
|
||||
|
||||
DataValue::Container(inner) => {
|
||||
let inner_formatted = format_data_container(inner, version.clone());
|
||||
|
|
@ -266,6 +266,31 @@ fn format_array(arr: Vec<DataValue>, version: Version) -> String {
|
|||
parts.join(", ")
|
||||
}
|
||||
|
||||
fn abbreviate_string(value: &str) -> String {
|
||||
const EDGE_LENGTH: usize = 4;
|
||||
|
||||
let chars: Vec<char> = value.chars().collect();
|
||||
if chars.len() <= EDGE_LENGTH * 2 {
|
||||
return value.to_string();
|
||||
}
|
||||
|
||||
let prefix: String = chars.iter().take(EDGE_LENGTH).collect();
|
||||
let suffix: String = chars.iter().rev().take(EDGE_LENGTH).rev().collect();
|
||||
format!("{prefix}...{suffix}")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::abbreviate_string;
|
||||
|
||||
#[test]
|
||||
fn abbreviates_only_strings_longer_than_eight_characters() {
|
||||
assert_eq!(abbreviate_string("12345678"), "12345678");
|
||||
assert_eq!(abbreviate_string("123456789"), "1234...6789");
|
||||
assert_eq!(abbreviate_string("YWJjZGVmZ2hpag=="), "YWJj...ag==");
|
||||
}
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! log_cv {
|
||||
($kind:expr, $cv:expr) => {
|
||||
|
|
|
|||
Loading…
Reference in a new issue