use std::fmt;
use crate::crypto::zeroize::Zeroizing;
use crate::json::JsonValue;
use crate::util::log::{info, warn};
#[derive(Debug, Clone, PartialEq, Eq)]
enum TokenResponseErrorKind {
InvalidJson,
OAuthError {
error: String,
description: Option<String>,
},
MissingAccessToken,
MissingTokenType,
InvalidExpiresIn,
}
#[doc(alias = "token_error")]
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TokenResponseError {
kind: TokenResponseErrorKind,
}
impl TokenResponseError {
const fn new(kind: TokenResponseErrorKind) -> Self {
Self { kind }
}
}
impl fmt::Display for TokenResponseError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.kind {
TokenResponseErrorKind::InvalidJson => {
write!(f, "token response: invalid JSON")
}
TokenResponseErrorKind::OAuthError { error, description } => {
write!(f, "token response: OAuth error: {error}")?;
if let Some(desc) = description {
write!(f, " ({desc})")?;
}
Ok(())
}
TokenResponseErrorKind::MissingAccessToken => {
write!(f, "token response: missing access_token")
}
TokenResponseErrorKind::MissingTokenType => {
write!(f, "token response: missing token_type")
}
TokenResponseErrorKind::InvalidExpiresIn => {
write!(f, "token response: invalid expires_in value")
}
}
}
}
impl std::error::Error for TokenResponseError {}
#[doc(alias = "token_response")]
pub struct TokenResponse {
access_token: Zeroizing<String>,
token_type: String,
expires_in: Option<u64>,
refresh_token: Option<Zeroizing<String>>,
scope: Option<String>,
id_token: Option<String>,
}
impl TokenResponse {
#[must_use]
#[inline]
pub fn access_token(&self) -> &str {
&self.access_token
}
#[must_use]
#[inline]
pub fn token_type(&self) -> &str {
&self.token_type
}
#[must_use]
#[inline]
pub fn expires_in(&self) -> Option<u64> {
self.expires_in
}
#[must_use]
#[inline]
pub fn refresh_token(&self) -> Option<&str> {
self.refresh_token.as_ref().map(|z| z.as_str())
}
#[must_use]
#[inline]
pub fn scope(&self) -> Option<&str> {
self.scope.as_deref()
}
#[must_use]
#[inline]
pub fn id_token(&self) -> Option<&str> {
self.id_token.as_deref()
}
}
impl TokenResponse {
#[must_use = "parsing may fail; handle the Result"]
pub fn parse(json: &str) -> Result<Self, TokenResponseError> {
let value = JsonValue::parse(json).map_err(|_| {
warn!("oauth: token response parse failed: invalid JSON");
TokenResponseError::new(TokenResponseErrorKind::InvalidJson)
})?;
if value.get("error").is_some() {
let error = value.get_str("error").unwrap_or("invalid_error").to_owned();
let description = value.get_str("error_description").map(String::from);
warn!(code = %error, "oauth: token endpoint error");
return Err(TokenResponseError::new(
TokenResponseErrorKind::OAuthError { error, description },
));
}
let access_token = Zeroizing::new(
value
.get_str("access_token")
.ok_or_else(|| {
warn!("oauth: token response parse failed: missing access_token");
TokenResponseError::new(TokenResponseErrorKind::MissingAccessToken)
})?
.to_owned(),
);
let token_type = value
.get_str("token_type")
.ok_or_else(|| {
warn!("oauth: token response parse failed: missing token_type");
TokenResponseError::new(TokenResponseErrorKind::MissingTokenType)
})?
.to_owned();
let expires_in = if let Some(val) = value.get("expires_in") {
if val.is_null() {
None
} else {
let secs = val.as_i64().ok_or_else(|| {
warn!("oauth: token response parse failed: invalid expires_in");
TokenResponseError::new(TokenResponseErrorKind::InvalidExpiresIn)
})?;
if secs < 0 {
warn!("oauth: token response parse failed: invalid expires_in");
return Err(TokenResponseError::new(
TokenResponseErrorKind::InvalidExpiresIn,
));
}
#[allow(clippy::cast_sign_loss)]
Some(secs as u64)
}
} else {
None
};
let refresh_token = value
.get_str("refresh_token")
.map(|s| Zeroizing::new(s.to_owned()));
let scope = value.get_str("scope").map(String::from);
let id_token = value.get_str("id_token").map(String::from);
info!(
token_type = %token_type,
expires_in = ?expires_in,
"oauth: token response parsed"
);
Ok(Self {
access_token,
token_type,
expires_in,
refresh_token,
scope,
id_token,
})
}
}
impl fmt::Debug for TokenResponse {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TokenResponse")
.field("access_token", &"[REDACTED]")
.field("token_type", &self.token_type)
.field("expires_in", &self.expires_in)
.field(
"refresh_token",
if self.refresh_token.is_some() {
&"Some([REDACTED])"
} else {
&"None"
},
)
.field("scope", &self.scope)
.field(
"id_token",
if self.id_token.is_some() {
&"Some([REDACTED])"
} else {
&"None"
},
)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_minimal_response() {
let json = r#"{"access_token":"abc123","token_type":"Bearer"}"#;
let resp = TokenResponse::parse(json).unwrap();
assert_eq!(resp.access_token(), "abc123");
assert_eq!(resp.token_type(), "Bearer");
assert_eq!(resp.expires_in(), None);
assert!(resp.refresh_token().is_none());
assert!(resp.scope().is_none());
assert!(resp.id_token().is_none());
}
#[test]
fn parse_oidc_response_exposes_id_token() {
let json = r#"{
"access_token": "ya29.a0Af...",
"token_type": "Bearer",
"expires_in": 3599,
"scope": "openid email profile",
"id_token": "eyJhbGciOiJSUzI1Ni'...header.payload.sig"
}"#;
let resp = TokenResponse::parse(json).unwrap();
assert_eq!(
resp.id_token(),
Some("eyJhbGciOiJSUzI1Ni'...header.payload.sig")
);
assert!(format!("{resp:?}").contains("id_token: \"Some([REDACTED])\""));
}
#[test]
fn parse_full_response() {
let json = r#"{
"access_token": "eyJhbGciOi...",
"token_type": "Bearer",
"expires_in": 3600,
"refresh_token": "tGzv3JOk...",
"scope": "openid profile email"
}"#;
let resp = TokenResponse::parse(json).unwrap();
assert_eq!(resp.access_token(), "eyJhbGciOi...");
assert_eq!(resp.token_type(), "Bearer");
assert_eq!(resp.expires_in(), Some(3600));
assert_eq!(resp.refresh_token(), Some("tGzv3JOk..."));
assert_eq!(resp.scope(), Some("openid profile email"));
}
#[test]
fn parse_response_with_zero_expires_in() {
let json = r#"{"access_token":"tok","token_type":"Bearer","expires_in":0}"#;
let resp = TokenResponse::parse(json).unwrap();
assert_eq!(resp.expires_in(), Some(0));
}
#[test]
fn parse_response_with_null_expires_in() {
let json = r#"{"access_token":"tok","token_type":"Bearer","expires_in":null}"#;
let resp = TokenResponse::parse(json).unwrap();
assert_eq!(resp.expires_in(), None);
}
#[test]
fn parse_missing_access_token() {
let json = r#"{"token_type":"Bearer"}"#;
let err = TokenResponse::parse(json).unwrap_err();
assert!(
err.to_string().contains("missing access_token"),
"error should mention missing access_token: {err}",
);
}
#[test]
fn parse_missing_token_type() {
let json = r#"{"access_token":"abc"}"#;
let err = TokenResponse::parse(json).unwrap_err();
assert!(
err.to_string().contains("missing token_type"),
"error should mention missing token_type: {err}",
);
}
#[test]
fn parse_invalid_expires_in() {
let json = r#"{"access_token":"tok","token_type":"Bearer","expires_in":"not-a-number"}"#;
let err = TokenResponse::parse(json).unwrap_err();
assert!(
err.to_string().contains("invalid expires_in"),
"error should mention invalid expires_in: {err}",
);
}
#[test]
fn parse_negative_expires_in() {
let json = r#"{"access_token":"tok","token_type":"Bearer","expires_in":-1}"#;
let err = TokenResponse::parse(json).unwrap_err();
assert!(
err.to_string().contains("invalid expires_in"),
"error should mention invalid expires_in: {err}",
);
}
#[test]
fn parse_fractional_expires_in() {
let json = r#"{"access_token":"tok","token_type":"Bearer","expires_in":3600.5}"#;
let err = TokenResponse::parse(json).unwrap_err();
assert!(
err.to_string().contains("invalid expires_in"),
"error should mention invalid expires_in: {err}",
);
}
#[test]
fn parse_oauth_error() {
let json = r#"{"error":"invalid_grant"}"#;
let err = TokenResponse::parse(json).unwrap_err();
assert!(
err.to_string().contains("invalid_grant"),
"error should contain the OAuth error code: {err}",
);
}
#[test]
fn parse_oauth_error_with_description() {
let json = r#"{"error":"invalid_grant","error_description":"The code has expired"}"#;
let err = TokenResponse::parse(json).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("invalid_grant"),
"error should contain the error code: {msg}",
);
assert!(
msg.contains("The code has expired"),
"error should contain the description: {msg}",
);
}
#[test]
fn parse_invalid_json() {
let err = TokenResponse::parse("not json at all").unwrap_err();
assert!(
err.to_string().contains("invalid JSON"),
"error should mention invalid JSON: {err}",
);
}
#[test]
fn parse_empty_string() {
let err = TokenResponse::parse("").unwrap_err();
assert!(
err.to_string().contains("invalid JSON"),
"error should mention invalid JSON: {err}",
);
}
#[test]
fn debug_redacts_tokens() {
let json = r#"{
"access_token": "secret-access-token",
"token_type": "Bearer",
"refresh_token": "secret-refresh-token"
}"#;
let resp = TokenResponse::parse(json).unwrap();
let debug_output = format!("{resp:?}");
assert!(
debug_output.contains("[REDACTED]"),
"debug should contain [REDACTED]: {debug_output}",
);
assert!(
!debug_output.contains("secret-access-token"),
"debug must not contain the access token",
);
assert!(
!debug_output.contains("secret-refresh-token"),
"debug must not contain the refresh token",
);
}
#[test]
fn error_display_messages() {
let err = TokenResponseError::new(TokenResponseErrorKind::InvalidJson);
assert_eq!(err.to_string(), "token response: invalid JSON");
let err = TokenResponseError::new(TokenResponseErrorKind::MissingAccessToken);
assert_eq!(err.to_string(), "token response: missing access_token");
let err = TokenResponseError::new(TokenResponseErrorKind::MissingTokenType);
assert_eq!(err.to_string(), "token response: missing token_type");
let err = TokenResponseError::new(TokenResponseErrorKind::InvalidExpiresIn);
assert_eq!(err.to_string(), "token response: invalid expires_in value");
}
#[test]
fn error_implements_std_error() {
let err: Box<dyn std::error::Error> =
Box::new(TokenResponseError::new(TokenResponseErrorKind::InvalidJson));
let _ = err.to_string();
}
}