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(¶ms);
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}