use super::*;
use std::{
io::{BufRead, BufReader},
sync::{Barrier, mpsc},
thread,
};
const WAIT: Duration = Duration::from_secs(5);
#[test]
fn shared_provider_concurrent_expired_access_refreshes_once() {
assert_shared_provider_refresh(false);
}
#[test]
fn shared_provider_access_waits_for_forced_refresh() {
assert_shared_provider_refresh(true);
}
fn assert_shared_provider_refresh(force_refresh: bool) {
let temp = tempfile::TempDir::new().unwrap();
let endpoint = RotatingTokenEndpoint::start();
let server_url = format!("{}/mcp", endpoint.base);
let now = Utc::now().timestamp();
write_token(
temp.path(),
"remote",
&StoredToken {
client_id: "fake-client".to_string(),
access_token: "fake-access-0".to_string(),
refresh_token: Some("fake-refresh-0".to_string()),
expires_at: Some(now + if force_refresh { 3600 } else { -60 }),
granted_scopes: Vec::new(),
client_secret: None,
authorization_server: None,
issuer: None,
token_endpoint: Some(format!("{}/token", endpoint.base)),
resource: None,
server_url: server_url.clone(),
token_received_at: now - 3600,
},
)
.unwrap();
let provider = TokenProvider::new(
temp.path().to_path_buf(),
"remote".to_string(),
server_url,
McpOAuthConfig::default(),
reqwest::blocking::Client::builder()
.no_proxy()
.timeout(WAIT)
.build()
.unwrap(),
);
let first_provider = provider.clone();
let first = thread::spawn(move || {
if force_refresh {
first_provider.force_refresh_access_token()
} else {
first_provider.access_token()
}
});
endpoint.request_started.recv_timeout(WAIT).unwrap();
let start = Arc::new(Barrier::new(5));
let (result_tx, result_rx) = mpsc::sync_channel(4);
let callers: Vec<_> = (0..4)
.map(|_| {
let provider = provider.clone();
let start = Arc::clone(&start);
let result_tx = result_tx.clone();
thread::spawn(move || {
start.wait();
result_tx.send(provider.access_token()).unwrap();
})
})
.collect();
drop(result_tx);
start.wait();
let early_result = result_rx.recv_timeout(Duration::from_millis(200));
endpoint.release_response.send(()).unwrap();
let first_result = first.join().unwrap();
for caller in callers {
caller.join().unwrap();
}
let results: Vec<_> = result_rx.try_iter().collect();
let later_result = provider.force_refresh_access_token();
endpoint.stop.send(()).unwrap();
let submitted_tokens = endpoint.handle.join().unwrap();
assert_eq!(submitted_tokens, ["fake-refresh-0", "fake-refresh-1"]);
assert!(matches!(early_result, Err(mpsc::RecvTimeoutError::Timeout)));
assert_eq!(first_result.unwrap(), "fake-access-1");
assert_eq!(results.len(), 4);
for result in results {
assert_eq!(result.unwrap(), "fake-access-1");
}
assert_eq!(later_result.unwrap(), "fake-access-2");
let stored = read_token(temp.path(), "remote").unwrap().unwrap();
assert_eq!(stored.access_token, "fake-access-2");
assert_eq!(stored.refresh_token.as_deref(), Some("fake-refresh-2"));
}
struct RotatingTokenEndpoint {
base: String,
request_started: mpsc::Receiver<()>,
release_response: mpsc::SyncSender<()>,
stop: mpsc::SyncSender<()>,
handle: thread::JoinHandle<Vec<String>>,
}
impl RotatingTokenEndpoint {
fn start() -> Self {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
listener.set_nonblocking(true).unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let (request_started_tx, request_started) = mpsc::sync_channel(1);
let (release_response, release_response_rx) = mpsc::sync_channel(1);
let (stop, stop_rx) = mpsc::sync_channel(1);
let handle = thread::spawn(move || {
let deadline = Instant::now() + WAIT * 4;
let mut generation = 0;
let mut submitted = Vec::new();
while Instant::now() < deadline {
match stop_rx.try_recv() {
Ok(()) | Err(mpsc::TryRecvError::Disconnected) => break,
Err(mpsc::TryRecvError::Empty) => {}
}
let (mut stream, _) = match listener.accept() {
Ok(connection) => connection,
Err(error) if error.kind() == io::ErrorKind::WouldBlock => {
thread::sleep(Duration::from_millis(5));
continue;
}
Err(error) => panic!("accept token request: {error}"),
};
stream.set_read_timeout(Some(WAIT)).unwrap();
stream.set_write_timeout(Some(WAIT)).unwrap();
let form = read_refresh_form(&mut stream);
let refresh = form.get("refresh_token").unwrap().clone();
let valid = refresh == format!("fake-refresh-{generation}");
submitted.push(refresh);
let (status, body) = if valid {
generation += 1;
(
"200 OK",
serde_json::json!({
"access_token": format!("fake-access-{generation}"),
"refresh_token": format!("fake-refresh-{generation}"),
"token_type": "Bearer",
"expires_in": 3600,
}),
)
} else {
(
"400 Bad Request",
serde_json::json!({"error": "invalid_grant"}),
)
};
if submitted.len() == 1 {
request_started_tx.send(()).unwrap();
release_response_rx.recv_timeout(WAIT).unwrap();
}
let body = body.to_string();
write!(
stream,
"HTTP/1.1 {status}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len(),
)
.unwrap();
}
submitted
});
Self {
base,
request_started,
release_response,
stop,
handle,
}
}
}
fn read_refresh_form(stream: &mut TcpStream) -> BTreeMap<String, String> {
let mut reader = BufReader::new(stream);
let mut line = String::new();
reader.read_line(&mut line).unwrap();
assert_eq!(line, "POST /token HTTP/1.1\r\n");
let mut content_length = None;
loop {
line.clear();
assert_ne!(reader.read_line(&mut line).unwrap(), 0);
if line == "\r\n" {
break;
}
if let Some((name, value)) = line.split_once(':')
&& name.eq_ignore_ascii_case("content-length")
{
content_length = Some(value.trim().parse::<usize>().unwrap());
}
}
let length = content_length.unwrap();
assert!(length <= 4096);
let mut body = vec![0; length];
reader.read_exact(&mut body).unwrap();
let mut url = reqwest::Url::parse("http://127.0.0.1/token").unwrap();
url.set_query(Some(std::str::from_utf8(&body).unwrap()));
let form: BTreeMap<_, _> = url.query_pairs().into_owned().collect();
assert_eq!(form.get("grant_type").unwrap(), "refresh_token");
assert_eq!(form.get("client_id").unwrap(), "fake-client");
form
}