use std::io::{Read, Write};
use std::net::TcpListener;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use crate::client::{RefreshOutcome, refresh_with_policy};
use crate::config::Configs;
use crate::oauth;
struct MockEndpoint {
base_url: String,
requests: Arc<std::sync::Mutex<Vec<String>>>,
}
impl MockEndpoint {
fn spawn(responses: Vec<(u16, String)>) -> Self {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let port = listener.local_addr().unwrap().port();
let requests = Arc::new(std::sync::Mutex::new(Vec::new()));
let requests_for_thread = Arc::clone(&requests);
std::thread::spawn(move || {
for stream in listener.incoming() {
let Ok(mut stream) = stream else { break };
let mut buf = Vec::new();
let mut tmp = [0u8; 1024];
let mut content_length = 0usize;
loop {
let Ok(read) = stream.read(&mut tmp) else {
break;
};
if read == 0 {
break;
}
buf.extend_from_slice(&tmp[..read]);
if let Some(pos) = find_headers_end(&buf) {
let headers = String::from_utf8_lossy(&buf[..pos]).to_lowercase();
for line in headers.lines() {
if let Some(v) = line.strip_prefix("content-length:") {
content_length = v.trim().parse().unwrap_or(0);
}
}
if buf.len() >= pos + 4 + content_length {
break;
}
}
}
let mut seen = requests_for_thread.lock().unwrap();
let idx = seen.len().min(responses.len().saturating_sub(1));
seen.push(String::from_utf8_lossy(&buf).to_string());
drop(seen);
let (status, body) = &responses[idx];
let reason = match status {
200 => "OK",
400 => "Bad Request",
500 => "Internal Server Error",
503 => "Service Unavailable",
_ => "Unknown",
};
let resp = format!(
"HTTP/1.1 {status} {reason}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
let _ = stream.write_all(resp.as_bytes());
let _ = stream.flush();
}
});
Self {
base_url: format!("http://127.0.0.1:{port}/oauth"),
requests,
}
}
fn hits(&self) -> usize {
self.requests.lock().unwrap().len()
}
fn auth_headers(&self) -> Vec<String> {
self.requests
.lock()
.unwrap()
.iter()
.map(|req| {
req.lines()
.find(|l| l.to_ascii_lowercase().starts_with("authorization:"))
.map(|l| l["authorization:".len()..].trim().to_string())
.unwrap_or_else(|| "(none)".to_string())
})
.collect()
}
}
fn find_headers_end(buf: &[u8]) -> Option<usize> {
buf.windows(4).position(|w| w == b"\r\n\r\n")
}
fn dead_grant() -> (u16, String) {
(
400,
r#"{"error":"invalid_grant","error_description":"grant request is invalid"}"#.to_string(),
)
}
fn server_error() -> (u16, String) {
(500, r#"{"error":"server_error"}"#.to_string())
}
fn fresh_tokens() -> (u16, String) {
(
200,
r#"{"access_token":"new-access","refresh_token":"new-refresh","expires_in":3600}"#
.to_string(),
)
}
fn ok_empty() -> (u16, String) {
(200, r#"{"data":{}}"#.to_string())
}
struct Fixture {
path: PathBuf,
_dir: tempfile::TempDir,
}
impl Fixture {
fn new() -> Self {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("config.json");
let mut configs = Configs::for_test(path.clone());
configs
.save_oauth_tokens("stale-access", Some("the-refresh-token"), 1)
.unwrap();
configs.root_config.user.token_expires_at = Some(0); configs.write().unwrap();
Self { path, _dir: dir }
}
fn load(&self) -> Configs {
let mut configs = Configs::for_test(self.path.clone());
configs.reload().unwrap();
configs
}
}
async fn legacy_policy(configs: &mut Configs, base_url: &str) -> Result<(), String> {
let refresh_token = match configs.get_refresh_token() {
Some(t) => t.to_owned(),
None => return Err("No refresh token available".to_string()),
};
let client = reqwest::Client::new();
match oauth::attempt_refresh(&client, base_url, &refresh_token).await {
Ok(resp) => {
configs
.save_oauth_tokens(
&resp.access_token,
resp.refresh_token.as_deref(),
resp.expires_in,
)
.unwrap();
Ok(())
}
Err(f) => Err(match f {
oauth::RefreshFailure::Terminal(e) | oauth::RefreshFailure::Transient(e) => {
e.to_string()
}
}),
}
}
const INVOCATIONS: usize = 20;
#[tokio::test]
async fn legacy_dead_grant_retries_forever_and_keeps_dead_credentials() {
let server = MockEndpoint::spawn(vec![dead_grant()]);
let fixture = Fixture::new();
for _ in 0..INVOCATIONS {
let mut configs = fixture.load();
let result = legacy_policy(&mut configs, &server.base_url).await;
assert!(result.is_err(), "dead grant should fail");
}
assert_eq!(
server.hits(),
INVOCATIONS,
"legacy policy re-presents a known-dead refresh token on every invocation"
);
let configs = fixture.load();
assert!(
configs.has_oauth_token(),
"legacy policy leaves the dead access token on disk"
);
assert_eq!(
configs.get_refresh_token(),
Some("the-refresh-token"),
"legacy policy leaves the dead refresh token on disk"
);
}
#[tokio::test]
async fn fixed_dead_grant_clears_credentials_and_stops_after_one_attempt() {
let server = MockEndpoint::spawn(vec![dead_grant()]);
let fixture = Fixture::new();
let mut outcomes = Vec::new();
for _ in 0..INVOCATIONS {
let mut configs = fixture.load();
outcomes.push(refresh_with_policy(&mut configs, &server.base_url, Duration::ZERO).await);
}
assert_eq!(
server.hits(),
1,
"fixed policy must present a dead refresh token exactly once, got {} hits",
server.hits()
);
assert!(matches!(outcomes[0], RefreshOutcome::SessionExpired(_)));
assert!(
outcomes[1..]
.iter()
.all(|o| matches!(o, RefreshOutcome::NoRefreshToken)),
"after clearing, later invocations have no token to retry"
);
let configs = fixture.load();
assert!(!configs.has_oauth_token());
assert_eq!(configs.get_refresh_token(), None);
}
#[tokio::test]
async fn legacy_transient_5xx_is_indistinguishable_from_a_dead_grant() {
let dead = MockEndpoint::spawn(vec![dead_grant()]);
let flaky = MockEndpoint::spawn(vec![server_error()]);
let f1 = Fixture::new();
let f2 = Fixture::new();
let dead_err = legacy_policy(&mut f1.load(), &dead.base_url)
.await
.unwrap_err();
let flaky_err = legacy_policy(&mut f2.load(), &flaky.base_url)
.await
.unwrap_err();
let dead_rendered = crate::errors::RailwayError::OAuthRefreshFailed(dead_err).to_string();
let flaky_rendered = crate::errors::RailwayError::OAuthRefreshFailed(flaky_err).to_string();
assert!(dead_rendered.contains("Couldn't refresh") || dead_rendered.contains("railway login"));
assert!(
flaky_rendered.contains("Couldn't refresh") || flaky_rendered.contains("railway login")
);
assert_eq!(
flaky.hits(),
1,
"legacy policy does not retry a transient 5xx"
);
}
#[tokio::test]
async fn fixed_transient_5xx_retries_and_preserves_credentials() {
let server = MockEndpoint::spawn(vec![server_error()]);
let fixture = Fixture::new();
let mut configs = fixture.load();
let outcome = refresh_with_policy(&mut configs, &server.base_url, Duration::ZERO).await;
assert!(
matches!(outcome, RefreshOutcome::Transient(_)),
"a 5xx must be transient, got {outcome:?}"
);
assert_eq!(
server.hits(),
oauth::REFRESH_MAX_ATTEMPTS as usize,
"fixed policy retries transient failures"
);
let reloaded = fixture.load();
assert_eq!(reloaded.get_refresh_token(), Some("the-refresh-token"));
assert!(reloaded.has_oauth_token());
}
#[tokio::test]
async fn fixed_policy_recovers_when_the_token_endpoint_comes_back() {
let server = MockEndpoint::spawn(vec![server_error(), server_error(), fresh_tokens()]);
let fixture = Fixture::new();
let mut configs = fixture.load();
let outcome = refresh_with_policy(&mut configs, &server.base_url, Duration::ZERO).await;
assert!(
matches!(outcome, RefreshOutcome::Refreshed),
"the user should never notice a brief outage, got {outcome:?}"
);
let reloaded = fixture.load();
assert_eq!(reloaded.get_refresh_token(), Some("new-refresh"));
assert!(
!reloaded.is_token_expired(),
"fresh token must not be expired"
);
}
#[tokio::test]
async fn only_invalid_grant_clears_credentials() {
for code in [
"invalid_client",
"unauthorized_client",
"invalid_scope",
"invalid_request",
"unsupported_grant_type",
"server_error",
"temporarily_unavailable",
"slow_down",
"unknown",
] {
let body = format!(r#"{{"error":"{code}","error_description":"boom"}}"#);
let server = MockEndpoint::spawn(vec![(400, body)]);
let fixture = Fixture::new();
let mut configs = fixture.load();
let outcome = refresh_with_policy(&mut configs, &server.base_url, Duration::ZERO).await;
assert!(
matches!(outcome, RefreshOutcome::Transient(_)),
"{code} must not be treated as a permanently dead credential, got {outcome:?}"
);
assert_eq!(
fixture.load().get_refresh_token(),
Some("the-refresh-token"),
"{code} must leave the refresh token on disk"
);
}
}
#[tokio::test]
async fn fixed_policy_treats_unparseable_4xx_as_transient() {
let server = MockEndpoint::spawn(vec![(400, "<html>blocked by proxy</html>".to_string())]);
let fixture = Fixture::new();
let mut configs = fixture.load();
let outcome = refresh_with_policy(&mut configs, &server.base_url, Duration::ZERO).await;
assert!(
matches!(outcome, RefreshOutcome::Transient(_)),
"unparseable 4xx must not clear credentials, got {outcome:?}"
);
assert_eq!(
fixture.load().get_refresh_token(),
Some("the-refresh-token")
);
}
#[tokio::test]
async fn stale_writer_cannot_resurrect_cleared_credentials() {
let server = MockEndpoint::spawn(vec![dead_grant()]);
let fixture = Fixture::new();
let mut stale_process = fixture.load();
assert!(stale_process.has_oauth_token());
let mut clearing_process = fixture.load();
let outcome =
refresh_with_policy(&mut clearing_process, &server.base_url, Duration::ZERO).await;
assert!(matches!(outcome, RefreshOutcome::SessionExpired(_)));
assert_eq!(fixture.load().get_refresh_token(), None, "clear persisted");
stale_process.root_config.user.id = Some("some-user".to_string());
stale_process.write().unwrap();
let after = fixture.load();
assert_eq!(
after.get_refresh_token(),
None,
"the stale writer must not resurrect the dead refresh token"
);
assert!(!after.has_oauth_token());
assert_eq!(after.root_config.user.id.as_deref(), Some("some-user"));
}
#[tokio::test]
async fn stale_writer_cannot_undo_a_refresh() {
let server = MockEndpoint::spawn(vec![fresh_tokens()]);
let fixture = Fixture::new();
let mut stale_process = fixture.load();
let mut refreshing = fixture.load();
assert!(matches!(
refresh_with_policy(&mut refreshing, &server.base_url, Duration::ZERO).await,
RefreshOutcome::Refreshed
));
stale_process.root_config.user.id = Some("some-user".to_string());
stale_process.write().unwrap();
let after = fixture.load();
assert_eq!(
after.get_refresh_token(),
Some("new-refresh"),
"the refreshed credentials must survive an unrelated concurrent write"
);
assert_eq!(
after.get_railway_auth_token().as_deref(),
Some("new-access")
);
}
#[tokio::test]
async fn logout_clears_credentials() {
let fixture = Fixture::new();
let mut configs = fixture.load();
assert!(configs.has_oauth_token());
configs.reset().unwrap();
configs.write_credentials().unwrap();
let after = fixture.load();
assert!(
!after.has_oauth_token(),
"logout must erase the access token"
);
assert_eq!(after.get_refresh_token(), None);
}
#[tokio::test]
async fn plain_write_cannot_erase_credentials() {
let fixture = Fixture::new();
let mut configs = fixture.load();
configs.reset().unwrap();
configs.write().unwrap();
let after = fixture.load();
assert_eq!(
after.get_refresh_token(),
Some("the-refresh-token"),
"a non-credential write must leave the stored credentials alone"
);
}
#[tokio::test]
async fn baked_in_bearer_ignores_new_credentials_on_disk() {
let backboard = MockEndpoint::spawn(vec![ok_empty()]);
let fixture = Fixture::new();
let startup_configs = fixture.load();
let frozen = crate::client::GQLClient::new_authorized(&startup_configs).unwrap();
let mut relogin = fixture.load();
relogin
.save_oauth_tokens("brand-new-access", Some("brand-new-refresh"), 3600)
.unwrap();
assert_eq!(
fixture.load().get_railway_auth_token().as_deref(),
Some("brand-new-access")
);
let _ = frozen.post(&backboard.base_url).json(&()).send().await;
assert_eq!(
backboard.auth_headers(),
vec!["Bearer stale-access".to_string()],
"EXPECTED DEFECT: the client built at startup keeps using the startup token"
);
let rebuilt = crate::client::GQLClient::new_authorized(&fixture.load()).unwrap();
let _ = rebuilt.post(&backboard.base_url).json(&()).send().await;
assert_eq!(
backboard.auth_headers()[1],
"Bearer brand-new-access",
"a per-request client adopts the new credentials"
);
}
#[tokio::test]
async fn expired_token_refreshes_and_new_bearer_reaches_the_wire() {
let token_endpoint = MockEndpoint::spawn(vec![fresh_tokens()]);
let backboard = MockEndpoint::spawn(vec![ok_empty()]);
let fixture = Fixture::new();
let mut configs = fixture.load();
assert!(configs.is_token_expired(), "fixture starts expired");
let outcome = refresh_with_policy(&mut configs, &token_endpoint.base_url, Duration::ZERO).await;
assert!(matches!(outcome, RefreshOutcome::Refreshed));
let client = crate::client::GQLClient::new_authorized(&fixture.load()).unwrap();
let _ = client.post(&backboard.base_url).json(&()).send().await;
assert_eq!(
backboard.auth_headers(),
vec!["Bearer new-access".to_string()],
"the refreshed token must be the one used for the request"
);
}
#[tokio::test]
async fn a_freshly_saved_token_is_not_considered_expired() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("config.json");
let mut configs = Configs::for_test(path.clone());
configs
.save_oauth_tokens("good-access", Some("good-refresh"), 3600)
.unwrap();
let mut reloaded = Configs::for_test(path);
reloaded.reload().unwrap();
assert!(
!reloaded.is_token_expired(),
"a token minted seconds ago must not be treated as expired, or every \
tool call would refresh"
);
reloaded.root_config.user.token_expires_at = Some(chrono::Utc::now().timestamp() + 30);
assert!(reloaded.is_token_expired());
}
#[tokio::test]
async fn dead_grant_in_a_long_session_refreshes_once_not_per_tool_call() {
let token_endpoint = MockEndpoint::spawn(vec![dead_grant()]);
let fixture = Fixture::new();
for _ in 0..INVOCATIONS {
let mut configs = fixture.load();
if configs.get_refresh_token().is_some() {
refresh_with_policy(&mut configs, &token_endpoint.base_url, Duration::ZERO).await;
}
}
assert_eq!(
token_endpoint.hits(),
1,
"a dead grant must be discovered once per session, not once per tool call"
);
}