1use lru::LruCache;
11use rsipstack::rsip::prelude::HeadersExt;
12use rsipstack::rsip::{Host, Method, SipMessage};
13use rsipstack::transaction::endpoint::MessageInspector;
14use std::collections::HashMap;
15use std::net::IpAddr;
16use std::num::NonZeroUsize;
17use std::sync::{Arc, Mutex};
18use std::time::{Duration, Instant};
19use tracing::trace;
20
21pub const LEARNED_PEERS_CAPACITY: usize = 1024;
23
24#[derive(Clone)]
27pub struct SharedLearnedPeers {
28 inner: Arc<Mutex<LruCache<IpAddr, Instant>>>,
29}
30
31impl SharedLearnedPeers {
32 pub fn new(capacity: usize) -> Self {
33 Self {
34 inner: Arc::new(Mutex::new(LruCache::new(
35 NonZeroUsize::new(capacity.max(1)).unwrap(),
36 ))),
37 }
38 }
39
40 pub fn learn(&self, ip: IpAddr) {
42 if let Ok(mut cache) = self.inner.lock() {
43 cache.put(ip, Instant::now());
44 }
45 }
46
47 pub fn contains_within(&self, ip: &IpAddr, ttl: Duration) -> bool {
50 let Ok(cache) = self.inner.lock() else {
51 return false;
52 };
53 match cache.peek(ip) {
54 Some(learned_at) => learned_at.elapsed() <= ttl,
55 None => false,
56 }
57 }
58
59 pub fn len(&self) -> usize {
60 self.inner.lock().map(|cache| cache.len()).unwrap_or(0)
61 }
62
63 pub fn is_empty(&self) -> bool {
64 self.len() == 0
65 }
66
67 pub fn snapshot(&self) -> HashMap<IpAddr, Instant> {
68 self.inner
69 .lock()
70 .map(|cache| cache.iter().map(|(ip, at)| (*ip, *at)).collect())
71 .unwrap_or_default()
72 }
73}
74
75pub struct PeerAddressLearner {
77 peers: SharedLearnedPeers,
78 next: Option<Box<dyn MessageInspector>>,
79}
80
81impl PeerAddressLearner {
82 pub fn new(peers: SharedLearnedPeers) -> Self {
83 Self { peers, next: None }
84 }
85
86 pub fn new_with_next(
87 peers: SharedLearnedPeers,
88 next: Option<Box<dyn MessageInspector>>,
89 ) -> Self {
90 Self { peers, next }
91 }
92
93 pub fn shared(&self) -> SharedLearnedPeers {
94 self.peers.clone()
95 }
96
97 fn learn_from(&self, msg: &SipMessage, from: Option<&rsipstack::transport::SipAddr>) {
98 let call_traffic = match msg {
99 SipMessage::Request(req) => req.method == Method::Invite,
100 SipMessage::Response(resp) => {
101 resp.cseq_header().and_then(|cseq| cseq.method()) == Ok(Method::Invite)
102 }
103 };
104 if !call_traffic {
105 return;
106 }
107 let Some(from) = from else {
108 return;
109 };
110 let Some(ip) = peer_ip(&from.addr) else {
111 trace!(peer = %from.addr, "peer address is not an IP literal, not learned");
112 return;
113 };
114 trace!(%ip, "learned peer address from call traffic");
115 self.peers.learn(ip);
116 }
117}
118
119fn peer_ip(host_with_port: &rsipstack::rsip::HostWithPort) -> Option<IpAddr> {
120 match &host_with_port.host {
121 Host::IpAddr(ip) => Some(*ip),
122 Host::Domain(_) => None,
123 }
124}
125
126impl MessageInspector for PeerAddressLearner {
127 fn before_send(
128 &self,
129 msg: SipMessage,
130 dest: Option<&rsipstack::transport::SipAddr>,
131 ) -> SipMessage {
132 match &self.next {
133 Some(next) => next.before_send(msg, dest),
134 None => msg,
135 }
136 }
137
138 fn after_received(
139 &self,
140 msg: SipMessage,
141 from: Option<&rsipstack::transport::SipAddr>,
142 ) -> SipMessage {
143 self.learn_from(&msg, from);
144 match &self.next {
145 Some(next) => next.after_received(msg, from),
146 None => msg,
147 }
148 }
149}
150
151#[cfg(test)]
152mod tests {
153 use super::*;
154 use rsipstack::rsip::{Headers, Header, HostWithPort, Uri, Version, Request, Response};
155 use rsipstack::sip::Transport;
156 use rsipstack::transport::SipAddr;
157
158 fn udp_addr(ip: &str) -> SipAddr {
159 SipAddr {
160 addr: HostWithPort::try_from(format!("{ip}:5060")).unwrap(),
161 r#type: Some(Transport::Udp),
162 }
163 }
164
165 fn invite_request(call_id: &str) -> SipMessage {
166 SipMessage::Request(Request {
167 method: Method::Invite,
168 uri: Uri::try_from("sip:bot@example.com").unwrap(),
169 headers: Headers::from(vec![
170 Header::CallId(call_id.into()),
171 Header::CSeq("1 INVITE".into()),
172 ]),
173 version: Version::V2,
174 body: vec![],
175 })
176 }
177
178 fn options_request() -> SipMessage {
179 SipMessage::Request(Request {
180 method: Method::Options,
181 uri: Uri::try_from("sip:bot@example.com").unwrap(),
182 headers: Headers::from(vec![
183 Header::CallId("probe".into()),
184 Header::CSeq("10 OPTIONS".into()),
185 ]),
186 version: Version::V2,
187 body: vec![],
188 })
189 }
190
191 fn response_with_cseq(cseq: &str) -> SipMessage {
192 SipMessage::Response(Response {
193 headers: Headers::from(vec![Header::CSeq(cseq.into())]),
194 ..Default::default()
195 })
196 }
197
198 #[test]
199 fn learns_only_call_traffic() {
200 let learner = PeerAddressLearner::new(SharedLearnedPeers::new(16));
201
202 learner.learn_from(&invite_request("c1"), Some(&udp_addr("1.1.1.1")));
203 learner.learn_from(&response_with_cseq("1 INVITE"), Some(&udp_addr("2.2.2.2")));
204 learner.learn_from(&options_request(), Some(&udp_addr("3.3.3.3")));
206 learner.learn_from(&response_with_cseq("1 REGISTER"), Some(&udp_addr("4.4.4.4")));
207
208 let ttl = Duration::from_secs(60);
209 assert!(learner.shared().contains_within(&"1.1.1.1".parse().unwrap(), ttl));
210 assert!(learner.shared().contains_within(&"2.2.2.2".parse().unwrap(), ttl));
211 assert!(!learner.shared().contains_within(&"3.3.3.3".parse().unwrap(), ttl));
212 assert!(!learner.shared().contains_within(&"4.4.4.4".parse().unwrap(), ttl));
213 }
214
215 #[test]
216 fn ttl_expiry_and_lru_bound() {
217 let learner = PeerAddressLearner::new(SharedLearnedPeers::new(2));
218
219 learner.learn_from(&invite_request("c1"), Some(&udp_addr("10.0.0.1")));
220 learner.learn_from(&invite_request("c2"), Some(&udp_addr("10.0.0.2")));
221 learner.learn_from(&invite_request("c3"), Some(&udp_addr("10.0.0.3")));
223
224 let ttl = Duration::from_secs(3600);
225 assert_eq!(learner.shared().len(), 2);
226 assert!(!learner
227 .shared()
228 .contains_within(&"10.0.0.1".parse().unwrap(), ttl));
229 assert!(learner
230 .shared()
231 .contains_within(&"10.0.0.3".parse().unwrap(), ttl));
232
233 assert!(!learner
235 .shared()
236 .contains_within(&"10.0.0.3".parse().unwrap(), Duration::ZERO));
237 }
238
239 #[test]
240 fn non_ip_peer_addresses_are_ignored() {
241 let learner = PeerAddressLearner::new(SharedLearnedPeers::new(16));
242 let domain_addr = SipAddr {
243 addr: HostWithPort::try_from("sip.example.com:5060").unwrap(),
244 r#type: Some(Transport::Udp),
245 };
246 learner.learn_from(&invite_request("c1"), Some(&domain_addr));
247 assert!(learner.shared().is_empty());
248 learner.learn_from(&invite_request("c2"), None);
250 assert!(learner.shared().is_empty());
251 }
252}