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;
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/// Shared by the streaming and buffered completion tasks so both debit quota
34/// and run the response-phase safety scan identically.
35#[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/// How the tapped stream is rendered back to the caller.
44///
45/// `stream_usage` is the caller's own `stream_options.include_usage`; it
46/// decides whether the closing frames carry a usage chunk.
47#[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}