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