1use std::{error::Error as StdError, fmt, num::ParseIntError, sync::Arc};
6
7use as_variant::as_variant;
8use bytes::{BufMut, Bytes};
9use serde::{Deserialize, Serialize};
10use serde_json::{Value as JsonValue, from_slice as from_json_slice};
11use thiserror::Error;
12
13mod kind;
14mod kind_serde;
15#[cfg(test)]
16mod tests;
17
18pub use self::kind::*;
19use super::{EndpointError, MatrixVersion, OutgoingResponse};
20
21#[derive(Clone, Debug)]
23#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
24pub struct Error {
25 pub status_code: http::StatusCode,
27
28 pub body: ErrorBody,
30}
31
32impl Error {
33 pub fn new(status_code: http::StatusCode, body: ErrorBody) -> Self {
37 Self { status_code, body }
38 }
39
40 pub fn error_kind(&self) -> Option<&ErrorKind> {
42 as_variant!(&self.body, ErrorBody::Standard(StandardErrorBody { kind, .. }) => kind)
43 }
44
45 pub fn is_endpoint_not_implemented(&self) -> bool {
53 self.status_code == http::StatusCode::NOT_FOUND
54 && self
55 .error_kind()
56 .is_some_and(|error_kind| matches!(error_kind, ErrorKind::Unrecognized))
57 }
58}
59
60impl fmt::Display for Error {
61 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
62 let status_code = self.status_code.as_u16();
63 match &self.body {
64 ErrorBody::Standard(StandardErrorBody { kind, message }) => {
65 let errcode = kind.errcode();
66 write!(f, "[{status_code} / {errcode}] {message}")
67 }
68 ErrorBody::Json(json) => write!(f, "[{status_code}] {json}"),
69 ErrorBody::NotJson { .. } => write!(f, "[{status_code}] <non-json bytes>"),
70 }
71 }
72}
73
74impl StdError for Error {}
75
76impl OutgoingResponse for Error {
77 fn try_into_http_response<T: Default + BufMut>(
78 self,
79 ) -> Result<http::Response<T>, IntoHttpError> {
80 let mut builder = http::Response::builder()
81 .header(http::header::CONTENT_TYPE, ruma_common::http_headers::APPLICATION_JSON)
82 .status(self.status_code);
83
84 if let Some(ErrorKind::LimitExceeded(LimitExceededErrorData {
86 retry_after: Some(retry_after),
87 })) = self.error_kind()
88 {
89 let header_value = http::HeaderValue::try_from(retry_after)?;
90 builder = builder.header(http::header::RETRY_AFTER, header_value);
91 }
92
93 builder
94 .body(match self.body {
95 ErrorBody::Standard(standard_body) => {
96 ruma_common::serde::json_to_buf(&standard_body)?
97 }
98 ErrorBody::Json(json) => ruma_common::serde::json_to_buf(&json)?,
99 ErrorBody::NotJson { .. } => {
100 return Err(IntoHttpError::Json(serde::ser::Error::custom(
101 "attempted to serialize ErrorBody::NotJson",
102 )));
103 }
104 })
105 .map_err(Into::into)
106 }
107}
108
109impl EndpointError for Error {
110 fn from_http_response(response: http::Response<&[u8]>) -> Self {
111 let status = response.status();
112
113 let body_bytes = response.body();
114 let error_body: ErrorBody = match from_json_slice::<StandardErrorBody>(body_bytes) {
115 Ok(mut standard_body) => {
116 let headers = response.headers();
117
118 if let ErrorKind::LimitExceeded(LimitExceededErrorData { retry_after }) =
119 &mut standard_body.kind
120 {
121 if let Some(Ok(retry_after_header)) =
124 headers.get(http::header::RETRY_AFTER).map(RetryAfter::try_from)
125 {
126 *retry_after = Some(retry_after_header);
127 }
128 }
129
130 ErrorBody::Standard(standard_body)
131 }
132 Err(_) => match from_json_slice(body_bytes) {
133 Ok(json) => ErrorBody::Json(json),
134 Err(error) => ErrorBody::NotJson {
135 bytes: Bytes::copy_from_slice(body_bytes),
136 deserialization_error: Arc::new(error),
137 },
138 },
139 };
140
141 error_body.into_error(status)
142 }
143}
144
145#[derive(Debug, Clone)]
147#[allow(clippy::exhaustive_enums)]
148pub enum ErrorBody {
149 Standard(StandardErrorBody),
151
152 Json(JsonValue),
154
155 NotJson {
157 bytes: Bytes,
159
160 deserialization_error: Arc<serde_json::Error>,
162 },
163}
164
165impl ErrorBody {
166 pub fn into_error(self, status_code: http::StatusCode) -> Error {
170 Error { status_code, body: self }
171 }
172}
173
174#[derive(Clone, Debug, Deserialize, Serialize)]
176#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
177pub struct StandardErrorBody {
178 #[serde(flatten)]
180 pub kind: ErrorKind,
181
182 #[serde(rename = "error")]
184 pub message: String,
185}
186
187impl StandardErrorBody {
188 pub fn new(kind: ErrorKind, message: String) -> Self {
190 Self { kind, message }
191 }
192}
193
194#[derive(Debug, Error)]
197#[non_exhaustive]
198pub enum IntoHttpError {
199 #[error("failed to add authentication scheme: {0}")]
201 Authentication(Box<dyn std::error::Error + Send + Sync + 'static>),
202
203 #[error(
208 "endpoint was not supported by server-reported versions, \
209 but no unstable path to fall back to was defined"
210 )]
211 NoUnstablePath,
212
213 #[error(
216 "could not create any path variant for endpoint, as it was removed in version {}",
217 .0.as_str().expect("no endpoint was removed in Matrix 1.0")
218 )]
219 EndpointRemoved(MatrixVersion),
220
221 #[error("JSON serialization failed: {0}")]
223 Json(#[from] serde_json::Error),
224
225 #[error("query parameter serialization failed: {0}")]
227 Query(#[from] serde_html_form::ser::Error),
228
229 #[error("header serialization failed: {0}")]
231 Header(#[from] HeaderSerializationError),
232
233 #[error("HTTP request construction failed: {0}")]
235 Http(#[from] http::Error),
236}
237
238impl IntoHttpError {
239 pub fn authentication(
241 error: impl Into<Box<dyn std::error::Error + Send + Sync + 'static>>,
242 ) -> Self {
243 Self::Authentication(error.into())
244 }
245}
246
247impl From<std::convert::Infallible> for IntoHttpError {
248 fn from(value: std::convert::Infallible) -> Self {
249 match value {}
250 }
251}
252
253impl From<http::header::InvalidHeaderValue> for IntoHttpError {
254 fn from(value: http::header::InvalidHeaderValue) -> Self {
255 Self::Header(value.into())
256 }
257}
258
259#[derive(Debug, Error)]
261#[non_exhaustive]
262pub enum FromHttpRequestError {
263 #[error("deserialization failed: {0}")]
265 Deserialization(DeserializationError),
266
267 #[error("http method mismatch: expected {expected}, received: {received}")]
269 MethodMismatch {
270 expected: http::method::Method,
272 received: http::method::Method,
274 },
275}
276
277impl<T> From<T> for FromHttpRequestError
278where
279 T: Into<DeserializationError>,
280{
281 fn from(err: T) -> Self {
282 Self::Deserialization(err.into())
283 }
284}
285
286#[derive(Debug)]
288#[non_exhaustive]
289pub enum FromHttpResponseError<E> {
290 Deserialization(DeserializationError),
292
293 Server(E),
295}
296
297impl<E> FromHttpResponseError<E> {
298 pub fn map<F>(self, f: impl FnOnce(E) -> F) -> FromHttpResponseError<F> {
301 match self {
302 Self::Deserialization(d) => FromHttpResponseError::Deserialization(d),
303 Self::Server(s) => FromHttpResponseError::Server(f(s)),
304 }
305 }
306}
307
308impl<E, F> FromHttpResponseError<Result<E, F>> {
309 pub fn transpose(self) -> Result<FromHttpResponseError<E>, F> {
311 match self {
312 Self::Deserialization(d) => Ok(FromHttpResponseError::Deserialization(d)),
313 Self::Server(s) => s.map(FromHttpResponseError::Server),
314 }
315 }
316}
317
318impl<E: fmt::Display> fmt::Display for FromHttpResponseError<E> {
319 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
320 match self {
321 Self::Deserialization(err) => write!(f, "deserialization failed: {err}"),
322 Self::Server(err) => write!(f, "the server returned an error: {err}"),
323 }
324 }
325}
326
327impl<E, T> From<T> for FromHttpResponseError<E>
328where
329 T: Into<DeserializationError>,
330{
331 fn from(err: T) -> Self {
332 Self::Deserialization(err.into())
333 }
334}
335
336impl<E: StdError> StdError for FromHttpResponseError<E> {}
337
338pub trait FromHttpResponseErrorExt {
340 fn error_kind(&self) -> Option<&ErrorKind>;
343}
344
345impl FromHttpResponseErrorExt for FromHttpResponseError<Error> {
346 fn error_kind(&self) -> Option<&ErrorKind> {
347 as_variant!(self, Self::Server)?.error_kind()
348 }
349}
350
351#[derive(Debug, Error)]
354#[non_exhaustive]
355pub enum DeserializationError {
356 #[error(transparent)]
358 Utf8(#[from] std::str::Utf8Error),
359
360 #[error(transparent)]
362 Json(#[from] serde_json::Error),
363
364 #[error(transparent)]
366 Query(#[from] serde_html_form::de::Error),
367
368 #[error(transparent)]
370 Ident(#[from] crate::IdParseError),
371
372 #[error(transparent)]
374 Header(#[from] HeaderDeserializationError),
375
376 #[error(transparent)]
378 MultipartMixed(#[from] MultipartMixedDeserializationError),
379}
380
381impl From<std::convert::Infallible> for DeserializationError {
382 fn from(err: std::convert::Infallible) -> Self {
383 match err {}
384 }
385}
386
387impl From<http::header::ToStrError> for DeserializationError {
388 fn from(err: http::header::ToStrError) -> Self {
389 Self::Header(HeaderDeserializationError::ToStrError(err))
390 }
391}
392
393#[derive(Debug, Error)]
395#[non_exhaustive]
396pub enum HeaderDeserializationError {
397 #[error("{0}")]
399 ToStrError(#[from] http::header::ToStrError),
400
401 #[error("{0}")]
403 ParseIntError(#[from] ParseIntError),
404
405 #[error("failed to parse HTTP date")]
407 InvalidHttpDate,
408
409 #[error("missing header `{0}`")]
411 MissingHeader(String),
412
413 #[error("invalid header: {0}")]
415 InvalidHeader(Box<dyn std::error::Error + Send + Sync + 'static>),
416
417 #[error(
419 "The {header} header was received with an unexpected value, \
420 expected {expected}, received {unexpected}"
421 )]
422 InvalidHeaderValue {
423 header: String,
425 expected: String,
427 unexpected: String,
429 },
430
431 #[error(
434 "The `Content-Type` header for a `multipart/mixed` response is missing the `boundary` attribute"
435 )]
436 MissingMultipartBoundary,
437}
438
439#[derive(Debug, Error)]
441#[non_exhaustive]
442pub enum MultipartMixedDeserializationError {
443 #[error(
445 "multipart/mixed response does not have enough body parts, \
446 expected {expected}, found {found}"
447 )]
448 MissingBodyParts {
449 expected: usize,
451 found: usize,
453 },
454
455 #[error("multipart/mixed body part is missing separator between headers and content")]
457 MissingBodyPartInnerSeparator,
458
459 #[error("multipart/mixed body part header is missing separator between name and value")]
461 MissingHeaderSeparator,
462
463 #[error("invalid multipart/mixed header: {0}")]
465 InvalidHeader(Box<dyn std::error::Error + Send + Sync + 'static>),
466}
467
468#[derive(Debug)]
470#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
471pub struct UnknownVersionError;
472
473impl fmt::Display for UnknownVersionError {
474 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
475 write!(f, "version string was unknown")
476 }
477}
478
479impl StdError for UnknownVersionError {}
480
481#[derive(Debug)]
486#[cfg_attr(not(ruma_unstable_exhaustive_types), non_exhaustive)]
487pub struct IncorrectArgumentCount {
488 pub expected: usize,
490
491 pub got: usize,
493}
494
495impl fmt::Display for IncorrectArgumentCount {
496 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
497 write!(f, "incorrect path argument count, expected {}, got {}", self.expected, self.got)
498 }
499}
500
501impl StdError for IncorrectArgumentCount {}
502
503#[derive(Debug, Error)]
505#[non_exhaustive]
506pub enum HeaderSerializationError {
507 #[error(transparent)]
509 ToHeaderValue(#[from] http::header::InvalidHeaderValue),
510
511 #[error("invalid HTTP date")]
516 InvalidHttpDate,
517}