use std::{fmt, sync::Arc};
use serde::Deserialize;
use zeroize::Zeroizing;
#[derive(Debug, Clone, Deserialize)]
pub struct OidcEndpoints {
pub authorization_endpoint: String,
pub token_endpoint: String,
}
#[derive(Debug, Deserialize)]
pub struct OidcTokenResponse {
pub access_token: String,
pub id_token: Option<String>,
pub expires_in: Option<u64>,
pub refresh_token: Option<String>,
}
pub struct OidcServerClient {
client_id: String,
pub(crate) client_secret: Zeroizing<String>,
server_redirect_uri: String,
authorization_endpoint: String,
token_endpoint: String,
}
#[allow(clippy::missing_fields_in_debug)] impl fmt::Debug for OidcServerClient {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OidcServerClient")
.field("client_id", &self.client_id)
.field("client_secret", &"[REDACTED]")
.field("authorization_endpoint", &self.authorization_endpoint)
.finish_non_exhaustive()
}
}
impl OidcServerClient {
pub(crate) const MAX_CODE_VERIFIER_BYTES: usize = 128;
pub(crate) const MAX_OIDC_RESPONSE_BYTES: usize = 1024 * 1024;
pub(crate) const MIN_CODE_VERIFIER_BYTES: usize = 43;
pub fn new(
client_id: impl Into<String>,
client_secret: impl Into<String>,
server_redirect_uri: impl Into<String>,
authorization_endpoint: impl Into<String>,
token_endpoint: impl Into<String>,
) -> Self {
Self {
client_id: client_id.into(),
client_secret: Zeroizing::new(client_secret.into()),
server_redirect_uri: server_redirect_uri.into(),
authorization_endpoint: authorization_endpoint.into(),
token_endpoint: token_endpoint.into(),
}
}
pub fn from_compiled_schema(schema_json: &serde_json::Value) -> Option<Arc<Self>> {
#[derive(Deserialize)]
struct AuthCfg {
client_id: String,
client_secret_env: String,
server_redirect_uri: String,
}
let auth_cfg: AuthCfg =
schema_json.get("auth").and_then(|v| serde_json::from_value(v.clone()).ok())?;
let Ok(client_secret) = std::env::var(&auth_cfg.client_secret_env) else {
tracing::error!(
env_var = %auth_cfg.client_secret_env,
"PKCE init failed: env var for OIDC client secret is not set"
);
return None;
};
let Some(endpoints): Option<OidcEndpoints> = schema_json
.get("auth_endpoints")
.and_then(|v| serde_json::from_value(v.clone()).ok())
else {
tracing::error!(
"PKCE init failed: 'auth_endpoints' not found in compiled schema. \
Re-compile the schema so that the CLI caches the OIDC discovery \
document (authorization_endpoint, token_endpoint)."
);
return None;
};
Some(Arc::new(Self {
client_id: auth_cfg.client_id,
client_secret: Zeroizing::new(client_secret),
server_redirect_uri: auth_cfg.server_redirect_uri,
authorization_endpoint: endpoints.authorization_endpoint,
token_endpoint: endpoints.token_endpoint,
}))
}
#[must_use]
pub fn authorization_url(
&self,
state: &str,
code_challenge: &str,
code_challenge_method: &str,
) -> String {
format!(
"{}?response_type=code\
&client_id={}\
&redirect_uri={}\
&scope=openid%20email%20profile\
&state={}\
&code_challenge={}\
&code_challenge_method={}",
self.authorization_endpoint,
urlencoding::encode(&self.client_id),
urlencoding::encode(&self.server_redirect_uri),
urlencoding::encode(state),
urlencoding::encode(code_challenge),
code_challenge_method,
)
}
pub async fn exchange_code(
&self,
code: &str,
code_verifier: &str,
http: &reqwest::Client,
) -> Result<OidcTokenResponse, anyhow::Error> {
anyhow::ensure!(
code_verifier.len() >= Self::MIN_CODE_VERIFIER_BYTES,
"code_verifier too short ({} bytes, min {})",
code_verifier.len(),
Self::MIN_CODE_VERIFIER_BYTES,
);
anyhow::ensure!(
code_verifier.len() <= Self::MAX_CODE_VERIFIER_BYTES,
"code_verifier too long ({} bytes, max {})",
code_verifier.len(),
Self::MAX_CODE_VERIFIER_BYTES,
);
let resp = http
.post(&self.token_endpoint)
.form(&[
("grant_type", "authorization_code"),
("code", code),
("code_verifier", code_verifier),
("redirect_uri", self.server_redirect_uri.as_str()),
("client_id", self.client_id.as_str()),
("client_secret", self.client_secret.as_str()),
])
.send()
.await?;
let status = resp.status();
let body_bytes = resp
.bytes()
.await
.map_err(|e| anyhow::anyhow!("Failed to read token response: {e}"))?;
anyhow::ensure!(
body_bytes.len() <= Self::MAX_OIDC_RESPONSE_BYTES,
"OIDC token response too large ({} bytes, max {})",
body_bytes.len(),
Self::MAX_OIDC_RESPONSE_BYTES
);
if !status.is_success() {
let body = String::from_utf8_lossy(&body_bytes);
anyhow::bail!("token endpoint returned {status}: {body}");
}
Ok(serde_json::from_slice::<OidcTokenResponse>(&body_bytes)?)
}
}