Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions crates/openshell-bootstrap/src/oidc_token.rs
Original file line number Diff line number Diff line change
Expand Up @@ -124,6 +124,21 @@ pub fn is_token_expired(bundle: &OidcTokenBundle) -> bool {
now + 30 >= expires_at
}

/// Check whether the stored access token has actually expired.
///
/// Unlike [`is_token_expired`], this does not include the 30-second refresh
/// window. Tokens without expiry metadata are treated as valid.
pub fn is_token_actually_expired(bundle: &OidcTokenBundle) -> bool {
let Some(expires_at) = bundle.expires_at else {
return false;
};
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
now >= expires_at
}

#[cfg(test)]
mod tests {
use super::*;
Expand Down Expand Up @@ -163,6 +178,30 @@ mod tests {
});
}

#[test]
fn expiry_checks_distinguish_refresh_window_from_actual_expiry() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let bundle = OidcTokenBundle {
access_token: "token".to_string(),
refresh_token: None,
expires_at: Some(now + 10),
issuer: "https://issuer.example.com".to_string(),
client_id: "openshell-cli".to_string(),
};

assert!(is_token_expired(&bundle));
assert!(!is_token_actually_expired(&bundle));

let expired = OidcTokenBundle {
expires_at: Some(now.saturating_sub(1)),
..bundle
};
assert!(is_token_actually_expired(&expired));
}

