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 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,
},
}
}
}
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"));
}
}