Skip to main content

systemprompt_api/services/gateway/stream_tap/
mod.rs

1//! Streaming response tap: re-renders upstream canonical events to the inbound
2//! wire format while accumulating a full response snapshot for the audit sink.
3//!
4//! Copyright (c) systemprompt.io — Business Source License 1.1.
5//! See <https://systemprompt.io> for licensing details.
6
7mod 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/// Shared by the streaming and buffered completion tasks so both debit quota
36/// and run the response-phase safety scan identically.
37#[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/// How the tapped stream is rendered back to the caller.
47///
48/// `stream_usage` is the caller's own `stream_options.include_usage`; it
49/// decides whether the closing frames carry a usage chunk.
50#[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}