Skip to main content

douyin_cli/
openapi.rs

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