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