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
20fn 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
29fn 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
51fn 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
66fn 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
85fn 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
127pub 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
165async 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
185async 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 let body: Value = res.json().await?;
218 check_version(body["version"].as_str());
219 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 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 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
253async 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
277async fn check_for_updates() {
279 let crate_name = env!("CARGO_PKG_NAME");
281 let crate_version = env!("CARGO_PKG_VERSION");
282 println!("{}", "\nChecking for updates...".blue());
283 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 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
327fn 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
336fn 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 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
369fn 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
382async 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
389fn exit_interrupted() -> ! {
391 println!(
392 "{}",
393 "\n\nAdGuardian setup interrupted by user, exiting...".yellow()
394 );
395 std::process::exit(0);
396}
397
398async 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
432pub async fn welcome() -> Result<(), Box<dyn std::error::Error>> {
443 print_ascii_art();
444
445 check_for_updates().await;
447
448 println!("{}", "\nStarting initialization checks...".blue());
449
450 let client = Client::new();
451
452 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 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 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 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 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}