Files
mcp-memory/server/src/api/rest.rs
T

175 lines
5.3 KiB
Rust

use crate::AppState;
use crate::error::AppError;
use crate::models::GateRecord;
use axum::{
Json,
extract::{Query, State},
response::IntoResponse,
};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(serde::Deserialize, serde::Serialize)]
pub struct GateVerifyReq {
pub action: String,
pub target: String,
pub namespace: Option<String>,
#[serde(default)]
pub params: HashMap<String, String>,
#[serde(default)]
pub consume: bool,
}
#[derive(serde::Deserialize, serde::Serialize)]
pub struct GateSetReq {
pub action: String,
pub target: String,
pub namespace: Option<String>,
#[serde(default)]
pub params: HashMap<String, String>,
pub authorize: Option<bool>,
pub block: Option<bool>,
pub reason: Option<String>,
}
pub async fn gate_verify_handler(
State(app_state): State<Arc<AppState>>,
Query(q): Query<GateVerifyReq>,
) -> Result<impl IntoResponse, AppError> {
let mut found = None;
let mut to_remove = None;
app_state.handler.state.env.gates.modify(|gates| {
if let Some(idx) = gates.iter().position(|g| {
g.action == q.action
&& g.target == q.target
&& g.namespace == q.namespace
&& g.params == q.params
}) {
found = Some(gates[idx].clone());
if q.consume {
to_remove = Some(idx);
}
}
if let Some(idx) = to_remove {
gates.remove(idx);
}
});
match found {
Some(record) => {
if record.status == "authorized" {
Ok((axum::http::StatusCode::OK, "Authorized"))
} else {
let msg = if let Some(r) = record.reason {
format!("Action blocked. Reason: {}", r)
} else {
"Action blocked.".to_string()
};
Err(AppError::Forbidden(msg))
}
}
None => Err(AppError::NotFound(
"Action not yet authorized (no gate record found).".to_string(),
)),
}
}
pub async fn gate_set_handler(
State(app_state): State<Arc<AppState>>,
Json(body): Json<GateSetReq>,
) -> Result<impl IntoResponse, AppError> {
let status = if body.block.unwrap_or(false) {
"blocked".to_string()
} else if body.authorize.unwrap_or(false) {
"authorized".to_string()
} else {
"pending".to_string()
};
let record = GateRecord {
id: uuid::Uuid::new_v4().to_string(),
action: body.action.clone(),
target: body.target.clone(),
namespace: body.namespace.clone(),
params: body.params.clone(),
status,
reason: body.reason.clone(),
timestamp: crate::handlers::utils::now_secs(),
};
app_state.handler.state.env.gates.modify(|gates| {
gates.retain(|g| !(g.action == record.action && g.target == record.target));
gates.push(record);
});
Ok((axum::http::StatusCode::OK, "Gate state updated."))
}
pub async fn health_handler() -> &'static str {
"OK"
}
#[cfg(test)]
mod tests {
use super::*;
use crate::router::MemoryHandler;
use crate::state::MemoryState;
use std::sync::RwLock;
use std::sync::atomic::AtomicUsize;
use tempfile::tempdir;
#[tokio::test]
async fn test_gate_handlers() {
let dir = tempdir().unwrap();
let state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let (shutdown_tx, _) = tokio::sync::oneshot::channel();
let app_state = Arc::new(AppState {
handler: Arc::new(MemoryHandler::new(state.clone())),
clients: RwLock::new(HashMap::new()),
next_id: AtomicUsize::new(1),
shutdown_tx: std::sync::Mutex::new(Some(shutdown_tx)),
});
// Set a gate to authorized
let set_req = GateSetReq {
action: "push".to_string(),
target: "main".to_string(),
namespace: Some("global".to_string()),
params: HashMap::new(),
authorize: Some(true),
block: None,
reason: None,
};
let res_set = gate_set_handler(State(app_state.clone()), Json(set_req))
.await
.unwrap();
assert_eq!(res_set.into_response().status(), axum::http::StatusCode::OK);
// Verify the gate (and consume it)
let verify_req = GateVerifyReq {
action: "push".to_string(),
target: "main".to_string(),
namespace: Some("global".to_string()),
params: HashMap::new(),
consume: true,
};
let res_verify = gate_verify_handler(State(app_state.clone()), Query(verify_req))
.await
.unwrap();
assert_eq!(
res_verify.into_response().status(),
axum::http::StatusCode::OK
);
// Verify again should fail since it was consumed
let verify_req2 = GateVerifyReq {
action: "push".to_string(),
target: "main".to_string(),
namespace: Some("global".to_string()),
params: HashMap::new(),
consume: false,
};
let res_verify2 = gate_verify_handler(State(app_state.clone()), Query(verify_req2)).await;
assert!(res_verify2.is_err());
}
}