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;