mod auth;
mod builder;
mod chunks;
mod response_shape;
mod retry;
mod v3;
mod walk;
pub use auth::{Auth, CommunityVersion};
pub use builder::{ClientBuilder, DEFAULT_CONSTRUCTION_TIMEOUT, Target, TargetClientBuilder};
pub use chunks::{FixedCardinalityChunk, FixedCardinalityChunkError, FixedCardinalityChunkStream};
pub use response_shape::{
BulkResponse, FixedCardinalityOperation, FixedCardinalityResponse, ResponseMetadata,
ResponseShapeAnomaly, ResponseShapePolicy,
};
pub use retry::{MAX_RETRIES, Retry, RetryBuilder, RetryConfigError};
impl Client<UdpHandle> {
pub fn builder(target: impl Into<Target>, auth: impl Into<Auth>) -> TargetClientBuilder {
ClientBuilder::new(auth).target(target)
}
#[must_use]
pub fn stats(&self) -> UdpStats {
self.inner.transport.stats()
}
}
#[cfg(test)]
use crate::error::ErrorStatus;
use crate::error::{Error, Result};
use crate::message::{CommunityMessage, Message, SecurityLevel};
use crate::oid::Oid;
use crate::pdu::{GetBulkPdu, NotificationPdu, Pdu, PduType, RequestPdu, TrapV1Notification};
use crate::transport::{Candidate, Transport, UdpHandle, UdpStats};
use crate::v3::{DesSaltState, EngineCache, EngineState, PrivProtocol, SaltCounter};
use crate::value::Value;
use crate::varbind::VarBind;
use crate::version::Version;
use response_shape::{RequestShape, classify};
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::RwLock;
use std::time::{Duration, Instant};
use tokio::sync::Mutex as AsyncMutex;
use tracing::{Span, instrument};
#[cfg(any(feature = "crypto-rustcrypto", feature = "crypto-fips"))]
pub use crate::v3::DerivedKeys;
#[cfg(not(any(feature = "crypto-rustcrypto", feature = "crypto-fips")))]
use crate::v3::DerivedKeys;
pub use crate::v3::UsmConfig;
pub use walk::{
OidOrdering, WalkCollection, WalkError, WalkItem, WalkMetadataStream, WalkMethod, WalkOptions,
WalkStream,
};
pub(crate) fn pdu_to_snmp_error(
pdu: &Pdu,
target: SocketAddr,
metadata: ResponseMetadata,
) -> Option<Box<Error>> {
if !pdu.is_error() {
return None;
}
let status = pdu.error_status_enum();
let oid = (pdu.error_index() as usize)
.checked_sub(1)
.and_then(|idx| pdu.varbinds.get(idx))
.map(|vb| Box::new(vb.oid.clone()));
Some(
Error::Snmp {
target,
status,
index: pdu.error_index().try_into().unwrap_or(0),
oid,
metadata: Box::new(metadata),
}
.boxed(),
)
}
pub const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(5);
pub const DEFAULT_SEND_TIMEOUT: Duration = Duration::from_secs(5);
pub const DEFAULT_MAX_OIDS_PER_REQUEST: usize = 10;
pub const DEFAULT_MAX_REPETITIONS: u32 = 25;
pub struct Client<T: Transport = UdpHandle> {
inner: Arc<ClientInner<T>>,
}
#[derive(Debug)]
pub(super) struct DecodedResponse {
pub(super) pdu: Pdu,
pub(super) decode_anomalies: Vec<crate::DecodeAnomaly>,
}
impl<T: Transport> Clone for Client<T> {
fn clone(&self) -> Self {
Self {
inner: Arc::clone(&self.inner),
}
}
}
struct ClientEngine {
state: EngineState,
derived_keys: DerivedKeys,
generation: Arc<()>,
}
#[derive(Debug)]
pub(crate) struct DiscoveryCoordinator {
flights: Mutex<HashMap<SocketAddr, Arc<DiscoveryFlight>>>,
}
impl DiscoveryCoordinator {
pub(crate) fn new() -> Self {
Self {
flights: Mutex::new(HashMap::new()),
}
}
fn lock_flights(&self) -> std::sync::MutexGuard<'_, HashMap<SocketAddr, Arc<DiscoveryFlight>>> {
match self.flights.lock() {
Ok(flights) => flights,
Err(poisoned) => {
let mut flights = poisoned.into_inner();
flights.clear();
self.flights.clear_poison();
flights
}
}
}
pub(crate) fn acquire(&self, target: SocketAddr) -> (Arc<DiscoveryFlight>, bool) {
let mut flights = self.lock_flights();
match flights.get(&target) {
Some(flight) => (Arc::clone(flight), false),
None => {
let flight = Arc::new(DiscoveryFlight::new());
flights.insert(target, Arc::clone(&flight));
(flight, true)
}
}
}
pub(crate) fn remove(&self, target: SocketAddr, flight: &Arc<DiscoveryFlight>) {
let mut flights = self.lock_flights();
if flights
.get(&target)
.is_some_and(|current| Arc::ptr_eq(current, flight))
{
flights.remove(&target);
}
}
#[cfg(test)]
pub(crate) fn flight_count(&self) -> usize {
self.lock_flights().len()
}
}
#[derive(Debug)]
pub(crate) struct DiscoveryFlight {
outcome: Mutex<Option<DiscoveryOutcome>>,
complete: tokio::sync::Notify,
retries: std::sync::atomic::AtomicU32,
}
impl DiscoveryFlight {
fn new() -> Self {
Self {
outcome: Mutex::new(None),
complete: tokio::sync::Notify::new(),
retries: std::sync::atomic::AtomicU32::new(0),
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct DiscoveredState {
state: EngineState,
metadata: ResponseMetadata,
}
#[derive(Debug, Clone)]
pub(crate) enum DiscoveryOutcome {
Success(DiscoveredState),
Timeout {
target: SocketAddr,
elapsed: Duration,
retries: u32,
},
Closed(SocketAddr),
Network {
target: SocketAddr,
kind: std::io::ErrorKind,
message: Arc<str>,
},
RequestIdInUse(i32),
OutboundMessageTooLarge {
size: usize,
limit: usize,
},
Auth(SocketAddr),
Decode(crate::DecodeError),
MalformedResponse(SocketAddr),
Config(Arc<str>),
InvalidMessage(Arc<str>),
InvalidOid(Arc<str>),
Failure(Arc<Error>),
}
impl DiscoveryOutcome {
fn share_result(result: Result<DiscoveredState>) -> (Self, Result<DiscoveredState>) {
match result {
Ok(discovered) => (Self::Success(discovered.clone()), Ok(discovered)),
Err(error) => {
let outcome = match &*error {
Error::Timeout {
target,
elapsed,
retries,
} => Self::Timeout {
target: *target,
elapsed: *elapsed,
retries: *retries,
},
Error::Closed { target } => Self::Closed(*target),
Error::Network { target, source } => Self::Network {
target: *target,
kind: source.kind(),
message: source.to_string().into(),
},
Error::RequestIdInUse { request_id } => Self::RequestIdInUse(*request_id),
Error::OutboundMessageTooLarge { size, limit } => {
Self::OutboundMessageTooLarge {
size: *size,
limit: *limit,
}
}
Error::Auth { target } => Self::Auth(*target),
Error::Decode(error) => Self::Decode(error.clone()),
Error::MalformedResponse { target } => Self::MalformedResponse(*target),
Error::Config(message) => Self::Config(message.as_ref().into()),
Error::InvalidMessage(message) => Self::InvalidMessage(message.as_ref().into()),
Error::InvalidOid(message) => Self::InvalidOid(message.as_ref().into()),
_ => {
let source: Arc<Error> = error.into();
return (
Self::Failure(Arc::clone(&source)),
Err(Error::SharedOperation { source }.boxed()),
);
}
};
(outcome, Err(error))
}
}
}
fn into_result(self) -> Result<DiscoveredState> {
match self {
Self::Success(metadata) => Ok(metadata),
Self::Timeout {
target,
elapsed,
retries,
} => Err(Error::Timeout {
target,
elapsed,
retries,
}
.boxed()),
Self::Closed(target) => Err(Error::Closed { target }.boxed()),
Self::Network {
target,
kind,
message,
} => Err(Error::Network {
target,
source: std::io::Error::new(kind, message.to_string()),
}
.boxed()),
Self::RequestIdInUse(request_id) => Err(Error::RequestIdInUse { request_id }.boxed()),
Self::OutboundMessageTooLarge { size, limit } => {
Err(Error::OutboundMessageTooLarge { size, limit }.boxed())
}
Self::Auth(target) => Err(Error::Auth { target }.boxed()),
Self::Decode(error) => Err(Error::Decode(error).boxed()),
Self::MalformedResponse(target) => Err(Error::MalformedResponse { target }.boxed()),
Self::Config(message) => Err(Error::Config(message.as_ref().into()).boxed()),
Self::InvalidMessage(message) => {
Err(Error::InvalidMessage(message.as_ref().into()).boxed())
}
Self::InvalidOid(message) => Err(Error::InvalidOid(message.as_ref().into()).boxed()),
Self::Failure(source) => Err(Error::SharedOperation { source }.boxed()),
}
}
}
impl ClientEngine {
fn new(state: EngineState, derived_keys: DerivedKeys) -> Self {
Self {
state,
derived_keys,
generation: Arc::new(()),
}
}
}
struct ClientInner<T: Transport> {
transport: T,
config: ClientConfig,
engine: RwLock<Option<ClientEngine>>,
salt_counter: Option<SaltCounter>,
engine_cache: Option<Arc<EngineCache>>,
discovery_lock: AsyncMutex<()>,
discovery_coordinator: Arc<DiscoveryCoordinator>,
local_derived_keys: RwLock<Option<DerivedKeys>>,
#[cfg(test)]
authenticated_response_validated_hook: RwLock<Option<Arc<dyn Fn() + Send + Sync>>>,
}
pub(crate) type LocalAuthoritativeTimeSource = Arc<dyn Fn() -> Result<(u32, u32)> + Send + Sync>;
#[derive(Clone)]
#[non_exhaustive]
pub struct ClientConfig {
pub auth: Auth,
pub decode_config: crate::DecodeConfig,
pub community_response_policy: crate::transport::CommunityResponsePolicy,
pub request_timeout: Duration,
pub exchange_timeout: Option<Duration>,
pub send_timeout: Duration,
pub retry: Retry,
pub max_oids_per_request: usize,
pub response_shape_policy: ResponseShapePolicy,
pub allow_unauthenticated_v3_time_correction: bool,
pub walk_options: WalkOptions,
pub local_authoritative_engine: Option<crate::v3::AuthoritativeEngine>,
pub des_salt_state: Option<DesSaltState>,
pub(crate) local_authoritative_time_source: Option<LocalAuthoritativeTimeSource>,
}
impl std::fmt::Debug for ClientConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ClientConfig")
.field("auth", &self.auth)
.field("decode_config", &self.decode_config)
.field("community_response_policy", &self.community_response_policy)
.field("request_timeout", &self.request_timeout)
.field("exchange_timeout", &self.exchange_timeout)
.field("send_timeout", &self.send_timeout)
.field("retry", &self.retry)
.field("max_oids_per_request", &self.max_oids_per_request)
.field("response_shape_policy", &self.response_shape_policy)
.field(
"allow_unauthenticated_v3_time_correction",
&self.allow_unauthenticated_v3_time_correction,
)
.field("walk_options", &self.walk_options)
.field(
"local_authoritative_engine",
&self.local_authoritative_engine,
)
.field("des_salt_state", &self.des_salt_state)
.field(
"local_authoritative_time_source",
&self
.local_authoritative_time_source
.as_ref()
.map(|_| "<callback>"),
)
.finish()
}
}
impl Default for ClientConfig {
fn default() -> Self {
Self {
auth: Auth::default(),
decode_config: crate::DecodeConfig::default(),
community_response_policy: crate::transport::CommunityResponsePolicy::Exact,
request_timeout: DEFAULT_REQUEST_TIMEOUT,
exchange_timeout: None,
send_timeout: DEFAULT_SEND_TIMEOUT,
retry: Retry::default(),
max_oids_per_request: DEFAULT_MAX_OIDS_PER_REQUEST,
response_shape_policy: ResponseShapePolicy::Compatible,
allow_unauthenticated_v3_time_correction: false,
walk_options: WalkOptions::default(),
local_authoritative_engine: None,
des_salt_state: None,
local_authoritative_time_source: None,
}
}
}
impl ClientConfig {
fn version(&self) -> Version {
self.auth.version()
}
fn community(&self) -> Result<crate::Community> {
self.auth
.community()
.cloned()
.ok_or_else(|| Error::Config("community authentication required".into()).boxed())
}
fn community_version(&self) -> Result<CommunityVersion> {
self.auth
.community_version()
.ok_or_else(|| Error::Config("community authentication required".into()).boxed())
}
fn usm_config(&self) -> Option<&UsmConfig> {
self.auth.usm_config()
}
pub(super) fn validate(&self) -> Result<()> {
crate::transport::checked_deadline(self.request_timeout, "request timeout")?;
crate::transport::checked_deadline(self.send_timeout, "send timeout")?;
if let Some(timeout) = self.exchange_timeout {
crate::transport::checked_deadline(timeout, "exchange timeout")?;
}
if self.max_oids_per_request == 0 {
return Err(
Error::Config("max_oids_per_request must be greater than 0".into()).boxed(),
);
}
if self.walk_options.max_repetitions > crate::pdu::MAX_GET_BULK_VALUE {
return Err(Error::Config("max_repetitions exceeds i32::MAX".into()).boxed());
}
self.walk_options.validate(self.version())?;
let uses_des = self
.usm_config()
.and_then(UsmConfig::priv_protocol)
.is_some_and(PrivProtocol::is_des_family);
if uses_des && self.des_salt_state.is_none() {
return Err(Error::Config(
"durable DES sender state is required for DES/3DES privacy".into(),
)
.boxed());
}
if uses_des
&& let (Some(engine), Some(state)) =
(&self.local_authoritative_engine, &self.des_salt_state)
&& engine.engine_boots() != state.engine_boots()
{
return Err(Error::Config(
"DES sender boots must match the local authoritative engine boots".into(),
)
.boxed());
}
Ok(())
}
fn validate_and_precompute(&mut self) -> Result<()> {
self.validate()?;
if let Auth::Usm(config) = &mut self.auth {
config.validate_and_precompute().map_err(|error| {
Error::Config(format!("invalid USM configuration: {error}").into()).boxed()
})?;
}
Ok(())
}
}
impl<T: Transport> Client<T> {
pub fn new(transport: T, config: ClientConfig) -> Result<Self> {
Self::with_optional_engine_cache(transport, config, None)
}
pub fn with_engine_cache(
transport: T,
config: ClientConfig,
engine_cache: Arc<EngineCache>,
) -> Result<Self> {
Self::with_optional_engine_cache(transport, config, Some(engine_cache))
}
fn with_optional_engine_cache(
transport: T,
mut config: ClientConfig,
engine_cache: Option<Arc<EngineCache>>,
) -> Result<Self> {
config.validate_and_precompute()?;
let salt_counter = config
.usm_config()
.filter(|security| {
security.security_level().requires_priv()
&& !security
.priv_protocol()
.is_some_and(PrivProtocol::is_des_family)
})
.map(|_| SaltCounter::new())
.transpose()?;
let discovery_coordinator = engine_cache.as_ref().map_or_else(
|| Arc::new(DiscoveryCoordinator::new()),
|cache| Arc::clone(&cache.discovery_coordinator),
);
Ok(Self {
inner: Arc::new(ClientInner {
transport,
config,
engine: RwLock::new(None),
salt_counter,
engine_cache,
discovery_lock: AsyncMutex::new(()),
discovery_coordinator,
local_derived_keys: RwLock::new(None),
#[cfg(test)]
authenticated_response_validated_hook: RwLock::new(None),
}),
})
}
#[must_use]
pub fn peer_addr(&self) -> SocketAddr {
self.inner.transport.peer_addr()
}
#[must_use]
pub fn version(&self) -> Version {
self.inner.config.version()
}
#[must_use]
pub fn decode_config(&self) -> crate::DecodeConfig {
self.inner.config.decode_config
}
#[must_use]
pub fn walk_options(&self) -> WalkOptions {
self.inner.config.walk_options
}
#[must_use]
pub fn security_level(&self) -> Option<SecurityLevel> {
self.inner
.config
.usm_config()
.map(UsmConfig::security_level)
}
fn next_request_id(&self) -> i32 {
self.inner.transport.alloc_request_id()
}
fn is_v3(&self) -> bool {
matches!(self.inner.config.auth, Auth::Usm(_))
}
pub(super) fn enforce_outbound_size(
&self,
encoded_size: usize,
remote_receive_capacity: Option<crate::MessageSize>,
) -> Result<()> {
let transport_capacity = self.inner.transport.send_capacity();
let effective_limit = remote_receive_capacity
.map(crate::MessageSize::as_usize)
.map_or(transport_capacity, |remote| transport_capacity.min(remote));
crate::message_size::enforce_outbound_size(encoded_size, effective_limit)
}
fn start_exchange_deadline(&self) -> Result<Option<tokio::time::Instant>> {
self.inner
.config
.exchange_timeout
.map(|timeout| {
tokio::time::Instant::now()
.checked_add(timeout)
.ok_or_else(|| {
Error::Config("exchange timeout exceeds the representable deadline".into())
.boxed()
})
})
.transpose()
}
fn transmission_deadline(
&self,
exchange_deadline: Option<tokio::time::Instant>,
) -> Result<tokio::time::Instant> {
let attempt = tokio::time::Instant::now()
.checked_add(self.inner.config.request_timeout)
.ok_or_else(|| {
Error::Config("request timeout exceeds the representable deadline".into()).boxed()
})?;
Ok(exchange_deadline.map_or(attempt, |deadline| deadline.min(attempt)))
}
fn retry_retention_deadline(
&self,
attempt_deadline: tokio::time::Instant,
retry_delay: Duration,
exchange_deadline: Option<tokio::time::Instant>,
) -> Result<tokio::time::Instant> {
if let Some(deadline) = exchange_deadline {
return Ok(deadline);
}
attempt_deadline
.checked_add(retry_delay)
.and_then(|deadline| deadline.checked_add(self.inner.config.request_timeout))
.ok_or_else(|| {
Error::Config("request retry schedule exceeds the representable deadline".into())
.boxed()
})
}
#[instrument(
level = "debug",
skip(self, data),
fields(
snmp.target = %self.peer_addr(),
snmp.request_id = request_id,
snmp.attempt = tracing::field::Empty,
snmp.elapsed_ms = tracing::field::Empty,
)
)]
async fn send_and_recv(&self, request_id: i32, data: &[u8]) -> Result<DecodedResponse> {
self.enforce_outbound_size(data.len(), None)?;
let start = Instant::now();
let exchange_deadline = self.start_exchange_deadline()?;
let max_attempts = if self.inner.transport.is_reliable() {
0
} else {
self.inner.config.retry.retries()
};
let mut retries = 0;
for attempt in 0..=max_attempts {
if attempt > 0 {
retries = attempt;
}
Span::current().record("snmp.attempt", attempt);
if attempt > 0 {
tracing::debug!(target: "async_snmp::client", "retrying request");
}
let version = self.inner.config.version();
let community_version = match version {
Version::V1 => CommunityVersion::V1,
Version::V2c => CommunityVersion::V2c,
Version::V3 => unreachable!("community request path cannot use SNMPv3"),
};
let community = self.inner.config.community()?;
let registration = crate::transport::RequestRegistration::community(
request_id,
self.transmission_deadline(exchange_deadline)?,
community_version,
community.clone(),
self.inner.config.community_response_policy,
)
.with_decode_config(self.inner.config.decode_config);
tracing::trace!(target: "async_snmp::client", { snmp.bytes = data.len() }, "sending request");
match self
.inner
.transport
.request_with(data, registration, |response_data, source| {
tracing::trace!(target: "async_snmp::client", { snmp.bytes = response_data.len() }, "received response candidate");
let Ok(decoded) = Message::decode_bounded_with_target(
response_data,
self.inner.transport.receive_limits().accepted(),
Some(source),
self.inner.config.decode_config,
) else {
return Ok(Candidate::Reject);
};
let response = decoded.value;
if response.version() != version {
return Ok(Candidate::Reject);
}
if let Message::Community(ref message) = response
&& !community.matches(message.community().as_bytes())
{
let accepted = match self.inner.config.community_response_policy {
crate::transport::CommunityResponsePolicy::Exact => false,
crate::transport::CommunityResponsePolicy::AllowMismatchFromTarget => {
source == self.peer_addr()
}
crate::transport::CommunityResponsePolicy::AllowMismatchFromAnySource => true,
};
if !accepted {
return Ok(Candidate::Reject);
}
}
let Some(response_pdu) = response.into_pdu() else {
return Ok(Candidate::Reject);
};
if response_pdu.pdu_type() != PduType::Response
|| response_pdu.request_id != request_id
{
return Ok(Candidate::Reject);
}
Ok(Candidate::Accept(DecodedResponse {
pdu: response_pdu,
decode_anomalies: decoded.anomalies,
}))
})
.await
{
Ok(response) => {
if let Some(err) = pdu_to_snmp_error(
&response.pdu,
self.peer_addr(),
ResponseMetadata::from_decode_anomalies(response.decode_anomalies.clone()),
) {
Span::current()
.record("snmp.elapsed_ms", start.elapsed().as_millis() as u64);
return Err(err);
}
Span::current().record("snmp.elapsed_ms", start.elapsed().as_millis() as u64);
return Ok(response);
}
Err(e) if matches!(*e, Error::Timeout { .. }) => {
if attempt < max_attempts {
let delay = self.inner.config.retry.compute_delay(attempt);
if !delay.is_zero() {
tracing::debug!(target: "async_snmp::client", { delay_ms = delay.as_millis() as u64 }, "backing off");
}
if !retry::wait_for_retry(delay, exchange_deadline).await {
break;
}
}
}
Err(e) => {
Span::current().record("snmp.elapsed_ms", start.elapsed().as_millis() as u64);
return Err(e);
}
}
}
let elapsed = start.elapsed();
Span::current().record("snmp.elapsed_ms", elapsed.as_millis() as u64);
tracing::debug!(target: "async_snmp::client", { request_id, peer = %self.peer_addr(), ?elapsed, retries }, "request timed out");
Err(Error::Timeout {
target: self.peer_addr(),
elapsed,
retries,
}
.boxed())
}
async fn send_request(&self, pdu: Pdu) -> Result<DecodedResponse> {
if self.is_v3() {
return self.send_v3_and_recv(pdu).await;
}
tracing::debug!(target: "async_snmp::client", { snmp.pdu_type = ?pdu.pdu_type(), snmp.varbind_count = pdu.varbinds.len() }, "sending {} request", pdu.pdu_type());
let request_id = pdu.request_id;
let message = CommunityMessage::new(
self.inner.config.community_version()?,
self.inner.config.community()?,
pdu,
)?;
let data = message.encode()?;
let response = self.send_and_recv(request_id, &data).await?;
tracing::debug!(target: "async_snmp::client", { snmp.pdu_type = ?response.pdu.pdu_type(), snmp.varbind_count = response.pdu.varbinds.len(), snmp.error_status = response.pdu.error_status(), snmp.error_index = response.pdu.error_index() }, "received {} response", response.pdu.pdu_type());
Ok(response)
}
fn apply_response_shape_policy(
&self,
response: FixedCardinalityResponse,
) -> Result<FixedCardinalityResponse> {
if self.inner.config.response_shape_policy == ResponseShapePolicy::Strict
&& !response.anomalies.is_empty()
{
return Err(Error::ResponseShape {
target: self.peer_addr(),
response,
}
.boxed());
}
Ok(response)
}
#[instrument(skip(self), err, fields(snmp.target = %self.peer_addr(), snmp.oid = %oid))]
pub async fn get(&self, oid: &Oid) -> Result<FixedCardinalityResponse> {
let request_id = self.next_request_id();
let pdu = RequestPdu::get(
self.inner.config.version(),
request_id,
std::slice::from_ref(oid),
)?
.into_raw();
let response = self.send_request(pdu).await?;
let mut classified = classify(
RequestShape::Get(std::slice::from_ref(oid)),
response.pdu.varbinds,
0,
0,
);
classified.metadata.decode_anomalies = response.decode_anomalies;
self.apply_response_shape_policy(classified)
}
#[instrument(skip(self, oids), err, fields(snmp.target = %self.peer_addr(), snmp.oid_count = oids.len()))]
pub async fn get_many(&self, oids: &[Oid]) -> Result<FixedCardinalityResponse>
where
T: 'static,
{
self.get_many_chunks(oids)?.collect_response().await
}
pub fn get_many_chunks(&self, oids: &[Oid]) -> Result<FixedCardinalityChunkStream<T>>
where
T: 'static,
{
FixedCardinalityChunkStream::new(self, oids, FixedCardinalityOperation::Get)
}
#[instrument(skip(self), err, fields(snmp.target = %self.peer_addr(), snmp.oid = %oid))]
pub async fn get_next(&self, oid: &Oid) -> Result<FixedCardinalityResponse> {
let request_id = self.next_request_id();
let pdu = RequestPdu::get_next(
self.inner.config.version(),
request_id,
std::slice::from_ref(oid),
)?
.into_raw();
let response = self.send_request(pdu).await?;
let mut classified = classify(
RequestShape::GetNext(std::slice::from_ref(oid)),
response.pdu.varbinds,
0,
0,
);
classified.metadata.decode_anomalies = response.decode_anomalies;
self.apply_response_shape_policy(classified)
}
#[instrument(skip(self, oids), err, fields(snmp.target = %self.peer_addr(), snmp.oid_count = oids.len()))]
pub async fn get_next_many(&self, oids: &[Oid]) -> Result<FixedCardinalityResponse>
where
T: 'static,
{
self.get_next_many_chunks(oids)?.collect_response().await
}
pub fn get_next_many_chunks(&self, oids: &[Oid]) -> Result<FixedCardinalityChunkStream<T>>
where
T: 'static,
{
FixedCardinalityChunkStream::new(self, oids, FixedCardinalityOperation::GetNext)
}
#[instrument(skip(self, value), err, fields(snmp.target = %self.peer_addr(), snmp.oid = %oid))]
pub async fn set(&self, oid: &Oid, value: Value) -> Result<FixedCardinalityResponse> {
let request_id = self.next_request_id();
let requested = [(oid.clone(), value)];
let pdu = RequestPdu::set(
self.inner.config.version(),
request_id,
vec![VarBind::new(requested[0].0.clone(), requested[0].1.clone())],
)?
.into_raw();
let response = self.send_request(pdu).await?;
let mut classified = classify(RequestShape::Set(&requested), response.pdu.varbinds, 0, 0);
classified.metadata.decode_anomalies = response.decode_anomalies;
self.apply_response_shape_policy(classified)
}
#[instrument(skip(self, varbinds), err, fields(snmp.target = %self.peer_addr(), snmp.oid_count = varbinds.len()))]
pub async fn set_many(&self, varbinds: &[(Oid, Value)]) -> Result<FixedCardinalityResponse> {
if varbinds.is_empty() {
return Ok(FixedCardinalityResponse::empty(
FixedCardinalityOperation::Set,
));
}
let max_per_request = self.inner.config.max_oids_per_request;
if varbinds.len() > max_per_request {
return Err(Error::Config(
format!(
"set_many: {} varbinds exceeds max_oids_per_request ({}); \
SET must be atomic and cannot be split across PDUs",
varbinds.len(),
max_per_request,
)
.into(),
)
.boxed());
}
let request_id = self.next_request_id();
let vbs: Vec<VarBind> = varbinds
.iter()
.map(|(oid, value)| VarBind::new(oid.clone(), value.clone()))
.collect();
let pdu = RequestPdu::set(self.inner.config.version(), request_id, vbs)?.into_raw();
let response = self.send_request(pdu).await?;
let mut classified = classify(RequestShape::Set(varbinds), response.pdu.varbinds, 0, 0);
classified.metadata.decode_anomalies = response.decode_anomalies;
self.apply_response_shape_policy(classified)
}
#[instrument(skip(self, varbinds), err, fields(snmp.target = %self.peer_addr(), snmp.trap_oid = %trap_oid))]
pub async fn send_trap(
&self,
trap_oid: &Oid,
uptime: u32,
varbinds: Vec<VarBind>,
) -> Result<()> {
if self.inner.config.version() == Version::V1 {
let local_ip = match self.inner.transport.local_addr().ip() {
std::net::IpAddr::V4(v4) => v4.octets(),
std::net::IpAddr::V6(_) => [0, 0, 0, 0],
};
let pdu = NotificationPdu::trap_v2(Version::V2c, 0, uptime, trap_oid, varbinds)?;
return self.send_v1_trap(pdu.to_v1_trap(local_ip)?).await;
}
let request_id = self.next_request_id();
let pdu = NotificationPdu::trap_v2(
self.inner.config.version(),
request_id,
uptime,
trap_oid,
varbinds,
)?;
if self.is_v3() {
self.ensure_local_keys_derived()?;
let msg_id = self.next_request_id();
let data = self.build_v3_trap_message(pdu.as_raw(), msg_id)?;
self.enforce_outbound_size(data.len(), None)?;
tracing::debug!(target: "async_snmp::client", { snmp.pdu_type = "TrapV2", snmp.varbind_count = pdu.as_raw().varbinds().len(), snmp.bytes = data.len() }, "sending V3 trap");
self.inner
.transport
.send_with_timeout(&data, self.inner.config.send_timeout)
.await?;
} else {
let message = CommunityMessage::new(
self.inner.config.community_version()?,
self.inner.config.community()?,
pdu,
)?;
let data = message.encode()?;
self.enforce_outbound_size(data.len(), None)?;
tracing::debug!(target: "async_snmp::client", { snmp.pdu_type = "TrapV2", snmp.bytes = data.len() }, "sending v2c trap");
self.inner
.transport
.send_with_timeout(&data, self.inner.config.send_timeout)
.await?;
}
Ok(())
}
#[instrument(skip(self, trap), err, fields(snmp.target = %self.peer_addr(), snmp.generic_trap = %trap.as_raw().generic_trap()))]
pub async fn send_v1_trap(&self, trap: TrapV1Notification) -> Result<()> {
if self.inner.config.version() != Version::V1 {
return Err(Error::Config("send_v1_trap requires a V1 client".into()).boxed());
}
let message = CommunityMessage::v1_trap(self.inner.config.community()?, trap.into_raw())?;
let data = message.encode()?;
self.enforce_outbound_size(data.len(), None)?;
tracing::debug!(target: "async_snmp::client", { snmp.pdu_type = "TrapV1", snmp.bytes = data.len() }, "sending v1 trap");
self.inner
.transport
.send_with_timeout(&data, self.inner.config.send_timeout)
.await?;
Ok(())
}
#[instrument(skip(self, varbinds), err, fields(snmp.target = %self.peer_addr(), snmp.trap_oid = %trap_oid))]
pub async fn send_inform(
&self,
trap_oid: &Oid,
uptime: u32,
varbinds: Vec<VarBind>,
) -> Result<()> {
self.send_inform_with_metadata(trap_oid, uptime, varbinds)
.await
.map(|_| ())
}
#[instrument(skip(self, varbinds), err, fields(snmp.target = %self.peer_addr(), snmp.trap_oid = %trap_oid))]
pub async fn send_inform_with_metadata(
&self,
trap_oid: &Oid,
uptime: u32,
varbinds: Vec<VarBind>,
) -> Result<ResponseMetadata> {
if self.inner.config.version() == Version::V1 {
return Err(Error::Config("v1 inform sending not supported".into()).boxed());
}
let request_id = self.next_request_id();
let pdu = NotificationPdu::inform(
self.inner.config.version(),
request_id,
uptime,
trap_oid,
varbinds,
)?;
let expected_varbinds = pdu.as_raw().varbinds().to_vec();
let response = self.send_request(pdu.into_raw()).await?;
if response.pdu.varbinds != expected_varbinds {
let metadata = ResponseMetadata::from_decode_anomalies(response.decode_anomalies);
return Err(Error::MalformedResponse {
target: self.peer_addr(),
}
.boxed()
.with_prior_response_metadata(&metadata));
}
Ok(ResponseMetadata::from_decode_anomalies(
response.decode_anomalies,
))
}
#[instrument(skip(self, oids), err, fields(
snmp.target = %self.peer_addr(),
snmp.oid_count = oids.len(),
snmp.non_repeaters = non_repeaters,
snmp.max_repetitions = max_repetitions
))]
pub async fn get_bulk(
&self,
oids: &[Oid],
non_repeaters: u32,
max_repetitions: u32,
) -> Result<Vec<VarBind>> {
Ok(self
.get_bulk_with_metadata(oids, non_repeaters, max_repetitions)
.await?
.varbinds)
}
pub async fn get_bulk_with_metadata(
&self,
oids: &[Oid],
non_repeaters: u32,
max_repetitions: u32,
) -> Result<BulkResponse> {
Pdu::checked_get_bulk_fields(non_repeaters, max_repetitions)?;
let request_id = self.next_request_id();
let pdu = GetBulkPdu::new(
self.inner.config.version(),
request_id,
non_repeaters,
max_repetitions,
oids.iter().map(|oid| VarBind::null(oid.clone())).collect(),
)?
.into_raw();
let response = self.send_request(pdu).await?;
Ok(BulkResponse {
varbinds: response.pdu.varbinds,
metadata: ResponseMetadata {
decode_anomalies: response.decode_anomalies,
},
})
}
#[instrument(skip(self), fields(snmp.target = %self.peer_addr(), snmp.oid = %oid))]
pub fn walk(&self, oid: Oid) -> Result<WalkStream<T>>
where
T: 'static,
{
self.walk_with(oid, self.inner.config.walk_options)
}
#[instrument(skip(self), fields(snmp.target = %self.peer_addr(), snmp.oid = %oid, snmp.walk_method = ?options.method))]
pub fn walk_with(&self, oid: Oid, options: WalkOptions) -> Result<WalkStream<T>>
where
T: 'static,
{
WalkStream::new(self.clone(), oid, self.inner.config.version(), options)
}
#[instrument(skip(self), fields(snmp.target = %self.peer_addr(), snmp.oid = %oid))]
pub fn walk_with_metadata(&self, oid: Oid) -> Result<WalkMetadataStream<T>>
where
T: 'static,
{
self.walk_with_metadata_and(oid, self.inner.config.walk_options)
}
#[instrument(skip(self), fields(snmp.target = %self.peer_addr(), snmp.oid = %oid, snmp.walk_method = ?options.method))]
pub fn walk_with_metadata_and(
&self,
oid: Oid,
options: WalkOptions,
) -> Result<WalkMetadataStream<T>>
where
T: 'static,
{
self.walk_with(oid, options).map(WalkMetadataStream::new)
}
pub fn walk_getnext(&self, oid: Oid) -> Result<WalkStream<T>>
where
T: 'static,
{
let mut options = self.inner.config.walk_options;
options.method = WalkMethod::GetNext;
self.walk_with(oid, options)
}
pub fn walk_getnext_with_metadata(&self, oid: Oid) -> Result<WalkMetadataStream<T>>
where
T: 'static,
{
let mut options = self.inner.config.walk_options;
options.method = WalkMethod::GetNext;
self.walk_with_metadata_and(oid, options)
}
pub fn bulk_walk(&self, oid: Oid, max_repetitions: u32) -> Result<WalkStream<T>>
where
T: 'static,
{
let mut options = self.inner.config.walk_options;
options.method = WalkMethod::GetBulk;
options.max_repetitions = max_repetitions;
self.walk_with(oid, options)
}
pub fn bulk_walk_with_metadata(
&self,
oid: Oid,
max_repetitions: u32,
) -> Result<WalkMetadataStream<T>>
where
T: 'static,
{
let mut options = self.inner.config.walk_options;
options.method = WalkMethod::GetBulk;
options.max_repetitions = max_repetitions;
self.walk_with_metadata_and(oid, options)
}
pub fn bulk_walk_default(&self, oid: Oid) -> Result<WalkStream<T>>
where
T: 'static,
{
let mut options = self.inner.config.walk_options;
options.method = WalkMethod::GetBulk;
self.walk_with(oid, options)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::message::CommunityMessage;
use crate::oid;
use crate::oid::Oid;
use crate::pdu::{Pdu, PduType};
use crate::varbind::VarBind;
use bytes::Bytes;
use std::collections::VecDeque;
use std::net::SocketAddr;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
#[derive(Clone)]
struct TruncatingTransport {
response_varbind_count: usize,
pending: Arc<Mutex<VecDeque<i32>>>,
}
impl TruncatingTransport {
fn new(response_varbind_count: usize) -> Self {
Self {
response_varbind_count,
pending: Arc::new(Mutex::new(VecDeque::new())),
}
}
}
impl TruncatingTransport {
fn recv(
&self,
_registration: crate::transport::RequestRegistration,
) -> impl std::future::Future<Output = Result<(Bytes, SocketAddr)>> + Send {
let request_id = {
let mut q = self.pending.lock().unwrap();
q.pop_front().unwrap_or(1)
};
let n = self.response_varbind_count;
let peer: SocketAddr = "127.0.0.1:161".parse().unwrap();
async move {
let varbinds: Vec<VarBind> = (0..n)
.map(|i| {
VarBind::new(
Oid::from_slice(&[1, 3, 6, 1, i as u32]),
crate::value::Value::Null,
)
})
.collect();
let pdu = Pdu::response(request_id, 0, 0, varbinds);
let msg = CommunityMessage::v2c(Bytes::from_static(b"public"), pdu).unwrap();
let encoded = msg.encode().unwrap();
Ok((encoded, peer))
}
}
}
impl Transport for TruncatingTransport {
fn send(&self, data: &[u8]) -> impl std::future::Future<Output = Result<()>> + Send {
let request_id = crate::transport::extract_request_id(data).unwrap_or(1);
{
let mut q = self.pending.lock().unwrap();
q.push_back(request_id);
}
async { Ok(()) }
}
fn request_with<T, F>(
&self,
data: &[u8],
registration: crate::transport::RequestRegistration,
validate: F,
) -> impl std::future::Future<Output = Result<T>> + Send
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
crate::transport::request_with_scripted(
self,
data,
registration,
move |registration| {
futures_util::stream::once(async move { self.recv(registration).await })
},
validate,
)
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
}
#[derive(Clone)]
struct ImmediateTimeoutTransport {
calls: Arc<AtomicUsize>,
}
impl Transport for ImmediateTimeoutTransport {
async fn send(&self, _data: &[u8]) -> Result<()> {
Ok(())
}
async fn request_with<T, F>(
&self,
_data: &[u8],
_registration: crate::transport::RequestRegistration,
_validate: F,
) -> Result<T>
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
self.calls.fetch_add(1, Ordering::Relaxed);
Err(Error::Timeout {
target: self.peer_addr(),
elapsed: Duration::ZERO,
retries: 0,
}
.boxed())
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
false
}
}
#[tokio::test]
async fn zero_backoff_community_retries_remain_cancellable() {
let calls = Arc::new(AtomicUsize::new(0));
let client = Client::new(
ImmediateTimeoutTransport {
calls: calls.clone(),
},
ClientConfig {
auth: Auth::v2c("public"),
retry: Retry::fixed(crate::MAX_RETRIES, Duration::ZERO).unwrap(),
..ClientConfig::default()
},
)
.unwrap();
let request =
tokio::spawn(async move { client.get(&oid!(1, 3, 6, 1, 2, 1, 1, 1, 0)).await });
for _ in 0..10 {
tokio::task::yield_now().await;
}
request.abort();
assert!(
calls.load(Ordering::Relaxed) < 10_000,
"retry loop did not yield to cancellation"
);
}
#[tokio::test(start_paused = true)]
async fn exchange_deadline_caps_retry_backoff() {
let calls = Arc::new(AtomicUsize::new(0));
let client = Client::new(
ImmediateTimeoutTransport {
calls: Arc::clone(&calls),
},
ClientConfig {
auth: Auth::v2c("public"),
request_timeout: Duration::from_secs(30),
exchange_timeout: Some(Duration::from_secs(5)),
retry: Retry::fixed(2, Duration::from_secs(10)).unwrap(),
..ClientConfig::default()
},
)
.unwrap();
let requested_oid = oid!(1, 3, 6, 1, 2, 1, 1, 1, 0);
let started = tokio::time::Instant::now();
let request = client.get(&requested_oid);
tokio::pin!(request);
assert!(futures::poll!(request.as_mut()).is_pending());
assert_eq!(calls.load(Ordering::Relaxed), 1);
tokio::time::advance(Duration::from_secs(5)).await;
let error = request.await.unwrap_err();
assert!(matches!(*error, Error::Timeout { retries: 0, .. }));
assert_eq!(
tokio::time::Instant::now() - started,
Duration::from_secs(5)
);
assert_eq!(calls.load(Ordering::Relaxed), 1);
}
#[test]
fn unrepresentable_exchange_timeout_is_rejected_at_construction() {
let calls = Arc::new(AtomicUsize::new(0));
let result = Client::new(
ImmediateTimeoutTransport { calls },
ClientConfig {
auth: Auth::v2c("public"),
exchange_timeout: Some(Duration::MAX),
..ClientConfig::default()
},
);
assert!(matches!(result, Err(error) if matches!(*error, Error::Config(_))));
}
fn metadata_client(auth: Auth) -> Client<TruncatingTransport> {
Client::new(
TruncatingTransport::new(0),
ClientConfig {
auth,
retry: Retry::none(),
..Default::default()
},
)
.expect("valid client config")
}
#[test]
fn client_protocol_metadata_covers_versions_and_security_levels() {
let v1 = metadata_client(Auth::v1("private"));
assert_eq!(v1.version(), Version::V1);
assert_eq!(v1.security_level(), None);
assert!(v1.inner.salt_counter.is_none());
let v2c = metadata_client(Auth::v2c("public"));
assert_eq!(v2c.version(), Version::V2c);
assert_eq!(v2c.security_level(), None);
assert!(v2c.inner.salt_counter.is_none());
let no_auth = metadata_client(Auth::usm("no-auth-user"));
assert_eq!(no_auth.version(), Version::V3);
assert_eq!(no_auth.security_level(), Some(SecurityLevel::NoAuthNoPriv));
assert!(no_auth.inner.salt_counter.is_none());
#[cfg(any(feature = "crypto-rustcrypto", feature = "crypto-fips"))]
{
let auth = metadata_client(
crate::UsmConfig::new("auth-user")
.auth(crate::AuthProtocol::Sha256, "authpassword")
.unwrap()
.into(),
);
assert_eq!(auth.version(), Version::V3);
assert_eq!(auth.security_level(), Some(SecurityLevel::AuthNoPriv));
assert!(auth.inner.salt_counter.is_none());
let auth_priv = metadata_client(
crate::UsmConfig::new("private-user")
.auth_priv(
crate::AuthProtocol::Sha256,
"authpassword",
crate::PrivProtocol::Aes128,
"privpassword",
)
.unwrap()
.into(),
);
assert_eq!(auth_priv.version(), Version::V3);
assert_eq!(auth_priv.security_level(), Some(SecurityLevel::AuthPriv));
assert!(auth_priv.inner.salt_counter.is_some());
}
}
#[cfg(feature = "crypto-rustcrypto")]
#[test]
fn independent_des_clients_require_and_share_caller_state() {
let auth = crate::UsmConfig::new("des-user")
.auth_priv(
crate::AuthProtocol::Sha1,
"auth-password",
crate::PrivProtocol::Des,
"priv-password",
)
.unwrap();
let without_state = Client::new(
TruncatingTransport::new(0),
ClientConfig {
auth: auth.clone().into(),
..ClientConfig::default()
},
);
assert!(matches!(without_state, Err(error) if matches!(*error, Error::Config(_))));
let state =
crate::DesSaltState::install(|_| Ok::<(), std::convert::Infallible>(())).unwrap();
let build = || {
Client::new(
TruncatingTransport::new(0),
ClientConfig {
auth: auth.clone().into(),
des_salt_state: Some(state.clone()),
..ClientConfig::default()
},
)
.unwrap()
};
let first = build();
let second = build();
assert!(first.inner.salt_counter.is_none());
assert!(second.inner.salt_counter.is_none());
assert_eq!(
first
.inner
.config
.des_salt_state
.as_ref()
.unwrap()
.reserve()
.unwrap()
.salt(),
1
);
assert_eq!(
second
.inner
.config
.des_salt_state
.as_ref()
.unwrap()
.reserve()
.unwrap()
.salt(),
2
);
}
#[tokio::test]
async fn client_protocol_metadata_is_transport_independent() {
let udp_transport = crate::UdpTransport::bind("127.0.0.1:0").await.unwrap();
let udp_handle = udp_transport
.handle("127.0.0.1:161".parse().unwrap())
.unwrap();
let udp_client = Client::new(
udp_handle,
ClientConfig {
auth: Auth::v1("private"),
..Default::default()
},
)
.unwrap();
assert_eq!(udp_client.version(), Version::V1);
assert_eq!(udp_client.security_level(), None);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let tcp_transport = crate::TcpTransport::connect(listener.local_addr().unwrap())
.await
.unwrap();
let tcp_client = Client::new(
tcp_transport,
ClientConfig {
auth: Auth::usm("no-auth-user"),
..Default::default()
},
)
.unwrap();
assert_eq!(tcp_client.version(), Version::V3);
assert_eq!(
tcp_client.security_level(),
Some(SecurityLevel::NoAuthNoPriv)
);
}
fn make_client(response_varbind_count: usize) -> Client<TruncatingTransport> {
make_client_with_policy(response_varbind_count, ResponseShapePolicy::Compatible)
}
fn make_client_with_policy(
response_varbind_count: usize,
response_shape_policy: ResponseShapePolicy,
) -> Client<TruncatingTransport> {
let transport = TruncatingTransport::new(response_varbind_count);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
max_oids_per_request: 10,
retry: crate::client::retry::Retry::none(),
response_shape_policy,
..Default::default()
};
Client::new(transport, config).expect("valid client config")
}
#[derive(Clone)]
struct ScriptedResponseTransport {
responses: Arc<Mutex<VecDeque<Vec<VarBind>>>>,
pending: Arc<Mutex<VecDeque<i32>>>,
}
impl ScriptedResponseTransport {
fn new(responses: Vec<Vec<VarBind>>) -> Self {
Self {
responses: Arc::new(Mutex::new(responses.into())),
pending: Arc::new(Mutex::new(VecDeque::new())),
}
}
}
impl ScriptedResponseTransport {
fn recv(
&self,
_registration: crate::transport::RequestRegistration,
) -> impl std::future::Future<Output = Result<(Bytes, SocketAddr)>> + Send {
let request_id = self.pending.lock().unwrap().pop_front().unwrap_or(1);
let varbinds = self
.responses
.lock()
.unwrap()
.pop_front()
.expect("missing scripted response");
async move {
let pdu = Pdu::response(request_id, 0, 0, varbinds);
let message = CommunityMessage::v2c(Bytes::from_static(b"public"), pdu).unwrap();
Ok((message.encode().unwrap(), "127.0.0.1:161".parse().unwrap()))
}
}
}
impl Transport for ScriptedResponseTransport {
fn send(&self, data: &[u8]) -> impl std::future::Future<Output = Result<()>> + Send {
let request_id = crate::transport::extract_request_id(data).unwrap_or(1);
self.pending.lock().unwrap().push_back(request_id);
async { Ok(()) }
}
fn request_with<T, F>(
&self,
data: &[u8],
registration: crate::transport::RequestRegistration,
validate: F,
) -> impl std::future::Future<Output = Result<T>> + Send
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
crate::transport::request_with_scripted(
self,
data,
registration,
move |registration| {
futures_util::stream::once(async move { self.recv(registration).await })
},
validate,
)
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
}
fn scripted_client(
responses: Vec<Vec<VarBind>>,
response_shape_policy: ResponseShapePolicy,
) -> Client<ScriptedResponseTransport> {
Client::new(
ScriptedResponseTransport::new(responses),
ClientConfig {
auth: crate::Auth::v2c("public"),
retry: crate::client::retry::Retry::none(),
response_shape_policy,
..Default::default()
},
)
.expect("valid client config")
}
fn response_shape_error(result: Result<FixedCardinalityResponse>) -> FixedCardinalityResponse {
match *result.expect_err("strict policy must reject the scripted anomaly") {
Error::ResponseShape { response, .. } => response,
ref other => panic!("expected ResponseShape, got {other:?}"),
}
}
#[derive(Clone)]
struct CountingTransport {
sends: Arc<AtomicUsize>,
allocations: Arc<AtomicUsize>,
}
#[derive(Clone)]
struct SendTimeoutProbe {
timeouts: Arc<Mutex<Vec<Duration>>>,
}
impl Transport for SendTimeoutProbe {
async fn send(&self, _data: &[u8]) -> Result<()> {
panic!("client trap sending must use the bounded send contract")
}
async fn send_with_timeout(&self, _data: &[u8], timeout: Duration) -> Result<()> {
self.timeouts.lock().unwrap().push(timeout);
Ok(())
}
async fn request_with<T, F>(
&self,
_data: &[u8],
_registration: crate::transport::RequestRegistration,
_validate: F,
) -> Result<T>
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
panic!("send-timeout probe does not receive responses")
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:162".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
}
#[tokio::test]
async fn v1_v2c_and_v3_traps_use_configured_send_timeout() {
let timeouts = Arc::new(Mutex::new(Vec::new()));
let trap_oid = crate::oid!(1, 3, 6, 1, 6, 3, 1, 1, 5, 1);
let send_timeout = Duration::from_secs(7);
for auth in [crate::Auth::v1("public"), crate::Auth::v2c("public")] {
let client = Client::new(
SendTimeoutProbe {
timeouts: Arc::clone(&timeouts),
},
ClientConfig {
auth,
send_timeout,
..Default::default()
},
)
.unwrap();
client.send_trap(&trap_oid, 0, vec![]).await.unwrap();
}
let authoritative_engine =
crate::AuthoritativeEngine::install(b"test-trap-engine".to_vec(), |_| {
Ok::<(), std::convert::Infallible>(())
})
.unwrap();
let v3_client = Client::new(
SendTimeoutProbe {
timeouts: Arc::clone(&timeouts),
},
ClientConfig {
auth: crate::Auth::usm("trapuser"),
send_timeout,
local_authoritative_engine: Some(authoritative_engine),
..Default::default()
},
)
.unwrap();
v3_client.send_trap(&trap_oid, 0, vec![]).await.unwrap();
assert_eq!(*timeouts.lock().unwrap(), [send_timeout; 3]);
}
#[derive(Clone)]
struct CapacityTransport {
capacity: usize,
requests: Arc<AtomicUsize>,
}
impl Transport for CapacityTransport {
async fn send(&self, _data: &[u8]) -> Result<()> {
Ok(())
}
async fn request_with<U, F>(
&self,
_data: &[u8],
_registration: crate::transport::RequestRegistration,
_validate: F,
) -> Result<U>
where
U: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<U>> + Send,
{
self.requests.fetch_add(1, Ordering::Relaxed);
Err(Error::Config("capacity boundary reached transport".into()).boxed())
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
fn send_capacity(&self) -> usize {
self.capacity
}
}
#[tokio::test]
async fn community_atomic_requests_enforce_exact_transport_boundary_before_send() {
for (community_version, auth) in [
(CommunityVersion::V1, crate::Auth::v1("public")),
(CommunityVersion::V2c, crate::Auth::v2c("public")),
] {
let pdu = Pdu::get_request(7, &[crate::oid!(1, 3, 6, 1, 2, 1, 1, 1, 0)]);
let exact_size = CommunityMessage::new(
community_version,
Bytes::from_static(b"public"),
pdu.clone(),
)
.unwrap()
.encode()
.unwrap()
.len();
let exact_requests = Arc::new(AtomicUsize::new(0));
let exact_client = Client::new(
CapacityTransport {
capacity: exact_size,
requests: Arc::clone(&exact_requests),
},
ClientConfig {
auth: auth.clone(),
retry: crate::client::retry::Retry::none(),
..Default::default()
},
)
.unwrap();
let error = exact_client.send_request(pdu.clone()).await.unwrap_err();
assert!(matches!(*error, Error::Config(_)));
assert_eq!(exact_requests.load(Ordering::Relaxed), 1);
let oversized_requests = Arc::new(AtomicUsize::new(0));
let oversized_client = Client::new(
CapacityTransport {
capacity: exact_size - 1,
requests: Arc::clone(&oversized_requests),
},
ClientConfig {
auth,
retry: crate::client::retry::Retry::none(),
..Default::default()
},
)
.unwrap();
let error = oversized_client.send_request(pdu).await.unwrap_err();
assert!(matches!(
*error,
Error::OutboundMessageTooLarge { size, limit }
if size == exact_size && limit == exact_size - 1
));
assert_eq!(oversized_requests.load(Ordering::Relaxed), 0);
}
}
impl Transport for CountingTransport {
fn send(&self, _data: &[u8]) -> impl std::future::Future<Output = Result<()>> + Send {
self.sends.fetch_add(1, Ordering::Relaxed);
async { Ok(()) }
}
async fn request_with<T, F>(
&self,
_data: &[u8],
_registration: crate::transport::RequestRegistration,
_validate: F,
) -> Result<T>
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
panic!("receive must not be reached after encode failure")
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn alloc_request_id(&self) -> i32 {
self.allocations.fetch_add(1, Ordering::Relaxed);
1
}
fn is_reliable(&self) -> bool {
true
}
}
#[tokio::test]
async fn get_bulk_parameter_ranges_are_checked_before_request_side_effects() {
let client = make_client(0);
for (non_repeaters, max_repetitions) in [
(0, 0),
(crate::pdu::MAX_GET_BULK_VALUE, 0),
(0, crate::pdu::MAX_GET_BULK_VALUE),
] {
assert!(
client
.get_bulk(&[], non_repeaters, max_repetitions)
.await
.is_ok()
);
}
let sends = Arc::new(AtomicUsize::new(0));
let allocations = Arc::new(AtomicUsize::new(0));
let transport = CountingTransport {
sends: Arc::clone(&sends),
allocations: Arc::clone(&allocations),
};
let client = Client::new(
transport,
ClientConfig {
auth: crate::Auth::v2c("public"),
retry: crate::client::retry::Retry::none(),
..Default::default()
},
)
.unwrap();
for (non_repeaters, max_repetitions) in [
(crate::pdu::MAX_GET_BULK_VALUE + 1, 0),
(0, crate::pdu::MAX_GET_BULK_VALUE + 1),
] {
let error = client
.get_bulk(&[], non_repeaters, max_repetitions)
.await
.unwrap_err();
assert!(matches!(*error, Error::InvalidMessage(_)));
}
assert_eq!(allocations.load(Ordering::Relaxed), 0);
assert_eq!(sends.load(Ordering::Relaxed), 0);
}
#[test]
fn bulk_walk_parameter_range_is_checked_before_stream_construction() {
let sends = Arc::new(AtomicUsize::new(0));
let allocations = Arc::new(AtomicUsize::new(0));
let transport = CountingTransport {
sends: Arc::clone(&sends),
allocations: Arc::clone(&allocations),
};
let client = Client::new(transport, ClientConfig::default()).unwrap();
let base = Oid::from_slice(&[1, 3, 6, 1]);
assert!(client.bulk_walk(base.clone(), 0).is_ok());
assert!(
client
.bulk_walk(base.clone(), crate::pdu::MAX_GET_BULK_VALUE)
.is_ok()
);
let error = client
.bulk_walk(base, crate::pdu::MAX_GET_BULK_VALUE + 1)
.err()
.expect("out-of-range bulk walk must not return a stream");
assert!(matches!(*error, Error::InvalidMessage(_)));
assert_eq!(allocations.load(Ordering::Relaxed), 0);
assert_eq!(sends.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn invalid_oid_is_not_sent() {
fn assert_invalid<T>(result: Result<T>) {
match result {
Err(error) => assert!(
matches!(&*error, Error::InvalidOid(_)),
"expected InvalidOid, got {error:?}"
),
Ok(_) => panic!("invalid OID operation succeeded"),
}
}
let sends = Arc::new(AtomicUsize::new(0));
let transport = CountingTransport {
sends: Arc::clone(&sends),
allocations: Arc::new(AtomicUsize::new(0)),
};
let client = Client::new(
transport.clone(),
ClientConfig {
auth: crate::Auth::v2c("public"),
retry: crate::client::retry::Retry::none(),
..Default::default()
},
)
.expect("valid client config");
let invalid = Oid::empty();
let valid = Oid::from_slice(&[1, 3, 6, 1]);
assert_invalid(client.get(&invalid).await);
assert_invalid(client.get_next(&invalid).await);
assert_invalid(client.get_bulk(std::slice::from_ref(&invalid), 0, 10).await);
assert_invalid(client.set(&invalid, Value::Integer(1)).await);
assert_invalid(
client
.set(&valid, Value::ObjectIdentifier(invalid.clone()))
.await,
);
assert_invalid(client.send_trap(&invalid, 1, vec![]).await);
assert_invalid(
client
.send_trap(&valid, 1, vec![VarBind::null(invalid.clone())])
.await,
);
assert_invalid(client.send_inform(&invalid, 1, vec![]).await);
assert_invalid(TrapV1Notification::new(
invalid.clone(),
[127, 0, 0, 1],
crate::pdu::GenericTrap::EnterpriseSpecific,
1,
1,
vec![],
));
let v3_client = Client::new(
transport,
ClientConfig {
auth: crate::Auth::Usm(crate::v3::UsmConfig::new("user")),
retry: crate::client::retry::Retry::none(),
..Default::default()
},
)
.expect("valid client config");
assert_invalid(v3_client.get(&invalid).await);
assert_invalid(v3_client.get_next(&invalid).await);
assert_invalid(
v3_client
.get_bulk(std::slice::from_ref(&invalid), 0, 10)
.await,
);
assert_invalid(v3_client.set(&invalid, Value::Integer(1)).await);
assert_invalid(
v3_client
.set(&valid, Value::ObjectIdentifier(invalid.clone()))
.await,
);
assert_invalid(v3_client.send_inform(&invalid, 1, vec![]).await);
assert_eq!(sends.load(Ordering::Relaxed), 0);
}
#[test]
fn client_config_validation() {
fn assert_config_error<T: Transport>(result: Result<Client<T>>) {
match result {
Err(error) => assert!(matches!(*error, Error::Config(_))),
Ok(_) => panic!("invalid client configuration was accepted"),
}
}
let sends = Arc::new(AtomicUsize::new(0));
let transport = CountingTransport {
sends: Arc::clone(&sends),
allocations: Arc::new(AtomicUsize::new(0)),
};
assert_config_error(Client::new(
transport.clone(),
ClientConfig {
request_timeout: Duration::MAX,
..ClientConfig::default()
},
));
assert_config_error(Client::new(
transport.clone(),
ClientConfig {
send_timeout: Duration::MAX,
..ClientConfig::default()
},
));
assert_config_error(Client::new(
transport.clone(),
ClientConfig {
max_oids_per_request: 0,
..ClientConfig::default()
},
));
assert_config_error(Client::new(
transport.clone(),
ClientConfig {
walk_options: WalkOptions {
max_repetitions: crate::pdu::MAX_GET_BULK_VALUE + 1,
..WalkOptions::default()
},
..ClientConfig::default()
},
));
assert_config_error(Client::new(
transport.clone(),
ClientConfig {
auth: Auth::v1("public"),
walk_options: WalkOptions {
method: WalkMethod::GetBulk,
..WalkOptions::default()
},
..ClientConfig::default()
},
));
assert_config_error(Client::new(
transport.clone(),
ClientConfig {
walk_options: WalkOptions {
ordering: OidOrdering::AllowNonIncreasing,
result_limit: None,
..WalkOptions::default()
},
..ClientConfig::default()
},
));
assert_config_error(Client::with_engine_cache(
transport.clone(),
ClientConfig {
max_oids_per_request: 0,
..ClientConfig::default()
},
Arc::new(EngineCache::new()),
));
Client::new(
transport.clone(),
ClientConfig {
request_timeout: Duration::ZERO,
..ClientConfig::default()
},
)
.expect("zero timeout remains an explicit immediate deadline");
Client::new(
transport.clone(),
ClientConfig {
send_timeout: Duration::ZERO,
..ClientConfig::default()
},
)
.expect("zero send timeout remains an explicit immediate deadline");
for auth in [
Auth::v1("public"),
Auth::v2c("public"),
Auth::Usm(UsmConfig::new("user")),
] {
Client::new(
transport.clone(),
ClientConfig {
auth,
..ClientConfig::default()
},
)
.expect("valid v1, v2c, and v3 configs must construct");
}
assert_eq!(sends.load(Ordering::Relaxed), 0);
}
#[tokio::test]
async fn single_operations_preserve_empty_and_excess_responses() {
let oid = Oid::from_slice(&[1, 3, 6, 1, 1]);
for response_count in [0, 2] {
let get = make_client(response_count).get(&oid).await.unwrap();
let get_next = make_client(response_count).get_next(&oid).await.unwrap();
let set = make_client(response_count)
.set(&oid, Value::Integer(1))
.await
.unwrap();
for response in [get, get_next, set] {
assert_eq!(response.varbinds.len(), response_count);
if response_count == 0 {
assert!(matches!(
response.anomalies[0],
ResponseShapeAnomaly::Truncated { .. }
));
} else {
assert!(matches!(
response.anomalies[0],
ResponseShapeAnomaly::Excess { .. }
));
}
}
}
}
#[tokio::test]
async fn strict_policy_returns_observable_shape_error() {
let oid = Oid::from_slice(&[1, 3, 6, 1, 1]);
let error = make_client_with_policy(2, ResponseShapePolicy::Strict)
.get(&oid)
.await
.unwrap_err();
match *error {
Error::ResponseShape { response, .. } => {
assert_eq!(response.varbinds.len(), 2);
assert!(matches!(
response.anomalies[0],
ResponseShapeAnomaly::Excess { .. }
));
}
other => panic!("expected ResponseShape, got {other:?}"),
}
}
#[tokio::test]
async fn inform_requires_exact_echoed_varbinds() {
let trap_oid = Oid::from_slice(&[1, 3, 6, 1, 6, 3, 1, 1, 5, 1]);
let additional = VarBind::new(
Oid::from_slice(&[1, 3, 6, 1, 4, 1, 9999, 1]),
Value::Integer(7),
);
let expected = Pdu::inform_request(1, 123, &trap_oid, vec![additional.clone()]).varbinds;
for policy in [ResponseShapePolicy::Compatible, ResponseShapePolicy::Strict] {
scripted_client(vec![expected.clone()], policy)
.send_inform(&trap_oid, 123, vec![additional.clone()])
.await
.expect("an exact Inform echo must be accepted");
let mut renamed = expected.clone();
renamed[0].oid = Oid::from_slice(&[1, 3, 6, 1, 2, 1, 1, 4, 0]);
let mut reordered = expected.clone();
reordered.swap(0, 1);
let mut changed = expected.clone();
changed[2].value = Value::Integer(8);
for response in [Vec::new(), renamed, reordered, changed] {
let error = scripted_client(vec![response], policy)
.send_inform(&trap_oid, 123, vec![additional.clone()])
.await
.expect_err("a malformed Inform acknowledgement must be rejected");
assert!(matches!(*error, Error::MalformedResponse { .. }));
}
}
}
#[tokio::test]
async fn compatible_policy_preserves_scripted_semantics_across_all_fixed_operations() {
let a = Oid::from_slice(&[1, 3, 6, 1, 1]);
let b = Oid::from_slice(&[1, 3, 6, 1, 2]);
let c = Oid::from_slice(&[1, 3, 6, 1, 3]);
let responses = vec![
vec![VarBind::new(c.clone(), Value::Integer(10))],
vec![
VarBind::new(b.clone(), Value::Integer(20)),
VarBind::new(a.clone(), Value::Integer(10)),
],
vec![VarBind::new(a.clone(), Value::Integer(10))],
vec![
VarBind::new(b.clone(), Value::Integer(20)),
VarBind::new(b.clone(), Value::EndOfMibView),
],
vec![VarBind::new(a.clone(), Value::Integer(2))],
vec![
VarBind::new(a.clone(), Value::Integer(1)),
VarBind::new(b.clone(), Value::Integer(3)),
],
];
let client = scripted_client(responses.clone(), ResponseShapePolicy::Compatible);
let get = client.get(&a).await.unwrap();
assert_eq!(get.varbinds, responses[0]);
assert!(matches!(
get.anomalies.as_slice(),
[ResponseShapeAnomaly::OidMismatch { .. }]
));
let get_many = client.get_many(&[a.clone(), b.clone()]).await.unwrap();
assert_eq!(get_many.varbinds, responses[1]);
assert!(matches!(
get_many.anomalies.as_slice(),
[ResponseShapeAnomaly::Reordered { .. }]
));
let get_next = client.get_next(&a).await.unwrap();
assert_eq!(get_next.varbinds, responses[2]);
assert!(matches!(
get_next.anomalies.as_slice(),
[ResponseShapeAnomaly::GetNextNotSuccessor { .. }]
));
let get_next_many = client.get_next_many(&[a.clone(), b.clone()]).await.unwrap();
assert_eq!(get_next_many.varbinds, responses[3]);
assert!(get_next_many.anomalies.is_empty());
let set = client.set(&a, Value::Integer(1)).await.unwrap();
assert_eq!(set.varbinds, responses[4]);
assert!(matches!(
set.anomalies.as_slice(),
[ResponseShapeAnomaly::SetValueMismatch { .. }]
));
let set_many = client
.set_many(&[
(a.clone(), Value::Integer(1)),
(b.clone(), Value::Integer(2)),
])
.await
.unwrap();
assert_eq!(set_many.varbinds, responses[5]);
assert!(matches!(
set_many.anomalies.as_slice(),
[ResponseShapeAnomaly::SetValueMismatch { .. }]
));
}
#[tokio::test]
async fn strict_policy_retains_scripted_evidence_across_all_fixed_operations() {
let a = Oid::from_slice(&[1, 3, 6, 1, 1]);
let b = Oid::from_slice(&[1, 3, 6, 1, 2]);
let c = Oid::from_slice(&[1, 3, 6, 1, 3]);
let responses = vec![
vec![VarBind::new(c.clone(), Value::Integer(10))],
vec![
VarBind::new(b.clone(), Value::Integer(20)),
VarBind::new(a.clone(), Value::Integer(10)),
],
vec![VarBind::new(a.clone(), Value::Integer(10))],
vec![
VarBind::new(b.clone(), Value::Integer(20)),
VarBind::new(c.clone(), Value::EndOfMibView),
],
vec![VarBind::new(a.clone(), Value::Integer(2))],
vec![
VarBind::new(a.clone(), Value::Integer(1)),
VarBind::new(b.clone(), Value::Integer(3)),
],
];
let client = scripted_client(responses.clone(), ResponseShapePolicy::Strict);
let errors = [
response_shape_error(client.get(&a).await),
response_shape_error(client.get_many(&[a.clone(), b.clone()]).await),
response_shape_error(client.get_next(&a).await),
response_shape_error(client.get_next_many(&[a.clone(), b.clone()]).await),
response_shape_error(client.set(&a, Value::Integer(1)).await),
response_shape_error(
client
.set_many(&[
(a.clone(), Value::Integer(1)),
(b.clone(), Value::Integer(2)),
])
.await,
),
];
for (response, expected) in errors.iter().zip(responses) {
assert_eq!(response.varbinds, expected);
assert!(!response.anomalies.is_empty());
}
assert!(matches!(
errors[3].anomalies.as_slice(),
[ResponseShapeAnomaly::GetNextEndOfMibNameMismatch { .. }]
));
}
#[tokio::test]
async fn walk_does_not_consume_an_anomalous_single_response() {
let oid = Oid::from_slice(&[1, 3, 6, 1, 1]);
let mut walk = make_client(2).walk_getnext(oid).unwrap();
let error = walk.next().await.unwrap().unwrap_err();
assert!(matches!(*error, Error::ResponseShape { .. }));
assert!(walk.next().await.is_none());
}
#[tokio::test]
async fn get_many_preserves_truncated_response() {
let client = make_client(1);
let oids = [
Oid::from_slice(&[1, 3, 6, 1, 1]),
Oid::from_slice(&[1, 3, 6, 1, 2]),
Oid::from_slice(&[1, 3, 6, 1, 3]),
];
let response = client.get_many(&oids).await.unwrap();
assert_eq!(response.varbinds.len(), 1);
assert!(matches!(
response.anomalies.as_slice(),
[ResponseShapeAnomaly::Truncated { .. }]
));
}
#[tokio::test]
async fn get_many_preserves_inflated_response() {
let client = make_client(5);
let oids = [
Oid::from_slice(&[1, 3, 6, 1, 1]),
Oid::from_slice(&[1, 3, 6, 1, 2]),
Oid::from_slice(&[1, 3, 6, 1, 3]),
];
let response = client.get_many(&oids).await.unwrap();
assert_eq!(response.varbinds.len(), 5);
assert!(matches!(
response.anomalies.as_slice(),
[ResponseShapeAnomaly::Excess { .. }]
));
}
#[tokio::test]
async fn get_many_accepts_correct_response_count() {
let client = make_client(3);
let oids = [
Oid::from_slice(&[1, 3, 6, 1, 1]),
Oid::from_slice(&[1, 3, 6, 1, 2]),
Oid::from_slice(&[1, 3, 6, 1, 3]),
];
let result = client.get_many(&oids).await;
assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
assert_eq!(result.unwrap().varbinds.len(), 3);
}
#[tokio::test]
async fn get_next_many_preserves_truncated_response() {
let client = make_client(1);
let oids = [
Oid::from_slice(&[1, 3, 6, 1, 1]),
Oid::from_slice(&[1, 3, 6, 1, 2]),
Oid::from_slice(&[1, 3, 6, 1, 3]),
];
let response = client.get_next_many(&oids).await.unwrap();
assert_eq!(response.varbinds.len(), 1);
assert!(matches!(
response.anomalies.as_slice(),
[ResponseShapeAnomaly::Truncated { .. }]
));
}
#[tokio::test]
async fn get_next_many_preserves_inflated_response() {
let client = make_client(5);
let oids = [
Oid::from_slice(&[1, 3, 6, 1, 1]),
Oid::from_slice(&[1, 3, 6, 1, 2]),
Oid::from_slice(&[1, 3, 6, 1, 3]),
];
let response = client.get_next_many(&oids).await.unwrap();
assert_eq!(response.varbinds.len(), 5);
assert!(matches!(
response.anomalies.as_slice(),
[ResponseShapeAnomaly::Excess { .. }]
));
}
#[tokio::test]
async fn get_next_many_accepts_correct_response_count() {
let client = make_client(3);
let oids = [
Oid::from_slice(&[1, 3, 6, 1, 1]),
Oid::from_slice(&[1, 3, 6, 1, 2]),
Oid::from_slice(&[1, 3, 6, 1, 3]),
];
let result = client.get_next_many(&oids).await;
assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
assert_eq!(result.unwrap().varbinds.len(), 3);
}
#[tokio::test]
async fn set_many_preserves_truncated_response() {
let client = make_client(1);
let varbinds = [
(
Oid::from_slice(&[1, 3, 6, 1, 1]),
crate::value::Value::Integer(1),
),
(
Oid::from_slice(&[1, 3, 6, 1, 2]),
crate::value::Value::Integer(2),
),
(
Oid::from_slice(&[1, 3, 6, 1, 3]),
crate::value::Value::Integer(3),
),
];
let result = client.set_many(&varbinds).await;
assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
assert_eq!(result.unwrap().varbinds.len(), 1);
}
#[tokio::test]
async fn set_many_preserves_inflated_response() {
let client = make_client(5);
let varbinds = [
(
Oid::from_slice(&[1, 3, 6, 1, 1]),
crate::value::Value::Integer(1),
),
(
Oid::from_slice(&[1, 3, 6, 1, 2]),
crate::value::Value::Integer(2),
),
(
Oid::from_slice(&[1, 3, 6, 1, 3]),
crate::value::Value::Integer(3),
),
];
let response = client.set_many(&varbinds).await.unwrap();
assert_eq!(response.varbinds.len(), 5);
assert!(matches!(
response.anomalies.as_slice(),
[ResponseShapeAnomaly::Excess { .. }]
));
}
#[tokio::test]
async fn set_many_accepts_correct_response_count() {
let client = make_client(3);
let varbinds = [
(
Oid::from_slice(&[1, 3, 6, 1, 1]),
crate::value::Value::Integer(1),
),
(
Oid::from_slice(&[1, 3, 6, 1, 2]),
crate::value::Value::Integer(2),
),
(
Oid::from_slice(&[1, 3, 6, 1, 3]),
crate::value::Value::Integer(3),
),
];
let result = client.set_many(&varbinds).await;
assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
assert_eq!(result.unwrap().varbinds.len(), 3);
}
#[derive(Clone)]
struct TooBigTransport {
max_varbinds: usize,
pending: Arc<Mutex<VecDeque<(i32, usize)>>>,
}
impl TooBigTransport {
fn new(max_varbinds: usize) -> Self {
Self {
max_varbinds,
pending: Arc::new(Mutex::new(VecDeque::new())),
}
}
}
impl TooBigTransport {
fn recv(
&self,
_registration: crate::transport::RequestRegistration,
) -> impl std::future::Future<Output = Result<(Bytes, SocketAddr)>> + Send {
let (request_id, varbind_count) = {
let mut q = self.pending.lock().unwrap();
q.pop_front().unwrap_or((1, 0))
};
let max = self.max_varbinds;
let peer: SocketAddr = "127.0.0.1:161".parse().unwrap();
async move {
let pdu = if varbind_count > max {
Pdu::response(request_id, ErrorStatus::TooBig.as_i32(), 0, vec![])
} else {
let varbinds: Vec<VarBind> = (0..varbind_count)
.map(|i| {
VarBind::new(
Oid::from_slice(&[1, 3, 6, 1, i as u32]),
crate::value::Value::Integer(i as i32),
)
})
.collect();
Pdu::response(request_id, 0, 0, varbinds)
};
let msg = CommunityMessage::v2c(Bytes::from_static(b"public"), pdu).unwrap();
Ok((msg.encode().unwrap(), peer))
}
}
}
impl Transport for TooBigTransport {
fn send(&self, data: &[u8]) -> impl std::future::Future<Output = Result<()>> + Send {
let request_id = crate::transport::extract_request_id(data).unwrap_or(1);
let msg = CommunityMessage::decode(
Bytes::copy_from_slice(data),
crate::DecodeConfig::default(),
)
.unwrap()
.value;
let varbind_count = msg.pdu().standard().unwrap().varbinds.len();
{
let mut q = self.pending.lock().unwrap();
q.push_back((request_id, varbind_count));
}
async { Ok(()) }
}
fn request_with<T, F>(
&self,
data: &[u8],
registration: crate::transport::RequestRegistration,
validate: F,
) -> impl std::future::Future<Output = Result<T>> + Send
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
crate::transport::request_with_scripted(
self,
data,
registration,
move |registration| {
futures_util::stream::once(async move { self.recv(registration).await })
},
validate,
)
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
}
#[tokio::test]
async fn inform_preserves_empty_varbind_too_big_response() {
let client = Client::new(
TooBigTransport::new(0),
ClientConfig {
auth: crate::Auth::v2c("public"),
retry: Retry::none(),
..Default::default()
},
)
.expect("valid client config");
let trap_oid = Oid::from_slice(&[1, 3, 6, 1, 6, 3, 1, 1, 5, 1]);
let error = client
.send_inform(&trap_oid, 123, Vec::new())
.await
.expect_err("tooBig must remain an SNMP protocol error");
assert!(matches!(
*error,
Error::Snmp {
status: ErrorStatus::TooBig,
..
}
));
}
#[derive(Clone)]
struct InformMetadataTransport {
pending: Arc<Mutex<VecDeque<Pdu>>>,
malformed_echo: bool,
}
impl InformMetadataTransport {
fn recv(
&self,
_registration: crate::transport::RequestRegistration,
) -> impl std::future::Future<Output = Result<(Bytes, SocketAddr)>> + Send {
let request = self.pending.lock().unwrap().pop_front().unwrap();
let malformed_echo = self.malformed_echo;
async move {
let mut varbinds = request.varbinds;
if malformed_echo {
varbinds.pop();
}
let response = Pdu::response(request.request_id, 0, 0, varbinds);
let message =
CommunityMessage::v2c(Bytes::from_static(b"public"), response).unwrap();
let mut encoded = message.encode().unwrap().to_vec();
encoded.extend_from_slice(&[0xaa, 0xbb]);
Ok((Bytes::from(encoded), "127.0.0.1:161".parse().unwrap()))
}
}
}
impl Transport for InformMetadataTransport {
fn send(&self, data: &[u8]) -> impl std::future::Future<Output = Result<()>> + Send {
let message = CommunityMessage::decode(
Bytes::copy_from_slice(data),
crate::DecodeConfig::default(),
)
.unwrap()
.value;
self.pending
.lock()
.unwrap()
.push_back(message.pdu().standard().unwrap().clone());
async { Ok(()) }
}
fn request_with<T, F>(
&self,
data: &[u8],
registration: crate::transport::RequestRegistration,
validate: F,
) -> impl std::future::Future<Output = Result<T>> + Send
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
crate::transport::request_with_scripted(
self,
data,
registration,
move |registration| {
futures_util::stream::once(async move { self.recv(registration).await })
},
validate,
)
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
}
#[tokio::test]
async fn inform_metadata_api_retains_acknowledgement_anomalies() {
let client = Client::new(
InformMetadataTransport {
pending: Arc::new(Mutex::new(VecDeque::new())),
malformed_echo: false,
},
ClientConfig {
auth: crate::Auth::v2c("public"),
retry: Retry::none(),
..Default::default()
},
)
.unwrap();
let metadata = client
.send_inform_with_metadata(
&Oid::from_slice(&[1, 3, 6, 1, 6, 3, 1, 1, 5, 1]),
123,
vec![],
)
.await
.unwrap();
assert_eq!(
metadata.decode_anomalies,
vec![crate::DecodeAnomaly::TrailingBytes {
original_length: 2,
canonical_length: 0,
}]
);
}
#[tokio::test]
async fn malformed_inform_acknowledgement_retains_decode_anomalies() {
let client = Client::new(
InformMetadataTransport {
pending: Arc::new(Mutex::new(VecDeque::new())),
malformed_echo: true,
},
ClientConfig {
auth: crate::Auth::v2c("public"),
retry: Retry::none(),
..Default::default()
},
)
.unwrap();
let error = client
.send_inform_with_metadata(
&Oid::from_slice(&[1, 3, 6, 1, 6, 3, 1, 1, 5, 1]),
123,
vec![],
)
.await
.expect_err("a malformed acknowledgement must be rejected");
assert_eq!(error.kind(), crate::ErrorKind::MalformedResponse);
assert_eq!(
error.response_metadata().unwrap().decode_anomalies,
vec![crate::DecodeAnomaly::TrailingBytes {
original_length: 2,
canonical_length: 0,
}]
);
}
#[tokio::test]
async fn get_many_bisects_on_too_big() {
let transport = TooBigTransport::new(3);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
max_oids_per_request: 10,
retry: crate::client::retry::Retry::none(),
..Default::default()
};
let client = Client::new(transport, config).expect("valid client config");
let oids: Vec<Oid> = (0..8u32)
.map(|i| Oid::from_slice(&[1, 3, 6, 1, i]))
.collect();
let result = client.get_many(&oids).await.unwrap();
assert_eq!(result.varbinds.len(), 8);
}
#[tokio::test]
async fn get_many_single_oid_too_big_is_unrecoverable() {
let transport = TooBigTransport::new(0);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
max_oids_per_request: 10,
retry: crate::client::retry::Retry::none(),
..Default::default()
};
let client = Client::new(transport, config).expect("valid client config");
let oids = [Oid::from_slice(&[1, 3, 6, 1, 1])];
let err = client.get_many(&oids).await.unwrap_err();
assert!(
matches!(
&*err,
Error::Snmp {
status: ErrorStatus::TooBig,
..
}
),
"expected TooBig, got: {err}"
);
}
#[tokio::test]
async fn get_next_many_bisects_on_too_big() {
let transport = TooBigTransport::new(3);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
max_oids_per_request: 10,
retry: crate::client::retry::Retry::none(),
..Default::default()
};
let client = Client::new(transport, config).expect("valid client config");
let oids: Vec<Oid> = (0..8u32)
.map(|i| Oid::from_slice(&[1, 3, 6, 1, i]))
.collect();
let result = client.get_next_many(&oids).await.unwrap();
assert_eq!(result.varbinds.len(), 8);
}
#[tokio::test]
async fn get_many_batched_preserves_truncated_response_offsets() {
let transport = TruncatingTransport::new(1);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
max_oids_per_request: 10,
retry: crate::client::retry::Retry::none(),
..Default::default()
};
let client = Client::new(transport, config).expect("valid client config");
let oids: Vec<Oid> = (0..12u32)
.map(|i| Oid::from_slice(&[1, 3, 6, 1, i]))
.collect();
let response = client.get_many(&oids).await.unwrap();
assert_eq!(response.varbinds.len(), 2);
assert!(matches!(
response.anomalies.as_slice(),
[
ResponseShapeAnomaly::Truncated { request_range, response_range, .. },
ResponseShapeAnomaly::Truncated { request_range: second_request, response_range: second_response, .. }
] if request_range == &(0..10)
&& response_range == &(0..1)
&& second_request == &(10..12)
&& second_response == &(1..2)
));
}
#[tokio::test]
async fn get_many_batched_preserves_inflated_response_offsets() {
let transport = TruncatingTransport::new(12);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
max_oids_per_request: 10,
retry: crate::client::retry::Retry::none(),
..Default::default()
};
let client = Client::new(transport, config).expect("valid client config");
let oids: Vec<Oid> = (0..12u32)
.map(|i| Oid::from_slice(&[1, 3, 6, 1, i]))
.collect();
let response = client.get_many(&oids).await.unwrap();
assert_eq!(response.varbinds.len(), 24);
assert!(matches!(
response.anomalies.as_slice(),
[
ResponseShapeAnomaly::Excess { request_range, response_range, .. },
ResponseShapeAnomaly::Excess { request_range: second_request, response_range: second_response, .. }
] if request_range == &(0..10)
&& response_range == &(0..12)
&& second_request == &(10..12)
&& second_response == &(12..24)
));
}
#[derive(Clone)]
struct AdversarialTransport {
pdu_type: PduType,
community: &'static [u8],
respond_as_v1: bool,
pending: Arc<Mutex<VecDeque<i32>>>,
}
impl AdversarialTransport {
fn new(pdu_type: PduType, community: &'static [u8], respond_as_v1: bool) -> Self {
Self {
pdu_type,
community,
respond_as_v1,
pending: Arc::new(Mutex::new(VecDeque::new())),
}
}
}
impl AdversarialTransport {
fn recv(
&self,
_registration: crate::transport::RequestRegistration,
) -> impl std::future::Future<Output = Result<(Bytes, SocketAddr)>> + Send {
let request_id = self.pending.lock().unwrap().pop_front().unwrap_or(1);
let peer: SocketAddr = "127.0.0.1:161".parse().unwrap();
let pdu = Pdu::standard(
crate::pdu::StandardPduType::try_from(self.pdu_type).unwrap(),
request_id,
0,
0,
vec![VarBind::new(
Oid::from_slice(&[1, 3, 6, 1, 1]),
crate::value::Value::Null,
)],
);
let community = Bytes::from_static(self.community);
let msg = if self.respond_as_v1 {
CommunityMessage::v1(community, pdu)
} else {
CommunityMessage::v2c(community, pdu)
}
.unwrap();
let encoded = msg.encode().unwrap();
async move { Ok((encoded, peer)) }
}
}
impl Transport for AdversarialTransport {
fn send(&self, data: &[u8]) -> impl std::future::Future<Output = Result<()>> + Send {
let request_id = crate::transport::extract_request_id(data).unwrap_or(1);
self.pending.lock().unwrap().push_back(request_id);
async { Ok(()) }
}
fn request_with<T, F>(
&self,
data: &[u8],
registration: crate::transport::RequestRegistration,
validate: F,
) -> impl std::future::Future<Output = Result<T>> + Send
where
T: Send,
F: FnMut(Bytes, SocketAddr) -> Result<crate::transport::Candidate<T>> + Send,
{
crate::transport::request_with_scripted(
self,
data,
registration,
move |registration| {
futures_util::stream::once(async move { self.recv(registration).await })
},
validate,
)
}
fn peer_addr(&self) -> SocketAddr {
"127.0.0.1:161".parse().unwrap()
}
fn local_addr(&self) -> SocketAddr {
"127.0.0.1:0".parse().unwrap()
}
fn is_reliable(&self) -> bool {
true
}
}
fn adversarial_client(
pdu_type: PduType,
community: &'static [u8],
respond_as_v1: bool,
) -> Client<AdversarialTransport> {
let transport = AdversarialTransport::new(pdu_type, community, respond_as_v1);
let config = ClientConfig {
auth: crate::Auth::v2c("public"),
retry: crate::client::retry::Retry::none(),
..Default::default()
};
Client::new(transport, config).expect("valid client config")
}
#[tokio::test]
async fn response_validation_accepts_well_formed_response() {
let client = adversarial_client(PduType::Response, b"public", false);
let result = client.get(&Oid::from_slice(&[1, 3, 6, 1, 1])).await;
assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
}
#[tokio::test(start_paused = true)]
async fn response_validation_rejects_echoed_request_pdu() {
let client = adversarial_client(PduType::GetRequest, b"public", false);
let err = client
.get(&Oid::from_slice(&[1, 3, 6, 1, 1]))
.await
.unwrap_err();
assert!(
matches!(*err, Error::Timeout { .. }),
"expected Timeout after the rejected candidate, got: {err}"
);
}
#[tokio::test(start_paused = true)]
async fn response_validation_rejects_community_mismatch() {
let client = adversarial_client(PduType::Response, b"other", false);
let err = client
.get(&Oid::from_slice(&[1, 3, 6, 1, 1]))
.await
.unwrap_err();
assert!(matches!(*err, Error::Timeout { .. }));
}
#[tokio::test]
async fn response_validation_accepts_explicit_any_source_rewrite() {
let transport = AdversarialTransport::new(PduType::Response, b"other", false);
let config = ClientConfig {
community_response_policy:
crate::transport::CommunityResponsePolicy::AllowMismatchFromAnySource,
retry: crate::client::retry::Retry::none(),
..Default::default()
};
let client = Client::new(transport, config).expect("valid client config");
let result = client.get(&Oid::from_slice(&[1, 3, 6, 1, 1])).await;
assert!(result.is_ok(), "expected Ok, got: {:?}", result.err());
}
#[tokio::test(start_paused = true)]
async fn response_validation_rejects_version_mismatch() {
let client = adversarial_client(PduType::Response, b"public", true);
let err = client
.get(&Oid::from_slice(&[1, 3, 6, 1, 1]))
.await
.unwrap_err();
assert!(
matches!(*err, Error::Timeout { .. }),
"expected Timeout after the rejected candidate, got: {err}"
);
}
}