use crate::security::plain::PlainAuthHandler;
use crate::security::zap::{ZAP_VERSION, ZapMechanism, ZapRequest, ZapResponse};
use crate::{DealerSocket, inproc_stream::InprocStream};
use monocoque_core::options::SocketOptions;
use std::io;
use std::sync::Arc;
#[async_trait::async_trait(?Send)]
pub trait ZapHandler {
async fn authenticate(&self, request: &ZapRequest) -> ZapResponse;
}
pub struct DefaultZapHandler<H: PlainAuthHandler> {
plain_handler: Arc<H>,
accept_curve: bool,
curve_key_whitelist: Option<std::collections::HashSet<[u8; 32]>>,
}
impl<H: PlainAuthHandler> DefaultZapHandler<H> {
pub const fn new(plain_handler: Arc<H>, accept_curve: bool) -> Self {
Self {
plain_handler,
accept_curve,
curve_key_whitelist: None,
}
}
pub fn with_curve_whitelist(mut self, keys: Vec<[u8; 32]>) -> Self {
self.curve_key_whitelist = Some(keys.into_iter().collect());
self
}
}
#[async_trait::async_trait(?Send)]
impl<H: PlainAuthHandler> ZapHandler for DefaultZapHandler<H> {
async fn authenticate(&self, request: &ZapRequest) -> ZapResponse {
if request.version != ZAP_VERSION {
return ZapResponse::failure(
request.request_id.clone(),
"Unsupported ZAP request version",
);
}
match request.mechanism {
ZapMechanism::Null => {
if !request.credentials.is_empty() {
return ZapResponse::failure(
request.request_id.clone(),
"Unexpected credentials",
);
}
ZapResponse::success(request.request_id.clone(), String::new())
}
ZapMechanism::Plain => {
if request.credentials.len() != 2 {
return ZapResponse::failure(request.request_id.clone(), "Missing credentials");
}
let username = match std::str::from_utf8(&request.credentials[0]) {
Ok(username) => username,
Err(_) => {
return ZapResponse::failure(
request.request_id.clone(),
"Invalid UTF-8 username",
);
}
};
let password = match std::str::from_utf8(&request.credentials[1]) {
Ok(password) => password,
Err(_) => {
return ZapResponse::failure(
request.request_id.clone(),
"Invalid UTF-8 password",
);
}
};
match self
.plain_handler
.authenticate(username, password, &request.domain, &request.address)
.await
{
Ok(user_id) => ZapResponse::success(request.request_id.clone(), user_id),
Err(err) => ZapResponse::failure(request.request_id.clone(), &err),
}
}
ZapMechanism::Curve => {
if !self.accept_curve {
return ZapResponse::failure(request.request_id.clone(), "CURVE not enabled");
}
if request.credentials.len() != 1 {
return ZapResponse::failure(
request.request_id.clone(),
"Missing CURVE public key",
);
}
let public_key = &request.credentials[0];
if public_key.len() != 32 {
return ZapResponse::failure(
request.request_id.clone(),
"Invalid CURVE key length",
);
}
if public_key.iter().all(|&byte| byte == 0) {
return ZapResponse::failure(
request.request_id.clone(),
"Invalid CURVE public key",
);
}
if let Some(ref whitelist) = self.curve_key_whitelist {
let mut key_arr = [0u8; 32];
key_arr.copy_from_slice(public_key);
if !whitelist.contains(&key_arr) {
return ZapResponse::failure(
request.request_id.clone(),
"CURVE key not in whitelist",
);
}
}
use std::fmt::Write as _;
let mut user_id = String::with_capacity(public_key.len() * 2);
for b in public_key {
write!(user_id, "{b:02x}").expect("write to String is infallible");
}
ZapResponse::success(request.request_id.clone(), user_id)
}
}
}
}
pub struct ZapServer<H: ZapHandler> {
socket: DealerSocket<InprocStream>,
handler: Arc<H>,
}
impl<H: ZapHandler> ZapServer<H> {
pub fn new(handler: Arc<H>) -> io::Result<Self> {
let socket =
DealerSocket::bind_inproc_bidi("inproc://zeromq.zap.01", SocketOptions::default())?;
Ok(Self { socket, handler })
}
pub async fn start(&mut self) -> io::Result<()> {
loop {
let Some(msg) = self.socket.recv().await? else {
continue;
};
let request = match ZapRequest::decode(&msg) {
Ok(req) => req,
Err(_e) => {
continue;
}
};
let response = self.handler.authenticate(&request).await;
let frames = response.encode();
if let Err(_e) = self.socket.send(frames).await {
}
}
}
}
pub fn spawn_zap_server<H: ZapHandler + 'static>(handler: Arc<H>) -> io::Result<()> {
let mut server = ZapServer::new(handler)?;
monocoque_core::rt::spawn_detached(async move {
let _ = server.start().await;
});
Ok(())
}
pub fn start_default_zap_server<H: PlainAuthHandler + 'static>(
plain_handler: Arc<H>,
accept_curve: bool,
) -> io::Result<()> {
let zap_handler = Arc::new(DefaultZapHandler::new(plain_handler, accept_curve));
spawn_zap_server(zap_handler)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::security::ZapStatus;
use crate::security::curve::CurveKeyPair;
use crate::security::plain::StaticPlainHandler;
use bytes::Bytes;
use monocoque_core::rt::LocalRuntime;
fn zap_request(mechanism: ZapMechanism, credentials: Vec<Bytes>) -> ZapRequest {
ZapRequest {
version: "1.0".to_string(),
request_id: "request".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism,
credentials,
}
}
fn plain_request(credentials: Vec<Bytes>) -> ZapRequest {
zap_request(ZapMechanism::Plain, credentials)
}
fn curve_request(credentials: Vec<Bytes>) -> ZapRequest {
zap_request(ZapMechanism::Curve, credentials)
}
fn default_handler(accept_curve: bool) -> DefaultZapHandler<StaticPlainHandler> {
DefaultZapHandler::new(Arc::new(StaticPlainHandler::new()), accept_curve)
}
fn default_plain_handler() -> DefaultZapHandler<StaticPlainHandler> {
let mut plain_handler = StaticPlainHandler::new();
plain_handler.add_user("admin", "secret");
DefaultZapHandler::new(Arc::new(plain_handler), true)
}
fn with_local_runtime<F>(f: F)
where
F: Future<Output = ()>,
{
LocalRuntime::new().unwrap().block_on(f);
}
#[test]
fn test_default_zap_handler_null() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, true);
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "1".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Success);
});
}
#[test]
fn test_default_zap_handler_rejects_unsupported_zap_version() {
with_local_runtime(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, true);
let request = ZapRequest {
version: "2.0".to_string(),
request_id: "bad-version".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Failure);
});
}
#[test]
fn test_default_zap_handler_plain_success() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let mut plain_handler = StaticPlainHandler::new();
plain_handler.add_user("admin", "secret");
let handler = DefaultZapHandler::new(Arc::new(plain_handler), true);
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "2".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Plain,
credentials: vec![Bytes::from("admin"), Bytes::from("secret")],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Success);
assert_eq!(response.user_id, "admin");
});
}
#[test]
fn test_default_zap_handler_rejects_null_credentials() {
with_local_runtime(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, true);
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "null-extra".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![Bytes::from("unexpected")],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Failure);
});
}
#[test]
fn test_default_zap_handler_rejects_plain_extra_credentials() {
with_local_runtime(async {
let mut plain_handler = StaticPlainHandler::new();
plain_handler.add_user("admin", "secret");
let handler = DefaultZapHandler::new(Arc::new(plain_handler), true);
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "plain-extra".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Plain,
credentials: vec![
Bytes::from("admin"),
Bytes::from("secret"),
Bytes::from("shadow"),
],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Failure);
});
}
#[test]
fn default_zap_handler_rejects_invalid_utf8_plain_credentials() {
with_local_runtime(async {
let handler = default_plain_handler();
let request = plain_request(vec![Bytes::from(vec![0xff]), Bytes::from("secret")]);
let response = handler.authenticate(&request).await;
assert_eq!(
response.status_code,
ZapStatus::Failure,
"Default ZAP handler authenticated invalid UTF-8 PLAIN credentials after lossy conversion"
);
});
}
#[test]
fn test_default_zap_handler_plain_failure() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, true);
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "3".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Plain,
credentials: vec![Bytes::from("admin"), Bytes::from("wrong")],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Failure);
});
}
#[test]
fn test_default_zap_handler_curve_success() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, true);
let public_key = [1u8; 32];
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "4".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Curve,
credentials: vec![Bytes::copy_from_slice(&public_key)],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Success);
});
}
#[test]
fn test_default_zap_handler_rejects_curve_extra_credentials() {
with_local_runtime(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, true);
let public_key = [0u8; 32];
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "curve-extra".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Curve,
credentials: vec![Bytes::copy_from_slice(&public_key), Bytes::from("shadow")],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Failure);
});
}
#[test]
fn default_zap_handler_rejects_curve_request_with_extra_credentials() {
with_local_runtime(async {
let handler = default_handler(true);
let public_key = CurveKeyPair::generate().public;
let request = curve_request(vec![
Bytes::copy_from_slice(public_key.as_bytes()),
Bytes::from("ignored-injected-frame"),
]);
let response = handler.authenticate(&request).await;
assert_eq!(
response.status_code,
ZapStatus::Failure,
"Default ZAP handler authenticated a malformed CURVE request with extra credential frames"
);
});
}
#[test]
fn test_default_zap_handler_curve_disabled() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let plain_handler = Arc::new(StaticPlainHandler::new());
let handler = DefaultZapHandler::new(plain_handler, false);
let public_key = [0u8; 32];
let request = ZapRequest {
version: "1.0".to_string(),
request_id: "5".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Curve,
credentials: vec![Bytes::copy_from_slice(&public_key)],
};
let response = handler.authenticate(&request).await;
assert_eq!(response.status_code, ZapStatus::Failure);
});
}
struct IpDenyListHandler {
denied_ips: Vec<String>,
}
impl IpDenyListHandler {
fn new(denied_ips: Vec<&str>) -> Self {
Self {
denied_ips: denied_ips.into_iter().map(str::to_string).collect(),
}
}
}
#[async_trait::async_trait(?Send)]
impl ZapHandler for IpDenyListHandler {
async fn authenticate(&self, request: &ZapRequest) -> ZapResponse {
if self
.denied_ips
.iter()
.any(|ip| request.address.starts_with(ip.as_str()))
{
return ZapResponse::failure(
request.request_id.clone(),
format!("Address {} is blocked", request.address),
);
}
ZapResponse::success(request.request_id.clone(), String::new())
}
}
#[test]
fn test_ip_based_rejection() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let handler = IpDenyListHandler::new(vec!["192.168.1.100", "10.0.0.1"]);
let denied_request = ZapRequest {
version: "1.0".to_string(),
request_id: "deny-1".to_string(),
domain: "global".to_string(),
address: "192.168.1.100".to_string(), identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let denied_response = handler.authenticate(&denied_request).await;
assert_eq!(
denied_response.status_code,
ZapStatus::Failure,
"connections from denied IPs must be rejected with status 400"
);
assert!(
denied_response.status_text.contains("192.168.1.100"),
"failure message should name the blocked address"
);
let denied_request2 = ZapRequest {
version: "1.0".to_string(),
request_id: "deny-2".to_string(),
domain: "global".to_string(),
address: "10.0.0.1".to_string(), identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let denied_response2 = handler.authenticate(&denied_request2).await;
assert_eq!(
denied_response2.status_code,
ZapStatus::Failure,
"10.0.0.1 is on the deny list and must be rejected"
);
let allowed_request = ZapRequest {
version: "1.0".to_string(),
request_id: "allow-1".to_string(),
domain: "global".to_string(),
address: "127.0.0.1".to_string(), identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let allowed_response = handler.authenticate(&allowed_request).await;
assert_eq!(
allowed_response.status_code,
ZapStatus::Success,
"connections from allowed IPs must succeed with status 200"
);
});
}
#[test]
fn test_ip_subnet_prefix_rejection() {
monocoque_core::rt::LocalRuntime::new()
.unwrap()
.block_on(async {
let handler = IpDenyListHandler::new(vec!["10.0.0."]);
let blocked = ZapRequest {
version: "1.0".to_string(),
request_id: "subnet-1".to_string(),
domain: "global".to_string(),
address: "10.0.0.55".to_string(),
identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let resp = handler.authenticate(&blocked).await;
assert_eq!(
resp.status_code,
ZapStatus::Failure,
"addresses matching a denied subnet prefix must be rejected"
);
let allowed = ZapRequest {
version: "1.0".to_string(),
request_id: "subnet-2".to_string(),
domain: "global".to_string(),
address: "10.0.1.1".to_string(), identity: Bytes::new(),
mechanism: ZapMechanism::Null,
credentials: vec![],
};
let resp2 = handler.authenticate(&allowed).await;
assert_eq!(
resp2.status_code,
ZapStatus::Success,
"addresses not matching a denied prefix must be accepted"
);
});
}
}