use std::collections::HashMap;
use std::io;
use std::net::SocketAddr;
use crate::Community;
use bytes::Bytes;
use socket2::{Domain, Protocol, Socket, Type};
use tokio::net::UdpSocket;
use crate::error::{Error, Result};
use crate::v3::{AuthoritativeEngine, UsmUser};
#[derive(Clone, Copy)]
pub(crate) enum EmptyCommunityPolicy {
Allow,
#[cfg(feature = "agent")]
Deny,
}
pub(crate) fn community_matches(
configured: &[Community],
community: &[u8],
empty_policy: EmptyCommunityPolicy,
) -> bool {
if configured.is_empty() {
return matches!(empty_policy, EmptyCommunityPolicy::Allow);
}
let mut valid = false;
for candidate in configured {
if candidate.matches(community) {
valid = true;
}
}
valid
}
pub(crate) struct PreparedAuthoritativeUsm {
pub(crate) users: HashMap<Bytes, UsmUser>,
pub(crate) authoritative_engine: Option<AuthoritativeEngine>,
pub(crate) engine_id: Bytes,
pub(crate) engine_boots: u32,
}
pub(crate) struct ValidatedAuthoritativeUsm {
users: HashMap<Bytes, UsmUser>,
authoritative_engine: Option<AuthoritativeEngine>,
configured_engine: Option<(Bytes, u32)>,
}
impl ValidatedAuthoritativeUsm {
pub(crate) fn prepare(
self,
generate_engine_id: impl FnOnce() -> Result<Bytes>,
) -> Result<PreparedAuthoritativeUsm> {
let (engine_id, engine_boots) = match self.configured_engine {
Some(configured) => configured,
None => match &self.authoritative_engine {
Some(engine) => {
let (engine_boots, _) = engine.current_boots_time()?;
(Bytes::copy_from_slice(engine.engine_id()), engine_boots)
}
None => (generate_engine_id()?, 1),
},
};
Ok(PreparedAuthoritativeUsm {
users: self.users,
authoritative_engine: self.authoritative_engine,
engine_id,
engine_boots,
})
}
}
pub(crate) fn validate_authoritative_usm(
users: HashMap<Bytes, UsmUser>,
authoritative_engine: Option<AuthoritativeEngine>,
requires_engine: bool,
invalid_user_context: &str,
missing_engine_context: &str,
) -> Result<ValidatedAuthoritativeUsm> {
validate_authoritative_usm_with(
users,
authoritative_engine,
requires_engine,
invalid_user_context,
missing_engine_context,
false,
)
}
#[cfg(feature = "agent")]
pub(crate) fn validate_authoritative_usm_deferred(
users: HashMap<Bytes, UsmUser>,
authoritative_engine: Option<AuthoritativeEngine>,
requires_engine: bool,
invalid_user_context: &str,
missing_engine_context: &str,
) -> Result<ValidatedAuthoritativeUsm> {
validate_authoritative_usm_with(
users,
authoritative_engine,
requires_engine,
invalid_user_context,
missing_engine_context,
true,
)
}
fn validate_authoritative_usm_with(
mut users: HashMap<Bytes, UsmUser>,
authoritative_engine: Option<AuthoritativeEngine>,
requires_engine: bool,
invalid_user_context: &str,
missing_engine_context: &str,
defer_engine_time: bool,
) -> Result<ValidatedAuthoritativeUsm> {
let mut validation_order: Vec<_> = users
.iter()
.map(|(key, config)| (config.username().clone(), key.clone()))
.collect();
validation_order.sort_unstable();
for (_, key) in validation_order {
let config = users
.get_mut(&key)
.expect("key was collected from the same map");
config.validate_and_precompute().map_err(|error| {
Error::Config(format!("{invalid_user_context}: {error}").into()).boxed()
})?;
}
let (authoritative_engine, configured_engine) = match authoritative_engine {
Some(engine) if defer_engine_time => (Some(engine), None),
Some(engine) => {
let (engine_boots, _) = engine.current_boots_time()?;
let engine_id = Bytes::copy_from_slice(engine.engine_id());
(Some(engine), Some((engine_id, engine_boots)))
}
None if requires_engine => {
return Err(Error::Config(missing_engine_context.into()).boxed());
}
None => (None, None),
};
Ok(ValidatedAuthoritativeUsm {
users,
authoritative_engine,
configured_engine,
})
}
pub(crate) async fn bind_udp_socket(
addr: SocketAddr,
recv_buffer_size: Option<usize>,
send_buffer_size: Option<usize>,
reuse_address: bool,
) -> io::Result<UdpSocket> {
let domain = if addr.is_ipv6() {
Domain::IPV6
} else {
Domain::IPV4
};
let socket = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP))?;
if addr.is_ipv6() {
socket.set_only_v6(false)?;
}
if reuse_address {
socket.set_reuse_address(true)?;
}
if let Some(size) = recv_buffer_size {
let _ = socket.set_recv_buffer_size(size);
}
if let Some(size) = send_buffer_size {
let _ = socket.set_send_buffer_size(size);
}
socket.set_nonblocking(true)?;
socket.bind(&addr.into())?;
UdpSocket::from_std(socket.into())
}
#[cfg(test)]
mod usm_validation_tests {
use super::*;
fn invalid_user_error_with_keys(users: impl IntoIterator<Item = (Bytes, UsmUser)>) -> String {
validate_authoritative_usm(
users.into_iter().collect(),
None,
false,
"invalid user",
"missing engine",
)
.err()
.expect("at least one user must be invalid")
.to_string()
}
fn invalid_user_error(users: impl IntoIterator<Item = Bytes>) -> String {
invalid_user_error_with_keys(
users
.into_iter()
.map(|username| (username.clone(), UsmUser::new(username))),
)
}
#[test]
fn authoritative_usm_validation_reports_first_invalid_username_by_octets() {
let lexically_first = Bytes::from(vec![0x01; 34]);
let lexically_second = Bytes::from(vec![0x02; 33]);
for users in [
vec![lexically_first.clone(), lexically_second.clone()],
vec![lexically_second.clone(), lexically_first.clone()],
] {
let error = invalid_user_error(users);
assert!(error.contains("got 34"), "unexpected error: {error}");
}
for users in [
vec![
(
Bytes::from_static(b"a-map-key"),
UsmUser::new(lexically_second.clone()),
),
(
Bytes::from_static(b"z-map-key"),
UsmUser::new(lexically_first.clone()),
),
],
vec![
(
Bytes::from_static(b"z-map-key"),
UsmUser::new(lexically_first.clone()),
),
(
Bytes::from_static(b"a-map-key"),
UsmUser::new(lexically_second.clone()),
),
],
] {
let error = invalid_user_error_with_keys(users);
assert!(error.contains("got 34"), "unexpected error: {error}");
}
}
#[test]
fn authoritative_usm_validation_order_is_stable_with_empty_and_overlong_users() {
let empty = Bytes::new();
let overlong = Bytes::from(vec![b'z'; 33]);
for users in [
vec![empty.clone(), overlong.clone()],
vec![overlong.clone(), empty.clone()],
] {
let error = invalid_user_error(users);
assert!(error.contains("got 0"), "unexpected error: {error}");
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_bind_udp_socket_ipv4() {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let socket = bind_udp_socket(addr, None, None, true).await.unwrap();
let local = socket.local_addr().unwrap();
assert!(local.is_ipv4());
assert_ne!(local.port(), 0);
}
#[tokio::test]
async fn test_bind_udp_socket_ipv6() {
let addr: SocketAddr = "[::1]:0".parse().unwrap();
let socket = bind_udp_socket(addr, None, None, true).await.unwrap();
let local = socket.local_addr().unwrap();
assert!(local.is_ipv6());
assert_ne!(local.port(), 0);
}
#[tokio::test]
async fn test_bind_udp_socket_with_buffer_size() {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let socket = bind_udp_socket(addr, Some(1024 * 1024), None, true)
.await
.unwrap();
let local = socket.local_addr().unwrap();
assert!(local.is_ipv4());
assert_ne!(local.port(), 0);
}
}