use serde_json::Value;
use super::errors::{ConduitError, ErrorKind};
use super::execution::LLMCore;
use super::results::ErrorPayload;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AttemptDecision {
RetrySameModel,
TryNextModel,
}
#[derive(Debug, Clone)]
pub struct AttemptOutcome {
pub error: ConduitError,
pub decision: AttemptDecision,
}
fn has_http_status_pattern(lower: &str, code: &str) -> bool {
lower.contains(&format!("status {code}"))
|| lower.contains(&format!("status: {code}"))
|| lower.contains(&format!("http {code}"))
|| lower.contains(&format!("http/{code}"))
|| lower.contains(&format!("code {code}"))
|| lower.contains(&format!("code: {code}"))
|| lower.contains(&format!("error {code}"))
}
pub fn classify_by_text_signature(message: &str) -> Option<ErrorKind> {
let lower = message.to_lowercase();
if lower.contains("auth")
|| lower.contains("unauthorized")
|| lower.contains("api key")
|| lower.contains("invalid key")
{
return Some(ErrorKind::Config);
}
if lower.contains("rate limit")
|| has_http_status_pattern(&lower, "429")
|| lower.contains("quota")
{
return Some(ErrorKind::Temporary);
}
if lower.contains("not found") || has_http_status_pattern(&lower, "404") {
return Some(ErrorKind::NotFound);
}
if lower.contains("timeout") || lower.contains("timed out") {
return Some(ErrorKind::Temporary);
}
if lower.contains("server error")
|| has_http_status_pattern(&lower, "500")
|| has_http_status_pattern(&lower, "502")
|| has_http_status_pattern(&lower, "503")
{
return Some(ErrorKind::Temporary);
}
None
}
fn mask_sensitive(text: &str) -> String {
let mut out = String::with_capacity(text.len());
let mut i = 0;
let bytes = text.as_bytes();
while i < bytes.len() {
if text[i..].starts_with("Bearer ") {
out.push_str("Bearer [MASKED]");
i += 7; while i < bytes.len() && !bytes[i].is_ascii_whitespace() {
i += 1;
}
continue;
}
let prefixes = ["sk-", "key-", "token-"];
let mut matched = false;
for prefix in &prefixes {
if text[i..].starts_with(prefix) {
let start = i + prefix.len();
let mut end = start;
while end < bytes.len()
&& (bytes[end].is_ascii_alphanumeric()
|| bytes[end] == b'_'
|| bytes[end] == b'-')
{
end += 1;
}
if end - start >= 20 {
out.push_str("[MASKED_KEY]");
i = end;
matched = true;
break;
}
}
}
if matched {
continue;
}
out.push(bytes[i] as char);
i += 1;
}
out
}
impl LLMCore {
pub fn log_error(&self, error: &ConduitError, provider: &str, model: &str, attempt: u32) {
if self.verbose() == 0 {
return;
}
let prefix = format!(
"[{}:{}] attempt {}/{}",
provider,
model,
attempt + 1,
self.max_attempts()
);
let sanitized = mask_sensitive(&error.to_string());
if let Some(ref cause) = error.cause {
let sanitized_cause = mask_sensitive(&format!("{cause:?}"));
tracing::warn!(
"{} failed: {} (cause={})",
prefix,
sanitized,
sanitized_cause
);
} else {
tracing::warn!("{} failed: {}", prefix, sanitized);
}
}
pub fn classify_error(&self, error: &ConduitError) -> ErrorKind {
if let Some(kind) = self.custom_classify(error) {
return kind;
}
if let Some(kind) = classify_by_text_signature(&error.message) {
return kind;
}
error.kind
}
pub fn classify_http_status(status: u16) -> Option<ErrorKind> {
match status {
401 | 403 => Some(ErrorKind::Config),
400 | 404 | 413 | 422 => Some(ErrorKind::InvalidInput),
408 | 409 | 425 | 429 => Some(ErrorKind::Temporary),
s if (500..600).contains(&s) => Some(ErrorKind::Provider),
_ => None,
}
}
pub fn should_retry(kind: ErrorKind) -> bool {
matches!(kind, ErrorKind::Temporary | ErrorKind::Provider)
}
pub fn wrap_error(
&self,
kind: ErrorKind,
provider: &str,
model: &str,
message: &str,
) -> ConduitError {
ConduitError::new(kind, format!("{}:{}: {}", provider, model, message))
}
pub fn handle_attempt_error(
&self,
error: ConduitError,
provider_name: &str,
model_id: &str,
attempt: u32,
) -> AttemptOutcome {
let kind = self.classify_error(&error);
self.log_error(&error, provider_name, model_id, attempt);
let can_retry = Self::should_retry(kind) && (attempt + 1) < self.max_attempts();
let decision = if can_retry {
AttemptDecision::RetrySameModel
} else {
AttemptDecision::TryNextModel
};
AttemptOutcome { error, decision }
}
pub fn build_error_payload(
&self,
error: &ConduitError,
provider_name: &str,
model_id: &str,
attempt: u32,
http_status: Option<u16>,
) -> ErrorPayload {
let mut details = serde_json::json!({
"provider": provider_name,
"model": model_id,
"attempt": attempt + 1,
"max_attempts": self.max_attempts(),
});
if let Some(status) = http_status
&& let Value::Object(obj) = &mut details
{
obj.insert("http_status".to_owned(), Value::Number(status.into()));
}
ErrorPayload::new(error.kind, &error.message).with_details(details)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mask_bearer_token() {
let input = "Authorization: Bearer sk-abc123def456ghi789jkl0 failed";
let masked = mask_sensitive(input);
assert!(masked.contains("Bearer [MASKED]"));
assert!(!masked.contains("sk-abc123"));
}
#[test]
fn mask_sk_key() {
let input = "invalid api key: sk-proj-abcdefghijklmnopqrstuvwx";
let masked = mask_sensitive(input);
assert!(masked.contains("[MASKED_KEY]"));
assert!(!masked.contains("sk-proj-abcdefgh"));
}
#[test]
fn no_mask_short_prefix() {
let input = "sk-short is fine";
let masked = mask_sensitive(input);
assert_eq!(masked, input);
}
#[test]
fn no_mask_normal_text() {
let input = "rate limit exceeded, please retry";
let masked = mask_sensitive(input);
assert_eq!(masked, input);
}
#[test]
fn mask_multiple_keys() {
let input = "key-aaaaaaaaaaaaaaaaaaaaaaaaa and token-bbbbbbbbbbbbbbbbbbbbbbbbb";
let masked = mask_sensitive(input);
assert_eq!(
masked.matches("[MASKED_KEY]").count(),
2,
"should mask both keys: {masked}"
);
}
#[test]
fn classify_auth_error() {
assert_eq!(
classify_by_text_signature("unauthorized access"),
Some(ErrorKind::Config)
);
}
#[test]
fn classify_rate_limit() {
assert_eq!(
classify_by_text_signature("rate limit exceeded"),
Some(ErrorKind::Temporary)
);
}
#[test]
fn classify_no_match() {
assert_eq!(
classify_by_text_signature("something random happened"),
None
);
}
}