Skip to main content

douyin_cli/
openapi.rs

1use std::collections::HashMap;
2use std::time::Duration;
3
4use reqwest::Url;
5use reqwest::blocking::{Client, RequestBuilder};
6use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
7use serde_json::{Map, Value, json};
8
9pub const BASE_URL: &str = "https://open.douyin.com";
10
11pub struct OpenApiClient {
12    base_url: Url,
13    client: Client,
14}
15
16impl OpenApiClient {
17    pub fn new() -> Result<Self, String> {
18        Self::with_base_url(BASE_URL)
19    }
20
21    pub fn with_base_url(base_url: &str) -> Result<Self, String> {
22        let mut base_url =
23            Url::parse(base_url).map_err(|error| format!("OpenAPI base URL 无效: {error}"))?;
24        if !matches!(base_url.scheme(), "http" | "https") || base_url.host_str().is_none() {
25            return Err("OpenAPI base URL 必须是有效的 HTTP(S) 地址".to_owned());
26        }
27        base_url.set_path("/");
28        base_url.set_query(None);
29        base_url.set_fragment(None);
30        let client = Client::builder()
31            .connect_timeout(Duration::from_secs(10))
32            .timeout(Duration::from_secs(30))
33            .build()
34            .map_err(|error| error.to_string())?;
35        Ok(Self { base_url, client })
36    }
37
38    pub fn authorize_url(
39        &self,
40        client_key: &str,
41        redirect_uri: &str,
42        scopes: &[String],
43        state: Option<&str>,
44    ) -> Result<String, String> {
45        let mut url = self.url("/platform/oauth/connect/")?;
46        {
47            let mut query = url.query_pairs_mut();
48            query.append_pair("client_key", client_key);
49            query.append_pair("response_type", "code");
50            query.append_pair("scope", &scopes.join(","));
51            query.append_pair("redirect_uri", redirect_uri);
52            if let Some(state) = state {
53                query.append_pair("state", state);
54            }
55        }
56        Ok(url.into())
57    }
58
59    pub fn client_token(&self, client_key: &str, client_secret: &str) -> Result<Value, String> {
60        self.request(RequestSpec {
61            method: "POST",
62            path: "/oauth/client_token/",
63            form: Some(HashMap::from([
64                ("client_key".to_owned(), client_key.to_owned()),
65                ("client_secret".to_owned(), client_secret.to_owned()),
66                ("grant_type".to_owned(), "client_credential".to_owned()),
67            ])),
68            auth_required: false,
69            ..RequestSpec::default()
70        })
71    }
72
73    pub fn access_token(
74        &self,
75        client_key: &str,
76        client_secret: &str,
77        code: &str,
78    ) -> Result<Value, String> {
79        self.request(RequestSpec {
80            method: "POST",
81            path: "/oauth/access_token/",
82            form: Some(HashMap::from([
83                ("client_key".to_owned(), client_key.to_owned()),
84                ("client_secret".to_owned(), client_secret.to_owned()),
85                ("code".to_owned(), code.to_owned()),
86                ("grant_type".to_owned(), "authorization_code".to_owned()),
87            ])),
88            auth_required: false,
89            ..RequestSpec::default()
90        })
91    }
92
93    pub fn refresh_token(&self, client_key: &str, refresh_token: &str) -> Result<Value, String> {
94        self.request(RequestSpec {
95            method: "POST",
96            path: "/oauth/refresh_token/",
97            form: Some(HashMap::from([
98                ("client_key".to_owned(), client_key.to_owned()),
99                ("grant_type".to_owned(), "refresh_token".to_owned()),
100                ("refresh_token".to_owned(), refresh_token.to_owned()),
101            ])),
102            auth_required: false,
103            ..RequestSpec::default()
104        })
105    }
106
107    pub fn renew_refresh_token(
108        &self,
109        client_key: &str,
110        refresh_token: &str,
111    ) -> Result<Value, String> {
112        self.request(RequestSpec {
113            method: "POST",
114            path: "/oauth/renew_refresh_token/",
115            form: Some(HashMap::from([
116                ("client_key".to_owned(), client_key.to_owned()),
117                ("refresh_token".to_owned(), refresh_token.to_owned()),
118            ])),
119            auth_required: false,
120            ..RequestSpec::default()
121        })
122    }
123
124    pub fn request(&self, spec: RequestSpec<'_>) -> Result<Value, String> {
125        if spec.auth_required && spec.token.is_none_or(str::is_empty) {
126            return Err("调用 OpenAPI 需要 access-token 或 client-token".to_owned());
127        }
128        let method = spec
129            .method
130            .parse::<reqwest::Method>()
131            .map_err(|error| format!("HTTP method 无效: {error}"))?;
132        let url = self.url(spec.path)?;
133        let mut request = self.client.request(method, url);
134        if let Some(token) = spec.token {
135            request = request.header("access-token", token);
136        }
137        request = add_headers(request, spec.headers)?;
138        if let Some(params) = spec.params {
139            request = request.query(&params);
140        }
141        if let Some(form) = spec.form {
142            request = request.form(&form);
143        } else if let Some(body) = spec.json_body {
144            request = request.json(&body);
145        }
146
147        let response = request.send().map_err(|error| error.to_string())?;
148        let status = response.status();
149        let text = response.text().map_err(|error| error.to_string())?;
150        if !status.is_success() {
151            return Err(format!(
152                "OpenAPI HTTP 请求失败: {status} {}",
153                body_excerpt(&text)
154            ));
155        }
156        let data: Value = serde_json::from_str(&text)
157            .map_err(|_| format!("OpenAPI 响应不是 JSON: {}", body_excerpt(&text)))?;
158        if !data.is_object() {
159            return Err("OpenAPI 响应不是 JSON object".to_owned());
160        }
161        Ok(data)
162    }
163
164    fn url(&self, path: &str) -> Result<Url, String> {
165        let resolved = self
166            .base_url
167            .join(path)
168            .map_err(|error| format!("OpenAPI path 无效: {error}"))?;
169        if !same_origin(&self.base_url, &resolved) {
170            return Err(format!(
171                "拒绝跨域 OpenAPI 请求: {}",
172                resolved.origin().ascii_serialization()
173            ));
174        }
175        Ok(resolved)
176    }
177}
178
179fn same_origin(left: &Url, right: &Url) -> bool {
180    left.scheme() == right.scheme()
181        && left.host_str() == right.host_str()
182        && left.port_or_known_default() == right.port_or_known_default()
183}
184
185fn body_excerpt(body: &str) -> String {
186    const MAX_CHARS: usize = 2_000;
187    let mut characters = body.chars();
188    let excerpt: String = characters.by_ref().take(MAX_CHARS).collect();
189    if characters.next().is_some() {
190        format!("{excerpt}…(响应已截断)")
191    } else {
192        excerpt
193    }
194}
195
196#[derive(Default)]
197pub struct RequestSpec<'a> {
198    pub method: &'a str,
199    pub path: &'a str,
200    pub token: Option<&'a str>,
201    pub params: Option<HashMap<String, String>>,
202    pub json_body: Option<Value>,
203    pub form: Option<HashMap<String, String>>,
204    pub headers: Option<HashMap<String, String>>,
205    pub auth_required: bool,
206}
207
208pub fn im_message_body(
209    to_user_id: &str,
210    message_type: &str,
211    content: Value,
212    persona_id: Option<&str>,
213    client_msg_id: Option<&str>,
214) -> Value {
215    let mut body = Map::from_iter([
216        ("to_user_id".to_owned(), json!(to_user_id)),
217        ("message_type".to_owned(), json!(message_type)),
218        ("content".to_owned(), json!(content.to_string())),
219    ]);
220    if let Some(value) = persona_id {
221        body.insert("persona_id".to_owned(), json!(value));
222    }
223    if let Some(value) = client_msg_id {
224        body.insert("client_msg_id".to_owned(), json!(value));
225    }
226    Value::Object(body)
227}
228
229fn add_headers(
230    mut request: RequestBuilder,
231    headers: Option<HashMap<String, String>>,
232) -> Result<RequestBuilder, String> {
233    let Some(headers) = headers else {
234        return Ok(request);
235    };
236    let mut values = HeaderMap::new();
237    for (key, value) in headers {
238        let key = HeaderName::try_from(key).map_err(|error| error.to_string())?;
239        let value = HeaderValue::try_from(value).map_err(|error| error.to_string())?;
240        values.insert(key, value);
241    }
242    request = request.headers(values);
243    Ok(request)
244}
245
246#[cfg(test)]
247mod tests {
248    use super::{OpenApiClient, RequestSpec, body_excerpt, im_message_body};
249    use serde_json::json;
250
251    #[test]
252    fn authorize_url_encodes_values() {
253        let client = OpenApiClient::new().unwrap();
254        let url = client
255            .authorize_url(
256                "client",
257                "https://example.com/callback",
258                &["user_info".to_owned(), "item.comment".to_owned()],
259                Some("state value"),
260            )
261            .unwrap();
262        assert!(url.starts_with("https://open.douyin.com/platform/oauth/connect/?"));
263        assert!(url.contains("scope=user_info%2Citem.comment"));
264        assert!(url.contains("redirect_uri=https%3A%2F%2Fexample.com%2Fcallback"));
265        assert!(url.contains("state=state+value"));
266    }
267
268    #[test]
269    fn request_rejects_missing_token_before_network() {
270        let client = OpenApiClient::new().unwrap();
271        let error = client
272            .request(RequestSpec {
273                method: "GET",
274                path: "/oauth/userinfo/",
275                auth_required: true,
276                ..RequestSpec::default()
277            })
278            .unwrap_err();
279        assert!(error.contains("access-token"));
280    }
281
282    #[test]
283    fn request_rejects_cross_origin_url_before_network() {
284        let client = OpenApiClient::new().unwrap();
285        let error = client
286            .request(RequestSpec {
287                method: "GET",
288                path: "https://example.com/collect",
289                token: Some("secret-token"),
290                auth_required: true,
291                ..RequestSpec::default()
292            })
293            .unwrap_err();
294        assert!(error.contains("拒绝跨域"));
295        assert!(!error.contains("secret-token"));
296    }
297
298    #[test]
299    fn validates_base_url_and_bounds_error_bodies() {
300        assert!(OpenApiClient::with_base_url("file:///tmp/api").is_err());
301        assert_eq!(body_excerpt("short"), "short");
302        let excerpt = body_excerpt(&"界".repeat(2_001));
303        assert!(excerpt.ends_with("…(响应已截断)"));
304        assert_eq!(excerpt.chars().count(), 2_008);
305    }
306
307    #[test]
308    fn im_body_serializes_content_as_compact_json_string() {
309        let body = im_message_body(
310            "user",
311            "text",
312            json!({"text": "你好"}),
313            None,
314            Some("client-msg"),
315        );
316        assert_eq!(body["content"], "{\"text\":\"你好\"}");
317        assert_eq!(body["client_msg_id"], "client-msg");
318    }
319}