use std::net::SocketAddr;
use bytes::Bytes;
use std::sync::Arc;
use std::sync::atomic::{AtomicU8, Ordering};
use subtle::ConstantTimeEq;
use tokio_util::sync::CancellationToken;
use crate::Community;
use crate::DecodeAnomaly;
use crate::message::SecurityLevel;
use crate::pdu::PduType;
use crate::version::{CommunityVersion, Version};
use super::SecurityModel;
#[derive(Debug, Clone)]
pub enum SecurityName {
Community(Community),
Usm(Bytes),
}
impl SecurityName {
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
match self {
Self::Community(community) => community.as_bytes(),
Self::Usm(username) => username,
}
}
#[must_use]
pub fn len(&self) -> usize {
self.as_bytes().len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.as_bytes().is_empty()
}
#[must_use]
pub fn matches(&self, candidate: &Self) -> bool {
match (self, candidate) {
(Self::Community(expected), Self::Community(actual)) => {
expected.matches(actual.as_bytes())
}
(Self::Usm(expected), Self::Usm(actual)) => {
expected.len() == actual.len()
&& bool::from(expected.as_ref().ct_eq(actual.as_ref()))
}
(Self::Community(_), Self::Usm(_)) | (Self::Usm(_), Self::Community(_)) => false,
}
}
}
#[derive(Debug, Clone)]
pub struct RequestContext {
received_at: tokio::time::Instant,
admitted_at: tokio::time::Instant,
deadline: Option<tokio::time::Instant>,
cancellation: CancellationToken,
phase: Arc<AtomicU8>,
source: SocketAddr,
version: Version,
security_model: SecurityModel,
security_name: SecurityName,
security_level: SecurityLevel,
context_name: Bytes,
request_id: i32,
pdu_type: PduType,
group_name: Option<Bytes>,
read_view: Option<Bytes>,
write_view: Option<Bytes>,
msg_max_size: Option<usize>,
decode_anomalies: Vec<DecodeAnomaly>,
}
impl RequestContext {
#[cfg(test)]
pub(crate) fn community(
source: SocketAddr,
version: CommunityVersion,
community: Community,
request_id: i32,
pdu_type: PduType,
decode_anomalies: Vec<DecodeAnomaly>,
) -> Self {
Self::community_with_lifecycle(
source,
version,
community,
request_id,
pdu_type,
decode_anomalies,
RequestLifecycle::for_test(pdu_type),
)
}
pub(crate) fn community_with_lifecycle(
source: SocketAddr,
version: CommunityVersion,
community: Community,
request_id: i32,
pdu_type: PduType,
decode_anomalies: Vec<DecodeAnomaly>,
lifecycle: RequestLifecycle,
) -> Self {
let (version, security_model) = match version {
CommunityVersion::V1 => (Version::V1, SecurityModel::V1),
CommunityVersion::V2c => (Version::V2c, SecurityModel::V2c),
};
lifecycle.classify(pdu_type);
Self {
received_at: lifecycle.received_at,
admitted_at: lifecycle.admitted_at,
deadline: lifecycle.deadline,
cancellation: lifecycle.cancellation,
phase: lifecycle.phase,
source,
version,
security_model,
security_name: SecurityName::Community(community),
security_level: SecurityLevel::NoAuthNoPriv,
context_name: Bytes::new(),
request_id,
pdu_type,
group_name: None,
read_view: None,
write_view: None,
msg_max_size: None,
decode_anomalies,
}
}
#[cfg(test)]
#[allow(clippy::too_many_arguments)]
pub(crate) fn usm(
source: SocketAddr,
username: Bytes,
security_level: SecurityLevel,
context_name: Bytes,
request_id: i32,
pdu_type: PduType,
msg_max_size: usize,
decode_anomalies: Vec<DecodeAnomaly>,
) -> Self {
Self::usm_with_lifecycle(
source,
username,
security_level,
context_name,
request_id,
pdu_type,
msg_max_size,
decode_anomalies,
RequestLifecycle::for_test(pdu_type),
)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn usm_with_lifecycle(
source: SocketAddr,
username: Bytes,
security_level: SecurityLevel,
context_name: Bytes,
request_id: i32,
pdu_type: PduType,
msg_max_size: usize,
decode_anomalies: Vec<DecodeAnomaly>,
lifecycle: RequestLifecycle,
) -> Self {
lifecycle.classify(pdu_type);
Self {
received_at: lifecycle.received_at,
admitted_at: lifecycle.admitted_at,
deadline: lifecycle.deadline,
cancellation: lifecycle.cancellation,
phase: lifecycle.phase,
source,
version: Version::V3,
security_model: SecurityModel::Usm,
security_name: SecurityName::Usm(username),
security_level,
context_name,
request_id,
pdu_type,
group_name: None,
read_view: None,
write_view: None,
msg_max_size: Some(msg_max_size),
decode_anomalies,
}
}
#[must_use]
pub const fn received_at(&self) -> tokio::time::Instant {
self.received_at
}
#[must_use]
pub const fn admitted_at(&self) -> tokio::time::Instant {
self.admitted_at
}
#[must_use]
pub const fn deadline(&self) -> Option<tokio::time::Instant> {
self.deadline
}
#[must_use]
pub fn cancellation_token(&self) -> CancellationToken {
self.cancellation.clone()
}
#[must_use]
pub fn is_cancelled(&self) -> bool {
self.cancellation.is_cancelled()
}
pub async fn cancelled(&self) {
self.cancellation.cancelled().await;
}
pub(crate) fn protect_set(&self) {
self.phase
.store(RequestTaskPhase::SetProtected as u8, Ordering::Release);
}
pub(crate) fn set_vacm_access(
&mut self,
group_name: Bytes,
read_view: Bytes,
write_view: Bytes,
) {
self.group_name = Some(group_name);
self.read_view = Some(read_view);
self.write_view = Some(write_view);
}
#[must_use]
pub const fn source(&self) -> SocketAddr {
self.source
}
#[must_use]
pub const fn version(&self) -> Version {
self.version
}
#[must_use]
pub const fn security_model(&self) -> SecurityModel {
self.security_model
}
#[must_use]
pub const fn security_name(&self) -> &SecurityName {
&self.security_name
}
#[must_use]
pub const fn security_level(&self) -> SecurityLevel {
self.security_level
}
#[must_use]
pub const fn context_name(&self) -> &Bytes {
&self.context_name
}
#[must_use]
pub const fn request_id(&self) -> i32 {
self.request_id
}
#[must_use]
pub const fn pdu_type(&self) -> PduType {
self.pdu_type
}
#[must_use]
pub const fn group_name(&self) -> Option<&Bytes> {
self.group_name.as_ref()
}
#[must_use]
pub const fn read_view(&self) -> Option<&Bytes> {
self.read_view.as_ref()
}
#[must_use]
pub const fn write_view(&self) -> Option<&Bytes> {
self.write_view.as_ref()
}
#[must_use]
pub const fn msg_max_size(&self) -> Option<usize> {
self.msg_max_size
}
#[must_use]
pub fn decode_anomalies(&self) -> &[DecodeAnomaly] {
&self.decode_anomalies
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub(crate) enum RequestTaskPhase {
Unclassified,
Retrieval,
SetPreCommit,
SetProtected,
Other,
}
#[derive(Debug, Clone)]
pub(crate) struct RequestLifecycle {
pub(crate) received_at: tokio::time::Instant,
pub(crate) admitted_at: tokio::time::Instant,
pub(crate) deadline: Option<tokio::time::Instant>,
pub(crate) cancellation: CancellationToken,
pub(crate) phase: Arc<AtomicU8>,
}
impl RequestLifecycle {
pub(crate) fn new(
received_at: tokio::time::Instant,
admitted_at: tokio::time::Instant,
deadline: Option<tokio::time::Instant>,
cancellation: CancellationToken,
phase: Arc<AtomicU8>,
) -> Self {
Self {
received_at,
admitted_at,
deadline,
cancellation,
phase,
}
}
#[cfg(test)]
fn for_test(pdu_type: PduType) -> Self {
let lifecycle = Self::standalone();
lifecycle.classify(pdu_type);
lifecycle
}
#[cfg(test)]
pub(crate) fn standalone() -> Self {
let now = tokio::time::Instant::now();
Self::new(
now,
now,
None,
CancellationToken::new(),
Arc::new(AtomicU8::new(RequestTaskPhase::Unclassified as u8)),
)
}
fn classify(&self, pdu_type: PduType) {
let phase = match pdu_type {
PduType::GetRequest | PduType::GetNextRequest | PduType::GetBulkRequest => {
RequestTaskPhase::Retrieval
}
PduType::SetRequest => RequestTaskPhase::SetPreCommit,
_ => RequestTaskPhase::Other,
};
self.phase.store(phase as u8, Ordering::Release);
}
}
#[cfg(test)]
mod tests {
use super::*;
use static_assertions::assert_impl_all;
assert_impl_all!(RequestContext: Send, Sync, Clone, std::fmt::Debug);
#[test]
fn community_accessors_cover_v1_v2c_and_non_utf8_names() {
let source = "192.0.2.55:6161".parse().unwrap();
let community = Community::from(b"reader\xff".as_slice());
for (community_version, version, model) in [
(CommunityVersion::V1, Version::V1, SecurityModel::V1),
(CommunityVersion::V2c, Version::V2c, SecurityModel::V2c),
] {
let context = RequestContext::community(
source,
community_version,
community.clone(),
416,
PduType::GetNextRequest,
Vec::new(),
);
assert_eq!(context.source(), source);
assert_eq!(context.version(), version);
assert_eq!(context.security_model(), model);
assert!(
context
.security_name()
.matches(&SecurityName::Community(community.clone()))
);
assert_eq!(context.security_level(), SecurityLevel::NoAuthNoPriv);
assert!(context.context_name().is_empty());
assert_eq!(context.request_id(), 416);
assert_eq!(context.pdu_type(), PduType::GetNextRequest);
assert_eq!(context.group_name(), None);
assert_eq!(context.read_view(), None);
assert_eq!(context.write_view(), None);
assert_eq!(context.msg_max_size(), None);
}
}
#[test]
fn usm_accessors_cover_every_security_level_and_octet_fields() {
let source = "[2001:db8::1]:6161".parse().unwrap();
let username = Bytes::from_static(b"operator\xff");
let context_name = Bytes::from_static(b"tenant\x80");
for level in [
SecurityLevel::NoAuthNoPriv,
SecurityLevel::AuthNoPriv,
SecurityLevel::AuthPriv,
] {
let mut context = RequestContext::usm(
source,
username.clone(),
level,
context_name.clone(),
-17,
PduType::SetRequest,
4096,
Vec::new(),
);
context.set_vacm_access(
Bytes::from_static(b"operators\xfe"),
Bytes::from_static(b"read\xfd"),
Bytes::from_static(b"write\xfc"),
);
assert_eq!(context.source(), source);
assert_eq!(context.version(), Version::V3);
assert_eq!(context.security_model(), SecurityModel::Usm);
assert_eq!(context.security_name().as_bytes(), username);
assert_eq!(context.security_level(), level);
assert_eq!(context.context_name(), &context_name);
assert_eq!(context.request_id(), -17);
assert_eq!(context.pdu_type(), PduType::SetRequest);
assert_eq!(context.group_name().unwrap(), b"operators\xfe".as_slice());
assert_eq!(context.read_view().unwrap(), b"read\xfd".as_slice());
assert_eq!(context.write_view().unwrap(), b"write\xfc".as_slice());
assert_eq!(context.msg_max_size(), Some(4096));
}
}
#[test]
fn debug_redacts_community_and_keeps_usm_username() {
let community = RequestContext::community(
"192.0.2.55:6161".parse().unwrap(),
CommunityVersion::V2c,
Community::from("community-redaction-sentinel-4d91"),
416,
PduType::GetRequest,
Vec::new(),
);
let rendered = format!("{community:?}");
assert!(rendered.contains("REDACTED"));
assert!(!rendered.contains("community-redaction-sentinel-4d91"));
assert!(rendered.contains("192.0.2.55:6161"));
assert!(rendered.contains("416"));
let usm = RequestContext::usm(
"192.0.2.55:6161".parse().unwrap(),
Bytes::from_static(b"visible-usm-user"),
SecurityLevel::AuthNoPriv,
Bytes::new(),
417,
PduType::GetRequest,
4096,
Vec::new(),
);
let rendered = format!("{usm:#?}");
assert!(rendered.contains("visible-usm-user"));
assert!(rendered.contains("192.0.2.55:6161"));
}
}