harness_gateway_client/wire/
classify.rs1use promptforge::model::{CompletionError, CompletionErrorKind};
22
23const OVERFLOW_PHRASES: &[&str] = &[
27 "context length",
28 "context window",
29 "context size",
30 "context_length_exceeded",
31 "maximum context length",
32 "prompt is too long",
33 "too many tokens",
34 "exceeds the available context size",
35 "exceed_context_size",
36 "input is too long",
37 "exceeds the maximum number of tokens",
38 "too large for model",
39];
40
41const QUOTA_WORDS: &[&str] = &["quota", "billing", "insufficient_quota", "credit"];
43
44const REFUSAL_WORDS: &[&str] = &["content_filter", "content policy", "safety", "refus"];
46
47const CREDENTIALS_PHRASE: &str = "the model backend did not accept the credentials";
50
51#[must_use]
68pub fn classify_http_failure(status: u16, body: &str) -> CompletionError {
69 let lower = body.to_lowercase();
70 let http = |kind: CompletionErrorKind| {
71 http_error(kind, kind.phrase(), status).with_detail(body.to_owned())
72 };
73 if matches!(status, 400 | 413) && names_any(&lower, OVERFLOW_PHRASES) {
74 let (prompt_tokens, window) = overflow_counts(&lower);
75 return CompletionError::context_overflow(
76 prompt_tokens,
77 window,
78 format!(
79 "{} (status {status})",
80 CompletionErrorKind::ContextOverflow.phrase()
81 ),
82 )
83 .with_detail(body.to_owned());
84 }
85 match status {
86 429 if names_any(&lower, QUOTA_WORDS) => http(CompletionErrorKind::QuotaExhausted),
87 429 => http(CompletionErrorKind::RateLimited),
88 503 | 529 => http(CompletionErrorKind::Overloaded),
89 500..=599 if lower.contains("overloaded") => http(CompletionErrorKind::Overloaded),
90 500..=599 => http(CompletionErrorKind::ServerError),
91 401 | 403 => http_error(CompletionErrorKind::Unavailable, CREDENTIALS_PHRASE, status)
92 .with_detail(body.to_owned()),
93 400 if names_any(&lower, REFUSAL_WORDS) => http(CompletionErrorKind::Refused),
94 _ => http(CompletionErrorKind::Rejected),
95 }
96}
97
98#[must_use]
111pub fn classify_stream_error(body: &str) -> CompletionError {
112 let lower = body.to_lowercase();
113 let error = if names_any(&lower, OVERFLOW_PHRASES) {
114 let (prompt_tokens, window) = overflow_counts(&lower);
115 CompletionError::context_overflow(
116 prompt_tokens,
117 window,
118 CompletionErrorKind::ContextOverflow.phrase(),
119 )
120 } else {
121 let kind = if names_any(&lower, QUOTA_WORDS) {
122 CompletionErrorKind::QuotaExhausted
123 } else if lower.contains("overloaded") {
124 CompletionErrorKind::Overloaded
125 } else if names_any(&lower, REFUSAL_WORDS) {
126 CompletionErrorKind::Refused
127 } else {
128 CompletionErrorKind::Transport
129 };
130 CompletionError::new(kind, kind.phrase())
131 };
132 error.with_detail(body.to_owned())
133}
134
135fn http_error(kind: CompletionErrorKind, phrase: &str, status: u16) -> CompletionError {
137 CompletionError::new(kind, format!("{phrase} (status {status})"))
138}
139
140fn names_any(lower: &str, words: &[&str]) -> bool {
141 words.iter().any(|word| lower.contains(word))
142}
143
144fn overflow_counts(lower: &str) -> (Option<u32>, Option<u32>) {
150 if let Some(counts) = counts_from_maximum(lower) {
151 return counts;
152 }
153 counts_from_greater_than(lower).unwrap_or((None, None))
154}
155
156fn counts_from_maximum(lower: &str) -> Option<(Option<u32>, Option<u32>)> {
157 const LEAD: &str = "maximum context length is ";
158 let start = lower.find(LEAD)? + LEAD.len();
159 let (window, rest) = leading_number(&lower[start..])?;
160 Some((first_count_before_tokens(rest), window))
161}
162
163fn counts_from_greater_than(lower: &str) -> Option<(Option<u32>, Option<u32>)> {
164 const MID: &str = " tokens > ";
165 let at = lower.find(MID)?;
166 let before = &lower[..at];
167 let prompt = &before[before.trim_end_matches(|c: char| c.is_ascii_digit()).len()..];
168 if prompt.is_empty() {
169 return None;
170 }
171 let (window, rest) = leading_number(&lower[at + MID.len()..])?;
172 rest.starts_with(" maximum")
173 .then(|| (prompt.parse().ok(), window))
174}
175
176fn leading_number(text: &str) -> Option<(Option<u32>, &str)> {
180 let end = text
181 .find(|c: char| !c.is_ascii_digit())
182 .unwrap_or(text.len());
183 if end == 0 {
184 return None;
185 }
186 Some((text[..end].parse().ok(), &text[end..]))
187}
188
189fn first_count_before_tokens(text: &str) -> Option<u32> {
191 let mut rest = text;
192 while let Some(start) = rest.find(|c: char| c.is_ascii_digit()) {
193 let (number, after) = leading_number(&rest[start..])?;
194 if after.starts_with(" tokens") {
195 return number;
196 }
197 rest = after;
198 }
199 None
200}
201
202#[cfg(test)]
203#[path = "classify-tests.rs"]
204mod tests;