use crate::error::{ApiError, ApiResponse, ApiResult, ErrorCode, LockConflictDetail};
use crate::models::user::{RefreshTokenRequest, Token};
use chrono::{DateTime, Duration, Utc};
use reqwest::{Client as HttpClient, Method};
use serde::de::DeserializeOwned;
use serde::Serialize;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::RwLock;
const API_PREFIX: &str = "/api/v4";
pub const CR_HEADER_PREFIX: &str = "X-Cr-";
#[derive(Debug, Clone)]
pub struct ClientConfig {
pub base_url: String,
pub timeout_seconds: u64,
pub client_id: String,
pub user_agent: Option<String>,
}
impl ClientConfig {
pub fn new(base_url: impl Into<String>) -> Self {
Self {
base_url: base_url.into(),
timeout_seconds: 60,
client_id: "".to_string(),
user_agent: None,
}
}
pub fn with_timeout(mut self, timeout_seconds: u64) -> Self {
self.timeout_seconds = timeout_seconds;
self
}
pub fn with_client_id(mut self, client_id: impl Into<String>) -> Self {
self.client_id = client_id.into();
self
}
pub fn with_user_agent(mut self, user_agent: impl Into<String>) -> Self {
self.user_agent = Some(user_agent.into());
self
}
}
#[derive(Debug, Clone)]
pub(crate) struct TokenStore {
access_token: Option<String>,
refresh_token: Option<String>,
access_token_expires: Option<DateTime<Utc>>,
refresh_token_expires: Option<DateTime<Utc>>,
}
impl TokenStore {
fn new() -> Self {
Self {
access_token: None,
refresh_token: None,
access_token_expires: None,
refresh_token_expires: None,
}
}
fn is_access_token_expired(&self) -> bool {
self.access_token_expires
.map(|exp| Utc::now() >= exp)
.unwrap_or(true)
}
fn is_refresh_token_expired(&self) -> bool {
self.refresh_token_expires
.map(|exp| Utc::now() >= exp)
.unwrap_or(true)
}
fn has_tokens(&self) -> bool {
self.access_token.is_some() && self.refresh_token.is_some()
}
}
#[derive(Debug, Clone, Default)]
pub struct RequestOptions {
pub no_credential: bool,
pub with_purchase_ticket: bool,
pub skip_batch_error: bool,
pub skip_lock_conflict: bool,
}
impl RequestOptions {
pub fn new() -> Self {
Self::default()
}
pub fn no_credential(mut self) -> Self {
self.no_credential = true;
self
}
pub fn with_purchase_ticket(mut self) -> Self {
self.with_purchase_ticket = true;
self
}
pub fn skip_batch_error(mut self) -> Self {
self.skip_batch_error = true;
self
}
pub fn skip_lock_conflict(mut self) -> Self {
self.skip_lock_conflict = true;
self
}
}
pub type OnCredentialRefreshed =
Arc<dyn Fn(Token) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
pub type OnCredentialInvalid =
Arc<dyn Fn() -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync>;
pub struct Client {
pub(crate) config: ClientConfig,
pub(crate) http_client: HttpClient,
pub(crate) tokens: Arc<RwLock<TokenStore>>,
pub(crate) purchase_ticket: Arc<RwLock<Option<String>>>,
on_credential_refreshed: Option<OnCredentialRefreshed>,
on_credential_invalid: Option<OnCredentialInvalid>,
}
impl Client {
pub fn new(config: ClientConfig) -> Self {
let mut builder = HttpClient::builder()
.connect_timeout(std::time::Duration::from_secs(config.timeout_seconds));
if let Some(ref user_agent) = config.user_agent {
builder = builder.user_agent(user_agent);
}
let http_client = builder.build().expect("Failed to create HTTP client");
Self {
config,
http_client,
tokens: Arc::new(RwLock::new(TokenStore::new())),
purchase_ticket: Arc::new(RwLock::new(None)),
on_credential_refreshed: None,
on_credential_invalid: None,
}
}
pub fn set_on_credential_refreshed(&mut self, callback: OnCredentialRefreshed) {
self.on_credential_refreshed = Some(callback);
}
pub fn clear_on_credential_refreshed(&mut self) {
self.on_credential_refreshed = None;
}
pub fn set_on_credential_invalid(&mut self, callback: OnCredentialInvalid) {
self.on_credential_invalid = Some(callback);
}
pub fn clear_on_credential_invalid(&mut self) {
self.on_credential_invalid = None;
}
async fn notify_credential_invalid(&self) {
if let Some(ref callback) = self.on_credential_invalid {
callback().await;
}
}
pub async fn set_tokens(&self, access_token: String, refresh_token: String) {
let mut store = self.tokens.write().await;
let access_expires = Utc::now() + Duration::hours(1);
let refresh_expires = Utc::now() + Duration::days(7);
store.access_token = Some(access_token);
store.refresh_token = Some(refresh_token);
store.access_token_expires = Some(access_expires);
store.refresh_token_expires = Some(refresh_expires);
}
pub async fn set_tokens_with_expiry(&self, token: &Token) -> ApiResult<()> {
let mut store = self.tokens.write().await;
store.access_token = Some(token.access_token.clone());
store.refresh_token = Some(token.refresh_token.clone());
if let Ok(exp) = DateTime::parse_from_rfc3339(&token.access_expires) {
store.access_token_expires = Some(exp.with_timezone(&Utc));
}
if let Ok(exp) = DateTime::parse_from_rfc3339(&token.refresh_expires) {
store.refresh_token_expires = Some(exp.with_timezone(&Utc));
}
Ok(())
}
pub async fn clear_tokens(&self) {
let mut store = self.tokens.write().await;
*store = TokenStore::new();
}
pub async fn set_purchase_ticket(&self, ticket: Option<String>) {
let mut pt = self.purchase_ticket.write().await;
*pt = ticket;
}
pub(crate) async fn get_access_token(&self) -> ApiResult<String> {
let store = self.tokens.read().await;
if !store.has_tokens() {
self.notify_credential_invalid().await;
return Err(ApiError::NoTokensAvailable);
}
if store.is_refresh_token_expired() {
self.notify_credential_invalid().await;
return Err(ApiError::RefreshTokenExpired);
}
if !store.is_access_token_expired() {
return Ok(store.access_token.clone().unwrap());
}
drop(store);
self.refresh_access_token().await
}
async fn refresh_access_token(&self) -> ApiResult<String> {
let refresh_token = {
let store = self.tokens.read().await;
store
.refresh_token
.clone()
.ok_or(ApiError::NoTokensAvailable)?
};
let url = self.build_url("/session/token/refresh");
let request = RefreshTokenRequest { refresh_token };
let response = self.http_client.post(&url).json(&request).send().await?;
let api_response: ApiResponse<Token> = response.json().await?;
if api_response.code != ErrorCode::Success as i32 {
if let Some(error_code) = ErrorCode::from_code(api_response.code) {
if error_code.is_credential_error() {
self.notify_credential_invalid().await;
}
}
return Err(ApiError::from_response(api_response));
}
let token = api_response
.data
.ok_or_else(|| ApiError::Other("No token in response".to_string()))?;
self.set_tokens_with_expiry(&token).await?;
if let Some(ref callback) = self.on_credential_refreshed {
callback(token.clone()).await;
}
Ok(token.access_token)
}
pub(crate) fn build_url(&self, path: &str) -> String {
format!("{}{}{}", self.config.base_url, API_PREFIX, path)
}
async fn send_internal<T, R>(
&self,
path: &str,
method: Method,
body: Option<&T>,
options: RequestOptions,
) -> ApiResult<R>
where
T: Serialize + ?Sized,
R: DeserializeOwned + Default,
{
let url = self.build_url(path);
let mut request = self.http_client.request(method, &url);
if !options.no_credential {
let token = self.get_access_token().await?;
request = request.header("Authorization", format!("Bearer {}", token));
}
if !self.config.client_id.is_empty() {
request = request.header(
format!("{}Client-Id", CR_HEADER_PREFIX),
self.config.client_id.clone(),
);
}
if options.with_purchase_ticket {
let ticket = self.purchase_ticket.read().await;
if let Some(t) = ticket.as_ref() {
request = request.header(format!("{}Purchase-Ticket", CR_HEADER_PREFIX), t);
}
}
if let Some(body) = body {
request = request.json(body);
}
let response = request.send().await?;
let response_text = response.text().await?;
let raw_value: serde_json::Value = serde_json::from_str(&response_text)?;
let code = raw_value.get("code").and_then(|c| c.as_i64()).unwrap_or(0) as i32;
if code == ErrorCode::LockConflict as i32 {
let msg = raw_value
.get("msg")
.and_then(|m| m.as_str())
.unwrap_or("")
.to_string();
let detail: Option<LockConflictDetail> = raw_value
.get("data")
.and_then(|d| serde_json::from_value(d.clone()).ok());
return Err(ApiError::LockConflict {
message: msg,
detail,
});
}
let api_response: ApiResponse<R> = serde_json::from_str(&response_text)?;
if api_response.code != ErrorCode::Success as i32 {
if let Some(error_code) = ErrorCode::from_code(api_response.code) {
if error_code.is_credential_error() {
self.notify_credential_invalid().await;
}
}
return Err(ApiError::from_response(api_response));
}
Ok(api_response.data.unwrap_or_default())
}
pub async fn send<T, R>(
&self,
path: &str,
method: Method,
body: Option<&T>,
options: RequestOptions,
) -> ApiResult<R>
where
T: Serialize + ?Sized,
R: DeserializeOwned + Default,
{
match self
.send_internal(path, method.clone(), body, options.clone())
.await
{
Ok(result) => Ok(result),
Err(ApiError::AccessTokenExpired) => {
self.refresh_access_token().await?;
self.send_internal(path, method, body, options).await
}
Err(e) => Err(e),
}
}
pub async fn get<R>(&self, path: &str, options: RequestOptions) -> ApiResult<R>
where
R: DeserializeOwned + Default,
{
self.send::<(), R>(path, Method::GET, None, options).await
}
pub async fn post<T, R>(&self, path: &str, body: &T, options: RequestOptions) -> ApiResult<R>
where
T: Serialize,
R: DeserializeOwned + Default,
{
self.send(path, Method::POST, Some(body), options).await
}
pub async fn put<T, R>(&self, path: &str, body: &T, options: RequestOptions) -> ApiResult<R>
where
T: Serialize,
R: DeserializeOwned + Default,
{
self.send(path, Method::PUT, Some(body), options).await
}
pub async fn delete<R>(&self, path: &str, options: RequestOptions) -> ApiResult<R>
where
R: DeserializeOwned + Default,
{
self.send::<(), R>(path, Method::DELETE, None, options)
.await
}
pub async fn delete_with_body<T, R>(
&self,
path: &str,
body: &T,
options: RequestOptions,
) -> ApiResult<R>
where
T: Serialize,
R: DeserializeOwned + Default,
{
self.send(path, Method::DELETE, Some(body), options).await
}
pub async fn patch<T, R>(&self, path: &str, body: &T, options: RequestOptions) -> ApiResult<R>
where
T: Serialize,
R: DeserializeOwned + Default,
{
self.send(path, Method::PATCH, Some(body), options).await
}
}