1use thiserror::Error;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub enum ErrorKind {
11 Transient,
13 Permanent,
15}
16
17#[derive(Error, Debug)]
19#[non_exhaustive]
20pub enum RetrievalError {
21 #[error("hnsw error: {0}")]
23 Hnsw(String),
24
25 #[error("bm25 error: {0}")]
27 Bm25(String),
28
29 #[error("fusion error: {0}")]
31 Fusion(String),
32
33 #[error("graph traversal error: {0}")]
35 GraphTraversal(String),
36
37 #[error("invalid query: {0}")]
39 InvalidQuery(String),
40
41 #[error("dimension mismatch: expected {expected}, got {actual}")]
43 DimensionMismatch {
44 expected: usize,
46 actual: usize,
48 },
49
50 #[error("configuration error: {0}")]
52 Configuration(String),
53
54 #[error("embedding store: {0}")]
56 EmbeddingStore(String),
57
58 #[error("link store: {0}")]
60 LinkStore(String),
61
62 #[error("index not initialized: {0}")]
64 IndexNotInitialized(String),
65
66 #[error("index rebuild required: {reason}")]
68 RebuildRequired {
69 reason: String,
71 },
72
73 #[error("query timed out after {elapsed_ms}ms")]
79 QueryTimeout {
80 elapsed_ms: u64,
82 },
83
84 #[error("query cancelled")]
89 QueryCancelled,
90
91 #[error("memory budget exceeded: current {current_usage} + item {item_size} > limit {limit}")]
97 BudgetExceeded {
98 current_usage: usize,
100 item_size: usize,
102 limit: usize,
104 },
105
106 #[error("rerank error: {0}")]
108 Rerank(String),
109 }
114
115impl RetrievalError {
116 pub fn kind(&self) -> ErrorKind {
141 match self {
142 RetrievalError::EmbeddingStore(_)
144 | RetrievalError::LinkStore(_)
145 | RetrievalError::QueryTimeout { .. }
146 | RetrievalError::QueryCancelled => ErrorKind::Transient,
147
148 RetrievalError::Hnsw(_)
150 | RetrievalError::Bm25(_)
151 | RetrievalError::Fusion(_)
152 | RetrievalError::GraphTraversal(_)
153 | RetrievalError::InvalidQuery(_)
154 | RetrievalError::DimensionMismatch { .. }
155 | RetrievalError::Configuration(_)
156 | RetrievalError::IndexNotInitialized(_)
157 | RetrievalError::RebuildRequired { .. }
158 | RetrievalError::BudgetExceeded { .. }
159 | RetrievalError::Rerank(_) => ErrorKind::Permanent,
160 }
164 }
165
166 #[inline]
168 pub fn is_transient(&self) -> bool {
169 self.kind() == ErrorKind::Transient
170 }
171
172 #[inline]
177 pub fn is_permanent(&self) -> bool {
178 self.kind() == ErrorKind::Permanent
179 }
180
181 #[inline]
185 pub fn is_retryable(&self) -> bool {
186 self.is_transient()
187 }
188
189 pub fn rerank(msg: impl Into<String>) -> Self {
191 Self::Rerank(msg.into())
192 }
193
194 pub fn hnsw(msg: impl Into<String>) -> Self {
196 Self::Hnsw(msg.into())
197 }
198
199 pub fn bm25(msg: impl Into<String>) -> Self {
201 Self::Bm25(msg.into())
202 }
203
204 pub fn fusion(msg: impl Into<String>) -> Self {
206 Self::Fusion(msg.into())
207 }
208
209 pub fn graph_traversal(msg: impl Into<String>) -> Self {
211 Self::GraphTraversal(msg.into())
212 }
213
214 pub fn invalid_query(msg: impl Into<String>) -> Self {
216 Self::InvalidQuery(msg.into())
217 }
218
219 pub fn dimension_mismatch(expected: usize, actual: usize) -> Self {
221 Self::DimensionMismatch { expected, actual }
222 }
223
224 pub fn configuration(msg: impl Into<String>) -> Self {
226 Self::Configuration(msg.into())
227 }
228
229 pub fn index_not_initialized(msg: impl Into<String>) -> Self {
231 Self::IndexNotInitialized(msg.into())
232 }
233
234 pub fn rebuild_required(reason: impl Into<String>) -> Self {
236 Self::RebuildRequired {
237 reason: reason.into(),
238 }
239 }
240
241 pub fn query_timeout(elapsed_ms: u64) -> Self {
243 Self::QueryTimeout { elapsed_ms }
244 }
245
246 pub fn query_cancelled() -> Self {
248 Self::QueryCancelled
249 }
250
251 pub fn budget_exceeded(current_usage: usize, item_size: usize, limit: usize) -> Self {
253 Self::BudgetExceeded {
254 current_usage,
255 item_size,
256 limit,
257 }
258 }
259}
260
261pub type Result<T> = std::result::Result<T, RetrievalError>;
263
264#[cfg(test)]
265mod tests {
266 use super::*;
267
268 #[test]
269 fn test_error_display() {
270 let err = RetrievalError::hnsw("connection failed");
271 assert_eq!(err.to_string(), "hnsw error: connection failed");
272 }
273
274 #[test]
275 fn test_dimension_mismatch() {
276 let err = RetrievalError::dimension_mismatch(768, 512);
277 assert_eq!(err.to_string(), "dimension mismatch: expected 768, got 512");
278 }
279
280 #[test]
281 fn test_is_retryable() {
282 assert!(!RetrievalError::hnsw("fail").is_retryable());
284 assert!(!RetrievalError::bm25("fail").is_retryable());
285 assert!(!RetrievalError::InvalidQuery("bad".into()).is_retryable());
286 assert!(!RetrievalError::dimension_mismatch(768, 512).is_retryable());
287 }
288
289 #[test]
292 fn test_error_kind_transient() {
293 }
297
298 #[test]
299 fn test_error_kind_permanent_all_variants() {
300 let permanent_errors: Vec<RetrievalError> = vec![
302 RetrievalError::hnsw("index corrupt"),
303 RetrievalError::bm25("tokenization failed"),
304 RetrievalError::fusion("incompatible scores"),
305 RetrievalError::graph_traversal("cycle detected"),
306 RetrievalError::invalid_query("empty query"),
307 RetrievalError::dimension_mismatch(768, 512),
308 RetrievalError::configuration("invalid k1 value"),
309 RetrievalError::index_not_initialized("HNSW index"),
310 RetrievalError::rebuild_required("version mismatch"),
311 RetrievalError::budget_exceeded(1000, 500, 1200),
312 ];
313
314 for err in permanent_errors {
315 assert!(err.is_permanent(), "Expected permanent: {err:?}");
316 assert!(!err.is_transient(), "Should not be transient: {err:?}");
317 assert_eq!(
318 err.kind(),
319 ErrorKind::Permanent,
320 "Kind mismatch for: {err:?}"
321 );
322 }
323 }
324
325 #[test]
326 fn test_is_transient_is_permanent_consistency() {
327 let test_errors: Vec<RetrievalError> = vec![
329 RetrievalError::hnsw("test"),
330 RetrievalError::bm25("test"),
331 RetrievalError::fusion("test"),
332 RetrievalError::invalid_query("test"),
333 RetrievalError::dimension_mismatch(1, 2),
334 RetrievalError::configuration("test"),
335 RetrievalError::budget_exceeded(100, 50, 120),
336 ];
337
338 for err in test_errors {
339 let transient = err.is_transient();
340 let permanent = err.is_permanent();
341
342 assert!(
344 transient ^ permanent,
345 "Error must be exactly transient OR permanent: {err:?} (transient={transient}, permanent={permanent})"
346 );
347
348 assert_eq!(
350 err.is_retryable(),
351 err.is_transient(),
352 "is_retryable should equal is_transient for: {err:?}"
353 );
354 }
355 }
356
357 #[test]
358 fn test_error_constructors_produce_correct_messages() {
359 assert_eq!(RetrievalError::hnsw("test").to_string(), "hnsw error: test");
360 assert_eq!(RetrievalError::bm25("test").to_string(), "bm25 error: test");
361 assert_eq!(
362 RetrievalError::fusion("test").to_string(),
363 "fusion error: test"
364 );
365 assert_eq!(
366 RetrievalError::graph_traversal("test").to_string(),
367 "graph traversal error: test"
368 );
369 assert_eq!(
370 RetrievalError::invalid_query("test").to_string(),
371 "invalid query: test"
372 );
373 assert_eq!(
374 RetrievalError::configuration("test").to_string(),
375 "configuration error: test"
376 );
377 assert_eq!(
378 RetrievalError::index_not_initialized("test").to_string(),
379 "index not initialized: test"
380 );
381 assert_eq!(
382 RetrievalError::rebuild_required("test").to_string(),
383 "index rebuild required: test"
384 );
385 assert_eq!(
386 RetrievalError::budget_exceeded(100, 50, 120).to_string(),
387 "memory budget exceeded: current 100 + item 50 > limit 120"
388 );
389 }
390
391 #[test]
392 fn test_error_kind_enum_debug() {
393 assert_eq!(format!("{:?}", ErrorKind::Transient), "Transient");
395 assert_eq!(format!("{:?}", ErrorKind::Permanent), "Permanent");
396 }
397
398 #[test]
399 fn test_error_kind_equality() {
400 assert_eq!(ErrorKind::Transient, ErrorKind::Transient);
402 assert_eq!(ErrorKind::Permanent, ErrorKind::Permanent);
403 assert_ne!(ErrorKind::Transient, ErrorKind::Permanent);
404 }
405
406 #[test]
407 fn test_query_timeout_error() {
408 let err = RetrievalError::query_timeout(5000);
409 assert_eq!(err.to_string(), "query timed out after 5000ms");
410 assert!(err.is_transient());
411 assert!(!err.is_permanent());
412 assert!(err.is_retryable());
413 assert_eq!(err.kind(), ErrorKind::Transient);
414 }
415
416 #[test]
417 fn test_query_cancelled_error() {
418 let err = RetrievalError::query_cancelled();
419 assert_eq!(err.to_string(), "query cancelled");
420 assert!(err.is_transient());
421 assert!(!err.is_permanent());
422 assert!(err.is_retryable());
423 assert_eq!(err.kind(), ErrorKind::Transient);
424 }
425
426 #[test]
427 fn test_transient_errors_classification() {
428 let transient_errors: Vec<RetrievalError> = vec![
430 RetrievalError::query_timeout(100),
431 RetrievalError::query_cancelled(),
432 ];
433
434 for err in transient_errors {
435 assert!(err.is_transient(), "Expected transient: {err:?}");
436 assert!(!err.is_permanent(), "Should not be permanent: {err:?}");
437 assert_eq!(
438 err.kind(),
439 ErrorKind::Transient,
440 "Kind mismatch for: {err:?}"
441 );
442 }
443 }
444}