use serde::Serialize;
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use tokio::sync::watch;
use tokio_util::sync::CancellationToken;
const DEFAULT_AUTH_URL: &str = "https://openrouter.ai/auth";
const DEFAULT_EXCHANGE_URL: &str = "https://openrouter.ai/api/v1/auth/keys";
pub const TEST_AUTHORIZATION_BASE_URL_ENV: &str = "OPENROUTER_AUTHORIZATION_BASE_URL";
pub const TEST_EXCHANGE_URL_ENV: &str = "OPENROUTER_EXCHANGE_URL";
const DEFAULT_TIMEOUT_SECS: u64 = 300;
const MAX_TIMEOUT_SECS: u64 = 900;
#[derive(Debug, Clone, Serialize)]
pub struct PendingStatus {
pub flow_id: String,
pub deadline_unix_ms: u64,
}
#[derive(Debug, Clone, Serialize)]
pub struct TerminalResult {
pub flow_id: String,
pub kind: String,
pub message: String,
}
#[derive(Debug, Clone, Serialize)]
pub struct Status {
pub authority_generation: u64,
pub state: String,
pub effective_source: String,
pub pasted_key_exists: bool,
pub oauth_key_exists: bool,
pub pending: Option<PendingStatus>,
pub last_result: Option<TerminalResult>,
}
#[derive(Debug, Clone, Serialize)]
pub struct StartResult {
pub authority_generation: u64,
pub authorize_url: String,
pub flow_id: String,
pub deadline_unix_ms: u64,
}
struct PendingFlow {
flow_id: String,
action_generation: u64,
deadline_unix_ms: u64,
cancel: CancellationToken,
}
#[derive(Clone, Copy, Default)]
struct CredentialPresence {
environment: bool,
pasted: bool,
oauth: bool,
}
struct Manager {
action_generation: u64,
pending: Option<PendingFlow>,
last_result: Option<TerminalResult>,
credentials: CredentialPresence,
credential_revision: u64,
}
impl Default for Manager {
fn default() -> Self {
let action_generation = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_nanos()
.try_into()
.unwrap_or(u64::MAX / 2);
Self {
action_generation,
pending: None,
last_result: None,
credentials: credential_presence_from_vault(),
credential_revision: 0,
}
}
}
impl Manager {
fn advance_action_generation(&mut self) -> u64 {
self.action_generation = self
.action_generation
.checked_add(1)
.expect("OpenRouter authority generation exhausted");
self.action_generation
}
fn set_pasted_exists(&mut self, exists: bool) {
if self.credentials.pasted != exists {
self.credentials.pasted = exists;
self.credential_revision = self.credential_revision.wrapping_add(1);
}
}
fn set_oauth_exists(&mut self, exists: bool) {
if self.credentials.oauth != exists {
self.credentials.oauth = exists;
self.credential_revision = self.credential_revision.wrapping_add(1);
}
}
fn replace_credential_presence(&mut self, presence: CredentialPresence) {
if self.credentials.environment != presence.environment
|| self.credentials.pasted != presence.pasted
|| self.credentials.oauth != presence.oauth
{
self.credentials = presence;
self.credential_revision = self.credential_revision.wrapping_add(1);
}
}
}
#[derive(Clone, Default)]
struct ManagerTransactionHook {
#[cfg(test)]
before_lock: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
inside_lock_entered: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
inside_lock_release: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
after_reserve_entered: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
after_reserve_release: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
before_return_entered: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
before_return_release: Option<Arc<std::sync::Barrier>>,
}
impl ManagerTransactionHook {
fn before_lock(&self) {
#[cfg(test)]
if let Some(barrier) = &self.before_lock {
barrier.wait();
}
}
fn inside_lock(&self) {
#[cfg(test)]
{
if let Some(barrier) = &self.inside_lock_entered {
barrier.wait();
}
if let Some(barrier) = &self.inside_lock_release {
barrier.wait();
}
}
}
fn after_reserve(&self) {
#[cfg(test)]
{
if let Some(barrier) = &self.after_reserve_entered {
barrier.wait();
}
if let Some(barrier) = &self.after_reserve_release {
barrier.wait();
}
}
}
fn before_return(&self) {
#[cfg(test)]
{
if let Some(barrier) = &self.before_return_entered {
barrier.wait();
}
if let Some(barrier) = &self.before_return_release {
barrier.wait();
}
}
}
}
#[derive(Clone, Default)]
struct DisconnectHook {
#[cfg(test)]
before_delete_entered: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
before_delete_release: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
forced_delete_error: Option<String>,
}
impl DisconnectHook {
fn before_delete(&self) {
#[cfg(test)]
{
if let Some(barrier) = &self.before_delete_entered {
barrier.wait();
}
if let Some(barrier) = &self.before_delete_release {
barrier.wait();
}
}
}
fn delete_oauth_key(&self) -> Result<(), String> {
self.before_delete();
#[cfg(test)]
if let Some(error) = &self.forced_delete_error {
return Err(error.clone());
}
delete_oauth_key()
}
}
#[derive(Clone, Default)]
struct CredentialFinishHook {
#[cfg(test)]
before_gate_entered: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
after_gate_entered: Option<Arc<std::sync::Barrier>>,
#[cfg(test)]
after_gate_release: Option<Arc<std::sync::Barrier>>,
}
impl CredentialFinishHook {
fn before_gate(&self) {
#[cfg(test)]
if let Some(barrier) = &self.before_gate_entered {
barrier.wait();
}
}
fn after_gate(&self) {
#[cfg(test)]
{
if let Some(barrier) = &self.after_gate_entered {
barrier.wait();
}
if let Some(barrier) = &self.after_gate_release {
barrier.wait();
}
}
}
}
#[derive(Clone, Default)]
struct PastedMutationHook {
#[cfg(test)]
before_gate_entered: Option<Arc<std::sync::Barrier>>,
}
impl PastedMutationHook {
fn before_gate(&self) {
#[cfg(test)]
if let Some(barrier) = &self.before_gate_entered {
barrier.wait();
}
}
}
fn manager() -> &'static Arc<Mutex<Manager>> {
static MANAGER: OnceLock<Arc<Mutex<Manager>>> = OnceLock::new();
MANAGER.get_or_init(|| Arc::new(Mutex::new(Manager::default())))
}
fn credential_io_gate() -> &'static tokio::sync::Mutex<()> {
static GATE: OnceLock<tokio::sync::Mutex<()>> = OnceLock::new();
GATE.get_or_init(|| tokio::sync::Mutex::new(()))
}
fn authority_updates() -> &'static watch::Sender<Option<Status>> {
static UPDATES: OnceLock<watch::Sender<Option<Status>>> = OnceLock::new();
UPDATES.get_or_init(|| watch::channel(None).0)
}
pub fn subscribe_authority_updates() -> watch::Receiver<Option<Status>> {
authority_updates().subscribe()
}
fn publish_status(status: &Status) {
authority_updates().send_replace(Some(status.clone()));
}
fn now_unix_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
fn terminal(flow_id: &str, kind: &str, message: &str) -> TerminalResult {
TerminalResult {
flow_id: flow_id.to_string(),
kind: kind.to_string(),
message: message.to_string(),
}
}
#[cfg(not(test))]
fn oauth_key_exists() -> bool {
car_inference::openrouter::oauth_key_exists()
}
#[cfg(test)]
fn test_oauth_slot() -> &'static Mutex<Option<String>> {
static SLOT: OnceLock<Mutex<Option<String>>> = OnceLock::new();
SLOT.get_or_init(|| Mutex::new(None))
}
#[cfg(test)]
fn oauth_key_exists() -> bool {
test_oauth_slot()
.lock()
.unwrap_or_else(|p| p.into_inner())
.is_some()
}
#[cfg(not(test))]
fn store_oauth_key(key: &str) -> Result<(), String> {
car_inference::openrouter::store_oauth_credential(key).map_err(|error| error.to_string())
}
#[cfg(test)]
fn store_oauth_key(key: &str) -> Result<(), String> {
*test_oauth_slot().lock().unwrap_or_else(|p| p.into_inner()) = Some(key.to_string());
Ok(())
}
#[cfg(not(test))]
fn delete_oauth_key() -> Result<(), String> {
car_inference::openrouter::delete_oauth_credential().map_err(|error| error.to_string())
}
#[cfg(test)]
fn delete_oauth_key() -> Result<(), String> {
*test_oauth_slot().lock().unwrap_or_else(|p| p.into_inner()) = None;
Ok(())
}
fn credential_presence_from_vault() -> CredentialPresence {
CredentialPresence {
environment: car_inference::openrouter::environment_key_exists(),
pasted: car_inference::openrouter::pasted_key_exists(),
oauth: oauth_key_exists(),
}
}
fn status_from_manager_with_presence(guard: &Manager, presence: CredentialPresence) -> Status {
let effective = if presence.environment {
Some(car_inference::openrouter::CredentialSource::Env)
} else if presence.pasted {
Some(car_inference::openrouter::CredentialSource::Pasted)
} else if presence.oauth {
Some(car_inference::openrouter::CredentialSource::Oauth)
} else {
None
};
let pending = guard.pending.as_ref().map(|flow| PendingStatus {
flow_id: flow.flow_id.clone(),
deadline_unix_ms: flow.deadline_unix_ms,
});
Status {
authority_generation: guard.action_generation,
state: if pending.is_some() {
"pending"
} else if effective.is_some() || presence.oauth {
"connected"
} else {
"idle"
}
.to_string(),
effective_source: effective
.map(|source| source.as_str().to_string())
.or_else(|| presence.oauth.then(|| "oauth".to_string()))
.unwrap_or_else(|| "none".to_string()),
pasted_key_exists: presence.pasted,
oauth_key_exists: presence.oauth,
pending,
last_result: guard.last_result.clone(),
}
}
fn status_from_manager(guard: &Manager) -> Status {
status_from_manager_with_presence(guard, guard.credentials)
}
pub fn status() -> Status {
let expected = {
let guard = manager().lock().unwrap_or_else(|p| p.into_inner());
(guard.action_generation, guard.credential_revision)
};
let observed = credential_presence_from_vault();
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
if (guard.action_generation, guard.credential_revision) == expected {
guard.replace_credential_presence(observed);
}
status_from_manager(&guard)
}
fn complete_pasted_change(reserved_generation: u64, exists_after_success: Option<bool>) -> Status {
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
let is_current = guard.action_generation == reserved_generation;
if is_current {
if let Some(exists) = exists_after_success {
guard.set_pasted_exists(exists);
}
}
let status = status_from_manager(&guard);
drop(guard);
if is_current {
publish_status(&status);
}
status
}
pub fn is_pending(flow_id: &str) -> bool {
manager()
.lock()
.unwrap_or_else(|p| p.into_inner())
.pending
.as_ref()
.is_some_and(|pending| pending.flow_id == flow_id)
}
pub fn terminal_status_for_flow(flow_id: &str) -> Option<Status> {
terminal_status_for_flow_with_hook(flow_id, ManagerTransactionHook::default())
}
fn terminal_status_for_flow_with_hook(
flow_id: &str,
hook: ManagerTransactionHook,
) -> Option<Status> {
hook.before_lock();
let guard = manager().lock().unwrap_or_else(|p| p.into_inner());
let matches_terminal = guard.pending.is_none()
&& guard
.last_result
.as_ref()
.is_some_and(|result| result.flow_id == flow_id);
if !matches_terminal {
return None;
}
hook.inside_lock();
Some(status_from_manager(&guard))
}
fn replace_pending_result(kind: &str, message: &str) -> Status {
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
guard.advance_action_generation();
if let Some(flow) = guard.pending.take() {
flow.cancel.cancel();
guard.last_result = Some(terminal(&flow.flow_id, kind, message));
}
let status = status_from_manager(&guard);
drop(guard);
publish_status(&status);
status
}
fn supersede_for_pasted_change() -> Status {
replace_pending_result(
"superseded",
"OpenRouter connect was superseded by a pasted-key change",
)
}
pub async fn mutate_pasted_credential<T, F>(
exists_on_success: bool,
operation: F,
) -> (Result<T, String>, Status)
where
F: FnOnce() -> Result<T, String>,
{
mutate_pasted_credential_with_hook(exists_on_success, operation, PastedMutationHook::default())
.await
}
async fn mutate_pasted_credential_with_hook<T, F>(
exists_on_success: bool,
operation: F,
hook: PastedMutationHook,
) -> (Result<T, String>, Status)
where
F: FnOnce() -> Result<T, String>,
{
hook.before_gate();
let _credential_io = credential_io_gate().lock().await;
let observed = credential_presence_from_vault();
manager()
.lock()
.unwrap_or_else(|p| p.into_inner())
.replace_credential_presence(observed);
let reserved = supersede_for_pasted_change();
let result = operation();
let final_status = complete_pasted_change(
reserved.authority_generation,
result.is_ok().then_some(exists_on_success),
);
(result, final_status)
}
pub fn cancel(flow_id: Option<&str>) -> Status {
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
let matches = guard
.pending
.as_ref()
.is_some_and(|flow| flow_id.is_none_or(|id| id == flow.flow_id));
if matches {
guard.advance_action_generation();
let flow = guard.pending.take().expect("checked above");
flow.cancel.cancel();
guard.last_result = Some(terminal(
&flow.flow_id,
"cancelled",
"OpenRouter connect was cancelled",
));
let status = status_from_manager(&guard);
drop(guard);
publish_status(&status);
return status;
}
let status = status_from_manager(&guard);
drop(guard);
status
}
pub async fn disconnect() -> Result<Status, String> {
disconnect_with_hook(DisconnectHook::default()).await
}
async fn disconnect_with_hook(hook: DisconnectHook) -> Result<Status, String> {
let _credential_io = credential_io_gate().lock().await;
let observed = credential_presence_from_vault();
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
guard.replace_credential_presence(observed);
let reserved_generation = guard.advance_action_generation();
if let Some(flow) = guard.pending.take() {
flow.cancel.cancel();
guard.last_result = Some(terminal(
&flow.flow_id,
"cancelled",
"OpenRouter connect was cancelled",
));
}
let reserved_status = status_from_manager(&guard);
drop(guard);
publish_status(&reserved_status);
let delete_result = hook
.delete_oauth_key()
.map_err(|error| format!("remove OpenRouter OAuth credential: {error}"));
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
let is_current = guard.action_generation == reserved_generation;
if is_current && delete_result.is_ok() {
guard.set_oauth_exists(false);
}
let final_status = status_from_manager(&guard);
drop(guard);
if is_current {
publish_status(&final_status);
}
delete_result.map(|()| final_status)
}
pub async fn start(
authorization_base_url: Option<&str>,
exchange_url: Option<&str>,
timeout_seconds: Option<u64>,
) -> Result<StartResult, String> {
start_with_hook(
authorization_base_url,
exchange_url,
timeout_seconds,
ManagerTransactionHook::default(),
)
.await
}
const SUPERSEDED_START_ERROR: &str =
"superseded: OpenRouter connect was replaced by a newer credential action";
fn test_mode_enabled() -> bool {
std::env::var("CAR_OPENROUTER_TEST_MODE").as_deref() == Ok("1")
}
fn non_empty_environment_value(key: &str) -> Option<String> {
std::env::var(key)
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
fn resolve_start_endpoints(
authorization_base_url: Option<&str>,
exchange_url: Option<&str>,
) -> Result<(String, String), String> {
let explicit_override = authorization_base_url.is_some() || exchange_url.is_some();
let test_mode = test_mode_enabled();
if explicit_override && !cfg!(test) && !test_mode {
return Err("OpenRouter endpoint overrides require CAR_OPENROUTER_TEST_MODE=1".into());
}
let authorization_base_url = authorization_base_url
.map(str::to_string)
.or_else(|| {
test_mode.then(|| non_empty_environment_value(TEST_AUTHORIZATION_BASE_URL_ENV))?
})
.unwrap_or_else(|| DEFAULT_AUTH_URL.to_string());
let exchange_url = exchange_url
.map(str::to_string)
.or_else(|| test_mode.then(|| non_empty_environment_value(TEST_EXCHANGE_URL_ENV))?)
.unwrap_or_else(|| DEFAULT_EXCHANGE_URL.to_string());
Ok((authorization_base_url, exchange_url))
}
fn start_reservation_is_current(action_generation: u64, flow_id: &str) -> bool {
let guard = manager().lock().unwrap_or_else(|p| p.into_inner());
guard.action_generation == action_generation
&& guard.pending.as_ref().is_some_and(|pending| {
pending.action_generation == action_generation && pending.flow_id == flow_id
})
}
fn fail_current_start_preparation(action_generation: u64, flow_id: &str, message: &str) {
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
let is_current = guard.action_generation == action_generation
&& guard.pending.as_ref().is_some_and(|pending| {
pending.action_generation == action_generation && pending.flow_id == flow_id
});
if !is_current {
return;
}
if let Some(flow) = guard.pending.take() {
flow.cancel.cancel();
}
guard.last_result = Some(terminal(flow_id, "exchange_rejected", message));
let status = status_from_manager(&guard);
drop(guard);
publish_status(&status);
}
async fn start_with_hook(
authorization_base_url: Option<&str>,
exchange_url: Option<&str>,
timeout_seconds: Option<u64>,
transaction_hook: ManagerTransactionHook,
) -> Result<StartResult, String> {
let (authorization_base_url, exchange_url) =
resolve_start_endpoints(authorization_base_url, exchange_url)?;
let flow_id = uuid::Uuid::new_v4().to_string();
let verifier = car_auth::pkce_verifier();
let challenge = car_auth::pkce_challenge(&verifier);
let timeout = timeout_seconds
.unwrap_or(DEFAULT_TIMEOUT_SECS)
.clamp(1, MAX_TIMEOUT_SECS);
let deadline_unix_ms = now_unix_ms().saturating_add(timeout * 1000);
let cancel = CancellationToken::new();
let (action_generation, reserved_status) = {
transaction_hook.before_lock();
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
transaction_hook.inside_lock();
let action_generation = guard.advance_action_generation();
if let Some(flow) = guard.pending.take() {
flow.cancel.cancel();
}
guard.pending = Some(PendingFlow {
flow_id: flow_id.clone(),
action_generation,
deadline_unix_ms,
cancel: cancel.clone(),
});
guard.last_result = None;
(action_generation, status_from_manager(&guard))
};
publish_status(&reserved_status);
transaction_hook.after_reserve();
if !start_reservation_is_current(action_generation, &flow_id) {
return Err(SUPERSEDED_START_ERROR.into());
}
let listener = match TcpListener::bind("127.0.0.1:0").await {
Ok(listener) => listener,
Err(error) => {
let message = format!("bind OpenRouter OAuth callback: {error}");
fail_current_start_preparation(action_generation, &flow_id, &message);
return Err(message);
}
};
let port = match listener.local_addr() {
Ok(address) => address.port(),
Err(error) => {
let message = format!("read OpenRouter OAuth callback address: {error}");
fail_current_start_preparation(action_generation, &flow_id, &message);
return Err(message);
}
};
let callback_url = format!("http://127.0.0.1:{port}/openrouter/callback/{flow_id}");
let mut auth_url = match reqwest::Url::parse(&authorization_base_url) {
Ok(url) => url,
Err(error) => {
let message = format!("invalid OpenRouter authorization URL: {error}");
fail_current_start_preparation(action_generation, &flow_id, &message);
return Err(message);
}
};
auth_url
.query_pairs_mut()
.append_pair("callback_url", &callback_url)
.append_pair("code_challenge", &challenge)
.append_pair("code_challenge_method", "S256");
if !start_reservation_is_current(action_generation, &flow_id) {
cancel.cancel();
return Err(SUPERSEDED_START_ERROR.into());
}
transaction_hook.before_return();
let task_flow_id = flow_id.clone();
tokio::spawn(async move {
run_flow(
listener,
task_flow_id,
verifier,
exchange_url,
Duration::from_secs(timeout),
cancel,
)
.await;
});
Ok(StartResult {
authority_generation: action_generation,
authorize_url: auth_url.to_string(),
flow_id,
deadline_unix_ms,
})
}
async fn run_flow(
listener: TcpListener,
flow_id: String,
verifier: String,
exchange_url: String,
timeout: Duration,
cancel: CancellationToken,
) {
let deadline = tokio::time::Instant::now() + timeout;
let accepted = tokio::select! {
_ = cancel.cancelled() => return,
_ = tokio::time::sleep_until(deadline) => {
finish(&flow_id, terminal(&flow_id, "timed_out", "OpenRouter connect timed out"), None);
return;
}
accepted = listener.accept() => accepted,
};
let Ok((mut stream, _)) = accepted else {
finish(
&flow_id,
terminal(&flow_id, "exchange_rejected", "OpenRouter callback failed"),
None,
);
return;
};
let callback_read_deadline = std::cmp::min(
deadline,
tokio::time::Instant::now() + Duration::from_secs(5),
);
let request = match read_callback_headers(&mut stream, callback_read_deadline, &cancel).await {
Ok(bytes) => String::from_utf8_lossy(&bytes).to_string(),
Err(CallbackReadError::Cancelled) => return,
Err(CallbackReadError::TimedOut) => {
finish(
&flow_id,
terminal(&flow_id, "timed_out", "OpenRouter connect timed out"),
None,
);
return;
}
Err(CallbackReadError::TooLarge) => {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter callback was too large",
),
None,
);
return;
}
Err(CallbackReadError::Incomplete) => {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter callback was incomplete",
),
None,
);
return;
}
Err(CallbackReadError::Io) => {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter callback was invalid",
),
None,
);
return;
}
};
let target = request
.lines()
.next()
.and_then(|line| line.split_whitespace().nth(1))
.unwrap_or("/");
let parsed = reqwest::Url::parse(&format!("http://127.0.0.1{target}"));
let (code, callback_error) = match parsed {
Ok(url) if url.path() == format!("/openrouter/callback/{flow_id}") => {
let pairs: std::collections::HashMap<_, _> = url.query_pairs().into_owned().collect();
(pairs.get("code").cloned(), pairs.get("error").cloned())
}
_ => (None, Some("invalid_callback".to_string())),
};
let response = if code.is_some() {
"HTTP/1.1 200 OK\r\nContent-Type: text/html; charset=utf-8\r\nConnection: close\r\n\r\n<!doctype html><title>CAR connected</title><p>OpenRouter connected. You can close this window.</p>"
} else {
"HTTP/1.1 400 Bad Request\r\nContent-Type: text/html; charset=utf-8\r\nConnection: close\r\n\r\n<!doctype html><title>CAR not connected</title><p>OpenRouter was not connected. Return to CarHost and retry.</p>"
};
let _ = stream.write_all(response.as_bytes()).await;
let _ = stream.shutdown().await;
if let Some(error) = callback_error {
let kind = if error == "access_denied" {
"cancelled"
} else {
"exchange_rejected"
};
finish(
&flow_id,
terminal(&flow_id, kind, "OpenRouter did not approve the connection"),
None,
);
return;
}
let Some(code) = code else {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter callback contained no code",
),
None,
);
return;
};
let exchange = reqwest::Client::new()
.post(&exchange_url)
.json(&serde_json::json!({
"code": code,
"code_verifier": verifier,
"code_challenge_method": "S256",
}))
.send();
let exchanged = tokio::select! {
_ = cancel.cancelled() => return,
_ = tokio::time::sleep_until(deadline) => {
finish(&flow_id, terminal(&flow_id, "timed_out", "OpenRouter connect timed out"), None);
return;
}
response = exchange => response,
};
let response = match exchanged {
Ok(response) if response.status().is_success() => response,
_ => {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter rejected the code exchange",
),
None,
);
return;
}
};
let payload: serde_json::Value = match tokio::select! {
_ = cancel.cancelled() => return,
_ = tokio::time::sleep_until(deadline) => {
finish(&flow_id, terminal(&flow_id, "timed_out", "OpenRouter connect timed out"), None);
return;
}
payload = response.json() => payload,
} {
Ok(payload) => payload,
Err(_) => {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter returned an invalid exchange response",
),
None,
);
return;
}
};
let Some(key) = payload
.get("key")
.and_then(|value| value.as_str())
.filter(|key| !key.trim().is_empty())
else {
finish(
&flow_id,
terminal(
&flow_id,
"exchange_rejected",
"OpenRouter exchange returned no key",
),
None,
);
return;
};
finish_with_key(
&flow_id,
terminal(&flow_id, "connected", "OpenRouter account connected"),
key,
)
.await;
}
const MAX_CALLBACK_HEADER_BYTES: usize = 8_192;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum CallbackReadError {
Cancelled,
TimedOut,
TooLarge,
Incomplete,
Io,
}
async fn read_callback_headers(
stream: &mut tokio::net::TcpStream,
deadline: tokio::time::Instant,
cancel: &CancellationToken,
) -> Result<Vec<u8>, CallbackReadError> {
let mut request = Vec::with_capacity(1024);
let mut chunk = [0u8; 1024];
loop {
let count = tokio::select! {
_ = cancel.cancelled() => return Err(CallbackReadError::Cancelled),
_ = tokio::time::sleep_until(deadline) => return Err(CallbackReadError::TimedOut),
read = stream.read(&mut chunk) => read.map_err(|_| CallbackReadError::Io)?,
};
if count == 0 {
return Err(CallbackReadError::Incomplete);
}
request.extend_from_slice(&chunk[..count]);
if request.len() > MAX_CALLBACK_HEADER_BYTES {
return Err(CallbackReadError::TooLarge);
}
if request.windows(4).any(|window| window == b"\r\n\r\n") {
return Ok(request);
}
}
}
fn finish(flow_id: &str, result: TerminalResult, key: Option<&str>) {
debug_assert!(key.is_none(), "credential finishes use finish_with_key");
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
if !guard.pending.as_ref().is_some_and(|flow| {
flow.flow_id == flow_id && flow.action_generation == guard.action_generation
}) {
return;
}
guard.pending = None;
guard.last_result = Some(result);
let status = status_from_manager(&guard);
drop(guard);
publish_status(&status);
}
async fn finish_with_key(flow_id: &str, result: TerminalResult, key: &str) {
finish_with_key_with_hook(flow_id, result, key, CredentialFinishHook::default()).await;
}
async fn finish_with_key_with_hook(
flow_id: &str,
result: TerminalResult,
key: &str,
hook: CredentialFinishHook,
) {
hook.before_gate();
let _credential_io = credential_io_gate().lock().await;
let observed = credential_presence_from_vault();
hook.after_gate();
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
if !guard.pending.as_ref().is_some_and(|flow| {
flow.flow_id == flow_id && flow.action_generation == guard.action_generation
}) {
return;
}
guard.replace_credential_presence(observed);
if store_oauth_key(key).is_err() {
guard.pending = None;
guard.last_result = Some(terminal(
flow_id,
"exchange_rejected",
"OpenRouter connected but the OS keychain rejected the credential",
));
let status = status_from_manager(&guard);
drop(guard);
publish_status(&status);
return;
}
guard.set_oauth_exists(true);
guard.pending = None;
guard.last_result = Some(result);
let status = status_from_manager(&guard);
drop(guard);
publish_status(&status);
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::net::TcpListener;
fn oauth_test_lock() -> &'static tokio::sync::Mutex<()> {
static LOCK: OnceLock<tokio::sync::Mutex<()>> = OnceLock::new();
LOCK.get_or_init(|| tokio::sync::Mutex::new(()))
}
fn reset_test_manager() {
delete_oauth_key().unwrap();
let mut guard = manager().lock().unwrap_or_else(|p| p.into_inner());
if let Some(flow) = guard.pending.take() {
flow.cancel.cancel();
}
guard.action_generation = 0;
guard.last_result = None;
guard.credentials.oauth = false;
guard.credential_revision = 0;
}
async fn exchange_server(
status: u16,
body: &'static str,
) -> (String, tokio::task::JoinHandle<String>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/exchange", listener.local_addr().unwrap());
let task = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut request = Vec::new();
let mut chunk = [0u8; 4096];
loop {
let count = stream.read(&mut chunk).await.unwrap();
if count == 0 {
break;
}
request.extend_from_slice(&chunk[..count]);
let text = String::from_utf8_lossy(&request);
let Some(header_end) = text.find("\r\n\r\n") else {
continue;
};
let content_length = text[..header_end]
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
if request.len() >= header_end + 4 + content_length {
break;
}
}
let reason = if status == 200 { "OK" } else { "Unauthorized" };
let response = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
stream.write_all(response.as_bytes()).await.unwrap();
String::from_utf8(request).unwrap()
});
(url, task)
}
async fn wait_for_terminal(flow_id: &str) -> Status {
for _ in 0..250 {
let current = status();
if !current
.pending
.as_ref()
.is_some_and(|pending| pending.flow_id == flow_id)
{
return current;
}
tokio::time::sleep(Duration::from_millis(25)).await;
}
panic!("OpenRouter flow {flow_id} did not finish");
}
fn spawn_paused_preparing_start() -> (
tokio::task::JoinHandle<Result<StartResult, String>>,
String,
Arc<std::sync::Barrier>,
) {
let reserved = Arc::new(std::sync::Barrier::new(2));
let release = Arc::new(std::sync::Barrier::new(2));
let hook = ManagerTransactionHook {
after_reserve_entered: Some(reserved.clone()),
after_reserve_release: Some(release.clone()),
..ManagerTransactionHook::default()
};
let task = tokio::spawn(async move {
start_with_hook(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
hook,
)
.await
});
tokio::task::block_in_place(|| reserved.wait());
let flow_id = status()
.pending
.expect("preparing start owns pending authority")
.flow_id;
(task, flow_id, release)
}
#[tokio::test]
async fn empty_param_start_uses_local_endpoints_only_in_explicit_test_mode() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let (exchange_url, exchange_task) =
exchange_server(200, r#"{"key":"local-oauth-test-key"}"#).await;
unsafe {
std::env::set_var("CAR_OPENROUTER_TEST_MODE", "1");
std::env::set_var(
"OPENROUTER_AUTHORIZATION_BASE_URL",
"http://127.0.0.1:9/local-openrouter-auth",
);
std::env::set_var("OPENROUTER_EXCHANGE_URL", &exchange_url);
}
let started = start(None, None, Some(5)).await.unwrap();
assert!(
started
.authorize_url
.starts_with("http://127.0.0.1:9/local-openrouter-auth?"),
"the exact native empty-params action must use the daemon test-mode override: {}",
started.authorize_url
);
let authorize_url = reqwest::Url::parse(&started.authorize_url).unwrap();
let callback_url = authorize_url
.query_pairs()
.find_map(|(name, value)| (name == "callback_url").then(|| value.into_owned()))
.unwrap();
reqwest::get(format!("{callback_url}?code=local-approved-code"))
.await
.unwrap();
let terminal = wait_for_terminal(&started.flow_id).await;
assert_eq!(terminal.state, "connected");
let exchange_request = exchange_task.await.unwrap();
assert!(exchange_request.starts_with("POST /exchange "));
assert!(exchange_request.contains("local-approved-code"));
assert!(exchange_request.contains("code_verifier"));
disconnect().await.unwrap();
unsafe {
std::env::remove_var("CAR_OPENROUTER_TEST_MODE");
}
let production = start(None, None, Some(5)).await.unwrap();
assert!(
production
.authorize_url
.starts_with("https://openrouter.ai/auth?"),
"unsafe endpoint env vars must be ignored outside explicit test mode: {}",
production.authorize_url
);
assert_eq!(
resolve_start_endpoints(None, None).unwrap().1,
DEFAULT_EXCHANGE_URL
);
cancel(Some(&production.flow_id));
unsafe {
std::env::remove_var("OPENROUTER_EXCHANGE_URL");
std::env::remove_var("OPENROUTER_AUTHORIZATION_BASE_URL");
}
}
#[tokio::test]
async fn pkce_flow_stores_real_key_and_terminal_states_are_bounded() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let marker = "sk-or-v1-test-marker-never-serialize";
let exchange_body: &'static str = Box::leak(
serde_json::json!({ "key": marker })
.to_string()
.into_boxed_str(),
);
let (exchange_url, exchange_task) = exchange_server(200, exchange_body).await;
let started = start(
Some("http://127.0.0.1:9/auth"),
Some(&exchange_url),
Some(5),
)
.await
.unwrap();
let url = reqwest::Url::parse(&started.authorize_url).unwrap();
let query: std::collections::HashMap<_, _> = url.query_pairs().into_owned().collect();
assert_eq!(
query.get("code_challenge_method").map(String::as_str),
Some("S256")
);
assert!(query.get("code_challenge").is_some_and(|v| v.len() >= 43));
assert!(query
.get("callback_url")
.is_some_and(|v| v.contains("127.0.0.1") && v.contains(&started.flow_id)));
assert_eq!(status().state, "pending");
let callback = reqwest::Url::parse(query.get("callback_url").unwrap()).unwrap();
let mut callback_stream = tokio::net::TcpStream::connect((
callback.host_str().unwrap(),
callback.port().unwrap(),
))
.await
.unwrap();
for fragment in [
"GET ".to_string(),
format!("{}?code=approved-code HTTP/1.1\r\n", callback.path()),
"Host: 127.0.0.1\r\n".to_string(),
"Connection: close\r\n\r\n".to_string(),
] {
callback_stream
.write_all(fragment.as_bytes())
.await
.unwrap();
tokio::task::yield_now().await;
}
let mut callback_response = Vec::new();
callback_stream
.read_to_end(&mut callback_response)
.await
.unwrap();
assert!(
String::from_utf8_lossy(&callback_response).starts_with("HTTP/1.1 200"),
"{}",
String::from_utf8_lossy(&callback_response)
);
let terminal = wait_for_terminal(&started.flow_id).await;
assert_eq!(
terminal.last_result.as_ref().unwrap().kind,
"connected",
"{}",
terminal.last_result.as_ref().unwrap().message
);
assert!(terminal.oauth_key_exists);
assert_eq!(
test_oauth_slot()
.lock()
.unwrap_or_else(|p| p.into_inner())
.as_deref(),
Some(marker)
);
assert!(!serde_json::to_string(&terminal).unwrap().contains(marker));
let exchange_request = exchange_task.await.unwrap();
assert!(exchange_request.contains("approved-code"));
assert!(exchange_request.contains("code_verifier"));
assert!(exchange_request.contains("S256"));
let disconnected = disconnect().await.unwrap();
assert!(!disconnected.oauth_key_exists);
assert!(test_oauth_slot()
.lock()
.unwrap_or_else(|p| p.into_inner())
.is_none());
let (reject_url, _reject_task) = exchange_server(401, "{}").await;
let rejected = start(Some("http://127.0.0.1:9/auth"), Some(&reject_url), Some(5))
.await
.unwrap();
let rejected_url = reqwest::Url::parse(&rejected.authorize_url).unwrap();
let rejected_callback = rejected_url
.query_pairs()
.find_map(|(name, value)| (name == "callback_url").then(|| value.into_owned()))
.unwrap();
reqwest::get(format!("{rejected_callback}?code=rejected-code"))
.await
.unwrap();
assert_eq!(
wait_for_terminal(&rejected.flow_id)
.await
.last_result
.unwrap()
.kind,
"exchange_rejected"
);
let timed = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(1),
)
.await
.unwrap();
assert_eq!(
wait_for_terminal(&timed.flow_id)
.await
.last_result
.unwrap()
.kind,
"timed_out"
);
let cancelled = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(5),
)
.await
.unwrap();
assert_eq!(
cancel(Some(&cancelled.flow_id)).last_result.unwrap().kind,
"cancelled"
);
let first = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(5),
)
.await
.unwrap();
let replacement = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(5),
)
.await
.unwrap();
let replacement_status = status();
assert_eq!(replacement_status.state, "pending");
assert_eq!(
replacement_status
.pending
.as_ref()
.map(|p| p.flow_id.as_str()),
Some(replacement.flow_id.as_str())
);
assert!(
terminal_status_for_flow(&first.flow_id).is_none(),
"superseded watcher A must not emit replacement flow B's pending state"
);
let replacement_terminal = cancel(Some(&replacement.flow_id));
assert_eq!(
replacement_terminal
.last_result
.as_ref()
.map(|r| r.flow_id.as_str()),
Some(replacement.flow_id.as_str())
);
assert!(terminal_status_for_flow(&replacement.flow_id).is_some());
let superseded = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(5),
)
.await
.unwrap();
supersede_for_pasted_change();
let superseded_status = wait_for_terminal(&superseded.flow_id).await;
assert_eq!(superseded_status.last_result.unwrap().kind, "superseded");
delete_oauth_key().unwrap();
}
#[tokio::test]
async fn callback_rejects_oversize_and_times_out_incomplete_headers() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let oversized = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(5),
)
.await
.unwrap();
let oversized_url = reqwest::Url::parse(&oversized.authorize_url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == "callback_url").then(|| value.into_owned()))
.unwrap();
let oversized_url = reqwest::Url::parse(&oversized_url).unwrap();
let mut stream = tokio::net::TcpStream::connect((
oversized_url.host_str().unwrap(),
oversized_url.port().unwrap(),
))
.await
.unwrap();
stream.write_all(&vec![b'x'; 8_193]).await.unwrap();
let oversized_terminal = wait_for_terminal(&oversized.flow_id).await;
let oversized_result = oversized_terminal.last_result.unwrap();
assert_eq!(oversized_result.kind, "exchange_rejected");
assert_eq!(
oversized_result.message,
"OpenRouter callback was too large"
);
let incomplete = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(10),
)
.await
.unwrap();
let incomplete_url = reqwest::Url::parse(&incomplete.authorize_url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == "callback_url").then(|| value.into_owned()))
.unwrap();
let incomplete_url = reqwest::Url::parse(&incomplete_url).unwrap();
let mut stream = tokio::net::TcpStream::connect((
incomplete_url.host_str().unwrap(),
incomplete_url.port().unwrap(),
))
.await
.unwrap();
stream
.write_all(b"GET /incomplete HTTP/1.1\r\n")
.await
.unwrap();
let incomplete_terminal = wait_for_terminal(&incomplete.flow_id).await;
assert_eq!(incomplete_terminal.last_result.unwrap().kind, "timed_out");
let half_closed = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(5),
)
.await
.unwrap();
let half_closed_url = reqwest::Url::parse(&half_closed.authorize_url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == "callback_url").then(|| value.into_owned()))
.unwrap();
let half_closed_url = reqwest::Url::parse(&half_closed_url).unwrap();
let mut stream = tokio::net::TcpStream::connect((
half_closed_url.host_str().unwrap(),
half_closed_url.port().unwrap(),
))
.await
.unwrap();
stream
.write_all(b"GET /half-close HTTP/1.1\r\nHost: localhost\r\n")
.await
.unwrap();
stream.shutdown().await.unwrap();
let half_closed_terminal = wait_for_terminal(&half_closed.flow_id).await;
let half_closed_result = half_closed_terminal.last_result.unwrap();
assert_eq!(half_closed_result.kind, "exchange_rejected");
assert_eq!(
half_closed_result.message,
"OpenRouter callback was incomplete"
);
}
#[tokio::test]
async fn callback_header_reader_rejects_half_closed_incomplete_request() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let client = tokio::spawn(async move {
let mut stream = tokio::net::TcpStream::connect(address).await.unwrap();
stream
.write_all(b"GET /callback?code=partial HTTP/1.1\r\nHost: localhost\r\n")
.await
.unwrap();
stream.shutdown().await.unwrap();
});
let (mut server, _) = listener.accept().await.unwrap();
let result = read_callback_headers(
&mut server,
tokio::time::Instant::now() + Duration::from_secs(1),
&CancellationToken::new(),
)
.await;
client.await.unwrap();
assert!(
result.is_err(),
"EOF before CRLFCRLF must be an invalid callback, not a parseable request"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn newer_start_supersedes_an_older_preparing_start() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let (first_task, first_flow_id, release_first) = spawn_paused_preparing_start();
let second = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
)
.await
.unwrap();
tokio::task::block_in_place(|| release_first.wait());
let first_error = first_task.await.unwrap().unwrap_err();
assert!(first_error.contains("superseded"));
let current = status();
assert_eq!(
current.pending.as_ref().map(|p| p.flow_id.as_str()),
Some(second.flow_id.as_str())
);
assert!(!is_pending(&first_flow_id));
assert!(is_pending(&second.flow_id));
assert!(terminal_status_for_flow(&first_flow_id).is_none());
cancel(Some(&second.flow_id));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn delayed_start_result_is_lower_authority_than_newer_broadcast() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let before_return = Arc::new(std::sync::Barrier::new(2));
let release_return = Arc::new(std::sync::Barrier::new(2));
let hook = ManagerTransactionHook {
before_return_entered: Some(before_return.clone()),
before_return_release: Some(release_return.clone()),
..ManagerTransactionHook::default()
};
let mut updates = subscribe_authority_updates();
let first_task = tokio::spawn(async move {
start_with_hook(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
hook,
)
.await
.unwrap()
});
tokio::task::block_in_place(|| before_return.wait());
let second = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
)
.await
.unwrap();
updates.changed().await.unwrap();
let broadcast = updates
.borrow_and_update()
.clone()
.expect("newer authority snapshot");
assert_eq!(broadcast.authority_generation, second.authority_generation);
assert_eq!(
broadcast
.pending
.as_ref()
.map(|pending| pending.flow_id.as_str()),
Some(second.flow_id.as_str())
);
tokio::task::block_in_place(|| release_return.wait());
let delayed_first = first_task.await.unwrap();
assert!(
delayed_first.authority_generation < broadcast.authority_generation,
"host must reject delayed flow A after receiving flow B broadcast"
);
assert_ne!(delayed_first.flow_id, second.flow_id);
cancel(Some(&second.flow_id));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn disconnect_publishes_before_delete_and_preserves_truth_on_failure() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
*test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some("existing-oauth".into());
manager()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.set_oauth_exists(true);
let before_return = Arc::new(std::sync::Barrier::new(2));
let release_return = Arc::new(std::sync::Barrier::new(2));
let start_hook = ManagerTransactionHook {
before_return_entered: Some(before_return.clone()),
before_return_release: Some(release_return.clone()),
..ManagerTransactionHook::default()
};
let mut updates = subscribe_authority_updates();
let start_task = tokio::spawn(async move {
start_with_hook(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
start_hook,
)
.await
.unwrap()
});
tokio::task::block_in_place(|| before_return.wait());
updates.changed().await.unwrap();
let start_snapshot = updates.borrow_and_update().clone().unwrap();
let before_delete = Arc::new(std::sync::Barrier::new(2));
let release_delete = Arc::new(std::sync::Barrier::new(2));
let disconnect_hook = DisconnectHook {
before_delete_entered: Some(before_delete.clone()),
before_delete_release: Some(release_delete.clone()),
..DisconnectHook::default()
};
let disconnect_task =
tokio::spawn(async move { disconnect_with_hook(disconnect_hook).await });
tokio::task::block_in_place(|| before_delete.wait());
updates.changed().await.unwrap();
let reserved_disconnect = updates.borrow_and_update().clone().unwrap();
assert!(reserved_disconnect.authority_generation > start_snapshot.authority_generation);
assert!(reserved_disconnect.pending.is_none());
assert!(reserved_disconnect.oauth_key_exists);
tokio::task::block_in_place(|| release_return.wait());
let delayed_start = start_task.await.unwrap();
assert!(
delayed_start.authority_generation < reserved_disconnect.authority_generation,
"a host must reject delayed flow A before disconnect enters the vault"
);
tokio::task::block_in_place(|| release_delete.wait());
let disconnected = disconnect_task.await.unwrap().unwrap();
assert_eq!(
disconnected.authority_generation,
reserved_disconnect.authority_generation
);
assert!(!disconnected.oauth_key_exists);
*test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some("existing-oauth".into());
manager()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.set_oauth_exists(true);
let generation_before_failure = status().authority_generation;
let error = disconnect_with_hook(DisconnectHook {
forced_delete_error: Some("permitted vault failure".into()),
..DisconnectHook::default()
})
.await
.unwrap_err();
assert!(error.contains("permitted vault failure"));
let failed_status = status();
assert_eq!(
failed_status.authority_generation,
generation_before_failure + 1
);
assert!(failed_status.oauth_key_exists);
assert!(failed_status.pending.is_none());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn disconnect_delete_finishes_before_newer_oauth_store() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
*test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some("old-oauth".into());
manager()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.set_oauth_exists(true);
let before_delete = Arc::new(std::sync::Barrier::new(2));
let release_delete = Arc::new(std::sync::Barrier::new(2));
let disconnect_hook = DisconnectHook {
before_delete_entered: Some(before_delete.clone()),
before_delete_release: Some(release_delete.clone()),
..DisconnectHook::default()
};
let disconnect_task =
tokio::spawn(async move { disconnect_with_hook(disconnect_hook).await });
tokio::task::block_in_place(|| before_delete.wait());
let newer = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
)
.await
.unwrap();
let finish_waiting = Arc::new(std::sync::Barrier::new(2));
let finish_entered = Arc::new(std::sync::Barrier::new(2));
let release_finish = Arc::new(std::sync::Barrier::new(2));
let finish_hook = CredentialFinishHook {
before_gate_entered: Some(finish_waiting.clone()),
after_gate_entered: Some(finish_entered.clone()),
after_gate_release: Some(release_finish.clone()),
};
let newer_flow_id = newer.flow_id.clone();
let finish_task = tokio::spawn(async move {
finish_with_key_with_hook(
&newer_flow_id,
terminal(&newer_flow_id, "connected", "newer OAuth connected"),
"new-oauth",
finish_hook,
)
.await;
});
tokio::task::block_in_place(|| finish_waiting.wait());
tokio::task::block_in_place(|| release_delete.wait());
tokio::task::block_in_place(|| finish_entered.wait());
let stale_disconnect = disconnect_task.await.unwrap().unwrap();
assert_eq!(
stale_disconnect.authority_generation,
newer.authority_generation
);
assert_eq!(
test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_deref(),
None,
"disconnect deletion must complete before newer OAuth store enters"
);
tokio::task::block_in_place(|| release_finish.wait());
finish_task.await.unwrap();
assert_eq!(
test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_deref(),
Some("new-oauth")
);
let final_status = status();
assert_eq!(
final_status.authority_generation,
newer.authority_generation
);
assert!(final_status.oauth_key_exists);
assert_eq!(
final_status
.last_result
.as_ref()
.map(|result| (result.flow_id.as_str(), result.kind.as_str())),
Some((newer.flow_id.as_str(), "connected"))
);
}
#[tokio::test]
async fn status_read_reconciles_external_oauth_slot_changes() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
*test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = Some("external-oauth".into());
let connected = status();
assert!(connected.oauth_key_exists);
*test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) = None;
let removed = status();
assert!(!removed.oauth_key_exists);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn cancel_supersedes_an_older_preparing_start() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let (start_task, flow_id, release_start) = spawn_paused_preparing_start();
let cancelled = cancel(Some(&flow_id));
tokio::task::block_in_place(|| release_start.wait());
let start_error = start_task.await.unwrap().unwrap_err();
assert!(start_error.contains("superseded"));
assert!(cancelled.pending.is_none());
assert_eq!(
cancelled
.last_result
.as_ref()
.map(|result| result.kind.as_str()),
Some("cancelled")
);
assert!(!is_pending(&flow_id));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn disconnect_supersedes_an_older_preparing_start() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let (start_task, flow_id, release_start) = spawn_paused_preparing_start();
let disconnected = disconnect().await.unwrap();
tokio::task::block_in_place(|| release_start.wait());
let start_error = start_task.await.unwrap().unwrap_err();
assert!(start_error.contains("superseded"));
assert!(disconnected.pending.is_none());
assert!(!disconnected.oauth_key_exists);
assert!(!is_pending(&flow_id));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn pasted_put_and_delete_reserve_before_vault_and_supersede_preparing_start() {
let _test_guard = oauth_test_lock().lock().await;
for operation in ["put", "delete"] {
reset_test_manager();
let (start_task, flow_id, release_start) = spawn_paused_preparing_start();
let generation_before_mutation = status().authority_generation;
let mut updates = subscribe_authority_updates();
let vault_entered = Arc::new(std::sync::Barrier::new(2));
let release_vault = Arc::new(std::sync::Barrier::new(2));
let vault_entered_task = vault_entered.clone();
let release_vault_task = release_vault.clone();
let mutation_task = tokio::spawn(async move {
mutate_pasted_credential(operation == "put", move || {
vault_entered_task.wait();
release_vault_task.wait();
Ok::<_, String>(())
})
.await
});
tokio::task::block_in_place(|| vault_entered.wait());
updates.changed().await.unwrap();
let reserved = updates.borrow_and_update().clone().unwrap();
assert!(
reserved.authority_generation > generation_before_mutation,
"{operation} must broadcast before entering the vault"
);
let finish_flow_id = flow_id.clone();
let finish_task = tokio::spawn(async move {
finish_with_key(
&finish_flow_id,
terminal(&finish_flow_id, "connected", "stale OAuth finish"),
"test-oauth-key",
)
.await;
});
tokio::task::block_in_place(|| release_start.wait());
let start_error = start_task.await.unwrap().unwrap_err();
assert!(start_error.contains("superseded"), "{operation}");
assert!(!is_pending(&flow_id), "{operation}");
assert!(
test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.is_none(),
"{operation} must reject A's OAuth key"
);
assert_eq!(
status()
.last_result
.as_ref()
.map(|result| result.kind.as_str()),
Some("superseded"),
"{operation}"
);
tokio::task::block_in_place(|| release_vault.wait());
let (mutation_result, _) = mutation_task.await.unwrap();
mutation_result.unwrap();
finish_task.await.unwrap();
assert!(
test_oauth_slot()
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.is_none(),
"{operation} must reject A's OAuth key after waiting for B's vault IO"
);
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn pasted_same_slot_generation_order_matches_physical_io_order() {
let _test_guard = oauth_test_lock().lock().await;
for (first_exists, second_exists) in [(false, true), (true, false)] {
reset_test_manager();
let physical_slot = Arc::new(Mutex::new(if first_exists {
None
} else {
Some("initial".to_string())
}));
let first_entered = Arc::new(std::sync::Barrier::new(2));
let release_first = Arc::new(std::sync::Barrier::new(2));
let first_slot = physical_slot.clone();
let first_entered_task = first_entered.clone();
let release_first_task = release_first.clone();
let first_task = tokio::spawn(async move {
mutate_pasted_credential(first_exists, move || {
first_entered_task.wait();
release_first_task.wait();
*first_slot
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) =
first_exists.then(|| "first".to_string());
Ok::<_, String>(())
})
.await
});
tokio::task::block_in_place(|| first_entered.wait());
let first_generation = status().authority_generation;
let second_waiting = Arc::new(std::sync::Barrier::new(2));
let second_slot = physical_slot.clone();
let second_waiting_task = second_waiting.clone();
let second_task = tokio::spawn(async move {
mutate_pasted_credential_with_hook(
second_exists,
move || {
*second_slot
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner()) =
second_exists.then(|| "second".to_string());
Ok::<_, String>(())
},
PastedMutationHook {
before_gate_entered: Some(second_waiting_task),
},
)
.await
});
tokio::task::block_in_place(|| second_waiting.wait());
assert_eq!(status().authority_generation, first_generation);
tokio::task::block_in_place(|| release_first.wait());
let (_, first_status) = first_task.await.unwrap();
let (_, second_status) = second_task.await.unwrap();
assert_eq!(first_status.authority_generation, first_generation);
assert!(second_status.authority_generation > first_generation);
assert_eq!(second_status.pasted_key_exists, second_exists);
assert_eq!(
physical_slot
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.as_deref(),
second_exists.then_some("second")
);
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn watcher_snapshot_never_turns_flow_a_terminal_into_flow_b_status() {
let _test_guard = oauth_test_lock().lock().await;
reset_test_manager();
let first = start(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
)
.await
.unwrap();
let cancelled = cancel(Some(&first.flow_id));
assert_eq!(
cancelled.last_result.as_ref().map(|r| r.flow_id.as_str()),
Some(first.flow_id.as_str())
);
let watcher_inside = Arc::new(std::sync::Barrier::new(2));
let release_watcher = Arc::new(std::sync::Barrier::new(2));
let watcher_hook = ManagerTransactionHook {
inside_lock_entered: Some(watcher_inside.clone()),
inside_lock_release: Some(release_watcher.clone()),
..ManagerTransactionHook::default()
};
let first_flow_id = first.flow_id.clone();
let watcher = tokio::task::spawn_blocking(move || {
terminal_status_for_flow_with_hook(&first_flow_id, watcher_hook)
});
tokio::task::block_in_place(|| watcher_inside.wait());
let second_ready = Arc::new(std::sync::Barrier::new(2));
let second_hook = ManagerTransactionHook {
before_lock: Some(second_ready.clone()),
..ManagerTransactionHook::default()
};
let second_task = tokio::spawn(async move {
start_with_hook(
Some("http://127.0.0.1:9/auth"),
Some("http://127.0.0.1:9/exchange"),
Some(30),
second_hook,
)
.await
.unwrap()
});
tokio::task::block_in_place(|| second_ready.wait());
tokio::task::block_in_place(|| release_watcher.wait());
let captured = watcher.await.unwrap().expect("flow A terminal snapshot");
let second = second_task.await.unwrap();
assert!(captured.pending.is_none());
assert_eq!(
captured.last_result.as_ref().map(|r| r.flow_id.as_str()),
Some(first.flow_id.as_str()),
"watcher A must carry captured A terminal state, never global B state"
);
assert_eq!(
status().pending.as_ref().map(|p| p.flow_id.as_str()),
Some(second.flow_id.as_str())
);
assert!(terminal_status_for_flow(&first.flow_id).is_none());
cancel(Some(&second.flow_id));
}
}