Skip to main content

douyin_cli/
auth.rs

1use std::collections::HashMap;
2use std::io::{Read, Write};
3use std::net::{TcpListener, TcpStream};
4use std::thread;
5use std::time::{Duration, Instant};
6
7use base64::Engine;
8use base64::engine::general_purpose::URL_SAFE_NO_PAD;
9use clap::{Args, Subcommand};
10use qrcode::{Color, QrCode};
11use ring::rand::{SecureRandom, SystemRandom};
12use serde_json::{Map, Value, json};
13
14use crate::cookie;
15use crate::openapi::{OpenApiClient, RequestSpec};
16use crate::settings;
17
18#[derive(Debug, Args)]
19pub struct AuthArgs {
20    #[command(subcommand)]
21    command: AuthCommand,
22}
23
24#[derive(Debug, Subcommand)]
25enum AuthCommand {
26    /// 通过官方 OAuth 授权接入账号
27    Login {
28        #[arg(long, env = "DOUYIN_CLIENT_KEY")]
29        client_key: Option<String>,
30        #[arg(long, env = "DOUYIN_CLIENT_SECRET")]
31        client_secret: Option<String>,
32        #[arg(long)]
33        redirect_uri: Option<String>,
34        #[arg(long)]
35        scope: Vec<String>,
36        #[arg(long)]
37        code: Option<String>,
38        #[arg(long, conflicts_with = "no_qr")]
39        qr: bool,
40        #[arg(long, conflicts_with = "qr")]
41        no_qr: bool,
42        #[arg(long)]
43        listen: bool,
44        #[arg(long, default_value = "127.0.0.1")]
45        callback_host: String,
46        #[arg(long, default_value_t = 8787, value_parser = clap::value_parser!(u16).range(1..))]
47        callback_port: u16,
48        #[arg(long, default_value_t = 300, value_parser = clap::value_parser!(u64).range(1..=3600))]
49        timeout: u64,
50    },
51    /// 用官方 OAuth code 换取并保存 token
52    Code {
53        #[arg(long)]
54        code: String,
55        #[arg(long, env = "DOUYIN_CLIENT_SECRET")]
56        client_secret: Option<String>,
57    },
58    /// 刷新已保存的官方 access_token
59    Refresh,
60    /// 检查官方授权状态
61    Status {
62        #[arg(long)]
63        json: bool,
64    },
65    /// 删除已保存的官方 OAuth token
66    Logout,
67    /// 保存网页端 Cookie,用于搜索、评论和下载等网页端采集
68    CookieLogin {
69        #[arg(long, env = "DOUYIN_COOKIE")]
70        cookie: String,
71    },
72    /// 检查已保存 Cookie 是否可用于网页端请求
73    CookieStatus {
74        #[arg(long)]
75        offline: bool,
76    },
77    /// 删除已保存的网页端 Cookie
78    CookieLogout,
79}
80
81pub fn run(args: AuthArgs) -> Result<(), String> {
82    match args.command {
83        AuthCommand::Login {
84            client_key,
85            client_secret,
86            redirect_uri,
87            scope,
88            code,
89            qr: _,
90            no_qr,
91            listen,
92            callback_host,
93            callback_port,
94            timeout,
95        } => login(LoginOptions {
96            client_key,
97            client_secret,
98            redirect_uri,
99            scopes: scope,
100            code,
101            show_qr: !no_qr,
102            listen,
103            callback_host,
104            callback_port,
105            timeout,
106        }),
107        AuthCommand::Code {
108            code,
109            client_secret,
110        } => exchange_code(&code, client_secret),
111        AuthCommand::Refresh => refresh(),
112        AuthCommand::Status { json } => status(json),
113        AuthCommand::Logout => logout(),
114        AuthCommand::CookieLogin { cookie } => cookie_login(&cookie),
115        AuthCommand::CookieStatus { offline } => cookie_status(offline),
116        AuthCommand::CookieLogout => cookie_logout(),
117    }
118}
119
120struct LoginOptions {
121    client_key: Option<String>,
122    client_secret: Option<String>,
123    redirect_uri: Option<String>,
124    scopes: Vec<String>,
125    code: Option<String>,
126    show_qr: bool,
127    listen: bool,
128    callback_host: String,
129    callback_port: u16,
130    timeout: u64,
131}
132
133fn login(options: LoginOptions) -> Result<(), String> {
134    let data = settings::load().map_err(|error| error.to_string())?;
135    let saved = settings::openapi(&data);
136    let client_key = options
137        .client_key
138        .or_else(|| saved_string(&saved, "clientKey"))
139        .ok_or_else(|| missing_client_key(options.show_qr))?;
140    let client_secret = options
141        .client_secret
142        .or_else(|| saved_string(&saved, "clientSecret"));
143    let mut redirect_uri = options
144        .redirect_uri
145        .or_else(|| saved_string(&saved, "redirectUri"));
146    let scopes = if options.scopes.is_empty() {
147        saved
148            .get("scopes")
149            .and_then(Value::as_array)
150            .map(|values| {
151                values
152                    .iter()
153                    .filter_map(Value::as_str)
154                    .map(str::to_owned)
155                    .collect()
156            })
157            .filter(|values: &Vec<String>| !values.is_empty())
158            .unwrap_or_else(|| vec!["user_info".to_owned()])
159    } else {
160        options.scopes
161    };
162    if options.listen {
163        redirect_uri = Some(format!(
164            "http://{}:{}/callback",
165            options.callback_host, options.callback_port
166        ));
167    }
168    let redirect_uri =
169        redirect_uri.ok_or_else(|| "缺少 redirect_uri,请传入 --redirect-uri".to_owned())?;
170    let state = (options.listen && options.code.is_none())
171        .then(random_state)
172        .transpose()?;
173    let client = OpenApiClient::new()?;
174    let url = client.authorize_url(&client_key, &redirect_uri, &scopes, state.as_deref())?;
175    println!("请在浏览器打开以下官方授权链接:\n{url}");
176    if options.show_qr {
177        print_qr(&url)?;
178    }
179
180    let mut code = options.code;
181    if options.listen && code.is_none() {
182        println!("正在等待授权回调: {redirect_uri}");
183        code = Some(wait_for_code(
184            &options.callback_host,
185            options.callback_port,
186            state.as_deref(),
187            Duration::from_secs(options.timeout),
188        )?);
189    }
190    let mut updates = Map::from_iter([
191        ("clientKey".to_owned(), json!(client_key)),
192        (
193            "clientSecret".to_owned(),
194            json!(client_secret.clone().unwrap_or_default()),
195        ),
196        ("redirectUri".to_owned(), json!(redirect_uri)),
197        ("scopes".to_owned(), json!(scopes)),
198    ]);
199    if let Some(code) = code {
200        let secret =
201            client_secret.ok_or_else(|| "使用 code 换 token 需要 --client-secret".to_owned())?;
202        let response = client.access_token(&client_key, &secret, &code)?;
203        updates.extend(extract_token_fields(&response));
204        print_json(&response)?;
205    } else {
206        println!("授权完成后运行:douyin auth code --code 授权码");
207    }
208    save_openapi(updates)?;
209    println!(
210        "官方授权配置已保存: {}",
211        settings::settings_file().display()
212    );
213    Ok(())
214}
215
216fn exchange_code(code: &str, client_secret: Option<String>) -> Result<(), String> {
217    let data = settings::load().map_err(|error| error.to_string())?;
218    let saved = settings::openapi(&data);
219    let client_key = saved_string(&saved, "clientKey")
220        .ok_or_else(|| "缺少 client_key,请先运行 douyin auth login".to_owned())?;
221    let secret = client_secret
222        .or_else(|| saved_string(&saved, "clientSecret"))
223        .ok_or_else(|| "缺少 client_secret,请传入 --client-secret".to_owned())?;
224    let response = OpenApiClient::new()?.access_token(&client_key, &secret, code)?;
225    let mut updates = extract_token_fields(&response);
226    updates.insert("clientSecret".to_owned(), json!(secret));
227    save_openapi(updates)?;
228    print_json(&response)?;
229    println!("官方 token 已保存: {}", settings::settings_file().display());
230    Ok(())
231}
232
233fn refresh() -> Result<(), String> {
234    let data = settings::load().map_err(|error| error.to_string())?;
235    let saved = settings::openapi(&data);
236    let client_key = saved_string(&saved, "clientKey");
237    let refresh_token = saved_string(&saved, "refreshToken");
238    let (Some(client_key), Some(refresh_token)) = (client_key, refresh_token) else {
239        return Err("缺少 client_key 或 refresh_token,请重新授权".to_owned());
240    };
241    let response = OpenApiClient::new()?.refresh_token(&client_key, &refresh_token)?;
242    save_openapi(extract_token_fields(&response))?;
243    print_json(&response)?;
244    println!("官方 token 已刷新");
245    Ok(())
246}
247
248fn status(json_output: bool) -> Result<(), String> {
249    let data = settings::load().map_err(|error| error.to_string())?;
250    let saved = settings::openapi(&data);
251    let token = saved_string(&saved, "accessToken");
252    let open_id = saved_string(&saved, "openId");
253    let authorized = token.is_some() && open_id.is_some();
254    let mut output = json!({
255        "authorized": authorized,
256        "connected": false,
257        "configFile": settings::settings_file(),
258        "openId": open_id.clone().unwrap_or_default(),
259        "scopes": saved.get("scopes").cloned().unwrap_or_else(|| json!([]))
260    });
261    let (Some(token), Some(open_id)) = (token, open_id) else {
262        if json_output {
263            return print_json(&output);
264        }
265        println!("未完成官方授权");
266        return Ok(());
267    };
268    if !json_output {
269        println!("已保存官方授权: {}", settings::settings_file().display());
270        println!("open_id: {open_id}");
271        println!("正在检查官方 OpenAPI 连通性...");
272    }
273    match OpenApiClient::new()?.request(RequestSpec {
274        method: "GET",
275        path: "/oauth/userinfo/",
276        token: Some(&token),
277        params: Some(HashMap::from([("open_id".to_owned(), open_id)])),
278        auth_required: true,
279        ..RequestSpec::default()
280    }) {
281        Ok(userinfo) => {
282            output["connected"] = json!(true);
283            output["userinfo"] = userinfo.clone();
284            print_json(if json_output { &output } else { &userinfo })
285        }
286        Err(error) if json_output => {
287            output["error"] = json!(error);
288            print_json(&output)?;
289            Err("官方 OpenAPI 连通性检查失败".to_owned())
290        }
291        Err(error) => Err(format!("官方 OpenAPI 连通性检查失败: {error}")),
292    }
293}
294
295fn logout() -> Result<(), String> {
296    save_openapi(Map::from_iter([
297        ("accessToken".to_owned(), json!("")),
298        ("refreshToken".to_owned(), json!("")),
299        ("openId".to_owned(), json!("")),
300        ("expiresIn".to_owned(), json!(0)),
301    ]))?;
302    println!("已清除官方授权 token");
303    Ok(())
304}
305
306fn cookie_login(value: &str) -> Result<(), String> {
307    let value = value.trim();
308    if !cookie::validate(value) {
309        return Err("Cookie 格式校验失败,未保存".to_owned());
310    }
311    let mut data = settings::load().map_err(|error| error.to_string())?;
312    data["cookie"] = json!(value);
313    settings::save(&data).map_err(|error| error.to_string())?;
314    println!("Cookie 已保存: {}", settings::settings_file().display());
315    Ok(())
316}
317
318fn cookie_status(offline: bool) -> Result<(), String> {
319    let data = settings::load().map_err(|error| error.to_string())?;
320    let value = data
321        .get("cookie")
322        .and_then(Value::as_str)
323        .unwrap_or("")
324        .trim();
325    if value.is_empty() {
326        println!("未保存 Cookie");
327        return Ok(());
328    }
329    if !cookie::validate(value) {
330        return Err(format!(
331            "已保存 Cookie,但格式无效: {}",
332            settings::settings_file().display()
333        ));
334    }
335    if offline {
336        println!("Cookie 格式有效: {}", settings::settings_file().display());
337        return Ok(());
338    }
339    println!("正在检查网页端请求连通性...");
340    if cookie::probe(value) {
341        println!(
342            "Cookie 可用于网页端请求: {}",
343            settings::settings_file().display()
344        );
345        Ok(())
346    } else {
347        Err("Cookie 已保存,但网页端请求连通性检查失败".to_owned())
348    }
349}
350
351fn cookie_logout() -> Result<(), String> {
352    let mut data = settings::load().map_err(|error| error.to_string())?;
353    data["cookie"] = json!("");
354    settings::save(&data).map_err(|error| error.to_string())?;
355    println!("已清除 Cookie");
356    Ok(())
357}
358
359fn save_openapi(updates: Map<String, Value>) -> Result<(), String> {
360    let mut data = settings::load().map_err(|error| error.to_string())?;
361    let openapi = data
362        .get_mut("openapi")
363        .and_then(Value::as_object_mut)
364        .ok_or_else(|| "openapi 配置格式无效".to_owned())?;
365    openapi.extend(updates);
366    settings::save(&data).map_err(|error| error.to_string())
367}
368
369fn extract_token_fields(data: &Value) -> Map<String, Value> {
370    let source = data
371        .get("data")
372        .filter(|value| value.is_object())
373        .unwrap_or(data);
374    [
375        ("access_token", "accessToken"),
376        ("refresh_token", "refreshToken"),
377        ("open_id", "openId"),
378        ("expires_in", "expiresIn"),
379    ]
380    .into_iter()
381    .filter_map(|(source_key, target_key)| {
382        source
383            .get(source_key)
384            .cloned()
385            .map(|value| (target_key.to_owned(), value))
386    })
387    .collect()
388}
389
390fn saved_string(values: &Map<String, Value>, key: &str) -> Option<String> {
391    values
392        .get(key)
393        .and_then(Value::as_str)
394        .filter(|value| !value.is_empty())
395        .map(str::to_owned)
396}
397
398fn missing_client_key(show_qr: bool) -> String {
399    let qr_line =
400        show_qr.then_some("--qr 只会把官方 OAuth 授权链接渲染成二维码,仍然需要 client_key。\n");
401    format!(
402        "当前命令是官方 OpenAPI OAuth 授权,需要开放平台 client_key。\n{}这不是网页端 Cookie 扫码登录,不能直接生成可保存 Cookie 的登录二维码。\n\n可选方案:\n  1. 官方 OpenAPI:传入 --client-key,或设置 DOUYIN_CLIENT_KEY。\n  2. 网页端采集:从浏览器复制 Cookie 后运行:\n     douyin auth cookie-login --cookie 'sessionid=...; ttwid=...'",
403        qr_line.unwrap_or("")
404    )
405}
406
407fn random_state() -> Result<String, String> {
408    let mut bytes = [0_u8; 18];
409    SystemRandom::new()
410        .fill(&mut bytes)
411        .map_err(|_| "无法生成 OAuth state".to_owned())?;
412    Ok(URL_SAFE_NO_PAD.encode(bytes))
413}
414
415fn print_qr(value: &str) -> Result<(), String> {
416    let code = QrCode::new(value.as_bytes()).map_err(|error| error.to_string())?;
417    let width = code.width();
418    println!();
419    for y in (0..width).step_by(2) {
420        let mut line = String::new();
421        for x in 0..width {
422            let top = code[(x, y)] == Color::Dark;
423            let bottom = y + 1 < width && code[(x, y + 1)] == Color::Dark;
424            line.push(match (top, bottom) {
425                (true, true) => '█',
426                (true, false) => '▀',
427                (false, true) => '▄',
428                (false, false) => ' ',
429            });
430        }
431        println!(" {line} ");
432    }
433    println!();
434    Ok(())
435}
436
437fn wait_for_code(
438    host: &str,
439    port: u16,
440    expected_state: Option<&str>,
441    timeout: Duration,
442) -> Result<String, String> {
443    let listener = TcpListener::bind((host, port))
444        .map_err(|_| format!("无法监听 {host}:{port},请换一个 --callback-port"))?;
445    listener
446        .set_nonblocking(true)
447        .map_err(|error| error.to_string())?;
448    let started = Instant::now();
449    while started.elapsed() < timeout {
450        match listener.accept() {
451            Ok((mut stream, _)) => return handle_callback(&mut stream, expected_state),
452            Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
453                thread::sleep(Duration::from_millis(50));
454            }
455            Err(error) => return Err(error.to_string()),
456        }
457    }
458    Err("等待授权回调超时,未获取到 code".to_owned())
459}
460
461fn handle_callback(stream: &mut TcpStream, expected_state: Option<&str>) -> Result<String, String> {
462    stream
463        .set_read_timeout(Some(Duration::from_secs(5)))
464        .map_err(|error| error.to_string())?;
465    let mut buffer = [0_u8; 8192];
466    let count = stream
467        .read(&mut buffer)
468        .map_err(|error| error.to_string())?;
469    let request = String::from_utf8_lossy(&buffer[..count]);
470    let target = request
471        .lines()
472        .next()
473        .and_then(|line| line.split_whitespace().nth(1))
474        .ok_or_else(|| "授权回调请求无效".to_owned())?;
475    let url = reqwest::Url::parse(&format!("http://localhost{target}"))
476        .map_err(|error| format!("授权回调 URL 无效: {error}"))?;
477    if url.path() != "/callback" {
478        send_html(stream, 404, "未找到回调路径")?;
479        return Err("授权回调路径无效".to_owned());
480    }
481    let params: HashMap<_, _> = url.query_pairs().into_owned().collect();
482    if let Some(error) = params
483        .get("error")
484        .or_else(|| params.get("error_description"))
485    {
486        send_html(stream, 400, "授权失败,可以关闭此页面并返回终端。")?;
487        return Err(format!("授权失败: {error}"));
488    }
489    if expected_state
490        .is_some_and(|expected| params.get("state").map(String::as_str) != Some(expected))
491    {
492        send_html(stream, 400, "state 不匹配。")?;
493        return Err("授权回调 state 不匹配,已拒绝".to_owned());
494    }
495    let Some(code) = params.get("code") else {
496        send_html(stream, 400, "回调缺少 code。")?;
497        return Err("授权回调缺少 code".to_owned());
498    };
499    send_html(stream, 200, "授权完成,可以关闭此页面并返回终端。")?;
500    Ok(code.to_owned())
501}
502
503fn send_html(stream: &mut TcpStream, status: u16, body: &str) -> Result<(), String> {
504    let content =
505        format!("<!doctype html><meta charset='utf-8'><title>Douyin CLI</title><p>{body}</p>");
506    let reason = if status == 200 { "OK" } else { "Bad Request" };
507    write!(
508        stream,
509        "HTTP/1.1 {status} {reason}\r\nContent-Type: text/html; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{content}",
510        content.len()
511    )
512    .map_err(|error| error.to_string())
513}
514
515fn print_json(value: &Value) -> Result<(), String> {
516    println!(
517        "{}",
518        serde_json::to_string_pretty(value).map_err(|error| error.to_string())?
519    );
520    Ok(())
521}
522
523#[cfg(test)]
524mod tests {
525    use super::{extract_token_fields, missing_client_key};
526    use serde_json::json;
527
528    #[test]
529    fn extracts_nested_token_fields() {
530        let fields = extract_token_fields(&json!({"data": {
531            "access_token": "access", "refresh_token": "refresh", "open_id": "open", "expires_in": 1
532        }}));
533        assert_eq!(fields["accessToken"], "access");
534        assert_eq!(fields["openId"], "open");
535    }
536
537    #[test]
538    fn missing_key_message_explains_cookie_alternative() {
539        let message = missing_client_key(true);
540        assert!(message.contains("官方 OpenAPI OAuth 授权"));
541        assert!(message.contains("--qr 只会把官方 OAuth 授权链接渲染成二维码"));
542        assert!(message.contains("douyin auth cookie-login --cookie"));
543    }
544}