use std::{collections::HashMap, path::Path, time::Duration};
use anyhow::{Context, Result, bail};
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
use chrono::{Duration as ChronoDuration, Utc};
use rand::RngCore;
use reqwest::blocking::Client;
use serde::Deserialize;
use sha2::{Digest, Sha256};
use tiny_http::{Header, Response, Server};
use url::Url;
use crate::{
api::LingvaApi,
credentials::{self, Credentials},
};
pub const DEFAULT_PLATFORM_URL: &str = "https://lingva.dev";
pub const DEFAULT_AUTH_PORT: u16 = 53_682;
pub const DEFAULT_AUTH_SCOPES: &str = "openid email profile";
#[derive(Debug)]
pub struct LoginOptions<'a> {
pub platform_url: &'a str,
pub issuer_url: Option<&'a str>,
pub authorize_url: Option<&'a str>,
pub token_url: Option<&'a str>,
pub client_id: Option<&'a str>,
pub scopes: &'a str,
pub port: u16,
pub open_browser: bool,
}
#[derive(Debug, Deserialize)]
struct TokenResponse {
access_token: String,
refresh_token: Option<String>,
id_token: Option<String>,
token_type: Option<String>,
scope: Option<String>,
expires_in: Option<i64>,
}
pub fn login(
api: &LingvaApi,
options: LoginOptions<'_>,
credentials_path: &Path,
) -> Result<Credentials> {
let discovery = if options.issuer_url.is_none() && options.authorize_url.is_none() {
Some(api.discover_auth(options.platform_url)?)
} else {
None
};
let issuer = options
.issuer_url
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| discovery.as_ref().map(|value| value.issuer_url.as_str()));
let authorize_url = options
.authorize_url
.map(str::to_owned)
.or_else(|| issuer.map(|value| format!("{}/oauth2/authorize", value.trim_end_matches('/'))))
.context(
"Lingva auth login requires --issuer-url or both --authorize-url and --token-url.",
)?;
let token_url = options
.token_url
.map(str::to_owned)
.or_else(|| issuer.map(|value| format!("{}/oauth2/token", value.trim_end_matches('/'))))
.context(
"Lingva auth login requires --issuer-url or both --authorize-url and --token-url.",
)?;
let client_id = options
.client_id
.map(str::to_owned)
.or_else(|| discovery.map(|value| value.client_id))
.context("Lingva auth login requires --client-id when discovery is disabled.")?;
let verifier = random_url_token(48);
let challenge = URL_SAFE_NO_PAD.encode(Sha256::digest(verifier.as_bytes()));
let expected_state = random_url_token(24);
let server = Server::http(("127.0.0.1", options.port))
.map_err(|error| anyhow::anyhow!(error.to_string()))?;
let actual_port = server
.server_addr()
.to_ip()
.map(|value| value.port())
.unwrap_or(options.port);
let redirect_uri = format!("http://127.0.0.1:{actual_port}/callback");
let mut login_url = Url::parse(&authorize_url)?;
login_url
.query_pairs_mut()
.append_pair("client_id", &client_id)
.append_pair("code_challenge", &challenge)
.append_pair("code_challenge_method", "S256")
.append_pair("redirect_uri", &redirect_uri)
.append_pair("response_type", "code")
.append_pair("scope", options.scopes)
.append_pair("state", &expected_state);
println!("Open this URL to authenticate Lingva CLI:\n{login_url}\n");
if options.open_browser {
open::that(login_url.as_str())
.context("Unable to open the browser for Lingva authentication.")?;
}
let request = server
.recv_timeout(Duration::from_secs(300))?
.context("Lingva authentication callback timed out.")?;
let callback = Url::parse(&format!("http://127.0.0.1{}", request.url()))?;
let query = callback
.query_pairs()
.into_owned()
.collect::<HashMap<_, _>>();
let valid = query.get("state") == Some(&expected_state)
&& query.get("code").is_some_and(|value| !value.is_empty());
let html = if valid {
"<h1>Lingva CLI is authenticated</h1><p>You can close this tab.</p>"
} else {
"<h1>Lingva CLI authentication failed</h1><p>Return to the terminal.</p>"
};
request
.respond(Response::from_string(html).with_header(
Header::from_bytes("content-type", "text/html; charset=utf-8").unwrap(),
))?;
if !valid {
bail!("Lingva authentication callback was invalid.");
}
let tokens = Client::new()
.post(&token_url)
.form(&[
("client_id", client_id.as_str()),
("code", query["code"].as_str()),
("code_verifier", verifier.as_str()),
("grant_type", "authorization_code"),
("redirect_uri", redirect_uri.as_str()),
])
.send()?
.error_for_status()?
.json::<TokenResponse>()?;
let now = Utc::now();
let credentials = Credentials {
access_token: Some(tokens.access_token),
refresh_token: tokens.refresh_token,
id_token: tokens.id_token,
token_type: Some(tokens.token_type.unwrap_or_else(|| "Bearer".into())),
scope: Some(tokens.scope.unwrap_or_else(|| options.scopes.into())),
client_id: Some(client_id),
token_url: Some(token_url),
obtained_at: now.to_rfc3339(),
expires_at: tokens
.expires_in
.map(|seconds| (now + ChronoDuration::seconds(seconds)).to_rfc3339()),
platform_url: Some(options.platform_url.into()),
..Default::default()
};
credentials::save(credentials_path, &credentials)?;
Ok(credentials)
}
fn random_url_token(size: usize) -> String {
let mut bytes = vec![0; size];
rand::rng().fill_bytes(&mut bytes);
URL_SAFE_NO_PAD.encode(bytes)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn random_tokens_are_url_safe_have_expected_size_and_do_not_repeat() {
let first = random_url_token(48);
let second = random_url_token(48);
assert_eq!(first.len(), 64);
assert_ne!(first, second);
assert!(
first
.chars()
.all(|value| value.is_ascii_alphanumeric() || matches!(value, '-' | '_'))
);
}
#[test]
fn default_auth_contract_is_explicit() {
assert_eq!(DEFAULT_PLATFORM_URL, "https://lingva.dev");
assert!(
DEFAULT_AUTH_SCOPES
.split_whitespace()
.all(|scope| { matches!(scope, "openid" | "email" | "profile") })
);
}
}