magi-code 0.82.1

Repository-aware CLI coding agent for terminal work
Documentation
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() {
    // Session DELETE uses access_token() too: it must wait rather than return the
    // old, locally unexpired token while another request is refreshing it.
    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();
    // All credentials in this test are synthetic and confined to loopback.
    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();
    // Hold the endpoint response open while the other callers enter token access.
    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();

    // A later refresh must use the newly persisted refresh token, not the old one.
    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
}