use std::sync::Arc;
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::sync::mpsc::{unbounded_channel, UnboundedReceiver, UnboundedSender};
use tokio_util::sync::CancellationToken;
use crate::dialog::{
dialog::{DialogInner, DialogState, TerminatedReason},
dialog_layer::DialogLayer,
invite_dialog::InviteDialog,
DialogId,
};
use crate::sip::{
headers::{CSeq, CallId, Contact, From, MaxForwards, Route, To, Via},
prelude::{HeadersExt, ToTypedHeader},
Host, HostWithPort, Method, Request, Response, SipMessage, StatusCode, Transport, Version,
};
use crate::transaction::{
endpoint::EndpointBuilder,
key::{TransactionKey, TransactionRole},
transaction::Transaction,
};
use crate::transport::{
channel::ChannelConnection,
udp::{UdpConnection, UdpInner},
SipAddr, SipConnection, TransportEvent, TransportLayer,
};
fn create_wss_invite_request(from_tag: &str, to_tag: &str, call_id: &str, branch: &str) -> Request {
Request {
method: Method::Invite,
uri: crate::sip::Uri::try_from("sip:+12761187472182@58.246.19.74:6988").unwrap(),
headers: vec![
Via::new(format!("SIP/2.0/WSS vfa095afrt10.invalid;branch={branch}").as_str()).into(),
CSeq::new("2421 INVITE").into(),
From::new(format!("<sip:agent@pbx.example.com>;tag={from_tag}").as_str()).into(),
To::new(format!("<sip:+12761187472182@58.246.19.74:6988>;tag={to_tag}").as_str())
.into(),
CallId::new(call_id).into(),
Contact::new("<sip:skpn0n1b@vfa095afrt10.invalid;transport=ws;ob>").into(),
MaxForwards::new("69").into(),
]
.into(),
version: Version::V2,
body: b"v=0\r\no=- 449672722 449672723 IN IP4 127.0.0.1\r\ns=-\r\nt=0 0\r\n".to_vec(),
}
}
fn plain_udp_invite(from_tag: &str, call_id: &str, via: &str) -> Request {
Request {
method: Method::Invite,
uri: crate::sip::Uri::try_from("sip:bob@example.com:5060").unwrap(),
headers: vec![
Via::new(via).into(),
CSeq::new("1 INVITE").into(),
From::new(format!("Alice <sip:alice@example.com>;tag={from_tag}").as_str()).into(),
To::new("Bob <sip:bob@example.com>").into(),
CallId::new(call_id).into(),
Contact::new("<sip:alice@alice.invalid:5060>").into(),
MaxForwards::new("70").into(),
]
.into(),
version: Version::V2,
body: vec![],
}
}
fn server_dialog(
endpoint_inner: crate::transaction::endpoint::EndpointInnerRef,
role: TransactionRole,
initial: Request,
local_contact: &str,
) -> DialogInner {
let id = DialogId::try_from((&initial, role)).unwrap();
let (tu_tx, _tu_rx) = unbounded_channel();
let (state_tx, _state_rx) = unbounded_channel();
let mut id_local = id;
if role == TransactionRole::Server && id_local.local_tag.is_empty() {
id_local.local_tag = "local-tag".into();
}
DialogInner::new(
role,
id_local,
initial,
endpoint_inner,
state_tx,
None,
Some(crate::sip::Uri::try_from(local_contact).unwrap()),
tu_tx,
)
.unwrap()
}
async fn create_endpoint_without_transports(
) -> crate::Result<crate::transaction::endpoint::Endpoint> {
create_endpoint_with_inspector(None).await
}
async fn create_endpoint_with_inspector(
inspector: Option<Box<dyn crate::transaction::endpoint::MessageInspector>>,
) -> crate::Result<crate::transaction::endpoint::Endpoint> {
let token = CancellationToken::new();
let tl = TransportLayer::new(token.child_token());
let local = tokio::net::UdpSocket::bind("127.0.0.1:0").await?;
let local_addr = {
let sock = local.local_addr()?;
SipAddr {
r#type: Some(Transport::Udp),
addr: crate::sip::HostWithPort {
host: Host::IpAddr(sock.ip()),
port: Some(sock.port().into()),
},
}
};
let udp_conn = UdpConnection::attach(
UdpInner {
conn: local,
addr: local_addr,
},
None,
Some(token.child_token()),
)
.await;
tl.inner.add_listener(udp_conn.into());
let mut builder = EndpointBuilder::new();
builder
.with_user_agent("rsipstack-affinity-test")
.with_transport_layer(tl);
if let Some(inspector) = inspector {
builder.with_inspector(inspector);
}
Ok(builder.build())
}
fn wss_flow_addr(port: u16) -> SipAddr {
SipAddr {
r#type: Some(Transport::Wss),
addr: HostWithPort {
host: Host::IpAddr(std::net::IpAddr::V4(std::net::Ipv4Addr::LOCALHOST)),
port: Some(port.into()),
},
}
}
struct WssFlow {
sip_conn: SipConnection,
wire_rx: UnboundedReceiver<TransportEvent>,
in_tx: UnboundedSender<TransportEvent>,
addr: SipAddr,
cancel: CancellationToken,
}
async fn create_wss_flow(
endpoint_inner: &crate::transaction::endpoint::EndpointInnerRef,
addr: SipAddr,
) -> WssFlow {
let cancel = CancellationToken::new();
let (in_tx, in_rx) = unbounded_channel::<TransportEvent>();
let (out_tx, out_rx) = unbounded_channel();
let conn =
ChannelConnection::create_connection(in_rx, out_tx, addr.clone(), Some(cancel.clone()))
.await
.expect("channel flow connection");
let sip_conn = SipConnection::Channel(conn);
endpoint_inner
.transport_layer
.add_connection(sip_conn.clone());
WssFlow {
sip_conn,
wire_rx: out_rx,
in_tx,
addr,
cancel,
}
}
#[tokio::test]
async fn test_server_dialog_bye_is_delivered_over_initial_connection() {
let endpoint = create_endpoint_without_transports().await.unwrap();
let endpoint_inner = endpoint.inner.clone();
let dialog_layer = DialogLayer::new(endpoint_inner.clone());
tokio::spawn(async move {
let _ = endpoint_inner.serve().await;
});
let mut flow = create_wss_flow(&endpoint.inner, wss_flow_addr(38_013)).await;
let invite =
create_wss_invite_request("h45rr596ko", "", "affinity-bye-callid", "z9hG4bK8544229");
let key = TransactionKey::from_request(&invite, TransactionRole::Server).unwrap();
let tx = Transaction::new_server(
key,
invite.clone(),
endpoint.inner.clone(),
Some(flow.sip_conn.clone()),
);
let (state_tx, mut state_rx) = unbounded_channel();
let dialog = dialog_layer
.get_or_create_server_invite(
&tx,
state_tx,
None,
Some(crate::sip::Uri::try_from("sip:pbx@pbx.example.com:5060").unwrap()),
)
.expect("dialog created");
assert!(
dialog.inner.server_connection.lock().is_some(),
"server dialog must remember the incoming connection"
);
assert!(dialog.inner.test_resolve_affinity_connection().is_some());
let id = dialog.id();
dialog
.inner
.transition(DialogState::Confirmed(id.clone(), Response::default()))
.expect("confirm dialog");
assert!(dialog.inner.is_confirmed());
let bye_dialog = dialog.clone();
let bye_handle = tokio::spawn(async move {
bye_dialog
.bye_with_headers(None)
.await
.expect("bye_with_headers must succeed")
});
let event = tokio::time::timeout(Duration::from_secs(3), flow.wire_rx.recv())
.await
.expect("BYE must be forwarded over the initial connection")
.expect("connection sender alive");
let bye_req = match event {
TransportEvent::Incoming(SipMessage::Request(req), _transport, _src) => {
assert_eq!(req.method, Method::Bye);
assert_eq!(
req.uri.host_with_port.host.to_string(),
"vfa095afrt10.invalid",
"remote target must remain the recorded browser Contact"
);
assert_eq!(
req.cseq_header()
.expect("cseq")
.typed()
.expect("typed cseq")
.method,
Method::Bye
);
let via = req.via_header().expect("via").typed().expect("typed via");
assert!(via.branch().is_some(), "outgoing BYE must carry a branch");
req
}
other => panic!("expected outgoing BYE request, got {other:?}"),
};
let ok_response = Response {
status_code: StatusCode::OK,
version: Version::V2,
headers: vec![
bye_req.via_header().unwrap().clone().into(),
bye_req.from_header().unwrap().clone().into(),
bye_req.to_header().unwrap().clone().into(),
bye_req.call_id_header().unwrap().clone().into(),
bye_req.cseq_header().unwrap().clone().into(),
crate::sip::headers::ContentLength::from(0u32).into(),
]
.into(),
body: vec![],
};
flow.in_tx
.send(TransportEvent::Incoming(
SipMessage::Response(ok_response),
flow.sip_conn.clone(),
flow.addr.clone(),
))
.expect("inbound injection");
let deadline = Duration::from_secs(2);
let mut terminated_reason = None;
while let Ok(Some(state)) = tokio::time::timeout(deadline, state_rx.recv()).await {
if let DialogState::Terminated(_, reason) = state {
terminated_reason = Some(reason);
break;
}
}
let reason = terminated_reason.expect("dialog must terminate after 200 OK for BYE");
assert!(matches!(reason, TerminatedReason::UasBye));
tokio::time::timeout(Duration::from_secs(1), bye_handle)
.await
.expect("bye task to finish")
.ok();
}
#[tokio::test]
async fn test_client_dialogs_use_recorded_flow_affinity() {
let endpoint = create_endpoint_without_transports().await.unwrap();
let initial = plain_udp_invite(
"alice-tag",
"client-flow-affinity",
"SIP/2.0/UDP alice.example.com:5060;branch=z9hG4bKclient1",
);
let inner = Arc::new(server_dialog(
endpoint.inner.clone(),
TransactionRole::Client,
initial,
"sip:uac@example.com",
));
let flow = create_wss_flow(&endpoint.inner, wss_flow_addr(38_014)).await;
inner.set_server_connection(Some(flow.sip_conn));
assert!(
inner.test_resolve_affinity_connection().is_some(),
"client dialogs must reuse the recorded reliable flow"
);
}
#[tokio::test]
async fn test_affinity_skips_unreliable_udp_connections() {
let endpoint = create_endpoint_without_transports().await.unwrap();
let initial = plain_udp_invite(
"alice-tag",
"udp-affinity-check",
"SIP/2.0/UDP alice.example.com:5060;branch=z9hG4bKudp1",
);
let inner = Arc::new(server_dialog(
endpoint.inner.clone(),
TransactionRole::Server,
initial,
"sip:bob@example.com",
));
let udp_conn = UdpConnection::create_connection("127.0.0.1:0".parse().unwrap(), None, None)
.await
.unwrap();
inner.set_server_connection(Some(SipConnection::Udp(udp_conn)));
assert!(
inner.test_resolve_affinity_connection().is_none(),
"UDP connections must not be used for affinity routing"
);
}
#[tokio::test]
async fn test_affinity_requires_recorded_reliable_server_connection() {
let endpoint = create_endpoint_without_transports().await.unwrap();
let initial = create_wss_invite_request("t1", "", "guards-callid", "z9hG4bKguard");
let inner = Arc::new(server_dialog(
endpoint.inner.clone(),
TransactionRole::Server,
initial,
"sip:bob@example.com",
));
assert!(inner.test_resolve_affinity_connection().is_none());
let flow = create_wss_flow(&endpoint.inner, wss_flow_addr(38_015)).await;
inner.set_server_connection(Some(flow.sip_conn));
assert!(inner.test_resolve_affinity_connection().is_some());
inner
.route_set
.lock()
.push(Route::from("sip:proxy.example.com;lr"));
assert!(
inner.test_resolve_affinity_connection().is_none(),
"route set must take precedence over raw flow affinity"
);
}
#[tokio::test]
async fn test_restored_dialog_falls_back_to_initial_via_dialback() {
let token = CancellationToken::new();
let tl = TransportLayer::new(token.child_token());
let probe = UdpSocket::bind("127.0.0.1:0").await.expect("bind probe");
let probe_port = probe.local_addr().unwrap().port();
let dummy = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let dummy_sock = dummy.local_addr().unwrap();
let udp_conn = UdpConnection::attach(
UdpInner {
conn: dummy,
addr: SipAddr {
r#type: Some(Transport::Udp),
addr: crate::sip::HostWithPort {
host: Host::IpAddr(dummy_sock.ip()),
port: Some(dummy_sock.port().into()),
},
},
},
None,
Some(token.child_token()),
)
.await;
tl.inner.add_listener(udp_conn.into());
let endpoint = EndpointBuilder::new()
.with_user_agent("rsipstack-fallback-test")
.with_transport_layer(tl)
.build();
let via = format!(
"SIP/2.0/UDP caller.invalid:5060;branch=z9hG4bKfallback1;received=127.0.0.1;rport={probe_port}"
);
let invite = plain_udp_invite("restored-tag", "fallback-dialback", &via);
let inner = Arc::new(server_dialog(
endpoint.inner.clone(),
TransactionRole::Server,
invite,
"sip:bob@bob.example.com:5060",
));
assert!(inner.test_resolve_affinity_connection().is_none());
let id = inner.id.lock().clone();
inner
.transition(DialogState::Confirmed(id.clone(), Response::default()))
.unwrap();
assert!(inner.is_confirmed());
let _bye_handle = tokio::spawn(async move {
let dialog = InviteDialog::from_inner(inner);
if let Err(e) = dialog.bye_with_headers(None).await {
std::eprintln!("[DBG] fallback bye error: {e}");
}
std::eprintln!("[DBG] fallback bye finished");
});
let deadline = Duration::from_secs(3);
let mut buf = [0u8; 2048];
loop {
let (len, _peer) = tokio::time::timeout(deadline, probe.recv_from(&mut buf))
.await
.expect("BYE dial-back datagram must reach the Via-received address")
.expect("recv ok");
let text = String::from_utf8_lossy(&buf[..len]);
if text.starts_with("BYE ") {
break;
}
}
}
#[tokio::test]
async fn test_bye_survives_cancelled_flow_without_received_via_param() {
let endpoint = create_endpoint_without_transports().await.unwrap();
let endpoint_inner = endpoint.inner.clone();
let dialog_layer = DialogLayer::new(endpoint_inner.clone());
tokio::spawn(async move {
let _ = endpoint_inner.serve().await;
});
let invite =
create_wss_invite_request("nocall123", "", "cancelled-flow-callid", "z9hG4bKnorecv1");
assert!(invite.via_header().unwrap().to_string().len() > 0);
let mut flow = create_wss_flow(&endpoint.inner, wss_flow_addr(38_016)).await;
let key = TransactionKey::from_request(&invite, TransactionRole::Server).unwrap();
let tx = Transaction::new_server(
key,
invite.clone(),
endpoint.inner.clone(),
Some(flow.sip_conn.clone()),
);
let (state_tx, _state_rx) = unbounded_channel();
let dialog = dialog_layer
.get_or_create_server_invite(
&tx,
state_tx,
None,
Some(crate::sip::Uri::try_from("sip:pbx@pbx.example.com:5060").unwrap()),
)
.expect("dialog created");
assert!(dialog.inner.server_connection.lock().is_some());
assert_eq!(
dialog.inner.dialback_target.lock().as_ref(),
Some(&flow.addr),
"dial-back target must be captured structurally at creation"
);
let id = dialog.id();
dialog
.inner
.transition(DialogState::Confirmed(id.clone(), Response::default()))
.expect("confirm dialog");
flow.cancel.cancel();
assert!(
dialog.inner.test_resolve_affinity_connection().is_none(),
"a cancelled flow must be rejected by the affinity filter"
);
assert!(
dialog.inner.dialback_target.lock().is_some(),
"dial-back target must outlive the connection"
);
let bye_dialog = dialog.clone();
tokio::spawn(async move {
let _ = bye_dialog.bye_with_headers(None).await;
});
let event = tokio::time::timeout(Duration::from_secs(3), flow.wire_rx.recv())
.await
.expect("BYE must reach the flow via tier-1 dial-back despite no Via params")
.expect("connection sender alive");
match event {
TransportEvent::Incoming(SipMessage::Request(req), _transport, _src) => {
assert_eq!(req.method, Method::Bye);
}
other => panic!("expected outgoing BYE request, got {other:?}"),
}
}
#[derive(Default, Clone)]
struct SentDests(Arc<std::sync::Mutex<Vec<(bool, Option<SipAddr>)>>>);
struct RecordingInspector(SentDests);
impl crate::transaction::endpoint::MessageInspector for RecordingInspector {
fn before_send(&self, msg: SipMessage, dest: Option<&SipAddr>) -> SipMessage {
self.0
.0
.lock()
.unwrap()
.push((msg.is_request(), dest.cloned()));
msg
}
fn after_received(&self, msg: SipMessage, _from: Option<&SipAddr>) -> SipMessage {
msg
}
}
#[tokio::test]
async fn test_bye_before_send_reports_flow_destination() {
let dests = SentDests::default();
let endpoint =
create_endpoint_with_inspector(Some(Box::new(RecordingInspector(dests.clone()))))
.await
.unwrap();
let endpoint_inner = endpoint.inner.clone();
let dialog_layer = DialogLayer::new(endpoint_inner.clone());
tokio::spawn(async move {
let _ = endpoint_inner.serve().await;
});
let mut flow = create_wss_flow(&endpoint.inner, wss_flow_addr(38_017)).await;
let invite = create_wss_invite_request(
"insp1234ko",
"",
"inspector-dest-callid",
"z9hG4bKinspdest1",
);
let key = TransactionKey::from_request(&invite, TransactionRole::Server).unwrap();
let tx = Transaction::new_server(
key,
invite.clone(),
endpoint.inner.clone(),
Some(flow.sip_conn.clone()),
);
let (state_tx, _state_rx) = unbounded_channel();
let dialog = dialog_layer
.get_or_create_server_invite(
&tx,
state_tx,
None,
Some(crate::sip::Uri::try_from("sip:pbx@pbx.example.com:5060").unwrap()),
)
.expect("dialog created");
let id = dialog.id();
dialog
.inner
.transition(DialogState::Confirmed(id.clone(), Response::default()))
.expect("confirm dialog");
let bye_dialog = dialog.clone();
tokio::spawn(async move {
let _ = bye_dialog.bye_with_headers(None).await;
});
let event = tokio::time::timeout(Duration::from_secs(3), flow.wire_rx.recv())
.await
.expect("BYE must be forwarded over the initial connection")
.expect("connection sender alive");
assert!(matches!(
event,
TransportEvent::Incoming(SipMessage::Request(ref req), _, _) if req.method == Method::Bye
));
let sent = dests.0.lock().unwrap().clone();
let bye_dest = sent
.iter()
.rev()
.find(|(is_request, _)| *is_request)
.map(|(_, dest)| dest.clone())
.flatten();
assert_eq!(
bye_dest,
Some(flow.addr),
"before_send must carry the affinity flow's real remote address"
);
}