Stage 3: RBAC-Middleware (Token -> Identity, tools/call gegen Policy)
CI / test (push) Successful in 13m14s

This commit is contained in:
2026-06-24 07:44:39 +00:00
parent fce25bfee1
commit 6bc8716f1a
+132
View File
@@ -0,0 +1,132 @@
use std::collections::HashSet;
use std::sync::Arc;
use axum::{
body::Body,
extract::{Request, State},
http::{header::AUTHORIZATION, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use dtrack_perms::Policy;
/// Maximale Body-Groesse, die die RBAC-Middleware puffert (1 MiB).
const MAX_BODY: usize = 1 << 20;
/// Zugriffskontrolle vor `/mcp`.
#[derive(Clone)]
pub enum Gate {
/// RBAC aus einer Policy-Datei (Token -> Rolle -> erlaubte Tools).
Rbac(Arc<Policy>),
/// Flache Bearer-Token-Liste (Rueckwaerts-Kompat); `allow_all` = ungeschuetzt.
Flat {
tokens: Arc<HashSet<String>>,
allow_all: bool,
},
}
impl Gate {
/// `DTRACK_PERMISSIONS` (Policy-Datei) hat Vorrang, sonst flache Liste aus
/// `DTRACK_HTTP_TOKENS`.
pub fn from_env() -> anyhow::Result<Self> {
if let Ok(path) = std::env::var("DTRACK_PERMISSIONS") {
let policy = Policy::load(&path)?;
tracing::info!("RBAC aktiv: {} Client(s) aus {path}", policy.client_count());
return Ok(Self::Rbac(Arc::new(policy)));
}
let raw = std::env::var("DTRACK_HTTP_TOKENS").unwrap_or_default();
let tokens: HashSet<String> = raw
.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(str::to_owned)
.collect();
let allow_all = tokens.is_empty();
if allow_all {
tracing::warn!(
"Weder DTRACK_PERMISSIONS noch DTRACK_HTTP_TOKENS gesetzt -- /mcp ist UNGESCHUETZT (nur lokal)."
);
}
Ok(Self::Flat {
tokens: Arc::new(tokens),
allow_all,
})
}
}
fn bearer(req: &Request) -> Option<String> {
req.headers()
.get(AUTHORIZATION)
.and_then(|h| h.to_str().ok())
.and_then(|h| h.strip_prefix("Bearer "))
.map(|s| s.trim().to_owned())
}
/// Sammelt die in einem (ggf. Batch-)JSON-RPC-Body aufgerufenen Tool-Namen.
fn called_tools(bytes: &[u8]) -> Vec<String> {
let Ok(value) = serde_json::from_slice::<serde_json::Value>(bytes) else {
return Vec::new();
};
let items: Vec<&serde_json::Value> = match &value {
serde_json::Value::Array(arr) => arr.iter().collect(),
single => vec![single],
};
items
.into_iter()
.filter(|item| item.get("method").and_then(|m| m.as_str()) == Some("tools/call"))
.filter_map(|item| {
item.get("params")
.and_then(|p| p.get("name"))
.and_then(|n| n.as_str())
.map(str::to_owned)
})
.collect()
}
/// Middleware vor `/mcp`: authentifiziert und prueft im RBAC-Modus die Tool-Calls.
pub async fn gate_mw(State(gate): State<Gate>, req: Request, next: Next) -> Response {
match &gate {
Gate::Flat { tokens, allow_all } => {
if *allow_all {
return next.run(req).await;
}
match bearer(&req) {
Some(tok) if tokens.contains(&tok) => next.run(req).await,
_ => StatusCode::UNAUTHORIZED.into_response(),
}
}
Gate::Rbac(policy) => {
let Some(token) = bearer(&req) else {
return StatusCode::UNAUTHORIZED.into_response();
};
let identity = match policy.identify(&token) {
Some(id) => id.clone(),
None => return StatusCode::UNAUTHORIZED.into_response(),
};
// Body puffern, Tool-Calls inspizieren, Body wieder einsetzen.
let (parts, body) = req.into_parts();
let bytes = match axum::body::to_bytes(body, MAX_BODY).await {
Ok(b) => b,
Err(_) => return StatusCode::PAYLOAD_TOO_LARGE.into_response(),
};
for tool in called_tools(&bytes) {
if !identity.can_use(&tool) {
return (
StatusCode::FORBIDDEN,
format!(
"Tool '{tool}' fuer Client '{}' (Rolle '{}') nicht erlaubt",
identity.client, identity.role
),
)
.into_response();
}
}
let req = Request::from_parts(parts, Body::from(bytes));
next.run(req).await
}
}
}