Skip to main content

aurum_core/remote/
limits.rs

1//! Response body and transcript expansion bounds (JOE-1588).
2
3use crate::error::{ProviderError, Result};
4use crate::providers::Segment;
5use futures_util::StreamExt;
6use reqwest::Response;
7
8/// Default hard caps for response bodies (bytes).
9pub const DEFAULT_STT_BODY_CAP: usize = 8 * 1024 * 1024;
10pub const DEFAULT_CHAT_BODY_CAP: usize = 4 * 1024 * 1024;
11pub const DEFAULT_CLEANUP_BODY_CAP: usize = 2 * 1024 * 1024;
12
13/// Limits applied after JSON parse to transcript structures.
14#[derive(Debug, Clone, Copy)]
15pub struct TranscriptLimits {
16    pub max_text_chars: usize,
17    pub max_segments: usize,
18    pub max_segment_chars: usize,
19    pub max_total_segment_chars: usize,
20    /// Reject output text that expands beyond this multiple of input size.
21    pub max_expansion_ratio: f64,
22}
23
24impl Default for TranscriptLimits {
25    fn default() -> Self {
26        Self {
27            max_text_chars: 500_000,
28            max_segments: 10_000,
29            max_segment_chars: 8_000,
30            max_total_segment_chars: 1_000_000,
31            max_expansion_ratio: 8.0,
32        }
33    }
34}
35
36/// Endpoint-specific body byte caps.
37#[derive(Debug, Clone, Copy)]
38pub struct RemoteBodyLimits {
39    pub max_bytes: usize,
40}
41
42impl RemoteBodyLimits {
43    pub fn stt() -> Self {
44        Self {
45            max_bytes: DEFAULT_STT_BODY_CAP,
46        }
47    }
48    pub fn chat() -> Self {
49        Self {
50            max_bytes: DEFAULT_CHAT_BODY_CAP,
51        }
52    }
53    pub fn cleanup() -> Self {
54        Self {
55            max_bytes: DEFAULT_CLEANUP_BODY_CAP,
56        }
57    }
58}
59
60/// Stream a response body under a hard byte cap.
61///
62/// `Content-Length` is an early rejection check only; it never raises the cap.
63pub async fn read_body_limited(
64    response: Response,
65    provider: &str,
66    limits: RemoteBodyLimits,
67) -> Result<Vec<u8>> {
68    if let Some(cl) = response.content_length() {
69        if cl as usize > limits.max_bytes {
70            return Err(ProviderError::ResponseTooLarge {
71                provider: provider.into(),
72                reason: format!("Content-Length {cl} exceeds cap {}", limits.max_bytes),
73            }
74            .into());
75        }
76    }
77
78    let mut out = Vec::new();
79    let mut stream = response.bytes_stream();
80    while let Some(chunk) = stream.next().await {
81        let chunk = chunk.map_err(|e| ProviderError::Network {
82            provider: provider.into(),
83            reason: e.to_string(),
84        })?;
85        if out.len().saturating_add(chunk.len()) > limits.max_bytes {
86            return Err(ProviderError::ResponseTooLarge {
87                provider: provider.into(),
88                reason: format!("stream exceeded body cap of {} bytes", limits.max_bytes),
89            }
90            .into());
91        }
92        out.extend_from_slice(&chunk);
93    }
94    Ok(out)
95}
96
97/// Validate transcript text size and optional expansion vs input.
98pub fn validate_text_bounds(
99    text: &str,
100    input_chars: Option<usize>,
101    limits: TranscriptLimits,
102    provider: &str,
103) -> Result<()> {
104    let n = text.chars().count();
105    if n > limits.max_text_chars {
106        return Err(ProviderError::LimitExceeded {
107            reason: format!(
108                "transcript text has {n} chars (limit {})",
109                limits.max_text_chars
110            ),
111        }
112        .into());
113    }
114    if let Some(input) = input_chars {
115        if input > 0 {
116            let ratio = n as f64 / input as f64;
117            if ratio > limits.max_expansion_ratio {
118                return Err(ProviderError::InvalidProviderPayload {
119                    provider: provider.into(),
120                    reason: format!(
121                        "implausible expansion ratio {ratio:.1}x (limit {}x)",
122                        limits.max_expansion_ratio
123                    ),
124                }
125                .into());
126            }
127        }
128    }
129    Ok(())
130}
131
132/// Validate segment list: counts, lengths, finite ordered timestamps.
133pub fn validate_segments(
134    segments: &[Segment],
135    media_duration: f64,
136    limits: TranscriptLimits,
137    provider: &str,
138) -> Result<()> {
139    if segments.len() > limits.max_segments {
140        return Err(ProviderError::LimitExceeded {
141            reason: format!(
142                "{} segments exceeds limit {}",
143                segments.len(),
144                limits.max_segments
145            ),
146        }
147        .into());
148    }
149    let mut total_chars = 0usize;
150    let mut prev_end = 0.0f64;
151    for (i, seg) in segments.iter().enumerate() {
152        if !seg.start().is_finite() || !seg.end().is_finite() {
153            return Err(ProviderError::InvalidProviderPayload {
154                provider: provider.into(),
155                reason: format!("segment {i} has non-finite timestamps"),
156            }
157            .into());
158        }
159        if seg.start() < 0.0 || seg.end() < 0.0 {
160            return Err(ProviderError::InvalidProviderPayload {
161                provider: provider.into(),
162                reason: format!("segment {i} has negative timestamps"),
163            }
164            .into());
165        }
166        if seg.end() < seg.start() {
167            return Err(ProviderError::InvalidProviderPayload {
168                provider: provider.into(),
169                reason: format!("segment {i} end < start"),
170            }
171            .into());
172        }
173        // Soft bound: allow small overshoot past media duration.
174        let bound = if media_duration.is_finite() && media_duration > 0.0 {
175            media_duration + 1.0
176        } else {
177            f64::MAX
178        };
179        if seg.start() > bound || seg.end() > bound {
180            return Err(ProviderError::InvalidProviderPayload {
181                provider: provider.into(),
182                reason: format!("segment {i} timestamps exceed media duration"),
183            }
184            .into());
185        }
186        // Ordered with limited overlap (start may equal previous end).
187        if seg.start() + 0.001 < prev_end - 30.0 {
188            // Allow modest overlap but reject severe reordering.
189            return Err(ProviderError::InvalidProviderPayload {
190                provider: provider.into(),
191                reason: format!("segment {i} severely out of order"),
192            }
193            .into());
194        }
195        prev_end = seg.end().max(prev_end);
196
197        let sc = seg.text().chars().count();
198        if sc > limits.max_segment_chars {
199            return Err(ProviderError::LimitExceeded {
200                reason: format!(
201                    "segment {i} has {sc} chars (limit {})",
202                    limits.max_segment_chars
203                ),
204            }
205            .into());
206        }
207        total_chars = total_chars.saturating_add(sc);
208        if total_chars > limits.max_total_segment_chars {
209            return Err(ProviderError::LimitExceeded {
210                reason: format!(
211                    "total segment text exceeds {} chars",
212                    limits.max_total_segment_chars
213                ),
214            }
215            .into());
216        }
217    }
218    Ok(())
219}
220
221#[cfg(test)]
222mod tests {
223    use super::*;
224
225    #[test]
226    fn rejects_expansion() {
227        let err = validate_text_bounds(
228            "x".repeat(100).as_str(),
229            Some(5),
230            TranscriptLimits::default(),
231            "t",
232        )
233        .unwrap_err();
234        assert!(err.to_string().contains("expansion") || err.to_string().contains("limit"));
235    }
236
237    #[test]
238    fn rejects_nonfinite_segment() {
239        let segs = vec![Segment::from_parts_unchecked(
240            f64::NAN,
241            1.0,
242            "hi".to_string(),
243        )];
244        assert!(validate_segments(&segs, 10.0, TranscriptLimits::default(), "t").is_err());
245    }
246
247    #[test]
248    fn accepts_normal_segments() {
249        let segs = vec![
250            Segment::from_parts_unchecked(0.0, 1.0, "a".to_string()),
251            Segment::from_parts_unchecked(1.0, 2.0, "b".to_string()),
252        ];
253        validate_segments(&segs, 2.0, TranscriptLimits::default(), "t").unwrap();
254    }
255}