use asupersync::Cx;
use super::{
AsyncOAuthRefreshStore, ManagedOAuthSession, OAuthError, OAuthRefreshRenewal,
OAuthRefreshRenewalCustody, SlotRevision, within,
};
pub use crate::http_auth::managed::logout::rotation::OAuthAccessRotationError;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct OAuthAccessInstallation {
session_generation: u64,
stored_revision: SlotRevision,
}
impl OAuthAccessInstallation {
pub fn session_generation(self) -> u64 {
self.session_generation
}
pub fn stored_revision(self) -> SlotRevision {
self.stored_revision
}
}
impl<A, P> OAuthRefreshRenewal<A, P> {
pub async fn install_managed_access(
&mut self,
observer: &Cx,
session: &ManagedOAuthSession,
expected_generation: u64,
) -> Result<(AsyncOAuthRefreshStore<A, P>, OAuthAccessInstallation), OAuthAccessRotationError>
{
self.check_installation_context(observer)?;
let OAuthRefreshRenewalCustody::Complete { credentials, .. } = &self.custody else {
return Err(OAuthAccessRotationError::NotComplete);
};
let origin = self.origin.clone();
let cancellation = self.cancellation.clone();
let deadline = self.deadline;
let observer_deadline = observer
.now()
.saturating_add_nanos(deadline.as_nanos().saturating_sub(origin.now().as_nanos()));
let reservation = within(observer, observer_deadline, async {
Ok(within(&origin, deadline, async {
Ok(session
.reserve_access_rotation(
observer,
&cancellation,
expected_generation,
credentials,
)
.await)
})
.await)
})
.await
.map_err(OAuthAccessRotationError::Context)?
.map_err(OAuthAccessRotationError::Context)??;
self.check_installation_context(observer)?;
let previous = std::mem::replace(
&mut self.custody,
OAuthRefreshRenewalCustody::Stopped {
store: None,
credentials: None,
},
);
let OAuthRefreshRenewalCustody::Complete {
store,
credentials,
revision,
} = previous
else {
self.custody = previous;
return Err(OAuthAccessRotationError::NotComplete);
};
match reservation.commit(credentials) {
Ok(session_generation) => Ok((
store,
OAuthAccessInstallation {
session_generation,
stored_revision: revision,
},
)),
Err((error, credentials)) => {
self.custody = OAuthRefreshRenewalCustody::Complete {
store,
credentials: *credentials,
revision,
};
Err(error)
}
}
}
fn check_installation_context(&self, observer: &Cx) -> Result<(), OAuthAccessRotationError> {
if self.cancellation.is_cancel_requested()
|| self.origin.checkpoint().is_err()
|| observer.checkpoint().is_err()
{
return Err(OAuthAccessRotationError::Context(OAuthError::Cancelled));
}
if self.origin.now() >= self.deadline
|| self
.origin
.budget()
.deadline
.is_some_and(|end| self.origin.now() >= end)
|| observer
.budget()
.deadline
.is_some_and(|end| observer.now() >= end)
{
return Err(OAuthAccessRotationError::Context(OAuthError::TimedOut));
}
Ok(())
}
}
#[cfg(test)]
mod tests;