Skip to main content

tmdb_rs/
client.rs

1use std::time::Duration;
2
3use reqwest::StatusCode;
4use serde::de::DeserializeOwned;
5use serde::Deserialize;
6
7use crate::{Error, Result};
8
9const BASE_V3: &str = "https://api.themoviedb.org/3";
10
11#[derive(Clone)]
12enum Auth {
13    /// v3 api key, sent as a query param
14    Key(String),
15    /// v4 read access token, sent as a bearer token
16    Token(String),
17}
18
19/// the entry point; clone freely, it shares one connection pool
20#[derive(Clone)]
21pub struct Client {
22    http: reqwest::Client,
23    auth: Auth,
24    base: String,
25}
26
27impl Client {
28    /// authenticate with a v4 read access token
29    pub fn with_read_token(read_access_token: impl Into<String>) -> Self {
30        Self::build(Auth::Token(read_access_token.into()))
31    }
32
33    /// authenticate with a v3 api key
34    pub fn with_api_key(api_key: impl Into<String>) -> Self {
35        Self::build(Auth::Key(api_key.into()))
36    }
37
38    fn build(auth: Auth) -> Self {
39        Self {
40            http: reqwest::Client::new(),
41            auth,
42            base: BASE_V3.into(),
43        }
44    }
45
46    /// the v4 api, authenticated with a user access token from the v4 auth
47    /// flow — a different credential than the read token, so mixing them is
48    /// a compile error
49    #[cfg(feature = "v4")]
50    pub fn v4(&self, access_token: &crate::AccessToken) -> crate::V4 {
51        let mut client = self.clone();
52        // a custom base (proxy, mock) serves both versions as-is
53        if let Some(origin) = client.base.strip_suffix("/3") {
54            client.base = format!("{origin}/4");
55        }
56        client.auth = Auth::Token(access_token.as_str().to_owned());
57        crate::V4::new(client)
58    }
59
60    /// bring your own reqwest client (proxies, timeouts, ...)
61    pub fn with_http_client(mut self, http: reqwest::Client) -> Self {
62        self.http = http;
63        self
64    }
65
66    /// point at another host (a proxy, or a mock in tests)
67    pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
68        self.base = url.into();
69        self
70    }
71
72    pub(crate) async fn get<T: DeserializeOwned>(
73        &self,
74        path: &str,
75        query: &[(&str, String)],
76    ) -> Result<T> {
77        self.request(reqwest::Method::GET, path, query, None).await
78    }
79
80    pub(crate) async fn post<T: DeserializeOwned>(
81        &self,
82        path: &str,
83        query: &[(&str, String)],
84        body: &serde_json::Value,
85    ) -> Result<T> {
86        self.request(reqwest::Method::POST, path, query, Some(body))
87            .await
88    }
89
90    #[cfg(feature = "v4")]
91    pub(crate) async fn put<T: DeserializeOwned>(
92        &self,
93        path: &str,
94        query: &[(&str, String)],
95        body: &serde_json::Value,
96    ) -> Result<T> {
97        self.request(reqwest::Method::PUT, path, query, Some(body))
98            .await
99    }
100
101    /// TMDB deletes carry a json body too (session, list items)
102    pub(crate) async fn delete<T: DeserializeOwned>(
103        &self,
104        path: &str,
105        query: &[(&str, String)],
106        body: &serde_json::Value,
107    ) -> Result<T> {
108        self.request(reqwest::Method::DELETE, path, query, Some(body))
109            .await
110    }
111
112    async fn request<T: DeserializeOwned>(
113        &self,
114        method: reqwest::Method,
115        path: &str,
116        query: &[(&str, String)],
117        body: Option<&serde_json::Value>,
118    ) -> Result<T> {
119        let mut retried = false;
120        loop {
121            let mut request = self
122                .http
123                .request(method.clone(), format!("{}{path}", self.base));
124            match &self.auth {
125                Auth::Key(key) => request = request.query(&[("api_key", key.as_str())]),
126                Auth::Token(token) => request = request.bearer_auth(token),
127            }
128            if let Some(body) = body {
129                request = request.json(body);
130            }
131            let response = request.query(query).send().await?;
132            let status = response.status();
133
134            if status == StatusCode::TOO_MANY_REQUESTS && !retried {
135                retried = true;
136                let wait = response
137                    .headers()
138                    .get("retry-after")
139                    .and_then(|value| value.to_str().ok())
140                    .and_then(|value| value.parse::<u64>().ok())
141                    .unwrap_or(1);
142                tokio::time::sleep(Duration::from_secs(wait)).await;
143                continue;
144            }
145
146            let body = response.text().await?;
147            if status.is_success() {
148                return Ok(serde_json::from_str(&body)?);
149            }
150
151            #[derive(Deserialize)]
152            struct TmdbError {
153                status_code: i32,
154                status_message: String,
155            }
156            return Err(match serde_json::from_str::<TmdbError>(&body) {
157                Ok(error) if error.status_code == 34 => Error::NotFound,
158                Ok(error) => Error::Tmdb {
159                    code: error.status_code,
160                    message: error.status_message,
161                },
162                Err(_) if status == StatusCode::NOT_FOUND => Error::NotFound,
163                Err(_) if status == StatusCode::TOO_MANY_REQUESTS => {
164                    Error::RateLimited { retry_after: None }
165                }
166                Err(_) => Error::Tmdb {
167                    code: status.as_u16().into(),
168                    message: body,
169                },
170            });
171        }
172    }
173}