use super::oauth::{is_expired, parse_token_response, Grant};
use super::wire::build_token_exchange_request;
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};
use crate::transport::{Transport, TransportResponse};
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 Some(fetched) = fetch_cred(store, auth) else {
return Err(auth_error(
"not logged in for this provider: run `bz --login --provider <id>` (or \
sign in to a tool whose ambient credential this row discovers)",
));
};
let borrowed = fetched.source == CredSource::Borrowed;
let Cred::OAuth2 {
access_token,
refresh_token,
expires_at,
scope,
account_id,
} = fetched.cred
else {
return Err(auth_error(
"not logged in for this provider: run `bz --login --provider <id>` (or \
sign in to a tool whose ambient credential this row discovers)",
));
};
let token = if is_expired(expires_at, clock.now()) {
if borrowed {
return Err(auth_error(
"borrowed OAuth credential is expired; refresh it with the tool that owns it",
));
}
let mut req = build_token_exchange_request(
cfg,
Grant::Refresh {
refresh_token: &refresh_token,
},
);
req.timeouts = wire.timeouts;
req.exec = wire.exec.clone();
let bytes = collect_body(transport.send(req)?)?;
let fresh = parse_token_response(&bytes, clock.now()).map_err(|_| {
auth_error(
"token refresh failed; re-run `bz --login --provider <id>` if this persists",
)
})?;
store
.put(
auth.store_key,
&fresh.as_cred(&refresh_token, &scope, &account_id),
)
.map_err(persist_failed)?;
fresh.access_token
} else {
access_token
};
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(())
}
}
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,
}
}