use serde::{Deserialize, Serialize};
use crate::error::{Result, TreetopError};
use super::policy::PermitPolicy;
use super::request::AuthorizeRequest;
use super::version::PolicyVersion;
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq, Hash)]
pub enum DecisionBrief {
Allow,
Deny,
}
impl std::fmt::Display for DecisionBrief {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
DecisionBrief::Allow => write!(f, "Allow"),
DecisionBrief::Deny => write!(f, "Deny"),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct AuthorizeDecisionBrief {
pub decision: DecisionBrief,
pub version: PolicyVersion,
pub policy_id: String,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct AuthorizeDecisionDetailed {
pub policy: Vec<PermitPolicy>,
pub decision: DecisionBrief,
pub version: PolicyVersion,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(tag = "status", rename_all = "lowercase")]
pub enum BatchResult<T> {
Success {
#[serde(rename = "result")]
data: T,
},
Failed {
#[serde(rename = "error")]
message: String,
},
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct IndexedResult<T> {
pub index: usize,
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(flatten)]
pub result: BatchResult<T>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct AuthorizeResponse<T> {
pub results: Vec<IndexedResult<T>>,
pub version: PolicyVersion,
pub successful: usize,
pub failed: usize,
}
pub(crate) trait ValidateDecision {
fn policy_version(&self) -> &PolicyVersion;
fn validate(&self) -> std::result::Result<(), &'static str>;
}
impl ValidateDecision for AuthorizeDecisionBrief {
fn policy_version(&self) -> &PolicyVersion {
&self.version
}
fn validate(&self) -> std::result::Result<(), &'static str> {
match self.decision {
DecisionBrief::Allow if self.policy_id.is_empty() => {
Err("an Allow decision has no matching policy ID")
}
DecisionBrief::Deny if !self.policy_id.is_empty() => {
Err("a Deny decision contains matching policy IDs")
}
DecisionBrief::Allow | DecisionBrief::Deny => Ok(()),
}
}
}
impl ValidateDecision for AuthorizeDecisionDetailed {
fn policy_version(&self) -> &PolicyVersion {
&self.version
}
fn validate(&self) -> std::result::Result<(), &'static str> {
match self.decision {
DecisionBrief::Allow if self.policy.is_empty() => {
Err("an Allow decision has no matching policies")
}
DecisionBrief::Deny if !self.policy.is_empty() => {
Err("a Deny decision contains matching policies")
}
DecisionBrief::Allow | DecisionBrief::Deny => Ok(()),
}
}
}
fn validate_response<T>(response: &AuthorizeResponse<T>, expected_results: usize) -> Result<()>
where
T: ValidateDecision,
{
if response.results.len() != expected_results {
return Err(TreetopError::InvalidResponse(format!(
"authorize returned {} results for {expected_results} requests",
response.results.len()
)));
}
let actual_successes = response
.results
.iter()
.filter(|result| matches!(&result.result, BatchResult::Success { .. }))
.count();
let actual_failures = response.results.len() - actual_successes;
if response.successful != actual_successes || response.failed != actual_failures {
return Err(TreetopError::InvalidResponse(format!(
"authorize result counts are inconsistent: declared {} successful and {} failed, observed {actual_successes} successful and {actual_failures} failed",
response.successful, response.failed
)));
}
for (position, result) in response.results.iter().enumerate() {
if result.index != position {
return Err(TreetopError::InvalidResponse(format!(
"authorize result at position {position} reports index {}; results must remain in request order",
result.index,
)));
}
if let BatchResult::Success { data } = &result.result {
if data.policy_version() != &response.version {
return Err(TreetopError::InvalidResponse(format!(
"authorize result index {} reports a different policy version than the batch",
result.index
)));
}
if let Err(message) = data.validate() {
return Err(TreetopError::InvalidResponse(format!(
"authorize result index {} is inconsistent: {message}",
result.index
)));
}
}
}
Ok(())
}
fn validate_response_against<T>(
response: &AuthorizeResponse<T>,
request: &AuthorizeRequest,
) -> Result<()>
where
T: ValidateDecision,
{
validate_response(response, request.len())?;
for (result, submitted) in response.results.iter().zip(request.requests()) {
if result.id.as_deref() != submitted.id() {
return Err(TreetopError::InvalidResponse(format!(
"authorize result index {} reports correlation ID {:?}, expected {:?}",
result.index,
result.id.as_deref(),
submitted.id()
)));
}
}
Ok(())
}
impl AuthorizeResponse<AuthorizeDecisionBrief> {
pub fn validate(&self, expected_results: usize) -> Result<()> {
validate_response(self, expected_results)
}
pub(crate) fn validate_against(&self, request: &AuthorizeRequest) -> Result<()> {
validate_response_against(self, request)
}
}
impl AuthorizeResponse<AuthorizeDecisionDetailed> {
pub fn validate(&self, expected_results: usize) -> Result<()> {
validate_response(self, expected_results)
}
pub(crate) fn validate_against(&self, request: &AuthorizeRequest) -> Result<()> {
validate_response_against(self, request)
}
}
impl<T> AuthorizeResponse<T> {
pub fn successes(&self) -> usize {
self.successful
}
pub fn failures(&self) -> usize {
self.failed
}
pub fn version(&self) -> &PolicyVersion {
&self.version
}
pub fn total(&self) -> usize {
self.results.len()
}
pub fn results(&self) -> &[IndexedResult<T>] {
&self.results
}
pub fn find_by_id(&self, id: &str) -> Option<&IndexedResult<T>> {
self.results.iter().find(|r| r.id.as_deref() == Some(id))
}
pub fn iter(&self) -> impl Iterator<Item = &IndexedResult<T>> {
self.results.iter()
}
pub fn into_results(self) -> Vec<IndexedResult<T>> {
self.results
}
}
impl<'a, T> IntoIterator for &'a AuthorizeResponse<T> {
type Item = &'a IndexedResult<T>;
type IntoIter = std::slice::Iter<'a, IndexedResult<T>>;
fn into_iter(self) -> Self::IntoIter {
self.results.iter()
}
}
pub type AuthorizeBriefResponse = AuthorizeResponse<AuthorizeDecisionBrief>;
pub type AuthorizeDetailedResponse = AuthorizeResponse<AuthorizeDecisionDetailed>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn brief_response_deserialization() {
let json = serde_json::json!({
"results": [{
"index": 0,
"id": "req-1",
"status": "success",
"result": {
"decision": "Allow",
"version": { "hash": "abc", "loaded_at": "2025-01-01T00:00:00Z", "label_set": null, "generation": 0},
"policy_id": "policy1"
}
}],
"version": { "hash": "abc", "loaded_at": "2025-01-01T00:00:00Z", "label_set": null, "generation": 0},
"successful": 1,
"failed": 0
});
let resp: AuthorizeBriefResponse = serde_json::from_value(json).unwrap();
assert_eq!(resp.successes(), 1);
assert_eq!(resp.failures(), 0);
assert_eq!(resp.total(), 1);
let result = resp.find_by_id("req-1").unwrap();
assert_eq!(result.index, 0);
match &result.result {
BatchResult::Success { data } => {
assert_eq!(data.decision, DecisionBrief::Allow);
assert_eq!(data.policy_id, "policy1");
}
_ => panic!("expected success"),
}
}
#[test]
fn failed_result_deserialization() {
let json = serde_json::json!({
"results": [{
"index": 0,
"status": "failed",
"error": "invalid principal"
}],
"version": { "hash": "abc", "loaded_at": "2025-01-01T00:00:00Z", "label_set": null, "generation": 0},
"successful": 0,
"failed": 1
});
let resp: AuthorizeBriefResponse = serde_json::from_value(json).unwrap();
assert_eq!(resp.failures(), 1);
match &resp.results()[0].result {
BatchResult::Failed { message } => assert_eq!(message, "invalid principal"),
_ => panic!("expected failure"),
}
}
#[test]
fn denied_response_cannot_report_matching_policy_ids() {
let json = serde_json::json!({
"results": [{
"index": 0,
"status": "success",
"result": {
"decision": "Deny",
"version": { "hash": "abc", "loaded_at": "2025-01-01T00:00:00Z", "label_set": null, "generation": 0},
"policy_id": "policy1"
}
}],
"version": { "hash": "abc", "loaded_at": "2025-01-01T00:00:00Z", "label_set": null, "generation": 0},
"successful": 1,
"failed": 0
});
let response: AuthorizeBriefResponse = serde_json::from_value(json).unwrap();
assert!(matches!(
response.validate(1),
Err(TreetopError::InvalidResponse(message))
if message.contains("Deny decision contains matching policy IDs")
));
}
}