mod error;
mod handlers;
use axum::Json;
use axum::extract::{Path, State};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use serde::{Deserialize, Serialize};
use serde_json::json;
use systemprompt_agent::repository::context::ContextRepository;
use systemprompt_identifiers::{ContextId, UserId};
use systemprompt_runtime::AppContext;
use crate::error::ApiHttpError;
use crate::routes::agent::parse_context_id;
use handlers::{
broadcast_notification, mark_notification_broadcasted, persist_notification,
process_notification,
};
#[derive(Debug, Deserialize, Serialize)]
pub struct A2aNotification {
pub jsonrpc: String,
pub method: String,
pub params: serde_json::Value,
}
pub async fn handle_context_notification(
Path(context_id): Path<String>,
State(app_context): State<AppContext>,
Json(notification): Json<A2aNotification>,
) -> Result<Response, ApiHttpError> {
let repos = app_context.a2a_repositories();
let ctx_repo = &repos.contexts;
let context_id = parse_context_id(&context_id)?;
tracing::debug!(context_id = %context_id, method = %notification.method, "Received notification for context");
let user_id = resolve_context_user(ctx_repo, &context_id).await?;
if notification.jsonrpc != "2.0" {
tracing::error!(jsonrpc_version = %notification.jsonrpc, "Invalid JSON-RPC version");
return Err(ApiHttpError::bad_request(
"Invalid JSON-RPC version, must be 2.0",
));
}
let agent_id = notification
.params
.get("agentId")
.and_then(|v| v.as_str())
.unwrap_or("unknown")
.to_owned();
let notification_id = persist_notification(
&repos.context_notifications,
&context_id,
&agent_id,
¬ification,
)
.await?;
tracing::debug!(notification_id = %notification_id, context_id = %context_id, "Persisted notification");
process_notification(app_context.clone(), ¬ification).await?;
broadcast_and_mark(
&app_context,
&context_id,
&user_id,
¬ification,
notification_id,
)
.await;
Ok((
StatusCode::OK,
Json(json!({
"status": "received",
"notification_id": notification_id
})),
)
.into_response())
}
async fn resolve_context_user(
ctx_repo: &ContextRepository,
context_id: &ContextId,
) -> Result<UserId, ApiHttpError> {
ctx_repo
.find_user_id_for_context(context_id)
.await?
.ok_or_else(|| ApiHttpError::not_found(format!("Context '{context_id}' not found")))
}
async fn broadcast_and_mark(
app_context: &AppContext,
context_id: &ContextId,
user_id: &UserId,
notification: &A2aNotification,
notification_id: i32,
) {
let broadcast_count = broadcast_notification(
app_context.event_router(),
context_id.as_str(),
user_id,
notification,
)
.await;
tracing::debug!(broadcast_count = %broadcast_count, context_id = %context_id, "Broadcasted notification to streams");
let notifications_repo = &app_context.a2a_repositories().context_notifications;
if let Err(e) = mark_notification_broadcasted(notifications_repo, notification_id).await {
tracing::error!(error = %e, notification_id = %notification_id, "Failed to mark notification as broadcasted");
}
}