use std::net::IpAddr;
use std::sync::{Arc, Mutex};
use sipx_sip::{Header, Message, Request};
use tokio::sync::mpsc;
use crate::counters::Meters;
use crate::{ConnectionKey, Target, TransportKind};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct SourcePrefix {
network: IpAddr,
bits: u8,
}
impl SourcePrefix {
#[must_use]
pub fn new(network: IpAddr, bits: u8) -> Option<Self> {
let max = if network.is_ipv4() { 32 } else { 128 };
(bits <= max).then_some(Self { network, bits })
}
#[must_use]
pub const fn address(address: IpAddr) -> Self {
let bits = if address.is_ipv4() { 32 } else { 128 };
Self {
network: address,
bits,
}
}
#[must_use]
pub fn contains(self, candidate: IpAddr) -> bool {
match (self.network, candidate) {
(IpAddr::V4(network), IpAddr::V4(candidate)) => prefix_matches(
u32::from(network).into(),
u32::from(candidate).into(),
self.bits,
32,
),
(IpAddr::V6(network), IpAddr::V6(candidate)) => {
prefix_matches(u128::from(network), u128::from(candidate), self.bits, 128)
}
_ => false,
}
}
}
fn prefix_matches(network: u128, candidate: u128, bits: u8, width: u8) -> bool {
if bits == 0 {
return true;
}
let shift = u32::from(width.saturating_sub(bits));
(network >> shift) == (candidate >> shift)
}
#[derive(Debug, Clone)]
struct AdmissionGeneration {
number: u64,
prefixes: Option<Arc<[SourcePrefix]>>,
}
#[derive(Debug)]
pub(crate) struct SourceAdmission {
current: Mutex<AdmissionGeneration>,
limit: usize,
}
impl Default for SourceAdmission {
fn default() -> Self {
Self::new(1024)
}
}
impl SourceAdmission {
pub(crate) fn new(limit: usize) -> Self {
Self {
current: Mutex::new(AdmissionGeneration {
number: 0,
prefixes: None,
}),
limit,
}
}
pub(crate) fn admit(&self, address: IpAddr) -> Option<u64> {
let current = self
.current
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let allowed = current
.prefixes
.as_ref()
.is_none_or(|prefixes| prefixes.iter().any(|prefix| prefix.contains(address)));
allowed.then_some(current.number)
}
pub(crate) fn replace(&self, prefixes: Vec<SourcePrefix>) -> crate::Result<u64> {
if prefixes.len() > self.limit {
return Err(crate::Error::SourceAdmissionCapacity {
max: self.limit,
attempted: prefixes.len(),
});
}
Ok(self.publish(Some(prefixes.into())))
}
pub(crate) fn clear(&self) -> u64 {
self.publish(None)
}
fn publish(&self, prefixes: Option<Arc<[SourcePrefix]>>) -> u64 {
let mut current = self
.current
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
current.number = current.number.wrapping_add(1).max(1);
current.prefixes = prefixes;
current.number
}
}
#[derive(Debug)]
pub enum RequestPolicyDecision {
Allow,
Reject(String),
AddHeaders(Vec<Header>),
}
pub trait RequestPolicy: Send + Sync {
fn decide(&self, request: &Request, target: &Target) -> RequestPolicyDecision;
}
#[derive(Clone)]
pub struct RequestPolicyRef(Arc<dyn RequestPolicy>);
impl RequestPolicyRef {
#[must_use]
pub fn new(policy: impl RequestPolicy + 'static) -> Self {
Self(Arc::new(policy))
}
pub(crate) fn decide(&self, request: &Request, target: &Target) -> RequestPolicyDecision {
self.0.decide(request, target)
}
}
impl std::fmt::Debug for RequestPolicyRef {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("RequestPolicyRef(..)")
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MessageDirection {
Inbound,
Outbound,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TransactionClass {
ServerCreated,
Matched,
Unmatched,
ClientCreated,
Direct,
}
#[derive(Debug, Clone)]
pub struct MessageObservation {
pub message: Message,
pub local: std::net::SocketAddr,
pub peer: std::net::SocketAddr,
pub transport: TransportKind,
pub direction: MessageDirection,
pub transaction: TransactionClass,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ConnectionId {
pub key: ConnectionKey,
pub generation: u64,
pub admission_generation: Option<u64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConnectionState {
Accepted,
Opened,
Authenticated,
Pooled,
Reused,
Failed,
Closed,
}
#[derive(Debug, Clone)]
pub struct ConnectionObservation {
pub connection: ConnectionId,
pub state: ConnectionState,
}
#[derive(Debug, Clone)]
pub enum EndpointObservation {
Message(Box<MessageObservation>),
Connection(ConnectionObservation),
}
#[derive(Debug)]
pub(crate) struct ObservationHub {
sink: Mutex<Option<mpsc::Sender<EndpointObservation>>>,
meters: Arc<Meters>,
}
impl ObservationHub {
pub(crate) fn new(meters: Arc<Meters>) -> Self {
Self {
sink: Mutex::new(None),
meters,
}
}
pub(crate) fn subscribe(&self, capacity: usize) -> mpsc::Receiver<EndpointObservation> {
let (sender, receiver) = mpsc::channel(capacity.max(1));
*self
.sink
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = Some(sender);
receiver
}
pub(crate) fn emit(&self, event: EndpointObservation) {
let mut sink = self
.sink
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let Some(sender) = sink.as_ref() else {
return;
};
match sender.try_send(event) {
Ok(()) => {}
Err(mpsc::error::TrySendError::Full(_)) => self.meters.observation_drop(),
Err(mpsc::error::TrySendError::Closed(_)) => *sink = None,
}
}
}
pub(crate) fn connection_event(
key: ConnectionKey,
generation: u64,
admission_generation: Option<u64>,
state: ConnectionState,
) -> EndpointObservation {
EndpointObservation::Connection(ConnectionObservation {
connection: ConnectionId {
key,
generation,
admission_generation,
},
state,
})
}
pub(crate) fn policy_header(name: &sipx_sip::HeaderName) -> (sipx_sip::HeaderName, bool) {
use sipx_sip::HeaderName;
let semantic = HeaderName::parse(&bytes::Bytes::copy_from_slice(name.canonical()));
let allowed = matches!(
semantic,
HeaderName::AlertInfo
| HeaderName::CallInfo
| HeaderName::Organization
| HeaderName::Priority
| HeaderName::Subject
| HeaderName::UserAgent
| HeaderName::Other(_)
);
(semantic, allowed)
}
pub(crate) fn duplicate_policy_header(request: &Request, semantic: &sipx_sip::HeaderName) -> bool {
!matches!(semantic, sipx_sip::HeaderName::Other(_)) && request.headers.get(semantic).is_some()
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
clippy::indexing_slicing
)]
mod tests {
use super::*;
use bytes::Bytes;
use sipx_sip::HeaderName;
use std::net::{Ipv4Addr, Ipv6Addr};
#[test]
fn prefixes_match_only_their_network_and_family() {
let v4 = SourcePrefix::new(Ipv4Addr::new(192, 0, 2, 0).into(), 24).unwrap();
assert!(v4.contains(Ipv4Addr::new(192, 0, 2, 42).into()));
assert!(!v4.contains(Ipv4Addr::new(192, 0, 3, 1).into()));
assert!(!v4.contains(Ipv6Addr::LOCALHOST.into()));
}
#[test]
fn replacement_publishes_complete_generations() {
let admission = SourceAdmission::default();
assert_eq!(admission.admit(IpAddr::V4(Ipv4Addr::LOCALHOST)), Some(0));
let one = admission
.replace(vec![SourcePrefix::address(IpAddr::V4(Ipv4Addr::LOCALHOST))])
.unwrap();
assert_eq!(admission.admit(IpAddr::V4(Ipv4Addr::LOCALHOST)), Some(one));
assert_eq!(admission.admit(IpAddr::V4(Ipv4Addr::UNSPECIFIED)), None);
let two = admission.clear();
assert!(two > one);
assert_eq!(
admission.admit(IpAddr::V4(Ipv4Addr::UNSPECIFIED)),
Some(two)
);
}
#[test]
fn oversized_replacement_preserves_the_old_generation() {
let admission = SourceAdmission::new(1);
let first = admission
.replace(vec![SourcePrefix::address(IpAddr::V4(Ipv4Addr::LOCALHOST))])
.unwrap();
let error = admission
.replace(vec![
SourcePrefix::address(IpAddr::V4(Ipv4Addr::LOCALHOST)),
SourcePrefix::address(IpAddr::V4(Ipv4Addr::UNSPECIFIED)),
])
.unwrap_err();
assert!(matches!(
error,
crate::Error::SourceAdmissionCapacity { .. }
));
assert_eq!(
admission.admit(IpAddr::V4(Ipv4Addr::LOCALHOST)),
Some(first)
);
assert_eq!(admission.admit(IpAddr::V4(Ipv4Addr::UNSPECIFIED)), None);
}
#[test]
fn request_policy_allows_only_application_fields_and_unknown_extensions() {
for name in [HeaderName::Subject, HeaderName::Organization] {
assert!(policy_header(&name).1);
}
assert!(policy_header(&HeaderName::Other(Bytes::from_static(b"X-Trace"))).1);
for name in [
HeaderName::Contact,
HeaderName::ContentType,
HeaderName::Event,
] {
assert!(!policy_header(&name).1);
}
for raw in [b"vIa".as_slice(), b"v".as_slice()] {
let (semantic, allowed) =
policy_header(&HeaderName::Other(Bytes::copy_from_slice(raw)));
assert_eq!(semantic, HeaderName::Via);
assert!(!allowed);
}
}
}