Skip to main content

harness_gateway_client/wire/
read.rs

1//! The transport-independent half of reading a completion off the wire:
2//! the byte cap on a body, the SSE loop to the `[DONE]` sentinel, and the
3//! client-side timing, over a caller-supplied [`ChunkSource`].
4//!
5//! No HTTP happens here and no clock is read. The transport supplies the
6//! chunks and the clock; this module applies the one rule set every
7//! transport shares, so transports differ only in how they send. A
8//! transport that grew its own copy of this loop would be one more place
9//! the byte cap, the sentinel rule, and the timing arithmetic could drift.
10
11use std::future::Future;
12use std::time::{Duration, Instant};
13
14use promptforge::metrics::ClientTiming;
15use promptforge::model::{Completion, CompletionError};
16use serde_json::Value;
17
18use super::delta::StreamDelta;
19use super::stream::{Applied, SseScanner, StreamAccumulator};
20use crate::failure::malformed;
21
22/// A response body that a transport supplies one chunk at a time.
23///
24/// A transport is the code that sends a model request and receives the
25/// reply, for example over HTTP. It answers a `Chat` effect, the Engine's
26/// request for one model reply, in four steps:
27///
28/// 1. It builds the request body with
29///    [`build_request_body`](crate::build_request_body). It reads its
30///    clock and then sends the request.
31/// 2. It wraps the response body in a `ChunkSource`.
32/// 3. On a status outside the 2xx range, it reads the error body whole with
33///    [`read_body_capped`]. It bounds the body's length and escapes its
34///    control characters with [`escape_controls`](crate::escape_controls).
35///    It then fails the round with the error that
36///    [`classify_http_failure`](crate::classify_http_failure) returns.
37/// 4. Otherwise, it passes the source to [`read_completion_stream`]. The
38///    returned [`Completion`] answers the effect.
39///
40/// The transport opens the connection and supplies every clock reading.
41/// The source is the only I/O that `read_body_capped` and
42/// `read_completion_stream` touch.
43pub trait ChunkSource {
44    /// One chunk of body bytes, in whatever buffer the transport yields.
45    type Chunk: AsRef<[u8]>;
46
47    /// Returns the next chunk, or `None` once the body is exhausted.
48    ///
49    /// When a read fails, the implementation returns the
50    /// [`CompletionError`] that the round fails with. It is a
51    /// `Timeout`-kind error when the read ran out of time and a
52    /// `Transport`-kind error otherwise. Build it with
53    /// [`CompletionError::new`] and attach the transport's own error with
54    /// [`CompletionError::with_source`].
55    fn next_chunk(
56        &mut self,
57    ) -> impl Future<Output = Result<Option<Self::Chunk>, CompletionError>> + Send;
58}
59
60/// Reads a whole response body from `source`, refusing it once it would
61/// exceed `cap` bytes.
62///
63/// Use it for a body the transport decodes whole, such as the error body
64/// of a status outside the 2xx range or a JSON document like the
65/// gateway's model list.
66///
67/// `content_length` is the length the response advertises, when the
68/// transport knows it. An advertised length over `cap` fails at once,
69/// before any chunk is read. The function also counts bytes as the chunks
70/// arrive, so a gateway that omits or misstates the length still cannot
71/// force an unbounded allocation before the body is decoded.
72///
73/// # Errors
74/// Returns a `MalformedResponse`-kind [`CompletionError`] when the body
75/// would exceed `cap`, and the source's own error when a read fails.
76pub async fn read_body_capped<S: ChunkSource>(
77    source: &mut S,
78    content_length: Option<u64>,
79    cap: u64,
80) -> Result<Vec<u8>, CompletionError> {
81    if let Some(len) = content_length
82        && len > cap
83    {
84        return Err(malformed(format!(
85            "response body of {len} bytes exceeds the {cap}-byte limit"
86        )));
87    }
88    let mut body: Vec<u8> = Vec::new();
89    while let Some(chunk) = source.next_chunk().await? {
90        let bytes = chunk.as_ref();
91        if body.len() as u64 + bytes.len() as u64 > cap {
92            return Err(malformed(format!(
93                "response body exceeds the {cap}-byte limit"
94            )));
95        }
96        body.extend_from_slice(bytes);
97    }
98    Ok(body)
99}
100
101/// Reads a streamed model reply from `source` and assembles it into a
102/// [`Completion`].
103///
104/// The reply arrives as server-sent events (SSE). The function reads
105/// events until the `[DONE]` sentinel and fails if the stream exceeds
106/// `max_bytes` bytes. It passes each [`StreamDelta`] to `on_delta` as soon
107/// as it is decoded, so a Host can show the reply as it arrives. The
108/// returned completion holds the whole turn either way.
109///
110/// `request_body` is the body the transport sent, as
111/// [`build_request_body`](crate::build_request_body) returned it. The
112/// completion carries it back, so a run's debug capture records exactly
113/// what was sent. The completion is labeled with the model that
114/// `request_body` names.
115///
116/// `started` is the transport's clock reading from just before it sent
117/// the request, and `now` reads that same clock. The completion's
118/// [`ClientTiming`] holds three figures measured against them: time to
119/// first token, mean inter-token latency, and end-to-end time. It takes
120/// every clock reading from `started` and `now`.
121///
122/// # Errors
123/// Returns a `MalformedResponse`-kind [`CompletionError`] when the stream
124/// exceeds `max_bytes` or ends before the sentinel, and the source's own
125/// error when a read fails. Also returns the error that reassembling the
126/// reply raises for a malformed chunk, a mid-stream error envelope, a
127/// truncated tool-call batch, or an empty turn.
128pub async fn read_completion_stream<S: ChunkSource>(
129    source: &mut S,
130    request_body: Value,
131    max_bytes: u64,
132    on_delta: impl Fn(StreamDelta),
133    started: Instant,
134    now: impl Fn() -> Instant,
135) -> Result<Completion, CompletionError> {
136    let mut scanner = SseScanner::new();
137    let mut accumulator = StreamAccumulator::new();
138    let mut received: u64 = 0;
139    let mut first_delta: Option<Instant> = None;
140    let mut last_delta: Option<Instant> = None;
141    let mut delta_chunks: u32 = 0;
142    let mut done = false;
143    'read: while let Some(chunk) = source.next_chunk().await? {
144        let bytes = chunk.as_ref();
145        received += bytes.len() as u64;
146        if received > max_bytes {
147            return Err(malformed(format!(
148                "response stream exceeds the {max_bytes}-byte limit"
149            )));
150        }
151        scanner.extend(bytes);
152        while let Some(data) = scanner.next_data() {
153            match accumulator.apply(&data, &on_delta)? {
154                Applied::Done => {
155                    done = true;
156                    break 'read;
157                }
158                Applied::Chunk { delta: true } => {
159                    let at = now();
160                    first_delta.get_or_insert(at);
161                    last_delta = Some(at);
162                    delta_chunks += 1;
163                }
164                Applied::Chunk { delta: false } => {}
165            }
166        }
167    }
168    // A stream that ends without the sentinel was cut off; its
169    // accumulation may be missing the tail, so it must never pass for a
170    // complete turn.
171    if !done {
172        return Err(malformed(
173            "completion stream ended without the [DONE] sentinel",
174        ));
175    }
176    let client_timing = ClientTiming {
177        ttft_ms: first_delta.map(|at| duration_ms(at.duration_since(started))),
178        mean_itl_ms: match (first_delta, last_delta) {
179            (Some(first), Some(last)) if delta_chunks >= 2 => Some(round_to_microsecond(
180                duration_ms(last.duration_since(first)) / f64::from(delta_chunks - 1),
181            )),
182            _ => None,
183        },
184        e2e_ms: duration_ms(now().duration_since(started)),
185    };
186    // The truncation rule, the strict turn normalizer, and the lenient
187    // metadata parser all run inside `finish`: one rule set for every
188    // transport.
189    accumulator.finish(request_body, Some(client_timing))
190}
191
192/// A duration as fractional milliseconds, rounded to a whole microsecond
193/// so the text the run log stores parses back exactly.
194fn duration_ms(duration: Duration) -> f64 {
195    round_to_microsecond(duration.as_secs_f64() * 1000.0)
196}
197
198/// Rounds fractional-millisecond `ms` to the nearest whole microsecond.
199///
200/// A run log stores a timing as JSON text and parses that text back. A
201/// whole-microsecond value has a short decimal form the parser reproduces
202/// exactly, so the parsed timing equals the recorded one.
203fn round_to_microsecond(ms: f64) -> f64 {
204    (ms * 1000.0).round() / 1000.0
205}
206
207#[cfg(test)]
208#[path = "read-tests.rs"]
209mod tests;