From f332ff9b114fe32aa167dec11f797ce211e184a5 Mon Sep 17 00:00:00 2001 From: elliottjiang Date: Sun, 6 Sep 2026 08:02:20 +0800 Subject: [PATCH] fix(locator): keep users online while active bindings remain - preserve authoritative registration identity on locator results - gate every binding-removal event on the final active binding - serialize locator mutations with local registration events - cover stale, fresh, concurrent, and wildcard registration paths --- src/call/mod.rs | 4 + src/call/queue_config.rs | 2 + src/proxy/locator.rs | 70 ++++++- src/proxy/locator_db.rs | 311 ++++++++++++++++++++++++------ src/proxy/presence.rs | 4 + src/proxy/registrar.rs | 55 +++--- src/proxy/server.rs | 12 +- src/proxy/tests/common.rs | 20 +- src/proxy/tests/test_auth.rs | 1 + src/proxy/tests/test_registrar.rs | 239 +++++++++++++++++------ 10 files changed, 568 insertions(+), 150 deletions(-) diff --git a/src/call/mod.rs b/src/call/mod.rs index a11b52c54..9021f59e2 100644 --- a/src/call/mod.rs +++ b/src/call/mod.rs @@ -161,6 +161,8 @@ pub struct Location { pub supports_webrtc: bool, pub credential: Option, pub headers: Option>, + pub registered_username: Option, + pub registered_realm: Option, pub registered_aor: Option, pub contact_raw: Option, pub contact_params: Option>, @@ -201,6 +203,8 @@ impl std::fmt::Debug for Location { .field("last_modified", &self.last_modified) .field("supports_webrtc", &self.supports_webrtc) .field("headers", &self.headers) + .field("registered_username", &self.registered_username) + .field("registered_realm", &self.registered_realm) .field("registered_aor", &self.registered_aor) .field("contact_raw", &self.contact_raw) .field("contact_params", &self.contact_params) diff --git a/src/call/queue_config.rs b/src/call/queue_config.rs index 1a7a2bbdb..c29916858 100644 --- a/src/call/queue_config.rs +++ b/src/call/queue_config.rs @@ -267,6 +267,8 @@ impl AgentConfig { supports_webrtc: false, credential: None, headers: None, + registered_username: None, + registered_realm: None, registered_aor: None, contact_raw: None, contact_params: None, diff --git a/src/proxy/locator.rs b/src/proxy/locator.rs index 463bebbe3..83cb410ac 100644 --- a/src/proxy/locator.rs +++ b/src/proxy/locator.rs @@ -21,7 +21,7 @@ use std::{ sync::Arc, time::{Duration, Instant}, }; -use tracing::{debug, info}; +use tracing::{debug, info, warn}; #[derive(Clone, Debug)] pub enum LocatorEvent { @@ -47,6 +47,7 @@ pub struct LocatorStats { pub type LocatorEventSender = tokio::sync::broadcast::Sender; pub type LocatorEventReceiver = tokio::sync::broadcast::Receiver; +pub type LocatorEventLock = Arc>; pub type LocatorCreationFuture = Pin>> + Send>>; pub type RealmChecker = Arc Pin + Send>> + Send + Sync>; @@ -85,6 +86,10 @@ pub trait Locator: Send + Sync { /// by the registrar, [`TransportInspectorLocator`] and the sweep task in /// `server.rs`; emitting there as well would produce duplicates. fn set_event_sender(&self, _sender: Option) {} + /// Share the lock that serializes locator mutations with their local + /// registration events. Backends use it only when removing bindings on + /// their own, so all event producers observe the same ordering boundary. + fn set_event_lock(&self, _lock: Option) {} async fn register(&self, username: &str, realm: Option<&str>, location: Location) -> Result<()>; async fn has_active_bindings(&self, username: &str, realm: Option<&str>) -> Result; @@ -101,6 +106,52 @@ pub trait Locator: Send + Sync { } } +/// Keep removed bindings only for identities that have no active binding left. +/// +/// Binding cleanup and user presence are different facts: a reconnect can add +/// a fresh Contact before an older Contact expires or its transport closes. +/// Every removal path must cross this boundary before publishing an offline or +/// unregistered event, otherwise the stale binding can overwrite the fresh +/// registration state. +pub async fn locations_without_active_bindings( + locator: &dyn Locator, + removed: Vec, +) -> Vec { + let mut grouped = HashMap::<(String, Option), Vec>::new(); + for location in removed { + let Some(username) = location.registered_username.clone() else { + warn!(binding = %location.aor, "Removed binding has no registered username"); + continue; + }; + grouped + .entry((username, location.registered_realm.clone())) + .or_default() + .push(location); + } + + let mut offline = Vec::new(); + for ((username, realm), locations) in grouped { + match locator + .has_active_bindings(&username, realm.as_deref()) + .await + { + Ok(true) => debug!(%username, "Binding removed while user remains registered"), + Ok(false) => offline.extend(locations), + Err(error) => warn!( + %username, + error = %error, + "Failed to verify remaining bindings after removal" + ), + } + } + offline +} + +pub async fn sweep_offline_locations(locator: &dyn Locator) -> Result> { + let removed = locator.sweep_expired().await?; + Ok(locations_without_active_bindings(locator, removed).await) +} + // ─────────────────────────────────────────────────────────────────────────── // Shared helpers for Locator backends (DbLocator, RedisLocator, …) // @@ -480,6 +531,7 @@ impl TargetLocator for DialogTargetLocator { pub struct TransportInspectorLocator { locator_events: LocatorEventSender, locator: Arc>, + locator_event_lock: LocatorEventLock, } impl TransportInspectorLocator { @@ -487,10 +539,12 @@ impl TransportInspectorLocator { pub fn new( locator: Arc>, locator_events: LocatorEventSender, + locator_event_lock: LocatorEventLock, ) -> Box { Box::new(Self { locator, locator_events, + locator_event_lock, }) as Box } } @@ -500,11 +554,15 @@ impl TransportEventInspector for TransportInspectorLocator { async fn handle(&self, event: TransportEvent) -> Option { if let TransportEvent::Closed(conn) = &event { let addr = conn.get_remote_addr().unwrap_or_else(|| conn.get_addr()); + let _event_guard = self.locator_event_lock.lock().await; match self.locator.unregister_with_address(addr).await { Ok(Some(removed)) => { - if !removed.is_empty() { + let offline = + locations_without_active_bindings(self.locator.as_ref().as_ref(), removed) + .await; + if !offline.is_empty() { self.locator_events - .send(LocatorEvent::Offline(removed)) + .send(LocatorEvent::Offline(offline)) .ok(); } } @@ -564,13 +622,15 @@ impl Locator for MemoryLocator { &self, username: &str, realm: Option<&str>, - location: Location, + mut location: Location, ) -> Result<()> { let identifier = self.get_identifier(username, realm).await; if identifier.is_empty() { debug!(%username, "skip registering location with empty identifier"); return Ok(()); } + location.registered_username = Some(username.to_string()); + location.registered_realm = realm.map(str::to_string); let binding_key = location.binding_key(); let now = Instant::now(); @@ -844,7 +904,7 @@ impl Locator for MemoryLocator { /// e.g. browsers that close without sending a REGISTER expires=0. async fn sweep_expired(&self) -> Result> { let now = Instant::now(); - let mut removed: Vec = Vec::new(); + let mut removed = Vec::new(); self.locations.retain(|_, map| { map.retain(|_, loc| { diff --git a/src/proxy/locator_db.rs b/src/proxy/locator_db.rs index 52d3d3c21..0471be4a9 100644 --- a/src/proxy/locator_db.rs +++ b/src/proxy/locator_db.rs @@ -1,21 +1,27 @@ use super::locator::{ - Locator, LocatorEvent, LocatorEventSender, RealmChecker, UNREGISTER_GRACE_SECS, - choose_registered_aor, invalid_host_fallback, is_local_realm, is_location_expired, - now_epoch_secs, sort_locations_by_recency, + Locator, LocatorEvent, LocatorEventLock, LocatorEventSender, RealmChecker, + UNREGISTER_GRACE_SECS, choose_registered_aor, invalid_host_fallback, is_local_realm, + is_location_expired, locations_without_active_bindings, now_epoch_secs, + sort_locations_by_recency, }; use crate::call::{LOCATOR_EXPIRE_GRACE_SECS, Location}; use anyhow::Result; use async_trait::async_trait; use rsipstack::transport::SipAddr; -use sea_orm::{ActiveModelTrait, Database, QueryOrder, Set, entity::prelude::*}; +use sea_orm::{ + ActiveModelTrait, Database, QueryOrder, QuerySelect, Set, TransactionTrait, entity::prelude::*, +}; pub use sea_orm_migration::prelude::*; use sea_orm_migration::schema::{ big_integer, boolean, string_len, string_len_null, timestamp_with_time_zone as timestamp, }; use sea_orm_migration::sea_query::ColumnDef as MigrationColumnDef; -use std::time::{Duration, Instant}; +use std::{ + collections::HashSet, + time::{Duration, Instant}, +}; use tokio::sync::Mutex; -use tracing::{info, warn}; +use tracing::{debug, info, warn}; #[derive(Clone, Debug, PartialEq, Eq, DeriveEntityModel)] // ... (rest of the model) @@ -47,6 +53,7 @@ pub struct DbLocator { db: DatabaseConnection, realm_checker: Mutex>, event_sender: Mutex>, + event_lock: Mutex>, } #[derive(DeriveMigrationName)] @@ -201,6 +208,7 @@ impl DbLocator { db, realm_checker: Mutex::new(None), event_sender: Mutex::new(None), + event_lock: Mutex::new(None), }; if migrate { info!("Creating DbLocator with migration"); @@ -410,6 +418,8 @@ fn model_to_location(model: &Model, now_epoch: i64, now_instant: Instant) -> Res last_modified: Some(last_modified_instant), supports_webrtc: model.supports_webrtc, transport: Some(transport), + registered_username: Some(model.username.clone()), + registered_realm: (!model.realm.is_empty()).then(|| model.realm.clone()), registered_aor: Some(registered_aor), user_agent, home_proxy, @@ -418,6 +428,81 @@ fn model_to_location(model: &Model, now_epoch: i64, now_instant: Instant) -> Res }) } +impl DbLocator { + /// Delete only the exact expired snapshots observed by the caller. + /// A concurrent REGISTER may refresh the same row between selection and + /// cleanup, including from another process that cannot share our mutex. + async fn delete_expired_snapshots( + &self, + candidates: Vec, + now_epoch: i64, + now_instant: Instant, + ) -> Result> { + let candidates: Vec<_> = candidates + .into_iter() + .filter(|model| is_location_expired(model.expires, model.last_modified, now_epoch)) + .collect(); + if candidates.is_empty() { + return Ok(Vec::new()); + } + + let mut snapshots = Condition::any(); + let candidate_ids: Vec<_> = candidates.iter().map(|model| model.id).collect(); + for model in &candidates { + snapshots = snapshots.add( + Condition::all() + .add(Column::Id.eq(model.id)) + .add(Column::LastModified.eq(model.last_modified)) + .add(Column::Expires.eq(model.expires)), + ); + } + + let transaction = self + .db + .begin() + .await + .map_err(|e| anyhow::anyhow!("Database error starting expired cleanup: {}", e))?; + Entity::delete_many() + .filter(snapshots) + .exec(&transaction) + .await + .map_err(|e| anyhow::anyhow!("Database error deleting expired locations: {}", e))?; + let survivors: HashSet = Entity::find() + .select_only() + .column(Column::Id) + .filter(Column::Id.is_in(candidate_ids)) + .into_tuple::() + .all(&transaction) + .await + .map_err(|e| anyhow::anyhow!("Database error verifying expired cleanup: {}", e))? + .into_iter() + .collect(); + transaction + .commit() + .await + .map_err(|e| anyhow::anyhow!("Database error committing expired cleanup: {}", e))?; + + let mut removed = Vec::new(); + for model in candidates { + if survivors.contains(&model.id) { + debug!( + identifier = %format!("{}/{}", model.username, model.realm), + binding = %model.aor, + "expired registration was refreshed before cleanup" + ); + continue; + } + match model_to_location(&model, now_epoch, now_instant) { + Ok(location) => removed.push(location), + Err(e) => { + warn!(error = %e, aor = %model.aor, "deleted unparsable expired location row"); + } + } + } + Ok(removed) + } +} + #[async_trait] impl Locator for DbLocator { async fn is_local_realm(&self, realm: &str) -> bool { @@ -445,6 +530,14 @@ impl Locator for DbLocator { *lock = sender; } + fn set_event_lock(&self, event_lock: Option) { + let mut lock = self + .event_lock + .try_lock() + .expect("failed to lock event_lock"); + *lock = event_lock; + } + /// Periodically sweep expired registrations that would otherwise linger — /// e.g. browsers that vanish without a REGISTER expires=0. Returns the /// removed bindings; the caller (the sweep task in `server.rs`) is @@ -465,37 +558,23 @@ impl Locator for DbLocator { .await .map_err(|e| anyhow::anyhow!("Database error on sweep_expired lookup: {}", e))?; - let mut expired_ids = Vec::new(); - let mut expired_locations = Vec::new(); + let mut expired = Vec::new(); for model in candidates { if is_location_expired(model.expires, model.last_modified, now_epoch) { - match model_to_location(&model, now_epoch, now_instant) { - Ok(location) => { - info!( - identifier = %format!("{}/{}", model.username, model.realm), - binding = %model.aor, - "swept expired registration" - ); - expired_ids.push(model.id); - expired_locations.push(location); - } - Err(e) => { - warn!(error = %e, aor = %model.aor, "skipping unparsable expired location row"); - } - } + expired.push(model); } } - if expired_ids.is_empty() { + if expired.is_empty() { return Ok(vec![]); } - Entity::delete_many() - .filter(Column::Id.is_in(expired_ids)) - .exec(&self.db) - .await - .map_err(|e| anyhow::anyhow!("Database error on sweep_expired delete: {}", e))?; - - Ok(expired_locations) + let removed = self + .delete_expired_snapshots(expired, now_epoch, now_instant) + .await?; + for location in &removed { + info!(binding = %location.aor, "swept expired registration"); + } + Ok(removed) } async fn register( @@ -812,18 +891,11 @@ impl Locator for DbLocator { } let mut locations = Vec::new(); - let mut expired_ids = Vec::new(); - let mut expired_locations = Vec::new(); + let mut expired = Vec::new(); let now_instant = Instant::now(); for model in models { if is_location_expired(model.expires, model.last_modified, now_epoch) { - expired_ids.push(model.id); - match model_to_location(&model, now_epoch, now_instant) { - Ok(location) => expired_locations.push(location), - Err(e) => { - warn!(error = %e, aor = %model.aor, "skipping unparsable expired location row"); - } - } + expired.push(model); continue; } locations.push(model_to_location(&model, now_epoch, now_instant)?); @@ -832,21 +904,29 @@ impl Locator for DbLocator { // Best-effort cleanup of expired bindings so they don't shadow live // registrations in subsequent .invalid username lookups (which order by // recency). Expired rows were previously only skipped, never deleted. - if !expired_ids.is_empty() { - if let Err(e) = Entity::delete_many() - .filter(Column::Id.is_in(expired_ids)) - .exec(&self.db) + if !expired.is_empty() { + let event_lock = self.event_lock.lock().await.clone(); + let _event_guard = match event_lock { + Some(lock) => Some(lock.lock_owned().await), + None => None, + }; + match self + .delete_expired_snapshots(expired, now_epoch, now_instant) .await { - warn!(error = %e, "Failed to delete expired location rows during lookup"); - } else if !expired_locations.is_empty() - && let Some(sender) = self.event_sender.lock().await.clone().as_ref() - { - // The backend removed these bindings on its own (no explicit - // unregister arrived). Broadcast Offline so downstream - // consumers (presence, CC agent state, cluster peers, - // locator_webhook) observe the transition. - let _ = sender.send(LocatorEvent::Offline(expired_locations)); + Err(e) => warn!(error = %e, "Failed to delete expired location rows during lookup"), + Ok(expired_locations) if !expired_locations.is_empty() => { + // The backend removed these bindings on its own (no explicit + // unregister arrived). Broadcast Offline so downstream + // consumers observe the transition. + let offline = locations_without_active_bindings(self, expired_locations).await; + if !offline.is_empty() + && let Some(sender) = self.event_sender.lock().await.clone().as_ref() + { + let _ = sender.send(LocatorEvent::Offline(offline)); + } + } + Ok(_) => {} } } @@ -857,6 +937,7 @@ impl Locator for DbLocator { #[cfg(test)] mod tests { use super::*; + use crate::proxy::locator::sweep_offline_locations; use rsipstack::sip::{Auth, HostWithPort, Scheme, transport::Transport}; use rsipstack::transport::SipAddr; use sea_orm::{DbBackend, MockDatabase, MockExecResult}; @@ -874,6 +955,7 @@ mod tests { db, realm_checker: Mutex::new(None), event_sender: Mutex::new(None), + event_lock: Mutex::new(None), }; let location = Location { aor: rsipstack::sip::Uri { @@ -1319,12 +1401,63 @@ mod tests { .await .expect("insert never-expire row"); - // Alice registered an hour ago with expires=60 — long expired. + // Alice's old binding expired, but a newer binding for the same + // registered identity is still active. Removing the old binding must + // not report Alice offline. backdate_rows(&locator, "alice", 3600).await; + let fresh_alice_aor: rsipstack::sip::Uri = + format!("sip:{}@{}", "alice-new", "pbx.example.com") + .try_into() + .expect("fresh alice aor"); + locator + .register( + "alice", + Some("pbx.example.com"), + Location { + aor: fresh_alice_aor.clone(), + expires: 60, + destination: Some(SipAddr { + r#type: Some(Transport::Udp), + addr: "192.0.2.11:5060".try_into().expect("fresh destination"), + }), + ..Default::default() + }, + ) + .await + .expect("register fresh alice binding"); + + // Bob has no replacement binding and therefore really becomes + // offline when his expired binding is swept. + let bob_aor: rsipstack::sip::Uri = format!("sip:{}@{}", "bob", "pbx.example.com") + .try_into() + .expect("bob aor"); + locator + .register( + "bob", + Some("pbx.example.com"), + Location { + aor: bob_aor.clone(), + expires: 60, + destination: Some(SipAddr { + r#type: Some(Transport::Udp), + addr: "192.0.2.20:5060".try_into().expect("bob destination"), + }), + ..Default::default() + }, + ) + .await + .expect("register bob"); + backdate_rows(&locator, "bob", 3600).await; - let swept = locator.sweep_expired().await.expect("sweep expired"); - assert_eq!(swept.len(), 1, "only alice's expired binding is swept"); - assert_eq!(swept[0].aor, aor); + let swept = sweep_offline_locations(&locator) + .await + .expect("sweep expired"); + assert_eq!( + swept.len(), + 1, + "only the fully offline identity is reported" + ); + assert_eq!(swept[0].aor, bob_aor); assert_eq!( swept[0].transport, Some(Transport::Udp), @@ -1337,9 +1470,19 @@ mod tests { "sweep_expired must not emit Offline events" ); - // Alice's row is gone; carol's never-expire row survives. - let alice_lookup = locator.lookup(&aor).await.expect("lookup alice"); - assert!(alice_lookup.is_empty(), "expired binding must be removed"); + // Alice's stale row is gone while her fresh binding and Carol's + // never-expire row survive. + let stale_alice = Entity::find() + .filter(Column::Aor.eq(aor.to_string())) + .one(&locator.db) + .await + .expect("query stale alice"); + assert!(stale_alice.is_none(), "expired binding must be removed"); + let fresh_alice_lookup = locator + .lookup(&fresh_alice_aor) + .await + .expect("lookup fresh alice"); + assert_eq!(fresh_alice_lookup.len(), 1, "fresh binding must survive"); assert!( locator .has_active_bindings("carol", Some("pbx.example.com")) @@ -1446,4 +1589,54 @@ mod tests { .expect("query dave rows"); assert!(remaining.is_empty(), "expired row must be deleted"); } + + #[tokio::test] + async fn expired_snapshot_cleanup_preserves_a_refreshed_binding() { + let locator = DbLocator::new_with_migrate("sqlite::memory:".to_string(), true) + .await + .expect("create db locator"); + let aor: rsipstack::sip::Uri = format!("sip:{}@{}", "alice", "pbx.example.com") + .try_into() + .expect("valid aor"); + let location = Location { + aor: aor.clone(), + expires: 60, + destination: Some(SipAddr { + r#type: Some(Transport::Udp), + addr: "192.0.2.10:5060".try_into().expect("destination"), + }), + ..Default::default() + }; + locator + .register("alice", Some("pbx.example.com"), location.clone()) + .await + .expect("register stale snapshot"); + backdate_rows(&locator, "alice", 3600).await; + let stale = Entity::find() + .filter(Column::Username.eq("alice")) + .one(&locator.db) + .await + .expect("read stale snapshot") + .expect("stale row"); + + locator + .register("alice", Some("pbx.example.com"), location) + .await + .expect("refresh binding"); + let removed = locator + .delete_expired_snapshots(vec![stale], now_epoch_secs(), Instant::now()) + .await + .expect("compare-and-delete stale snapshot"); + + assert!(removed.is_empty(), "the refreshed row was not removed"); + assert_eq!( + locator + .lookup(&aor) + .await + .expect("lookup refreshed binding") + .len(), + 1, + "the refreshed binding must remain routable" + ); + } } diff --git a/src/proxy/presence.rs b/src/proxy/presence.rs index 22b9e2458..7341ad102 100644 --- a/src/proxy/presence.rs +++ b/src/proxy/presence.rs @@ -1652,6 +1652,8 @@ mod tests { supports_webrtc: true, credential: None, headers: None, + registered_username: None, + registered_realm: None, registered_aor: Some(registered), contact_raw: None, contact_params: None, @@ -1698,6 +1700,8 @@ mod tests { supports_webrtc: true, credential: None, headers: None, + registered_username: None, + registered_realm: None, registered_aor: Some(watcher), contact_raw: None, contact_params: None, diff --git a/src/proxy/registrar.rs b/src/proxy/registrar.rs index 907b9ee10..85c58fb1b 100644 --- a/src/proxy/registrar.rs +++ b/src/proxy/registrar.rs @@ -3,7 +3,7 @@ use crate::call::user::SipUser; use crate::call::{Location, TransactionCookie}; use crate::config::ProxyConfig; use crate::metrics; -use crate::proxy::locator::LocatorEvent; +use crate::proxy::locator::{LocatorEvent, locations_without_active_bindings}; use anyhow::{Result, anyhow}; use async_trait::async_trait; use rsipstack::sip::prelude::HeadersExt; @@ -11,7 +11,7 @@ use rsipstack::sip::{Header, Param, Transport, Uri}; use rsipstack::{transaction::transaction::Transaction, transport::SipAddr}; use std::{collections::HashMap, sync::Arc, time::Instant}; use tokio_util::sync::CancellationToken; -use tracing::{debug, info, warn}; +use tracing::{debug, info}; #[derive(Clone)] pub struct RegistrarModule { @@ -515,6 +515,7 @@ impl ProxyModule for RegistrarModule { return Ok(ProxyAction::Abort); } + let event_guard = self.server.locator_event_lock.lock().await; self.server .locator .unregister(user.username.as_str(), user.realm.as_deref()) @@ -525,14 +526,24 @@ impl ProxyModule for RegistrarModule { metrics::sip::unregistration(&realm); if let Some(locator_events) = &self.server.locator_events { - locator_events - .send(LocatorEvent::Unregistered(Location { + let removed = locations_without_active_bindings( + self.server.locator.as_ref().as_ref(), + vec![Location { aor: registered_aor.clone(), + registered_username: Some(user.username.clone()), + registered_realm: user.realm.clone(), registered_aor: Some(registered_aor), ..Default::default() - })) - .ok(); + }], + ) + .await; + if let Some(location) = removed.into_iter().next() { + locator_events + .send(LocatorEvent::Unregistered(location)) + .ok(); + } } + drop(event_guard); let headers = Vec::new(); tx.reply_with(rsipstack::sip::StatusCode::OK, headers, None) @@ -619,6 +630,8 @@ impl ProxyModule for RegistrarModule { }, credential: None, headers: Some(headers), + registered_username: Some(user.username.clone()), + registered_realm: user.realm.clone(), registered_aor: Some(registered_aor.clone()), contact_raw: Some(rendered_contact.clone()), contact_params: Some(entry.param_map()), @@ -640,6 +653,7 @@ impl ProxyModule for RegistrarModule { location.transport = destination.r#type; } + let _event_guard = self.server.locator_event_lock.lock().await; match self .server .locator @@ -654,26 +668,15 @@ impl ProxyModule for RegistrarModule { metrics::sip::registration_succeeded(&realm); if let Some(locator_events) = &self.server.locator_events { if location.expires == 0 { - match self - .server - .locator - .has_active_bindings(user.username.as_str(), user.realm.as_deref()) - .await - { - Ok(true) => debug!( - username = %user.username, - "Binding removed while user remains registered" - ), - Ok(false) => { - locator_events - .send(LocatorEvent::Unregistered(location)) - .ok(); - } - Err(error) => warn!( - username = %user.username, - error = %error, - "Failed to verify remaining bindings after unregister" - ), + let removed = locations_without_active_bindings( + self.server.locator.as_ref().as_ref(), + vec![location], + ) + .await; + if let Some(location) = removed.into_iter().next() { + locator_events + .send(LocatorEvent::Unregistered(location)) + .ok(); } } else { locator_events.send(LocatorEvent::Registered(location)).ok(); diff --git a/src/proxy/server.rs b/src/proxy/server.rs index b341861af..2a531ae7a 100644 --- a/src/proxy/server.rs +++ b/src/proxy/server.rs @@ -23,7 +23,8 @@ use crate::{ call::{CallRouter, DialplanInspector}, cluster_event::ClusterEventHub, locator::{ - DialogTargetLocator, LocatorEvent, LocatorEventSender, TransportInspectorLocator, + DialogTargetLocator, LocatorEvent, LocatorEventLock, LocatorEventSender, + TransportInspectorLocator, sweep_offline_locations, }, presence::PresenceManager, }, @@ -94,6 +95,7 @@ pub struct SipServerInner { pub create_route_invites: Vec, pub ignore_out_of_dialog_request: bool, pub locator_events: Option, + pub locator_event_lock: LocatorEventLock, pub sipflow_config: ArcSwap>, pub recording_policy: ArcSwap>, pub sip_flow: Option, @@ -912,9 +914,11 @@ impl SipServerBuilder { let (tx, _) = tokio::sync::broadcast::channel(12); tx }); + let locator_event_lock = Arc::new(tokio::sync::Mutex::new(())); // Let the backend report bindings it removes on its own (e.g. expired // rows pruned during `lookup`) as LocatorEvent::Offline. locator.set_event_sender(Some(locator_events.clone())); + locator.set_event_lock(Some(locator_event_lock.clone())); let locator_local_addrs = endpoint_local_addrs; let cluster_enabled = !self.cluster_peers.is_empty(); @@ -928,6 +932,7 @@ impl SipServerBuilder { .with_transport_inspector(TransportInspectorLocator::new( locator.clone(), locator_events.clone(), + locator_event_lock.clone(), )); let endpoint = endpoint_builder.build(); @@ -1036,6 +1041,7 @@ impl SipServerBuilder { { let locator_for_sweep = locator.clone(); let locator_events_for_sweep = locator_events.clone(); + let locator_event_lock_for_sweep = locator_event_lock.clone(); let sweep_token = cancel_token.child_token(); tokio::spawn(async move { // Run roughly every quarter of the shortest typical registrar @@ -1055,7 +1061,8 @@ impl SipServerBuilder { biased; _ = sweep_token.cancelled() => break, _ = ticker.tick() => { - match locator_for_sweep.sweep_expired().await { + let _event_guard = locator_event_lock_for_sweep.lock().await; + match sweep_offline_locations(locator_for_sweep.as_ref().as_ref()).await { Ok(removed) if !removed.is_empty() => { info!( count = removed.len(), @@ -1166,6 +1173,7 @@ impl SipServerBuilder { create_route_invites: self.create_route_invites, ignore_out_of_dialog_request: self.ignore_out_of_dialog_request, locator_events: Some(locator_events), + locator_event_lock, sipflow_config: ArcSwap::new(Arc::new(self.sipflow_config.clone())), recording_policy: ArcSwap::new(Arc::new(self.config.recording.clone())), sip_flow, diff --git a/src/proxy/tests/common.rs b/src/proxy/tests/common.rs index 53503d3f2..a50a44400 100644 --- a/src/proxy/tests/common.rs +++ b/src/proxy/tests/common.rs @@ -37,8 +37,24 @@ pub async fn create_test_server_with_config( } pub async fn create_test_server_with_config_and_sipflow_backend( + config: ProxyConfig, + sipflow_backend: Option>, +) -> (Arc, Arc) { + let locator = Arc::new(Box::new(MemoryLocator::new()) as Box); + create_test_server_with_dependencies(config, sipflow_backend, locator).await +} + +pub async fn create_test_server_with_config_and_locator( + config: ProxyConfig, + locator: Arc>, +) -> (Arc, Arc) { + create_test_server_with_dependencies(config, None, locator).await +} + +async fn create_test_server_with_dependencies( mut config: ProxyConfig, sipflow_backend: Option>, + locator: Arc>, ) -> (Arc, Arc) { // Add rustpbx.com to the allowed realms for testing if config.realms.is_none() { @@ -51,7 +67,6 @@ pub async fn create_test_server_with_config_and_sipflow_backend( .push("rustpbx.com".to_string()); let user_backend = Box::new(MemoryUserBackend::new(None)); - let locator = Arc::new(Box::new(MemoryLocator::new()) as Box); let config = Arc::new(config); let endpoint = rsipstack::EndpointBuilder::new().build(); @@ -81,6 +96,8 @@ pub async fn create_test_server_with_config_and_sipflow_backend( ); let (locator_events_tx, _) = tokio::sync::broadcast::channel(100); + let locator_event_lock = Arc::new(tokio::sync::Mutex::new(())); + locator.set_event_lock(Some(locator_event_lock.clone())); // Share ONE ConferenceManager between the conference_manager field and the // ConferenceServer so tests that cross between them observe the same state. @@ -116,6 +133,7 @@ pub async fn create_test_server_with_config_and_sipflow_backend( create_route_invites: Vec::new(), ignore_out_of_dialog_request: true, locator_events: Some(locator_events_tx), + locator_event_lock, sipflow_config: ArcSwap::new(Arc::new(None)), sip_flow: sipflow_backend .map(|backend| crate::callrecord::sipflow::SipFlow::new(Some(backend), Vec::new())), diff --git a/src/proxy/tests/test_auth.rs b/src/proxy/tests/test_auth.rs index 3818994d4..beb261442 100644 --- a/src/proxy/tests/test_auth.rs +++ b/src/proxy/tests/test_auth.rs @@ -515,6 +515,7 @@ async fn test_guest_call_allowed_extension() { create_route_invites: Vec::new(), ignore_out_of_dialog_request: true, locator_events: None, + locator_event_lock: Arc::new(tokio::sync::Mutex::new(())), sipflow_config: ArcSwap::new(Arc::new(None)), sip_flow: None, active_call_registry: Arc::new(ActiveProxyCallRegistry::new()), diff --git a/src/proxy/tests/test_registrar.rs b/src/proxy/tests/test_registrar.rs index c0911fca7..93d4c7061 100644 --- a/src/proxy/tests/test_registrar.rs +++ b/src/proxy/tests/test_registrar.rs @@ -1,13 +1,92 @@ use super::common::{ create_register_request, create_test_request, create_test_server, - create_test_server_with_config, create_transaction, + create_test_server_with_config, create_test_server_with_config_and_locator, create_transaction, }; use crate::call::{Location, TransactionCookie}; use crate::config::ProxyConfig; +use crate::proxy::locator::{Locator, LocatorStats, MemoryLocator, RealmChecker}; use crate::proxy::registrar::RegistrarModule; use crate::proxy::{ProxyAction, ProxyModule}; +use anyhow::Result; +use async_trait::async_trait; +use rsipstack::sip::Header; +use rsipstack::transport::SipAddr; +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use tokio::sync::Notify; use tokio_util::sync::CancellationToken; +#[derive(Clone)] +struct PausingLocator { + inner: Arc, + pause_has_active: Arc, + pause_unregister: Arc, + checked: Arc, + resume: Arc, +} + +impl PausingLocator { + fn new() -> Self { + Self { + inner: Arc::new(MemoryLocator::new()), + pause_has_active: Arc::new(AtomicBool::new(false)), + pause_unregister: Arc::new(AtomicBool::new(false)), + checked: Arc::new(Notify::new()), + resume: Arc::new(Notify::new()), + } + } +} + +#[async_trait] +impl Locator for PausingLocator { + fn set_realm_checker(&self, checker: RealmChecker) { + self.inner.set_realm_checker(checker); + } + + async fn register( + &self, + username: &str, + realm: Option<&str>, + location: Location, + ) -> Result<()> { + self.inner.register(username, realm, location).await + } + + async fn has_active_bindings(&self, username: &str, realm: Option<&str>) -> Result { + let active = self.inner.has_active_bindings(username, realm).await?; + if !active && self.pause_has_active.swap(false, Ordering::SeqCst) { + self.checked.notify_one(); + self.resume.notified().await; + } + Ok(active) + } + + async fn unregister(&self, username: &str, realm: Option<&str>) -> Result<()> { + self.inner.unregister(username, realm).await?; + if self.pause_unregister.swap(false, Ordering::SeqCst) { + self.checked.notify_one(); + self.resume.notified().await; + } + Ok(()) + } + + async fn unregister_with_address(&self, addr: &SipAddr) -> Result>> { + self.inner.unregister_with_address(addr).await + } + + async fn lookup(&self, uri: &rsipstack::sip::Uri) -> Result> { + self.inner.lookup(uri).await + } + + async fn sweep_expired(&self) -> Result> { + self.inner.sweep_expired().await + } + + async fn online_stats(&self) -> Result { + self.inner.online_stats().await + } +} + #[tokio::test] async fn test_registrar_register_success() { // Create test server with user backend and locator @@ -116,62 +195,108 @@ async fn test_registrar_unregister() { #[tokio::test] async fn test_registrar_unregister_keeps_user_online_with_another_binding() { - let config = ProxyConfig { - realms: Some(vec!["example.com".to_string()]), - ..Default::default() - }; - let (server_inner, config) = create_test_server_with_config(config).await; - let module = RegistrarModule::new(server_inner.clone(), config); - - let register_request = create_register_request("agent-a", "example.com", Some(60)); - let registered_aor = register_request.uri.clone(); - let registered_realm = registered_aor.host_with_port.to_string(); - let (mut tx, _) = create_transaction(register_request).await; - module - .on_transaction_begin( - CancellationToken::new(), - &mut tx, - TransactionCookie::default(), - ) - .await - .unwrap(); - - let second_aor = create_register_request("new-device", "client.invalid", None).uri; - server_inner - .locator - .register( - "agent-a", - Some(®istered_realm), - Location { - aor: second_aor.clone(), - expires: 60, - registered_aor: Some(registered_aor.clone()), - instance_id: Some("new-binding".to_string()), - ..Default::default() - }, - ) - .await - .unwrap(); - - let mut events = server_inner.locator_events.as_ref().unwrap().subscribe(); - let unregister_request = create_register_request("agent-a", "example.com", Some(0)); - let (mut tx, _) = create_transaction(unregister_request).await; - module - .on_transaction_begin( - CancellationToken::new(), - &mut tx, - TransactionCookie::default(), - ) - .await - .unwrap(); - - let locations = server_inner.locator.lookup(®istered_aor).await.unwrap(); - assert_eq!(locations.len(), 1); - assert_eq!(locations[0].aor, second_aor); - assert!(matches!( - events.try_recv(), - Err(tokio::sync::broadcast::error::TryRecvError::Empty) - )); + let mut stale_event_scenarios = Vec::new(); + for wildcard in [false, true] { + let locator = PausingLocator::new(); + let locator_trait = Arc::new(Box::new(locator.clone()) as Box); + let config = ProxyConfig { + realms: Some(vec!["example.com".to_string()]), + ..Default::default() + }; + let (server_inner, config) = + create_test_server_with_config_and_locator(config, locator_trait).await; + let module = RegistrarModule::new(server_inner.clone(), config); + + let register_request = create_register_request("agent-a", "example.com", Some(60)); + let registered_aor = register_request.uri.clone(); + let (mut tx, _) = create_transaction(register_request).await; + module + .on_transaction_begin( + CancellationToken::new(), + &mut tx, + TransactionCookie::default(), + ) + .await + .unwrap(); + + let mut events = server_inner.locator_events.as_ref().unwrap().subscribe(); + if wildcard { + locator.pause_unregister.store(true, Ordering::SeqCst); + } else { + locator.pause_has_active.store(true, Ordering::SeqCst); + } + + let mut unregister_request = create_register_request("agent-a", "example.com", Some(0)); + if wildcard { + unregister_request + .headers + .retain(|header| !matches!(header, Header::Contact(_))); + unregister_request + .headers + .push(Header::Other("Contact".into(), "*".into())); + } + let (mut unregister_tx, _) = create_transaction(unregister_request).await; + let unregister_module = module.clone(); + let unregister_task = tokio::spawn(async move { + unregister_module + .on_transaction_begin( + CancellationToken::new(), + &mut unregister_tx, + TransactionCookie::default(), + ) + .await + .unwrap(); + }); + + locator.checked.notified().await; + + let register_request = create_register_request("agent-a", "example.com", Some(60)); + let (mut register_tx, _) = create_transaction(register_request).await; + let register_module = module.clone(); + let register_task = tokio::spawn(async move { + register_module + .on_transaction_begin( + CancellationToken::new(), + &mut register_tx, + TransactionCookie::default(), + ) + .await + .unwrap(); + }); + + let first_event = + tokio::time::timeout(std::time::Duration::from_millis(100), events.recv()) + .await + .ok() + .and_then(|result| result.ok()); + locator.resume.notify_one(); + unregister_task.await.unwrap(); + register_task.await.unwrap(); + + let mut observed = first_event.into_iter().collect::>(); + while let Ok(event) = events.try_recv() { + observed.push(event); + } + if !matches!( + observed.last(), + Some(crate::proxy::locator::LocatorEvent::Registered(_)) + ) { + stale_event_scenarios.push((wildcard, observed)); + } + assert_eq!( + server_inner + .locator + .lookup(®istered_aor) + .await + .unwrap() + .len(), + 1 + ); + } + assert!( + stale_event_scenarios.is_empty(), + "stale unregister events followed concurrent registrations: {stale_event_scenarios:?}" + ); } #[tokio::test]