use thiserror::Error;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ErrorKind {
Transient,
Permanent,
}
#[derive(Error, Debug)]
#[non_exhaustive]
pub enum RetrievalError {
#[error("hnsw error: {0}")]
Hnsw(String),
#[error("bm25 error: {0}")]
Bm25(String),
#[error("fusion error: {0}")]
Fusion(String),
#[error("graph traversal error: {0}")]
GraphTraversal(String),
#[error("invalid query: {0}")]
InvalidQuery(String),
#[error("dimension mismatch: expected {expected}, got {actual}")]
DimensionMismatch {
expected: usize,
actual: usize,
},
#[error("configuration error: {0}")]
Configuration(String),
#[error("embedding store: {0}")]
EmbeddingStore(String),
#[error("link store: {0}")]
LinkStore(String),
#[error("index not initialized: {0}")]
IndexNotInitialized(String),
#[error("index rebuild required: {reason}")]
RebuildRequired {
reason: String,
},
#[error("query timed out after {elapsed_ms}ms")]
QueryTimeout {
elapsed_ms: u64,
},
#[error("query cancelled")]
QueryCancelled,
#[error("memory budget exceeded: current {current_usage} + item {item_size} > limit {limit}")]
BudgetExceeded {
current_usage: usize,
item_size: usize,
limit: usize,
},
#[error("rerank error: {0}")]
Rerank(String),
}
impl RetrievalError {
pub fn kind(&self) -> ErrorKind {
match self {
RetrievalError::EmbeddingStore(_)
| RetrievalError::LinkStore(_)
| RetrievalError::QueryTimeout { .. }
| RetrievalError::QueryCancelled => ErrorKind::Transient,
RetrievalError::Hnsw(_)
| RetrievalError::Bm25(_)
| RetrievalError::Fusion(_)
| RetrievalError::GraphTraversal(_)
| RetrievalError::InvalidQuery(_)
| RetrievalError::DimensionMismatch { .. }
| RetrievalError::Configuration(_)
| RetrievalError::IndexNotInitialized(_)
| RetrievalError::RebuildRequired { .. }
| RetrievalError::BudgetExceeded { .. }
| RetrievalError::Rerank(_) => ErrorKind::Permanent,
}
}
#[inline]
pub fn is_transient(&self) -> bool {
self.kind() == ErrorKind::Transient
}
#[inline]
pub fn is_permanent(&self) -> bool {
self.kind() == ErrorKind::Permanent
}
#[inline]
pub fn is_retryable(&self) -> bool {
self.is_transient()
}
pub fn rerank(msg: impl Into<String>) -> Self {
Self::Rerank(msg.into())
}
pub fn hnsw(msg: impl Into<String>) -> Self {
Self::Hnsw(msg.into())
}
pub fn bm25(msg: impl Into<String>) -> Self {
Self::Bm25(msg.into())
}
pub fn fusion(msg: impl Into<String>) -> Self {
Self::Fusion(msg.into())
}
pub fn graph_traversal(msg: impl Into<String>) -> Self {
Self::GraphTraversal(msg.into())
}
pub fn invalid_query(msg: impl Into<String>) -> Self {
Self::InvalidQuery(msg.into())
}
pub fn dimension_mismatch(expected: usize, actual: usize) -> Self {
Self::DimensionMismatch { expected, actual }
}
pub fn configuration(msg: impl Into<String>) -> Self {
Self::Configuration(msg.into())
}
pub fn index_not_initialized(msg: impl Into<String>) -> Self {
Self::IndexNotInitialized(msg.into())
}
pub fn rebuild_required(reason: impl Into<String>) -> Self {
Self::RebuildRequired {
reason: reason.into(),
}
}
pub fn query_timeout(elapsed_ms: u64) -> Self {
Self::QueryTimeout { elapsed_ms }
}
pub fn query_cancelled() -> Self {
Self::QueryCancelled
}
pub fn budget_exceeded(current_usage: usize, item_size: usize, limit: usize) -> Self {
Self::BudgetExceeded {
current_usage,
item_size,
limit,
}
}
}
pub type Result<T> = std::result::Result<T, RetrievalError>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_error_display() {
let err = RetrievalError::hnsw("connection failed");
assert_eq!(err.to_string(), "hnsw error: connection failed");
}
#[test]
fn test_dimension_mismatch() {
let err = RetrievalError::dimension_mismatch(768, 512);
assert_eq!(err.to_string(), "dimension mismatch: expected 768, got 512");
}
#[test]
fn test_is_retryable() {
assert!(!RetrievalError::hnsw("fail").is_retryable());
assert!(!RetrievalError::bm25("fail").is_retryable());
assert!(!RetrievalError::InvalidQuery("bad".into()).is_retryable());
assert!(!RetrievalError::dimension_mismatch(768, 512).is_retryable());
}
#[test]
fn test_error_kind_transient() {
}
#[test]
fn test_error_kind_permanent_all_variants() {
let permanent_errors: Vec<RetrievalError> = vec![
RetrievalError::hnsw("index corrupt"),
RetrievalError::bm25("tokenization failed"),
RetrievalError::fusion("incompatible scores"),
RetrievalError::graph_traversal("cycle detected"),
RetrievalError::invalid_query("empty query"),
RetrievalError::dimension_mismatch(768, 512),
RetrievalError::configuration("invalid k1 value"),
RetrievalError::index_not_initialized("HNSW index"),
RetrievalError::rebuild_required("version mismatch"),
RetrievalError::budget_exceeded(1000, 500, 1200),
];
for err in permanent_errors {
assert!(err.is_permanent(), "Expected permanent: {err:?}");
assert!(!err.is_transient(), "Should not be transient: {err:?}");
assert_eq!(
err.kind(),
ErrorKind::Permanent,
"Kind mismatch for: {err:?}"
);
}
}
#[test]
fn test_is_transient_is_permanent_consistency() {
let test_errors: Vec<RetrievalError> = vec![
RetrievalError::hnsw("test"),
RetrievalError::bm25("test"),
RetrievalError::fusion("test"),
RetrievalError::invalid_query("test"),
RetrievalError::dimension_mismatch(1, 2),
RetrievalError::configuration("test"),
RetrievalError::budget_exceeded(100, 50, 120),
];
for err in test_errors {
let transient = err.is_transient();
let permanent = err.is_permanent();
assert!(
transient ^ permanent,
"Error must be exactly transient OR permanent: {err:?} (transient={transient}, permanent={permanent})"
);
assert_eq!(
err.is_retryable(),
err.is_transient(),
"is_retryable should equal is_transient for: {err:?}"
);
}
}
#[test]
fn test_error_constructors_produce_correct_messages() {
assert_eq!(RetrievalError::hnsw("test").to_string(), "hnsw error: test");
assert_eq!(RetrievalError::bm25("test").to_string(), "bm25 error: test");
assert_eq!(
RetrievalError::fusion("test").to_string(),
"fusion error: test"
);
assert_eq!(
RetrievalError::graph_traversal("test").to_string(),
"graph traversal error: test"
);
assert_eq!(
RetrievalError::invalid_query("test").to_string(),
"invalid query: test"
);
assert_eq!(
RetrievalError::configuration("test").to_string(),
"configuration error: test"
);
assert_eq!(
RetrievalError::index_not_initialized("test").to_string(),
"index not initialized: test"
);
assert_eq!(
RetrievalError::rebuild_required("test").to_string(),
"index rebuild required: test"
);
assert_eq!(
RetrievalError::budget_exceeded(100, 50, 120).to_string(),
"memory budget exceeded: current 100 + item 50 > limit 120"
);
}
#[test]
fn test_error_kind_enum_debug() {
assert_eq!(format!("{:?}", ErrorKind::Transient), "Transient");
assert_eq!(format!("{:?}", ErrorKind::Permanent), "Permanent");
}
#[test]
fn test_error_kind_equality() {
assert_eq!(ErrorKind::Transient, ErrorKind::Transient);
assert_eq!(ErrorKind::Permanent, ErrorKind::Permanent);
assert_ne!(ErrorKind::Transient, ErrorKind::Permanent);
}
#[test]
fn test_query_timeout_error() {
let err = RetrievalError::query_timeout(5000);
assert_eq!(err.to_string(), "query timed out after 5000ms");
assert!(err.is_transient());
assert!(!err.is_permanent());
assert!(err.is_retryable());
assert_eq!(err.kind(), ErrorKind::Transient);
}
#[test]
fn test_query_cancelled_error() {
let err = RetrievalError::query_cancelled();
assert_eq!(err.to_string(), "query cancelled");
assert!(err.is_transient());
assert!(!err.is_permanent());
assert!(err.is_retryable());
assert_eq!(err.kind(), ErrorKind::Transient);
}
#[test]
fn test_transient_errors_classification() {
let transient_errors: Vec<RetrievalError> = vec![
RetrievalError::query_timeout(100),
RetrievalError::query_cancelled(),
];
for err in transient_errors {
assert!(err.is_transient(), "Expected transient: {err:?}");
assert!(!err.is_permanent(), "Should not be permanent: {err:?}");
assert_eq!(
err.kind(),
ErrorKind::Transient,
"Kind mismatch for: {err:?}"
);
}
}
}