use std::sync::Arc;
use async_trait::async_trait;
use crate::handler::{PermissionHandler, PermissionResult, permission_handler_failure};
use crate::types::{PermissionRequestData, RequestId, SessionId};
pub fn approve_all() -> Arc<dyn PermissionHandler> {
Arc::new(PolicyHandler {
policy: Policy::ApproveAll,
})
}
pub fn deny_all() -> Arc<dyn PermissionHandler> {
Arc::new(PolicyHandler {
policy: Policy::DenyAll,
})
}
pub fn approve_if<F>(predicate: F) -> Arc<dyn PermissionHandler>
where
F: Fn(&PermissionRequestData) -> bool + Send + Sync + 'static,
{
Arc::new(PolicyHandler {
policy: Policy::Predicate(Arc::new(predicate)),
})
}
#[derive(Clone)]
pub(crate) enum Policy {
ApproveAll,
DenyAll,
Predicate(Arc<dyn Fn(&PermissionRequestData) -> bool + Send + Sync>),
}
impl std::fmt::Debug for Policy {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ApproveAll => f.write_str("Policy::ApproveAll"),
Self::DenyAll => f.write_str("Policy::DenyAll"),
Self::Predicate(_) => f.write_str("Policy::Predicate(<fn>)"),
}
}
}
pub(crate) fn resolve_handler(
handler: Option<Arc<dyn PermissionHandler>>,
policy: Option<Policy>,
) -> Option<Arc<dyn PermissionHandler>> {
match (handler, policy) {
(_, Some(policy)) => Some(Arc::new(PolicyHandler { policy })),
(handler, None) => handler,
}
}
struct PolicyHandler {
policy: Policy,
}
#[async_trait]
impl PermissionHandler for PolicyHandler {
async fn handle(
&self,
_session_id: SessionId,
_request_id: RequestId,
data: PermissionRequestData,
) -> PermissionResult {
let approved = match &self.policy {
Policy::ApproveAll => true,
Policy::DenyAll => false,
Policy::Predicate(f) => f(&data),
};
if approved {
if matches!(self.policy, Policy::ApproveAll) && data.managed_settings_enabled {
permission_handler_failure(
"approve-all policy cannot be used when managed settings are enabled",
)
} else if data.managed_approval_required == Some(true) {
PermissionResult::no_result()
} else {
PermissionResult::approve_once()
}
} else {
PermissionResult::reject(None)
}
}
}
#[cfg(test)]
mod tests;