use anyhow::{Context, Result};
use serde::{Deserialize, Serialize};
pub struct HeadscaleClient {
base_url: String,
api_key: String,
http: reqwest::Client,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PreauthKey {
pub key: String,
pub acl_tags: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NodeInfo {
pub id: String,
pub name: String,
pub ip_addresses: Vec<String>,
pub online: bool,
}
#[derive(Debug, Clone)]
pub struct HeadscaleHealth {
pub reachable: bool,
pub status_code: u16,
}
impl HeadscaleClient {
pub fn new(base_url: impl Into<String>, api_key: impl Into<String>) -> Result<Self> {
let base_url = base_url.into().trim_end_matches('/').to_string();
let http = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(30))
.build()
.context("building Headscale HTTP client")?;
Ok(Self {
base_url,
api_key: api_key.into(),
http,
})
}
pub fn from_vault_or_env() -> Result<Option<Self>> {
let api_key = match fob::get_or_env("headscale-api-key", "HEADSCALE_API_KEY")? {
Some(k) => k,
None => return Ok(None),
};
let url = match fob::get_or_env("mesh-url", "HEADSCALE_URL")? {
Some(u) => u,
None => return Ok(None),
};
Ok(Some(Self::new(url, api_key)?))
}
pub async fn create_preauth_key(&self, user: &str, tags: &[String]) -> Result<PreauthKey> {
let expiration = chrono::Utc::now()
+ chrono::TimeDelta::try_hours(1)
.ok_or_else(|| anyhow::anyhow!("overflow computing 1-hour expiry"))?;
let body = serde_json::json!({
"user": user,
"expiration": expiration.to_rfc3339(),
"reusable": false,
"ephemeral": false,
"aclTags": tags,
});
let resp = self
.http
.post(format!("{}/api/v1/preauthkey", self.base_url))
.bearer_auth(&self.api_key)
.json(&body)
.send()
.await
.context("POST /api/v1/preauthkey")?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Headscale API {status} on POST /api/v1/preauthkey: {text}");
}
let resp_json: serde_json::Value =
resp.json().await.context("parsing preauthkey response")?;
let key = resp_json["preAuthKey"]["key"]
.as_str()
.ok_or_else(|| anyhow::anyhow!("missing preAuthKey.key in response: {resp_json}"))?;
Ok(PreauthKey {
key: key.to_string(),
acl_tags: tags.to_vec(),
})
}
pub async fn list_nodes(&self) -> Result<Vec<NodeInfo>> {
let resp = self
.http
.get(format!("{}/api/v1/node", self.base_url))
.bearer_auth(&self.api_key)
.send()
.await
.context("GET /api/v1/node")?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Headscale API {status} on GET /api/v1/node: {text}");
}
let resp_json: serde_json::Value = resp.json().await.context("parsing node list")?;
let nodes = resp_json["nodes"].as_array().cloned().unwrap_or_default();
Ok(nodes
.iter()
.filter_map(|m| {
Some(NodeInfo {
id: m["id"].as_str()?.to_string(),
name: m["name"].as_str()?.to_string(),
ip_addresses: m["ipAddresses"]
.as_array()
.map(|arr| {
arr.iter()
.filter_map(|v| v.as_str().map(str::to_string))
.collect()
})
.unwrap_or_default(),
online: m["online"].as_bool().unwrap_or(false),
})
})
.collect())
}
pub async fn list_users(&self) -> Result<Vec<String>> {
let resp = self
.http
.get(format!("{}/api/v1/user", self.base_url))
.bearer_auth(&self.api_key)
.send()
.await
.context("GET /api/v1/user")?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Headscale API {status} on GET /api/v1/user: {text}");
}
let resp_json: serde_json::Value = resp.json().await.context("parsing user list")?;
Ok(resp_json["users"]
.as_array()
.cloned()
.unwrap_or_default()
.iter()
.filter_map(|u| u["name"].as_str().map(str::to_string))
.collect())
}
pub async fn resolve_user(&self, preferred: Option<&str>) -> Result<String> {
let users = self.list_users().await?;
choose_user(&self.base_url, &users, preferred)
}
pub async fn health(&self) -> HeadscaleHealth {
match self
.http
.get(format!("{}/health", self.base_url))
.send()
.await
{
Ok(r) => HeadscaleHealth {
reachable: r.status().is_success(),
status_code: r.status().as_u16(),
},
Err(_) => HeadscaleHealth {
reachable: false,
status_code: 0,
},
}
}
}
fn choose_user(base_url: &str, users: &[String], preferred: Option<&str>) -> Result<String> {
if users.is_empty() {
anyhow::bail!(
"Headscale at {base_url} has no users — preauth keys are scoped to one. \
Create it with `headscale users create yah`."
);
}
if let Some(want) = preferred {
if users.iter().any(|u| u == want) {
return Ok(want.to_string());
}
anyhow::bail!(
"Headscale at {base_url} has no user '{want}' — preauth keys are scoped to \
a user, and minting against a missing one fails with an opaque 500. \
Existing users: {}.",
users.join(", ")
);
}
if let [only] = users {
return Ok(only.clone());
}
anyhow::bail!(
"Headscale at {base_url} has several users ({}) — pass one explicitly rather \
than letting this pick.",
users.join(", ")
)
}
pub fn generate_headscale_config(server_url: &str, data_dir: &std::path::Path) -> String {
let private_key = data_dir.join("private.key").display().to_string();
let noise_key = data_dir.join("noise_private.key").display().to_string();
let db_path = data_dir.join("headscale.db").display().to_string();
let socket_path = data_dir.join("headscale.sock").display().to_string();
let acl_path = data_dir.join("acls.yaml").display().to_string();
format!(
r#"---
server_url: {server_url}
listen_addr: 127.0.0.1:8080
grpc_listen_addr: 127.0.0.1:50443
metrics_listen_addr: 127.0.0.1:9090
private_key_path: {private_key}
noise:
private_key_path: {noise_key}
database:
type: sqlite
sqlite:
path: {db_path}
unix_socket: {socket_path}
unix_socket_permission: "0770"
dns:
magic_dns: true
base_domain: mesh.internal
nameservers:
global:
- 1.1.1.1
- 8.8.8.8
log:
level: info
prefixes:
v4: 100.64.0.0/10
v6: fd7a:115c:a1e0::/48
allocation: sequential
policy:
mode: file
path: {acl_path}
derp:
server:
enabled: false
urls:
- https://controlplane.tailscale.com/derpmap/default
auto_update_enabled: false
update_frequency: 24h
"#
)
}
pub const DEFAULT_ACL_POLICY: &str = r#"{
"acls": [
{ "action": "accept", "src": ["*"], "dst": ["*:*"] }
]
}
"#;
pub const HEADSCALE_VERSION: &str = "0.23.0";
pub fn headscale_download_url() -> Result<String> {
let os = match std::env::consts::OS {
"macos" => "darwin",
"linux" => "linux",
other => anyhow::bail!(
"unsupported OS '{other}' for `yah mesh start`; \
install headscale manually from https://headscale.net"
),
};
let arch = match std::env::consts::ARCH {
"x86_64" => "amd64",
"aarch64" => "arm64",
other => anyhow::bail!(
"unsupported architecture '{other}' for `yah mesh start`; \
install headscale manually from https://headscale.net"
),
};
Ok(format!(
"https://github.com/juanfont/headscale/releases/download/v{HEADSCALE_VERSION}/headscale_{HEADSCALE_VERSION}_{os}_{arch}"
))
}
pub async fn update_cloudflare_dns(
api_token: &str,
zone_id: &str,
record_name: &str,
new_ip: &str,
) -> Result<()> {
let http = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(30))
.build()
.context("building Cloudflare HTTP client")?;
let list_url = format!("https://api.cloudflare.com/client/v4/zones/{zone_id}/dns_records");
let resp = http
.get(&list_url)
.bearer_auth(api_token)
.query(&[("name", record_name), ("type", "A")])
.send()
.await
.context("GET Cloudflare DNS records")?;
if !resp.status().is_success() {
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Cloudflare list DNS records failed: {text}");
}
let body: serde_json::Value = resp
.json()
.await
.context("parsing Cloudflare records list")?;
let records = body["result"]
.as_array()
.ok_or_else(|| anyhow::anyhow!("unexpected Cloudflare response shape: missing 'result'"))?;
let record_id = match records.len() {
0 => {
println!(" A record '{record_name}' absent — creating it (→ {new_ip}) ...");
let create_body = serde_json::json!({
"type": "A",
"name": record_name,
"content": new_ip,
"ttl": 120,
"proxied": false
});
let resp = http
.post(&list_url)
.bearer_auth(api_token)
.json(&create_body)
.send()
.await
.context("POST Cloudflare DNS record (create-if-missing)")?;
if !resp.status().is_success() {
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Cloudflare create DNS record failed: {text}");
}
println!(" DNS record created.");
return Ok(());
}
1 => records[0]["id"]
.as_str()
.ok_or_else(|| anyhow::anyhow!("missing 'id' in Cloudflare record"))?
.to_string(),
n => anyhow::bail!(
"{n} A records named '{record_name}' found — expected exactly one; \
resolve the ambiguity in your Cloudflare dashboard"
),
};
let patch_url = format!("{list_url}/{record_id}");
let patch_body = serde_json::json!({
"type": "A",
"name": record_name,
"content": new_ip,
"ttl": 120,
"proxied": false
});
let resp = http
.patch(&patch_url)
.bearer_auth(api_token)
.json(&patch_body)
.send()
.await
.context("PATCH Cloudflare DNS record")?;
if !resp.status().is_success() {
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Cloudflare PATCH DNS record failed: {text}");
}
Ok(())
}
pub fn cloudflare_credentials() -> Result<Option<(String, String)>> {
let token = fob::get_or_env("cloudflare-api-token", "CLOUDFLARE_API_TOKEN")?;
let zone_id = fob::get_or_env("cloudflare-zone-id", "CLOUDFLARE_ZONE_ID")?;
match (token, zone_id) {
(Some(t), Some(z)) => Ok(Some((t, z))),
_ => Ok(None),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn config_contains_server_url() {
let dir = std::path::PathBuf::from("/tmp/test-mesh");
let config = generate_headscale_config("https://mesh.example.com", &dir);
assert!(config.contains("server_url: https://mesh.example.com"));
assert!(config.contains("127.0.0.1:8080"));
assert!(config.contains("base_domain: mesh.internal"));
assert!(config.contains("acls.yaml"));
let cfg2 = generate_headscale_config("https://mesh.yah.dev", &dir);
assert!(!cfg2.contains("base_domain: mesh.yah\n"));
}
#[test]
fn config_all_paths_in_data_dir() {
let dir = std::path::PathBuf::from("/home/user/.yah/mesh");
let config = generate_headscale_config("https://mesh.example.com", &dir);
assert!(config.contains("/home/user/.yah/mesh/private.key"));
assert!(config.contains("/home/user/.yah/mesh/headscale.db"));
}
#[test]
fn download_url_current_platform() {
let result = headscale_download_url();
assert!(result.is_ok(), "unsupported platform: {result:?}");
let url = result.unwrap();
assert!(url.contains(HEADSCALE_VERSION));
assert!(url.starts_with("https://github.com"));
}
fn users(names: &[&str]) -> Vec<String> {
names.iter().map(|s| s.to_string()).collect()
}
#[test]
fn a_lone_user_is_chosen_without_being_named() {
assert_eq!(
choose_user("https://mesh.test", &users(&["yah"]), None).unwrap(),
"yah"
);
assert_eq!(
choose_user("https://mesh.test", &users(&["default"]), None).unwrap(),
"default"
);
}
#[test]
fn a_preferred_user_that_exists_wins() {
assert_eq!(
choose_user("https://mesh.test", &users(&["yah", "ops"]), Some("ops")).unwrap(),
"ops"
);
}
#[test]
fn a_preferred_user_that_does_not_exist_is_refused_by_name() {
let err = choose_user("https://mesh.test", &users(&["yah"]), Some("default"))
.expect_err("a missing user must not be requested from the API");
let msg = err.to_string();
assert!(msg.contains("default"), "names what was asked for: {msg}");
assert!(msg.contains("yah"), "names what actually exists: {msg}");
}
#[test]
fn several_users_and_no_preference_refuses_rather_than_guessing() {
let err = choose_user("https://mesh.test", &users(&["yah", "default"]), None)
.expect_err("an ambiguous mint target must not be guessed");
let msg = err.to_string();
assert!(msg.contains("yah") && msg.contains("default"), "{msg}");
}
#[test]
fn no_users_at_all_says_how_to_create_one() {
let err = choose_user("https://mesh.test", &[], None).expect_err("no users is fatal");
assert!(err.to_string().contains("headscale users create"));
}
}