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 super::protocol::outbound::anthropic::streaming::SseDecoder;
29
30pub use self::finalize::{FailCause, FinalizeDecision, classify};
31
32pub use self::abort::STREAM_ABORT_MESSAGE;
33
34/// Shared by the streaming and buffered completion tasks so both debit quota
35/// and run the response-phase safety scan identically.
36#[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/// How the tapped stream is rendered back to the caller.
46///
47/// `stream_usage` is the caller's own `stream_options.include_usage`; it
48/// decides whether the closing frames carry a usage chunk.
49#[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}