diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 98dd34b..8956669 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -78,7 +78,7 @@ src/ ├── cli.rs clap argument structs + Stage + slug validation ├── config.rs Config / Profile / AppContext + on-disk persistence ├── api.rs stage → API URL resolution -├── auth_cache.rs AuthCache (Bearer | Basic) + 0600 on-disk store +├── auth_cache.rs AuthCache (ClientCredentials | Bearer | Basic) + 0600 on-disk store ├── http.rs ApiClient: auth-attached reqwest wrapper ├── jsonapi.rs generic Document / Resource / Single / List envelopes ├── prompt.rs interactive yes_no / text / stage / organization prompts diff --git a/README.md b/README.md index b96ec1d..4fa5021 100644 --- a/README.md +++ b/README.md @@ -46,7 +46,7 @@ rw config profile rm mercy # Remove the "mercy" profile (prompts fo rw config profile rm mercy --yes # Remove without prompting rw config profile set mercy -o new-org # Update organization for a profile rw config profile set mercy -g sandbox # Update stage for a profile -rw config profile auth mercy # Save basic auth credentials for a profile (see below) +rw config profile auth mercy # Save credentials for a profile (see below) ``` #### Overriding the stage @@ -131,6 +131,26 @@ rw config profile auth mercy --username alice \ | `--username` | `-u` | Username for basic auth | | `--password` | `-P` | Password for basic auth | +#### Using Client Credentials + +For non-interactive use (CI, scripts), store a client ID and secret instead of logging in: + +```sh +rw config profile auth mercy --client-id client_123 # Secret prompted securely +rw config profile auth mercy --client-id client_123 \ + --client-secret secret # Fully non-interactive +``` + +| Flag | Description | +|-------------------|---------------------------------------------------------------| +| `--client-id` | Client ID (cannot be combined with `--username`/`--password`) | +| `--client-secret` | Client secret (prompted securely if not provided) | + +`rw` exchanges the credentials for an access token on first use and caches it in +`~/.config/rw/auth/{profile}.json`, renewing it automatically. `rw auth login` on such a +profile forces a fresh exchange instead of opening a browser; `rw auth logout` only drops the +cached access token and keeps the client credentials (use `rw config profile rm` to remove them). + #### Diagnostics `rw config doctor` runs a fixed sequence of checks against the active profile (or `--profile`): @@ -154,10 +174,10 @@ Each check reports `pass`, `warn`, `fail`, `skip`, or `info`. A later check is s ### Authentication ```sh -rw auth login # Open browser and authenticate via WorkOS +rw auth login # Open browser and authenticate via WorkOS (profiles with client credentials exchange them instead) rw auth status # Show authentication status for current profile rw auth header # Show the authentication header for current profile -rw auth logout # Remove stored credentials for current profile +rw auth logout # Remove stored credentials for current profile (client credentials: cached token only) # Use a named profile rw auth login --profile mercy diff --git a/docs/config.md b/docs/config.md index 156994a..a99d15b 100644 --- a/docs/config.md +++ b/docs/config.md @@ -56,3 +56,19 @@ Basic credentials (written using `rw config profile auth `): "password": "" } ``` + +Client credentials (written using `rw config profile auth --client-id … --client-secret …`): + +```json +{ + "client_id": "", + "client_secret": "", + "access_token": "", + "expires_at": 1234567890 +} +``` + +`access_token` and `expires_at` are absent until the first operation that needs a token (an +API call, `rw auth login` or `rw config doctor`). `rw` then exchanges the credentials for an +access token (OAuth `client_credentials` grant) and writes it back here, re-exchanging when +it is within 60 seconds of expiry. `rw auth logout` removes only the cached token. diff --git a/skills/rw-skill.md b/skills/rw-skill.md index 463aaeb..2752975 100644 --- a/skills/rw-skill.md +++ b/skills/rw-skill.md @@ -43,6 +43,15 @@ rw config profile list # List all configured profiles rw config profile show # Show the active profile ``` +**Non-interactive credentials (client credentials grant):** + +```sh +rw config profile auth --client-id # Secret prompted securely +rw config profile auth --client-id --client-secret # Fully non-interactive +``` + +`--client-id` / `--client-secret` cannot be combined with `--username` / `--password`, and with `--json` both are required. `rw` exchanges them for an access token on the first command that needs one and caches it, renewing automatically. `rw auth login` on such a profile forces a fresh exchange instead of opening a browser; `rw auth logout` drops only the cached token and keeps the credentials. A `--client-secret` passed on the command line is visible in shell history and process listings, so prefer the prompt unless the value comes from a secret store. + ### `rw clinicians` — Clinician Management All targets accept a UUID or email address. Roles accept a UUID or name. Teams accept a UUID or abbreviation. diff --git a/src/auth_cache.rs b/src/auth_cache.rs index 447d1ce..1d874a1 100644 --- a/src/auth_cache.rs +++ b/src/auth_cache.rs @@ -6,6 +6,17 @@ use std::path::{Path, PathBuf}; #[derive(Debug, Serialize, Deserialize, Clone)] #[serde(untagged)] pub enum AuthCache { + /// Declared before `Bearer`: a cached token gives it the same `access_token` + + /// `expires_at` fields, and untagged enums take the first variant that matches. + ClientCredentials { + client_id: String, + client_secret: String, + /// Access token from the last exchange, if any. + #[serde(default, skip_serializing_if = "Option::is_none")] + access_token: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + expires_at: Option, + }, Bearer { access_token: String, #[serde(skip_serializing_if = "Option::is_none")] @@ -20,10 +31,17 @@ pub enum AuthCache { } impl AuthCache { - /// Returns true if this is a bearer token that is expired or expires within 60 seconds. + /// Returns true if this is a bearer or client-credentials token that is missing, + /// expired, or expires within 60 seconds. pub fn is_expired(&self) -> bool { match self { - AuthCache::Bearer { expires_at, .. } => unix_now() >= expires_at - 60, + AuthCache::ClientCredentials { + access_token: Some(_), + expires_at: Some(expires_at), + .. + } + | AuthCache::Bearer { expires_at, .. } => unix_now() >= expires_at - 60, + AuthCache::ClientCredentials { .. } => true, AuthCache::Basic { .. } => false, } } @@ -219,6 +237,70 @@ mod tests { } } + fn client_credentials(token: Option<&str>, expires_at: Option) -> AuthCache { + AuthCache::ClientCredentials { + client_id: "id".to_string(), + client_secret: "sec".to_string(), + access_token: token.map(str::to_string), + expires_at, + } + } + + #[test] + fn test_client_credentials_without_token_is_expired() { + assert!(client_credentials(None, None).is_expired()); + } + + #[test] + fn test_client_credentials_fresh_token_not_expired() { + assert!(!client_credentials(Some("t"), Some(unix_now() + 3600)).is_expired()); + } + + #[test] + fn test_client_credentials_token_in_grace_period_is_expired() { + assert!(client_credentials(Some("t"), Some(unix_now() + 30)).is_expired()); + } + + #[test] + fn test_client_credentials_with_cached_token_roundtrips_as_client_credentials() { + // Must not be swallowed by `Bearer`, which also has access_token + expires_at. + let json = serde_json::to_string(&client_credentials(Some("t"), Some(9999999999))).unwrap(); + match serde_json::from_str::(&json).unwrap() { + AuthCache::ClientCredentials { + client_id, + client_secret, + access_token, + expires_at, + } => { + assert_eq!(client_id, "id"); + assert_eq!(client_secret, "sec"); + assert_eq!(access_token.as_deref(), Some("t")); + assert_eq!(expires_at, Some(9999999999)); + } + other => panic!("expected client credentials, got {:?}", other), + } + } + + #[test] + fn test_client_credentials_without_token_omits_token_fields() { + let json = serde_json::to_string(&client_credentials(None, None)).unwrap(); + assert!(!json.contains("access_token")); + assert!(!json.contains("expires_at")); + assert!(matches!( + serde_json::from_str::(&json).unwrap(), + AuthCache::ClientCredentials { .. } + )); + } + + #[test] + fn test_bearer_json_still_parses_as_bearer() { + let json = r#"{"access_token":"a","refresh_token":"r","expires_at":9999999999}"#; + assert!(matches!( + serde_json::from_str::(json).unwrap(), + AuthCache::Bearer { .. } + )); + } + #[test] fn test_expires_at_from_duration() { let before = unix_now(); diff --git a/src/cli.rs b/src/cli.rs index de839e8..3015df1 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -273,7 +273,7 @@ pub enum ConfigProfileCommands { Rm(ConfigProfileRmArgs), /// Add a new profile. Add(ConfigProfileAddArgs), - /// Save basic auth credentials for a profile. + /// Save credentials (basic auth or client credentials) for a profile. Auth(ConfigProfileAuthArgs), } @@ -339,6 +339,12 @@ pub struct ConfigProfileAuthArgs { /// Password (prompted securely if not provided). #[arg(short = 'P', long)] pub password: Option, + /// Client ID for the client credentials grant (selects client credentials over basic auth). + #[arg(long, conflicts_with_all = ["username", "password"])] + pub client_id: Option, + /// Client secret for the client credentials grant (prompted securely if not provided). + #[arg(long, conflicts_with_all = ["username", "password"])] + pub client_secret: Option, } /// Arguments for `config updates`. @@ -513,6 +519,52 @@ pub struct SkillsInstallArgs { mod tests { use super::*; + #[test] + fn test_profile_auth_client_flags_conflict_with_basic_flags() { + use clap::Parser; + assert!(Cli::try_parse_from([ + "rw", + "config", + "profile", + "auth", + "demo", + "--client-id", + "a", + "--username", + "b", + ]) + .is_err()); + assert!(Cli::try_parse_from([ + "rw", + "config", + "profile", + "auth", + "demo", + "--client-secret", + "a", + "--password", + "b", + ]) + .is_err()); + } + + #[test] + fn test_profile_auth_client_flags_parse() { + use clap::Parser; + assert!(Cli::try_parse_from([ + "rw", + "config", + "profile", + "auth", + "demo", + "--client-id", + "a", + "--client-secret", + "b", + ]) + .is_ok()); + } + #[test] fn test_auth_flag_long() { use clap::Parser; diff --git a/src/commands/auth.rs b/src/commands/auth.rs index 0c0d87a..a309af8 100644 --- a/src/commands/auth.rs +++ b/src/commands/auth.rs @@ -1,5 +1,6 @@ use anyhow::{bail, Context, Result}; use serde::Serialize; +use std::path::Path; use std::time::Duration; use tokio::time::sleep; @@ -72,6 +73,8 @@ pub struct StatusOutput { pub expired: bool, #[serde(skip_serializing_if = "Option::is_none")] pub username: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub client_id: Option, #[serde(skip)] pub profile: String, } @@ -101,6 +104,11 @@ impl CommandOutput for StatusOutput { format!("✓ Authenticated using profile '{}' (basic).", self.profile) } } + (true, Some("client_credentials"), _) => format!( + "✓ Authenticated using profile '{}' (client credentials, client: {}).", + self.profile, + self.client_id.as_deref().unwrap_or("unknown") + ), _ => format!("✓ Authenticated using profile '{}'.", self.profile), } } @@ -120,12 +128,24 @@ impl CommandOutput for HeaderOutput { // --- Command implementations --- /// Run `rw auth login` – use the OAuth Device Authorization Flow to authenticate -/// via WorkOS AuthKit, poll for a token, and persist credentials. +/// via WorkOS AuthKit, poll for a token, and persist credentials. Profiles with +/// client credentials exchange them instead. pub async fn login(ctx: &AppContext, out: &Output) -> Result<()> { if out.json { anyhow::bail!("`rw auth login` is interactive and cannot be used with --json"); } + if login_with_client_credentials( + &ctx.config_dir, + &ctx.profile, + workos_config(&ctx.auth_stage).token_url, + out, + ) + .await? + { + return Ok(()); + } + let wos = workos_config(&ctx.stage); let client = reqwest::Client::new(); @@ -240,6 +260,43 @@ pub async fn login(ctx: &AppContext, out: &Output) -> Result<()> { } } +/// If the profile holds client credentials, exchanges them for a fresh access token +/// (even when a cached one is still valid) instead of starting an interactive login, +/// which would replace the stored secret. Returns `false` when the profile has no +/// client credentials, leaving the caller to run the device flow. +async fn login_with_client_credentials( + config_dir: &Path, + profile: &str, + token_url: &str, + out: &Output, +) -> Result { + let Some(AuthCache::ClientCredentials { + client_id, + client_secret, + .. + }) = load_auth_cache(config_dir, profile)? + else { + return Ok(false); + }; + + // Dropping the cached token forces an exchange. + let cache = AuthCache::ClientCredentials { + client_id, + client_secret, + access_token: None, + expires_at: None, + }; + client_credentials_access_token(config_dir, profile, token_url, cache).await?; + + out.print(&MessageOutput { + message: format!( + "✓ Authenticated successfully using client credentials for profile '{}'.", + profile + ), + }); + Ok(true) +} + /// Run `rw auth status` – report whether stored credentials exist. pub fn status(ctx: &AppContext, out: &Output) -> Result<()> { match load_auth_cache(&ctx.config_dir, &ctx.auth_profile)? { @@ -249,6 +306,17 @@ pub fn status(ctx: &AppContext, out: &Output) -> Result<()> { authenticated: true, expired: cache.is_expired(), username: None, + client_id: None, + profile: ctx.auth_profile.clone(), + }); + } + Some(ref cache @ AuthCache::ClientCredentials { ref client_id, .. }) => { + out.print(&StatusOutput { + auth_type: Some("client_credentials".to_string()), + authenticated: true, + expired: cache.is_expired(), + username: None, + client_id: Some(client_id.clone()), profile: ctx.auth_profile.clone(), }); } @@ -258,6 +326,7 @@ pub fn status(ctx: &AppContext, out: &Output) -> Result<()> { authenticated: true, expired: false, username: Some(username.clone()), + client_id: None, profile: ctx.auth_profile.clone(), }); } @@ -267,6 +336,7 @@ pub fn status(ctx: &AppContext, out: &Output) -> Result<()> { authenticated: false, expired: false, username: None, + client_id: None, profile: ctx.auth_profile.clone(), }); } @@ -303,9 +373,33 @@ pub async fn header(ctx: &AppContext, out: &Output) -> Result<()> { Ok(()) } -/// Run `rw auth logout` – remove stored credentials for the profile. +/// Run `rw auth logout` – remove stored credentials for the profile. Client credentials +/// are configuration rather than a session, so only the cached access token is dropped. pub fn logout(ctx: &AppContext, out: &Output) -> Result<()> { - if delete_auth_cache(&ctx.config_dir, &ctx.profile)? { + // An unreadable cache is not an error here: logout must still be able to remove it. + if let Ok(Some(AuthCache::ClientCredentials { + client_id, + client_secret, + .. + })) = load_auth_cache(&ctx.config_dir, &ctx.profile) + { + save_auth_cache( + &ctx.config_dir, + &ctx.profile, + &AuthCache::ClientCredentials { + client_id, + client_secret, + access_token: None, + expires_at: None, + }, + )?; + out.print(&MessageOutput { + message: format!( + "✓ Cached access token for profile '{}' removed. Client credentials kept.", + ctx.profile + ), + }); + } else if delete_auth_cache(&ctx.config_dir, &ctx.profile)? { out.print(&MessageOutput { message: format!("✓ Credentials for profile '{}' removed.", ctx.profile), }); @@ -324,7 +418,8 @@ pub enum ResolvedAuth { } /// Resolves auth credentials for the given organization+stage, loading the cache once. -/// For bearer tokens, automatically refreshes if expired. +/// Bearer tokens are refreshed when expired; client credentials are exchanged for an +/// access token when none is cached or it is expired. /// Returns `None` if no credentials are stored. pub async fn resolve_auth(ctx: &AppContext) -> Result> { let Some(cache) = load_auth_cache(&ctx.config_dir, &ctx.auth_profile)? else { @@ -332,6 +427,16 @@ pub async fn resolve_auth(ctx: &AppContext) -> Result> { }; match cache { + AuthCache::ClientCredentials { .. } => { + let token = client_credentials_access_token( + &ctx.config_dir, + &ctx.auth_profile, + workos_config(&ctx.auth_stage).token_url, + cache, + ) + .await?; + Ok(Some(ResolvedAuth::Bearer(token))) + } AuthCache::Basic { username, password } => { Ok(Some(ResolvedAuth::Basic { username, password })) } @@ -401,6 +506,74 @@ async fn try_refresh(stage: &Stage, refresh_token: &str) -> Result { }) } +/// Requests an access token with the OAuth `client_credentials` grant. +async fn fetch_client_credentials_token( + token_url: &str, + client_id: &str, + client_secret: &str, +) -> Result { + let resp = reqwest::Client::new() + .post(token_url) + .form(&[ + ("grant_type", "client_credentials"), + ("client_id", client_id), + ("client_secret", client_secret), + ]) + .send() + .await + .context("failed to reach WorkOS token endpoint")?; + + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + bail!("token endpoint returned {}: {}", status, body); + } + resp.json() + .await + .context("failed to parse client credentials token response") +} + +/// Returns an access token for a `ClientCredentials` cache. A cached token that is not +/// within the 60s expiry grace period is reused; otherwise a new one is exchanged and +/// written back to the profile's auth file alongside the credentials. +pub async fn client_credentials_access_token( + config_dir: &Path, + profile: &str, + token_url: &str, + cache: AuthCache, +) -> Result { + let expired = cache.is_expired(); + let AuthCache::ClientCredentials { + client_id, + client_secret, + access_token, + .. + } = cache + else { + bail!("not a client credentials cache"); + }; + + if let Some(token) = access_token.filter(|_| !expired) { + return Ok(token); + } + + let token = fetch_client_credentials_token(token_url, &client_id, &client_secret) + .await + .context("client credentials exchange failed; check `rw config profile auth`")?; + + save_auth_cache( + config_dir, + profile, + &AuthCache::ClientCredentials { + client_id, + client_secret, + access_token: Some(token.access_token.clone()), + expires_at: Some(expires_at_from_duration(token.expires_in)), + }, + )?; + Ok(token.access_token) +} + /// Returns the Authorization header value, or fails with a friendly message if /// no credentials are stored. pub async fn require_auth(ctx: &AppContext) -> Result { @@ -474,6 +647,7 @@ mod tests { authenticated: true, expired: false, username: None, + client_id: None, profile: "demo".to_string(), }; let json = serde_json::to_value(&output).unwrap(); @@ -492,6 +666,7 @@ mod tests { authenticated: true, expired: false, username: Some("alice".to_string()), + client_id: None, profile: "demo".to_string(), }; let json = serde_json::to_value(&output).unwrap(); @@ -506,6 +681,7 @@ mod tests { authenticated: false, expired: false, username: None, + client_id: None, profile: "demo".to_string(), }; let json = serde_json::to_value(&output).unwrap(); @@ -520,6 +696,7 @@ mod tests { authenticated: true, expired: false, username: None, + client_id: None, profile: "demo".to_string(), }; assert_eq!( @@ -535,6 +712,7 @@ mod tests { authenticated: true, expired: false, username: Some("alice".to_string()), + client_id: None, profile: "demo".to_string(), }; assert_eq!( @@ -550,6 +728,7 @@ mod tests { authenticated: false, expired: false, username: None, + client_id: None, profile: "demo".to_string(), }; assert!(output.plain().contains("✗ Not authenticated")); @@ -682,6 +861,345 @@ mod tests { status(&ctx, &out).unwrap(); } + fn cc_cache(token: Option<&str>, expires_at: Option) -> AuthCache { + AuthCache::ClientCredentials { + client_id: "id".to_string(), + client_secret: "sec".to_string(), + access_token: token.map(str::to_string), + expires_at, + } + } + + #[tokio::test] + async fn test_client_credentials_fresh_token_skips_exchange() { + let dir = tempfile::TempDir::new().unwrap(); + let mut server = mockito::Server::new_async().await; + let mock = server.mock("POST", "/token").expect(0).create_async().await; + + let cache = cc_cache(Some("cached"), Some(expires_at_from_duration(3600))); + let token = client_credentials_access_token( + dir.path(), + "demo", + &format!("{}/token", server.url()), + cache, + ) + .await + .unwrap(); + + assert_eq!(token, "cached"); + mock.assert_async().await; + } + + #[tokio::test] + async fn test_client_credentials_missing_token_exchanges_and_saves() { + use mockito::Matcher; + let dir = tempfile::TempDir::new().unwrap(); + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/token") + .match_body(Matcher::AllOf(vec![ + Matcher::UrlEncoded("grant_type".into(), "client_credentials".into()), + Matcher::UrlEncoded("client_id".into(), "id".into()), + Matcher::UrlEncoded("client_secret".into(), "sec".into()), + ])) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(r#"{"access_token":"fresh","expires_in":3600}"#) + .create_async() + .await; + + let token = client_credentials_access_token( + dir.path(), + "demo", + &format!("{}/token", server.url()), + cc_cache(None, None), + ) + .await + .unwrap(); + + assert_eq!(token, "fresh"); + mock.assert_async().await; + // Saved back as client credentials (secret retained), not downgraded to Bearer. + match load_auth_cache(dir.path(), "demo").unwrap().unwrap() { + AuthCache::ClientCredentials { + client_id, + client_secret, + access_token, + expires_at, + } => { + assert_eq!(client_id, "id"); + assert_eq!(client_secret, "sec"); + assert_eq!(access_token.as_deref(), Some("fresh")); + assert!(expires_at.unwrap() >= expires_at_from_duration(3600) - 5); + } + other => panic!("expected client credentials, got {:?}", other), + } + } + + #[tokio::test] + async fn test_client_credentials_expired_token_is_reexchanged() { + let dir = tempfile::TempDir::new().unwrap(); + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/token") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(r#"{"access_token":"fresh","expires_in":3600}"#) + .create_async() + .await; + + // 30s left: inside the 60s grace period. + let cache = cc_cache(Some("stale"), Some(expires_at_from_duration(30))); + let token = client_credentials_access_token( + dir.path(), + "demo", + &format!("{}/token", server.url()), + cache, + ) + .await + .unwrap(); + + assert_eq!(token, "fresh"); + mock.assert_async().await; + } + + #[tokio::test] + async fn test_client_credentials_exchange_failure_reports_body_and_keeps_cache() { + let dir = tempfile::TempDir::new().unwrap(); + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/token") + .with_status(401) + .with_body(r#"{"error":"invalid_client"}"#) + .create_async() + .await; + + let err = client_credentials_access_token( + dir.path(), + "demo", + &format!("{}/token", server.url()), + cc_cache(None, None), + ) + .await + .unwrap_err(); + + let msg = format!("{:#}", err); + assert!(msg.contains("401"), "{}", msg); + assert!(msg.contains("invalid_client"), "{}", msg); + mock.assert_async().await; + assert!(load_auth_cache(dir.path(), "demo").unwrap().is_none()); + } + + #[tokio::test] + async fn test_resolve_auth_client_credentials_fresh_token_is_bearer() { + use crate::cli::Stage; + use std::collections::BTreeMap; + + let dir = tempfile::TempDir::new().unwrap(); + save_auth_cache( + dir.path(), + "demo", + &cc_cache(Some("cached"), Some(expires_at_from_duration(3600))), + ) + .unwrap(); + let ctx = AppContext { + config_dir: dir.path().to_path_buf(), + profile: "demo".to_string(), + auth_profile: "demo".to_string(), + stage: Stage::Dev, + auth_stage: Stage::Dev, + base_url: "http://example".to_string(), + defaults: BTreeMap::new(), + }; + + match resolve_auth(&ctx).await.unwrap() { + Some(ResolvedAuth::Bearer(t)) => assert_eq!(t, "cached"), + _ => panic!("expected bearer"), + } + } + + #[tokio::test] + async fn test_login_with_client_credentials_forces_exchange_and_saves() { + let dir = tempfile::TempDir::new().unwrap(); + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/token") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(r#"{"access_token":"fresh","expires_in":3600}"#) + .create_async() + .await; + // A still-valid cached token must not suppress the exchange on explicit login. + save_auth_cache( + dir.path(), + "m2m", + &cc_cache(Some("cached"), Some(expires_at_from_duration(3600))), + ) + .unwrap(); + + let out = Output { json: false }; + let handled = login_with_client_credentials( + dir.path(), + "m2m", + &format!("{}/token", server.url()), + &out, + ) + .await + .unwrap(); + + assert!(handled); + mock.assert_async().await; + match load_auth_cache(dir.path(), "m2m").unwrap().unwrap() { + AuthCache::ClientCredentials { + client_secret, + access_token, + .. + } => { + assert_eq!(client_secret, "sec"); + assert_eq!(access_token.as_deref(), Some("fresh")); + } + other => panic!("expected client credentials, got {:?}", other), + } + } + + #[tokio::test] + async fn test_login_with_client_credentials_ignores_other_profiles() { + let dir = tempfile::TempDir::new().unwrap(); + let mut server = mockito::Server::new_async().await; + let mock = server.mock("POST", "/token").expect(0).create_async().await; + let out = Output { json: false }; + let url = format!("{}/token", server.url()); + + // No credentials stored. + assert!( + !login_with_client_credentials(dir.path(), "demo", &url, &out) + .await + .unwrap() + ); + // Basic credentials stored. + save_auth_cache( + dir.path(), + "demo", + &AuthCache::Basic { + username: "alice".to_string(), + password: "secret".to_string(), + }, + ) + .unwrap(); + assert!( + !login_with_client_credentials(dir.path(), "demo", &url, &out) + .await + .unwrap() + ); + mock.assert_async().await; + } + + fn logout_ctx(dir: &Path) -> AppContext { + use crate::cli::Stage; + use std::collections::BTreeMap; + AppContext { + config_dir: dir.to_path_buf(), + profile: "m2m".to_string(), + auth_profile: "m2m".to_string(), + stage: Stage::Dev, + auth_stage: Stage::Dev, + base_url: "http://example".to_string(), + defaults: BTreeMap::new(), + } + } + + #[test] + fn test_logout_client_credentials_keeps_credentials_and_drops_token() { + let dir = tempfile::TempDir::new().unwrap(); + save_auth_cache( + dir.path(), + "m2m", + &cc_cache(Some("cached"), Some(expires_at_from_duration(3600))), + ) + .unwrap(); + + logout(&logout_ctx(dir.path()), &Output { json: false }).unwrap(); + + match load_auth_cache(dir.path(), "m2m").unwrap().unwrap() { + AuthCache::ClientCredentials { + client_id, + client_secret, + access_token, + expires_at, + } => { + assert_eq!(client_id, "id"); + assert_eq!(client_secret, "sec"); + assert!(access_token.is_none()); + assert!(expires_at.is_none()); + } + other => panic!("expected client credentials, got {:?}", other), + } + } + + #[test] + fn test_logout_deletes_malformed_auth_cache() { + let dir = tempfile::TempDir::new().unwrap(); + let path = crate::auth_cache::auth_cache_path(dir.path(), "m2m"); + std::fs::create_dir_all(path.parent().unwrap()).unwrap(); + std::fs::write(&path, "not json").unwrap(); + + logout(&logout_ctx(dir.path()), &Output { json: false }).unwrap(); + + assert!(!path.exists()); + } + + #[test] + fn test_logout_basic_still_deletes_credentials() { + let dir = tempfile::TempDir::new().unwrap(); + save_auth_cache( + dir.path(), + "m2m", + &AuthCache::Basic { + username: "alice".to_string(), + password: "secret".to_string(), + }, + ) + .unwrap(); + + logout(&logout_ctx(dir.path()), &Output { json: false }).unwrap(); + + assert!(load_auth_cache(dir.path(), "m2m").unwrap().is_none()); + } + + #[test] + fn test_status_output_json_client_credentials() { + let output = StatusOutput { + auth_type: Some("client_credentials".to_string()), + authenticated: true, + expired: true, + username: None, + client_id: Some("client_123".to_string()), + profile: "demo".to_string(), + }; + let json = serde_json::to_value(&output).unwrap(); + assert_eq!(json["type"], "client_credentials"); + assert_eq!(json["client_id"], "client_123"); + assert!(json.get("client_secret").is_none()); + assert!(json.get("username").is_none()); + } + + #[test] + fn test_status_output_plain_client_credentials() { + let output = StatusOutput { + auth_type: Some("client_credentials".to_string()), + authenticated: true, + expired: false, + username: None, + client_id: Some("client_123".to_string()), + profile: "demo".to_string(), + }; + assert_eq!( + output.plain(), + "✓ Authenticated using profile 'demo' (client credentials, client: client_123)." + ); + } + #[tokio::test] async fn test_attach_auth_uses_override_credentials() { use crate::auth_cache::{save_auth_cache, AuthCache}; diff --git a/src/commands/config/doctor.rs b/src/commands/config/doctor.rs index 22b3f83..a94cdb6 100644 --- a/src/commands/config/doctor.rs +++ b/src/commands/config/doctor.rs @@ -83,7 +83,8 @@ pub async fn doctor( } /// Run all checks and assemble the report. Pure-ish: side effects are limited -/// to filesystem reads (auth cache) and one HTTP request. +/// to filesystem reads (auth cache) and one HTTP request, plus — for client +/// credentials — a token exchange that writes the new token back to the auth cache. pub(crate) async fn run_checks( config: &Config, config_dir: &Path, @@ -117,7 +118,7 @@ pub(crate) async fn run_checks( // 3. API reachability. let api = match (&profile_ctx, auth_ok, cache_load.as_ref()) { - (Some(ctx), true, Some(Ok(Some(cache)))) => check_api(ctx, cache).await, + (Some(ctx), true, Some(Ok(Some(cache)))) => check_api(ctx, cache, config_dir).await, _ => skip("api", "auth check failed"), }; checks.push(api); @@ -138,6 +139,8 @@ pub(crate) async fn run_checks( struct ProfileCtx { profile: String, organization: String, + /// The profile's saved stage (not the `-g` override): where its credentials are issued. + auth_stage: Stage, base_url: String, } @@ -149,9 +152,11 @@ fn resolve_profile_ctx( let (profile, organization, stage) = crate::config::resolve_profile(config, profile_override, stage_override).ok()?; let base_url = resolve_api(&organization, &stage); + let auth_stage = config.profiles.get(&profile)?.stage.clone(); Some(ProfileCtx { profile, organization, + auth_stage, base_url, }) } @@ -250,6 +255,17 @@ fn check_auth(loaded: &Result, anyhow::Error>) -> CheckResult } } } + AuthCache::ClientCredentials { client_id, .. } => { + let mut details = BTreeMap::new(); + details.insert("type".to_string(), serde_json::json!("client_credentials")); + details.insert("client_id".to_string(), serde_json::json!(client_id)); + CheckResult { + name: "auth".to_string(), + status: CheckStatus::Pass, + message: format!("client credentials (client: {})", client_id), + details, + } + } AuthCache::Basic { username, .. } => { let mut details = BTreeMap::new(); details.insert("type".to_string(), serde_json::json!("basic")); @@ -264,7 +280,7 @@ fn check_auth(loaded: &Result, anyhow::Error>) -> CheckResult } } -async fn check_api(ctx: &ProfileCtx, auth: &AuthCache) -> CheckResult { +async fn check_api(ctx: &ProfileCtx, auth: &AuthCache, config_dir: &Path) -> CheckResult { let url = format!("{}/clinicians/me", ctx.base_url.trim_end_matches('/')); let client = reqwest::Client::new(); let mut req = client.get(&url); @@ -273,6 +289,28 @@ async fn check_api(ctx: &ProfileCtx, auth: &AuthCache) -> CheckResult { reqwest::header::AUTHORIZATION, format!("Bearer {}", access_token), ), + AuthCache::ClientCredentials { .. } => { + match crate::commands::auth::client_credentials_access_token( + config_dir, + &ctx.profile, + ctx.auth_stage.workos_config().token_url, + auth.clone(), + ) + .await + { + Ok(token) => { + req.header(reqwest::header::AUTHORIZATION, format!("Bearer {}", token)) + } + Err(e) => { + return CheckResult { + name: "api".to_string(), + status: CheckStatus::Fail, + message: format!("could not obtain access token: {:#}", e), + details: BTreeMap::new(), + } + } + } + } AuthCache::Basic { username, password } => req.basic_auth(username, Some(password)), }; @@ -501,6 +539,32 @@ mod tests { assert!(!json.contains("\"rt\"")); } + #[test] + fn test_check_auth_client_credentials_passes() { + let r = check_auth(&Ok(Some(AuthCache::ClientCredentials { + client_id: "client_123".to_string(), + client_secret: "supersecret".to_string(), + access_token: None, + expires_at: None, + }))); + assert_eq!(r.status, CheckStatus::Pass); + assert!(r.message.contains("client credentials")); + assert!(r.message.contains("client_123")); + } + + #[test] + fn test_check_auth_client_credentials_details_omit_secret() { + let r = check_auth(&Ok(Some(AuthCache::ClientCredentials { + client_id: "client_123".to_string(), + client_secret: "supersecret".to_string(), + access_token: Some("tok-secret".to_string()), + expires_at: Some(unix_now() + 3600), + }))); + let json = serde_json::to_string(&r).unwrap(); + assert!(!json.contains("supersecret")); + assert!(!json.contains("tok-secret")); + } + // --- check_defaults --- #[test] @@ -545,6 +609,7 @@ mod tests { let ctx = ProfileCtx { profile: "demo".to_string(), organization: "demonstration".to_string(), + auth_stage: Stage::Prod, base_url: server.url(), }; let auth = AuthCache::Bearer { @@ -552,7 +617,7 @@ mod tests { refresh_token: None, expires_at: unix_now() + 3600, }; - let r = check_api(&ctx, &auth).await; + let r = check_api(&ctx, &auth, Path::new("/tmp")).await; assert_eq!(r.status, CheckStatus::Pass); assert!(r.message.contains("200")); mock.assert_async().await; @@ -571,6 +636,7 @@ mod tests { let ctx = ProfileCtx { profile: "demo".to_string(), organization: "demonstration".to_string(), + auth_stage: Stage::Prod, base_url: server.url(), }; let auth = AuthCache::Bearer { @@ -578,7 +644,7 @@ mod tests { refresh_token: None, expires_at: unix_now() + 3600, }; - let r = check_api(&ctx, &auth).await; + let r = check_api(&ctx, &auth, Path::new("/tmp")).await; assert_eq!(r.status, CheckStatus::Fail); assert!(r.message.contains("401")); mock.assert_async().await; @@ -599,13 +665,14 @@ mod tests { let ctx = ProfileCtx { profile: "demo".to_string(), organization: "demonstration".to_string(), + auth_stage: Stage::Prod, base_url: server.url(), }; let auth = AuthCache::Basic { username: "alice".to_string(), password: "secret".to_string(), }; - let r = check_api(&ctx, &auth).await; + let r = check_api(&ctx, &auth, Path::new("/tmp")).await; assert_eq!(r.status, CheckStatus::Pass); mock.assert_async().await; } @@ -623,6 +690,7 @@ mod tests { let ctx = ProfileCtx { profile: "demo".to_string(), organization: "demonstration".to_string(), + auth_stage: Stage::Prod, base_url: format!("http://{}", addr), }; let auth = AuthCache::Bearer { @@ -630,7 +698,7 @@ mod tests { refresh_token: None, expires_at: unix_now() + 3600, }; - let r = check_api(&ctx, &auth).await; + let r = check_api(&ctx, &auth, Path::new("/tmp")).await; assert_eq!(r.status, CheckStatus::Fail); assert!(r.message.contains("could not reach")); } @@ -723,6 +791,15 @@ mod tests { assert_eq!(ctx.base_url, "http://localhost:8080"); } + #[test] + fn test_resolve_profile_ctx_auth_stage_ignores_stage_override() { + // Credentials belong to the profile's saved stage, as in real commands. + let config = cfg_with_default(Stage::Prod); + let ctx = resolve_profile_ctx(&config, None, Some(&Stage::Dev)).unwrap(); + assert_eq!(ctx.auth_stage, Stage::Prod); + assert_eq!(ctx.base_url, "https://demonstration.roundingwell.dev/api"); + } + #[test] fn test_resolve_profile_ctx_without_override_uses_configured_stage() { let config = cfg_with_default(Stage::Prod); diff --git a/src/commands/config/profile.rs b/src/commands/config/profile.rs index 1f6cbef..5a3c1e0 100644 --- a/src/commands/config/profile.rs +++ b/src/commands/config/profile.rs @@ -115,11 +115,29 @@ impl CommandOutput for ProfileRmOutput { #[derive(Serialize)] pub struct ProfileAuthOutput { pub name: String, + #[serde(skip)] + pub kind: &'static str, } impl CommandOutput for ProfileAuthOutput { fn plain(&self) -> String { - format!("Basic auth credentials saved for profile '{}'.", self.name) + format!( + "{} credentials saved for profile '{}'.", + self.kind, self.name + ) + } +} + +/// Returns `value` if non-empty, prompts when absent, and rejects an empty value. +fn required( + value: Option, + label: &str, + prompt: impl FnOnce() -> Result, +) -> Result { + match value { + Some(v) if v.is_empty() => anyhow::bail!("{} cannot be empty", label), + Some(v) => Ok(v), + None => prompt(), } } @@ -147,6 +165,7 @@ pub fn profile_show(config: &Config, config_dir: &Path, out: &Output) -> Result< let auth = match load_auth_cache(config_dir, name)? { Some(AuthCache::Basic { .. }) => Some("basic".to_string()), Some(AuthCache::Bearer { .. }) => Some("bearer".to_string()), + Some(AuthCache::ClientCredentials { .. }) => Some("client_credentials".to_string()), None => None, }; @@ -296,40 +315,40 @@ pub fn profile_auth( anyhow::bail!("profile '{}' does not exist", args.name); } - if out.json && (args.username.is_none() || args.password.is_none()) { - anyhow::bail!("cannot use interactive mode with --json; provide --username and --password"); - } - - let username = args - .username - .map(|u| { - if u.is_empty() { - anyhow::bail!("username cannot be empty") - } else { - Ok(u) - } - }) - .unwrap_or_else(|| p::text("Username"))?; - - let pw = args - .password - .map(|p| { - if p.is_empty() { - anyhow::bail!("password cannot be empty") - } else { - Ok(p) - } - }) - .unwrap_or_else(p::password)?; + let client_credentials = args.client_id.is_some() || args.client_secret.is_some(); - let cache = AuthCache::Basic { - username, - password: pw, + let (cache, kind) = if client_credentials { + if out.json && (args.client_id.is_none() || args.client_secret.is_none()) { + anyhow::bail!( + "cannot use interactive mode with --json; provide --client-id and --client-secret" + ); + } + let cache = AuthCache::ClientCredentials { + client_id: required(args.client_id, "client id", || p::text("Client ID"))?, + client_secret: required(args.client_secret, "client secret", || { + p::secret("Client secret") + })?, + access_token: None, + expires_at: None, + }; + (cache, "Client") + } else { + if out.json && (args.username.is_none() || args.password.is_none()) { + anyhow::bail!( + "cannot use interactive mode with --json; provide --username and --password" + ); + } + let cache = AuthCache::Basic { + username: required(args.username, "username", || p::text("Username"))?, + password: required(args.password, "password", p::password)?, + }; + (cache, "Basic auth") }; save_auth_cache(config_dir, &args.name, &cache)?; out.print(&ProfileAuthOutput { name: args.name.clone(), + kind, }); Ok(()) } @@ -487,6 +506,26 @@ mod tests { assert_eq!(output.auth, Some("bearer".to_string())); } + #[test] + fn test_profile_show_auth_type_client_credentials() { + let dir = tempfile::TempDir::new().unwrap(); + let mut config = config_with_profile("demo", "mercy", Stage::Prod); + config.default = Some("demo".to_string()); + save_auth_cache( + dir.path(), + "demo", + &AuthCache::ClientCredentials { + client_id: "id".to_string(), + client_secret: "sec".to_string(), + access_token: None, + expires_at: None, + }, + ) + .unwrap(); + let output = profile_show(&config, dir.path(), &out_plain()).unwrap(); + assert_eq!(output.auth, Some("client_credentials".to_string())); + } + #[test] fn test_profile_use_sets_default() { let (_tmp, path) = tmp_path(); @@ -711,6 +750,8 @@ mod tests { name: "demo".to_string(), username: Some("alice".to_string()), password: Some("secret".to_string()), + client_id: None, + client_secret: None, }; profile_auth(args, &config, dir.path(), &out_plain()).unwrap(); let cache = load_auth_cache(dir.path(), "demo").unwrap().unwrap(); @@ -723,6 +764,85 @@ mod tests { } } + fn client_args(id: Option<&str>, secret: Option<&str>) -> ConfigProfileAuthArgs { + ConfigProfileAuthArgs { + name: "demo".to_string(), + username: None, + password: None, + client_id: id.map(str::to_string), + client_secret: secret.map(str::to_string), + } + } + + #[test] + fn test_profile_auth_saves_client_credentials_cache() { + let dir = tempfile::TempDir::new().unwrap(); + let config = config_with_profile("demo", "mercy", Stage::Prod); + profile_auth( + client_args(Some("client_123"), Some("sec")), + &config, + dir.path(), + &out_plain(), + ) + .unwrap(); + match load_auth_cache(dir.path(), "demo").unwrap().unwrap() { + AuthCache::ClientCredentials { + client_id, + client_secret, + access_token, + expires_at, + } => { + assert_eq!(client_id, "client_123"); + assert_eq!(client_secret, "sec"); + assert!(access_token.is_none()); + assert!(expires_at.is_none()); + } + _ => panic!("expected client credentials cache"), + } + } + + #[test] + fn test_profile_auth_client_json_mode_requires_both() { + let dir = tempfile::TempDir::new().unwrap(); + let config = config_with_profile("demo", "mercy", Stage::Prod); + let err = profile_auth( + client_args(Some("client_123"), None), + &config, + dir.path(), + &out_json(), + ) + .unwrap_err(); + assert!(err.to_string().contains("--client-id and --client-secret")); + } + + #[test] + fn test_profile_auth_rejects_empty_client_id() { + let dir = tempfile::TempDir::new().unwrap(); + let config = config_with_profile("demo", "mercy", Stage::Prod); + let err = profile_auth( + client_args(Some(""), Some("sec")), + &config, + dir.path(), + &out_plain(), + ) + .unwrap_err(); + assert!(err.to_string().contains("client id cannot be empty")); + } + + #[test] + fn test_profile_auth_rejects_empty_client_secret() { + let dir = tempfile::TempDir::new().unwrap(); + let config = config_with_profile("demo", "mercy", Stage::Prod); + let err = profile_auth( + client_args(Some("client_123"), Some("")), + &config, + dir.path(), + &out_plain(), + ) + .unwrap_err(); + assert!(err.to_string().contains("client secret cannot be empty")); + } + #[test] fn test_profile_auth_errors_when_profile_not_found() { let dir = tempfile::TempDir::new().unwrap(); @@ -731,6 +851,8 @@ mod tests { name: "missing".to_string(), username: Some("alice".to_string()), password: Some("secret".to_string()), + client_id: None, + client_secret: None, }; let err = profile_auth(args, &config, dir.path(), &out_plain()).unwrap_err(); assert!(err.to_string().contains("does not exist")); @@ -744,6 +866,8 @@ mod tests { name: "demo".to_string(), username: None, password: Some("secret".to_string()), + client_id: None, + client_secret: None, }; let err = profile_auth(args, &config, dir.path(), &out_json()).unwrap_err(); assert!(err.to_string().contains("--json")); @@ -757,6 +881,8 @@ mod tests { name: "demo".to_string(), username: Some("".to_string()), password: Some("secret".to_string()), + client_id: None, + client_secret: None, }; let err = profile_auth(args, &config, dir.path(), &out_plain()).unwrap_err(); assert!(err.to_string().contains("username cannot be empty")); @@ -770,6 +896,8 @@ mod tests { name: "demo".to_string(), username: Some("alice".to_string()), password: Some("".to_string()), + client_id: None, + client_secret: None, }; let err = profile_auth(args, &config, dir.path(), &out_plain()).unwrap_err(); assert!(err.to_string().contains("password cannot be empty")); diff --git a/src/prompt.rs b/src/prompt.rs index 1d6f13e..ef8f73d 100644 --- a/src/prompt.rs +++ b/src/prompt.rs @@ -127,18 +127,23 @@ pub fn stage() -> Result { stage_with(std::io::stdin().lock(), std::io::stderr().lock()) } -/// Reads a password from the terminal without echoing it. +/// Reads a secret from the terminal without echoing it. /// Re-prompts on empty input. Backed by the `rpassword` crate. -pub fn password() -> Result { +pub fn secret(label: &str) -> Result { loop { - let pw = rpassword::prompt_password("Password: ")?; - if !pw.is_empty() { - return Ok(pw); + let value = rpassword::prompt_password(format!("{}: ", label))?; + if !value.is_empty() { + return Ok(value); } - eprintln!("Password cannot be empty"); + eprintln!("{} cannot be empty", label); } } +/// Reads a password from the terminal without echoing it. +pub fn password() -> Result { + secret("Password") +} + #[cfg(test)] mod tests { use super::*;