diff --git a/FEATURE_PARITY.md b/FEATURE_PARITY.md index b6265291..cda8dfd1 100644 --- a/FEATURE_PARITY.md +++ b/FEATURE_PARITY.md @@ -133,7 +133,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O |---------|----------|----------|-------| | Pi agent runtime | ✅ | ➖ | IronClaw uses custom runtime | | RPC-based execution | ✅ | ✅ | Orchestrator/worker pattern | -| Multi-provider failover | ✅ | ❌ | Provider fallback chains | +| Multi-provider failover | ✅ | ✅ | `FailoverProvider` tries providers sequentially on retryable errors | | Per-sender sessions | ✅ | ✅ | | | Global sessions | ✅ | ❌ | Optional shared context | | Session pruning | ✅ | ❌ | Auto cleanup old sessions | @@ -173,7 +173,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O | Feature | OpenClaw | IronClaw | Notes | |---------|----------|----------|-------| | Auto-discovery | ✅ | ❌ | | -| Failover chains | ✅ | ❌ | Provider fallback | +| Failover chains | ✅ | ✅ | `FailoverProvider` with configurable `fallback_model` | | Cooldown management | ✅ | ❌ | Skip failed providers | | Per-session model override | ✅ | ✅ | Model selector in TUI | | Model selection UI | ✅ | ✅ | TUI keyboard shortcut | @@ -419,7 +419,7 @@ This document tracks feature parity between IronClaw (Rust implementation) and O - ❌ Slack channel (real implementation) - ✅ Telegram channel (WASM, DM pairing, caption, /start) - ❌ WhatsApp channel -- ❌ Multi-provider failover +- ✅ Multi-provider failover (`FailoverProvider` with retryable error classification) - ❌ Hooks system (beforeInbound, beforeToolCall, etc.) ### P2 - Medium Priority diff --git a/src/config.rs b/src/config.rs index a6ddb158..35c89604 100644 --- a/src/config.rs +++ b/src/config.rs @@ -397,6 +397,15 @@ pub struct NearAiConfig { pub api_mode: NearAiApiMode, /// API key for cloud-api (required for chat_completions mode) pub api_key: Option, + /// Optional fallback model for failover (default: None). + /// When set, a secondary provider is created with this model and wrapped + /// in a `FailoverProvider` so transient errors on the primary model + /// automatically fall through to the fallback. + pub fallback_model: Option, + /// Maximum number of retries for transient errors (default: 3). + /// With the default of 3, the provider makes up to 4 total attempts + /// (1 initial + 3 retries) before giving up. + pub max_retries: u32, } impl LlmConfig { @@ -441,6 +450,8 @@ impl LlmConfig { .unwrap_or_else(default_session_path), api_mode, api_key: nearai_api_key, + fallback_model: optional_env("NEARAI_FALLBACK_MODEL")?, + max_retries: parse_optional_env("NEARAI_MAX_RETRIES", 3)?, }; // Resolve provider-specific configs based on backend diff --git a/src/llm/failover.rs b/src/llm/failover.rs new file mode 100644 index 00000000..6bc5c239 --- /dev/null +++ b/src/llm/failover.rs @@ -0,0 +1,483 @@ +//! Multi-provider LLM failover. +//! +//! Wraps multiple LlmProvider instances and tries each in sequence +//! until one succeeds. Transparent to callers --- same LlmProvider trait. + +use std::future::Future; +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use async_trait::async_trait; +use rust_decimal::Decimal; + +use crate::error::LlmError; +use crate::llm::provider::{ + CompletionRequest, CompletionResponse, LlmProvider, ToolCompletionRequest, + ToolCompletionResponse, +}; + +/// Returns `true` if the error is transient and the request should be retried +/// on the next provider in the failover chain. +/// +/// Retryable: `RequestFailed`, `RateLimited`, `InvalidResponse`, +/// `SessionRenewalFailed`, `ModelNotAvailable`, `Http`, `Io`. +/// +/// `ModelNotAvailable` is retryable because the next provider in the chain may +/// offer a different model, so it's worth trying. +/// +/// Non-retryable errors (`AuthFailed`, `SessionExpired`, `ContextLengthExceeded`) +/// propagate immediately because a different provider won't fix them. +fn is_retryable(err: &LlmError) -> bool { + matches!( + err, + LlmError::RequestFailed { .. } + | LlmError::RateLimited { .. } + | LlmError::InvalidResponse { .. } + | LlmError::SessionRenewalFailed { .. } + // ModelNotAvailable is retryable: the next provider may offer a different model. + | LlmError::ModelNotAvailable { .. } + | LlmError::Http(_) + | LlmError::Io(_) + ) +} + +/// An LLM provider that wraps multiple providers and tries each in sequence +/// on transient failures. +/// +/// The first provider in the list is the primary. If it fails with a retryable +/// error, the next provider is tried, and so on. Non-retryable errors +/// (e.g. `AuthFailed`, `ContextLengthExceeded`) propagate immediately. +pub struct FailoverProvider { + providers: Vec>, + /// Index of the provider that last handled a request successfully. + /// Used by `model_name()` and `cost_per_token()` so downstream cost + /// tracking reflects the provider that actually served the request. + last_used: AtomicUsize, +} + +impl FailoverProvider { + /// Create a new failover provider. + /// + /// Returns an error if `providers` is empty. + pub fn new(providers: Vec>) -> Result { + if providers.is_empty() { + return Err(LlmError::RequestFailed { + provider: "failover".to_string(), + reason: "FailoverProvider requires at least one provider".to_string(), + }); + } + Ok(Self { + providers, + last_used: AtomicUsize::new(0), + }) + } + + /// Try each provider in sequence until one succeeds or all fail. + async fn try_providers(&self, mut call: F) -> Result + where + F: FnMut(Arc) -> Fut, + Fut: Future>, + { + let mut last_error: Option = None; + + for (i, provider) in self.providers.iter().enumerate() { + let result = call(Arc::clone(provider)).await; + match result { + Ok(response) => { + self.last_used.store(i, Ordering::Relaxed); + return Ok(response); + } + Err(err) => { + if !is_retryable(&err) { + return Err(err); + } + if i + 1 < self.providers.len() { + tracing::warn!( + provider = %provider.model_name(), + error = %err, + next_provider = %self.providers[i + 1].model_name(), + "Provider failed with retryable error, trying next provider" + ); + } + last_error = Some(err); + } + } + } + + // SAFETY: providers is non-empty (checked in `new`), so at least one + // iteration ran and `last_error` is `Some`. + Err(last_error.expect("providers list is non-empty")) + } +} + +#[async_trait] +impl LlmProvider for FailoverProvider { + fn model_name(&self) -> &str { + self.providers[self.last_used.load(Ordering::Relaxed)].model_name() + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + self.providers[self.last_used.load(Ordering::Relaxed)].cost_per_token() + } + + async fn complete(&self, request: CompletionRequest) -> Result { + self.try_providers(|provider| { + let req = request.clone(); + async move { provider.complete(req).await } + }) + .await + } + + async fn complete_with_tools( + &self, + request: ToolCompletionRequest, + ) -> Result { + self.try_providers(|provider| { + let req = request.clone(); + async move { provider.complete_with_tools(req).await } + }) + .await + } + + async fn list_models(&self) -> Result, LlmError> { + let mut all_models = Vec::new(); + + for provider in &self.providers { + match provider.list_models().await { + Ok(models) => all_models.extend(models), + Err(err) => { + tracing::warn!( + provider = %provider.model_name(), + error = %err, + "Failed to list models from provider, skipping" + ); + } + } + } + + all_models.sort(); + all_models.dedup(); + Ok(all_models) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + use std::sync::Mutex; + use std::time::Duration; + + use crate::llm::provider::{CompletionResponse, FinishReason, ToolCompletionResponse}; + + /// A mock LLM provider that returns a predetermined result. + struct MockProvider { + name: String, + input_cost: Decimal, + output_cost: Decimal, + complete_result: Mutex>>, + tool_complete_result: Mutex>>, + } + + impl MockProvider { + fn succeeding(name: &str, content: &str) -> Self { + Self { + name: name.to_string(), + input_cost: Decimal::ZERO, + output_cost: Decimal::ZERO, + complete_result: Mutex::new(Some(Ok(CompletionResponse { + content: content.to_string(), + input_tokens: 10, + output_tokens: 5, + finish_reason: FinishReason::Stop, + response_id: None, + }))), + tool_complete_result: Mutex::new(Some(Ok(ToolCompletionResponse { + content: Some(content.to_string()), + tool_calls: vec![], + input_tokens: 10, + output_tokens: 5, + finish_reason: FinishReason::Stop, + response_id: None, + }))), + } + } + + fn succeeding_with_cost( + name: &str, + content: &str, + input_cost: Decimal, + output_cost: Decimal, + ) -> Self { + Self { + input_cost, + output_cost, + ..Self::succeeding(name, content) + } + } + + fn failing_retryable(name: &str) -> Self { + Self { + name: name.to_string(), + input_cost: Decimal::ZERO, + output_cost: Decimal::ZERO, + complete_result: Mutex::new(Some(Err(LlmError::RequestFailed { + provider: name.to_string(), + reason: "server error".to_string(), + }))), + tool_complete_result: Mutex::new(Some(Err(LlmError::RequestFailed { + provider: name.to_string(), + reason: "server error".to_string(), + }))), + } + } + + fn failing_non_retryable(name: &str) -> Self { + Self { + name: name.to_string(), + input_cost: Decimal::ZERO, + output_cost: Decimal::ZERO, + complete_result: Mutex::new(Some(Err(LlmError::AuthFailed { + provider: name.to_string(), + }))), + tool_complete_result: Mutex::new(Some(Err(LlmError::AuthFailed { + provider: name.to_string(), + }))), + } + } + + fn failing_rate_limited(name: &str) -> Self { + Self { + name: name.to_string(), + input_cost: Decimal::ZERO, + output_cost: Decimal::ZERO, + complete_result: Mutex::new(Some(Err(LlmError::RateLimited { + provider: name.to_string(), + retry_after: Some(Duration::from_secs(30)), + }))), + tool_complete_result: Mutex::new(Some(Err(LlmError::RateLimited { + provider: name.to_string(), + retry_after: Some(Duration::from_secs(30)), + }))), + } + } + } + + #[async_trait] + impl LlmProvider for MockProvider { + fn model_name(&self) -> &str { + &self.name + } + + fn cost_per_token(&self) -> (Decimal, Decimal) { + (self.input_cost, self.output_cost) + } + + async fn complete( + &self, + _request: CompletionRequest, + ) -> Result { + self.complete_result + .lock() + .unwrap() + .take() + .expect("MockProvider::complete called more than once") + } + + async fn complete_with_tools( + &self, + _request: ToolCompletionRequest, + ) -> Result { + self.tool_complete_result + .lock() + .unwrap() + .take() + .expect("MockProvider::complete_with_tools called more than once") + } + + async fn list_models(&self) -> Result, LlmError> { + Ok(vec![self.name.clone()]) + } + } + + fn make_request() -> CompletionRequest { + CompletionRequest::new(vec![crate::llm::ChatMessage::user("hello")]) + } + + fn make_tool_request() -> ToolCompletionRequest { + ToolCompletionRequest::new(vec![crate::llm::ChatMessage::user("hello")], vec![]) + } + + // Test 1: Primary succeeds, no failover occurs. + #[tokio::test] + async fn primary_succeeds_no_failover() { + let primary = Arc::new(MockProvider::succeeding("primary", "primary response")); + let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response")); + + let failover = FailoverProvider::new(vec![primary, fallback]).unwrap(); + + let response = failover.complete(make_request()).await.unwrap(); + assert_eq!(response.content, "primary response"); + } + + // Test 2: Primary fails with retryable error, fallback succeeds. + #[tokio::test] + async fn primary_fails_retryable_fallback_succeeds() { + let primary = Arc::new(MockProvider::failing_retryable("primary")); + let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response")); + + let failover = FailoverProvider::new(vec![primary, fallback]).unwrap(); + + let response = failover.complete(make_request()).await.unwrap(); + assert_eq!(response.content, "fallback response"); + } + + // Test 3: All providers fail, returns last error. + #[tokio::test] + async fn all_providers_fail_returns_last_error() { + let primary = Arc::new(MockProvider::failing_retryable("primary")); + let fallback = Arc::new(MockProvider::failing_retryable("fallback")); + + let failover = FailoverProvider::new(vec![primary, fallback]).unwrap(); + + let err = failover.complete(make_request()).await.unwrap_err(); + match err { + LlmError::RequestFailed { provider, .. } => { + assert_eq!(provider, "fallback"); + } + other => panic!("expected RequestFailed, got: {other:?}"), + } + } + + // Test 4: Non-retryable error fails immediately, no failover. + #[tokio::test] + async fn non_retryable_error_fails_immediately() { + let primary = Arc::new(MockProvider::failing_non_retryable("primary")); + let fallback = Arc::new(MockProvider::succeeding("fallback", "fallback response")); + + let failover = FailoverProvider::new(vec![primary, fallback]).unwrap(); + + let err = failover.complete(make_request()).await.unwrap_err(); + match err { + LlmError::AuthFailed { provider } => { + assert_eq!(provider, "primary"); + } + other => panic!("expected AuthFailed, got: {other:?}"), + } + } + + // Test 5: Three providers, first two fail (retryable), third succeeds. + #[tokio::test] + async fn three_providers_first_two_fail_third_succeeds() { + let p1 = Arc::new(MockProvider::failing_retryable("provider-1")); + let p2 = Arc::new(MockProvider::failing_rate_limited("provider-2")); + let p3 = Arc::new(MockProvider::succeeding("provider-3", "third time lucky")); + + let failover = FailoverProvider::new(vec![p1, p2, p3]).unwrap(); + + let response = failover.complete(make_request()).await.unwrap(); + assert_eq!(response.content, "third time lucky"); + } + + // Test: complete_with_tools follows same failover logic. + #[tokio::test] + async fn complete_with_tools_failover() { + let primary = Arc::new(MockProvider::failing_retryable("primary")); + let fallback = Arc::new(MockProvider::succeeding("fallback", "tools fallback")); + + let failover = FailoverProvider::new(vec![primary, fallback]).unwrap(); + + let response = failover + .complete_with_tools(make_tool_request()) + .await + .unwrap(); + assert_eq!(response.content.as_deref(), Some("tools fallback")); + } + + // Test: model_name and cost_per_token reflect the last-used provider. + #[tokio::test] + async fn model_name_and_cost_track_last_used_provider() { + let fallback_cost = Decimal::new(15, 6); // 0.000015 + + let primary = Arc::new(MockProvider::failing_retryable("primary-model")); + let fallback = Arc::new(MockProvider::succeeding_with_cost( + "fallback-model", + "ok", + fallback_cost, + fallback_cost, + )); + + let failover = FailoverProvider::new(vec![primary, fallback]).unwrap(); + + // Before any call, defaults to primary (index 0). + assert_eq!(failover.model_name(), "primary-model"); + assert_eq!(failover.cost_per_token(), (Decimal::ZERO, Decimal::ZERO)); + + // After failover, should reflect the fallback provider. + let _ = failover.complete(make_request()).await.unwrap(); + assert_eq!(failover.model_name(), "fallback-model"); + assert_eq!(failover.cost_per_token(), (fallback_cost, fallback_cost)); + } + + // Test: list_models aggregates from all providers. + #[tokio::test] + async fn list_models_aggregates_all() { + let p1 = Arc::new(MockProvider::succeeding("model-a", "ok")); + let p2 = Arc::new(MockProvider::succeeding("model-b", "ok")); + + let failover = FailoverProvider::new(vec![p1, p2]).unwrap(); + + let models = failover.list_models().await.unwrap(); + assert!(models.contains(&"model-a".to_string())); + assert!(models.contains(&"model-b".to_string())); + } + + // Test: is_retryable correctly classifies errors. + #[test] + fn retryable_classification() { + // Retryable + assert!(is_retryable(&LlmError::RequestFailed { + provider: "p".into(), + reason: "err".into(), + })); + assert!(is_retryable(&LlmError::RateLimited { + provider: "p".into(), + retry_after: None, + })); + assert!(is_retryable(&LlmError::InvalidResponse { + provider: "p".into(), + reason: "bad json".into(), + })); + assert!(is_retryable(&LlmError::SessionRenewalFailed { + provider: "p".into(), + reason: "timeout".into(), + })); + assert!(is_retryable(&LlmError::Io(std::io::Error::new( + std::io::ErrorKind::ConnectionReset, + "reset" + )))); + assert!(is_retryable(&LlmError::ModelNotAvailable { + provider: "p".into(), + model: "m".into(), + })); + + // Non-retryable + assert!(!is_retryable(&LlmError::AuthFailed { + provider: "p".into(), + })); + assert!(!is_retryable(&LlmError::SessionExpired { + provider: "p".into(), + })); + assert!(!is_retryable(&LlmError::ContextLengthExceeded { + used: 100_000, + limit: 50_000, + })); + } + + // Test: empty providers list returns error (not panic). + #[test] + fn empty_providers_returns_error() { + let result = FailoverProvider::new(vec![]); + assert!(result.is_err()); + } +} diff --git a/src/llm/mod.rs b/src/llm/mod.rs index ee801df7..95b9981a 100644 --- a/src/llm/mod.rs +++ b/src/llm/mod.rs @@ -8,13 +8,16 @@ //! - **OpenAI-compatible**: Any endpoint that speaks the OpenAI API mod costs; +pub mod failover; mod nearai; mod nearai_chat; mod provider; mod reasoning; +mod retry; mod rig_adapter; pub mod session; +pub use failover::FailoverProvider; pub use nearai::{ModelInfo, NearAiProvider}; pub use nearai_chat::NearAiChatProvider; pub use provider::{ @@ -33,7 +36,7 @@ use std::sync::Arc; use rig::client::CompletionClient; use secrecy::ExposeSecret; -use crate::config::{LlmBackend, LlmConfig, NearAiApiMode}; +use crate::config::{LlmBackend, LlmConfig, NearAiApiMode, NearAiConfig}; use crate::error::LlmError; /// Create an LLM provider based on configuration. @@ -46,7 +49,7 @@ pub fn create_llm_provider( session: Arc, ) -> Result, LlmError> { match config.backend { - LlmBackend::NearAi => create_nearai_provider(config, session), + LlmBackend::NearAi => create_llm_provider_with_config(&config.nearai, session), LlmBackend::OpenAi => create_openai_provider(config), LlmBackend::Anthropic => create_anthropic_provider(config), LlmBackend::Ollama => create_ollama_provider(config), @@ -54,21 +57,28 @@ pub fn create_llm_provider( } } -fn create_nearai_provider( - config: &LlmConfig, +/// Create an LLM provider from a `NearAiConfig` directly. +/// +/// This is useful when constructing additional providers for failover, +/// where only the model name differs from the primary config. +pub fn create_llm_provider_with_config( + config: &NearAiConfig, session: Arc, ) -> Result, LlmError> { - match config.nearai.api_mode { + match config.api_mode { NearAiApiMode::Responses => { - tracing::info!("Using NEAR AI Responses API (chat-api) with session auth"); - Ok(Arc::new(NearAiProvider::new( - config.nearai.clone(), - session, - ))) + tracing::info!( + model = %config.model, + "Using Responses API (chat-api) with session auth" + ); + Ok(Arc::new(NearAiProvider::new(config.clone(), session))) } NearAiApiMode::ChatCompletions => { - tracing::info!("Using NEAR AI Chat Completions API (cloud-api) with API key auth"); - Ok(Arc::new(NearAiChatProvider::new(config.nearai.clone())?)) + tracing::info!( + model = %config.model, + "Using Chat Completions API (cloud-api) with API key auth" + ); + Ok(Arc::new(NearAiChatProvider::new(config.clone())?)) } } } diff --git a/src/llm/nearai.rs b/src/llm/nearai.rs index ed32d2fc..70c966e0 100644 --- a/src/llm/nearai.rs +++ b/src/llm/nearai.rs @@ -19,6 +19,7 @@ use crate::llm::provider::{ ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, }; +use crate::llm::retry::{is_retryable_status, retry_backoff_delay}; use crate::llm::session::SessionManager; /// Information about an available model from NEAR AI API. @@ -270,88 +271,139 @@ impl NearAiProvider { } } - /// Inner request implementation without retry logic. + /// Inner request implementation with retry logic for transient errors. + /// + /// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff. + /// Does not retry on client errors (400, 401, 403, 404) or parse errors. async fn send_request_inner Deserialize<'de>>( &self, path: &str, body: &T, ) -> Result { let url = self.api_url(path); - let token = self.session.get_token().await?; + let max_retries = self.config.max_retries; - tracing::debug!("Sending request to NEAR AI: {}", url); - tracing::debug!("Request body: {:?}", body); + for attempt in 0..=max_retries { + let token = self.session.get_token().await?; - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", token.expose_secret())) - .header("Content-Type", "application/json") - .json(body) - .send() - .await - .map_err(|e| { - tracing::error!("NEAR AI request failed: {}", e); - e - })?; + tracing::debug!( + "Sending request to NEAR AI: {} (attempt {})", + url, + attempt + 1 + ); + tracing::debug!("Request body: {:?}", body); - let status = response.status(); - let response_text = response.text().await.unwrap_or_default(); + let response = self + .client + .post(&url) + .header("Authorization", format!("Bearer {}", token.expose_secret())) + .header("Content-Type", "application/json") + .json(body) + .send() + .await; - tracing::debug!("NEAR AI response status: {}", status); - tracing::debug!("NEAR AI response body: {}", response_text); + let response = match response { + Ok(r) => r, + Err(e) => { + tracing::error!("NEAR AI request failed: {}", e); + // Network errors (timeout, connection refused) are transient + if attempt < max_retries { + let delay = retry_backoff_delay(attempt); + tracing::warn!( + "NEAR AI request error (attempt {}/{}), retrying in {:?}: {}", + attempt + 1, + max_retries + 1, + delay, + e, + ); + tokio::time::sleep(delay).await; + continue; + } + return Err(e.into()); + } + }; - if !status.is_success() { - // Check for session expiration (401 with specific message patterns) - if status.as_u16() == 401 { - let is_session_expired = response_text.to_lowercase().contains("session") - && (response_text.to_lowercase().contains("expired") - || response_text.to_lowercase().contains("invalid")); + let status = response.status(); + let response_text = response.text().await.unwrap_or_default(); - if is_session_expired { - return Err(LlmError::SessionExpired { + tracing::debug!("NEAR AI response status: {}", status); + tracing::debug!("NEAR AI response body: {}", response_text); + + if !status.is_success() { + let status_code = status.as_u16(); + + // Check for session expiration (401 with specific message patterns) + if status_code == 401 { + let lower = response_text.to_lowercase(); + let is_session_expired = lower.contains("session") + && (lower.contains("expired") || lower.contains("invalid")); + + if is_session_expired { + return Err(LlmError::SessionExpired { + provider: "nearai".to_string(), + }); + } + + // Generic 401 -- not retryable + return Err(LlmError::AuthFailed { provider: "nearai".to_string(), }); } - // Generic 401 without session expiration indication - return Err(LlmError::AuthFailed { - provider: "nearai".to_string(), - }); - } + // Check if this is a transient error worth retrying + if is_retryable_status(status_code) && attempt < max_retries { + let delay = retry_backoff_delay(attempt); + tracing::warn!( + "NEAR AI returned HTTP {} (attempt {}/{}), retrying in {:?}", + status_code, + attempt + 1, + max_retries + 1, + delay, + ); + tokio::time::sleep(delay).await; + continue; + } - // Try to parse as JSON error - if let Ok(error) = serde_json::from_str::(&response_text) { - if status.as_u16() == 429 { - return Err(LlmError::RateLimited { + // Non-retryable error or exhausted retries + if let Ok(error) = serde_json::from_str::(&response_text) { + if status_code == 429 { + return Err(LlmError::RateLimited { + provider: "nearai".to_string(), + retry_after: None, + }); + } + return Err(LlmError::RequestFailed { provider: "nearai".to_string(), - retry_after: None, + reason: error.error, }); } + return Err(LlmError::RequestFailed { provider: "nearai".to_string(), - reason: error.error, + reason: format!("HTTP {}: {}", status, response_text), }); } - return Err(LlmError::RequestFailed { - provider: "nearai".to_string(), - reason: format!("HTTP {}: {}", status, response_text), - }); + // Success -- parse the response + return match serde_json::from_str::(&response_text) { + Ok(parsed) => Ok(parsed), + Err(e) => { + tracing::debug!("Response is not expected JSON format: {}", e); + tracing::debug!("Will try alternative parsing in caller"); + Err(LlmError::InvalidResponse { + provider: "nearai".to_string(), + reason: format!("Parse error: {}. Raw: {}", e, response_text), + }) + } + }; } - // Try to parse as our expected type - match serde_json::from_str::(&response_text) { - Ok(parsed) => Ok(parsed), - Err(e) => { - tracing::debug!("Response is not expected JSON format: {}", e); - tracing::debug!("Will try alternative parsing in caller"); - Err(LlmError::InvalidResponse { - provider: "nearai".to_string(), - reason: format!("Parse error: {}. Raw: {}", e, response_text), - }) - } - } + // This is unreachable because the loop always returns, but the compiler + // cannot prove that. Return a generic error as a safety net. + Err(LlmError::RequestFailed { + provider: "nearai".to_string(), + reason: "retry loop exited unexpectedly".to_string(), + }) } } diff --git a/src/llm/nearai_chat.rs b/src/llm/nearai_chat.rs index 13a36a91..dbe51bc7 100644 --- a/src/llm/nearai_chat.rs +++ b/src/llm/nearai_chat.rs @@ -16,6 +16,7 @@ use crate::llm::provider::{ ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, ModelMetadata, Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, }; +use crate::llm::retry::{is_retryable_status, retry_backoff_delay}; /// NEAR AI Chat Completions API provider. pub struct NearAiChatProvider { @@ -62,64 +63,116 @@ impl NearAiChatProvider { .unwrap_or_default() } - /// Send a request to the chat completions API. + /// Send a request to the chat completions API with retry on transient errors. + /// + /// Retries on HTTP 429, 500, 502, 503, 504 with exponential backoff. + /// Does not retry on client errors (400, 401, 403, 404) or parse errors. async fn send_request Deserialize<'de>>( &self, body: &T, ) -> Result { let url = self.api_url("chat/completions"); + let max_retries = self.config.max_retries; - tracing::debug!("Sending request to NEAR AI Chat: {}", url); + for attempt in 0..=max_retries { + tracing::debug!( + "Sending request to NEAR AI Chat: {} (attempt {})", + url, + attempt + 1, + ); - if tracing::enabled!(tracing::Level::DEBUG) - && let Ok(json) = serde_json::to_string(body) - { - tracing::debug!("NEAR AI Chat request body: {}", json); - } + if tracing::enabled!(tracing::Level::DEBUG) + && let Ok(json) = serde_json::to_string(body) + { + tracing::debug!("NEAR AI Chat request body: {}", json); + } - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key())) - .header("Content-Type", "application/json") - .json(body) - .send() - .await - .map_err(|e| { - tracing::error!("NEAR AI Chat request failed: {}", e); - LlmError::RequestFailed { - provider: "nearai_chat".to_string(), - reason: e.to_string(), + let response = self + .client + .post(&url) + .header("Authorization", format!("Bearer {}", self.api_key())) + .header("Content-Type", "application/json") + .json(body) + .send() + .await; + + let response = match response { + Ok(r) => r, + Err(e) => { + tracing::error!("NEAR AI Chat request failed: {}", e); + if attempt < max_retries { + let delay = retry_backoff_delay(attempt); + tracing::warn!( + "NEAR AI Chat request error (attempt {}/{}), retrying in {:?}: {}", + attempt + 1, + max_retries + 1, + delay, + e, + ); + tokio::time::sleep(delay).await; + continue; + } + return Err(LlmError::RequestFailed { + provider: "nearai_chat".to_string(), + reason: e.to_string(), + }); } - })?; + }; - let status = response.status(); - let response_text = response.text().await.unwrap_or_default(); + let status = response.status(); + let response_text = response.text().await.unwrap_or_default(); - tracing::debug!("NEAR AI Chat response status: {}", status); - tracing::debug!("NEAR AI Chat response body: {}", response_text); + tracing::debug!("NEAR AI Chat response status: {}", status); + tracing::debug!("NEAR AI Chat response body: {}", response_text); - if !status.is_success() { - if status.as_u16() == 401 { - return Err(LlmError::AuthFailed { + if !status.is_success() { + let status_code = status.as_u16(); + + // Auth errors are not retryable + if status_code == 401 { + return Err(LlmError::AuthFailed { + provider: "nearai_chat".to_string(), + }); + } + + // Transient errors: retry with backoff + if is_retryable_status(status_code) && attempt < max_retries { + let delay = retry_backoff_delay(attempt); + tracing::warn!( + "NEAR AI Chat returned HTTP {} (attempt {}/{}), retrying in {:?}", + status_code, + attempt + 1, + max_retries + 1, + delay, + ); + tokio::time::sleep(delay).await; + continue; + } + + // Non-retryable or exhausted retries + if status_code == 429 { + return Err(LlmError::RateLimited { + provider: "nearai_chat".to_string(), + retry_after: None, + }); + } + return Err(LlmError::RequestFailed { provider: "nearai_chat".to_string(), + reason: format!("HTTP {}: {}", status, response_text), }); } - if status.as_u16() == 429 { - return Err(LlmError::RateLimited { - provider: "nearai_chat".to_string(), - retry_after: None, - }); - } - return Err(LlmError::RequestFailed { + + // Success — parse the response + return serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse { provider: "nearai_chat".to_string(), - reason: format!("HTTP {}: {}", status, response_text), + reason: format!("JSON parse error: {}. Raw: {}", e, response_text), }); } - serde_json::from_str(&response_text).map_err(|e| LlmError::InvalidResponse { + // Safety net: unreachable because the loop always returns + Err(LlmError::RequestFailed { provider: "nearai_chat".to_string(), - reason: format!("JSON parse error: {}. Raw: {}", e, response_text), + reason: "retry loop exited unexpectedly".to_string(), }) } diff --git a/src/llm/retry.rs b/src/llm/retry.rs new file mode 100644 index 00000000..2dced0b0 --- /dev/null +++ b/src/llm/retry.rs @@ -0,0 +1,96 @@ +//! Shared retry helpers for LLM providers. +//! +//! Provides exponential backoff with jitter and retryable status classification +//! used by both `NearAiProvider` and `NearAiChatProvider`. + +use std::time::Duration; + +use rand::Rng; + +/// Returns `true` if the HTTP status code is transient and worth retrying. +pub(crate) fn is_retryable_status(status: u16) -> bool { + matches!(status, 429 | 500 | 502 | 503 | 504) +} + +/// Calculate exponential backoff delay with random jitter. +/// +/// Base delay is 1 second, doubled each attempt, with +/-25% jitter. +/// - attempt 0: ~1s (0.75s - 1.25s) +/// - attempt 1: ~2s (1.5s - 2.5s) +/// - attempt 2: ~4s (3.0s - 5.0s) +pub(crate) fn retry_backoff_delay(attempt: u32) -> Duration { + let base_ms: u64 = 1000u64.saturating_mul(2u64.saturating_pow(attempt)); + let jitter_range = base_ms / 4; // 25% + let jitter = if jitter_range > 0 { + let offset = rand::thread_rng().gen_range(0..=jitter_range * 2); + offset as i64 - jitter_range as i64 + } else { + 0 + }; + let delay_ms = (base_ms as i64 + jitter).max(100) as u64; + Duration::from_millis(delay_ms) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_is_retryable_status() { + // Transient errors should be retryable + assert!(is_retryable_status(429)); + assert!(is_retryable_status(500)); + assert!(is_retryable_status(502)); + assert!(is_retryable_status(503)); + assert!(is_retryable_status(504)); + + // Client errors should not be retryable + assert!(!is_retryable_status(400)); + assert!(!is_retryable_status(401)); + assert!(!is_retryable_status(403)); + assert!(!is_retryable_status(404)); + assert!(!is_retryable_status(422)); + + // Success codes should not be retryable + assert!(!is_retryable_status(200)); + assert!(!is_retryable_status(201)); + } + + #[test] + fn test_retry_backoff_delay_exponential_growth() { + // Run multiple samples to verify the range, accounting for jitter + for _ in 0..20 { + let d0 = retry_backoff_delay(0); + let d1 = retry_backoff_delay(1); + let d2 = retry_backoff_delay(2); + + // Attempt 0: base 1000ms, jitter +/-250ms -> [750, 1250] + assert!(d0.as_millis() >= 750, "attempt 0 too low: {:?}", d0); + assert!(d0.as_millis() <= 1250, "attempt 0 too high: {:?}", d0); + + // Attempt 1: base 2000ms, jitter +/-500ms -> [1500, 2500] + assert!(d1.as_millis() >= 1500, "attempt 1 too low: {:?}", d1); + assert!(d1.as_millis() <= 2500, "attempt 1 too high: {:?}", d1); + + // Attempt 2: base 4000ms, jitter +/-1000ms -> [3000, 5000] + assert!(d2.as_millis() >= 3000, "attempt 2 too low: {:?}", d2); + assert!(d2.as_millis() <= 5000, "attempt 2 too high: {:?}", d2); + } + } + + #[test] + fn test_retry_backoff_delay_minimum() { + // Even at attempt 0, delay should be at least 100ms (the minimum floor) + for _ in 0..20 { + let delay = retry_backoff_delay(0); + assert!(delay.as_millis() >= 100); + } + } + + #[test] + fn test_retry_backoff_delay_no_overflow() { + // Very high attempt numbers should not panic from overflow + let delay = retry_backoff_delay(30); + assert!(delay.as_millis() >= 100); + } +} diff --git a/src/main.rs b/src/main.rs index 97355aaa..ca9754ad 100644 --- a/src/main.rs +++ b/src/main.rs @@ -22,7 +22,10 @@ use ironclaw::{ config::Config, context::ContextManager, extensions::ExtensionManager, - llm::{SessionConfig, create_llm_provider, create_session_manager}, + llm::{ + FailoverProvider, LlmProvider, SessionConfig, create_llm_provider, + create_llm_provider_with_config, create_session_manager, + }, orchestrator::{ ContainerJobConfig, ContainerJobManager, OrchestratorApi, TokenStore, api::OrchestratorState, @@ -447,6 +450,27 @@ async fn main() -> anyhow::Result<()> { let llm = create_llm_provider(&config.llm, session.clone())?; tracing::info!("LLM provider initialized: {}", llm.model_name()); + // Wrap in failover if a fallback model is configured + let llm: Arc = + if let Some(fallback_model) = config.llm.nearai.fallback_model.as_ref() { + if fallback_model == &config.llm.nearai.model { + tracing::warn!( + "fallback_model is the same as primary model, failover may not be effective" + ); + } + let mut fallback_config = config.llm.nearai.clone(); + fallback_config.model = fallback_model.clone(); + let fallback = create_llm_provider_with_config(&fallback_config, session.clone())?; + tracing::info!( + primary = %llm.model_name(), + fallback = %fallback.model_name(), + "LLM failover enabled" + ); + Arc::new(FailoverProvider::new(vec![llm, fallback])?) + } else { + llm + }; + // Initialize safety layer let safety = Arc::new(SafetyLayer::new(&config.safety)); tracing::info!("Safety layer initialized"); diff --git a/src/setup/wizard.rs b/src/setup/wizard.rs index 696a7a08..16be4d30 100644 --- a/src/setup/wizard.rs +++ b/src/setup/wizard.rs @@ -653,6 +653,8 @@ impl SetupWizard { session_path: crate::llm::session::default_session_path(), api_mode: crate::config::NearAiApiMode::Responses, api_key: None, + fallback_model: None, + max_retries: 3, }, openai: None, anthropic: None,