use ares_store::oauth_credentials::{OAuthCredential, OAuthCredentialStore};
use ares_store::{decrypt_api_key, MasterKey};
use ares_types::types::{AppError, Result};
use serde::Deserialize;
use sqlx::PgPool;
use std::sync::Arc;
use std::time::Duration;
pub mod google;
pub mod hubspot;
pub mod linkedin;
pub mod salesforce;
pub mod slack;
#[derive(Debug, thiserror::Error)]
pub enum ConnectorError {
#[error("Auth failed for {provider}: {message}")]
Auth { provider: String, message: String },
#[error("HTTP error for {provider}: {status} — {message}")]
Http {
provider: String,
status: u16,
message: String,
},
#[error("Rate limited by {provider}: {message}")]
RateLimited { provider: String, message: String },
#[error("Configuration error for {provider}: {message}")]
Config { provider: String, message: String },
}
impl From<ConnectorError> for AppError {
fn from(e: ConnectorError) -> Self {
match e {
ConnectorError::Auth { provider, message } => {
AppError::Auth(format!("{provider} auth failed: {message}"))
}
ConnectorError::Http {
provider,
status,
message,
} => AppError::External(format!("{provider} HTTP {status}: {message}")),
ConnectorError::RateLimited { provider, message } => {
AppError::RateLimited(format!("{provider}: {message}"))
}
ConnectorError::Config { provider, message } => {
AppError::Configuration(format!("{provider}: {message}"))
}
}
}
}
pub async fn get_access_token(
pool: &PgPool,
master_key: &MasterKey,
tenant_id: &str,
provider: &str,
connector_type: &str,
) -> Result<String> {
let store = OAuthCredentialStore::new(pool);
let cred = store
.get(tenant_id, provider, connector_type)
.await?
.ok_or_else(|| {
AppError::Auth(format!(
"No OAuth credential found for tenant '{tenant_id}' and provider '{provider}'"
))
})?;
let now = chrono::Utc::now().timestamp();
if cred.expires_at.map(|e| e <= now).unwrap_or(false) {
return Err(AppError::Auth(format!(
"OAuth token for {provider} has expired (expires_at={expires_at:?}, now={now}). Re-authentication required.",
provider = provider,
expires_at = cred.expires_at,
now = now
)));
}
let at_payload = cred.access_token.ok_or_else(|| {
AppError::Auth(format!(
"OAuth credential for {provider} has no access token"
))
})?;
let token = decrypt_api_key(&at_payload, master_key)
.map_err(|e| AppError::Auth(format!("Failed to decrypt access token: {e}")))?;
Ok(token)
}
pub async fn get_oauth_credential(
pool: &PgPool,
_master_key: &MasterKey,
tenant_id: &str,
provider: &str,
connector_type: &str,
) -> Result<OAuthCredential> {
let store = OAuthCredentialStore::new(pool);
store
.get(tenant_id, provider, connector_type)
.await?
.ok_or_else(|| {
AppError::Auth(format!(
"No OAuth credential found for tenant '{tenant_id}' and provider '{provider}'"
))
})
}
pub const MAX_RETRIES: u32 = 3;
pub const BASE_RETRY_DELAY_MS: u64 = 1000;
pub async fn execute_with_retry(
_client: &reqwest::Client,
request: reqwest::RequestBuilder,
provider: &str,
) -> std::result::Result<reqwest::Response, ConnectorError> {
let mut delay = Duration::from_millis(BASE_RETRY_DELAY_MS);
let mut attempt = 0;
loop {
let response = request
.try_clone()
.ok_or_else(|| ConnectorError::Config {
provider: provider.to_string(),
message: "request body is a stream and cannot be cloned for retry".to_string(),
})?
.send()
.await
.map_err(|e| ConnectorError::Http {
provider: provider.to_string(),
status: 0,
message: format!("reqwest error: {e}"),
})?;
let status = response.status().as_u16();
if response.status().is_success() {
return Ok(response);
}
if response.status() == reqwest::StatusCode::TOO_MANY_REQUESTS && attempt < MAX_RETRIES {
let sleep_dur = response
.headers()
.get("retry-after")
.and_then(|h| h.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
.map(Duration::from_secs)
.unwrap_or(delay);
tokio::time::sleep(sleep_dur).await;
delay *= 2;
attempt += 1;
continue;
}
let body = response
.text()
.await
.unwrap_or_else(|_| "<unreadable body>".to_string());
return Err(ConnectorError::Http {
provider: provider.to_string(),
status,
message: body,
});
}
}
#[derive(Debug, Clone)]
pub struct ConnectorConfig {
pub base_url: String,
pub version: String,
}
#[derive(Debug, Deserialize)]
pub struct ApiErrorBody {
pub error: Option<ApiErrorDetail>,
pub message: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct ApiErrorDetail {
pub message: Option<String>,
pub code: Option<String>,
}
pub fn extract_error_message(body: &str) -> String {
serde_json::from_str::<ApiErrorBody>(body)
.ok()
.and_then(|e| e.error.and_then(|d| d.message).or(e.message))
.unwrap_or_else(|| body.to_string())
}
pub async fn refresh_oauth2_token(
pool: &PgPool,
master_key: &MasterKey,
tenant_id: &str,
provider: &str,
connector_type: &str,
token_url: &str,
) -> Result<String> {
let cred = get_oauth_credential(pool, master_key, tenant_id, provider, connector_type).await?;
let refresh_token = cred
.refresh_token
.as_ref()
.ok_or_else(|| AppError::Auth(format!("{provider} credential has no refresh token")))?;
let refresh_token_plain = decrypt_api_key(refresh_token, master_key)
.map_err(|e| AppError::Auth(format!("Failed to decrypt refresh token: {e}")))?;
let client_secret_plain = decrypt_api_key(&cred.client_secret, master_key)
.map_err(|e| AppError::Auth(format!("Failed to decrypt client secret: {e}")))?;
let client = reqwest::Client::new();
let response = client
.post(token_url)
.form(&[
("grant_type", "refresh_token"),
("refresh_token", &refresh_token_plain),
("client_id", &cred.client_id),
("client_secret", &client_secret_plain),
])
.send()
.await
.map_err(|e| AppError::External(format!("{provider} refresh token request failed: {e}")))?;
if !response.status().is_success() {
let body = response.text().await.unwrap_or_default();
return Err(AppError::Auth(format!(
"{provider} token refresh failed: {body}"
)));
}
#[derive(Debug, Deserialize)]
struct RefreshResponse {
access_token: String,
expires_in: i64,
#[serde(default)]
refresh_token: Option<String>,
}
let refresh_data: RefreshResponse = response
.json()
.await
.map_err(|e| AppError::External(format!("{provider} refresh token parse failed: {e}")))?;
let now = chrono::Utc::now().timestamp();
let expires_at = now + refresh_data.expires_in;
let new_refresh_token = refresh_data
.refresh_token
.as_deref()
.unwrap_or(&refresh_token_plain);
let store = OAuthCredentialStore::new(pool);
store
.update_tokens(
&cred.id,
&refresh_data.access_token,
Some(new_refresh_token),
expires_at,
)
.await?;
Ok(refresh_data.access_token)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum AccessTokenPlan {
UseStored,
Refresh,
ExpiredWithoutRefresh,
}
fn access_token_plan(
expires_at: Option<i64>,
has_refresh_token: bool,
now: i64,
) -> AccessTokenPlan {
match expires_at {
Some(expires_at) if expires_at <= now + 300 && has_refresh_token => {
AccessTokenPlan::Refresh
}
Some(expires_at) if expires_at <= now + 300 => AccessTokenPlan::ExpiredWithoutRefresh,
_ => AccessTokenPlan::UseStored,
}
}
pub async fn get_valid_access_token(
pool: &PgPool,
master_key: &MasterKey,
tenant_id: &str,
provider: &str,
connector_type: &str,
token_url: &str,
) -> Result<String> {
let store = OAuthCredentialStore::new(pool);
let cred = store
.get(tenant_id, provider, connector_type)
.await?
.ok_or_else(|| {
AppError::Auth(format!(
"No OAuth credential found for tenant '{tenant_id}' and provider '{provider}'"
))
})?;
let at_payload = cred.access_token.ok_or_else(|| {
AppError::Auth(format!(
"OAuth credential for {provider} has no access token"
))
})?;
match access_token_plan(
cred.expires_at,
cred.refresh_token.is_some(),
chrono::Utc::now().timestamp(),
) {
AccessTokenPlan::UseStored => decrypt_api_key(&at_payload, master_key),
AccessTokenPlan::Refresh => {
return refresh_oauth2_token(
pool,
master_key,
tenant_id,
provider,
connector_type,
token_url,
)
.await;
}
AccessTokenPlan::ExpiredWithoutRefresh => {
return Err(AppError::Auth(format!(
"OAuth credential for {provider} is expired and has no refresh token"
)));
}
}
.map_err(|e| AppError::Auth(format!("Failed to decrypt access token: {e}")))
}
pub fn require_tenant_id(args: &serde_json::Value) -> Result<String> {
args.get("tenant_id")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.ok_or_else(|| AppError::InvalidInput("tenant_id is required".to_string()))
}
pub(crate) fn register_prebuilt_connector_tools(
registry: &mut crate::registry::ToolRegistry,
pool: PgPool,
master_key: MasterKey,
) {
let google_calendar = google::GoogleClient::calendar(pool.clone(), master_key.clone());
registry.register(Arc::new(google::calendar::GoogleCalendarListEvents::new(
google_calendar.clone(),
)));
registry.register(Arc::new(google::calendar::GoogleCalendarCreateEvent::new(
google_calendar.clone(),
)));
registry.register(Arc::new(google::calendar::GoogleCalendarDeleteEvent::new(
google_calendar.clone(),
)));
registry.register(Arc::new(google::calendar::GoogleCalendarGetFreeBusy::new(
google_calendar,
)));
let gmail = google::GoogleClient::gmail(pool.clone(), master_key.clone());
registry.register(Arc::new(google::gmail::GmailSendEmail::new(gmail.clone())));
registry.register(Arc::new(google::gmail::GmailListMessages::new(
gmail.clone(),
)));
registry.register(Arc::new(google::gmail::GmailGetMessage::new(gmail)));
let hubspot = hubspot::HubSpotClient::new(pool.clone(), master_key.clone());
registry.register(Arc::new(hubspot::HubSpotGetContact::new(hubspot.clone())));
registry.register(Arc::new(hubspot::HubSpotCreateContact::new(
hubspot.clone(),
)));
registry.register(Arc::new(hubspot::HubSpotListDeals::new(hubspot.clone())));
registry.register(Arc::new(hubspot::HubSpotCreateDeal::new(hubspot)));
let linkedin = linkedin::LinkedInClient::new(pool.clone(), master_key.clone());
registry.register(Arc::new(linkedin::LinkedInCreateShare::new(
linkedin.clone(),
)));
registry.register(Arc::new(linkedin::LinkedInGetCompanyUpdates::new(linkedin)));
let salesforce = salesforce::SalesforceClient::new(pool.clone(), master_key.clone());
registry.register(Arc::new(salesforce::SalesforceSoqlQuery::new(
salesforce.clone(),
)));
registry.register(Arc::new(salesforce::SalesforceGetRecord::new(
salesforce.clone(),
)));
registry.register(Arc::new(salesforce::SalesforceCreateRecord::new(
salesforce,
)));
let slack = slack::SlackClient::new(pool, master_key);
registry.register(Arc::new(slack::SlackSendMessage::new(slack.clone())));
registry.register(Arc::new(slack::SlackListChannels::new(slack.clone())));
registry.register(Arc::new(slack::SlackUploadFile::new(slack)));
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn access_token_plan_treats_missing_expiry_as_stored_token() {
assert_eq!(
access_token_plan(None, false, 1_700_000_000),
AccessTokenPlan::UseStored
);
assert_eq!(
access_token_plan(None, true, 1_700_000_000),
AccessTokenPlan::UseStored
);
}
#[test]
fn access_token_plan_refreshes_only_expiring_credentials_with_refresh_token() {
let now = 1_700_000_000;
assert_eq!(
access_token_plan(Some(now + 299), true, now),
AccessTokenPlan::Refresh
);
assert_eq!(
access_token_plan(Some(now + 299), false, now),
AccessTokenPlan::ExpiredWithoutRefresh
);
assert_eq!(
access_token_plan(Some(now + 301), true, now),
AccessTokenPlan::UseStored
);
}
#[tokio::test]
async fn register_prebuilt_connector_tools_registers_all_bundled_tools() {
let pool = sqlx::postgres::PgPoolOptions::new()
.connect_lazy("postgres://dirmacs@localhost/ares_test")
.expect("lazy pool");
let master_key = MasterKey::from_secret("test-only-master-key");
let mut registry = crate::registry::ToolRegistry::new();
register_prebuilt_connector_tools(&mut registry, pool, master_key);
for name in [
"google_calendar_list_events",
"google_calendar_create_event",
"google_calendar_delete_event",
"google_calendar_get_free_busy",
"gmail_send_email",
"gmail_list_messages",
"gmail_get_message",
"hubspot_get_contact",
"hubspot_create_contact",
"hubspot_list_deals",
"hubspot_create_deal",
"linkedin_create_share",
"linkedin_get_company_updates",
"salesforce_soql_query",
"salesforce_get_record",
"salesforce_create_record",
"slack_send_message",
"slack_list_channels",
"slack_upload_file",
] {
assert!(registry.get(name).is_some(), "missing {name}");
}
}
}