use std::collections::HashSet;
use chrono::{DateTime, Utc};
use sqlx::{PgPool, Row};
use tracing::{debug, warn};
use uuid::Uuid;
use backbone_orm::org_scope;
use crate::application::service::integrations_oauth_ports::{
OAuthCredentialFailure, OAuthCredentialStore, TokenBundle, PURPOSE_OAUTH_TOKEN,
};
use crate::infrastructure::http::endpoint_guard::{
EndpointOverrides, OAuthClientConfigs, OAuthTransport, ProviderRegistry, TokenRequestForm,
TransportFailureKind, ValidatedEndpoints,
};
#[derive(Debug, Clone, PartialEq)]
pub struct RefreshSchedule {
pub refresh_window_seconds: i64,
pub refresh_batch_size: i64,
}
impl Default for RefreshSchedule {
fn default() -> Self {
Self { refresh_window_seconds: 600, refresh_batch_size: 100 }
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct RefreshReport {
pub refreshed: usize,
pub expired: usize,
pub skipped: usize,
}
struct DueAccount {
id: Uuid,
provider: String,
account_ref: String,
}
fn refresh_form(
client_id: &str,
client_secret: Option<&str>,
refresh_token: &str,
) -> TokenRequestForm {
TokenRequestForm {
grant_type: "refresh_token".into(),
code: None,
refresh_token: Some(refresh_token.to_string()),
redirect_uri: None,
code_verifier: None,
client_id: client_id.to_string(),
client_secret: client_secret.map(str::to_string),
scope: None,
}
}
pub async fn refresh_oauth_credentials(
pool: &PgPool,
company_id: Uuid,
registry: &ProviderRegistry,
overrides: &EndpointOverrides,
clients: &OAuthClientConfigs,
store: &dyn OAuthCredentialStore,
transport: &dyn OAuthTransport,
schedule: &RefreshSchedule,
) -> Result<RefreshReport, sqlx::Error> {
let mut report = RefreshReport::default();
let mut attempted: Vec<Uuid> = Vec::new();
for _ in 0..schedule.refresh_batch_size.max(0) {
let mut tx = pool.begin().await?;
if let Some(scope) = org_scope::current_org_scope() {
org_scope::bind_org_scope_on(&mut tx, &scope).await?;
}
let claimed = sqlx::query(
r#"SELECT id, provider::text AS provider, account_ref
FROM integrations.integration_accounts
WHERE status = 'active'
AND expires_at IS NOT NULL
AND expires_at < now() + make_interval(secs => $1)
AND id <> ALL($2::uuid[])
ORDER BY expires_at
LIMIT 1
FOR UPDATE SKIP LOCKED"#,
)
.bind(schedule.refresh_window_seconds)
.bind(&attempted)
.fetch_optional(&mut *tx)
.await?;
let row = match claimed {
Some(row) => row,
None => {
tx.rollback().await?;
break;
}
};
let account = DueAccount {
id: row.get("id"),
provider: row.get("provider"),
account_ref: row.get("account_ref"),
};
attempted.push(account.id);
refresh_one(
tx, company_id, registry, overrides, clients, store, transport, &account, &mut report,
)
.await?;
}
Ok(report)
}
pub async fn refresh_oauth_credentials_for_companies(
pool: &PgPool,
registry: &ProviderRegistry,
overrides: &EndpointOverrides,
clients: &OAuthClientConfigs,
store: &dyn OAuthCredentialStore,
transport: &dyn OAuthTransport,
schedule: &RefreshSchedule,
companies: &[Uuid],
) -> Vec<(Uuid, Result<RefreshReport, sqlx::Error>)> {
let mut out = Vec::with_capacity(companies.len());
for company_id in companies {
let r = refresh_oauth_credentials(
pool, *company_id, registry, overrides, clients, store, transport, schedule,
)
.await;
out.push((*company_id, r));
}
out
}
async fn refresh_one(
mut tx: sqlx::Transaction<'_, sqlx::Postgres>,
company_id: Uuid,
registry: &ProviderRegistry,
overrides: &EndpointOverrides,
clients: &OAuthClientConfigs,
store: &dyn OAuthCredentialStore,
transport: &dyn OAuthTransport,
account: &DueAccount,
report: &mut RefreshReport,
) -> Result<(), sqlx::Error> {
let endpoints: ValidatedEndpoints = match ValidatedEndpoints::resolve(
registry,
&account.provider,
overrides,
) {
Ok(endpoints) => endpoints,
Err(e) => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"endpoint guard refused the provider config; account left due: {e}"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
};
let client = match clients.get(&account.provider) {
Some(client) => client,
None => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"no OAuth client configured for provider; account left due"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
};
let bundle: TokenBundle = match store
.read_token(company_id, &account.provider, &account.account_ref)
.await
{
Ok(bundle) => bundle,
Err(e) if e.is_transport() => {
debug!(
target: "integrations.oauth.refresh",
account_id = %account.id,
"credential store unreachable; account left due ({e})"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
Err(OAuthCredentialFailure { code, message }) => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"credential unreadable ({code}: {message}); account moves to expired (reconnect required)"
);
expire_account(tx, account.id).await?;
report.expired += 1;
return Ok(());
}
};
let refresh_token = match bundle.refresh_token() {
Some(token) => token.to_string(),
None => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"credential carries no refresh token; account moves to expired (reconnect required)"
);
expire_account(tx, account.id).await?;
report.expired += 1;
return Ok(());
}
};
let form = refresh_form(&client.client_id, client.client_secret.as_deref(), &refresh_token);
let now = Utc::now();
let response = match transport.exchange(&endpoints.token, &form).await {
Ok(response) => response,
Err(e) if e.kind == TransportFailureKind::InvalidGrant => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"provider rejected the refresh grant (invalid_grant); account moves to expired (reconnect required)"
);
expire_account(tx, account.id).await?;
report.expired += 1;
return Ok(());
}
Err(e) => {
debug!(
target: "integrations.oauth.refresh",
account_id = %account.id,
"refresh exchange failed; account left due ({e})"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
};
let expires_at: DateTime<Utc> = match response.expires_at(now) {
Some(expires_at) => expires_at,
None => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"provider returned no expires_in; unstoreable response refused, account left due"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
};
let successor = TokenBundle::new(
response.access_token.clone(),
response
.refresh_token
.clone()
.or_else(|| bundle.refresh_token().map(str::to_string)),
expires_at,
response.scope.clone().or_else(|| bundle.scope().map(str::to_string)),
);
match store
.rotate(
company_id,
&account.provider,
&account.account_ref,
successor,
expires_at,
)
.await
{
Ok(_) => {}
Err(e) if e.is_transport() => {
debug!(
target: "integrations.oauth.refresh",
account_id = %account.id,
"credential store unreachable at rotate; account left due ({e})"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
Err(OAuthCredentialFailure { code, message }) => {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
"rotate refused ({code}: {message}); account moves to expired"
);
expire_account(tx, account.id).await?;
report.expired += 1;
return Ok(());
}
}
let mirrored = sqlx::query(
r#"UPDATE integrations.integration_accounts
SET expires_at = $2,
last_refreshed_at = now()
WHERE id = $1"#,
)
.bind(account.id)
.bind(expires_at)
.execute(&mut *tx)
.await?;
if mirrored.rows_affected() != 1 {
warn!(
target: "integrations.oauth.refresh",
account_id = %account.id,
"account mirror affected no rows; rolled back"
);
report.skipped += 1;
tx.rollback().await?;
return Ok(());
}
tx.commit().await?;
report.refreshed += 1;
debug!(
target: "integrations.oauth.refresh",
account_id = %account.id,
provider = %account.provider,
purpose = PURPOSE_OAUTH_TOKEN,
"refreshed before expiry; successor mirrored"
);
Ok(())
}
async fn expire_account(
mut tx: sqlx::Transaction<'_, sqlx::Postgres>,
account_id: Uuid,
) -> Result<(), sqlx::Error> {
sqlx::query(
r#"UPDATE integrations.integration_accounts
SET status = 'expired'
WHERE id = $1"#,
)
.bind(account_id)
.execute(&mut *tx)
.await?;
tx.commit().await?;
Ok(())
}