use crate::{SlackApiError, SlackError};
use bytes::Bytes;
use reqwest::header::{CONTENT_TYPE, RETRY_AFTER};
use reqwest::{Response, StatusCode};
use serde::de::DeserializeOwned;
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use std::time::Duration;
const DEFAULT_BASE_URL: &str = "https://slack.com/api/";
const DEFAULT_MAX_RETRIES: u32 = 0;
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
const DEFAULT_RETRY_AFTER: Duration = Duration::from_secs(1);
const FORM_CONTENT_TYPE: &str = "application/x-www-form-urlencoded";
pub trait SlackApiMethod: Serialize {
const METHOD: &'static str;
type Response: DeserializeOwned;
}
pub trait CursorPaginated: SlackApiMethod<Response: NextCursor> {
fn set_cursor(&mut self, cursor: String);
}
pub trait NextCursor {
fn next_cursor(&self) -> Option<&str>;
}
#[derive(Clone)]
pub struct SlackClient {
inner: Arc<Inner>,
}
struct Inner {
http: reqwest::Client,
token: Option<String>,
base_url: String,
max_retries: u32,
}
impl std::fmt::Debug for SlackClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SlackClient")
.field("token", &self.inner.token.as_ref().map(|_| "***"))
.field("base_url", &self.inner.base_url)
.field("max_retries", &self.inner.max_retries)
.finish()
}
}
impl SlackClient {
pub fn new(token: impl Into<String>) -> Self {
Self::builder().token(token).build()
}
pub fn builder() -> SlackClientBuilder {
SlackClientBuilder::default()
}
pub fn with_token(&self, token: impl Into<String>) -> Self {
Self {
inner: Arc::new(Inner {
http: self.inner.http.clone(),
token: Some(token.into()),
base_url: self.inner.base_url.clone(),
max_retries: self.inner.max_retries,
}),
}
}
pub async fn call<M: SlackApiMethod>(&self, request: &M) -> Result<M::Response, SlackError> {
let (status, body) = self.post(M::METHOD, request).await?;
decode(status, &body)
}
pub async fn call_raw<P: Serialize + ?Sized>(
&self,
method: &str,
params: &P,
) -> Result<serde_json::Value, SlackError> {
let (status, body) = self.post(method, params).await?;
decode(status, &body)
}
pub fn pages<M: CursorPaginated>(&self, request: M) -> Pages<'_, M> {
Pages {
client: self,
request: Some(request),
}
}
pub(crate) async fn call_bytes<P: Serialize + ?Sized>(
&self,
method: &str,
params: &P,
) -> Result<Bytes, SlackError> {
let (status, body) = self.post(method, params).await?;
if is_json(&body) {
decode::<serde::de::IgnoredAny>(status, &body)?;
}
if !status.is_success() {
return Err(SlackError::Http {
status: status.as_u16(),
body: String::from_utf8_lossy(&body).into_owned(),
});
}
Ok(body)
}
pub(crate) fn http(&self) -> &reqwest::Client {
&self.inner.http
}
async fn post<P: Serialize + ?Sized>(
&self,
method: &str,
params: &P,
) -> Result<(StatusCode, Bytes), SlackError> {
let form = crate::form::to_form(params).map_err(|e| SlackError::Encode(e.to_string()))?;
let body = Bytes::from(form);
let url = format!("{}{}", self.inner.base_url, method);
let mut attempt = 0;
loop {
let mut request = self
.inner
.http
.post(&url)
.header(CONTENT_TYPE, FORM_CONTENT_TYPE)
.body(body.clone());
if let Some(token) = &self.inner.token {
request = request.bearer_auth(token);
}
let response = request.send().await?;
let status = response.status();
if status == StatusCode::TOO_MANY_REQUESTS {
let wait = retry_after(&response);
if attempt >= self.inner.max_retries {
return Err(SlackError::RateLimited { retry_after: wait });
}
attempt += 1;
tokio::time::sleep(wait.unwrap_or(DEFAULT_RETRY_AFTER)).await;
continue;
}
return Ok((status, response.bytes().await?));
}
}
}
#[derive(Default)]
pub struct SlackClientBuilder {
http: Option<reqwest::Client>,
token: Option<String>,
base_url: Option<String>,
max_retries: Option<u32>,
}
impl SlackClientBuilder {
pub fn token(mut self, token: impl Into<String>) -> Self {
self.token = Some(token.into());
self
}
pub fn http_client(mut self, http: reqwest::Client) -> Self {
self.http = Some(http);
self
}
pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
let mut base_url = base_url.into();
if !base_url.ends_with('/') {
base_url.push('/');
}
self.base_url = Some(base_url);
self
}
pub fn max_retries(mut self, max_retries: u32) -> Self {
self.max_retries = Some(max_retries);
self
}
pub fn build(self) -> SlackClient {
let http = self.http.unwrap_or_else(|| {
reqwest::Client::builder()
.timeout(DEFAULT_TIMEOUT)
.connect_timeout(DEFAULT_CONNECT_TIMEOUT)
.build()
.expect("failed to initialize the TLS backend for reqwest")
});
SlackClient {
inner: Arc::new(Inner {
http,
token: self.token,
base_url: self.base_url.unwrap_or_else(|| DEFAULT_BASE_URL.to_owned()),
max_retries: self.max_retries.unwrap_or(DEFAULT_MAX_RETRIES),
}),
}
}
}
pub struct Pages<'a, M> {
client: &'a SlackClient,
request: Option<M>,
}
impl<M: CursorPaginated> Pages<'_, M> {
pub async fn next_page(&mut self) -> Option<Result<M::Response, SlackError>> {
let mut request = self.request.take()?;
let result = self.client.call(&request).await;
if let Ok(response) = &result {
if let Some(cursor) = response.next_cursor().filter(|c| !c.is_empty()) {
request.set_cursor(cursor.to_owned());
self.request = Some(request);
}
}
Some(result)
}
}
#[derive(Deserialize)]
struct Envelope {
ok: bool,
#[serde(default)]
error: Option<String>,
#[serde(default)]
warning: Option<String>,
#[serde(default)]
response_metadata: Option<crate::ResponseMetadata>,
}
fn decode<R: DeserializeOwned>(status: StatusCode, body: &[u8]) -> Result<R, SlackError> {
let envelope: Envelope = match serde_json::from_slice(body) {
Ok(envelope) => envelope,
Err(source) => {
let body = String::from_utf8_lossy(body).into_owned();
if !status.is_success() {
return Err(SlackError::Http {
status: status.as_u16(),
body,
});
}
return Err(SlackError::Decode { source, body });
}
};
if !envelope.ok {
return Err(SlackError::Api(Box::new(SlackApiError {
error: envelope.error.unwrap_or_default(),
warning: envelope.warning,
response_metadata: envelope.response_metadata,
})));
}
serde_json::from_slice(body).map_err(|source| SlackError::Decode {
source,
body: String::from_utf8_lossy(body).into_owned(),
})
}
fn is_json(body: &[u8]) -> bool {
body.iter().find(|b| !b.is_ascii_whitespace()) == Some(&b'{')
}
fn retry_after(response: &Response) -> Option<Duration> {
let secs = response
.headers()
.get(RETRY_AFTER)?
.to_str()
.ok()?
.trim()
.parse::<u64>()
.ok()?;
Some(Duration::from_secs(secs))
}