use std::time::Duration;
use crate::types::RequestMetadata;
pub type Result<T> = std::result::Result<T, Error>;
#[allow(missing_docs)]
pub mod codes {
pub const LORA_LOADING: &str = "LORA_LOADING";
pub const MODEL_LOADING: &str = "MODEL_LOADING";
pub const PROVISIONING: &str = "PROVISIONING";
pub const MODEL_LOAD_FAILED: &str = "MODEL_LOAD_FAILED";
pub const INPUT_TOO_LONG: &str = "INPUT_TOO_LONG";
pub const RESOURCE_EXHAUSTED: &str = "RESOURCE_EXHAUSTED";
pub const QUEUE_UNAVAILABLE: &str = "QUEUE_UNAVAILABLE";
pub const INTERNAL_ERROR: &str = "INTERNAL_ERROR";
pub const ENCODE_RESULT_COUNT_MISMATCH: &str = "ENCODE_RESULT_COUNT_MISMATCH";
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelLoadErrorClass {
Gated,
Oom,
Dependency,
NotFound,
Network,
Unknown,
}
impl ModelLoadErrorClass {
fn from_wire(value: &str) -> Self {
match value {
"GATED" => Self::Gated,
"OOM" => Self::Oom,
"DEPENDENCY" => Self::Dependency,
"NOT_FOUND" => Self::NotFound,
"NETWORK" => Self::Network,
_ => Self::Unknown,
}
}
pub(crate) fn parse(value: Option<&str>) -> Self {
value.map_or(Self::Unknown, Self::from_wire)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TransportErrorKind {
Connect,
MidFlight,
Timeout,
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
#[allow(missing_docs)]
pub enum Error {
#[error("{message}")]
Connection {
message: String,
kind: TransportErrorKind,
#[source]
source: Option<Box<dyn std::error::Error + Send + Sync>>,
},
#[error("{message}")]
Request {
message: String,
code: Option<String>,
status: u16,
request: Option<Box<RequestMetadata>>,
},
#[error("{message}")]
Server {
message: String,
code: Option<String>,
status: u16,
request: Option<Box<RequestMetadata>>,
},
#[error("{message}")]
InputTooLong {
message: String,
model: Option<String>,
request: Option<Box<RequestMetadata>>,
},
#[error("{message}")]
ModelLoadFailed {
message: String,
model: Option<String>,
error_class: ModelLoadErrorClass,
permanent: bool,
attempts: u32,
request: Option<Box<RequestMetadata>>,
},
#[error("{message}")]
ResourceExhausted {
message: String,
model: Option<String>,
retries: u32,
request: Option<Box<RequestMetadata>>,
},
#[error("{message}")]
EstimateUnroutable {
message: String,
code: Option<String>,
request: Option<Box<RequestMetadata>>,
},
#[error("{message}")]
Provisioning {
message: String,
gpu: Option<String>,
retry_after: Option<Duration>,
},
#[error("{message}")]
ModelLoading {
message: String,
model: Option<String>,
},
#[error("{message}")]
LoraLoading {
message: String,
lora: Option<String>,
model: Option<String>,
},
#[error("{message}")]
Pool {
message: String,
pool_name: Option<String>,
state: Option<String>,
},
#[error("{0}")]
Decode(String),
#[error("{0}")]
InvalidRequest(String),
#[error(transparent)]
Io(#[from] std::io::Error),
}
impl Error {
pub(crate) fn decode(message: impl Into<String>) -> Self {
Self::Decode(message.into())
}
pub(crate) fn invalid(message: impl Into<String>) -> Self {
Self::InvalidRequest(message.into())
}
pub(crate) fn connection(
kind: TransportErrorKind,
message: impl Into<String>,
source: impl std::error::Error + Send + Sync + 'static,
) -> Self {
Self::Connection {
message: message.into(),
kind,
source: Some(Box::new(source)),
}
}
pub fn status(&self) -> Option<u16> {
match self {
Self::Request { status, .. } | Self::Server { status, .. } => Some(*status),
Self::InputTooLong { .. } => Some(400),
Self::ModelLoadFailed { .. } => Some(502),
Self::ResourceExhausted { .. }
| Self::EstimateUnroutable { .. }
| Self::Provisioning { .. }
| Self::ModelLoading { .. } => Some(503),
_ => None,
}
}
pub fn code(&self) -> Option<&str> {
match self {
Self::Request { code, .. }
| Self::Server { code, .. }
| Self::EstimateUnroutable { code, .. } => code.as_deref(),
Self::InputTooLong { .. } => Some(codes::INPUT_TOO_LONG),
Self::ModelLoadFailed { .. } => Some(codes::MODEL_LOAD_FAILED),
Self::ResourceExhausted { .. } => Some(codes::RESOURCE_EXHAUSTED),
Self::Provisioning { .. } => Some(codes::PROVISIONING),
Self::ModelLoading { .. } => Some(codes::MODEL_LOADING),
Self::LoraLoading { .. } => Some(codes::LORA_LOADING),
_ => None,
}
}
pub fn request_metadata(&self) -> Option<&RequestMetadata> {
#[allow(clippy::borrowed_box)]
match self {
Self::Request { request, .. }
| Self::Server { request, .. }
| Self::InputTooLong { request, .. }
| Self::ModelLoadFailed { request, .. }
| Self::ResourceExhausted { request, .. }
| Self::EstimateUnroutable { request, .. } => request.as_deref(),
_ => None,
}
}
pub fn retry_after(&self) -> Option<Duration> {
match self {
Self::Provisioning { retry_after, .. } => *retry_after,
_ => None,
}
}
pub fn is_server_error(&self) -> bool {
matches!(
self,
Self::Server { .. }
| Self::ModelLoadFailed { .. }
| Self::ResourceExhausted { .. }
| Self::EstimateUnroutable { .. }
)
}
pub fn is_request_error(&self) -> bool {
matches!(self, Self::Request { .. } | Self::InputTooLong { .. })
}
pub fn is_capacity_error(&self) -> bool {
matches!(
self,
Self::Provisioning { .. }
| Self::ModelLoading { .. }
| Self::LoraLoading { .. }
| Self::ResourceExhausted { .. }
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn predicates_follow_the_python_hierarchy() {
let exhausted = Error::ResourceExhausted {
message: "boom".into(),
model: None,
retries: 3,
request: None,
};
assert!(exhausted.is_server_error());
assert!(exhausted.is_capacity_error());
assert_eq!(exhausted.status(), Some(503));
assert_eq!(exhausted.code(), Some(codes::RESOURCE_EXHAUSTED));
let too_long = Error::InputTooLong {
message: "too long".into(),
model: Some("m".into()),
request: None,
};
assert!(too_long.is_request_error());
assert!(!too_long.is_server_error());
assert_eq!(too_long.status(), Some(400));
}
#[test]
fn model_load_error_class_defaults_to_unknown() {
assert_eq!(
ModelLoadErrorClass::parse(Some("GATED")),
ModelLoadErrorClass::Gated
);
assert_eq!(
ModelLoadErrorClass::parse(Some("nonsense")),
ModelLoadErrorClass::Unknown
);
assert_eq!(
ModelLoadErrorClass::parse(None),
ModelLoadErrorClass::Unknown
);
}
}