#[test]
fn oidc_login_prompt_marker_is_per_gateway() {
let tmp = tempfile::tempdir().unwrap();
Expand Down
64 changes: 36 additions & 28 deletions crates/openshell-cli/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -156,8 +156,11 @@ fn resolve_gateway_name(gateway_flag: &Option<String>) -> Option<String> {
/// Handles Cloudflare Access and OIDC auth modes by loading the stored token
/// and setting it on `TlsOptions`. For OIDC, automatically refreshes the token
/// if it's near expiry.
fn apply_auth(tls: &mut TlsOptions, gateway_name: &str) {
let _ = apply_auth_with_status(tls, gateway_name);
fn apply_auth(tls: &mut TlsOptions, gateway_name: &str) -> Result<()> {
if let Some(error) = apply_auth_with_status(tls, gateway_name) {
return Err(miette::miette!(error));
}
Ok(())
}

/// Apply stored authentication and return a user-facing preparation failure,
Expand Down Expand Up @@ -206,12 +209,17 @@ fn apply_auth_with_status(tls: &mut TlsOptions, gateway_name: &str) -> Option<St
}
Err(e) => {
tracing::warn!("OIDC token refresh failed: {e}");
// Use the expired token anyway — server will reject it
// with a clear error prompting re-login.
tls.oidc_token = Some(bundle.access_token);
Some(format!(
"OIDC token refresh failed; run `openshell gateway login {gateway_name}`"
))
if openshell_bootstrap::oidc_token::is_token_actually_expired(&bundle) {
Some(format!(
"OIDC token refresh failed: {e}\ncached OIDC token has expired; run `openshell gateway login {gateway_name}`"
))
} else {
tls.oidc_token = Some(bundle.access_token);
eprintln!(
"Warning: OIDC token refresh failed; continuing with the cached token: {e}"
);
None
}
}
}
} else {
Expand Down Expand Up @@ -2506,7 +2514,7 @@ async fn run_async() -> Result<()> {
GatewayCommands::Info { output } => {
if let Ok(ctx) = resolve_gateway(&cli.gateway, &cli.gateway_endpoint) {
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
run::gateway_info(&ctx.name, &ctx.endpoint, &tls, output.as_str()).await?;
} else {
run::gateway_info_not_configured()?;
Expand Down Expand Up @@ -2574,7 +2582,7 @@ async fn run_async() -> Result<()> {
Some(Commands::Whoami { output }) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
run::whoami(&ctx.endpoint, &tls, output.as_str()).await?;
}

Expand Down Expand Up @@ -2677,7 +2685,7 @@ async fn run_async() -> Result<()> {
} => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let name = resolve_sandbox_name(name, &ctx.name, &cli.workspace)?;
let local = local.unwrap_or_else(|| target_port.to_string());
run::service_forward_tcp(
Expand All @@ -2699,7 +2707,7 @@ async fn run_async() -> Result<()> {
let spec = openshell_core::forward::ForwardSpec::parse(&port)?;
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let name = resolve_sandbox_name(name, &ctx.name, &cli.workspace)?;
run::sandbox_forward(
&ctx.endpoint,
Expand Down Expand Up @@ -2730,7 +2738,7 @@ async fn run_async() -> Result<()> {
}) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
match command {
ServiceCommands::Expose {
sandbox,
Expand Down Expand Up @@ -2792,7 +2800,7 @@ async fn run_async() -> Result<()> {
}) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let name = resolve_sandbox_name(name, &ctx.name, &cli.workspace)?;
run::sandbox_logs(
&ctx.endpoint,
Expand All @@ -2816,7 +2824,7 @@ async fn run_async() -> Result<()> {
}) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
match policy_cmd {
PolicyCommands::Set {
name,
Expand Down Expand Up @@ -2970,7 +2978,7 @@ async fn run_async() -> Result<()> {
}) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;

match settings_cmd {
SettingsCommands::Get { name, global, json } => {
Expand Down Expand Up @@ -3049,7 +3057,7 @@ async fn run_async() -> Result<()> {
}) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
match draft_cmd {
DraftCommands::Get { name, status } => {
let name = resolve_sandbox_name(name, &ctx.name, &cli.workspace)?;
Expand Down Expand Up @@ -3124,7 +3132,7 @@ async fn run_async() -> Result<()> {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let endpoint = &ctx.endpoint;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
match command {
InferenceCommands::Set {
provider,
Expand Down Expand Up @@ -3275,7 +3283,7 @@ async fn run_async() -> Result<()> {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let endpoint = &ctx.endpoint;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let exit_code = Box::pin(run::sandbox_create(
endpoint,
&ctx.name,
Expand Down Expand Up @@ -3318,7 +3326,7 @@ async fn run_async() -> Result<()> {
} => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let local = std::path::Path::new(&local_path);
run::sandbox_upload(
&ctx.endpoint,
Expand All @@ -3338,7 +3346,7 @@ async fn run_async() -> Result<()> {
} => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let local_dest = dest.as_deref().unwrap_or(".");
eprintln!("Downloading sandbox:{sandbox_path} -> {local_dest}");
run::sandbox_sync_down(
Expand All @@ -3356,7 +3364,7 @@ async fn run_async() -> Result<()> {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let endpoint = &ctx.endpoint;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
match other {
SandboxCommands::Create { .. }
| SandboxCommands::Upload { .. }
Expand Down Expand Up @@ -3615,7 +3623,7 @@ async fn run_async() -> Result<()> {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let endpoint = &ctx.endpoint;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;

match command {
WorkspaceCommands::Create { name, labels } => {
Expand Down Expand Up @@ -3681,7 +3689,7 @@ async fn run_async() -> Result<()> {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let endpoint = &ctx.endpoint;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;

match command {
ProviderCommands::Create {
Expand Down Expand Up @@ -3896,7 +3904,7 @@ async fn run_async() -> Result<()> {
Some(Commands::Term { theme }) => {
let ctx = resolve_gateway(&cli.gateway, &cli.gateway_endpoint)?;
let mut tls = tls.with_gateway_name(&ctx.name);
apply_auth(&mut tls, &ctx.name);
apply_auth(&mut tls, &ctx.name)?;
let channel = openshell_cli::tls::build_channel(&ctx.endpoint, &tls).await?;
let interceptor = openshell_core::auth::EdgeAuthInterceptor::new(
tls.oidc_token.as_deref(),
Expand Down Expand Up @@ -3940,7 +3948,7 @@ async fn run_async() -> Result<()> {
None => tls,
};
if let Some(ref g) = gateway_name_opt {
apply_auth(&mut effective_tls, g);
apply_auth(&mut effective_tls, g)?;
}
run::sandbox_ssh_proxy(&gw, &sid, &tok, &effective_tls).await?;
}
Expand All @@ -3959,7 +3967,7 @@ async fn run_async() -> Result<()> {
meta.gateway_endpoint
};
let mut tls = tls.with_gateway_name(&g);
apply_auth(&mut tls, &g);
apply_auth(&mut tls, &g)?;
run::sandbox_ssh_proxy_by_name(&endpoint, &n, &tls, &cli.workspace).await?;
}
// Legacy name mode with --server only (no --gateway-name).
Expand Down Expand Up @@ -4645,7 +4653,7 @@ mod tests {
store_edge_token("edge-gateway", "token-123").unwrap();

let mut tls = TlsOptions::default();
apply_auth(&mut tls, "edge-gateway");
apply_auth(&mut tls, "edge-gateway").unwrap();

assert_eq!(tls.edge_token.as_deref(), Some("token-123"));
});
Expand Down
Loading
Loading