use std::{
alloc::LayoutError,
array::TryFromSliceError,
fmt::{Debug, Display},
io,
num::TryFromIntError,
sync::mpsc,
};
use crate::always_escalate;
pub type ANNResult<T> = Result<T, ANNError>;
#[derive(Debug)]
pub struct ANNError {
error: anyhow::Error,
}
impl ANNError {
#[track_caller]
#[inline(never)]
pub fn new<E>(err: E) -> Self
where
E: std::error::Error + Send + Sync + 'static,
{
Self {
error: anyhow::Error::new(Located::new(err)),
}
}
#[track_caller]
#[inline(never)]
pub fn message<D>(display: D) -> Self
where
D: Display + Debug + Send + Sync + 'static,
{
Self {
error: anyhow::Error::msg(Located::new(display)),
}
}
#[must_use]
pub fn is<E>(&self) -> bool
where
E: Display + Debug + Send + Sync + 'static,
{
self.error.is::<Located<E>>() || self.error.is::<E>()
}
pub fn downcast<E>(self) -> Result<E, Self>
where
E: Display + Debug + Send + Sync + 'static,
{
match self.error.downcast::<E>() {
Ok(value) => Ok(value),
Err(error) => match error.downcast::<Located<E>>() {
Ok(value) => Ok(value.err),
Err(error) => Err(Self { error }),
},
}
}
pub fn downcast_ref<E>(&self) -> Option<&E>
where
E: Display + Debug + Send + Sync + 'static,
{
match self.error.downcast_ref::<E>() {
Some(err) => Some(err),
None => self.error.downcast_ref::<Located<E>>().map(|e| &e.err),
}
}
pub fn downcast_mut<E>(&mut self) -> Option<&mut E>
where
E: Display + Debug + Send + Sync + 'static,
{
if self.error.is::<E>() {
self.error.downcast_mut::<E>()
} else {
self.error.downcast_mut::<Located<E>>().map(|e| &mut e.err)
}
}
#[track_caller]
#[inline(never)]
pub fn context<C>(self, context: C) -> Self
where
C: Display + Debug + Send + Sync + 'static,
{
Self {
error: self.error.context(Located::new(context)),
}
}
}
impl Display for ANNError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
write!(formatter, "ANNError\n\n{:?}", self.error)
}
}
impl std::error::Error for ANNError {
}
always_escalate!(ANNError);
#[macro_export]
macro_rules! convert_error {
($T:ty) => {
impl From<$T> for $crate::ANNError {
#[track_caller]
fn from(e: $T) -> $crate::ANNError {
$crate::ANNError::new(e)
}
}
};
}
impl From<std::convert::Infallible> for ANNError {
#[track_caller]
fn from(_: std::convert::Infallible) -> Self {
unreachable!("Infallible is an unconstructible type");
}
}
convert_error!(io::Error);
convert_error!(LayoutError);
convert_error!(TryFromIntError);
convert_error!(TryFromSliceError);
convert_error!(diskann_utils::io::ReadBinError);
convert_error!(diskann_utils::io::SaveBinError);
convert_error!(diskann_utils::views::TryFromErrorLight);
impl<T> From<mpsc::SendError<T>> for ANNError
where
T: Send + Sync + 'static,
{
#[track_caller]
fn from(err: mpsc::SendError<T>) -> Self {
ANNError::new(err)
}
}
impl<T, U> From<diskann_utils::io::MetadataError<T, U>> for ANNError
where
T: std::error::Error + Send + Sync + 'static,
U: std::error::Error + Send + Sync + 'static,
{
#[track_caller]
fn from(err: diskann_utils::io::MetadataError<T, U>) -> Self {
ANNError::new(err)
}
}
impl<T> From<diskann_utils::views::TryFromError<T>> for ANNError
where
T: diskann_utils::views::DenseData,
{
#[track_caller]
fn from(err: diskann_utils::views::TryFromError<T>) -> Self {
Self::from(err.as_static())
}
}
#[derive(Debug)]
struct Located<T>
where
T: Debug,
{
err: T,
location: &'static std::panic::Location<'static>,
}
impl<T> Located<T>
where
T: Debug,
{
#[track_caller]
fn new(err: T) -> Self {
Self {
err,
location: std::panic::Location::caller(),
}
}
}
impl<T> Display for Located<T>
where
T: Display + Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
write!(
f,
"{} -- ({}:{})",
self.err,
self.location.file(),
self.location.line()
)
}
}
impl<T> std::error::Error for Located<T>
where
T: std::error::Error + Debug,
{
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
self.err.source()
}
}
pub trait ErrorContext<T> {
fn context<C>(self, context: C) -> Result<T, ANNError>
where
C: Display + Debug + Send + Sync + 'static;
fn with_context<F, C>(self, f: F) -> Result<T, ANNError>
where
C: Display + Debug + Send + Sync + 'static,
F: FnOnce() -> C;
}
impl<T, E> ErrorContext<T> for Result<T, E>
where
ANNError: From<E>,
{
#[track_caller]
fn context<C>(self, context: C) -> Result<T, ANNError>
where
C: Display + Debug + Send + Sync + 'static,
{
match self {
Ok(value) => Ok(value),
Err(error) => Err(ANNError::from(error).context(context)),
}
}
#[track_caller]
fn with_context<F, C>(self, f: F) -> Result<T, ANNError>
where
C: Display + Debug + Send + Sync + 'static,
F: FnOnce() -> C,
{
match self {
Ok(value) => Ok(value),
Err(error) => Err(ANNError::from(error).context(f())),
}
}
}
pub trait IntoANNResult<T> {
fn into_ann_result(self) -> Result<T, ANNError>;
}
impl<T, E> IntoANNResult<T> for Result<T, E>
where
E: Into<ANNError>,
{
#[inline(always)]
#[track_caller]
fn into_ann_result(self) -> Result<T, ANNError> {
match self {
Ok(v) => Ok(v),
Err(e) => Err(e.into()),
}
}
}
#[cfg(test)]
mod ann_result_test {
use super::*;
#[test]
fn ann_err_is_send_and_sync() {
fn assert_send_and_sync<T: Send + Sync>() {}
assert_send_and_sync::<ANNError>();
}
#[test]
fn check_struct_size() {
assert_eq!(std::mem::size_of::<ANNError>(), 8);
assert_eq!(std::mem::size_of::<Option<ANNError>>(), 8);
assert_eq!(std::mem::size_of::<Result<f32, ANNError>>(), 16);
}
#[derive(Debug, Clone)]
struct SampleError {
value: usize,
}
impl SampleError {
fn new(value: usize) -> Self {
Self { value }
}
}
impl Display for SampleError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
write!(f, "SampleError {{ {} }}", self.value)
}
}
impl std::error::Error for SampleError {}
convert_error!(SampleError);
#[derive(Debug, Clone)]
struct SampleChainedError {
value: usize,
source: SampleError,
}
impl SampleChainedError {
fn new(value: usize, source: SampleError) -> Self {
Self { value, source }
}
}
impl Display for SampleChainedError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> Result<(), std::fmt::Error> {
write!(f, "SampleChainedError {{ {} }}", self.value)
}
}
impl std::error::Error for SampleChainedError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
Some(&self.source)
}
}
#[test]
fn check_downcasting() {
let err = SampleError::new(10);
let base_error = err.to_string();
{
let mut ann = ANNError::from(err.clone());
assert!(format!("{}", ann).contains(&base_error));
assert!(ann.is::<SampleError>());
let r = ann.downcast_ref::<SampleError>().unwrap();
assert_eq!(r.value, 10);
let r = ann.downcast_mut::<SampleError>().unwrap();
r.value = 100;
let r = ann.downcast_ref::<SampleError>().unwrap();
assert_eq!(r.value, 100);
let r = ann.downcast::<SampleError>().unwrap();
assert_eq!(r.value, 100);
}
{
let mut ann = ANNError::from(err.clone())
.context("some context here")
.context("more context");
assert!(ann.is::<SampleError>());
let formatted = ann.to_string();
assert!(formatted.contains(&base_error));
assert!(formatted.contains("some context here"));
assert!(formatted.contains("more context"));
let r = ann.downcast_ref::<SampleError>().unwrap();
assert_eq!(r.value, 10);
let r = ann.downcast_mut::<SampleError>().unwrap();
r.value = 100;
let r = ann.downcast_ref::<SampleError>().unwrap();
assert_eq!(r.value, 100);
let r = ann.downcast::<SampleError>().unwrap();
assert_eq!(r.value, 100);
}
{
let ann = ANNError::from(err.clone())
.context("some context here")
.context("more context");
assert!(!ann.is::<usize>());
println!("{}", ann);
let formatted = ann.to_string();
let mut ann = ann.downcast::<usize>().unwrap_err();
assert_eq!(formatted, ann.to_string());
assert!(ann.downcast_ref::<usize>().is_none());
assert!(ann.downcast_mut::<usize>().is_none());
}
}
#[test]
fn context_chaining() {
let sample = SampleError::new(5).to_string();
fn err() -> Result<usize, ANNError> {
Err(ANNError::new(SampleError::new(5)))
}
fn ok() -> Result<usize, ANNError> {
Ok(77)
}
{
let propagates = || err().context("with context");
let chained = propagates().unwrap_err();
let message = chained.to_string();
assert!(message.contains("with context"), "got: {}", message);
assert!(message.contains(&sample), "got: {}", message);
assert_eq!(chained.downcast_ref::<SampleError>().unwrap().value, 5);
assert!(chained.is::<SampleError>());
}
{
let propagates = || ok().context("with context");
let fine = propagates().unwrap();
assert_eq!(fine, 77);
}
{
let mut called = false;
let mut propagates = || {
err().with_context(|| {
assert!(!called);
called = true;
"with context"
})
};
let chained = propagates().unwrap_err();
assert!(called);
let message = chained.to_string();
assert!(message.contains("with context"), "got: {}", message);
assert!(message.contains(&sample), "got: {}", message);
assert_eq!(chained.downcast_ref::<SampleError>().unwrap().value, 5);
}
{
let propagates = || ok().with_context(|| -> ! { panic!("should not be called") });
let fine = propagates().unwrap();
assert_eq!(fine, 77);
}
}
#[test]
fn full_formatting() {
let sample = SampleError::new(5);
let file = file!();
let l0 = line!() + 1;
let err = ANNError::from(sample);
let l1 = line!() + 1;
let err = err.context("some context");
let l2 = line!() + 1;
let err = err.context("more context");
let expected = format!(
"ANNError
more context -- ({}:{})
Caused by:
0: some context -- ({}:{})
1: SampleError {{ {} }} -- ({}:{})",
file, l2, file, l1, 5, file, l0
);
let got = err.to_string();
assert!(
got.starts_with(&expected),
"got:\n{}\n\nexpected:\n{}",
got,
expected
);
}
#[test]
fn full_formatting_with_cause() {
let sample = SampleChainedError::new(10, SampleError::new(5));
let file = file!();
let l0 = line!() + 1;
let err = ANNError::new(sample);
let l1 = line!() + 1;
let err = err.context("some context");
let l2 = line!() + 1;
let err = err.context("more context");
let expected = format!(
"ANNError
more context -- ({}:{})
Caused by:
0: some context -- ({}:{})
1: SampleChainedError {{ 10 }} -- ({}:{})
2: SampleError {{ 5 }}",
file, l2, file, l1, file, l0
);
let got = err.to_string();
assert!(
got.starts_with(&expected),
"got:\n{}\n\nexpected:\n{}",
got,
expected
);
}
#[test]
fn full_formatting_with_cause_no_context() {
let sample = SampleChainedError::new(10, SampleError::new(5));
let file = file!();
let l0 = line!() + 1;
let err = ANNError::new(sample);
let expected = format!(
"ANNError
SampleChainedError {{ 10 }} -- ({}:{})
Caused by:
SampleError {{ 5 }}",
file, l0
);
let got = err.to_string();
assert!(
got.starts_with(&expected),
"got:\n{}\n\nexpected:\n{}",
got,
expected
);
}
}