systemprompt_api/routes/stream/
mod.rs1use axum::Router;
13use axum::extract::Extension;
14use axum::response::IntoResponse;
15use axum::response::sse::Sse;
16use axum::routing::get;
17use std::convert::Infallible;
18use std::sync::{Arc, LazyLock};
19use systemprompt_agent::services::ContextProviderService;
20use systemprompt_events::{
21 A2A_BROADCASTER, AGUI_BROADCASTER, Broadcaster, ConnectionGuard, GenericBroadcaster, ToSse,
22 standard_keep_alive,
23};
24use systemprompt_models::RequestContext;
25use systemprompt_models::api::ApiError;
26use systemprompt_runtime::AppContext;
27use tokio::sync::mpsc;
28use tokio_stream::wrappers::ReceiverStream;
29
30use crate::error::ApiHttpError;
31
32pub mod contexts;
33
34#[derive(Clone, Debug)]
35pub struct StreamState {
36 pub context_provider: Arc<ContextProviderService>,
37}
38
39pub fn stream_router(ctx: &AppContext) -> Router {
40 let context_provider = ContextProviderService::new(ctx.a2a_repositories().contexts.clone());
41 let state = StreamState {
42 context_provider: Arc::new(context_provider),
43 };
44
45 Router::new()
46 .route("/contexts", get(contexts::stream_context_state))
47 .route("/agui", get(stream_agui_events))
48 .route("/a2a", get(stream_a2a_events))
49 .with_state(state)
50}
51
52pub async fn stream_a2a_events(
53 Extension(request_context): Extension<RequestContext>,
54) -> impl IntoResponse {
55 create_sse_stream(request_context, &A2A_BROADCASTER, "A2A").await
56}
57
58pub async fn stream_agui_events(
59 Extension(request_context): Extension<RequestContext>,
60) -> impl IntoResponse {
61 create_sse_stream(request_context, &AGUI_BROADCASTER, "AgUI").await
62}
63
64#[derive(Debug)]
65pub struct StreamWithGuard<E: ToSse + Clone + Send + Sync + 'static> {
66 stream: ReceiverStream<Result<axum::response::sse::Event, Infallible>>,
67 _cleanup_guard: ConnectionGuard<E>,
68}
69
70impl<E: ToSse + Clone + Send + Sync + 'static> StreamWithGuard<E> {
71 pub const fn new(
72 stream: ReceiverStream<Result<axum::response::sse::Event, Infallible>>,
73 cleanup_guard: ConnectionGuard<E>,
74 ) -> Self {
75 Self {
76 stream,
77 _cleanup_guard: cleanup_guard,
78 }
79 }
80}
81
82impl<E: ToSse + Clone + Send + Sync + 'static> futures_util::Stream for StreamWithGuard<E> {
83 type Item = Result<axum::response::sse::Event, Infallible>;
84
85 fn poll_next(
86 mut self: std::pin::Pin<&mut Self>,
87 cx: &mut std::task::Context<'_>,
88 ) -> std::task::Poll<Option<Self::Item>> {
89 std::pin::Pin::new(&mut self.stream).poll_next(cx)
90 }
91}
92
93pub async fn create_sse_stream<E: ToSse + Clone + Send + Sync + 'static>(
94 request_context: RequestContext,
95 broadcaster: &'static LazyLock<GenericBroadcaster<E>>,
96 stream_name: &str,
97) -> impl IntoResponse {
98 let user_id = request_context.user_id().clone();
99 let user_id_str = user_id.to_string();
100 let conn_id = systemprompt_identifiers::ConnectionId::generate();
101 let conn_id_str = conn_id.as_str().to_owned();
102
103 tracing::info!(user_id = %user_id_str, conn_id = %conn_id_str, stream = %stream_name, "SSE stream opened");
104
105 let (tx, rx) = mpsc::channel(1024);
106
107 if !broadcaster.register(&user_id, &conn_id, tx.clone()).await {
108 tracing::warn!(user_id = %user_id_str, stream = %stream_name, "SSE stream rejected: per-user connection cap reached");
109 return ApiHttpError::from(ApiError::rate_limited(
110 "Per-user stream connection limit reached",
111 ))
112 .into_response();
113 }
114
115 let cleanup_guard = ConnectionGuard::new(broadcaster, user_id, conn_id);
116 let stream = ReceiverStream::new(rx);
117 let stream_with_guard = StreamWithGuard::<E>::new(stream, cleanup_guard);
118
119 tracing::info!(user_id = %user_id_str, conn_id = %conn_id_str, stream = %stream_name, "SSE stream ready");
120
121 Sse::new(stream_with_guard)
122 .keep_alive(standard_keep_alive())
123 .into_response()
124}