mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-01 00:39:24 +00:00
fix(web): fence routes to the current session (#2522)
* fix(web): publish routable sessions before connect events * fix(web): fence routes to the current session
This commit is contained in:
@@ -5,7 +5,7 @@ pub mod storage;
|
||||
|
||||
use std::sync::{
|
||||
Arc,
|
||||
atomic::{AtomicU32, Ordering},
|
||||
atomic::{AtomicU32, AtomicU64, Ordering},
|
||||
};
|
||||
use std::time::Duration;
|
||||
|
||||
@@ -62,6 +62,7 @@ pub struct ClientManager {
|
||||
tasks: JoinSet<()>,
|
||||
|
||||
listeners_cnt: Arc<AtomicU32>,
|
||||
next_session_epoch: Arc<AtomicU64>,
|
||||
|
||||
client_sessions: Arc<DashMap<url::Url, Arc<Session>>>,
|
||||
storage: Storage,
|
||||
@@ -94,6 +95,7 @@ impl ClientManager {
|
||||
tasks,
|
||||
|
||||
listeners_cnt: Arc::new(AtomicU32::new(0)),
|
||||
next_session_epoch: Arc::new(AtomicU64::new(0)),
|
||||
|
||||
client_sessions,
|
||||
storage: Storage::new(db),
|
||||
@@ -115,6 +117,7 @@ impl ClientManager {
|
||||
let sessions = self.client_sessions.clone();
|
||||
let storage = self.storage.weak_ref();
|
||||
let listeners_cnt = self.listeners_cnt.clone();
|
||||
let next_session_epoch = self.next_session_epoch.clone();
|
||||
let geoip_db = self.geoip_db.clone();
|
||||
let heartbeat_min_response_delay = self.heartbeat_min_response_delay;
|
||||
let feature_flags = self.feature_flags.clone();
|
||||
@@ -148,9 +151,12 @@ impl ClientManager {
|
||||
heartbeat_min_response_delay,
|
||||
feature_flags.clone(),
|
||||
webhook_config.clone(),
|
||||
next_session_epoch.fetch_add(1, Ordering::Relaxed) + 1,
|
||||
);
|
||||
session.serve(tunnel).await;
|
||||
sessions.insert(client_url, Arc::new(session));
|
||||
let session = Arc::new(session);
|
||||
sessions.insert(client_url, session.clone());
|
||||
session.mark_route_ready();
|
||||
}
|
||||
listeners_cnt.fetch_sub(1, Ordering::Relaxed);
|
||||
});
|
||||
@@ -427,6 +433,9 @@ mod tests {
|
||||
validate_count: Arc<AtomicUsize>,
|
||||
block_second_validate: Arc<AtomicBool>,
|
||||
allow_second_validate: Arc<AtomicBool>,
|
||||
connected_count: Arc<AtomicUsize>,
|
||||
block_connected: Arc<AtomicBool>,
|
||||
allow_connected: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
impl TestWebhookState {
|
||||
@@ -438,6 +447,9 @@ mod tests {
|
||||
validate_count: Arc::new(AtomicUsize::new(0)),
|
||||
block_second_validate: Arc::new(AtomicBool::new(false)),
|
||||
allow_second_validate: Arc::new(AtomicBool::new(true)),
|
||||
connected_count: Arc::new(AtomicUsize::new(0)),
|
||||
block_connected: Arc::new(AtomicBool::new(false)),
|
||||
allow_connected: Arc::new(AtomicBool::new(true)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -450,6 +462,13 @@ mod tests {
|
||||
state
|
||||
}
|
||||
|
||||
fn with_blocked_connected(validate_responses: impl IntoIterator<Item = bool>) -> Self {
|
||||
let state = Self::new(validate_responses);
|
||||
state.block_connected.store(true, Ordering::Release);
|
||||
state.allow_connected.store(false, Ordering::Release);
|
||||
state
|
||||
}
|
||||
|
||||
fn allow_second_validate(&self) {
|
||||
self.allow_second_validate.store(true, Ordering::Release);
|
||||
}
|
||||
@@ -457,6 +476,14 @@ mod tests {
|
||||
fn validate_count(&self) -> usize {
|
||||
self.validate_count.load(Ordering::Acquire)
|
||||
}
|
||||
|
||||
fn allow_connected(&self) {
|
||||
self.allow_connected.store(true, Ordering::Release);
|
||||
}
|
||||
|
||||
fn connected_count(&self) -> usize {
|
||||
self.connected_count.load(Ordering::Acquire)
|
||||
}
|
||||
}
|
||||
|
||||
async fn validate_token_handler(
|
||||
@@ -489,6 +516,18 @@ mod tests {
|
||||
Json(json!({}))
|
||||
}
|
||||
|
||||
async fn node_connected_handler(
|
||||
State(state): State<TestWebhookState>,
|
||||
) -> Json<serde_json::Value> {
|
||||
state.connected_count.fetch_add(1, Ordering::AcqRel);
|
||||
while state.block_connected.load(Ordering::Acquire)
|
||||
&& !state.allow_connected.load(Ordering::Acquire)
|
||||
{
|
||||
tokio::time::sleep(Duration::from_millis(20)).await;
|
||||
}
|
||||
Json(json!({}))
|
||||
}
|
||||
|
||||
async fn test_webhook_config() -> (
|
||||
crate::webhook::SharedWebhookConfig,
|
||||
tokio::task::JoinHandle<()>,
|
||||
@@ -507,7 +546,7 @@ mod tests {
|
||||
) {
|
||||
let app = Router::new()
|
||||
.route("/validate-token", post(validate_token_handler))
|
||||
.route("/webhook/node-connected", post(webhook_ack_handler))
|
||||
.route("/webhook/node-connected", post(node_connected_handler))
|
||||
.route("/webhook/node-disconnected", post(webhook_ack_handler))
|
||||
.with_state(state.clone());
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
@@ -541,6 +580,42 @@ mod tests {
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn connected_webhook_observes_a_routable_current_session() {
|
||||
let webhook_state = TestWebhookState::with_blocked_connected([true]);
|
||||
let (webhook_config, webhook_server, webhook_state) =
|
||||
test_webhook_config_with_state(webhook_state).await;
|
||||
let mut mgr = ClientManager::new(
|
||||
Db::memory_db().await,
|
||||
None,
|
||||
Duration::ZERO,
|
||||
Arc::new(FeatureFlags::default()),
|
||||
webhook_config,
|
||||
);
|
||||
let config_server_addr = add_random_udp_listener(&mut mgr).await;
|
||||
let machine_id = uuid::Uuid::new_v4();
|
||||
let _client = start_web_client_for_test(
|
||||
config_server_addr,
|
||||
machine_id,
|
||||
Arc::new(native_instance_manager()),
|
||||
)
|
||||
.await;
|
||||
|
||||
wait_for_condition(
|
||||
|| async { webhook_state.connected_count() == 1 },
|
||||
Duration::from_secs(12),
|
||||
)
|
||||
.await;
|
||||
let user_id = wait_for_validated_user(&mgr, machine_id).await;
|
||||
let session = mgr
|
||||
.get_session_by_machine_id(user_id, &machine_id)
|
||||
.expect("connected target must already resolve to a session");
|
||||
assert!(session.is_running());
|
||||
|
||||
webhook_state.allow_connected();
|
||||
webhook_server.abort();
|
||||
}
|
||||
|
||||
async fn wait_for_validated_user(mgr: &ClientManager, machine_id: uuid::Uuid) -> i32 {
|
||||
tokio::time::timeout(Duration::from_secs(12), async {
|
||||
loop {
|
||||
|
||||
@@ -67,6 +67,7 @@ pub struct SessionData {
|
||||
webhook_connected_binding_version: Option<u64>,
|
||||
webhook_validation_dirty: bool,
|
||||
webhook_validation_notify: Arc<Notify>,
|
||||
session_epoch: u64,
|
||||
}
|
||||
|
||||
impl SessionData {
|
||||
@@ -96,6 +97,7 @@ impl SessionData {
|
||||
webhook_connected_binding_version: None,
|
||||
webhook_validation_dirty: false,
|
||||
webhook_validation_notify: Arc::new(Notify::new()),
|
||||
session_epoch: 0,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -252,7 +254,7 @@ impl Drop for SessionData {
|
||||
if let Ok(storage) = Storage::try_from(self.storage.clone())
|
||||
&& let Some(token) = self.storage_token.as_ref()
|
||||
{
|
||||
storage.remove_client(token);
|
||||
storage.remove_session_client(token, self.session_epoch);
|
||||
|
||||
// Notify the webhook receiver when a node disconnects.
|
||||
if self.webhook_config.is_enabled()
|
||||
@@ -426,7 +428,12 @@ impl SessionRpcService {
|
||||
let authorized = data.auth_state.is_authorized();
|
||||
if let Some(storage_token) = data.storage_token.clone() {
|
||||
let report_time = Self::heartbeat_report_timestamp(&runtime_req);
|
||||
storage.update_client(storage_token, report_time, authorized);
|
||||
storage.update_session_client(
|
||||
storage_token,
|
||||
report_time,
|
||||
authorized,
|
||||
data.session_epoch,
|
||||
);
|
||||
}
|
||||
let runtime_notify = (authorized && data.storage_token.is_some())
|
||||
.then(|| (data.notifier.clone(), runtime_req));
|
||||
@@ -498,7 +505,7 @@ impl SessionRpcService {
|
||||
}
|
||||
};
|
||||
|
||||
let (storage_token, notifier, runtime_req) = {
|
||||
let (storage_token, notifier, runtime_req, session_epoch) = {
|
||||
let mut data = self.data.write().await;
|
||||
let is_new_storage_token = data.storage_token.is_none();
|
||||
let runtime_req = Self::store_latest_heartbeat_req(&mut data, req.clone());
|
||||
@@ -519,11 +526,16 @@ impl SessionRpcService {
|
||||
tracing::error!("Heartbeat succeeded before session token was initialized");
|
||||
return Ok(HeartbeatResponse {});
|
||||
};
|
||||
(storage_token, data.notifier.clone(), runtime_req)
|
||||
(
|
||||
storage_token,
|
||||
data.notifier.clone(),
|
||||
runtime_req,
|
||||
data.session_epoch,
|
||||
)
|
||||
};
|
||||
|
||||
let report_time = Self::heartbeat_report_timestamp(&runtime_req);
|
||||
storage.update_client(storage_token, report_time, true);
|
||||
storage.update_session_client(storage_token, report_time, true, session_epoch);
|
||||
let _ = notifier.send(runtime_req);
|
||||
Ok(HeartbeatResponse {})
|
||||
}
|
||||
@@ -575,6 +587,7 @@ pub struct Session {
|
||||
|
||||
webhook_validation_task: Option<AbortOnDropHandle<()>>,
|
||||
config_reconcile_task: Option<AbortOnDropHandle<()>>,
|
||||
route_ready: Arc<Notify>,
|
||||
}
|
||||
|
||||
impl Debug for Session {
|
||||
@@ -594,9 +607,11 @@ impl Session {
|
||||
heartbeat_min_response_delay: Duration,
|
||||
feature_flags: Arc<FeatureFlags>,
|
||||
webhook_config: SharedWebhookConfig,
|
||||
session_epoch: u64,
|
||||
) -> Self {
|
||||
let session_data =
|
||||
let mut session_data =
|
||||
SessionData::new(storage, client_url, location, feature_flags, webhook_config);
|
||||
session_data.session_epoch = session_epoch;
|
||||
let data = Arc::new(RwLock::new(session_data));
|
||||
|
||||
let rpc_mgr =
|
||||
@@ -615,6 +630,7 @@ impl Session {
|
||||
data,
|
||||
webhook_validation_task: None,
|
||||
config_reconcile_task: None,
|
||||
route_ready: Arc::new(Notify::new()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -623,10 +639,13 @@ impl Session {
|
||||
|
||||
let data = self.data.read().await;
|
||||
if data.webhook_config.is_enabled() {
|
||||
let route_ready = self.route_ready.clone();
|
||||
let session_data = Arc::downgrade(&self.data);
|
||||
self.webhook_validation_task
|
||||
.replace(AbortOnDropHandle::new(tokio::spawn(
|
||||
webhook_validation::run_worker(Arc::downgrade(&self.data)),
|
||||
)));
|
||||
.replace(AbortOnDropHandle::new(tokio::spawn(async move {
|
||||
route_ready.notified().await;
|
||||
webhook_validation::run_worker(session_data).await;
|
||||
})));
|
||||
}
|
||||
self.config_reconcile_task
|
||||
.replace(AbortOnDropHandle::new(tokio::spawn(
|
||||
@@ -640,6 +659,10 @@ impl Session {
|
||||
)));
|
||||
}
|
||||
|
||||
pub fn mark_route_ready(&self) {
|
||||
self.route_ready.notify_one();
|
||||
}
|
||||
|
||||
pub fn is_running(&self) -> bool {
|
||||
self.rpc_mgr.is_running()
|
||||
}
|
||||
@@ -796,6 +819,8 @@ mod tests {
|
||||
struct ValidateWebhookTestState {
|
||||
received: Arc<Mutex<Option<oneshot::Sender<()>>>>,
|
||||
release: Arc<Notify>,
|
||||
connected_received: Arc<Mutex<Option<oneshot::Sender<()>>>>,
|
||||
connected_release: Option<Arc<Notify>>,
|
||||
}
|
||||
|
||||
async fn valid_validate_token_handler(
|
||||
@@ -813,11 +838,25 @@ mod tests {
|
||||
}))
|
||||
}
|
||||
|
||||
async fn node_connected_handler(
|
||||
State(state): State<ValidateWebhookTestState>,
|
||||
) -> Json<serde_json::Value> {
|
||||
if let Some(sender) = state.connected_received.lock().await.take() {
|
||||
let _ = sender.send(());
|
||||
}
|
||||
if let Some(release) = state.connected_release {
|
||||
release.notified().await;
|
||||
}
|
||||
|
||||
Json(json!({}))
|
||||
}
|
||||
|
||||
async fn test_webhook_config(
|
||||
state: ValidateWebhookTestState,
|
||||
) -> (SharedWebhookConfig, tokio::task::JoinHandle<()>) {
|
||||
let app = Router::new()
|
||||
.route("/validate-token", post(valid_validate_token_handler))
|
||||
.route("/webhook/node-connected", post(node_connected_handler))
|
||||
.with_state(state);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
@@ -907,10 +946,14 @@ mod tests {
|
||||
let req = heartbeat_request("token", machine_id);
|
||||
let storage = Storage::new(crate::db::Db::memory_db().await);
|
||||
let (received_tx, received_rx) = oneshot::channel();
|
||||
let (connected_tx, connected_rx) = oneshot::channel();
|
||||
let release = Arc::new(Notify::new());
|
||||
let connected_release = Arc::new(Notify::new());
|
||||
let (webhook_config, server) = test_webhook_config(ValidateWebhookTestState {
|
||||
received: Arc::new(Mutex::new(Some(received_tx))),
|
||||
release: release.clone(),
|
||||
connected_received: Arc::new(Mutex::new(Some(connected_tx))),
|
||||
connected_release: Some(connected_release.clone()),
|
||||
})
|
||||
.await;
|
||||
let mut session = SessionData::new(
|
||||
@@ -938,6 +981,18 @@ mod tests {
|
||||
));
|
||||
received_rx.await.unwrap();
|
||||
release.notify_waiters();
|
||||
connected_rx.await.unwrap();
|
||||
let user_id = storage
|
||||
.db()
|
||||
.get_user_id_by_token("token")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
storage.get_client_url_by_machine_id(user_id, &machine_id),
|
||||
Some(url::Url::parse("http://127.0.0.1").unwrap())
|
||||
);
|
||||
connected_release.notify_waiters();
|
||||
validation.await.unwrap().unwrap();
|
||||
server.abort();
|
||||
|
||||
@@ -970,6 +1025,8 @@ mod tests {
|
||||
let (webhook_config, server) = test_webhook_config(ValidateWebhookTestState {
|
||||
received: Arc::new(Mutex::new(Some(received_tx))),
|
||||
release,
|
||||
connected_received: Arc::new(Mutex::new(None)),
|
||||
connected_release: None,
|
||||
})
|
||||
.await;
|
||||
let mut data = SessionData::new(
|
||||
|
||||
@@ -228,7 +228,7 @@ pub(super) async fn apply_rejected(
|
||||
let Some(session_data) = session_data.upgrade() else {
|
||||
return;
|
||||
};
|
||||
let (storage_token, disconnect_notification) = {
|
||||
let (storage_token, disconnect_notification, session_epoch) = {
|
||||
let mut data = session_data.write().await;
|
||||
if !data.req.as_ref().is_some_and(|req| {
|
||||
SessionRpcService::heartbeat_matches_identity(
|
||||
@@ -258,13 +258,13 @@ pub(super) async fn apply_rejected(
|
||||
binding_version,
|
||||
})
|
||||
});
|
||||
(storage_token, disconnect_notification)
|
||||
(storage_token, disconnect_notification, data.session_epoch)
|
||||
};
|
||||
if let Some(storage_token) = storage_token {
|
||||
let report_time = SessionRpcService::heartbeat_report_timestamp(&input.req);
|
||||
input
|
||||
.storage
|
||||
.update_client(storage_token, report_time, false);
|
||||
.update_session_client(storage_token, report_time, false, session_epoch);
|
||||
}
|
||||
if disconnect_notification.is_some() {
|
||||
wait_webhook_connection_transition(
|
||||
@@ -290,7 +290,14 @@ pub(super) async fn apply_success(
|
||||
let Some(session_data) = session_data.upgrade() else {
|
||||
return;
|
||||
};
|
||||
let (storage_token, notifier, disconnect_notification, connect_notification, runtime_req) = {
|
||||
let (
|
||||
storage_token,
|
||||
notifier,
|
||||
disconnect_notification,
|
||||
connect_notification,
|
||||
runtime_req,
|
||||
session_epoch,
|
||||
) = {
|
||||
let mut data = session_data.write().await;
|
||||
let Some(runtime_req) = data.req.clone() else {
|
||||
return;
|
||||
@@ -367,9 +374,15 @@ pub(super) async fn apply_success(
|
||||
disconnect_notification,
|
||||
connect_notification,
|
||||
runtime_req,
|
||||
data.session_epoch,
|
||||
)
|
||||
};
|
||||
|
||||
let report_time = SessionRpcService::heartbeat_report_timestamp(&runtime_req);
|
||||
input
|
||||
.storage
|
||||
.update_session_client(storage_token, report_time, true, session_epoch);
|
||||
|
||||
if disconnect_notification.is_some() || connect_notification.is_some() {
|
||||
wait_webhook_connection_transition(
|
||||
Arc::downgrade(&session_data),
|
||||
@@ -379,10 +392,6 @@ pub(super) async fn apply_success(
|
||||
.await;
|
||||
}
|
||||
|
||||
let report_time = SessionRpcService::heartbeat_report_timestamp(&runtime_req);
|
||||
input
|
||||
.storage
|
||||
.update_client(storage_token, report_time, true);
|
||||
let _ = notifier.send(runtime_req);
|
||||
}
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ struct ClientInfo {
|
||||
storage_token: StorageToken,
|
||||
report_time: i64,
|
||||
authorized: bool,
|
||||
session_epoch: u64,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
@@ -46,10 +47,15 @@ impl Storage {
|
||||
}))
|
||||
}
|
||||
|
||||
fn remove_client_info_map(map: &DashMap<uuid::Uuid, ClientInfo>, stoken: &StorageToken) {
|
||||
fn remove_client_info_map(
|
||||
map: &DashMap<uuid::Uuid, ClientInfo>,
|
||||
stoken: &StorageToken,
|
||||
session_epoch: u64,
|
||||
) {
|
||||
map.remove_if(&stoken.machine_id, |_, v| {
|
||||
v.storage_token.client_url == stoken.client_url
|
||||
&& v.storage_token.user_id == stoken.user_id
|
||||
&& v.session_epoch == session_epoch
|
||||
});
|
||||
}
|
||||
|
||||
@@ -59,7 +65,9 @@ impl Storage {
|
||||
let same_client = e.storage_token.client_url
|
||||
== client_info.storage_token.client_url
|
||||
&& e.storage_token.user_id == client_info.storage_token.user_id;
|
||||
let should_replace = if (same_client && e.authorized != client_info.authorized)
|
||||
let should_replace = if e.session_epoch != client_info.session_epoch {
|
||||
e.session_epoch < client_info.session_epoch
|
||||
} else if (same_client && e.authorized != client_info.authorized)
|
||||
|| (!e.authorized && client_info.authorized)
|
||||
{
|
||||
true
|
||||
@@ -79,22 +87,38 @@ impl Storage {
|
||||
.or_insert(client_info.clone());
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub fn update_client(&self, stoken: StorageToken, report_time: i64, authorized: bool) {
|
||||
self.update_session_client(stoken, report_time, authorized, 0);
|
||||
}
|
||||
|
||||
pub(super) fn update_session_client(
|
||||
&self,
|
||||
stoken: StorageToken,
|
||||
report_time: i64,
|
||||
authorized: bool,
|
||||
session_epoch: u64,
|
||||
) {
|
||||
let inner = self.0.user_clients_map.entry(stoken.user_id).or_default();
|
||||
|
||||
let client_info = ClientInfo {
|
||||
storage_token: stoken.clone(),
|
||||
report_time,
|
||||
authorized,
|
||||
session_epoch,
|
||||
};
|
||||
Self::update_client_info_map(&inner, &client_info);
|
||||
}
|
||||
|
||||
pub fn remove_client(&self, stoken: &StorageToken) {
|
||||
self.remove_session_client(stoken, 0);
|
||||
}
|
||||
|
||||
pub(super) fn remove_session_client(&self, stoken: &StorageToken, session_epoch: u64) {
|
||||
self.0
|
||||
.user_clients_map
|
||||
.remove_if(&stoken.user_id, |_, set| {
|
||||
Self::remove_client_info_map(set, stoken);
|
||||
Self::remove_client_info_map(set, stoken, session_epoch);
|
||||
set.is_empty()
|
||||
});
|
||||
}
|
||||
@@ -238,6 +262,31 @@ mod tests {
|
||||
assert_eq!(storage.get_client_url_by_machine_id(2, &machine_id), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn newer_session_epoch_owns_route_until_it_is_removed() {
|
||||
let storage = Storage::new(Db::memory_db().await);
|
||||
let machine_id = uuid::Uuid::new_v4();
|
||||
let old = make_storage_token(1, machine_id, "tcp://127.0.0.1:1001");
|
||||
let current = make_storage_token(1, machine_id, "tcp://127.0.0.1:1002");
|
||||
|
||||
storage.update_session_client(old.clone(), 20, true, 1);
|
||||
storage.update_session_client(current.clone(), 20, true, 2);
|
||||
storage.update_session_client(old.clone(), 30, true, 1);
|
||||
|
||||
assert_eq!(
|
||||
storage.get_client_url_by_machine_id(1, &machine_id),
|
||||
Some(current.client_url.clone())
|
||||
);
|
||||
|
||||
storage.remove_session_client(&old, 1);
|
||||
assert_eq!(
|
||||
storage.get_client_url_by_machine_id(1, &machine_id),
|
||||
Some(current.client_url.clone())
|
||||
);
|
||||
storage.remove_session_client(¤t, 2);
|
||||
assert_eq!(storage.get_client_url_by_machine_id(1, &machine_id), None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_clients_returns_current_storage_tokens() {
|
||||
let storage = Storage::new(Db::memory_db().await);
|
||||
|
||||
Reference in New Issue
Block a user