use axum::{
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
Query, State,
},
response::{
sse::{Event, KeepAlive, Sse},
IntoResponse,
},
};
use futures::{SinkExt, StreamExt};
use std::convert::Infallible;
use tokio_stream::wrappers::BroadcastStream;
use super::state::MultiUserMemoryManager;
use crate::relevance;
use crate::streaming;
use crate::validation;
#[derive(Debug, serde::Deserialize)]
pub struct SseQuery {
pub user_id: Option<String>,
}
pub type AppState = std::sync::Arc<MultiUserMemoryManager>;
pub async fn context_status_sse(
State(state): State<AppState>,
) -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
let receiver = state.context_broadcaster.subscribe();
let stream = BroadcastStream::new(receiver);
let event_stream = stream.filter_map(|result| async move {
match result {
Ok(status) => {
let data = serde_json::to_string(&status).ok()?;
Some(Ok(Event::default().event("context").data(data)))
}
Err(_) => None,
}
});
Sse::new(event_stream).keep_alive(
KeepAlive::new()
.interval(std::time::Duration::from_secs(15))
.text("ping"),
)
}
pub async fn memory_events_sse(
State(state): State<AppState>,
Query(params): Query<SseQuery>,
) -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
let receiver = state.subscribe_events();
let stream = BroadcastStream::new(receiver);
let filter_user_id = params.user_id;
let event_stream = stream.filter_map(move |result| {
let filter_uid = filter_user_id.clone();
async move {
match result {
Ok(event) => {
if let Some(ref uid) = filter_uid {
if event.user_id != *uid {
return None;
}
} else {
return None;
}
let json = serde_json::to_string(&event).ok()?;
Some(Ok(Event::default().event(&event.event_type).data(json)))
}
Err(_) => None,
}
}
});
Sse::new(event_stream).keep_alive(
KeepAlive::new()
.interval(std::time::Duration::from_secs(15))
.text("heartbeat"),
)
}
pub async fn streaming_memory_ws(
ws: WebSocketUpgrade,
State(state): State<AppState>,
) -> impl IntoResponse {
ws.on_upgrade(|socket| handle_streaming_socket(socket, state))
}
async fn handle_streaming_socket(socket: WebSocket, state: AppState) {
let (mut sender, mut receiver) = socket.split();
let mut session_id: Option<String> = None;
while let Some(msg) = receiver.next().await {
let msg = match msg {
Ok(Message::Text(text)) => text,
Ok(Message::Close(_)) => {
tracing::debug!("WebSocket closed before handshake");
return;
}
Ok(_) => continue, Err(e) => {
tracing::warn!("WebSocket error before handshake: {}", e);
return;
}
};
let handshake: streaming::StreamHandshake = match serde_json::from_str(&msg) {
Ok(h) => h,
Err(e) => {
let error = streaming::ExtractionResult::Error {
code: "INVALID_HANDSHAKE".to_string(),
message: format!("Failed to parse handshake: {}", e),
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
};
if let Err(e) = validation::validate_user_id(&handshake.user_id) {
let error = streaming::ExtractionResult::Error {
code: "INVALID_USER_ID".to_string(),
message: format!("Invalid user_id: {}", e),
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
{
let config = &handshake.extraction_config;
if config.checkpoint_interval_ms > 0 && config.checkpoint_interval_ms < 1000 {
let error = streaming::ExtractionResult::Error {
code: "INVALID_CONFIG".to_string(),
message: "checkpoint_interval_ms must be 0 (disabled) or >= 1000ms".to_string(),
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
if config.max_buffer_size > 1000 {
let error = streaming::ExtractionResult::Error {
code: "INVALID_CONFIG".to_string(),
message: "max_buffer_size must be <= 1000".to_string(),
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
}
let id = match state
.streaming_extractor
.create_session(handshake.clone())
.await
{
Ok(id) => id,
Err(e) => {
let error = streaming::ExtractionResult::Error {
code: "SESSION_LIMIT_REACHED".to_string(),
message: e,
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
};
session_id = Some(id.clone());
let ack = streaming::ExtractionResult::Ack {
message_type: "handshake".to_string(),
timestamp: chrono::Utc::now(),
};
if sender
.send(Message::Text(
serde_json::to_string(&ack)
.unwrap_or_else(|_| r#"{"ack":true}"#.to_string())
.into(),
))
.await
.is_err()
{
return;
}
tracing::info!(
"Streaming session {} created for user {} in {:?} mode",
id,
handshake.user_id,
handshake.mode
);
break;
}
let session_id = match session_id {
Some(id) => id,
None => return,
};
let user_memory = {
let stats = state
.streaming_extractor
.get_session_stats(&session_id)
.await;
match stats {
Some(s) => match state.get_user_memory(&s.user_id) {
Ok(m) => m,
Err(e) => {
tracing::error!("Failed to get user memory: {}", e);
return;
}
},
None => return,
}
};
while let Some(msg) = receiver.next().await {
let text = match msg {
Ok(Message::Text(text)) => text,
Ok(Message::Close(_)) => {
let _ = state.streaming_extractor.close_session(&session_id).await;
return;
}
Ok(Message::Ping(data)) => {
let _ = sender.send(Message::Pong(data)).await;
continue;
}
Ok(_) => continue,
Err(e) => {
tracing::warn!("WebSocket error: {}", e);
break;
}
};
let stream_msg: streaming::StreamMessage = match serde_json::from_str(&text) {
Ok(m) => m,
Err(e) => {
let error = streaming::ExtractionResult::Error {
code: "INVALID_MESSAGE".to_string(),
message: format!("Failed to parse message: {}", e),
fatal: false,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
continue;
}
};
let result = state
.streaming_extractor
.process_message(&session_id, stream_msg, user_memory.clone())
.await;
let response = serde_json::to_string(&result)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string());
if sender.send(Message::Text(response.into())).await.is_err() {
break;
}
if matches!(result, streaming::ExtractionResult::Closed { .. }) {
break;
}
}
if let Some(total) = state.streaming_extractor.close_session(&session_id).await {
tracing::info!(
"Streaming session {} closed. Total memories created: {}",
session_id,
total
);
}
}
pub async fn context_monitor_ws(
ws: WebSocketUpgrade,
State(state): State<AppState>,
) -> impl IntoResponse {
ws.on_upgrade(|socket| handle_context_monitor_socket(socket, state))
}
async fn handle_context_monitor_socket(socket: WebSocket, state: AppState) {
let (mut sender, mut receiver) = socket.split();
let mut user_id: Option<String> = None;
let mut config = relevance::RelevanceConfig::default();
let mut _debounce_ms: u64 = 100;
let mut last_surface_time = std::time::Instant::now();
while let Some(msg) = receiver.next().await {
let text = match msg {
Ok(Message::Text(text)) => text,
Ok(Message::Close(_)) => {
tracing::debug!("Context monitor WebSocket closed before handshake");
return;
}
Ok(_) => continue,
Err(e) => {
tracing::warn!("Context monitor WebSocket error before handshake: {}", e);
return;
}
};
let handshake: relevance::ContextMonitorHandshake = match serde_json::from_str(&text) {
Ok(h) => h,
Err(e) => {
let error = relevance::ContextMonitorResponse::Error {
code: "INVALID_HANDSHAKE".to_string(),
message: format!("Failed to parse handshake: {}", e),
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
};
if let Err(e) = validation::validate_user_id(&handshake.user_id) {
let error = relevance::ContextMonitorResponse::Error {
code: "INVALID_USER_ID".to_string(),
message: format!("Invalid user_id: {}", e),
fatal: true,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
return;
}
user_id = Some(handshake.user_id.clone());
if let Some(cfg) = handshake.config {
config = cfg;
}
_debounce_ms = handshake.debounce_ms;
let ack = relevance::ContextMonitorResponse::Ack {
timestamp: chrono::Utc::now(),
};
if sender
.send(Message::Text(
serde_json::to_string(&ack)
.unwrap_or_else(|_| r#"{"ack":true}"#.to_string())
.into(),
))
.await
.is_err()
{
return;
}
tracing::info!(
"Context monitor session started for user {}",
handshake.user_id
);
break;
}
let user_id = match user_id {
Some(id) => id,
None => return,
};
let memory_sys = match state.get_user_memory(&user_id) {
Ok(m) => m,
Err(e) => {
tracing::error!("Failed to get user memory: {}", e);
return;
}
};
let graph_memory = match state.get_user_graph(&user_id) {
Ok(g) => g,
Err(e) => {
tracing::error!("Failed to get user graph: {}", e);
return;
}
};
let ner = state.get_neural_ner();
let engine = std::sync::Arc::new(relevance::RelevanceEngine::new(ner));
while let Some(msg) = receiver.next().await {
let text = match msg {
Ok(Message::Text(text)) => text,
Ok(Message::Close(_)) => {
tracing::debug!("Context monitor closed by client");
return;
}
Ok(Message::Ping(data)) => {
let _ = sender.send(Message::Pong(data)).await;
continue;
}
Ok(_) => continue,
Err(e) => {
tracing::warn!("Context monitor WebSocket error: {}", e);
break;
}
};
let update: relevance::ContextUpdate = match serde_json::from_str(&text) {
Ok(u) => u,
Err(e) => {
let error = relevance::ContextMonitorResponse::Error {
code: "INVALID_MESSAGE".to_string(),
message: format!("Failed to parse context update: {}", e),
fatal: false,
timestamp: chrono::Utc::now(),
};
let _ = sender
.send(Message::Text(
serde_json::to_string(&error)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string())
.into(),
))
.await;
continue;
}
};
let elapsed = last_surface_time.elapsed().as_millis() as u64;
if elapsed < _debounce_ms {
continue;
}
last_surface_time = std::time::Instant::now();
let effective_config = update.config.unwrap_or_else(|| config.clone());
let response = {
let memory_sys = memory_sys.clone();
let graph_memory = graph_memory.clone();
let engine = engine.clone();
let context = update.context.clone();
tokio::task::spawn_blocking(move || {
let memory_guard = memory_sys.read();
let graph_guard = graph_memory.read();
engine.surface_relevant(
&context,
&memory_guard,
Some(&*graph_guard),
&effective_config,
None,
)
})
.await
};
let response = match response {
Ok(Ok(r)) => relevance::ContextMonitorResponse::Relevant {
memories: r.memories,
detected_entities: r.detected_entities,
latency_ms: r.latency_ms,
timestamp: chrono::Utc::now(),
},
Ok(Err(e)) => relevance::ContextMonitorResponse::Error {
code: "SURFACE_ERROR".to_string(),
message: e.to_string(),
fatal: false,
timestamp: chrono::Utc::now(),
},
Err(e) => relevance::ContextMonitorResponse::Error {
code: "TASK_PANIC".to_string(),
message: e.to_string(),
fatal: false,
timestamp: chrono::Utc::now(),
},
};
let response_json = serde_json::to_string(&response)
.unwrap_or_else(|_| r#"{"error":"serialization_failed"}"#.to_string());
if sender
.send(Message::Text(response_json.into()))
.await
.is_err()
{
break;
}
}
}