Skip to main content

adguardian/
welcome.rs

1use base64::{engine::general_purpose::STANDARD, Engine as _};
2use colored::*;
3use crossterm::{
4  event::{self, Event, KeyCode, KeyEventKind, KeyModifiers},
5  terminal::{disable_raw_mode, enable_raw_mode},
6};
7use reqwest::{header::LOCATION, redirect::Policy, Client, Error, StatusCode};
8use std::{
9  cmp::Ordering,
10  env,
11  fmt::Display,
12  io::{self, IsTerminal, Write},
13  time::Duration,
14};
15
16use semver::Version;
17use serde::Deserialize;
18use serde_json::Value;
19
20/// Reusable function that just prints success messages to the console
21fn print_info(text: &str, is_secondary: bool) {
22  if is_secondary {
23    println!("{}", text.green().italic().dimmed());
24  } else {
25    println!("{}", text.green());
26  };
27}
28
29/// Prints the AdGuardian ASCII art to console
30fn print_ascii_art() {
31  let art = r"
32 █████╗ ██████╗  ██████╗ ██╗   ██╗ █████╗ ██████╗ ██████╗ ██╗ █████╗ ███╗   ██╗
33██╔══██╗██╔══██╗██╔════╝ ██║   ██║██╔══██╗██╔══██╗██╔══██╗██║██╔══██╗████╗  ██║
34███████║██║  ██║██║  ███╗██║   ██║███████║██████╔╝██║  ██║██║███████║██╔██╗ ██║
35██╔══██║██║  ██║██║   ██║██║   ██║██╔══██║██╔══██╗██║  ██║██║██╔══██║██║╚██╗██║
36██║  ██║██████╔╝╚██████╔╝╚██████╔╝██║  ██║██║  ██║██████╔╝██║██║  ██║██║ ╚████║
37╚═╝  ╚═╝╚═════╝  ╚═════╝  ╚═════╝ ╚═╝  ╚═╝╚═╝  ╚═╝╚═════╝ ╚═╝╚═╝  ╚═╝╚═╝  ╚═══╝
38";
39  print_info(art, false);
40  print_info("\nWelcome to AdGuardian Terminal Edition!", false);
41  print_info(
42    "Terminal-based, real-time traffic monitoring and statistics for your AdGuard Home instance",
43    true,
44  );
45  print_info(
46    "For documentation and support, please visit: https://github.com/lissy93/adguardian-term",
47    true,
48  );
49}
50
51/// Print error message, along with (optional) stack trace, then exit
52fn print_error(message: &str, sub_message: &str, error: Option<&Error>) -> ! {
53  eprintln!(
54    "{}{}{}",
55    message.red(),
56    match error {
57      Some(err) => format!("\n{}", err).red().dimmed(),
58      None => "".red().dimmed(),
59    },
60    format!("\n{}", sub_message).yellow(),
61  );
62
63  std::process::exit(1);
64}
65
66/// Given a key, get the value from the environmental variables, and print it to the console
67fn get_env(key: &str) -> Result<String, env::VarError> {
68  env::var(key).inspect(|v| {
69    println!(
70      "{}",
71      format!(
72        "{} is set to {}",
73        key.bold(),
74        if key.contains("PASSWORD") {
75          "******"
76        } else {
77          v
78        }
79      )
80      .green()
81    );
82  })
83}
84
85/// Given a possibly undefined version number, check if it's present and supported
86fn check_version(version: Option<&str>) {
87  let min_version = Version::parse("0.107.29").unwrap();
88
89  match version {
90    Some(version_str) => {
91      match Version::parse(version_str.strip_prefix('v').unwrap_or(version_str)) {
92        Ok(adguard_version) if adguard_version < min_version => print_error(
93          "AdGuard Home version is too old, and is now unsupported",
94          format!(
95            "You're running AdGuard {}. Please upgrade to v{} or later.",
96            version_str, min_version
97          )
98          .as_str(),
99          None,
100        ),
101        Ok(_) => {}
102        Err(_) => print_error(
103          "Unsupported AdGuard Home version",
104          "Couldn't parse the version number reported by your AdGuard Home instance.",
105          None,
106        ),
107      }
108    }
109    None => {
110      print_error(
111        "Unsupported AdGuard Home version",
112        format!(
113          concat!(
114            "Failed to get the version number of your AdGuard Home instance.\n",
115            "This usually means you're running an old, and unsupported version.\n",
116            "Please upgrade to v{} or later."
117          ),
118          min_version
119        )
120        .as_str(),
121        None,
122      );
123    }
124  }
125}
126
127/// Run an async operation, retrying on error up to `attempts` times, `delay` apart.
128/// Each failure is reported; the last error is returned once attempts are exhausted.
129pub async fn with_retries<T, E, F, Fut>(
130  attempts: u32,
131  delay: Duration,
132  label: &str,
133  mut operation: F,
134) -> Result<T, E>
135where
136  F: FnMut() -> Fut,
137  Fut: std::future::Future<Output = Result<T, E>>,
138  E: Display,
139{
140  let mut attempt = 1;
141  loop {
142    match operation().await {
143      Ok(value) => return Ok(value),
144      Err(e) if attempt < attempts => {
145        println!(
146          "{}",
147          format!(
148            "{} failed (attempt {}/{}): {}\nRetrying in {}s...",
149            label,
150            attempt,
151            attempts,
152            e,
153            delay.as_secs()
154          )
155          .yellow()
156        );
157        tokio::time::sleep(delay).await;
158        attempt += 1;
159      }
160      Err(e) => return Err(e),
161    }
162  }
163}
164
165/// Check if AdGuard Home runs in GL.iNet mode, which redirects its web UI to the router's login
166async fn is_glinet_mode(ip: &str, port: &str, protocol: &str) -> bool {
167  let Ok(client) = Client::builder().redirect(Policy::none()).build() else {
168    return false;
169  };
170  let url = format!("{}://{}:{}/", protocol, ip, port);
171  let router_url = format!("http://{}", ip);
172  client
173    .get(&url)
174    .timeout(Duration::from_secs(2))
175    .send()
176    .await
177    .is_ok_and(|res| {
178      res
179        .headers()
180        .get(LOCATION)
181        .is_some_and(|l| l == &router_url)
182    })
183}
184
185/// With the users specified AdGuard details, verify the connection.
186/// Returns `Err` on a failed connection (so the caller can retry); exits on
187/// rejected auth or an unsupported version, which retrying wouldn't fix.
188async fn verify_connection(
189  client: &Client,
190  ip: &str,
191  port: &str,
192  protocol: &str,
193  username: &str,
194  password: &str,
195) -> Result<(), Box<dyn std::error::Error>> {
196  println!(
197    "{}",
198    "\nVerifying connection to your AdGuard instance...".blue()
199  );
200
201  let auth_string = format!("{}:{}", username, password);
202  let auth_header_value = format!("Basic {}", STANDARD.encode(&auth_string));
203  let mut headers = reqwest::header::HeaderMap::new();
204  headers.insert("Authorization", auth_header_value.parse()?);
205
206  let url = format!("{}://{}:{}/control/status", protocol, ip, port);
207
208  match client
209    .get(&url)
210    .headers(headers)
211    .timeout(Duration::from_secs(2))
212    .send()
213    .await
214  {
215    Ok(res) if res.status().is_success() => {
216      // Get version string (if present), and check if valid - exit if not
217      let body: Value = res.json().await?;
218      check_version(body["version"].as_str());
219      // All good! Print success message :)
220      let safe_version = body["version"].as_str().unwrap_or("mystery version");
221      println!(
222        "{}",
223        format!("AdGuard ({}) connection successful!\n", safe_version).green()
224      );
225      Ok(())
226    }
227    // Connection failed to authenticate. Print error and exit
228    Ok(res) => print_error(
229      &format!("Authentication with AdGuard at {}:{} failed", ip, port),
230      if res.status() == StatusCode::UNAUTHORIZED && is_glinet_mode(ip, port, protocol).await {
231        "AdGuard Home is running in GL.iNet mode (--glinet), which doesn't accept a username and password."
232      } else {
233        "Check the credentials you passed as environmental variables and try again."
234      },
235      None,
236    ),
237    // Connection failed to establish - return so the caller can retry
238    Err(e) => Err(e.into()),
239  }
240}
241
242#[derive(Deserialize)]
243struct CratesIoResponse {
244  #[serde(rename = "crate")]
245  krate: Crate,
246}
247
248#[derive(Deserialize)]
249struct Crate {
250  max_version: String,
251}
252
253/// Gets the latest version of the crate from crates.io
254async fn get_latest_version(crate_name: &str) -> Result<String, Box<dyn std::error::Error>> {
255  let url = format!("https://crates.io/api/v1/crates/{}", crate_name);
256  let client = reqwest::Client::new();
257  let res = client
258    .get(&url)
259    .header(
260      reqwest::header::USER_AGENT,
261      "version_check (adguardian.as93.net)",
262    )
263    .timeout(Duration::from_secs(2))
264    .send()
265    .await?;
266
267  if res.status().is_success() {
268    let response: CratesIoResponse = res.json().await?;
269    Ok(response.krate.max_version)
270  } else {
271    let status = res.status();
272    let body = res.text().await?;
273    Err(format!("Request failed with status {}: body: {}", status, body).into())
274  }
275}
276
277/// Checks for updates to the crate, and prints a message if an update is available
278async fn check_for_updates() {
279  // Get crate name and version from Cargo.toml
280  let crate_name = env!("CARGO_PKG_NAME");
281  let crate_version = env!("CARGO_PKG_VERSION");
282  println!("{}", "\nChecking for updates...".blue());
283  // Parse the current version, and fetch and parse the latest version
284  let zero = Version::new(0, 0, 0);
285  let current_version = Version::parse(crate_version).unwrap_or_else(|_| zero.clone());
286  let latest_version = Version::parse(
287    &get_latest_version(crate_name)
288      .await
289      .unwrap_or_else(|_| "0.0.0".to_string()),
290  )
291  .unwrap_or_else(|_| zero.clone());
292
293  // Compare the current and latest versions, and print the appropriate message
294  if current_version == zero || latest_version == zero {
295    println!("{}", "Unable to check for updates".yellow());
296    return;
297  }
298  match current_version.cmp(&latest_version) {
299    Ordering::Less => println!(
300      "{}",
301      format!(
302        "A new version of AdGuardian is available.\nUpdate from {} to {} for the best experience",
303        current_version.to_string().bold(),
304        latest_version.to_string().bold()
305      )
306      .yellow()
307    ),
308    Ordering::Equal => println!(
309      "{}",
310      format!(
311        "AdGuardian is up-to-date, running version {}",
312        current_version.to_string().bold()
313      )
314      .green()
315    ),
316    Ordering::Greater => println!(
317      "{}",
318      format!(
319        "Running a pre-released edition of AdGuardian, version {}",
320        current_version.to_string().bold()
321      )
322      .green()
323    ),
324  }
325}
326
327/// The value to pre-fill for a field's interactive prompt, where a sensible one exists
328fn default_for(key: &str) -> Option<&'static str> {
329  match key {
330    "ADGUARD_IP" => Some("127.0.0.1"),
331    "ADGUARD_PORT" => Some("3000"),
332    _ => None,
333  }
334}
335
336/// Read a line from the terminal in raw mode, echoing nothing. Ctrl-C cancels.
337fn read_masked() -> io::Result<String> {
338  enable_raw_mode()?;
339  let result = masked_loop();
340  let _ = disable_raw_mode();
341  if result.is_ok() {
342    println!();
343  }
344  result
345}
346
347fn masked_loop() -> io::Result<String> {
348  let mut value = String::new();
349  loop {
350    if let Event::Key(key) = event::read()? {
351      // Windows also reports key releases, which would double each keypress
352      if key.kind == KeyEventKind::Release {
353        continue;
354      }
355      let ctrl = key.modifiers.contains(KeyModifiers::CONTROL);
356      match key.code {
357        KeyCode::Enter => return Ok(value),
358        KeyCode::Char('c') if ctrl => return Err(io::ErrorKind::Interrupted.into()),
359        KeyCode::Char(c) if !ctrl => value.push(c),
360        KeyCode::Backspace => {
361          value.pop();
362        }
363        _ => {}
364      }
365    }
366  }
367}
368
369/// Print the prompt and read a value, masking secret fields on an interactive terminal
370fn read_field(prompt: &ColoredString, secret: bool) -> io::Result<String> {
371  print!("{}", prompt);
372  io::stdout().flush()?;
373  if secret && io::stdin().is_terminal() {
374    read_masked()
375  } else {
376    let mut value = String::new();
377    io::stdin().read_line(&mut value)?;
378    Ok(value)
379  }
380}
381
382/// Read a field off the async runtime threads
383async fn read_input(prompt: ColoredString, secret: bool) -> io::Result<String> {
384  tokio::task::spawn_blocking(move || read_field(&prompt, secret))
385    .await
386    .expect("input task panicked")
387}
388
389/// Print the cancellation notice and exit cleanly
390fn exit_interrupted() -> ! {
391  println!(
392    "{}",
393    "\n\nAdGuardian setup interrupted by user, exiting...".yellow()
394  );
395  std::process::exit(0);
396}
397
398/// Prompt for a single field, re-prompting until the input is valid.
399/// Masks passwords, applies the field's default on empty input, validates the
400/// port is numeric, and exits cleanly if the user interrupts with Ctrl-C.
401async fn prompt_for(key: &str) -> Result<String, Box<dyn std::error::Error>> {
402  let default = default_for(key);
403  let secret = key.contains("PASSWORD");
404  loop {
405    let hint = default.map(|d| format!(" [{}]", d)).unwrap_or_default();
406    let prompt = format!("› Enter a value for {}{}: ", key, hint)
407      .blue()
408      .bold();
409
410    let input = tokio::select! {
411      res = read_input(prompt, secret) => match res {
412        Ok(value) => value,
413        Err(e) if e.kind() == io::ErrorKind::Interrupted => exit_interrupted(),
414        Err(e) => return Err(e.into()),
415      },
416      _ = tokio::signal::ctrl_c() => exit_interrupted(),
417    };
418
419    let value = match input.trim() {
420      "" => default.unwrap_or_default(),
421      trimmed => trimmed,
422    };
423
424    if key == "ADGUARD_PORT" && value.parse::<u16>().is_err() {
425      println!("{}", "Port must be a number, and a valid port".yellow());
426      continue;
427    }
428    return Ok(value.to_string());
429  }
430}
431
432/// Initiate the welcome script
433/// This function will:
434/// - Print the AdGuardian ASCII art
435/// - Check if there's an update available
436/// - Check for the required environmental variables
437/// - Prompt the user to enter any missing variables
438/// - Verify the connection to the AdGuard instance
439/// - Verify authentication is successful
440/// - Verify the AdGuard Home version is supported
441/// - Then either print a success message, or show instructions to fix and exit
442pub async fn welcome() -> Result<(), Box<dyn std::error::Error>> {
443  print_ascii_art();
444
445  // Check for updates
446  check_for_updates().await;
447
448  println!("{}", "\nStarting initialization checks...".blue());
449
450  let client = Client::new();
451
452  // List of available flags, ant their associated env vars
453  let flags = [
454    ("--adguard-ip", "ADGUARD_IP"),
455    ("--adguard-port", "ADGUARD_PORT"),
456    ("--adguard-username", "ADGUARD_USERNAME"),
457    ("--adguard-password", "ADGUARD_PASSWORD"),
458  ];
459
460  let protocol: String = env::var("ADGUARD_PROTOCOL")
461    .unwrap_or_else(|_| "http".into())
462    .parse()?;
463  env::set_var("ADGUARD_PROTOCOL", protocol);
464
465  // Parse command line arguments
466  let mut args = std::env::args().peekable();
467  while let Some(arg) = args.next() {
468    for &(flag, var) in &flags {
469      if arg == flag {
470        if let Some(value) = args.peek().filter(|v| !v.starts_with("--")) {
471          env::set_var(var, value);
472          args.next();
473        }
474      }
475    }
476  }
477
478  // If any of the env variables or flags are not yet set, prompt the user to enter them
479  for &key in &[
480    "ADGUARD_IP",
481    "ADGUARD_PORT",
482    "ADGUARD_USERNAME",
483    "ADGUARD_PASSWORD",
484  ] {
485    if env::var(key).is_err() {
486      println!(
487        "{}",
488        format!("The {} environmental variable is not yet set", key.bold()).yellow()
489      );
490      env::set_var(key, prompt_for(key).await?);
491    }
492  }
493
494  // Grab the values of the (now set) environmental variables
495  let ip = get_env("ADGUARD_IP")?;
496  let port = get_env("ADGUARD_PORT")?;
497  let protocol = get_env("ADGUARD_PROTOCOL")?;
498  let username = get_env("ADGUARD_USERNAME")?;
499  let password = get_env("ADGUARD_PASSWORD")?;
500
501  // Verify we can connect, authenticate, and that the version is supported
502  let connected = with_retries(3, Duration::from_secs(5), "AdGuard connection", || {
503    verify_connection(&client, &ip, &port, &protocol, &username, &password)
504  })
505  .await;
506
507  if connected.is_err() {
508    print_error(
509      &format!(
510        "Could not reach AdGuard at {}:{} after 3 attempts",
511        ip, port
512      ),
513      "Please check that AdGuard Home is running and your settings are correct.",
514      None,
515    );
516  }
517
518  Ok(())
519}