From 31df1a720dcd88fb4786a56526fbc4547ce9108f Mon Sep 17 00:00:00 2001 From: Alex Clarke Date: Mon, 20 Jul 2026 13:12:38 -0600 Subject: [PATCH] test: unit tests for OAuthConfig merge + get_oauth_provider_for_client --- src/client/oauth.rs | 245 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 245 insertions(+) diff --git a/src/client/oauth.rs b/src/client/oauth.rs index a9e1fe1..90b8873 100644 --- a/src/client/oauth.rs +++ b/src/client/oauth.rs @@ -651,3 +651,248 @@ pub(crate) fn client_config_info( ClientConfig::Unknown => ("unknown", "unknown", None), } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::client::{ModelData, ProviderModels}; + use crate::client::openai_compatible::OpenAICompatibleConfig; + + fn base_config() -> OAuthConfig { + OAuthConfig { + client_id: "base-id".into(), + token_url: "https://base.example/token".into(), + flow: OAuthFlow::Pkce, + client_secret: Some("base-secret".into()), + authorize_url: Some("https://base.example/authorize".into()), + redirect_uri: None, + redirect_port: Some(1234), + scopes: vec!["a".into(), "b".into()], + token_request_format: Some(TokenRequestFormat::FormUrlEncoded), + extra_authorize_params: IndexMap::from([("plan".into(), "base".into())]), + extra_token_headers: IndexMap::new(), + extra_request_headers: IndexMap::new(), + echo_pkce_in_token_exchange: false, + include_state_in_token_exchange: true, + } + } + + fn empty_user_override(client_id: &str, token_url: &str) -> OAuthConfig { + OAuthConfig { + client_id: client_id.into(), + token_url: token_url.into(), + flow: OAuthFlow::Pkce, + client_secret: None, + authorize_url: None, + redirect_uri: None, + redirect_port: None, + scopes: vec![], + token_request_format: None, + extra_authorize_params: IndexMap::new(), + extra_token_headers: IndexMap::new(), + extra_request_headers: IndexMap::new(), + echo_pkce_in_token_exchange: false, + include_state_in_token_exchange: true, + } + } + + #[test] + fn oauth_config_merge_user_wins_per_field() { + let base = base_config(); + let mut user = empty_user_override("user-id", "https://user.example/token"); + user.client_secret = Some("user-secret".into()); + user.scopes = vec!["c".into()]; + user.extra_authorize_params = IndexMap::from([("plan".into(), "user".into())]); + + let merged = base.merge(user); + + assert_eq!(merged.client_id, "user-id"); + assert_eq!(merged.token_url, "https://user.example/token"); + assert_eq!(merged.client_secret.as_deref(), Some("user-secret")); + assert_eq!(merged.authorize_url.as_deref(), Some("https://base.example/authorize")); + assert_eq!(merged.redirect_port, Some(1234)); + assert_eq!(merged.scopes, vec!["c"]); + assert_eq!(merged.extra_authorize_params.get("plan").map(String::as_str), Some("user")); + } + + #[test] + fn oauth_config_merge_empty_user_keeps_base_optionals() { + let base = base_config(); + let user = empty_user_override("user-id", "https://user.example/token"); + + let merged = base.merge(user); + + assert_eq!(merged.client_id, "user-id"); + assert_eq!(merged.token_url, "https://user.example/token"); + assert_eq!(merged.client_secret.as_deref(), Some("base-secret")); + assert_eq!(merged.authorize_url.as_deref(), Some("https://base.example/authorize")); + assert_eq!(merged.redirect_port, Some(1234)); + assert_eq!(merged.scopes, vec!["a", "b"]); + assert!(matches!(merged.token_request_format, Some(TokenRequestFormat::FormUrlEncoded))); + assert_eq!(merged.extra_authorize_params.get("plan").map(String::as_str), Some("base")); + } + + #[test] + fn oauth_config_serde_roundtrip_from_yaml() { + let yaml = r#" +client_id: xai-client +token_url: https://auth.x.ai/oauth2/token +authorize_url: https://auth.x.ai/oauth2/authorize +scopes: + - openid + - profile + - api:access +redirect_port: 56121 +flow: pkce +token_request_format: form_url_encoded +extra_authorize_params: + plan: generic + referrer: coyote +echo_pkce_in_token_exchange: true +"#; + let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); + assert_eq!(cfg.client_id, "xai-client"); + assert_eq!(cfg.token_url, "https://auth.x.ai/oauth2/token"); + assert_eq!(cfg.scopes.len(), 3); + assert_eq!(cfg.redirect_port, Some(56121)); + assert!(matches!(cfg.flow, OAuthFlow::Pkce)); + assert!(matches!(cfg.token_request_format, Some(TokenRequestFormat::FormUrlEncoded))); + assert!(cfg.echo_pkce_in_token_exchange); + assert!(cfg.include_state_in_token_exchange); + assert_eq!(cfg.extra_authorize_params.get("plan").map(String::as_str), Some("generic")); + } + + #[test] + fn oauth_flow_defaults_to_pkce_when_missing() { + let yaml = "client_id: x\ntoken_url: y"; + let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); + assert!(matches!(cfg.flow, OAuthFlow::Pkce)); + } + + #[test] + fn oauth_flow_client_credentials_parses() { + let yaml = "client_id: x\ntoken_url: y\nflow: client_credentials"; + let cfg: OAuthConfig = serde_yaml::from_str(yaml).unwrap(); + assert!(matches!(cfg.flow, OAuthFlow::ClientCredentials)); + } + + fn make_provider_models(provider: &str, oauth: Option) -> ProviderModels { + ProviderModels { + provider: provider.into(), + oauth, + models: vec![ModelData::new("some-model")], + } + } + + fn make_openai_compat_client( + name: &str, + auth: Option<&str>, + oauth: Option, + ) -> ClientConfig { + ClientConfig::OpenAICompatibleConfig(OpenAICompatibleConfig { + name: Some(name.into()), + api_base: Some("https://api.example/v1".into()), + api_key: None, + auth: auth.map(str::to_string), + oauth: oauth.map(Box::new), + models: vec![], + patch: None, + extra: None, + }) + } + + #[test] + fn get_oauth_provider_for_client_merges_defaults_with_user_override() { + let base = base_config(); + let mut user = empty_user_override("user-id", "https://user.example/token"); + user.echo_pkce_in_token_exchange = true; + let models = vec![make_provider_models("acme", Some(base))]; + let cc = make_openai_compat_client("acme", Some("oauth"), Some(user)); + + let provider = get_oauth_provider_for_client(&cc, &models).unwrap(); + assert_eq!(provider.client_id(), "user-id"); + assert_eq!(provider.token_url(), "https://user.example/token"); + assert!(provider.echo_pkce_in_token_exchange()); + assert_eq!( + provider.fixed_redirect_uri().as_deref(), + Some("http://127.0.0.1:1234/callback") + ); + } + + #[test] + fn get_oauth_provider_for_client_uses_bundled_defaults_only() { + let base = base_config(); + let models = vec![make_provider_models("bundled-only", Some(base))]; + let cc = make_openai_compat_client("bundled-only", Some("oauth"), None); + + let provider = get_oauth_provider_for_client(&cc, &models).unwrap(); + assert_eq!(provider.client_id(), "base-id"); + assert_eq!(provider.token_url(), "https://base.example/token"); + } + + #[test] + fn get_oauth_provider_for_client_uses_inline_only() { + let user = OAuthConfig { + client_id: "inline-id".into(), + token_url: "https://inline.example/token".into(), + ..empty_user_override("inline-id", "https://inline.example/token") + }; + let cc = make_openai_compat_client("inline-only", Some("oauth"), Some(user)); + + let provider = get_oauth_provider_for_client(&cc, &[]).unwrap(); + assert_eq!(provider.client_id(), "inline-id"); + } + + #[test] + fn get_oauth_provider_for_client_returns_none_when_no_config_anywhere() { + let cc = make_openai_compat_client("nothing", Some("oauth"), None); + assert!(get_oauth_provider_for_client(&cc, &[]).is_none()); + } + + #[test] + fn get_oauth_provider_for_client_returns_none_when_auth_not_oauth() { + let base = base_config(); + let models = vec![make_provider_models("api-key-client", Some(base))]; + let cc = make_openai_compat_client("api-key-client", None, None); + assert!(get_oauth_provider_for_client(&cc, &models).is_none()); + } + + #[test] + fn openai_compatible_provider_joins_scopes_with_spaces() { + let mut cfg = base_config(); + cfg.scopes = vec!["one".into(), "two".into(), "three".into()]; + let provider = super::super::openai_compatible_oauth::OpenAICompatibleOAuthProvider { + config: cfg, + client_name: "test".into(), + }; + assert_eq!(provider.scopes(), "one two three"); + } + + #[test] + fn openai_compatible_provider_prefers_redirect_uri_over_port() { + let mut cfg = base_config(); + cfg.redirect_uri = Some("https://custom.example/cb".into()); + cfg.redirect_port = Some(9999); + let provider = super::super::openai_compatible_oauth::OpenAICompatibleOAuthProvider { + config: cfg, + client_name: "test".into(), + }; + assert_eq!( + provider.fixed_redirect_uri().as_deref(), + Some("https://custom.example/cb") + ); + } + + #[test] + fn openai_compatible_provider_ephemeral_when_no_redirect() { + let mut cfg = base_config(); + cfg.redirect_uri = None; + cfg.redirect_port = None; + let provider = super::super::openai_compatible_oauth::OpenAICompatibleOAuthProvider { + config: cfg, + client_name: "test".into(), + }; + assert!(provider.uses_localhost_redirect()); + assert!(provider.fixed_redirect_uri().is_none()); + } +}