Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
164 changes: 158 additions & 6 deletions src/auth/google_oauth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,22 @@ use tokio::net::TcpListener;

use super::oauth::OAuthCredentials;

/// Desktop app client (loopback redirect flow).
const CLIENT_ID: &str = "701529528334-otljpqp2bjvhm7lp2eqktu5ja8uo05g6.apps.googleusercontent.com";
const CLIENT_SECRET: &str = "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";

/// TV / Limited Input client (device code flow — works over SSH).
const DEVICE_CLIENT_ID: &str =
"701529528334-7buapusrvqo9ogqio29gd8i3ka96j3qg.apps.googleusercontent.com";
const DEVICE_CLIENT_SECRET: &str = "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";

const AUTHORIZE_URL: &str = "https://accounts.google.com/o/oauth2/v2/auth";
const TOKEN_URL: &str = "https://oauth2.googleapis.com/token";
const SCOPES: &str = "https://www.googleapis.com/auth/generative-language";
const DEVICE_CODE_URL: &str = "https://oauth2.googleapis.com/device/code";
/// OAuth scopes for both flows. The Gemini API's generateContent has no scope
/// requirements (confirmed via Google's discovery doc), so any valid OAuth
/// token works. We request only the minimum needed for authentication.
const SCOPES: &str = "openid email";

/// 5-minute buffer (in ms) subtracted from token expiry.
const EXPIRY_BUFFER_MS: u64 = 5 * 60 * 1000;
Expand Down Expand Up @@ -221,15 +232,25 @@ pub async fn exchange_code(code: &str, verifier: &str, port: u16) -> Result<OAut
access: data.access_token,
refresh,
expires: expiry_with_buffer(data.expires_in),
client_hint: None,
})
}

/// Refresh an expired access token.
pub async fn refresh_token(refresh: &str) -> Result<OAuthCredentials> {
///
/// `client_hint` selects which OAuth client to use:
/// - `Some("device")` → TV / Limited Input client (device code flow)
/// - anything else → Desktop client (loopback flow)
pub async fn refresh_token(refresh: &str, client_hint: Option<&str>) -> Result<OAuthCredentials> {
let (cid, csecret) = match client_hint {
Some("device") => (DEVICE_CLIENT_ID, DEVICE_CLIENT_SECRET),
_ => (CLIENT_ID, CLIENT_SECRET),
};

let params = [
("grant_type", "refresh_token"),
("client_id", CLIENT_ID),
("client_secret", CLIENT_SECRET),
("client_id", cid),
("client_secret", csecret),
("refresh_token", refresh),
];

Expand All @@ -248,6 +269,7 @@ pub async fn refresh_token(refresh: &str) -> Result<OAuthCredentials> {
access: data.access_token,
refresh: data.refresh_token.unwrap_or_else(|| refresh.to_string()),
expires: expiry_with_buffer(data.expires_in),
client_hint: client_hint.map(String::from),
})
}

Expand All @@ -258,6 +280,136 @@ struct TokenResponse {
expires_in: u64,
}

// ---------------------------------------------------------------------------
// Device code flow (for SSH / headless environments)
// ---------------------------------------------------------------------------

/// Response from Google's device code endpoint.
#[derive(serde::Deserialize)]
struct DeviceCodeResponse {
device_code: String,
user_code: String,
verification_url: String,
expires_in: u64,
interval: u64,
}

/// Response while polling — may be a pending status or final tokens.
#[derive(serde::Deserialize)]
struct DevicePollResponse {
/// Present on error (e.g. "authorization_pending", "slow_down", "access_denied").
error: Option<String>,
/// Present on success.
access_token: Option<String>,
refresh_token: Option<String>,
expires_in: Option<u64>,
}

/// Initiate the device code flow. Returns the user code and verification URL
/// for display, plus the device code for polling.
pub async fn device_code_authorize() -> Result<DeviceAuth> {
let params = [("client_id", DEVICE_CLIENT_ID), ("scope", SCOPES)];

let client = reqwest::Client::new();
let resp = client.post(DEVICE_CODE_URL).form(&params).send().await?;

if !resp.status().is_success() {
let text = resp.text().await.unwrap_or_default();
bail!("Google device code request failed: {text}");
}

let data: DeviceCodeResponse = resp.json().await?;

Ok(DeviceAuth {
device_code: data.device_code,
user_code: data.user_code,
verification_url: data.verification_url,
expires_in: data.expires_in,
interval: data.interval,
})
}

