1use anyhow::{Context, Result};
14use serde::{Deserialize, Serialize};
15
16pub struct HeadscaleClient {
27 base_url: String,
28 api_key: String,
29 http: reqwest::Client,
30}
31
32#[derive(Debug, Clone, Serialize, Deserialize)]
34pub struct PreauthKey {
35 pub key: String,
37 pub acl_tags: Vec<String>,
39}
40
41#[derive(Debug, Clone, Serialize, Deserialize)]
43pub struct NodeInfo {
44 pub id: String,
45 pub name: String,
46 pub ip_addresses: Vec<String>,
47 pub online: bool,
48}
49
50#[derive(Debug, Clone)]
52pub struct HeadscaleHealth {
53 pub reachable: bool,
54 pub status_code: u16,
55}
56
57impl HeadscaleClient {
58 pub fn new(base_url: impl Into<String>, api_key: impl Into<String>) -> Result<Self> {
60 let base_url = base_url.into().trim_end_matches('/').to_string();
61 let http = reqwest::Client::builder()
62 .timeout(std::time::Duration::from_secs(30))
63 .build()
64 .context("building Headscale HTTP client")?;
65 Ok(Self {
66 base_url,
67 api_key: api_key.into(),
68 http,
69 })
70 }
71
72 pub fn from_vault_or_env() -> Result<Option<Self>> {
78 let api_key = match fob::get_or_env("headscale-api-key", "HEADSCALE_API_KEY")? {
79 Some(k) => k,
80 None => return Ok(None),
81 };
82 let url = match fob::get_or_env("mesh-url", "HEADSCALE_URL")? {
83 Some(u) => u,
84 None => return Ok(None),
85 };
86 Ok(Some(Self::new(url, api_key)?))
87 }
88
89 pub async fn create_preauth_key(&self, user: &str, tags: &[String]) -> Result<PreauthKey> {
95 let expiration = chrono::Utc::now()
96 + chrono::TimeDelta::try_hours(1)
97 .ok_or_else(|| anyhow::anyhow!("overflow computing 1-hour expiry"))?;
98 let body = serde_json::json!({
99 "user": user,
100 "expiration": expiration.to_rfc3339(),
101 "reusable": false,
102 "ephemeral": false,
103 "aclTags": tags,
104 });
105 let resp = self
106 .http
107 .post(format!("{}/api/v1/preauthkey", self.base_url))
108 .bearer_auth(&self.api_key)
109 .json(&body)
110 .send()
111 .await
112 .context("POST /api/v1/preauthkey")?;
113 let status = resp.status();
114 if !status.is_success() {
115 let text = resp.text().await.unwrap_or_default();
116 anyhow::bail!("Headscale API {status} on POST /api/v1/preauthkey: {text}");
117 }
118 let resp_json: serde_json::Value =
119 resp.json().await.context("parsing preauthkey response")?;
120 let key = resp_json["preAuthKey"]["key"]
121 .as_str()
122 .ok_or_else(|| anyhow::anyhow!("missing preAuthKey.key in response: {resp_json}"))?;
123 Ok(PreauthKey {
124 key: key.to_string(),
125 acl_tags: tags.to_vec(),
126 })
127 }
128
129 pub async fn list_nodes(&self) -> Result<Vec<NodeInfo>> {
136 let resp = self
137 .http
138 .get(format!("{}/api/v1/node", self.base_url))
139 .bearer_auth(&self.api_key)
140 .send()
141 .await
142 .context("GET /api/v1/node")?;
143 let status = resp.status();
144 if !status.is_success() {
145 let text = resp.text().await.unwrap_or_default();
146 anyhow::bail!("Headscale API {status} on GET /api/v1/node: {text}");
147 }
148 let resp_json: serde_json::Value = resp.json().await.context("parsing node list")?;
149 let nodes = resp_json["nodes"].as_array().cloned().unwrap_or_default();
150 Ok(nodes
151 .iter()
152 .filter_map(|m| {
153 Some(NodeInfo {
154 id: m["id"].as_str()?.to_string(),
155 name: m["name"].as_str()?.to_string(),
156 ip_addresses: m["ipAddresses"]
157 .as_array()
158 .map(|arr| {
159 arr.iter()
160 .filter_map(|v| v.as_str().map(str::to_string))
161 .collect()
162 })
163 .unwrap_or_default(),
164 online: m["online"].as_bool().unwrap_or(false),
165 })
166 })
167 .collect())
168 }
169
170 pub async fn health(&self) -> HeadscaleHealth {
172 match self
173 .http
174 .get(format!("{}/health", self.base_url))
175 .send()
176 .await
177 {
178 Ok(r) => HeadscaleHealth {
179 reachable: r.status().is_success(),
180 status_code: r.status().as_u16(),
181 },
182 Err(_) => HeadscaleHealth {
183 reachable: false,
184 status_code: 0,
185 },
186 }
187 }
188}
189
190pub fn generate_headscale_config(server_url: &str, data_dir: &std::path::Path) -> String {
202 let private_key = data_dir.join("private.key").display().to_string();
203 let noise_key = data_dir.join("noise_private.key").display().to_string();
204 let db_path = data_dir.join("headscale.db").display().to_string();
205 let socket_path = data_dir.join("headscale.sock").display().to_string();
206 let acl_path = data_dir.join("acls.yaml").display().to_string();
207
208 format!(
211 r#"---
212server_url: {server_url}
213listen_addr: 127.0.0.1:8080
214grpc_listen_addr: 127.0.0.1:50443
215metrics_listen_addr: 127.0.0.1:9090
216private_key_path: {private_key}
217noise:
218 private_key_path: {noise_key}
219database:
220 type: sqlite
221 sqlite:
222 path: {db_path}
223unix_socket: {socket_path}
224unix_socket_permission: "0770"
225dns:
226 magic_dns: true
227 base_domain: mesh.internal
228 nameservers:
229 global:
230 - 1.1.1.1
231 - 8.8.8.8
232log:
233 level: info
234prefixes:
235 v4: 100.64.0.0/10
236 v6: fd7a:115c:a1e0::/48
237 allocation: sequential
238policy:
239 mode: file
240 path: {acl_path}
241derp:
242 server:
243 enabled: false
244 urls:
245 - https://controlplane.tailscale.com/derpmap/default
246 auto_update_enabled: false
247 update_frequency: 24h
248"#
249 )
250}
251
252pub const DEFAULT_ACL_POLICY: &str = r#"{
259 "acls": [
260 { "action": "accept", "src": ["*"], "dst": ["*:*"] }
261 ]
262}
263"#;
264
265pub const HEADSCALE_VERSION: &str = "0.23.0";
271
272pub fn headscale_download_url() -> Result<String> {
277 let os = match std::env::consts::OS {
278 "macos" => "darwin",
279 "linux" => "linux",
280 other => anyhow::bail!(
281 "unsupported OS '{other}' for `yah mesh start`; \
282 install headscale manually from https://headscale.net"
283 ),
284 };
285 let arch = match std::env::consts::ARCH {
286 "x86_64" => "amd64",
287 "aarch64" => "arm64",
288 other => anyhow::bail!(
289 "unsupported architecture '{other}' for `yah mesh start`; \
290 install headscale manually from https://headscale.net"
291 ),
292 };
293 Ok(format!(
294 "https://github.com/juanfont/headscale/releases/download/v{HEADSCALE_VERSION}/headscale_{HEADSCALE_VERSION}_{os}_{arch}"
295 ))
296}
297
298pub async fn update_cloudflare_dns(
310 api_token: &str,
311 zone_id: &str,
312 record_name: &str,
313 new_ip: &str,
314) -> Result<()> {
315 let http = reqwest::Client::builder()
316 .timeout(std::time::Duration::from_secs(30))
317 .build()
318 .context("building Cloudflare HTTP client")?;
319
320 let list_url = format!("https://api.cloudflare.com/client/v4/zones/{zone_id}/dns_records");
321
322 let resp = http
323 .get(&list_url)
324 .bearer_auth(api_token)
325 .query(&[("name", record_name), ("type", "A")])
326 .send()
327 .await
328 .context("GET Cloudflare DNS records")?;
329
330 if !resp.status().is_success() {
331 let text = resp.text().await.unwrap_or_default();
332 anyhow::bail!("Cloudflare list DNS records failed: {text}");
333 }
334
335 let body: serde_json::Value = resp
336 .json()
337 .await
338 .context("parsing Cloudflare records list")?;
339 let records = body["result"]
340 .as_array()
341 .ok_or_else(|| anyhow::anyhow!("unexpected Cloudflare response shape: missing 'result'"))?;
342
343 let record_id = match records.len() {
344 0 => {
351 println!(" A record '{record_name}' absent — creating it (→ {new_ip}) ...");
352 let create_body = serde_json::json!({
353 "type": "A",
354 "name": record_name,
355 "content": new_ip,
356 "ttl": 120,
357 "proxied": false
358 });
359 let resp = http
360 .post(&list_url)
361 .bearer_auth(api_token)
362 .json(&create_body)
363 .send()
364 .await
365 .context("POST Cloudflare DNS record (create-if-missing)")?;
366 if !resp.status().is_success() {
367 let text = resp.text().await.unwrap_or_default();
368 anyhow::bail!("Cloudflare create DNS record failed: {text}");
369 }
370 println!(" DNS record created.");
371 return Ok(());
372 }
373 1 => records[0]["id"]
374 .as_str()
375 .ok_or_else(|| anyhow::anyhow!("missing 'id' in Cloudflare record"))?
376 .to_string(),
377 n => anyhow::bail!(
378 "{n} A records named '{record_name}' found — expected exactly one; \
379 resolve the ambiguity in your Cloudflare dashboard"
380 ),
381 };
382
383 let patch_url = format!("{list_url}/{record_id}");
384 let patch_body = serde_json::json!({
385 "type": "A",
386 "name": record_name,
387 "content": new_ip,
388 "ttl": 120,
389 "proxied": false
390 });
391
392 let resp = http
393 .patch(&patch_url)
394 .bearer_auth(api_token)
395 .json(&patch_body)
396 .send()
397 .await
398 .context("PATCH Cloudflare DNS record")?;
399
400 if !resp.status().is_success() {
401 let text = resp.text().await.unwrap_or_default();
402 anyhow::bail!("Cloudflare PATCH DNS record failed: {text}");
403 }
404
405 Ok(())
406}
407
408pub fn cloudflare_credentials() -> Result<Option<(String, String)>> {
417 let token = fob::get_or_env("cloudflare-api-token", "CLOUDFLARE_API_TOKEN")?;
418 let zone_id = fob::get_or_env("cloudflare-zone-id", "CLOUDFLARE_ZONE_ID")?;
419 match (token, zone_id) {
420 (Some(t), Some(z)) => Ok(Some((t, z))),
421 _ => Ok(None),
422 }
423}
424
425#[cfg(test)]
430mod tests {
431 use super::*;
432
433 #[test]
434 fn config_contains_server_url() {
435 let dir = std::path::PathBuf::from("/tmp/test-mesh");
436 let config = generate_headscale_config("https://mesh.example.com", &dir);
437 assert!(config.contains("server_url: https://mesh.example.com"));
438 assert!(config.contains("127.0.0.1:8080"));
439 assert!(config.contains("base_domain: mesh.internal"));
440 assert!(config.contains("acls.yaml"));
441 let cfg2 = generate_headscale_config("https://mesh.yah.dev", &dir);
446 assert!(!cfg2.contains("base_domain: mesh.yah\n"));
447 }
448
449 #[test]
450 fn config_all_paths_in_data_dir() {
451 let dir = std::path::PathBuf::from("/home/user/.yah/mesh");
452 let config = generate_headscale_config("https://mesh.example.com", &dir);
453 assert!(config.contains("/home/user/.yah/mesh/private.key"));
454 assert!(config.contains("/home/user/.yah/mesh/headscale.db"));
455 }
456
457 #[test]
458 fn download_url_current_platform() {
459 let result = headscale_download_url();
461 assert!(result.is_ok(), "unsupported platform: {result:?}");
462 let url = result.unwrap();
463 assert!(url.contains(HEADSCALE_VERSION));
464 assert!(url.starts_with("https://github.com"));
465 }
466}