diff --git a/crates/goose-provider-types/src/errors.rs b/crates/goose-provider-types/src/errors.rs index c8d099c29abc..cbe038447ce1 100644 --- a/crates/goose-provider-types/src/errors.rs +++ b/crates/goose-provider-types/src/errors.rs @@ -6,6 +6,9 @@ use crate::request_log::LogError; #[derive(Error, Debug, Clone, PartialEq)] pub enum ProviderError { + #[error("Provider is not configured")] + NotConfigured, + #[error("Authentication error: {0}")] Authentication(String), @@ -59,6 +62,7 @@ impl ProviderError { pub fn telemetry_type(&self) -> &'static str { match self { + ProviderError::NotConfigured => "not_configured", ProviderError::Authentication(_) => "auth", ProviderError::ContextLengthExceeded(_) => "context_length", ProviderError::RateLimitExceeded { .. } => "rate_limit", @@ -131,16 +135,23 @@ fn provider_error_from_reqwest(error: &reqwest::Error) -> ProviderError { impl From for ProviderError { fn from(error: anyhow::Error) -> Self { - if let Some(provider_error) = error.downcast_ref::() { + if let Some(provider_error) = error + .chain() + .find_map(|cause| cause.downcast_ref::()) + { return provider_error.clone(); } - if let Some(reqwest_err) = error.downcast_ref::() { + if let Some(reqwest_err) = error + .chain() + .find_map(|cause| cause.downcast_ref::()) + { return provider_error_from_reqwest(reqwest_err); } - if error - .downcast_ref::() - .is_some() - { + if error.chain().any(|cause| { + cause + .downcast_ref::() + .is_some() + }) { return ProviderError::NetworkError( "Request timed out — check your network connection and try again.".to_string(), ); diff --git a/crates/goose/src/acp/server/providers.rs b/crates/goose/src/acp/server/providers.rs index d1a496137f0c..4c8219d8f1a4 100644 --- a/crates/goose/src/acp/server/providers.rs +++ b/crates/goose/src/acp/server/providers.rs @@ -478,10 +478,21 @@ impl GooseAcpAgent { .create_provider(&req.provider_id, Vec::new(), None) .await .internal_err_ctx("Failed to initialize provider")?; - let models = provider - .fetch_supported_models() - .await - .internal_err_ctx("Failed to fetch provider supported models")?; + let models = match provider.fetch_supported_models().await { + Ok(models) => models, + Err(goose_providers::errors::ProviderError::Authentication(error)) => { + return Err(agent_client_protocol::Error::auth_required().data(error)); + } + Err(goose_providers::errors::ProviderError::NotConfigured) => { + return Err(agent_client_protocol::Error::invalid_params() + .data(format!("Provider is not configured: {}", req.provider_id))); + } + Err(error) => { + return Err(agent_client_protocol::Error::internal_error().data(format!( + "Failed to fetch provider supported models: {error}" + ))); + } + }; Ok(ProviderSupportedModelsListResponse { provider_id: req.provider_id, diff --git a/crates/goose/src/doctor.rs b/crates/goose/src/doctor.rs index cd7f91acff9c..7f794634d475 100644 --- a/crates/goose/src/doctor.rs +++ b/crates/goose/src/doctor.rs @@ -243,6 +243,9 @@ async fn try_other_providers( fn describe_error(e: &ProviderError) -> String { match e { + ProviderError::NotConfigured => { + "Provider is not configured. Run `goose configure` to set it up.".to_string() + } ProviderError::Authentication(_) => { "Authentication failed — check your API key. Run `goose configure` to update it." .to_string() diff --git a/crates/goose/src/providers/githubcopilot.rs b/crates/goose/src/providers/githubcopilot.rs index d1550afcae9b..81bb6ad858e2 100644 --- a/crates/goose/src/providers/githubcopilot.rs +++ b/crates/goose/src/providers/githubcopilot.rs @@ -5,7 +5,7 @@ use crate::providers::openai_compatible::{ handle_status, stream_openai_compat, stream_responses_compat, }; use crate::providers::private_file::write_private_file; -use anyhow::{anyhow, Context, Result}; +use anyhow::{anyhow, Result}; use async_trait::async_trait; use axum::http; use chrono::{DateTime, Utc}; @@ -288,7 +288,7 @@ impl GithubCopilotProvider { .map_err(|e| e.into()) } - async fn get_api_info(&self) -> Result<(String, String)> { + async fn get_api_info(&self) -> Result<(String, String), ProviderError> { let guard = self.mu.lock().await; if let Some(state) = guard.borrow().as_ref() { @@ -306,53 +306,64 @@ impl GithubCopilotProvider { } } + let config = Config::global(); + let github_token = match config.get_secret::("GITHUB_COPILOT_TOKEN") { + Ok(token) => token, + Err(ConfigError::NotFound(_)) => return Err(ProviderError::NotConfigured), + Err(error) => return Err(ProviderError::ExecutionError(error.to_string())), + }; + const MAX_ATTEMPTS: i32 = 3; + let mut last_error = None; for attempt in 0..MAX_ATTEMPTS { tracing::trace!("attempt {} to refresh api info", attempt + 1); - let info = match self.refresh_api_info().await { + let info = match self.refresh_api_info(&github_token).await { Ok(data) => data, Err(err) => { tracing::warn!("failed to refresh api info: {}", err); + last_error = Some(err); continue; } }; let expires_at = Utc::now() + chrono::Duration::seconds(info.refresh_in); let new_state = CopilotState { info, expires_at }; - self.cache.save(&new_state).await?; + self.cache + .save(&new_state) + .await + .map_err(ProviderError::from)?; guard.replace(Some(new_state.clone())); return Ok((new_state.info.endpoints.api, new_state.info.token)); } - Err(anyhow!("failed to get api info after 3 attempts")) + Err(last_error.unwrap()) } - async fn refresh_api_info(&self) -> Result { - let config = Config::global(); - let token = match config.get_secret::("GITHUB_COPILOT_TOKEN") { - Ok(token) => token, - Err(err) => match err { - ConfigError::NotFound(_) => { - let token = self - .get_access_token() - .await - .context("unable to login into github")?; - config.set_secret("GITHUB_COPILOT_TOKEN", &token)?; - token - } - _ => return Err(err.into()), - }, - }; - let resp = self + async fn refresh_api_info( + &self, + github_token: &str, + ) -> Result { + let response = self .client .get(&self.urls.copilot_token_url) .headers(self.get_github_headers()) - .header(http::header::AUTHORIZATION, format!("bearer {}", &token)) + .header( + http::header::AUTHORIZATION, + format!("bearer {github_token}"), + ) .send() - .await? - .error_for_status()? - .text() .await?; + if matches!( + response.status(), + reqwest::StatusCode::UNAUTHORIZED | reqwest::StatusCode::FORBIDDEN + ) { + return Err(ProviderError::Authentication(format!( + "GitHub Copilot token request failed ({})", + response.status() + ))); + } + let resp = response.error_for_status()?.text().await?; tracing::trace!("copilot token response: {}", resp); - let info: CopilotTokenInfo = serde_json::from_str(&resp)?; + let info: CopilotTokenInfo = serde_json::from_str(&resp) + .map_err(|error| ProviderError::RequestFailed(error.to_string()))?; Ok(info) } @@ -659,8 +670,8 @@ impl Provider for GithubCopilotProvider { async fn configure_oauth(&self) -> Result<(), ProviderError> { let config = Config::global(); - if config.get_secret::("GITHUB_COPILOT_TOKEN").is_ok() { - match self.refresh_api_info().await { + if let Ok(github_token) = config.get_secret::("GITHUB_COPILOT_TOKEN") { + match self.refresh_api_info(&github_token).await { Ok(_) => return Ok(()), Err(_) => { tracing::debug!("Existing token is invalid, starting OAuth flow"); @@ -720,6 +731,8 @@ fn promote_tool_choice(response: Value) -> Value { mod tests { use super::*; use serde_json::json; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; #[cfg(unix)] #[tokio::test] @@ -757,6 +770,74 @@ mod tests { assert_eq!(saved.info.token, "copilot-secret"); } + #[tokio::test] + async fn get_api_info_uses_valid_cache_without_github_token() { + let directory = tempfile::tempdir().unwrap(); + let cache = DiskCache { + cache_path: directory.path().join("info.json"), + }; + let state = CopilotState { + expires_at: Utc::now() + chrono::Duration::minutes(10), + info: CopilotTokenInfo { + token: "copilot-secret".to_string(), + expires_at: 1, + refresh_in: 600, + endpoints: CopilotTokenEndpoints { + api: "https://api.githubcopilot.com".to_string(), + _extra: HashMap::new(), + }, + _extra: HashMap::new(), + }, + }; + cache.save(&state).await.unwrap(); + let provider = GithubCopilotProvider { + client: Client::new(), + cache, + mu: tokio::sync::Mutex::new(RefCell::new(None)), + urls: GithubCopilotUrls::new("github.com", None), + client_id: DEFAULT_GITHUB_COPILOT_CLIENT_ID.to_string(), + name: GITHUB_COPILOT_PROVIDER_NAME.to_string(), + tls_config: None, + }; + + let (endpoint, token) = provider.get_api_info().await.unwrap(); + + assert_eq!(endpoint, "https://api.githubcopilot.com"); + assert_eq!(token, "copilot-secret"); + } + + #[tokio::test] + async fn refresh_api_info_returns_authentication_for_rejected_token() { + for status in [401, 403] { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/copilot-token")) + .respond_with(ResponseTemplate::new(status)) + .mount(&server) + .await; + let directory = tempfile::tempdir().unwrap(); + let provider = GithubCopilotProvider { + client: Client::new(), + cache: DiskCache { + cache_path: directory.path().join("info.json"), + }, + mu: tokio::sync::Mutex::new(RefCell::new(None)), + urls: GithubCopilotUrls { + device_code_url: String::new(), + access_token_url: String::new(), + copilot_token_url: format!("{}/copilot-token", server.uri()), + }, + client_id: DEFAULT_GITHUB_COPILOT_CLIENT_ID.to_string(), + name: GITHUB_COPILOT_PROVIDER_NAME.to_string(), + tls_config: None, + }; + + let error = provider.refresh_api_info("rejected").await.unwrap_err(); + + assert!(matches!(error, ProviderError::Authentication(_))); + } + } + #[test] fn responses_models_routed_correctly() { assert!(is_openai_responses_model("gpt-5.5")); diff --git a/crates/goose/src/providers/inventory/registrations.rs b/crates/goose/src/providers/inventory/registrations.rs index b27679fa6d8d..c42e6660c28d 100644 --- a/crates/goose/src/providers/inventory/registrations.rs +++ b/crates/goose/src/providers/inventory/registrations.rs @@ -15,7 +15,7 @@ use crate::providers::gemini_oauth::TokenCache as GeminiOAuthTokenCache; use crate::providers::google::{GOOGLE_API_HOST, GOOGLE_PROVIDER_NAME}; use crate::providers::huggingface::HuggingFaceProvider; use crate::providers::huggingface_auth; -use crate::providers::kimicode::KIMI_CONFIGURED_MARKER; +use crate::providers::kimicode; use crate::providers::ollama::OLLAMA_PROVIDER_NAME; use crate::providers::openai::{OPEN_AI_DEFAULT_BASE_PATH, OPEN_AI_PROVIDER_NAME}; use crate::providers::pi_acp::{PI_ACP_BINARY, PI_ACP_PROVIDER_NAME}; @@ -196,11 +196,7 @@ pub fn refresh_only() -> InventoryRegistration { } pub fn kimi_code_inventory() -> InventoryRegistration { - refresh_only().with_configured(|| { - Config::global() - .get_param::(KIMI_CONFIGURED_MARKER) - .unwrap_or(false) - }) + refresh_only().with_configured(kimicode::has_configured_token) } pub fn chatgpt_codex_inventory() -> InventoryRegistration { @@ -325,4 +321,34 @@ mod tests { assert!(configured()); } + + #[test] + #[serial_test::serial] + fn kimi_code_inventory_configured_uses_token_cache() { + let root = tempfile::tempdir().unwrap(); + let root_path = root.path().to_string_lossy().to_string(); + let _guard = env_lock::lock_env([("GOOSE_PATH_ROOT", Some(root_path.as_str()))]); + + let registration = kimi_code_inventory(); + let configured = registration + .configured + .expect("Kimi Code should define configured resolver"); + + assert!(!configured()); + + let cache_path = Paths::in_config_dir("kimicode/token.json"); + std::fs::create_dir_all(cache_path.parent().unwrap()).unwrap(); + std::fs::write( + cache_path, + serde_json::to_string(&serde_json::json!({ + "access_token": "access", + "refresh_token": "refresh", + "expires_at": (Utc::now() + chrono::Duration::hours(1)).to_rfc3339(), + })) + .unwrap(), + ) + .unwrap(); + + assert!(configured()); + } } diff --git a/crates/goose/src/providers/kimicode.rs b/crates/goose/src/providers/kimicode.rs index 3e0a38e39194..b36a8f4b33ce 100644 --- a/crates/goose/src/providers/kimicode.rs +++ b/crates/goose/src/providers/kimicode.rs @@ -1,5 +1,4 @@ use crate::config::paths::Paths; -use crate::config::Config; use anyhow::Result; use async_stream::try_stream; use async_trait::async_trait; @@ -23,7 +22,8 @@ use super::base::{ }; use super::formats::anthropic::{create_request, response_to_streaming_message}; use super::oauth_device_flow::{ - refresh_device_flow_token, run_device_flow, DeviceFlowConfig, DeviceFlowTokens, RequestEncoding, + refresh_device_flow_token, run_device_flow, DeviceFlowConfig, DeviceFlowTokenRefreshError, + DeviceFlowTokens, RequestEncoding, }; use super::openai_compatible::handle_status; use super::retry::ProviderRetry; @@ -60,14 +60,9 @@ const REFRESH_THRESHOLD_SECS: i64 = 300; /// Fallback access-token lifetime when the server omits `expires_in`. const DEFAULT_TOKEN_LIFETIME_SECS: i64 = 3600; -/// Marker key written to the user config when OAuth completes successfully. -/// `check_provider_configured` (server) keys off this when an OAuth-flow -/// provider has no required secret env var. -pub(crate) const KIMI_CONFIGURED_MARKER: &str = "kimi_code_configured"; - // ── Token persistence ──────────────────────────────────────────────────────── -#[derive(Debug, Serialize, Deserialize, Clone)] +#[derive(Debug, Serialize, Deserialize, Clone, PartialEq)] struct KimiToken { access_token: String, refresh_token: String, @@ -97,6 +92,13 @@ struct TokenCache { path: std::path::PathBuf, } +pub(crate) fn has_configured_token() -> bool { + std::fs::read_to_string(TokenCache::new().path) + .ok() + .and_then(|raw| serde_json::from_str::(&raw).ok()) + .is_some() +} + impl TokenCache { fn new() -> Self { Self { @@ -226,57 +228,59 @@ impl KimiCodeProvider { // ── Token management ───────────────────────────────────────────────────── - /// Returns a valid access token, refreshing or re-authenticating as needed. - async fn get_access_token(&self) -> Result { + async fn get_access_token(&self) -> Result { Ok(self.ensure_token().await?.access_token) } - /// Ensures we have a usable token, walking the cache → refresh → device-flow ladder. - async fn ensure_token(&self) -> Result { + async fn ensure_token(&self) -> Result { let mut guard = self.cached_token.lock().await; if let Some(token) = guard.clone() { - if let Some(usable) = self.use_or_refresh(token).await { - *guard = Some(usable.clone()); - return Ok(usable); - } + let usable = self.use_or_refresh(token).await?; + *guard = Some(usable.clone()); + return Ok(usable); } if let Some(token) = self.token_cache.load().await { - if let Some(usable) = self.use_or_refresh(token).await { - *guard = Some(usable.clone()); - return Ok(usable); - } + let usable = self.use_or_refresh(token).await?; + *guard = Some(usable.clone()); + return Ok(usable); } - tracing::info!("kimicode: starting OAuth device-flow login"); - let token = self.device_flow_login().await?; - self.token_cache.save(&token).await?; - *guard = Some(token.clone()); - Ok(token) + Err(ProviderError::NotConfigured) } - /// Returns a usable token derived from `token`, or `None` if it is unusable. - /// On a successful refresh, the new token is also persisted to disk. - async fn use_or_refresh(&self, token: KimiToken) -> Option { - if token.expires_at - Utc::now() > Duration::seconds(REFRESH_THRESHOLD_SECS) { - return Some(token); - } - match self.do_refresh_token(&token.refresh_token).await { - Ok(refreshed) => { - tracing::debug!("kimicode: token refreshed"); - if let Err(e) = self.token_cache.save(&refreshed).await { - tracing::warn!("failed to persist refreshed kimicode token: {}", e); - } - Some(refreshed) + async fn use_or_refresh(&self, mut token: KimiToken) -> Result { + let mut reloaded = false; + + loop { + if token.expires_at - Utc::now() > Duration::seconds(REFRESH_THRESHOLD_SECS) { + return Ok(token); } - Err(e) => { - tracing::debug!("kimicode: token refresh failed: {}", e); - if token.expires_at > Utc::now() { - tracing::debug!("kimicode: falling back to still-unexpired token"); - Some(token) - } else { - None + match self.do_refresh_token(&token.refresh_token).await { + Ok(refreshed) => { + tracing::debug!("kimicode: token refreshed"); + if let Err(e) = self.token_cache.save(&refreshed).await { + tracing::warn!("failed to persist refreshed kimicode token: {}", e); + } + return Ok(refreshed); + } + Err(error) => { + tracing::debug!("kimicode: token refresh failed: {}", error); + if !reloaded { + reloaded = true; + if let Some(persisted) = self.token_cache.load().await { + if persisted != token { + token = persisted; + continue; + } + } + } + if token.expires_at > Utc::now() { + tracing::debug!("kimicode: falling back to still-unexpired token"); + return Ok(token); + } + return Err(kimi_refresh_error(error)); } } } @@ -316,9 +320,7 @@ impl KimiCodeProvider { // ── HTTP ───────────────────────────────────────────────────────────────── async fn post(&self, payload: &Value) -> Result { - let access_token = self.get_access_token().await.map_err(|e| { - ProviderError::Authentication(format!("Failed to get Kimi access token: {}", e)) - })?; + let access_token = self.get_access_token().await?; let builder = self .client @@ -337,6 +339,33 @@ impl KimiCodeProvider { } } +fn kimi_refresh_error(error: anyhow::Error) -> ProviderError { + let refresh_error = error + .chain() + .find_map(|cause| cause.downcast_ref::()); + let status = refresh_error.map(|error| error.status).or_else(|| { + error + .chain() + .find_map(|cause| cause.downcast_ref::()) + .and_then(reqwest::Error::status) + }); + let details = error.to_string(); + + if refresh_error.and_then(|error| error.error.as_deref()) == Some("invalid_grant") { + return ProviderError::Authentication(details); + } + + match status { + Some(reqwest::StatusCode::TOO_MANY_REQUESTS) => ProviderError::RateLimitExceeded { + details, + retry_delay: None, + }, + Some(status) if status.is_server_error() => ProviderError::ServerError(details), + Some(_) => ProviderError::RequestFailed(details), + _ => ProviderError::from(error), + } +} + // ── ProviderDef ─────────────────────────────────────────────────────────────── impl goose_providers::base::ProviderDescriptor for KimiCodeProvider { @@ -348,9 +377,6 @@ impl goose_providers::base::ProviderDescriptor for KimiCodeProvider { KIMI_CODE_DEFAULT_MODEL, KIMI_CODE_KNOWN_MODELS.to_vec(), KIMI_CODE_DOC_URL, - // Marker key — the actual token lives in ~/.config/goose/kimicode/token.json. - // `oauth_flow=true` routes config through `configure_oauth`; - // readiness is tracked via the `kimi_code_configured` param. vec![ConfigKey::new_oauth_device_code( "KIMI_CODE_TOKEN", true, @@ -450,9 +476,7 @@ impl Provider for KimiCodeProvider { data: Vec, } - let access_token = self.get_access_token().await.map_err(|e| { - ProviderError::Authentication(format!("Failed to get Kimi access token: {}", e)) - })?; + let access_token = self.get_access_token().await?; let resp = self .client @@ -474,18 +498,19 @@ impl Provider for KimiCodeProvider { } async fn configure_oauth(&self) -> Result<(), ProviderError> { - self.ensure_token() - .await - .map_err(|e| ProviderError::Authentication(format!("OAuth flow failed: {}", e)))?; - - Config::global() - .set_param(KIMI_CONFIGURED_MARKER, Value::Bool(true)) - .map_err(|e| { - ProviderError::ExecutionError(format!( - "Failed to record kimi_code configured state: {}", - e - )) - })?; + match self.ensure_token().await { + Ok(_) => {} + Err(ProviderError::NotConfigured | ProviderError::Authentication(_)) => { + let token = self.device_flow_login().await.map_err(|e| { + ProviderError::Authentication(format!("OAuth flow failed: {}", e)) + })?; + self.token_cache.save(&token).await.map_err(|e| { + ProviderError::Authentication(format!("Failed to save OAuth token: {}", e)) + })?; + *self.cached_token.lock().await = Some(token); + } + Err(error) => return Err(error), + } Ok(()) } @@ -636,6 +661,97 @@ mod tests { assert_eq!(usable.access_token, "still_good"); } + #[tokio::test] + async fn use_or_refresh_preserves_transient_error_for_expired_token() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/api/oauth/token")) + .respond_with(ResponseTemplate::new(503)) + .mount(&server) + .await; + + let provider = test_provider(&server.uri(), "abc"); + let expired = KimiToken { + access_token: "expired".to_string(), + refresh_token: "ref".to_string(), + expires_at: Utc::now() - Duration::seconds(1), + }; + + let error = provider.use_or_refresh(expired).await.unwrap_err(); + + assert!( + matches!(error, ProviderError::ServerError(_)), + "expected ServerError, got {error:?}" + ); + } + + #[tokio::test] + async fn use_or_refresh_only_authenticates_for_invalid_grant() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/api/oauth/token")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "invalid_grant", + }))) + .mount(&server) + .await; + + let provider = test_provider(&server.uri(), "abc"); + let expired = KimiToken { + access_token: "expired".to_string(), + refresh_token: "rejected".to_string(), + expires_at: Utc::now() - Duration::seconds(1), + }; + + let error = provider.use_or_refresh(expired).await.unwrap_err(); + + assert!(matches!(error, ProviderError::Authentication(_))); + } + + #[tokio::test] + async fn use_or_refresh_does_not_authenticate_for_invalid_client() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/api/oauth/token")) + .respond_with(ResponseTemplate::new(401).set_body_json(json!({ + "error": "invalid_client", + }))) + .mount(&server) + .await; + + let provider = test_provider(&server.uri(), "abc"); + let expired = KimiToken { + access_token: "expired".to_string(), + refresh_token: "still-valid".to_string(), + expires_at: Utc::now() - Duration::seconds(1), + }; + + let error = provider.use_or_refresh(expired).await.unwrap_err(); + + assert!(matches!(error, ProviderError::RequestFailed(_))); + } + + #[tokio::test] + async fn use_or_refresh_does_not_authenticate_for_proxy_rejection() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/api/oauth/token")) + .respond_with(ResponseTemplate::new(403).set_body_string("request rejected by proxy")) + .mount(&server) + .await; + + let provider = test_provider(&server.uri(), "abc"); + let expired = KimiToken { + access_token: "expired".to_string(), + refresh_token: "still-valid".to_string(), + expires_at: Utc::now() - Duration::seconds(1), + }; + + let error = provider.use_or_refresh(expired).await.unwrap_err(); + + assert!(matches!(error, ProviderError::RequestFailed(_))); + } + #[tokio::test] async fn use_or_refresh_returns_new_token_on_successful_refresh() { let server = MockServer::start().await; @@ -662,6 +778,45 @@ mod tests { assert_eq!(usable.refresh_token, "new_refresh"); } + #[tokio::test] + async fn use_or_refresh_reloads_token_rotated_by_another_provider() { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/api/oauth/token")) + .and(body_string_contains("refresh_token=old_refresh")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "access_token": "new_access", + "refresh_token": "new_refresh", + "expires_in": 3600, + }))) + .up_to_n_times(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/api/oauth/token")) + .and(body_string_contains("refresh_token=old_refresh")) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({ + "error": "invalid_grant", + }))) + .mount(&server) + .await; + + let first = test_provider(&server.uri(), "first"); + let mut second = test_provider(&server.uri(), "second"); + second.token_cache.path = first.token_cache.path.clone(); + let expired = KimiToken { + access_token: "old_access".to_string(), + refresh_token: "old_refresh".to_string(), + expires_at: Utc::now() - Duration::seconds(1), + }; + + let refreshed = first.use_or_refresh(expired.clone()).await.unwrap(); + let reloaded = second.use_or_refresh(expired).await.unwrap(); + + assert_eq!(refreshed.access_token, "new_access"); + assert_eq!(reloaded, refreshed); + } + // NOTE: RFC 8628 polling behavior (authorization_pending, slow_down, missing // refresh_token, HTTP errors during polling) is covered by // `providers::oauth_device_flow` tests. Tests here focus on Kimi-specific @@ -749,4 +904,15 @@ mod tests { err ); } + + #[tokio::test] + async fn fetch_supported_models_does_not_authenticate_when_unconfigured() { + let server = MockServer::start().await; + let provider = test_provider(&server.uri(), "abc"); + + let err = provider.fetch_supported_models().await.unwrap_err(); + + assert_eq!(err, ProviderError::NotConfigured); + assert!(server.received_requests().await.unwrap().is_empty()); + } } diff --git a/crates/goose/src/providers/oauth_device_flow.rs b/crates/goose/src/providers/oauth_device_flow.rs index 6baff9a54108..30c49a50953e 100644 --- a/crates/goose/src/providers/oauth_device_flow.rs +++ b/crates/goose/src/providers/oauth_device_flow.rs @@ -80,6 +80,14 @@ pub struct DeviceFlowTokens { pub expires_at: Option>, } +#[derive(Debug, thiserror::Error)] +#[error("token refresh failed ({status}): {body}")] +pub struct DeviceFlowTokenRefreshError { + pub status: reqwest::StatusCode, + pub error: Option, + body: String, +} + // ── Public entry points ────────────────────────────────────────────────────── /// Request a device code from the authorization server. @@ -202,14 +210,25 @@ pub async fn refresh_device_flow_token( refresh_token, }; - let raw: TokenResponseBody = send_request(client, cfg, cfg.token_url, &req) + let response = send_request(client, cfg, cfg.token_url, &req) .await - .context("failed to refresh token")? - .error_for_status() - .context("token refresh failed")? - .json() + .context("failed to refresh token")?; + let status = response.status(); + let bytes = response + .bytes() .await - .context("failed to parse token refresh response")?; + .context("failed to read token refresh response")?; + let raw = serde_json::from_slice::(&bytes); + + if !status.is_success() { + return Err(anyhow::Error::new(DeviceFlowTokenRefreshError { + status, + error: raw.ok().and_then(|body| body.error), + body: String::from_utf8_lossy(&bytes).into_owned(), + })); + } + + let raw = raw.context("failed to parse token refresh response")?; let access_token = raw .access_token diff --git a/crates/goose/src/providers/xai_oauth.rs b/crates/goose/src/providers/xai_oauth.rs index 5e67c02eb093..4155cf63b591 100644 --- a/crates/goose/src/providers/xai_oauth.rs +++ b/crates/goose/src/providers/xai_oauth.rs @@ -221,7 +221,7 @@ async fn exchange_code_for_tokens(code: &str, pkce: &PkceChallenge) -> Result Result { +async fn refresh_access_token(refresh_token: &str) -> Result { let client = reqwest::Client::new(); let params = [ ("grant_type", "refresh_token"), @@ -235,15 +235,38 @@ async fn refresh_access_token(refresh_token: &str) -> Result { .header("Accept", "application/json") .form(¶ms) .send() - .await?; + .await + .map_err(ProviderError::from)?; if !resp.status().is_success() { let status = resp.status(); let text = resp.text().await.unwrap_or_default(); - return Err(anyhow!("xAI token refresh failed ({}): {}", status, text)); + return Err(token_refresh_error(status, text)); } - Ok(resp.json().await?) + resp.json() + .await + .map_err(|error| ProviderError::RequestFailed(error.to_string())) +} + +fn token_refresh_error(status: reqwest::StatusCode, body: String) -> ProviderError { + let details = format!("xAI token refresh failed ({status}): {body}"); + let oauth_error = serde_json::from_str::(&body) + .ok() + .and_then(|value| value.get("error")?.as_str().map(str::to_owned)); + + if oauth_error.as_deref() == Some("invalid_grant") { + return ProviderError::Authentication(details); + } + + match status { + reqwest::StatusCode::TOO_MANY_REQUESTS => ProviderError::RateLimitExceeded { + details, + retry_delay: None, + }, + _ if status.is_server_error() => ProviderError::ServerError(details), + _ => ProviderError::RequestFailed(details), + } } #[derive(Debug, Deserialize)] @@ -606,7 +629,7 @@ impl XaiOAuthAuthProvider { } } - async fn get_valid_token(&self) -> Result { + async fn get_valid_token(&self) -> Result { if let Some(mut token_data) = self.cache.load() { if token_data.expires_at > Utc::now() + chrono::Duration::seconds(ACCESS_TOKEN_REFRESH_SKEW_SECS) @@ -638,30 +661,25 @@ impl XaiOAuthAuthProvider { } token_data.expires_at = Utc::now() + chrono::Duration::seconds(new_tokens.expires_in.unwrap_or(3600)); - self.cache.save(&token_data)?; + self.cache.save(&token_data).map_err(ProviderError::from)?; tracing::info!("xAI access token refreshed"); return Ok(token_data); } - Err(e) => { - tracing::warn!("xAI token refresh failed, will re-authenticate: {}", e); + Err(error @ ProviderError::Authentication(_)) => { + tracing::warn!("xAI token refresh rejected: {}", error); self.cache.clear(); + return Err(error); + } + Err(error) => { + if token_data.expires_at > Utc::now() { + return Ok(token_data); + } + return Err(error); } } } - tracing::info!("Starting xAI OAuth flow (SuperGrok subscription)"); - let token_data = match perform_loopback_oauth_flow(self.state.as_ref()).await { - Ok(td) => td, - Err(e) => { - tracing::warn!( - "xAI loopback OAuth failed ({}); falling back to device-code flow", - e - ); - perform_device_code_flow().await? - } - }; - self.cache.save(&token_data)?; - Ok(token_data) + Err(ProviderError::NotConfigured) } } @@ -708,12 +726,14 @@ impl Provider for XaiOAuthProvider { messages: &[Message], tools: &[Tool], ) -> Result { + self.auth_provider.get_valid_token().await?; self.inner .stream(model_config, system, messages, tools) .await } async fn fetch_supported_models(&self) -> Result, ProviderError> { + self.auth_provider.get_valid_token().await?; self.inner.fetch_supported_models().await } @@ -872,6 +892,84 @@ mod tests { assert!(s.ends_with("tokens.json")); } + #[tokio::test] + async fn missing_token_does_not_start_oauth() { + let directory = tempfile::tempdir().unwrap(); + let auth_provider = XaiOAuthAuthProvider { + cache: TokenCache { + cache_path: directory.path().join("missing.json"), + }, + state: XaiAuthState::instance(), + }; + + let error = auth_provider.get_valid_token().await.unwrap_err(); + + assert_eq!(error, ProviderError::NotConfigured); + } + + #[tokio::test] + async fn stream_preserves_not_configured_error() { + let directory = tempfile::tempdir().unwrap(); + let auth_provider = Arc::new(XaiOAuthAuthProvider { + cache: TokenCache { + cache_path: directory.path().join("missing.json"), + }, + state: XaiAuthState::instance(), + }); + let api_client = + ApiClient::new_with_tls("http://127.0.0.1:1".to_string(), AuthMethod::NoAuth, None) + .unwrap(); + let provider = XaiOAuthProvider { + inner: OpenAiCompatibleProvider::new( + XAI_OAUTH_PROVIDER_NAME.to_string(), + api_client, + String::new(), + ), + auth_provider, + }; + + let error = provider + .stream(&ModelConfig::new(XAI_DEFAULT_MODEL), "", &[], &[]) + .await + .err() + .unwrap(); + + assert_eq!(error, ProviderError::NotConfigured); + } + + #[test] + fn token_refresh_errors_distinguish_rejected_and_transient_requests() { + assert!(matches!( + token_refresh_error( + reqwest::StatusCode::BAD_REQUEST, + r#"{"error":"invalid_grant"}"#.to_string() + ), + ProviderError::Authentication(_) + )); + assert!(matches!( + token_refresh_error( + reqwest::StatusCode::UNAUTHORIZED, + r#"{"error":"invalid_client"}"#.to_string() + ), + ProviderError::RequestFailed(_) + )); + assert!(matches!( + token_refresh_error( + reqwest::StatusCode::FORBIDDEN, + "request rejected by proxy".to_string() + ), + ProviderError::RequestFailed(_) + )); + assert!(matches!( + token_refresh_error(reqwest::StatusCode::TOO_MANY_REQUESTS, String::new()), + ProviderError::RateLimitExceeded { .. } + )); + assert!(matches!( + token_refresh_error(reqwest::StatusCode::SERVICE_UNAVAILABLE, String::new()), + ProviderError::ServerError(_) + )); + } + #[cfg(unix)] #[test] fn token_cache_replaces_loose_file_with_owner_only_permissions() { diff --git a/crates/goose/tests/acp_custom_requests_test.rs b/crates/goose/tests/acp_custom_requests_test.rs index 5f1c126ab637..661624d51f2a 100644 --- a/crates/goose/tests/acp_custom_requests_test.rs +++ b/crates/goose/tests/acp_custom_requests_test.rs @@ -48,7 +48,7 @@ fn write_acp_global_config(contents: &str) -> PathBuf { struct MockProvider { name: String, recommended_models: Vec, - supported_models: Vec, + supported_models: Result, ProviderError>, } #[async_trait::async_trait] @@ -75,7 +75,7 @@ impl Provider for MockProvider { } async fn fetch_supported_models(&self) -> Result, ProviderError> { - Ok(self.supported_models.clone()) + self.supported_models.clone() } } @@ -1179,10 +1179,10 @@ fn test_custom_provider_supported_models_lists_raw_provider_models() { Ok(Arc::new(MockProvider { name: provider_name, recommended_models: vec!["canonical-filtered-model".to_string()], - supported_models: vec![ + supported_models: Ok(vec![ "goose-claude-opus-4-8".to_string(), "raw-databricks-endpoint".to_string(), - ], + ]), }) as Arc) }) }); @@ -1216,3 +1216,79 @@ fn test_custom_provider_supported_models_lists_raw_provider_models() { ); }); } + +#[test] +#[serial] +fn test_custom_provider_supported_models_maps_not_configured_error() { + write_acp_global_config(DEFAULT_ACP_TEST_CONFIG); + run_test(async move { + let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; + let provider_factory: AcpProviderFactory = Arc::new(|provider_name, _, _| { + Box::pin(async move { + Ok(Arc::new(MockProvider { + name: provider_name, + recommended_models: Vec::new(), + supported_models: Err(ProviderError::NotConfigured), + }) as Arc) + }) + }); + let conn = AcpServerConnection::new( + TestConnectionConfig { + provider_factory: Some(provider_factory), + ..Default::default() + }, + openai, + ) + .await; + + let error = send_custom( + conn.cx(), + "_goose/unstable/providers/supported-models/list", + serde_json::json!({ "providerId": "openai" }), + ) + .await + .expect_err("not configured should be returned to the client"); + + assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); + assert!(error.to_string().contains("Provider is not configured")); + }); +} + +#[test] +#[serial] +fn test_custom_provider_supported_models_maps_authentication_error() { + write_acp_global_config(DEFAULT_ACP_TEST_CONFIG); + run_test(async move { + let openai = OpenAiFixture::new(vec![], Arc::new(EnforceSessionId::default())).await; + let provider_factory: AcpProviderFactory = Arc::new(|provider_name, _, _| { + Box::pin(async move { + Ok(Arc::new(MockProvider { + name: provider_name, + recommended_models: Vec::new(), + supported_models: Err(ProviderError::Authentication( + "credentials rejected".to_string(), + )), + }) as Arc) + }) + }); + let conn = AcpServerConnection::new( + TestConnectionConfig { + provider_factory: Some(provider_factory), + ..Default::default() + }, + openai, + ) + .await; + + let error = send_custom( + conn.cx(), + "_goose/unstable/providers/supported-models/list", + serde_json::json!({ "providerId": "openai" }), + ) + .await + .expect_err("authentication failure should be returned to the client"); + + assert_eq!(error.code, agent_client_protocol::ErrorCode::AuthRequired); + assert!(error.to_string().contains("credentials rejected")); + }); +}