use base64::{Engine, engine::general_purpose::STANDARD};
use http::HeaderValue;
use std::{sync::Arc, time::Duration};
use serde::{Deserialize, Serialize};
use volga_oauth_core::{AuthorizationServerMetadata, OAuthErrorCode};
use crate::{
ClientConfig, ClientError, Pkce, TokenResponse, TokenSet, TokenStore,
pkce::{PKCE_METHOD, random_urlsafe},
transport::Transport,
};
const EXPIRY_LEEWAY: Duration = Duration::from_secs(30);
const TOKEN_STORE_NOT_CONFIGURED: &str =
"OAuth client: token store is not configured; attach one with with_token_store(..)";
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[non_exhaustive]
pub enum ClientAuthMethod {
#[default]
Basic,
Post,
}
pub struct OAuthClient {
transport: Transport,
client_id: String,
client_secret: Option<String>,
auth_method: ClientAuthMethod,
redirect_uri: Option<String>,
store: Option<Arc<dyn TokenStore>>,
}
impl OAuthClient {
pub fn new(client_id: impl Into<String>) -> Self {
Self {
transport: Transport::new(ClientConfig::new()),
client_id: client_id.into(),
client_secret: None,
auth_method: ClientAuthMethod::default(),
redirect_uri: None,
store: None,
}
}
pub fn from_registration(
response: &volga_oauth_core::ClientRegistrationResponse,
) -> Result<Self, ClientError> {
let mut client = Self::new(response.client_id.clone());
match response
.metadata
.token_endpoint_auth_method
.as_deref()
.unwrap_or("client_secret_basic")
{
"none" => {}
"client_secret_basic" => {
if let Some(secret) = &response.client_secret {
client = client.with_secret(secret.clone());
}
}
"client_secret_post" => {
if let Some(secret) = &response.client_secret {
client = client
.with_secret(secret.clone())
.with_auth_method(ClientAuthMethod::Post);
}
}
unsupported => {
return Err(ClientError::validation(format!(
"registered token_endpoint_auth_method '{unsupported}' is not supported; \
this client supports client_secret_basic, client_secret_post and none"
)));
}
}
if let [redirect_uri] = response.metadata.redirect_uris.as_slice() {
client = client.with_redirect_uri(redirect_uri.clone());
}
Ok(client)
}
pub fn with_config(mut self, config: ClientConfig) -> Self {
self.transport = Transport::new(config);
self
}
pub fn with_secret(mut self, client_secret: impl Into<String>) -> Self {
self.client_secret = Some(client_secret.into());
self
}
pub fn with_auth_method(mut self, method: ClientAuthMethod) -> Self {
self.auth_method = method;
self
}
pub fn with_redirect_uri(mut self, redirect_uri: impl Into<String>) -> Self {
self.redirect_uri = Some(redirect_uri.into());
self
}
pub fn with_token_store(mut self, store: Arc<dyn TokenStore>) -> Self {
self.store = Some(store);
self
}
pub fn authorization_request<'a>(
&'a self,
metadata: &'a AuthorizationServerMetadata,
) -> AuthorizationRequestBuilder<'a> {
AuthorizationRequestBuilder {
client: self,
metadata,
scopes: Vec::new(),
resources: Vec::new(),
state: None,
extra: Vec::new(),
}
}
pub async fn exchange_code(
&self,
metadata: &AuthorizationServerMetadata,
code: &str,
request: &AuthorizationRequest,
) -> Result<TokenSet, ClientError> {
let endpoint = token_endpoint(metadata)?;
let (body, authorization) = {
let mut form = form_urlencoded::Serializer::new(String::new());
form.append_pair("grant_type", "authorization_code")
.append_pair("code", code)
.append_pair("code_verifier", request.pkce.verifier());
if let Some(redirect_uri) = &self.redirect_uri {
form.append_pair("redirect_uri", redirect_uri);
}
for resource in &request.resources {
form.append_pair("resource", resource);
}
let authorization = self.apply_client_auth(&mut form);
(form.finish(), authorization)
};
self.request_tokens(endpoint, body, authorization).await
}
pub async fn refresh(
&self,
metadata: &AuthorizationServerMetadata,
refresh_token: &str,
) -> Result<TokenSet, ClientError> {
let endpoint = token_endpoint(metadata)?;
let (body, authorization) = {
let mut form = form_urlencoded::Serializer::new(String::new());
form.append_pair("grant_type", "refresh_token")
.append_pair("refresh_token", refresh_token);
let authorization = self.apply_client_auth(&mut form);
(form.finish(), authorization)
};
self.request_tokens(endpoint, body, authorization).await
}
pub async fn token(
&self,
key: &str,
metadata: &AuthorizationServerMetadata,
) -> Result<Option<TokenSet>, ClientError> {
let store = self.store.as_deref().expect(TOKEN_STORE_NOT_CONFIGURED);
let Some(tokens) = store.get(key) else {
return Ok(None);
};
if !tokens.expires_within(EXPIRY_LEEWAY) {
return Ok(Some(tokens));
}
let Some(refresh_token) = tokens.refresh_token else {
store.remove(key);
return Ok(None);
};
match self.refresh(metadata, &refresh_token).await {
Ok(mut fresh) => {
if fresh.refresh_token.is_none() {
fresh.refresh_token = Some(refresh_token);
}
store.put(key, &fresh);
Ok(Some(fresh))
}
Err(ClientError::Protocol(err)) if err.error == OAuthErrorCode::InvalidGrant => {
store.remove(key);
Ok(None)
}
Err(err) => Err(err),
}
}
pub fn store_tokens(&self, key: &str, tokens: &TokenSet) {
self.store
.as_deref()
.expect(TOKEN_STORE_NOT_CONFIGURED)
.put(key, tokens);
}
async fn request_tokens(
&self,
endpoint: &str,
body: String,
authorization: Option<HeaderValue>,
) -> Result<TokenSet, ClientError> {
let value = self
.transport
.post_form(endpoint, body, authorization)
.await?;
let response: TokenResponse = serde_json::from_value(value)?;
Ok(response.into())
}
fn apply_client_auth(
&self,
form: &mut form_urlencoded::Serializer<'_, String>,
) -> Option<HeaderValue> {
match (&self.client_secret, self.auth_method) {
(Some(secret), ClientAuthMethod::Basic) => {
Some(basic_credentials(&self.client_id, secret))
}
(Some(secret), ClientAuthMethod::Post) => {
form.append_pair("client_id", &self.client_id)
.append_pair("client_secret", secret);
None
}
(None, _) => {
form.append_pair("client_id", &self.client_id);
None
}
}
}
}
impl std::fmt::Debug for OAuthClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OAuthClient")
.field("transport", &self.transport)
.field("client_id", &self.client_id)
.field(
"client_secret",
&self.client_secret.as_ref().map(|_| "[redacted]"),
)
.field("auth_method", &self.auth_method)
.field("redirect_uri", &self.redirect_uri)
.field("store", &self.store.as_ref().map(|_| "dyn TokenStore"))
.finish()
}
}
pub struct AuthorizationRequestBuilder<'a> {
client: &'a OAuthClient,
metadata: &'a AuthorizationServerMetadata,
scopes: Vec<String>,
resources: Vec<String>,
state: Option<String>,
extra: Vec<(String, String)>,
}
impl AuthorizationRequestBuilder<'_> {
pub fn with_scopes<I, S>(mut self, scopes: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.scopes = scopes.into_iter().map(Into::into).collect();
self
}
pub fn with_resource(mut self, resource: impl Into<String>) -> Self {
self.resources.push(resource.into());
self
}
pub fn with_state(mut self, state: impl Into<String>) -> Self {
self.state = Some(state.into());
self
}
pub fn with_param(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
self.extra.push((name.into(), value.into()));
self
}
pub fn build(self) -> Result<AuthorizationRequest, ClientError> {
let endpoint = self
.metadata
.authorization_endpoint
.as_deref()
.ok_or_else(|| {
ClientError::validation("server metadata declares no authorization_endpoint")
})?;
self.client.transport.check_scheme(endpoint)?;
let methods = &self.metadata.code_challenge_methods_supported;
if !methods.is_empty() && !methods.iter().any(|method| method == PKCE_METHOD) {
return Err(ClientError::validation(format!(
"authorization server does not support the {PKCE_METHOD} PKCE method"
)));
}
let pkce = Pkce::new();
let state = self.state.unwrap_or_else(|| random_urlsafe(16));
let mut query = form_urlencoded::Serializer::new(String::new());
query
.append_pair("response_type", "code")
.append_pair("client_id", &self.client.client_id)
.append_pair("state", &state)
.append_pair("code_challenge", pkce.challenge())
.append_pair("code_challenge_method", PKCE_METHOD);
if let Some(redirect_uri) = &self.client.redirect_uri {
query.append_pair("redirect_uri", redirect_uri);
}
if !self.scopes.is_empty() {
query.append_pair("scope", &self.scopes.join(" "));
}
for resource in &self.resources {
query.append_pair("resource", resource);
}
for (name, value) in &self.extra {
query.append_pair(name, value);
}
let query = query.finish();
let separator = if endpoint.contains('?') { '&' } else { '?' };
Ok(AuthorizationRequest {
url: format!("{endpoint}{separator}{query}"),
state,
pkce,
resources: self.resources,
})
}
}
impl std::fmt::Debug for AuthorizationRequestBuilder<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AuthorizationRequestBuilder")
.field("scopes", &self.scopes)
.field("resources", &self.resources)
.field("state", &self.state)
.field("extra", &self.extra)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct AuthorizationRequest {
pub url: String,
pub state: String,
pub pkce: Pkce,
pub resources: Vec<String>,
}
impl AuthorizationRequest {
#[inline]
pub fn matches_state(&self, state: &str) -> bool {
self.state == state
}
pub fn validate_callback(
&self,
metadata: &AuthorizationServerMetadata,
state: &str,
iss: Option<&str>,
) -> Result<(), ClientError> {
if !self.matches_state(state) {
return Err(ClientError::validation(
"authorization response `state` does not match the request",
));
}
match (iss, metadata.authorization_response_iss_parameter_supported) {
(Some(iss), _) if iss != metadata.issuer => Err(ClientError::validation(format!(
"authorization response `iss` mismatch: expected {}, got {iss}",
metadata.issuer
))),
(None, true) => Err(ClientError::validation(
"authorization server advertises RFC 9207 but the response carries no `iss`",
)),
_ => Ok(()),
}
}
}
fn basic_credentials(client_id: &str, client_secret: &str) -> HeaderValue {
let encode =
|value: &str| -> String { form_urlencoded::byte_serialize(value.as_bytes()).collect() };
let credentials = STANDARD.encode(format!("{}:{}", encode(client_id), encode(client_secret)));
HeaderValue::from_str(&format!("Basic {credentials}"))
.expect("base64 output is always a valid header value")
}
fn token_endpoint(metadata: &AuthorizationServerMetadata) -> Result<&str, ClientError> {
metadata
.token_endpoint
.as_deref()
.ok_or_else(|| ClientError::validation("server metadata declares no token_endpoint"))
}
#[cfg(test)]
mod tests {
use super::*;
fn metadata() -> AuthorizationServerMetadata {
let mut metadata = AuthorizationServerMetadata::new("https://auth.example.com");
metadata.authorization_endpoint = Some("https://auth.example.com/authorize".into());
metadata.token_endpoint = Some("https://auth.example.com/token".into());
metadata
}
fn query_pairs(url: &str) -> Vec<(String, String)> {
let query = url.split_once('?').unwrap().1;
form_urlencoded::parse(query.as_bytes())
.into_owned()
.collect()
}
#[test]
fn it_builds_a_spec_compliant_authorization_url() {
let client =
OAuthClient::new("my-client").with_redirect_uri("https://app.example.com/callback");
let request = client
.authorization_request(&metadata())
.with_scopes(["read", "write"])
.with_resource("https://api.example.com")
.with_param("nonce", "n-1")
.build()
.unwrap();
assert!(
request
.url
.starts_with("https://auth.example.com/authorize?")
);
let pairs = query_pairs(&request.url);
let get = |name: &str| {
pairs
.iter()
.find(|(key, _)| key == name)
.map(|(_, value)| value.as_str())
};
assert_eq!(get("response_type"), Some("code"));
assert_eq!(get("client_id"), Some("my-client"));
assert_eq!(
get("redirect_uri"),
Some("https://app.example.com/callback")
);
assert_eq!(get("scope"), Some("read write"));
assert_eq!(get("resource"), Some("https://api.example.com"));
assert_eq!(get("code_challenge"), Some(request.pkce.challenge()));
assert_eq!(get("code_challenge_method"), Some("S256"));
assert_eq!(get("state"), Some(request.state.as_str()));
assert_eq!(get("nonce"), Some("n-1"));
assert!(request.matches_state(&request.state.clone()));
assert!(!request.matches_state("other"));
}
#[test]
fn it_appends_to_an_existing_query_and_respects_custom_state() {
let mut metadata = metadata();
metadata.authorization_endpoint =
Some("https://auth.example.com/authorize?tenant=t1".into());
let request = OAuthClient::new("my-client")
.authorization_request(&metadata)
.with_state("custom-state")
.build()
.unwrap();
assert!(request.url.contains("tenant=t1&response_type=code"));
assert_eq!(request.state, "custom-state");
}
#[test]
fn it_validates_metadata_before_building_requests() {
let client = OAuthClient::new("my-client");
let mut incomplete = metadata();
incomplete.authorization_endpoint = None;
assert!(matches!(
client.authorization_request(&incomplete).build(),
Err(ClientError::Validation(reason)) if reason.contains("authorization_endpoint")
));
let mut plain_only = metadata();
plain_only.code_challenge_methods_supported = vec!["plain".into()];
assert!(matches!(
client.authorization_request(&plain_only).build(),
Err(ClientError::Validation(reason)) if reason.contains("S256")
));
let mut insecure = metadata();
insecure.authorization_endpoint = Some("http://auth.example.com/authorize".into());
assert!(matches!(
client.authorization_request(&insecure).build(),
Err(ClientError::InsecureUrl(_))
));
}
#[test]
fn it_encodes_basic_credentials_per_rfc6749() {
let header = basic_credentials("client with space", "s&cret");
let encoded = header
.to_str()
.unwrap()
.strip_prefix("Basic ")
.unwrap()
.to_owned();
let decoded = String::from_utf8(STANDARD.decode(encoded).unwrap()).unwrap();
assert_eq!(decoded, "client+with+space:s%26cret");
}
#[test]
fn it_applies_the_configured_client_authentication() {
let public = OAuthClient::new("my-client");
let mut form = form_urlencoded::Serializer::new(String::new());
assert!(public.apply_client_auth(&mut form).is_none());
assert_eq!(form.finish(), "client_id=my-client");
let basic = OAuthClient::new("my-client").with_secret("s3cret");
let mut form = form_urlencoded::Serializer::new(String::new());
assert!(basic.apply_client_auth(&mut form).is_some());
assert_eq!(form.finish(), "");
let post = OAuthClient::new("my-client")
.with_secret("s3cret")
.with_auth_method(ClientAuthMethod::Post);
let mut form = form_urlencoded::Serializer::new(String::new());
assert!(post.apply_client_auth(&mut form).is_none());
assert_eq!(form.finish(), "client_id=my-client&client_secret=s3cret");
}
#[test]
fn it_adopts_registered_credentials_per_auth_method() {
let registration = |auth_method: serde_json::Value| {
serde_json::from_value::<volga_oauth_core::ClientRegistrationResponse>(
serde_json::json!({
"client_id": "generated-id",
"client_secret": "generated-secret",
"token_endpoint_auth_method": auth_method,
"redirect_uris": ["https://app.example.com/callback"]
}),
)
.unwrap()
};
let client =
OAuthClient::from_registration(®istration(serde_json::Value::Null)).unwrap();
assert_eq!(client.auth_method, ClientAuthMethod::Basic);
assert_eq!(client.client_secret.as_deref(), Some("generated-secret"));
assert_eq!(
client.redirect_uri.as_deref(),
Some("https://app.example.com/callback")
);
let client =
OAuthClient::from_registration(®istration("client_secret_post".into())).unwrap();
assert_eq!(client.auth_method, ClientAuthMethod::Post);
assert_eq!(client.client_secret.as_deref(), Some("generated-secret"));
let client = OAuthClient::from_registration(®istration("none".into())).unwrap();
assert_eq!(client.client_secret, None);
let err =
OAuthClient::from_registration(®istration("client_secret_jwt".into())).unwrap_err();
assert!(matches!(
err,
ClientError::Validation(reason) if reason.contains("client_secret_jwt")
));
}
#[test]
fn it_returns_send_token_endpoint_futures() {
fn assert_send(_: impl Send) {}
let client = OAuthClient::new("my-client")
.with_token_store(Arc::new(crate::InMemoryTokenStore::default()));
let metadata = metadata();
let request = client.authorization_request(&metadata).build().unwrap();
assert_send(client.exchange_code(&metadata, "the-code", &request));
assert_send(client.refresh(&metadata, "the-refresh-token"));
assert_send(client.token("alice", &metadata));
}
#[test]
fn it_validates_the_authorization_callback() {
let client = OAuthClient::new("my-client");
let metadata = metadata();
let request = client.authorization_request(&metadata).build().unwrap();
let state = request.state.clone();
assert!(request.validate_callback(&metadata, &state, None).is_ok());
assert!(
request
.validate_callback(&metadata, &state, Some(&metadata.issuer))
.is_ok()
);
let err = request
.validate_callback(&metadata, "forged", None)
.unwrap_err();
assert!(matches!(err, ClientError::Validation(reason) if reason.contains("state")));
let err = request
.validate_callback(&metadata, &state, Some("https://evil.example.com"))
.unwrap_err();
assert!(
matches!(err, ClientError::Validation(reason) if reason.contains("`iss` mismatch"))
);
let advertised = metadata.with_authorization_response_iss_parameter(true);
assert!(
request
.validate_callback(&advertised, &state, Some(&advertised.issuer))
.is_ok()
);
let err = request
.validate_callback(&advertised, &state, None)
.unwrap_err();
assert!(matches!(err, ClientError::Validation(reason) if reason.contains("RFC 9207")));
}
#[test]
#[should_panic(expected = "token store is not configured")]
fn it_panics_on_store_access_without_a_store() {
OAuthClient::new("my-client").store_tokens(
"alice",
&TokenSet {
access_token: "at".into(),
token_type: "Bearer".into(),
refresh_token: None,
scope: None,
id_token: None,
expires_at: None,
},
);
}
}