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