systemprompt_api/services/gateway/stream_tap/
mod.rs1mod abort;
8pub mod accumulator;
9mod finalize;
10
11use std::pin::Pin;
12use std::sync::{Arc, Mutex};
13use std::task::{Context, Poll};
14
15use axum::body::Body;
16use bytes::Bytes;
17use futures_util::stream::{BoxStream, Stream};
18use systemprompt_database::DbPool;
19use systemprompt_identifiers::AiRequestId;
20
21use self::accumulator::{Summary, TapState, accumulate_event, extract_summary, snapshot};
22use self::finalize::finalize;
23use super::audit::GatewayAudit;
24use super::policy::GatewayPolicySpec;
25use super::protocol::canonical_response::CanonicalEvent;
26use super::protocol::inbound::InboundAdapter;
27use super::protocol::outbound::anthropic::streaming::SseDecoder;
28
29pub use self::finalize::{FailCause, FinalizeDecision, classify};
30
31pub use self::abort::STREAM_ABORT_MESSAGE;
32
33#[derive(Debug)]
36pub struct TapFinalizeCtx {
37 pub db: DbPool,
38 pub repos: crate::services::gateway::GatewayRepositories,
39 pub policy: GatewayPolicySpec,
40 pub ai_request_id: AiRequestId,
41}
42
43#[derive(Debug)]
48pub struct TapRender {
49 pub inbound: Arc<dyn InboundAdapter>,
50 pub request_model: String,
51 pub stream_usage: bool,
52}
53
54pub fn tap(
55 upstream: BoxStream<'static, Result<CanonicalEvent, String>>,
56 render: TapRender,
57 audit: Arc<GatewayAudit>,
58 finalize_ctx: TapFinalizeCtx,
59) -> Body {
60 let TapRender {
61 inbound,
62 request_model,
63 stream_usage,
64 } = render;
65 let state = Arc::new(Mutex::new(TapState::default()));
66 let tapped = TappedStream {
67 inner: upstream,
68 state: Arc::clone(&state),
69 inbound,
70 request_model,
71 stream_usage,
72 audit,
73 finalize_ctx: Some(finalize_ctx),
74 message_stop_rendered: false,
75 ended: false,
76 };
77 Body::from_stream(tapped)
78}
79
80pub fn tap_raw(
81 upstream: BoxStream<'static, Result<Bytes, String>>,
82 inbound: Arc<dyn InboundAdapter>,
83 audit: Arc<GatewayAudit>,
84 finalize_ctx: TapFinalizeCtx,
85) -> Body {
86 Body::from_stream(RawTappedStream {
87 inner: upstream,
88 state: Arc::new(Mutex::new(TapState::default())),
89 decoder: SseDecoder::default(),
90 inbound,
91 audit,
92 finalize_ctx: Some(finalize_ctx),
93 ended: false,
94 })
95}
96
97struct RawTappedStream {
98 inner: BoxStream<'static, Result<Bytes, String>>,
99 state: Arc<Mutex<TapState>>,
100 decoder: SseDecoder,
101 inbound: Arc<dyn InboundAdapter>,
102 audit: Arc<GatewayAudit>,
103 finalize_ctx: Option<TapFinalizeCtx>,
104 ended: bool,
105}
106
107impl RawTappedStream {
108 fn take_summary(&mut self) -> Option<(Summary, TapFinalizeCtx)> {
109 let ctx = self.finalize_ctx.take()?;
110 self.state.lock().ok().and_then(|mut s| {
111 if s.finalized {
112 return None;
113 }
114 s.finalized = true;
115 Some((extract_summary(&mut s), ctx))
116 })
117 }
118}
119
120impl Stream for RawTappedStream {
121 type Item = Result<Bytes, std::io::Error>;
122
123 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
124 if self.ended {
125 return Poll::Ready(None);
126 }
127 match self.inner.as_mut().poll_next(cx) {
128 Poll::Pending => Poll::Pending,
129 Poll::Ready(None) => {
130 self.ended = true;
131 let Some((summary, ctx)) = self.take_summary() else {
132 return Poll::Ready(None);
133 };
134 let aborted = abort::is_abort(&summary);
135 finalize(Arc::clone(&self.audit), summary, ctx, "eof");
136 if !aborted {
137 return Poll::Ready(None);
138 }
139 Poll::Ready(abort::abort_frame(&self.inbound, "").map(Ok))
140 },
141 Poll::Ready(Some(Err(e))) => {
142 if let Ok(mut s) = self.state.lock() {
143 s.error = Some(e.clone());
144 }
145 Poll::Ready(Some(Err(std::io::Error::new(
146 std::io::ErrorKind::BrokenPipe,
147 e,
148 ))))
149 },
150 Poll::Ready(Some(Ok(bytes))) => {
151 let events = self.decoder.push(&bytes);
152 if let Ok(mut s) = self.state.lock() {
153 for event in &events {
154 accumulate_event(&mut s, event);
155 }
156 s.final_bytes.extend_from_slice(&bytes);
157 }
158 Poll::Ready(Some(Ok(bytes)))
159 },
160 }
161 }
162}
163
164impl Drop for RawTappedStream {
165 fn drop(&mut self) {
166 let Some((summary, ctx)) = self.take_summary() else {
167 return;
168 };
169 finalize(Arc::clone(&self.audit), summary, ctx, "drop");
170 }
171}
172
173struct TappedStream {
174 inner: BoxStream<'static, Result<CanonicalEvent, String>>,
175 state: Arc<Mutex<TapState>>,
176 inbound: Arc<dyn InboundAdapter>,
177 request_model: String,
178 stream_usage: bool,
179 audit: Arc<GatewayAudit>,
180 finalize_ctx: Option<TapFinalizeCtx>,
181 message_stop_rendered: bool,
182 ended: bool,
183}
184
185impl Stream for TappedStream {
186 type Item = Result<Bytes, std::io::Error>;
187
188 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
189 if self.ended {
190 return Poll::Ready(None);
191 }
192 loop {
193 match self.inner.as_mut().poll_next(cx) {
194 Poll::Pending => return Poll::Pending,
195 Poll::Ready(None) => {
196 self.ended = true;
197 return self.finalize_on_eof();
198 },
199 Poll::Ready(Some(Err(e))) => {
200 if let Ok(mut s) = self.state.lock() {
201 s.error = Some(e.clone());
202 }
203 let err = std::io::Error::new(std::io::ErrorKind::BrokenPipe, e);
204 return Poll::Ready(Some(Err(err)));
205 },
206 Poll::Ready(Some(Ok(event))) => {
207 let is_message_stop = matches!(event, CanonicalEvent::MessageStop { .. });
208 let terminal = matches!(event, CanonicalEvent::ContentBlockStop { .. })
209 || (is_message_stop && !self.message_stop_rendered);
210 let snap = self.state.lock().map_or(None, |mut s| {
211 accumulate_event(&mut s, &event);
212 terminal.then(|| snapshot(&s))
213 });
214 let terminal_suppressed = is_message_stop && self.message_stop_rendered;
215 if is_message_stop {
216 self.message_stop_rendered = true;
217 }
218 let rendered = snap
219 .as_ref()
220 .and_then(|snapshot| {
221 self.inbound.render_terminal_event(
222 &event,
223 snapshot,
224 &self.request_model,
225 )
226 })
227 .or_else(|| {
228 (!terminal_suppressed)
229 .then(|| self.inbound.render_event(&event, &self.request_model))
230 .flatten()
231 });
232 if let Some(bytes) = rendered {
233 if let Ok(mut s) = self.state.lock() {
234 s.final_bytes.extend_from_slice(&bytes);
235 }
236 return Poll::Ready(Some(Ok(bytes)));
237 }
238 },
239 }
240 }
241 }
242}
243
244impl TappedStream {
245 fn take_summary(&mut self) -> Option<(Summary, TapFinalizeCtx)> {
246 let ctx = self.finalize_ctx.take()?;
247 self.state.lock().ok().and_then(|mut s| {
248 if s.finalized {
249 return None;
250 }
251 s.finalized = true;
252 Some((extract_summary(&mut s), ctx))
253 })
254 }
255
256 fn finalize_on_eof(&mut self) -> Poll<Option<Result<Bytes, std::io::Error>>> {
257 let Some((summary, ctx)) = self.take_summary() else {
258 return Poll::Ready(None);
259 };
260 let aborted = abort::is_abort(&summary);
261 let tail = (!aborted)
262 .then(|| abort::tail_frames(&self.inbound, &summary.response, self.stream_usage))
263 .flatten();
264 finalize(Arc::clone(&self.audit), summary, ctx, "eof");
265 if !aborted {
266 return Poll::Ready(tail.map(Ok));
267 }
268 Poll::Ready(abort::abort_frame(&self.inbound, &self.request_model).map(Ok))
269 }
270}
271
272impl Drop for TappedStream {
273 fn drop(&mut self) {
274 let Some((summary, ctx)) = self.take_summary() else {
275 return;
276 };
277 finalize(Arc::clone(&self.audit), summary, ctx, "drop");
278 }
279}