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

105 lines
3.6 KiB
Rust

use crate::AppState;
use crate::state::GenericEvent;
use axum::extract::{Query, State};
use axum::response::IntoResponse;
use std::collections::HashMap;
use std::sync::Arc;
pub async fn wait_for_event_handler(
State(state): State<Arc<AppState>>,
Query(params): Query<HashMap<String, String>>,
) -> impl IntoResponse {
let topic = params.get("topic").cloned();
let session_id = params.get("session_id").cloned();
let mut rx = state.handler.state.event_bus_tx.subscribe();
loop {
match rx.recv().await {
Ok(event) => {
let topic_matches = topic.as_ref().is_none_or(|t| t == &event.topic);
let session_matches = session_id
.as_ref()
.is_none_or(|s| Some(s) == event.session_id.as_ref());
if topic_matches && session_matches {
return axum::Json(event);
}
}
Err(_) => {
return axum::Json(GenericEvent {
topic: "error".to_string(),
session_id: None,
payload: serde_json::json!({"error": "Event bus lagged or closed"}),
});
}
}
}
}
pub async fn post_event_handler(
State(state): State<Arc<AppState>>,
axum::Json(event): axum::Json<GenericEvent>,
) -> impl IntoResponse {
let _ = state.handler.state.event_bus_tx.send(event);
axum::Json(serde_json::json!({"status": "ok"}))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::router::MemoryHandler;
use crate::state::MemoryState;
use axum::extract::Query;
use axum::extract::State;
use std::collections::HashMap;
use std::sync::RwLock;
use std::sync::atomic::AtomicUsize;
use tempfile::tempdir;
#[tokio::test]
async fn test_events_wait_and_post() {
let dir = tempdir().unwrap();
let mem_state = Arc::new(MemoryState::new(dir.path().to_str().unwrap()));
let app_state = Arc::new(AppState {
handler: Arc::new(MemoryHandler::new(mem_state.clone())),
clients: RwLock::new(HashMap::new()),
next_id: AtomicUsize::new(1),
});
// Start wait_for_event in a background task
let app_state_clone = app_state.clone();
let mut params = HashMap::new();
params.insert("topic".to_string(), "test_topic".to_string());
params.insert("session_id".to_string(), "123".to_string());
let wait_task = tokio::spawn(async move {
let res = wait_for_event_handler(State(app_state_clone), Query(params)).await;
// axum::Json<GenericEvent> is returned, we need to extract it somehow, but just returning is enough for testing
res
});
// Yield slightly to ensure the wait task has subscribed
tokio::time::sleep(tokio::time::Duration::from_millis(50)).await;
// Post an event that shouldn't match
let unmatched_event = GenericEvent {
topic: "wrong_topic".to_string(),
session_id: Some("123".to_string()),
payload: serde_json::json!({}),
};
post_event_handler(State(app_state.clone()), axum::Json(unmatched_event)).await;
// Post the matching event
let matched_event = GenericEvent {
topic: "test_topic".to_string(),
session_id: Some("123".to_string()),
payload: serde_json::json!({"foo": "bar"}),
};
post_event_handler(State(app_state.clone()), axum::Json(matched_event)).await;
// Wait for the wait task to complete
let _ = wait_task.await.unwrap();
}
}