[Fix] Limit hosted-user MTP sessions
This commit is contained in:
parent
e19c3c3d12
commit
6fa8ac8b5f
4 changed files with 49 additions and 4 deletions
|
|
@ -4,7 +4,12 @@ use iota_identity::{LocalDescriptorPublisher, LocalUserId, LocalUserStore, Signe
|
|||
use iota_logger::log;
|
||||
use mtp::host::HostConfig;
|
||||
use mtp::webserver::{HttpRequest, HttpResponse, MTPWebServer, WebServerConfig};
|
||||
use std::{net::IpAddr, path::PathBuf, sync::Arc};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
net::IpAddr,
|
||||
path::PathBuf,
|
||||
sync::{Arc, Mutex as StdMutex},
|
||||
};
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
|
@ -41,6 +46,7 @@ pub struct WebConfig {
|
|||
pub tls: Option<TlsConfig>,
|
||||
pub required: bool,
|
||||
pub max_mtp_sessions: usize,
|
||||
pub max_mtp_sessions_per_user: usize,
|
||||
pub authority_discovery: iota_identity::AuthorityDiscoveryDocument,
|
||||
pub local_users: Arc<dyn LocalUserStore>,
|
||||
pub descriptor_publisher: Arc<dyn LocalDescriptorPublisher>,
|
||||
|
|
@ -64,6 +70,23 @@ impl std::fmt::Display for WebServerError {
|
|||
}
|
||||
impl std::error::Error for WebServerError {}
|
||||
|
||||
struct UserSessionGuard {
|
||||
sessions: Arc<StdMutex<HashMap<u64, usize>>>,
|
||||
user_id: u64,
|
||||
}
|
||||
|
||||
impl Drop for UserSessionGuard {
|
||||
fn drop(&mut self) {
|
||||
let mut sessions = self.sessions.lock().unwrap();
|
||||
if let Some(count) = sessions.get_mut(&self.user_id) {
|
||||
*count -= 1;
|
||||
if *count == 0 {
|
||||
sessions.remove(&self.user_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WebServerHandle {
|
||||
cancellation: CancellationToken,
|
||||
join: Mutex<Option<JoinHandle<()>>>,
|
||||
|
|
@ -292,6 +315,8 @@ pub async fn start(
|
|||
let cancellation = parent.child_token();
|
||||
let task_cancellation = cancellation.clone();
|
||||
let session_limit = Arc::new(tokio::sync::Semaphore::new(config.max_mtp_sessions));
|
||||
let user_sessions = Arc::new(StdMutex::new(HashMap::<u64, usize>::new()));
|
||||
let max_user_sessions = config.max_mtp_sessions_per_user;
|
||||
let join = tokio::spawn(async move {
|
||||
loop {
|
||||
tokio::select! {
|
||||
|
|
@ -301,9 +326,22 @@ pub async fn start(
|
|||
log!("Rejected MTP connection: session limit reached");
|
||||
continue;
|
||||
};
|
||||
let user_guard = if connection.client_id & (1_u64 << 63) == 0 {
|
||||
let mut sessions = user_sessions.lock().unwrap();
|
||||
let count = sessions.entry(connection.client_id).or_default();
|
||||
if *count >= max_user_sessions {
|
||||
log!("Rejected MTP connection: per-user session limit reached");
|
||||
continue;
|
||||
}
|
||||
*count += 1;
|
||||
Some(UserSessionGuard { sessions: user_sessions.clone(), user_id: connection.client_id })
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let handler = mtp_handler.clone();
|
||||
tokio::spawn(async move {
|
||||
let _permit = permit;
|
||||
let _user_guard = user_guard;
|
||||
handler.accept(connection).await;
|
||||
});
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue