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