use std::sync::Arc;
use std::time::{Duration, Instant};
use serde::Serialize;
use serde::de::DeserializeOwned;
use serde_json::Value;
use tokio::sync::Semaphore;
use crate::credentials::{Credentials, DEFAULT_TAILNET};
use crate::error::{ApiError, Idempotence, describe};
use crate::token::Tokens;
pub const DEFAULT_BASE_URL: &str = "https://api.tailscale.com";
pub const DEFAULT_BUDGET: Duration = Duration::from_secs(30);
pub const DEFAULT_CONCURRENCY: usize = 8;
const UNAUTHORIZED: u16 = 401;
const TOO_MANY_REQUESTS: u16 = 429;
const MAX_ATTEMPTS: u32 = 4;
const BASE_BACKOFF: Duration = Duration::from_millis(250);
const MAX_BACKOFF: Duration = Duration::from_secs(20);
const MAX_ERROR_BYTES: usize = 8 * 1024;
#[derive(Debug, Clone)]
pub struct ClientConfig {
pub base_url: String,
pub tailnet: String,
pub credentials: Credentials,
pub budget: Duration,
pub concurrency: usize,
pub max_response_bytes: usize,
pub user_agent: String,
}
impl ClientConfig {
pub fn new(credentials: Credentials) -> Self {
Self {
base_url: DEFAULT_BASE_URL.to_owned(),
tailnet: DEFAULT_TAILNET.to_owned(),
credentials,
budget: DEFAULT_BUDGET,
concurrency: DEFAULT_CONCURRENCY,
max_response_bytes: 1 << 20,
user_agent: format!("tailscale-mcp/{}", env!("CARGO_PKG_VERSION")),
}
}
}
#[derive(Debug, Clone)]
pub struct Client {
inner: Arc<Inner>,
}
#[derive(Debug)]
struct Inner {
http: reqwest::Client,
tokens: Tokens,
base_url: String,
tailnet: String,
budget: Duration,
max_response_bytes: usize,
in_flight: Semaphore,
}
impl Client {
pub fn new(config: ClientConfig) -> Result<Self, ApiError> {
let base_url = checked_base_url(&config.base_url)?;
if config.concurrency == 0 {
return Err(ApiError::Config(
"at least one call has to be allowed in flight".to_owned(),
));
}
if config.max_response_bytes == 0 {
return Err(ApiError::Config(
"a response size cap of zero would reject every answer".to_owned(),
));
}
let http = reqwest::Client::builder()
.user_agent(config.user_agent)
.timeout(config.budget)
.build()
.map_err(|source| {
ApiError::Config(format!("the HTTP client could not be built: {source}"))
})?;
Ok(Self {
inner: Arc::new(Inner {
tokens: Tokens::new(config.credentials, &base_url, http.clone()),
http,
base_url,
tailnet: config.tailnet,
budget: config.budget,
max_response_bytes: config.max_response_bytes,
in_flight: Semaphore::new(config.concurrency),
}),
})
}
pub fn tailnet(&self) -> &str {
&self.inner.tailnet
}
pub fn tailnet_path(&self, tailnet: Option<&str>, rest: &str) -> String {
let tailnet = tailnet.map_or(self.tailnet(), str::trim);
let tailnet = if tailnet.is_empty() {
self.tailnet()
} else {
tailnet
};
format!("/api/v2/tailnet/{}{rest}", escape(tailnet))
}
pub fn get(&self, path: impl Into<String>) -> RequestBuilder<'_> {
self.request(reqwest::Method::GET, path)
}
pub fn post(&self, path: impl Into<String>) -> RequestBuilder<'_> {
self.request(reqwest::Method::POST, path)
}
pub fn put(&self, path: impl Into<String>) -> RequestBuilder<'_> {
self.request(reqwest::Method::PUT, path)
}
pub fn patch(&self, path: impl Into<String>) -> RequestBuilder<'_> {
self.request(reqwest::Method::PATCH, path)
}
pub fn delete(&self, path: impl Into<String>) -> RequestBuilder<'_> {
self.request(reqwest::Method::DELETE, path)
}
fn request(&self, method: reqwest::Method, path: impl Into<String>) -> RequestBuilder<'_> {
RequestBuilder {
client: self,
method,
path: path.into(),
query: Vec::new(),
headers: Vec::new(),
body: None,
budget: self.inner.budget,
broken: None,
}
}
}
#[derive(Debug)]
pub struct RequestBuilder<'a> {
client: &'a Client,
method: reqwest::Method,
path: String,
query: Vec<(String, String)>,
headers: Vec<(String, String)>,
body: Option<Body>,
budget: Duration,
broken: Option<ApiError>,
}
#[derive(Debug, Clone)]
enum Body {
Json(Value),
Text { content_type: String, text: String },
}
impl RequestBuilder<'_> {
#[must_use]
pub fn query(mut self, name: &str, value: impl std::fmt::Display) -> Self {
self.query.push((name.to_owned(), value.to_string()));
self
}
#[must_use]
pub fn maybe_query(self, name: &str, value: Option<impl std::fmt::Display>) -> Self {
match value {
Some(value) => self.query(name, value),
None => self,
}
}
#[must_use]
pub fn header(mut self, name: &str, value: impl Into<String>) -> Self {
self.headers.push((name.to_owned(), value.into()));
self
}
#[must_use]
pub fn json(mut self, body: &impl Serialize) -> Self {
match serde_json::to_value(body) {
Ok(value) => self.body = Some(Body::Json(value)),
Err(source) => {
self.broken.get_or_insert(ApiError::Config(format!(
"the request body could not be built: {source}"
)));
}
}
self
}
#[must_use]
pub fn text(mut self, content_type: &str, body: impl Into<String>) -> Self {
self.body = Some(Body::Text {
content_type: content_type.to_owned(),
text: body.into(),
});
self
}
#[must_use]
pub fn budget(mut self, budget: Duration) -> Self {
self.budget = budget;
self
}
pub async fn send(self) -> Result<Value, ApiError> {
let request = self.describe_request();
let answer = self.send_raw().await?;
parse(&answer.bytes, &request)
}
pub async fn send_as<T: DeserializeOwned>(self) -> Result<T, ApiError> {
Ok(self.send_answer().await?.value)
}
pub async fn send_answer<T: DeserializeOwned>(self) -> Result<Answer<T>, ApiError> {
let request = self.describe_request();
let answer = self.send_raw().await?;
let raw = parse(&answer.bytes, &request)?;
let value =
T::deserialize(&raw).map_err(|source| ApiError::Malformed { request, source })?;
Ok(Answer {
value,
raw,
etag: answer.etag,
})
}
pub async fn send_text(self) -> Result<TextBody, ApiError> {
let answer = self.send_raw().await?;
Ok(TextBody {
text: String::from_utf8_lossy(&answer.bytes).into_owned(),
etag: answer.etag,
})
}
fn describe_request(&self) -> String {
format!("{} {}", self.method, self.path)
}
async fn send_raw(self) -> Result<RawBody, ApiError> {
let request = self.describe_request();
let budget = self.budget;
match tokio::time::timeout(budget, self.attempts()).await {
Ok(answer) => answer,
Err(_) => Err(ApiError::Timeout { request, budget }),
}
}
async fn attempts(self) -> Result<RawBody, ApiError> {
if let Some(broken) = self.broken {
return Err(broken);
}
let request = self.describe_request();
let idempotence = idempotence(&self.method);
let deadline = Instant::now() + self.budget;
let inner = &self.client.inner;
let url = format!("{}{}", inner.base_url, self.path);
let mut attempt = 0;
let mut refreshed = false;
loop {
attempt += 1;
let outcome = self.attempt(&url, &request).await;
let error = match outcome {
Ok(answer) => return Ok(answer),
Err(error) => error,
};
if error.status() == Some(UNAUTHORIZED)
&& inner.tokens.can_refresh()
&& !refreshed
&& attempt < MAX_ATTEMPTS
{
refreshed = true;
tracing::debug!(request = %request, "the token was refused; minting another");
continue;
}
let repeatable =
idempotence == Idempotence::Repeatable || error.status() == Some(TOO_MANY_REQUESTS);
if !error.is_transient() || !repeatable || attempt >= MAX_ATTEMPTS {
return Err(error);
}
let delay = backoff(attempt, &error);
if Instant::now() + delay >= deadline {
return Err(error);
}
tracing::debug!(
request = %request,
attempt,
delay_ms = delay.as_millis(),
because = %error,
"retrying a control-plane call"
);
tokio::time::sleep(delay).await;
}
}
async fn attempt(&self, url: &str, request: &str) -> Result<RawBody, ApiError> {
let inner = &self.client.inner;
let _permit = inner
.in_flight
.acquire()
.await
.map_err(|_| ApiError::Config("the client has been shut down".to_owned()))?;
let bearer = inner.tokens.bearer().await?;
let mut sending = inner
.http
.request(self.method.clone(), url)
.bearer_auth(bearer.value.expose())
.query(&self.query);
for (name, value) in &self.headers {
sending = sending.header(name, value);
}
match &self.body {
Some(Body::Json(value)) => sending = sending.json(value),
Some(Body::Text { content_type, text }) => {
sending = sending
.header(reqwest::header::CONTENT_TYPE, content_type)
.body(text.clone());
}
None => {}
}
let response = sending.send().await.map_err(|source| {
if source.is_timeout() {
ApiError::Timeout {
request: request.to_owned(),
budget: self.budget,
}
} else {
ApiError::Transport {
request: request.to_owned(),
source,
}
}
})?;
let status = response.status();
if status.is_success() {
return read_body(response, request, inner.max_response_bytes).await;
}
if status.as_u16() == UNAUTHORIZED
&& let Some(generation) = bearer.generation
{
inner.tokens.evict(generation).await;
}
let retry_after = retry_after(&response);
let body = read_body(response, request, MAX_ERROR_BYTES)
.await
.map(|raw| String::from_utf8_lossy(&raw.bytes).into_owned())
.unwrap_or_default();
Err(ApiError::Status {
request: request.to_owned(),
status: status.as_u16(),
message: describe(status, &body),
retry_after,
})
}
}
#[derive(Debug)]
struct RawBody {
bytes: Vec<u8>,
etag: Option<String>,
}
fn parse(bytes: &[u8], request: &str) -> Result<Value, ApiError> {
if bytes.iter().all(u8::is_ascii_whitespace) {
return Ok(Value::Null);
}
serde_json::from_slice(bytes).map_err(|source| ApiError::Malformed {
request: request.to_owned(),
source,
})
}
#[derive(Debug, Clone)]
pub struct Answer<T> {
pub value: T,
pub raw: Value,
pub etag: Option<String>,
}
#[derive(Debug, Clone)]
pub struct TextBody {
pub text: String,
pub etag: Option<String>,
}
async fn read_body(
mut response: reqwest::Response,
request: &str,
cap: usize,
) -> Result<RawBody, ApiError> {
let etag = response
.headers()
.get(reqwest::header::ETAG)
.and_then(|value| value.to_str().ok())
.map(str::to_owned);
let too_large = || ApiError::TooLarge {
request: request.to_owned(),
cap,
};
if response
.content_length()
.is_some_and(|len| len > cap as u64)
{
return Err(too_large());
}
let mut bytes = Vec::new();
while let Some(chunk) = response
.chunk()
.await
.map_err(|source| ApiError::Transport {
request: request.to_owned(),
source,
})?
{
if bytes.len() + chunk.len() > cap {
return Err(too_large());
}
bytes.extend_from_slice(&chunk);
}
Ok(RawBody { bytes, etag })
}
fn idempotence(method: &reqwest::Method) -> Idempotence {
match *method {
reqwest::Method::POST | reqwest::Method::PATCH => Idempotence::Once,
_ => Idempotence::Repeatable,
}
}
fn backoff(attempt: u32, error: &ApiError) -> Duration {
if let ApiError::Status {
retry_after: Some(asked),
..
} = error
{
return (*asked).min(MAX_BACKOFF);
}
(BASE_BACKOFF * 2u32.saturating_pow(attempt - 1)).min(MAX_BACKOFF)
}
fn retry_after(response: &reqwest::Response) -> Option<Duration> {
response
.headers()
.get(reqwest::header::RETRY_AFTER)?
.to_str()
.ok()?
.trim()
.parse::<u64>()
.ok()
.map(Duration::from_secs)
}
pub fn checked_base_url(base_url: &str) -> Result<String, ApiError> {
let trimmed = base_url.trim().trim_end_matches('/');
let parsed = reqwest::Url::parse(trimmed)
.map_err(|source| ApiError::Config(format!("`{base_url}` is not a URL: {source}")))?;
let loopback = parsed.host_str().is_some_and(|host| {
let host = host.trim_start_matches('[').trim_end_matches(']');
host.eq_ignore_ascii_case("localhost")
|| host
.parse::<std::net::IpAddr>()
.is_ok_and(|address| address.is_loopback())
});
if parsed.scheme() != "https" && !loopback {
return Err(ApiError::Config(format!(
"`{base_url}` is neither https nor a loopback address, and a \
control-plane credential is not sent anywhere else"
)));
}
if !parsed.path().is_empty() && parsed.path() != "/" {
return Err(ApiError::Config(format!(
"`{base_url}` has a path; the base URL is a host and nothing more"
)));
}
if !parsed.username().is_empty() || parsed.password().is_some() {
return Err(ApiError::Config(
"the base URL carries a username or password; a control-plane \
credential is sent as a header and never in a URL"
.to_owned(),
));
}
Ok(trimmed.to_owned())
}
fn escape(segment: &str) -> String {
let mut out = String::with_capacity(segment.len());
for byte in segment.bytes() {
if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'~' | b'@') {
out.push(char::from(byte));
} else {
out.push_str(&format!("%{byte:02X}"));
}
}
out
}
#[cfg(test)]
mod tests {
use std::collections::BTreeSet;
use std::path::PathBuf;
use serde_json::json;
use super::*;
use crate::fake::{FakeControlPlane, Response};
use crate::secret::Secret;
const KEY: &str = "tskey-api-redacted-example";
const DEVICES: &str = "/api/v2/tailnet/-/devices";
async fn fake() -> FakeControlPlane {
FakeControlPlane::start()
.await
.expect("a loopback socket is available")
}
fn client(fake: &FakeControlPlane, credentials: Credentials) -> Client {
client_with(fake, credentials, |_| {})
}
fn client_with(
fake: &FakeControlPlane,
credentials: Credentials,
adjust: impl FnOnce(&mut ClientConfig),
) -> Client {
let mut config = ClientConfig::new(credentials);
config.base_url = fake.base_url().to_owned();
adjust(&mut config);
Client::new(config).expect("the fake answers on a loopback address")
}
fn api_key() -> Credentials {
Credentials::ApiKey(Secret::new(KEY))
}
fn oauth() -> Credentials {
Credentials::OauthClient {
client_id: "kExAmPlE1CNTRL".to_owned(),
client_secret: Secret::new("tskey-client-redacted-example"),
scopes: vec!["devices:read".to_owned(), "dns".to_owned()],
}
}
fn token(value: &str, seconds: u64) -> Response {
Response::json(json!({
"access_token": value,
"token_type": "Bearer",
"expires_in": seconds,
}))
}
fn bearers(fake: &FakeControlPlane) -> Vec<String> {
fake.recorded()
.into_iter()
.filter_map(|r| r.authorization().map(str::to_owned))
.collect()
}
#[tokio::test]
async fn an_api_key_is_the_bearer_token_itself() {
let fake = fake()
.await
.on("GET", DEVICES, Response::json(json!({"devices": []})));
let client = client(&fake, api_key());
let answer = client.get(DEVICES).send().await.expect("the fake answers");
assert_eq!(answer, json!({"devices": []}));
let request = fake.only_request();
assert_eq!(
request.authorization(),
Some(format!("Bearer {KEY}").as_str())
);
}
#[tokio::test]
async fn an_oauth_client_is_exchanged_for_a_token() {
let fake = fake()
.await
.on("POST", crate::token::TOKEN_PATH, token("minted-1", 3600))
.on("GET", DEVICES, Response::json(json!({"devices": []})));
let client = client(&fake, oauth());
client.get(DEVICES).send().await.expect("the fake answers");
let recorded = fake.recorded();
assert_eq!(
recorded.len(),
2,
"an exchange and then the call: {recorded:#?}"
);
let exchange = &recorded[0];
assert_eq!(exchange.path, crate::token::TOKEN_PATH);
for expected in [
"grant_type=client_credentials",
"client_id=kExAmPlE1CNTRL",
"client_secret=tskey-client-redacted-example",
"scope=devices%3Aread+dns",
] {
assert!(
exchange.body.contains(expected),
"the exchange did not send `{expected}`: {}",
exchange.body
);
}
assert_eq!(recorded[1].authorization(), Some("Bearer minted-1"));
}
#[tokio::test]
async fn a_federated_identity_signs_with_the_jwt_on_disk() {
let directory = tempfile::tempdir().expect("a temporary directory");
let jwt_file = directory.path().join("token");
std::fs::write(&jwt_file, "header.payload.signature\n").expect("the file is written");
let fake = fake()
.await
.on(
"POST",
crate::token::TOKEN_PATH,
token("minted-federated", 3600),
)
.on("GET", DEVICES, Response::json(json!({"devices": []})));
let client = client(
&fake,
Credentials::Federated {
client_id: Some("kExAmPlE1CNTRL".to_owned()),
jwt_file,
scopes: Vec::new(),
},
);
client.get(DEVICES).send().await.expect("the fake answers");
let exchange = &fake.recorded()[0];
assert!(
exchange
.body
.contains("client_assertion=header.payload.signature"),
"the JWT was not sent, or was sent with its trailing newline: {}",
exchange.body
);
assert!(
exchange.body.contains("client_assertion_type=urn%3Aietf"),
"the assertion type was not sent: {}",
exchange.body
);
}
#[tokio::test]
async fn a_missing_jwt_file_says_which_file_it_was() {
let fake = fake().await;
let client = client(
&fake,
Credentials::Federated {
client_id: None,
jwt_file: PathBuf::from("/nonexistent/identity/token"),
scopes: Vec::new(),
},
);
let error = client
.get(DEVICES)
.send()
.await
.expect_err("there is no file");
assert!(
matches!(&error, ApiError::JwtFile { path, .. } if path.ends_with("token")),
"unexpected error: {error:?}"
);
assert_eq!(fake.request_count(), 0, "nothing should have been sent");
}
#[tokio::test]
async fn the_credential_with_precedence_is_the_one_that_is_used() {
let environment = |key: &str| match key {
crate::credentials::API_KEY_ENV => Some(KEY.to_owned()),
crate::credentials::OAUTH_CLIENT_ID_ENV => Some("kExAmPlE1CNTRL".to_owned()),
crate::credentials::OAUTH_CLIENT_SECRET_ENV => Some("unused".to_owned()),
_ => None,
};
let credentials = Credentials::from_source(environment).expect("both are set");
let fake = fake().await.on("GET", DEVICES, Response::json(json!({})));
let client = client(&fake, credentials);
client.get(DEVICES).send().await.expect("the fake answers");
let request = fake.only_request();
assert_eq!(
request.authorization(),
Some(format!("Bearer {KEY}").as_str())
);
}
#[tokio::test]
async fn a_token_is_minted_once_and_reused() {
let fake = fake()
.await
.on("POST", crate::token::TOKEN_PATH, token("minted-1", 3600))
.on("GET", DEVICES, Response::json(json!({})));
let client = client(&fake, oauth());
for _ in 0..3 {
client.get(DEVICES).send().await.expect("the fake answers");
}
let exchanges = fake
.recorded()
.iter()
.filter(|r| r.path == crate::token::TOKEN_PATH)
.count();
assert_eq!(exchanges, 1, "the token should have been minted once");
assert_eq!(bearers(&fake), vec!["Bearer minted-1".to_owned(); 3]);
}
#[tokio::test]
async fn a_token_near_its_expiry_is_minted_again() {
let remaining = crate::token::REFRESH_SKEW.as_secs() / 2;
let fake = fake()
.await
.on(
"POST",
crate::token::TOKEN_PATH,
token("minted-1", remaining),
)
.on("GET", DEVICES, Response::json(json!({})));
let client = client(&fake, oauth());
for _ in 0..2 {
client.get(DEVICES).send().await.expect("the fake answers");
}
let exchanges = fake
.recorded()
.iter()
.filter(|r| r.path == crate::token::TOKEN_PATH)
.count();
assert_eq!(exchanges, 2, "a token inside the skew should not be reused");
}
#[tokio::test]
async fn a_refused_token_is_replaced_exactly_once() {
let fake = fake()
.await
.once("POST", crate::token::TOKEN_PATH, token("stale", 3600))
.on("POST", crate::token::TOKEN_PATH, token("fresh", 3600))
.once(
"GET",
DEVICES,
Response::status(401, json!({"message": "expired"})),
)
.on("GET", DEVICES, Response::json(json!({"devices": []})));
let client = client(&fake, oauth());
let answer = client
.get(DEVICES)
.send()
.await
.expect("the second try works");
assert_eq!(answer, json!({"devices": []}));
assert_eq!(
bearers(&fake),
vec!["Bearer stale".to_owned(), "Bearer fresh".to_owned()],
"the refused token should have been replaced, once"
);
}
#[tokio::test]
async fn a_token_refused_twice_is_the_credential_being_wrong() {
let fake = fake()
.await
.on("POST", crate::token::TOKEN_PATH, token("minted", 3600))
.on(
"GET",
DEVICES,
Response::status(401, json!({"message": "no"})),
);
let client = client(&fake, oauth());
let error = client
.get(DEVICES)
.send()
.await
.expect_err("it is always refused");
assert_eq!(error.status(), Some(401));
let calls = fake.recorded().iter().filter(|r| r.path == DEVICES).count();
assert_eq!(calls, 2, "one retry with a fresh token, and then no more");
}
#[tokio::test]
async fn a_refused_api_key_is_not_replaced_because_there_is_nothing_to_mint() {
let fake = fake().await.on(
"GET",
DEVICES,
Response::status(401, json!({"message": "no"})),
);
let client = client(&fake, api_key());
let error = client.get(DEVICES).send().await.expect_err("it is refused");
assert_eq!(error.status(), Some(401));
assert_eq!(fake.request_count(), 1);
}
#[tokio::test]
async fn a_transient_failure_on_a_repeatable_method_is_retried() {
let fake = fake()
.await
.once(
"GET",
DEVICES,
Response::status(503, json!({"message": "later"})),
)
.on("GET", DEVICES, Response::json(json!({"devices": []})));
let client = client(&fake, api_key());
let answer = client
.get(DEVICES)
.send()
.await
.expect("the second try works");
assert_eq!(answer, json!({"devices": []}));
assert_eq!(fake.request_count(), 2);
}
#[tokio::test]
async fn a_transient_failure_on_an_unsafe_method_is_not_retried() {
let keys = "/api/v2/tailnet/-/keys";
let fake = fake().await.on(
"POST",
keys,
Response::status(503, json!({"message": "later"})),
);
let client = client(&fake, api_key());
let error = client
.post(keys)
.json(&json!({"capabilities": {}}))
.send()
.await
.expect_err("the fake never succeeds");
assert_eq!(error.status(), Some(503));
assert_eq!(fake.request_count(), 1, "a POST must not be sent twice");
}
#[tokio::test]
async fn a_rate_limit_is_retried_even_on_an_unsafe_method() {
let keys = "/api/v2/tailnet/-/keys";
let fake = fake()
.await
.once(
"POST",
keys,
Response::status(429, json!({"message": "slow down"}))
.with_header("retry-after", "0"),
)
.on(
"POST",
keys,
Response::json(json!({"key": "tskey-auth-redacted-example"})),
);
let client = client(&fake, api_key());
let answer = client
.post(keys)
.json(&json!({"capabilities": {}}))
.send()
.await
.expect("the second try works");
assert_eq!(answer["key"], json!("tskey-auth-redacted-example"));
assert_eq!(fake.request_count(), 2);
}
#[tokio::test]
async fn a_permanent_failure_is_not_retried() {
let fake = fake().await.on(
"GET",
DEVICES,
Response::status(404, json!({"message": "no such tailnet"})),
);
let client = client(&fake, api_key());
let error = client
.get(DEVICES)
.send()
.await
.expect_err("there is nothing there");
assert!(
matches!(&error, ApiError::Status { message, .. } if message == "no such tailnet"),
"the API's own message should be passed on: {error:?}"
);
assert_eq!(fake.request_count(), 1);
}
#[tokio::test]
async fn a_call_stops_after_a_bounded_number_of_attempts() {
let fake = fake().await.on(
"GET",
DEVICES,
Response::status(503, json!({"message": "later"})).with_header("retry-after", "0"),
);
let client = client(&fake, api_key());
let error = client
.get(DEVICES)
.send()
.await
.expect_err("it never works");
assert_eq!(error.status(), Some(503));
assert_eq!(fake.request_count(), MAX_ATTEMPTS as usize);
}
#[tokio::test]
async fn retrying_stops_when_the_budget_would_not_cover_the_wait() {
let fake = fake().await.on(
"GET",
DEVICES,
Response::status(503, json!({"message": "later"})),
);
let client = client(&fake, api_key());
let error = client
.get(DEVICES)
.budget(Duration::from_millis(50))
.send()
.await
.expect_err("it never works");
assert_eq!(error.status(), Some(503), "not a bare timeout: {error:?}");
assert_eq!(
fake.request_count(),
1,
"the first backoff is longer than the budget"
);
}
#[tokio::test]
async fn the_wait_a_server_asks_for_is_read_off_the_wire() {
let budget = Duration::from_secs(1);
let refusal = |wait| {
Response::status(503, json!({"message": "later"})).with_header("retry-after", wait)
};
let patient = fake().await.on("GET", DEVICES, refusal("300"));
let error = client(&patient, api_key())
.get(DEVICES)
.budget(budget)
.send()
.await
.expect_err("it never works");
assert_eq!(error.status(), Some(503), "not a bare timeout: {error:?}");
assert_eq!(
patient.request_count(),
1,
"the server asked for longer than the budget, so there was no second try; \
ignoring the header would have waited {BASE_BACKOFF:?} and tried again"
);
let impatient = fake().await.on("GET", DEVICES, refusal("0"));
client(&impatient, api_key())
.get(DEVICES)
.budget(budget)
.send()
.await
.expect_err("it never works");
assert_eq!(
impatient.request_count(),
MAX_ATTEMPTS as usize,
"a server asking for no wait at all should be believed too"
);
}
#[test]
fn the_server_is_believed_about_when_to_come_back() {
let asked = |seconds| ApiError::Status {
request: "GET /x".to_owned(),
status: 429,
message: String::new(),
retry_after: Some(Duration::from_secs(seconds)),
};
let guessed = ApiError::Status {
request: "GET /x".to_owned(),
status: 503,
message: String::new(),
retry_after: None,
};
assert_eq!(backoff(1, &asked(5)), Duration::from_secs(5));
assert_eq!(backoff(1, &asked(600)), MAX_BACKOFF);
assert_eq!(backoff(1, &guessed), BASE_BACKOFF);
assert_eq!(backoff(2, &guessed), BASE_BACKOFF * 2);
assert_eq!(backoff(30, &guessed), MAX_BACKOFF);
}
#[test]
fn only_the_methods_http_calls_idempotent_may_be_repeated() {
for method in [
reqwest::Method::GET,
reqwest::Method::HEAD,
reqwest::Method::PUT,
reqwest::Method::DELETE,
] {
assert_eq!(idempotence(&method), Idempotence::Repeatable, "{method}");
}
for method in [reqwest::Method::POST, reqwest::Method::PATCH] {
assert_eq!(idempotence(&method), Idempotence::Once, "{method}");
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn no_more_calls_are_in_flight_than_the_limit_allows() {
const LIMIT: usize = 2;
let fake = fake().await.on(
"GET",
DEVICES,
Response::json(json!({})).slow(Duration::from_millis(80)),
);
let client = client_with(&fake, api_key(), |config| config.concurrency = LIMIT);
let calls: Vec<_> = (0..8)
.map(|_| {
let client = client.clone();
tokio::spawn(async move { client.get(DEVICES).send().await })
})
.collect();
for call in calls {
call.await
.expect("the task finished")
.expect("the fake answers");
}
assert_eq!(fake.request_count(), 8);
let peak = fake.peak_concurrency();
assert!(
(1..=LIMIT).contains(&peak),
"{peak} calls were in flight at once, and the limit is {LIMIT}"
);
}
#[tokio::test]
async fn an_answer_over_the_cap_is_refused_rather_than_truncated() {
let big = json!({"devices": vec![json!({"name": "x".repeat(200)})]});
let fake = fake().await.on("GET", DEVICES, Response::json(&big));
let client = client_with(&fake, api_key(), |config| config.max_response_bytes = 64);
let error = client
.get(DEVICES)
.send()
.await
.expect_err("it is too large");
assert!(
matches!(error, ApiError::TooLarge { cap: 64, .. }),
"a truncated body would have failed to parse instead: {error:?}"
);
}
#[tokio::test]
async fn an_answer_with_no_stated_length_is_refused_while_it_is_read() {
let big = json!({"devices": vec![json!({"name": "x".repeat(200)})]});
let fake = fake()
.await
.on("GET", DEVICES, Response::json(&big).chunked());
let client = client_with(&fake, api_key(), |config| config.max_response_bytes = 64);
let error = client
.get(DEVICES)
.send()
.await
.expect_err("it is too large");
assert!(
matches!(error, ApiError::TooLarge { cap: 64, .. }),
"unexpected error: {error:?}"
);
}
#[tokio::test]
async fn an_answer_under_the_cap_arrives_whole_however_it_is_framed() {
let body = json!({"devices": [{"name": "workstation"}]});
let fake = fake()
.await
.on("GET", DEVICES, Response::json(&body).chunked());
let client = client(&fake, api_key());
let answer = client.get(DEVICES).send().await.expect("the fake answers");
assert_eq!(answer, body);
}
#[tokio::test]
async fn an_empty_body_is_an_answer_rather_than_a_parse_failure() {
let device = "/api/v2/device/n1111111CNTRL";
let fake = fake().await.on("DELETE", device, Response::empty());
let client = client(&fake, api_key());
let answer = client
.delete(device)
.send()
.await
.expect("the fake answers");
assert_eq!(answer, Value::Null, "a deletion answers with nothing");
}
#[tokio::test]
async fn an_answer_is_read_both_ways_from_one_parse() {
let body = json!({
"id": "kExAmPlE",
"description": "a key",
"invented": {"by": "a later control plane"},
});
let keys = "/api/v2/tailnet/-/keys/kExAmPlE";
let fake = fake().await.on("GET", keys, Response::json(&body));
let client = client(&fake, api_key());
let answer = client
.get(keys)
.send_answer::<crate::models::key::Key>()
.await
.expect("the fake answers");
assert_eq!(answer.value.id.as_deref(), Some("kExAmPlE"));
assert_eq!(
answer.value.unknown.get("invented"),
Some(&json!({"by": "a later control plane"})),
"the typed half keeps what it had no field for"
);
assert_eq!(answer.raw, body, "and the raw half is the body, untouched");
}
#[tokio::test]
async fn an_answer_carries_the_etag_that_versions_it() {
let acl = "/api/v2/tailnet/-/acl";
let fake = fake().await.on(
"GET",
acl,
Response::json(json!({"acls": []})).with_header("ETag", "\"abc123\""),
);
let client = client(&fake, api_key());
let answer = client
.get(acl)
.send_answer::<Value>()
.await
.expect("the fake answers");
assert_eq!(answer.etag.as_deref(), Some("\"abc123\""));
}
#[tokio::test]
async fn an_empty_body_answers_as_nothing_rather_than_failing_to_parse() {
let device = "/api/v2/device/n1111111CNTRL";
let fake = fake().await.on("DELETE", device, Response::empty());
let client = client(&fake, api_key());
let answer = client
.delete(device)
.send_answer::<Value>()
.await
.expect("the fake answers");
assert_eq!(answer.value, Value::Null);
assert_eq!(answer.raw, Value::Null, "both halves agree about nothing");
}
#[tokio::test]
async fn the_query_and_the_body_reach_the_control_plane_as_written() {
let fake = fake().await.on("POST", DEVICES, Response::json(json!({})));
let client = client(&fake, api_key());
client
.post(DEVICES)
.query("fields", "all")
.maybe_query("since", Some(7))
.maybe_query("until", Option::<u8>::None)
.header("If-Match", "\"v1\"")
.json(&json!({"name": "workstation"}))
.send()
.await
.expect("the fake answers");
let request = fake.only_request();
assert_eq!(
request
.query
.keys()
.map(String::as_str)
.collect::<BTreeSet<_>>(),
BTreeSet::from(["fields", "since"]),
"an absent parameter should not be sent"
);
assert_eq!(request.query["fields"], "all");
assert_eq!(
request.headers.get("if-match").map(String::as_str),
Some("\"v1\"")
);
assert_eq!(request.json(), json!({"name": "workstation"}));
}
#[tokio::test]
async fn text_comes_back_with_the_version_it_was_read_at() {
let policy = "/api/v2/tailnet/-/acl";
let hujson = "{\n // a comment, which JSON does not have\n \"acls\": [],\n}";
let fake = fake().await.on(
"GET",
policy,
Response {
status: 200,
headers: vec![("content-type".to_owned(), "application/hujson".to_owned())],
body: hujson.to_owned(),
delay: Duration::ZERO,
chunked: false,
}
.with_header("etag", "\"abc123\""),
);
let client = client(&fake, api_key());
let answer = client
.get(policy)
.send_text()
.await
.expect("the fake answers");
assert_eq!(answer.text, hujson);
assert_eq!(answer.etag.as_deref(), Some("\"abc123\""));
}
#[tokio::test]
async fn a_body_that_is_not_what_was_asked_for_says_so() {
#[derive(Debug, serde::Deserialize)]
struct Devices {
#[allow(dead_code)]
devices: Vec<String>,
}
let fake = fake()
.await
.on("GET", DEVICES, Response::json(json!({"devices": 7})));
let client = client(&fake, api_key());
let error = client
.get(DEVICES)
.send_as::<Devices>()
.await
.expect_err("seven is not a list");
assert!(
matches!(&error, ApiError::Malformed { request, .. } if request == "GET /api/v2/tailnet/-/devices"),
"unexpected error: {error:?}"
);
}
#[test]
fn a_base_url_is_an_encrypted_host_and_nothing_more() {
for allowed in [
DEFAULT_BASE_URL,
"https://api.example.com",
"https://example.com",
"http://127.0.0.1:8080",
"http://localhost:9999",
"http://[::1]:1234",
] {
assert!(
checked_base_url(allowed).is_ok(),
"{allowed} should have been accepted"
);
}
for refused in [
"http://api.tailscale.com",
"http://evil.example.com",
"api.tailscale.com",
"ftp://api.tailscale.com",
"https://api.tailscale.com/api/v2",
"https://user:pass@example.com",
"https://token@example.com",
] {
assert!(
checked_base_url(refused).is_err(),
"{refused} should have been refused"
);
}
assert_eq!(
checked_base_url("https://api.tailscale.com/").expect("a valid URL"),
DEFAULT_BASE_URL
);
}
#[test]
fn a_name_in_a_path_cannot_reach_into_the_path_around_it() {
let fake_config = ClientConfig::new(api_key());
let client = Client::new(fake_config).expect("the default base URL is valid");
assert_eq!(
client.tailnet_path(None, "/devices"),
"/api/v2/tailnet/-/devices"
);
assert_eq!(
client.tailnet_path(Some("example.com"), "/dns/nameservers"),
"/api/v2/tailnet/example.com/dns/nameservers"
);
assert_eq!(
client.tailnet_path(Some(" "), "/devices"),
"/api/v2/tailnet/-/devices"
);
assert_eq!(
client.tailnet_path(Some("../../device/n1111111CNTRL"), "/devices"),
"/api/v2/tailnet/..%2F..%2Fdevice%2Fn1111111CNTRL/devices"
);
}
#[test]
fn a_client_that_could_not_work_is_refused_at_the_start() {
for (what, adjust) in [
(
"no calls in flight",
Box::new(|c: &mut ClientConfig| c.concurrency = 0) as Box<dyn FnOnce(&mut _)>,
),
(
"no bytes allowed back",
Box::new(|c: &mut ClientConfig| c.max_response_bytes = 0),
),
] {
let mut config = ClientConfig::new(api_key());
adjust(&mut config);
assert!(
matches!(Client::new(config), Err(ApiError::Config(_))),
"{what} should have been refused"
);
}
}
}