use super::oauth::{is_expired, parse_token_response, Grant};
use super::wire::build_token_exchange_request;
use super::OAuthConfig;
use super::{auth_error, fetch_cred, require_header, set_auth_header, Auth, AuthCtx, CredSource};
use crate::canonical::{CanonicalError, ErrorKind};
use crate::protocol::{ProviderCtx, WireRequest};
use crate::store::{Clock, Cred, CredStore, Secret};
use crate::transport::{Transport, TransportResponse};
type Bearer = (Secret, Option<String>);
struct Prior {
refresh_token: Secret,
scope: Option<String>,
account_id: Option<String>,
}
pub struct OAuth2Auth;
impl Auth for OAuth2Auth {
fn apply(
&self,
wire: &mut WireRequest,
_ctx: &ProviderCtx,
auth: &AuthCtx,
store: &dyn CredStore,
clock: &dyn Clock,
transport: &dyn Transport,
) -> Result<(), CanonicalError> {
let cfg = auth.oauth.ok_or_else(oauth_row_misconfigured)?;
let (token, account_id) = bearer(wire, auth, store, clock, transport, cfg)?;
set_auth_header(wire, require_header(auth)?, &token);
for (name, value) in &cfg.beta_headers {
wire.set_header(name, value);
}
if let (Some(name), Some(id)) = (cfg.account_header.as_deref(), account_id.as_deref()) {
wire.set_header(name, id);
}
Ok(())
}
}
fn bearer(
wire: &WireRequest,
auth: &AuthCtx,
store: &dyn CredStore,
clock: &dyn Clock,
transport: &dyn Transport,
cfg: &OAuthConfig,
) -> Result<Bearer, CanonicalError> {
let Some(fetched) = fetch_cred(store, auth) else {
return Err(not_logged_in(auth.store_key));
};
let Cred::OAuth2 {
access_token,
refresh_token,
expires_at,
scope,
account_id,
} = fetched.cred
else {
return Err(not_logged_in(auth.store_key));
};
if !is_expired(expires_at, clock.now()) {
return Ok((access_token, account_id));
}
match fetched.source {
CredSource::Borrowed(path) => Err(borrowed_expired(auth.store_key, &path)),
CredSource::Owned => {
let prior = Prior {
refresh_token,
scope,
account_id,
};
refresh(wire, auth, store, clock, transport, cfg, &prior)
}
}
}
fn refresh(
wire: &WireRequest,
auth: &AuthCtx,
store: &dyn CredStore,
clock: &dyn Clock,
transport: &dyn Transport,
cfg: &OAuthConfig,
prior: &Prior,
) -> Result<Bearer, CanonicalError> {
let mut req = build_token_exchange_request(
cfg,
Grant::Refresh {
refresh_token: &prior.refresh_token,
},
);
req.timeouts = wire.timeouts;
req.exec = wire.exec.clone();
let bytes = collect_body(transport.send(req)?)?;
let Ok(fresh) = parse_token_response(&bytes, clock.now()) else {
return spent(store, auth, clock);
};
store
.put(
auth.store_key,
&fresh.as_cred(&prior.refresh_token, &prior.scope, &prior.account_id),
)
.map_err(persist_failed)?;
Ok((fresh.access_token, prior.account_id.clone()))
}
fn spent(
store: &dyn CredStore,
auth: &AuthCtx,
clock: &dyn Clock,
) -> Result<Bearer, CanonicalError> {
let row = auth.store_key;
let Some(spec) = auth.ambient else {
return Err(auth_error(&format!(
"token refresh failed for provider `{row}`; re-run `bz --login --provider {row}` \
if this persists"
)));
};
if let Some(Cred::OAuth2 {
access_token,
expires_at,
account_id,
..
}) = store.discover(spec)
{
if !is_expired(expires_at, clock.now()) {
return Ok((access_token, account_id));
}
}
Err(auth_error(&format!(
"token refresh failed for provider `{row}`, and the ambient credential at {path} is \
absent or expired; sign in with the tool that owns that file, or re-run \
`bz --login --provider {row}`",
path = spec.path,
)))
}
fn not_logged_in(row: &str) -> CanonicalError {
auth_error(&format!(
"not logged in for provider `{row}`: run `bz --login --provider {row}` (or sign in \
to a tool whose ambient credential that row discovers)"
))
}
fn borrowed_expired(row: &str, path: &str) -> CanonicalError {
auth_error(&format!(
"provider `{row}`: the ambient credential at {path} is expired; refresh it with the \
tool that owns that file, or run `bz --login --provider {row}` to hold one of \
brazen's own"
))
}
pub(crate) fn collect_body(resp: TransportResponse) -> Result<Vec<u8>, CanonicalError> {
let mut out = Vec::new();
for chunk in resp.body {
let bytes = chunk.map_err(|e| CanonicalError {
kind: ErrorKind::Transport,
message: format!("transport error reading token response: {e}"),
provider_detail: None,
retry_after_seconds: None,
})?;
out.extend_from_slice(&bytes);
}
Ok(out)
}
fn oauth_row_misconfigured() -> CanonicalError {
CanonicalError {
kind: ErrorKind::Config,
message: "oauth2 provider row has no oauth config (should be caught at resolve)".to_owned(),
provider_detail: None,
retry_after_seconds: None,
}
}
fn persist_failed(e: std::io::Error) -> CanonicalError {
CanonicalError {
kind: ErrorKind::Auth,
message: format!("could not persist refreshed credential: {e}"),
provider_detail: None,
retry_after_seconds: None,
}
}