/// Everything needed to complete the device code flow.
pub struct DeviceAuth {
pub device_code: String,
pub user_code: String,
pub verification_url: String,
pub expires_in: u64,
pub interval: u64,
}

/// Poll Google's token endpoint until the user approves (or the code expires).
pub async fn poll_device_token(auth: &DeviceAuth) -> Result<OAuthCredentials> {
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(auth.expires_in);
let mut interval = std::time::Duration::from_secs(auth.interval.max(5));

let client = reqwest::Client::new();

loop {
tokio::time::sleep(interval).await;

if std::time::Instant::now() > deadline {
bail!("device code expired — please try again");
}

let params = [
("client_id", DEVICE_CLIENT_ID),
("client_secret", DEVICE_CLIENT_SECRET),
("device_code", auth.device_code.as_str()),
("grant_type", "urn:ietf:params:oauth:grant-type:device_code"),
];

let resp = client.post(TOKEN_URL).form(&params).send().await?;
let data: DevicePollResponse = resp.json().await?;

match data.error.as_deref() {
Some("authorization_pending") => continue,
Some("slow_down") => {
// Back off by 5 seconds as required by Google
interval += std::time::Duration::from_secs(5);
continue;
}
Some(err) => bail!("Google device auth failed: {err}"),
None => {
// Success — tokens present
let access = data
.access_token
.ok_or_else(|| anyhow::anyhow!("missing access_token in device response"))?;
let refresh = data.refresh_token.ok_or_else(|| {
anyhow::anyhow!(
"Google did not return a refresh token. \
Try revoking access at https://myaccount.google.com/permissions \
and logging in again."
)
})?;
let expires_in = data.expires_in.unwrap_or(3600);

return Ok(OAuthCredentials {
access,
refresh,
expires: expiry_with_buffer(expires_in),
client_hint: Some("device".into()),
});
}
}
}
}

/// Detect whether we're in a headless / SSH environment where loopback
/// redirect won't work.
pub fn is_headless() -> bool {
// SSH session — browser redirect to 127.0.0.1 on remote won't work
if std::env::var("SSH_CONNECTION").is_ok() || std::env::var("SSH_TTY").is_ok() {
return true;
}
// No display server on Linux
#[cfg(target_os = "linux")]
if std::env::var("DISPLAY").is_err() && std::env::var("WAYLAND_DISPLAY").is_err() {
return true;
}
false
}

/// Decode a percent-encoded string (e.g. `hello%20world` → `hello world`).
fn urldecode(s: &str) -> String {
let mut out = Vec::with_capacity(s.len());
Expand Down Expand Up @@ -421,8 +573,8 @@ mod tests {
#[test]
fn urlencoded_encodes_slashes() {
assert_eq!(
urlencoded("https://www.googleapis.com/auth/generative-language"),
"https%3A%2F%2Fwww.googleapis.com%2Fauth%2Fgenerative-language"
urlencoded("https://example.com/foo/bar"),
"https%3A%2F%2Fexample.com%2Ffoo%2Fbar"
);
}

Expand Down
11 changes: 11 additions & 0 deletions src/auth/oauth.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,12 @@ pub struct OAuthCredentials {
pub refresh: String,
/// Expiration timestamp in milliseconds since epoch.
pub expires: u64,
/// Optional hint identifying which OAuth client issued these tokens.
/// Used by providers with multiple OAuth clients to select the correct
/// client configuration during token refresh. `None` means the default
/// client configuration will be used.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub client_hint: Option<String>,
}

impl OAuthCredentials {
Expand Down Expand Up @@ -114,6 +120,7 @@ pub async fn exchange_code(auth_code_raw: &str, verifier: &str) -> Result<OAuthC
access: data.access_token,
refresh: data.refresh_token,
expires,
client_hint: None,
})
}

