use crate::access_control::traits::AccessController;
use crate::guardian::error::{GuardianError, Result};
use iroh::endpoint::{Connection, Endpoint};
use iroh::protocol::{AcceptError, ProtocolHandler};
use iroh::{EndpointId as NodeId, PublicKey};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use tracing::{debug, warn};
pub const TICKET_ALPN: &[u8] = b"/guardian-db/ticket/1";
const RESP_GRANTED: u8 = 1;
const RESP_DENIED: u8 = 0;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum GrantedMode {
Read,
Write,
}
#[derive(Clone)]
pub struct TicketProvider {
pub read_ticket: String,
pub write_ticket: String,
pub access_controller: Arc<dyn AccessController>,
}
pub type TicketRegistry = Arc<RwLock<HashMap<String, TicketProvider>>>;
pub fn new_registry() -> TicketRegistry {
Arc::new(RwLock::new(HashMap::new()))
}
#[derive(Clone)]
pub struct TicketProtocolHandler {
registry: TicketRegistry,
}
impl std::fmt::Debug for TicketProtocolHandler {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TicketProtocolHandler")
.finish_non_exhaustive()
}
}
impl TicketProtocolHandler {
pub fn new(registry: TicketRegistry) -> Self {
Self { registry }
}
async fn resolve(&self, address: &str, requester: PublicKey) -> Vec<u8> {
let provider = {
let reg = self.registry.read().await;
reg.get(address).cloned()
};
let Some(provider) = provider else {
debug!(address, "Ticket requested for unknown store — denied");
return vec![RESP_DENIED];
};
match authorized_mode(&*provider.access_controller, requester).await {
Some(mode) => {
let ticket = match mode {
GrantedMode::Write => &provider.write_ticket,
GrantedMode::Read => &provider.read_ticket,
};
debug!(address, peer = %requester.fmt_short(), ?mode, "Ticket granted");
let mut out = Vec::with_capacity(ticket.len() + 1);
out.push(RESP_GRANTED);
out.extend_from_slice(ticket.as_bytes());
out
}
None => {
warn!(address, peer = %requester.fmt_short(), "Ticket denied by the access controller");
vec![RESP_DENIED]
}
}
}
}
impl ProtocolHandler for TicketProtocolHandler {
async fn accept(&self, connection: Connection) -> std::result::Result<(), AcceptError> {
let requester = connection.remote_id();
let (mut send, mut recv) = connection.accept_bi().await?;
let req = recv
.read_to_end(4096)
.await
.map_err(AcceptError::from_err)?;
let address = String::from_utf8_lossy(&req).to_string();
let response = self.resolve(&address, requester).await;
send.write_all(&response)
.await
.map_err(AcceptError::from_err)?;
send.finish().map_err(AcceptError::from_err)?;
connection.closed().await;
Ok(())
}
}
async fn authorized_mode(acl: &dyn AccessController, requester: PublicKey) -> Option<GrantedMode> {
let requester_hex = hex::encode(requester.as_bytes());
let role_grants = |keys: Vec<String>| {
keys.iter().any(|k| k == "*") || keys.iter().any(|k| k == &requester_hex)
};
if let Ok(keys) = acl.get_authorized_by_role("write").await
&& role_grants(keys)
{
return Some(GrantedMode::Write);
}
if let Ok(keys) = acl.get_authorized_by_role("read").await
&& role_grants(keys)
{
return Some(GrantedMode::Read);
}
None
}
pub async fn request_ticket(
endpoint: &Endpoint,
peer: NodeId,
address: &str,
) -> Result<Option<String>> {
let connection = endpoint
.connect(peer, TICKET_ALPN)
.await
.map_err(|e| GuardianError::Other(format!("Failed to connect for ticket: {}", e)))?;
let (mut send, mut recv) = connection
.open_bi()
.await
.map_err(|e| GuardianError::Other(format!("Failed to open ticket stream: {}", e)))?;
send.write_all(address.as_bytes())
.await
.map_err(|e| GuardianError::Other(format!("Failed to send ticket request: {}", e)))?;
send.finish()
.map_err(|e| GuardianError::Other(format!("Failed to finish ticket stream: {}", e)))?;
let resp = recv
.read_to_end(64 * 1024)
.await
.map_err(|e| GuardianError::Other(format!("Failed to read ticket response: {}", e)))?;
connection.close(0u32.into(), b"done");
match resp.first() {
Some(&RESP_GRANTED) if resp.len() > 1 => {
Ok(Some(String::from_utf8_lossy(&resp[1..]).to_string()))
}
_ => Ok(None),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::access_control::acl_simple::SimpleAccessController;
use std::collections::HashMap;
fn random_public_key() -> PublicKey {
iroh::SecretKey::generate().public()
}
fn acl_with(role: &str, keys: Vec<&str>) -> Arc<dyn AccessController> {
let mut map = HashMap::new();
map.insert(
role.to_string(),
keys.into_iter().map(String::from).collect(),
);
Arc::new(SimpleAccessController::new(Some(map))) as Arc<dyn AccessController>
}
#[tokio::test]
async fn wildcard_write_grants_write_to_any_peer() {
let acl = acl_with("write", vec!["*"]);
assert_eq!(
authorized_mode(&*acl, random_public_key()).await,
Some(GrantedMode::Write)
);
}
#[tokio::test]
async fn wildcard_read_grants_read_to_any_peer() {
let acl = acl_with("read", vec!["*"]);
assert_eq!(
authorized_mode(&*acl, random_public_key()).await,
Some(GrantedMode::Read)
);
}
#[tokio::test]
async fn read_only_peer_never_gets_write() {
let peer = random_public_key();
let peer_hex = hex::encode(peer.as_bytes());
let mut map = HashMap::new();
map.insert(
"write".to_string(),
vec![hex::encode(random_public_key().as_bytes())],
);
map.insert("read".to_string(), vec![peer_hex]);
let acl = Arc::new(SimpleAccessController::new(Some(map))) as Arc<dyn AccessController>;
assert_eq!(authorized_mode(&*acl, peer).await, Some(GrantedMode::Read));
}
#[tokio::test]
async fn write_precedence_over_read() {
let peer = random_public_key();
let peer_hex = hex::encode(peer.as_bytes());
let mut map = HashMap::new();
map.insert("write".to_string(), vec![peer_hex.clone()]);
map.insert("read".to_string(), vec![peer_hex]);
let acl = Arc::new(SimpleAccessController::new(Some(map))) as Arc<dyn AccessController>;
assert_eq!(authorized_mode(&*acl, peer).await, Some(GrantedMode::Write));
}
#[tokio::test]
async fn specific_authorized_key_gets_write() {
let peer = random_public_key();
let peer_hex = hex::encode(peer.as_bytes());
let acl = acl_with("write", vec![peer_hex.as_str()]);
assert_eq!(authorized_mode(&*acl, peer).await, Some(GrantedMode::Write));
}
#[tokio::test]
async fn unknown_key_is_denied_when_no_wildcard() {
let other_hex = hex::encode(random_public_key().as_bytes());
let acl = acl_with("write", vec![other_hex.as_str()]);
assert_eq!(authorized_mode(&*acl, random_public_key()).await, None);
}
#[tokio::test]
async fn empty_acl_denies() {
let acl = acl_with("write", vec![]);
assert_eq!(authorized_mode(&*acl, random_public_key()).await, None);
}
#[tokio::test]
async fn resolve_unknown_store_is_denied() {
let handler = TicketProtocolHandler::new(new_registry());
let resp = handler.resolve("does-not-exist", random_public_key()).await;
assert_eq!(resp, vec![RESP_DENIED]);
}
#[tokio::test]
async fn resolve_grants_write_ticket_to_write_peer() {
let registry = new_registry();
registry.write().await.insert(
"shared-kv".to_string(),
TicketProvider {
read_ticket: "read-ticket-xyz".to_string(),
write_ticket: "write-ticket-xyz".to_string(),
access_controller: acl_with("write", vec!["*"]),
},
);
let handler = TicketProtocolHandler::new(registry);
let resp = handler.resolve("shared-kv", random_public_key()).await;
assert_eq!(resp.first(), Some(&RESP_GRANTED));
assert_eq!(&resp[1..], b"write-ticket-xyz");
}
#[tokio::test]
async fn resolve_grants_read_ticket_to_read_only_peer() {
let peer = random_public_key();
let peer_hex = hex::encode(peer.as_bytes());
let mut map = HashMap::new();
map.insert("read".to_string(), vec![peer_hex]);
let acl = Arc::new(SimpleAccessController::new(Some(map))) as Arc<dyn AccessController>;
let registry = new_registry();
registry.write().await.insert(
"shared-kv".to_string(),
TicketProvider {
read_ticket: "read-ticket-xyz".to_string(),
write_ticket: "write-ticket-xyz".to_string(),
access_controller: acl,
},
);
let handler = TicketProtocolHandler::new(registry);
let resp = handler.resolve("shared-kv", peer).await;
assert_eq!(resp.first(), Some(&RESP_GRANTED));
assert_eq!(&resp[1..], b"read-ticket-xyz");
}
#[tokio::test]
async fn resolve_denies_unauthorized_peer() {
let registry = new_registry();
let other_hex = hex::encode(random_public_key().as_bytes());
registry.write().await.insert(
"private-kv".to_string(),
TicketProvider {
read_ticket: "secret-read-ticket".to_string(),
write_ticket: "secret-write-ticket".to_string(),
access_controller: acl_with("write", vec![other_hex.as_str()]),
},
);
let handler = TicketProtocolHandler::new(registry);
let resp = handler.resolve("private-kv", random_public_key()).await;
assert_eq!(resp, vec![RESP_DENIED]);
}
}