use std::fmt;
use std::time::Duration;
use super::checks::{Context, Outcome, stage};
use super::payloads;
use super::report::Check;
use super::{ProviderFactory, Scenario, WireFixtures};
use crate::error::{ProviderError, RetryClass};
use crate::provider::ModelProvider;
use crate::request::ModelRequest;
use crate::response::{FinishReason, ModelResponse};
use crate::secret::ApiKey;
const RESET_DEADLINE: Duration = Duration::from_secs(10);
const ACCEPT_WAIT: Duration = Duration::from_secs(1);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
#[non_exhaustive]
pub enum StatusRow {
Unauthorized,
Forbidden,
NotFound,
RequestTimeout,
ConnectionReset,
TooManyRequests,
ContextLength,
BadRequest,
InternalServerError,
ServiceUnavailable,
ContentFilter,
ExpiredCredential,
QuotaExhausted,
}
impl StatusRow {
pub const ALL: [Self; 13] = [
Self::Unauthorized,
Self::Forbidden,
Self::NotFound,
Self::RequestTimeout,
Self::ConnectionReset,
Self::TooManyRequests,
Self::ContextLength,
Self::BadRequest,
Self::InternalServerError,
Self::ServiceUnavailable,
Self::ContentFilter,
Self::ExpiredCredential,
Self::QuotaExhausted,
];
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::Unauthorized => "status_401_authentication",
Self::Forbidden => "status_403_authorization",
Self::NotFound => "status_404_model_not_found",
Self::RequestTimeout => "status_408_timeout",
Self::ConnectionReset => "connection_reset_timeout",
Self::TooManyRequests => "status_429_rate_limited",
Self::ContextLength => "status_400_context_overflow",
Self::BadRequest => "status_400_invalid_request",
Self::InternalServerError => "status_500_server",
Self::ServiceUnavailable => "status_503_server",
Self::ContentFilter => "content_filter_kind",
Self::ExpiredCredential => "expired_credential_kind",
Self::QuotaExhausted => "quota_exhausted_kind",
}
}
#[must_use]
pub const fn http_status(self) -> Option<u16> {
match self {
Self::Unauthorized => Some(401),
Self::Forbidden => Some(403),
Self::NotFound => Some(404),
Self::RequestTimeout => Some(408),
Self::TooManyRequests => Some(429),
Self::ContextLength | Self::BadRequest => Some(400),
Self::InternalServerError => Some(500),
Self::ServiceUnavailable => Some(503),
Self::ConnectionReset
| Self::ContentFilter
| Self::ExpiredCredential
| Self::QuotaExhausted => None,
}
}
#[must_use]
pub const fn scenario(self) -> Option<Scenario> {
match self {
Self::Unauthorized => Some(Scenario::Authentication),
Self::Forbidden => Some(Scenario::Authorization),
Self::NotFound => Some(Scenario::ModelNotFound),
Self::RequestTimeout => Some(Scenario::RequestTimeout),
Self::TooManyRequests => Some(Scenario::RateLimited),
Self::ContextLength => Some(Scenario::ContextOverflow),
Self::BadRequest => Some(Scenario::InvalidRequest),
Self::InternalServerError => Some(Scenario::ServerError),
Self::ServiceUnavailable => Some(Scenario::ServiceUnavailable),
Self::ContentFilter => Some(Scenario::ContentFilter),
Self::ExpiredCredential => Some(Scenario::ExpiredCredential),
Self::QuotaExhausted => Some(Scenario::QuotaExhausted),
Self::ConnectionReset => None,
}
}
#[must_use]
pub const fn expected_kinds(self) -> &'static [&'static str] {
match self {
Self::Unauthorized => &["authentication"],
Self::Forbidden => &["authorization"],
Self::NotFound => &["model_not_found"],
Self::RequestTimeout => &["timeout"],
Self::ConnectionReset => &["timeout", "transport"],
Self::TooManyRequests => &["rate_limited"],
Self::ContextLength => &["context_overflow"],
Self::BadRequest => &["invalid_request"],
Self::InternalServerError | Self::ServiceUnavailable => &["server"],
Self::ContentFilter => &["content_filter", "refusal"],
Self::ExpiredCredential => &["credential_expired"],
Self::QuotaExhausted => &["quota_exhausted"],
}
}
#[must_use]
pub const fn expected_retry_class(self) -> RetryClass {
match self {
Self::RequestTimeout | Self::ConnectionReset => RetryClass::Retry,
Self::InternalServerError | Self::ServiceUnavailable => RetryClass::Retry,
Self::TooManyRequests => RetryClass::RetryAfter,
Self::Unauthorized
| Self::Forbidden
| Self::NotFound
| Self::ExpiredCredential
| Self::QuotaExhausted => RetryClass::Fallback,
Self::ContextLength | Self::BadRequest | Self::ContentFilter => RetryClass::Fatal,
}
}
#[must_use]
pub const fn wire_description(self) -> &'static str {
match self {
Self::Unauthorized => "HTTP 401",
Self::Forbidden => "HTTP 403",
Self::NotFound => "HTTP 404",
Self::RequestTimeout => "HTTP 408",
Self::ConnectionReset => "a connection reset",
Self::TooManyRequests => "HTTP 429",
Self::ContextLength => "HTTP 400 (context length)",
Self::BadRequest => "HTTP 400 (bad request)",
Self::InternalServerError => "HTTP 500",
Self::ServiceUnavailable => "HTTP 503",
Self::ContentFilter => "the vendor's content-filter shape",
Self::ExpiredCredential => "the vendor's expired-credential signal",
Self::QuotaExhausted => "the vendor's quota or billing signal",
}
}
#[must_use]
pub const fn accepts_finish(self, finish: FinishReason) -> bool {
matches!(self, Self::ContentFilter)
&& matches!(finish, FinishReason::ContentFilter | FinishReason::Refusal)
}
fn request(self) -> ModelRequest {
match self {
Self::ContextLength | Self::ContentFilter => payloads::structured_request(),
Self::ExpiredCredential | Self::QuotaExhausted => payloads::structured_request(),
_ => payloads::narration_request(),
}
}
fn expected_label(self) -> String {
self.expected_kinds().join("|")
}
}
impl fmt::Display for StatusRow {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum RowSupport {
Mounted,
NotProducible {
reason: String,
},
}
impl RowSupport {
#[must_use]
pub fn not_producible(reason: impl Into<String>) -> Self {
Self::NotProducible {
reason: reason.into(),
}
}
#[must_use]
pub fn reason(&self) -> Option<&str> {
match self {
Self::Mounted => None,
Self::NotProducible { reason } => Some(reason.trim()),
}
}
}
pub type StatusSupport = RowSupport;
pub(super) async fn status_mapping<F: ProviderFactory, W: WireFixtures>(
factory: &F,
fixtures: &W,
row: StatusRow,
context: &mut Context,
) -> Outcome {
if let Some(outcome) = super::declared(
&fixtures.status_support(row),
row.wire_description(),
&format!("so {} is unproven", row.expected_label()),
) {
return outcome;
}
if row == StatusRow::ConnectionReset {
return connection_reset(factory, context).await;
}
let Some(scenario) = row.scenario() else {
return Outcome::Failed(format!(
"{} has no fixture scenario, which is a defect in the suite",
row.wire_description()
));
};
let (_server, provider) = match stage(factory, fixtures, scenario).await {
Ok(staged) => staged,
Err(reason) => return Outcome::Failed(reason),
};
let result = provider.generate(row.request()).await;
judge(row, result, context)
}
async fn connection_reset<F: ProviderFactory>(factory: &F, context: &mut Context) -> Outcome {
let row = StatusRow::ConnectionReset;
let listener = match tokio::net::TcpListener::bind(("127.0.0.1", 0)).await {
Ok(listener) => listener,
Err(error) => {
return Outcome::Failed(format!(
"the suite could not bind a local socket for {row}: {error}"
));
}
};
let address = match listener.local_addr() {
Ok(address) => address,
Err(error) => {
return Outcome::Failed(format!("the suite's local socket has no address: {error}"));
}
};
let closer = tokio::spawn(async move {
while let Ok((stream, _peer)) = listener.accept().await {
let mut scratch = [0_u8; 1];
let _ = tokio::time::timeout(ACCEPT_WAIT, stream.peek(&mut scratch)).await;
drop(stream);
}
});
let outcome = match factory.build(
&format!("http://{address}"),
ApiKey::new(payloads::DUMMY_API_KEY),
) {
Err(error) => Outcome::Failed(format!("adapter could not be built: {error}")),
Ok(provider) => {
match tokio::time::timeout(RESET_DEADLINE, provider.generate(row.request())).await {
Err(_elapsed) => Outcome::Failed(format!(
"{} left the call unresolved after {}s, expected {}",
row.wire_description(),
RESET_DEADLINE.as_secs(),
row.expected_label()
)),
Ok(result) => judge(row, result, context),
}
}
};
closer.abort();
outcome
}
fn confusion_hint(row: StatusRow, observed: &str) -> &'static str {
match (row, observed) {
(StatusRow::ExpiredCredential, "authentication") => {
"; an expired credential is not a rejected one — a caller holding a \
refresher could have refreshed and continued, and this tells it the key is bad"
}
(StatusRow::QuotaExhausted, "rate_limited") => {
"; a spent quota is not a rate limit — waiting will not refill it, so this \
makes the runtime sleep through its deadline for nothing, whatever status \
the provider used"
}
(StatusRow::ContextLength, "invalid_request") => {
"; a context overflow is fixed by shrinking the prompt, and a generic bad \
request is not fixed at all"
}
_ => "",
}
}
fn judge(
row: StatusRow,
result: Result<ModelResponse, ProviderError>,
context: &mut Context,
) -> Outcome {
match result {
Ok(response) => {
if row.accepts_finish(response.finish) {
Outcome::Passed
} else {
Outcome::Failed(format!(
"{} produced a response finishing as {}, expected {}",
row.wire_description(),
response.finish,
row.expected_label()
))
}
}
Err(error) => {
let kind = error.kind().clone();
context
.observed
.push((Check::StatusMapping(row), kind.clone()));
if !row.expected_kinds().contains(&kind.as_str()) {
let hint = confusion_hint(row, kind.as_str());
return Outcome::Failed(format!(
"{} mapped to {kind}, expected {}{hint}",
row.wire_description(),
row.expected_label()
));
}
let class = kind.retry_class();
if class != row.expected_retry_class() {
return Outcome::Failed(format!(
"{} mapped to {kind} classified as {class}, expected {} classified as {}",
row.wire_description(),
row.expected_label(),
row.expected_retry_class()
));
}
if row == StatusRow::TooManyRequests
&& kind.retry_after() != Some(Duration::from_secs(payloads::RETRY_AFTER_SECONDS))
{
return Outcome::Failed(format!(
"{} mapped to {kind}, expected rate_limited(retry_after={}s): \
the response advertised the delay and it must survive the mapping",
row.wire_description(),
payloads::RETRY_AFTER_SECONDS
));
}
Outcome::Passed
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::capabilities::ProviderCapabilities;
use crate::error::{ProviderErrorKind, RetryClass};
use crate::ids::RequestId;
fn context() -> Context {
Context::new(ProviderCapabilities::minimal())
}
#[test]
fn every_row_has_a_unique_label_and_an_expectation() {
let mut labels: Vec<&str> = StatusRow::ALL.iter().map(|row| row.as_str()).collect();
assert_eq!(labels.len(), 13);
labels.sort_unstable();
labels.dedup();
assert_eq!(labels.len(), 13, "labels must be unique");
for row in StatusRow::ALL {
assert!(!row.expected_kinds().is_empty(), "{row} expects nothing");
assert!(!row.wire_description().is_empty());
assert_eq!(row.scenario().is_none(), row == StatusRow::ConnectionReset);
}
assert_eq!(StatusRow::Forbidden.to_string(), "status_403_authorization");
}
#[test]
fn the_two_four_hundreds_are_told_apart_by_expectation_not_by_status() {
assert_eq!(StatusRow::ContextLength.http_status(), Some(400));
assert_eq!(StatusRow::BadRequest.http_status(), Some(400));
assert_ne!(
StatusRow::ContextLength.expected_kinds(),
StatusRow::BadRequest.expected_kinds()
);
assert_eq!(StatusRow::ConnectionReset.http_status(), None);
}
#[test]
fn every_expected_kind_keeps_the_retry_class_the_row_needs() {
let cases = [
(StatusRow::Unauthorized, RetryClass::Fallback),
(StatusRow::Forbidden, RetryClass::Fallback),
(StatusRow::NotFound, RetryClass::Fallback),
(StatusRow::RequestTimeout, RetryClass::Retry),
(StatusRow::TooManyRequests, RetryClass::RetryAfter),
(StatusRow::InternalServerError, RetryClass::Retry),
(StatusRow::ContextLength, RetryClass::Fatal),
(StatusRow::BadRequest, RetryClass::Fatal),
];
for (row, expected) in cases {
let kind = match row.expected_kinds()[0] {
"authentication" => ProviderErrorKind::Authentication,
"authorization" => ProviderErrorKind::Authorization,
"model_not_found" => ProviderErrorKind::ModelNotFound,
"timeout" => ProviderErrorKind::Timeout,
"rate_limited" => ProviderErrorKind::RateLimited { retry_after: None },
"server" => ProviderErrorKind::Server { status: None },
"context_overflow" => ProviderErrorKind::ContextOverflow {
needed_tokens: None,
limit_tokens: None,
},
"invalid_request" => ProviderErrorKind::InvalidRequest,
other => panic!("unmapped expectation {other}"),
};
assert_eq!(kind.retry_class(), expected, "{row}");
}
}
#[test]
fn a_wrong_mapping_names_the_status_the_expectation_and_what_arrived() {
let mut context = context();
let outcome = judge(
StatusRow::Forbidden,
Err(ProviderError::authentication()),
&mut context,
);
let Outcome::Failed(detail) = outcome else {
panic!("a 403 mapped to authentication must fail");
};
assert!(detail.contains("HTTP 403"), "{detail}");
assert!(detail.contains("expected authorization"), "{detail}");
assert!(detail.contains("mapped to authentication"), "{detail}");
assert_eq!(context.observed.len(), 1);
assert_eq!(
context.observed[0].0,
Check::StatusMapping(StatusRow::Forbidden)
);
}
#[test]
fn a_dropped_retry_after_fails_the_rate_limit_row() {
let mut context = context();
let kept = judge(
StatusRow::TooManyRequests,
Err(ProviderError::rate_limited(Some(Duration::from_secs(
payloads::RETRY_AFTER_SECONDS,
)))),
&mut context,
);
assert!(matches!(kept, Outcome::Passed));
let dropped = judge(
StatusRow::TooManyRequests,
Err(ProviderError::rate_limited(None)),
&mut context,
);
let Outcome::Failed(detail) = dropped else {
panic!("a dropped Retry-After must fail");
};
assert!(detail.contains("HTTP 429"), "{detail}");
assert!(detail.contains("retry_after=3s"), "{detail}");
}
#[test]
fn a_response_where_a_failure_was_required_fails_the_row() {
let mut context = context();
let response = ModelResponse::new(RequestId::nil(), "p", "m").with_text("hello");
let outcome = judge(StatusRow::InternalServerError, Ok(response), &mut context);
let Outcome::Failed(detail) = outcome else {
panic!("a 500 that produced an answer must fail");
};
assert!(detail.contains("HTTP 500"), "{detail}");
assert!(detail.contains("finishing as stop"), "{detail}");
assert!(context.observed.is_empty(), "no error, nothing to classify");
}
#[test]
fn a_filtered_answer_passes_whether_it_arrives_as_a_kind_or_as_a_finish() {
let mut context = context();
let filtered =
ModelResponse::new(RequestId::nil(), "p", "m").with_finish(FinishReason::ContentFilter);
assert!(matches!(
judge(StatusRow::ContentFilter, Ok(filtered), &mut context),
Outcome::Passed
));
assert!(matches!(
judge(
StatusRow::ContentFilter,
Err(ProviderError::content_filter()),
&mut context
),
Outcome::Passed
));
let as_server = judge(
StatusRow::ContentFilter,
Err(ProviderError::server(Some(500))),
&mut context,
);
let Outcome::Failed(detail) = as_server else {
panic!("a filter reported as a server error must fail");
};
assert!(detail.contains("content_filter|refusal"), "{detail}");
}
#[test]
fn a_reset_accepts_a_timeout_and_a_transport_failure_but_nothing_else() {
let mut context = context();
for error in [ProviderError::timeout(), ProviderError::transport("reset")] {
assert!(matches!(
judge(StatusRow::ConnectionReset, Err(error), &mut context),
Outcome::Passed
));
}
let wrong = judge(
StatusRow::ConnectionReset,
Err(ProviderError::other("mystery")),
&mut context,
);
let Outcome::Failed(detail) = wrong else {
panic!("an unclassified reset must fail");
};
assert!(detail.contains("a connection reset"), "{detail}");
assert!(detail.contains("timeout|transport"), "{detail}");
}
#[test]
fn every_rows_expected_kind_really_carries_the_class_the_row_declares() {
let representative = |label: &str| match label {
"authentication" => ProviderErrorKind::Authentication,
"authorization" => ProviderErrorKind::Authorization,
"model_not_found" => ProviderErrorKind::ModelNotFound,
"timeout" => ProviderErrorKind::Timeout,
"transport" => ProviderErrorKind::Transport,
"rate_limited" => ProviderErrorKind::RateLimited { retry_after: None },
"server" => ProviderErrorKind::Server { status: None },
"context_overflow" => ProviderErrorKind::ContextOverflow {
needed_tokens: None,
limit_tokens: None,
},
"invalid_request" => ProviderErrorKind::InvalidRequest,
"content_filter" => ProviderErrorKind::ContentFilter,
"refusal" => ProviderErrorKind::Refusal,
"credential_expired" => ProviderErrorKind::CredentialExpired,
"quota_exhausted" => ProviderErrorKind::QuotaExhausted { scope: None },
other => panic!("unmapped expectation {other}"),
};
for row in StatusRow::ALL {
for label in row.expected_kinds() {
assert_eq!(
representative(label).retry_class(),
row.expected_retry_class(),
"{row} accepts {label}, whose class disagrees with the row"
);
}
}
}
#[test]
fn an_expired_credential_reported_as_a_bad_key_fails_and_says_why() {
let mut context = context();
let outcome = judge(
StatusRow::ExpiredCredential,
Err(ProviderError::authentication()),
&mut context,
);
let Outcome::Failed(detail) = outcome else {
panic!("an expired credential reported as authentication must fail");
};
assert!(detail.contains("expired-credential signal"), "{detail}");
assert!(detail.contains("expected credential_expired"), "{detail}");
assert!(detail.contains("mapped to authentication"), "{detail}");
assert!(detail.contains("refresher"), "{detail}");
assert!(matches!(
judge(
StatusRow::ExpiredCredential,
Err(ProviderError::credential_expired()),
&mut context
),
Outcome::Passed
));
}
#[test]
fn a_spent_quota_reported_as_a_rate_limit_fails_and_says_why() {
let mut context = context();
let outcome = judge(
StatusRow::QuotaExhausted,
Err(ProviderError::rate_limited(Some(Duration::from_secs(60)))),
&mut context,
);
let Outcome::Failed(detail) = outcome else {
panic!("a quota reported as a rate limit must fail");
};
assert!(detail.contains("quota or billing signal"), "{detail}");
assert!(detail.contains("expected quota_exhausted"), "{detail}");
assert!(detail.contains("waiting will not refill it"), "{detail}");
assert!(matches!(
judge(
StatusRow::QuotaExhausted,
Err(ProviderError::quota_exhausted(Some("credit_balance"))),
&mut context
),
Outcome::Passed
));
}
#[test]
fn a_right_kind_with_a_drifted_class_still_fails() {
assert_ne!(
ProviderErrorKind::CredentialExpired.retry_class(),
RetryClass::Retry,
"an expired credential must never be a plain retry"
);
assert_eq!(
StatusRow::ExpiredCredential.expected_retry_class(),
RetryClass::Fallback
);
assert_eq!(
StatusRow::QuotaExhausted.expected_retry_class(),
RetryClass::Fallback
);
assert_ne!(
StatusRow::QuotaExhausted.expected_retry_class(),
StatusRow::TooManyRequests.expected_retry_class(),
"a quota must not be classified like a rate limit"
);
}
#[test]
fn not_producible_needs_a_reason() {
let explained = StatusSupport::not_producible("the endpoint never answers 408");
assert_eq!(
explained,
RowSupport::NotProducible {
reason: "the endpoint never answers 408".to_owned()
}
);
assert!(format!("{explained:?}").contains("NotProducible"));
assert_eq!(explained.reason(), Some("the endpoint never answers 408"));
assert_eq!(RowSupport::not_producible(" \n ").reason(), Some(""));
assert_eq!(RowSupport::Mounted.reason(), None);
}
}