Expand Down Expand Up @@ -146,6 +153,7 @@ pub async fn refresh_token(refresh: &str) -> Result<OAuthCredentials> {
access: data.access_token,
refresh: data.refresh_token,
expires,
client_hint: None,
})
}

Expand Down Expand Up @@ -289,6 +297,7 @@ mod tests {
access: "token".to_string(),
refresh: "refresh".to_string(),
expires: now_ms() + 3_600_000, // 1 hour from now
client_hint: None,
};
assert!(!creds.is_expired());
}
Expand All @@ -299,6 +308,7 @@ mod tests {
access: "token".to_string(),
refresh: "refresh".to_string(),
expires: 1000, // epoch + 1 second
client_hint: None,
};
assert!(creds.is_expired());
}
Expand All @@ -309,6 +319,7 @@ mod tests {
access: "token".to_string(),
refresh: "refresh".to_string(),
expires: 0,
client_hint: None,
};
assert!(creds.is_expired());
}
Expand Down
8 changes: 7 additions & 1 deletion src/auth/storage.rs
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,13 @@ impl AuthStorage {
Credential::OAuth(mut oauth) => {
if oauth.is_expired() {
let refreshed = match provider {
"google" => super::google_oauth::refresh_token(&oauth.refresh).await?,
"google" => {
super::google_oauth::refresh_token(
&oauth.refresh,
oauth.client_hint.as_deref(),
)
.await?
}
_ => super::oauth::refresh_token(&oauth.refresh).await?,
};
oauth = refreshed.clone();
Expand Down
30 changes: 30 additions & 0 deletions src/provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,7 @@ mod anthropic_provider {
mod google_provider {
use super::*;
use crate::auth::google_oauth;
use crate::auth::storage::Credential;
use crate::thinker::gemini::GeminiThinker;

pub struct Google;
Expand Down Expand Up @@ -131,6 +132,17 @@ mod google_provider {
}

async fn login(&self, db_path: &str) -> Result<()> {
if google_oauth::is_headless() {
self.login_device_code(db_path).await
} else {
self.login_loopback(db_path).await
}
}
}

impl Google {
/// Loopback redirect flow — opens browser, Google redirects to localhost.
async fn login_loopback(&self, db_path: &str) -> Result<()> {
let (auth_result, listener) = google_oauth::prepare_authorize().await?;

let _ = open::that(&auth_result.url);
Expand All @@ -154,6 +166,24 @@ mod google_provider {
.await?;
Ok(())
}

/// Device code flow — works over SSH / headless.
async fn login_device_code(&self, db_path: &str) -> Result<()> {
println!("Headless environment detected — using device code flow.\n");

let auth = google_oauth::device_code_authorize().await?;

println!("Go to: {}\n", auth.verification_url);
println!("Enter code: {}\n", auth.user_code);
println!("Waiting for approval...");

let creds = google_oauth::poll_device_token(&auth).await?;

let storage = AuthStorage::open(db_path)?;
storage.set(self.id(), Credential::OAuth(creds))?;

Ok(())
}
}
}

Expand Down
2 changes: 2 additions & 0 deletions tests/auth_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ fn set_and_get_oauth() {
access: "access-token".to_string(),
refresh: "refresh-token".to_string(),
expires: 9999999999999,
client_hint: None,
};
storage.set("anthropic", Credential::OAuth(oauth)).unwrap();

Expand Down Expand Up @@ -190,6 +191,7 @@ async fn get_api_key_from_oauth_non_expired() {
access: "sk-ant-oat01-valid".to_string(),
refresh: "refresh".to_string(),
expires: u64::MAX, // far future
client_hint: None,
};
storage.set("anthropic", Credential::OAuth(oauth)).unwrap();

Expand Down
1 change: 1 addition & 0 deletions tests/login_test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,7 @@ fn build_provider_detects_oauth_credentials() {
access: "token".to_string(),
refresh: "refresh".to_string(),
expires: u64::MAX,
client_hint: None,
}),
)
.unwrap();
Expand Down