Skip to main content

ohcldiag_dev/
main.rs

1// Copyright (c) Microsoft Corporation.
2// Licensed under the MIT License.
3
4//! Host-side development CLI for a running OpenHCL diagnostics server.
5//!
6//! The tool connects to an OpenHCL/Underhill instance and provides process
7//! execution, interactive shells, inspect queries and updates, kernel logs,
8//! core dumps, debugger relays, saved-state dumps, packet capture, and other
9//! diagnostics. The available transport identifies the target VM or local
10//! diagnostic endpoint.
11//!
12//! `ohcldiag-dev` is intended for interactive investigation on development
13//! builds. It deliberately provides no long-term guarantees for its command
14//! syntax, output, or inspect-tree paths and must not be treated as a stable
15//! automation API.
16
17#![expect(missing_docs)]
18#![forbid(unsafe_code)]
19
20mod completions;
21
22use anyhow::Context;
23use clap::ArgGroup;
24use clap::Args;
25use clap::Parser;
26use clap::Subcommand;
27use diag_client::DiagClient;
28use diag_client::PacketCaptureOperation;
29use futures::StreamExt;
30use futures::io::AllowStdIo;
31use futures_concurrency::future::Race;
32use pal_async::DefaultPool;
33use pal_async::driver::Driver;
34use pal_async::socket::PolledSocket;
35use pal_async::task::Spawn;
36use pal_async::timer::PolledTimer;
37use std::ffi::OsStr;
38use std::io::ErrorKind;
39use std::io::IsTerminal;
40use std::io::Write;
41use std::net::TcpListener;
42use std::path::Path;
43use std::path::PathBuf;
44use std::str::FromStr;
45use std::sync::Arc;
46use std::time::Duration;
47use thiserror::Error;
48use tracing_subscriber::layer::SubscriberExt;
49use tracing_subscriber::util::SubscriberInitExt;
50use unicycle::FuturesUnordered;
51
52#[derive(Parser)]
53#[clap(about = "(dev) CLI to interact with the Underhill diagnostics server")]
54#[clap(long_about = r#"
55CLI to interact with the Underhill diagnostics server.
56
57DISCLAIMER:
58    `ohcldiag-dev` does not make ANY stability guarantees regarding the layout of
59    the CLI, the syntax that is emitted via stdout/stderr, the location of nodes
60    in the `inspect` graph, etc...
61
62        !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!
63        !! ANY AUTOMATION THAT USES ohcldiag-dev WILL EVENTUALLY BREAK !!
64        !!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!!
65"#)]
66struct Options {
67    #[clap(flatten)]
68    vm: VmArg,
69
70    #[clap(subcommand)]
71    command: Command,
72}
73
74#[derive(Subcommand)]
75enum Command {
76    #[clap(hide = true)]
77    Complete(clap_dyn_complete::Complete),
78    Completions(completions::Completions),
79    /// Starts an interactive terminal in VTL2.
80    Shell {
81        /// The shell process to start.
82        #[clap(default_value = "/bin/sh")]
83        shell: String,
84        /// The arguments to pass to the shell process.
85        args: Vec<String>,
86    },
87    /// Runs a process in VTL2.
88    Run {
89        /// The command to run.
90        command: String,
91        /// Arguments to pass to the command.
92        args: Vec<String>,
93    },
94    /// Inspects the Underhill state.
95    #[clap(visible_alias = "i")]
96    Inspect {
97        /// Recursively enumerate child nodes.
98        #[clap(short)]
99        recursive: bool,
100        /// Limit the recursive inspection depth.
101        #[clap(short, long, requires("recursive"))]
102        limit: Option<usize>,
103        /// Output in JSON format.
104        #[clap(short, long)]
105        json: bool,
106        /// Poll periodically.
107        #[clap(short)]
108        poll: bool,
109        /// The poll period in seconds.
110        #[clap(long, default_value = "1", requires("poll"))]
111        period: f64,
112        /// The count of polls
113        #[clap(long, requires("poll"))]
114        count: Option<usize>,
115        /// The path to inspect.
116        path: Option<String>,
117        /// Update the path with a new value.
118        #[clap(short, long, conflicts_with("recursive"))]
119        update: Option<String>,
120        /// Timeout to wait for the inspection. 0 means no timeout.
121        #[clap(short, default_value = "1", conflicts_with("update"))]
122        timeout: u64,
123    },
124    /// Updates an inspectable value.
125    #[clap(hide = true)]
126    Update {
127        /// The path.
128        path: String,
129        /// The new value.
130        value: String,
131    },
132    /// Starts the VM if it's waiting for the signal to start.
133    ///
134    /// Underhill must have been started with --wait-for-start or
135    /// OPENHCL_WAIT_FOR_START set.
136    Start {
137        /// Environment variables to set, in the form X=Y
138        #[clap(short, long)]
139        env: Vec<EnvString>,
140        /// Environment variables to clear
141        #[clap(short, long)]
142        unset: Vec<String>,
143        /// Extra command line arguments to append.
144        args: Vec<String>,
145    },
146    /// Writes the contents of the kernel message buffer, /dev/kmsg.
147    Kmsg {
148        /// Keep waiting for and writing new data as its logged.
149        #[clap(short, long)]
150        follow: bool,
151        /// Reconnect (retrying indefinitely) whenever the connection is lost.
152        #[clap(short, long)]
153        reconnect: bool,
154        /// Write verbose information about the connection state.
155        #[clap(short, long)]
156        verbose: bool,
157        /// Read kmsg from the VM's serial port.
158        ///
159        /// This only works on Hyper-V.
160        #[cfg(windows)]
161        #[clap(long, conflicts_with = "reconnect")]
162        serial: bool,
163        /// Pipe to read from for the serial port (or any other pipe)
164        ///
165        /// This only works on Hyper-V.
166        #[cfg(windows)]
167        #[clap(long, requires = "serial")]
168        pipe_path: Option<String>,
169    },
170    /// Writes the contents of the file.
171    File {
172        /// Keep waiting for and writing new data as its logged.
173        #[clap(short, long)]
174        follow: bool,
175        #[clap(short('p'), long)]
176        file_path: String,
177    },
178    /// Starts GDB server on stdio.
179    ///
180    /// Use this with gdb's target command:
181    ///
182    ///     target remote |ohcldiag-dev.exe gdbserver my-vm
183    ///
184    /// Or for multi-process debugging:
185    ///
186    ///     target extended-remote |ohcldiag-dev.exe gdbserver --multi my-vm
187    Gdbserver {
188        /// The pid to attach to. Defaults to Underhill's.
189        #[clap(long)]
190        pid: Option<i32>,
191        /// Use multi-process debugging, for use with gdb's extended-remote.
192        #[clap(long, conflicts_with("pid"))]
193        multi: bool,
194    },
195    /// Starts the GDB stub for debugging the guest on stdio.
196    ///
197    /// Use this with gdb's target command:
198    ///
199    ///     target remote |ohcldiag-dev.exe gdbstub my-vm
200    ///
201    Gdbstub {
202        /// The vsock prot to connect to.
203        #[clap(short, long, default_value = "4")]
204        port: u32,
205    },
206    /// Crashes the VM.
207    ///
208    /// Must specify the VM name, as well as the crash type.
209    #[clap(group(
210        ArgGroup::new("process")
211            .required(true)
212            .args(&["pid", "name"]),
213    ))]
214    Crash {
215        /// Type of crash.
216        ///
217        /// Current crash types supported: "panic"
218        crash_type: CrashType,
219        /// PID of underhill process to crash
220        #[clap(short, long)]
221        pid: Option<i32>,
222        /// Name of underhill process to crash
223        #[clap(short, long)]
224        name: Option<String>,
225    },
226    /// Streams the ELF core dump file of a process to the host.
227    ///
228    /// Streams the core dump file of a process to the host where the file
229    /// is saved as `dst`.
230    #[clap(group(
231        ArgGroup::new("process")
232        .required(true)
233        .args(&["pid", "name"]),
234    ))]
235    CoreDump {
236        /// Enable verbose output.
237        #[clap(short, long)]
238        verbose: bool,
239        /// PID of process to dump
240        #[clap(short, long)]
241        pid: Option<i32>,
242        /// Name of underhill process to dump
243        #[clap(short, long)]
244        name: Option<String>,
245        /// Destination file path. If omitted, the data is written to the standard
246        /// output unless it is a terminal. In that case, an error is returned.
247        dst: Option<PathBuf>,
248    },
249    /// Restarts the Underhill worker process, keeping VTL0 running.
250    Restart,
251    /// Get the current contents of the performance trace buffer, for use with
252    /// <https://ui.perfetto.dev>.
253    PerfTrace {
254        /// The output file. Defaults to stdout.
255        #[clap(short)]
256        output: Option<PathBuf>,
257    },
258    /// Sets up a relay between a virtual socket and a TCP client on the host.
259    VsockTcpRelay {
260        vsock_port: u32,
261        tcp_port: u16,
262        #[clap(long)]
263        allow_remote: bool,
264        /// Reconnect (retrying indefinitely) whenever either side of the
265        /// connection is lost.
266        ///
267        /// NOTE: Today, this does not handle the case where the vsock side is
268        /// not ready to connect. That will cause the relay to terminate.
269        #[clap(short, long)]
270        reconnect: bool,
271    },
272    /// Pause the VM (including all devices)
273    Pause,
274    /// Resume the VM
275    Resume,
276    /// Dumps the VM's VTL2 state without servicing or tearing down Underhill.
277    DumpSavedState {
278        /// The output file. Defaults to stdout.
279        #[clap(short)]
280        output: Option<PathBuf>,
281    },
282    /// Starts a network packet capture trace.
283    PacketCapture {
284        /// Destination file path. nic index is appended to the file name.
285        #[clap(short('w'), default_value = "nic")]
286        output: PathBuf,
287        /// Number of seconds for which to capture packets.
288        #[clap(short('G'), long, default_value = "60", value_parser = |arg: &str| -> Result<Duration, std::num::ParseIntError> {Ok(Duration::from_secs(arg.parse()?))})]
289        seconds: Duration,
290        /// Length of the packet to capture.
291        #[clap(short('s'), long, default_value = "65535", value_parser = clap::value_parser!(u16).range(1..))]
292        snaplen: u16,
293    },
294    /// Memory usage profile tracing.
295    MemoryProfileTrace {
296        /// PID of process to collect the trace for
297        #[clap(short, long)]
298        pid: Option<i32>,
299        /// Name of underhill process to dump
300        #[clap(short, long)]
301        name: Option<String>,
302        /// The output file. Defaults to stdout.
303        #[clap(short)]
304        output: Option<PathBuf>,
305    },
306    /// Processes EFI diagnostics from guest memory and outputs the logs.
307    ///
308    /// The log level filter controls which UEFI log entries are emitted.
309    /// The buffer already contains all log levels; this filter selects
310    /// which ones to display.
311    EfiDiagnostics {
312        /// The log level filter to apply.
313        ///
314        /// Accepted values: "default" (errors+warnings), "info" (errors+warnings+info),
315        /// "full" (all levels).
316        log_level: EfiDiagnosticsLogLevel,
317        /// The output destination.
318        ///
319        /// Accepted values: "stdout", "tracing".
320        output: EfiDiagnosticsOutput,
321    },
322}
323
324#[derive(Debug, Clone, Args)]
325pub struct VmArg {
326    #[doc = r#"VM identifier.
327
328    This can be one of:
329
330    * vsock:PATH - A path to a hybrid vsock Unix socket for a VM, as used by OpenVMM
331
332    * unix:PATH - A path to a Unix socket for connecting to the control plane
333
334    "#]
335    #[cfg_attr(
336        windows,
337        doc = "* hyperv:NAME - A Hyper-V VM name
338
339    "
340    )]
341    #[cfg_attr(
342        windows,
343        doc = "* hyperv-id:GUID - A Hyper-V VM ID (with or without braces)
344
345    "
346    )]
347    #[cfg_attr(
348        windows,
349        doc = "* NAME_OR_PATH - Either a Hyper-V VM name, or a path as in vsock:PATH"
350    )]
351    #[cfg_attr(not(windows), doc = "* PATH - A path as in vsock:PATH")]
352    #[clap(name = "VM")]
353    id: VmId,
354}
355
356#[derive(Debug, Clone)]
357enum VmId {
358    #[cfg(windows)]
359    HyperV(String),
360    #[cfg(windows)]
361    HyperVId(guid::Guid),
362    HybridVsock(PathBuf),
363}
364
365impl FromStr for VmId {
366    type Err = ParseVmIdError;
367
368    fn from_str(s: &str) -> Result<Self, Self::Err> {
369        if let Some(s) = s.strip_prefix("vsock:") {
370            Ok(Self::HybridVsock(Path::new(s).to_owned()))
371        } else {
372            #[cfg(windows)]
373            {
374                if let Some(rest) = s.strip_prefix("hyperv-id:") {
375                    let guid = rest
376                        .parse::<guid::Guid>()
377                        .map_err(|_| ParseVmIdError::InvalidGuid(rest.to_owned()))?;
378                    return Ok(Self::HyperVId(guid));
379                }
380
381                if let Some(name) = s.strip_prefix("hyperv:") {
382                    return Ok(Self::HyperV(name.to_owned()));
383                }
384
385                if !pal::windows::fs::is_unix_socket(s.as_ref()).unwrap_or(false) {
386                    return Ok(Self::HyperV(s.to_owned()));
387                }
388            }
389            // Default to hybrid vsock since this is what OpenVMM supports for
390            // Underhill.
391            Ok(Self::HybridVsock(Path::new(s).to_owned()))
392        }
393    }
394}
395
396/// Error parsing a [`VmId`].
397#[derive(Debug, Error)]
398enum ParseVmIdError {
399    #[cfg(windows)]
400    #[error("invalid VM ID GUID '{0}' (expected a GUID, with or without braces)")]
401    InvalidGuid(String),
402}
403
404#[derive(Clone)]
405struct EnvString {
406    name: String,
407    value: String,
408}
409
410#[derive(Clone, clap::ValueEnum)]
411enum CrashType {
412    #[clap(name = "panic")]
413    UhPanic,
414}
415
416#[derive(Clone, clap::ValueEnum)]
417enum EfiDiagnosticsLogLevel {
418    /// Errors and warnings only
419    Default,
420    /// Errors, warnings, and info
421    Info,
422    /// All log levels
423    Full,
424}
425
426impl EfiDiagnosticsLogLevel {
427    fn as_inspect_value(&self) -> &'static str {
428        match self {
429            EfiDiagnosticsLogLevel::Default => "default",
430            EfiDiagnosticsLogLevel::Info => "info",
431            EfiDiagnosticsLogLevel::Full => "full",
432        }
433    }
434}
435
436#[derive(Clone, clap::ValueEnum)]
437enum EfiDiagnosticsOutput {
438    /// Emit to stdout
439    Stdout,
440    /// Emit to tracing
441    Tracing,
442}
443
444impl EfiDiagnosticsOutput {
445    fn as_inspect_value(&self) -> &'static str {
446        match self {
447            EfiDiagnosticsOutput::Stdout => "stdout",
448            EfiDiagnosticsOutput::Tracing => "tracing",
449        }
450    }
451}
452
453#[derive(Debug, Error)]
454#[error("bad environment variable, expected VAR=value")]
455struct BadEnvString;
456
457impl FromStr for EnvString {
458    type Err = BadEnvString;
459
460    fn from_str(s: &str) -> Result<Self, Self::Err> {
461        let (name, value) = s.split_once('=').ok_or(BadEnvString)?;
462        Ok(Self {
463            name: name.to_owned(),
464            value: value.to_owned(),
465        })
466    }
467}
468
469// N.B. this exits after a successful completion.
470async fn run(
471    client: &DiagClient,
472    command: impl AsRef<str>,
473    args: impl IntoIterator<Item = impl AsRef<str>>,
474) -> anyhow::Result<()> {
475    // TODO: if stdout and stderr of this process are backed by the
476    // same thing, then pass combine_stderr instead.
477    let mut process = client
478        .exec(&command)
479        .args(args)
480        .stdin(true)
481        .stdout(true)
482        .stderr(true)
483        .spawn()
484        .await?;
485
486    let mut stdin = process.stdin.take().unwrap();
487    let mut stdout = process.stdout.take().unwrap();
488    let mut stderr = process.stderr.take().unwrap();
489
490    std::thread::spawn({
491        move || {
492            let _ = std::io::copy(&mut std::io::stdin(), &mut stdin);
493        }
494    });
495
496    let stderr_thread =
497        std::thread::spawn(move || std::io::copy(&mut stderr, &mut term::raw_stderr()));
498
499    std::io::copy(&mut stdout, &mut term::raw_stdout()).context("failed stdout copy")?;
500
501    stderr_thread
502        .join()
503        .unwrap()
504        .context("failed stderr thread")?;
505
506    let status = process.wait().await?;
507    std::process::exit(status.exit_code());
508}
509
510fn new_client(driver: impl Driver + Spawn + Clone, input: &VmArg) -> anyhow::Result<DiagClient> {
511    let client = match &input.id {
512        #[cfg(windows)]
513        VmId::HyperV(name) => DiagClient::from_hyperv_name(driver, name)?,
514        #[cfg(windows)]
515        VmId::HyperVId(guid) => DiagClient::from_hyperv_id(driver, *guid),
516        VmId::HybridVsock(path) => DiagClient::from_hybrid_vsock(driver, path),
517    };
518    Ok(client)
519}
520
521pub fn main() -> anyhow::Result<()> {
522    tracing_subscriber::registry()
523        .with(tracing_subscriber::fmt::layer())
524        .with(tracing_subscriber::EnvFilter::from_default_env())
525        .init();
526
527    term::enable_vt_and_utf8();
528    DefaultPool::run_with(async |driver| {
529        let Options { vm, command } = Options::parse();
530
531        match command {
532            Command::Complete(cmd) => {
533                cmd.println_to_stub_script::<Options>(
534                    None,
535                    completions::OhcldiagDevCompleteFactory {
536                        driver: driver.clone(),
537                    },
538                )
539                .await
540            }
541            Command::Completions(cmd) => cmd.run()?,
542            Command::Shell { shell, args } => {
543                let client = new_client(driver.clone(), &vm)?;
544
545                // Set TERM to ensure function keys and other characters work.
546                let term = std::env::var("TERM");
547                let term = term.as_deref().unwrap_or("xterm-256color");
548
549                let mut process = client
550                    .exec(&shell)
551                    .args(&args)
552                    .tty(true)
553                    .stdin(true)
554                    .stdout(true)
555                    .env("TERM", term)
556                    .spawn()
557                    .await?;
558
559                let mut stdin = process.stdin.take().unwrap();
560                let mut stdout = process.stdout.take().unwrap();
561
562                crossterm::terminal::enable_raw_mode().expect("failed to set raw console mode");
563                std::thread::spawn({
564                    move || {
565                        let _ = std::io::copy(&mut std::io::stdin(), &mut stdin);
566                    }
567                });
568
569                std::io::copy(&mut stdout, &mut term::raw_stdout()).context("failed copy")?;
570
571                let status = process.wait().await?;
572
573                if !status.success() {
574                    eprintln!(
575                        "shell exited with non-zero exit code: {}",
576                        status.exit_code()
577                    );
578                }
579            }
580            Command::Run { command, args } => {
581                let client = new_client(driver.clone(), &vm)?;
582                run(&client, command, &args).await?;
583            }
584            Command::Inspect {
585                recursive,
586                limit,
587                json,
588                poll,
589                period,
590                count,
591                timeout,
592
593                path,
594                update,
595            } => {
596                let client = new_client(driver.clone(), &vm)?;
597
598                if let Some(update) = update {
599                    let Some(path) = path else {
600                        anyhow::bail!("must provide path for update")
601                    };
602
603                    let value = client.update(path, update).await?;
604                    match value.kind {
605                        inspect::ValueKind::String(s) => println!("{s}"),
606                        _ => println!("{value}"),
607                    }
608                } else {
609                    let timeout = if timeout == 0 {
610                        None
611                    } else {
612                        Some(Duration::from_secs(timeout))
613                    };
614                    let query = async || {
615                        client
616                            .inspect(
617                                path.as_deref().unwrap_or(""),
618                                if recursive { limit } else { Some(0) },
619                                timeout,
620                            )
621                            .await
622                    };
623
624                    if poll {
625                        let mut timer = PolledTimer::new(&driver);
626                        let period = Duration::from_secs_f64(period);
627                        let mut last_time = pal_async::timer::Instant::now();
628                        let mut last = query().await?;
629                        let mut count = count;
630
631                        loop {
632                            match count.as_mut() {
633                                Some(count) if *count == 0 => break,
634                                Some(count) => *count -= 1,
635                                None => {}
636                            }
637                            timer.sleep_until(last_time + period).await;
638                            let now = pal_async::timer::Instant::now();
639                            let this = query().await?;
640                            let diff = this.since(&last, now - last_time);
641                            if json {
642                                println!("{}", diff.json());
643                            } else {
644                                println!("{diff:#}");
645                            }
646                            last = this;
647                            last_time = now;
648                        }
649                    } else {
650                        let node = query().await?;
651                        if json {
652                            println!("{}", node.json());
653                        } else {
654                            println!("{node:#}");
655                        }
656                    }
657                }
658            }
659            Command::Update { path, value } => {
660                eprintln!(
661                    "`update` is deprecated - please use `ohcldiag-dev inspect <path> -u <new value>`"
662                );
663                let client = new_client(driver.clone(), &vm)?;
664                let value = client.update(path, value).await?;
665                match value.kind {
666                    inspect::ValueKind::String(s) => println!("{s}"),
667                    _ => println!("{value}"),
668                }
669            }
670            Command::Start { env, unset, args } => {
671                let client = new_client(driver.clone(), &vm)?;
672
673                let env = env
674                    .into_iter()
675                    .map(|EnvString { name, value }| (name, Some(value)))
676                    .chain(unset.into_iter().map(|name| (name, None)));
677
678                client.start(env, args).await?;
679            }
680            Command::Kmsg {
681                follow,
682                reconnect,
683                verbose,
684                #[cfg(windows)]
685                serial,
686                #[cfg(windows)]
687                pipe_path,
688            } => {
689                let is_terminal = std::io::stdout().is_terminal();
690
691                #[cfg(windows)]
692                if serial {
693                    use diag_client::hyperv::ComPortAccessInfo;
694                    use futures::AsyncBufReadExt;
695
696                    let port_access_info = if let Some(pipe_path) = pipe_path.as_ref() {
697                        ComPortAccessInfo::PortPipePath(pipe_path)
698                    } else {
699                        match &vm.id {
700                            VmId::HyperV(name) => ComPortAccessInfo::NameAndPortNumber(name, 3),
701                            #[cfg(windows)]
702                            VmId::HyperVId(guid) => ComPortAccessInfo::IdAndPortNumber(*guid, 3),
703                            _ => anyhow::bail!("--serial is only supported for Hyper-V VMs"),
704                        }
705                    };
706
707                    let pipe =
708                        diag_client::hyperv::open_serial_port(&driver, port_access_info).await?;
709                    let pipe = pal_async::pipe::PolledPipe::new(&driver, pipe)
710                        .context("failed to make a polled pipe")?;
711                    let pipe = futures::io::BufReader::new(pipe);
712
713                    let mut lines = pipe.lines();
714                    while let Some(line) = lines.next().await {
715                        let line = line?;
716                        if let Some(message) = kmsg::SyslogParsedEntry::new(&line) {
717                            println!("{}", message.display(is_terminal));
718                        } else {
719                            println!("{line}");
720                        }
721                    }
722
723                    return Ok(());
724                }
725
726                if verbose {
727                    eprintln!("Connecting to the diagnostics server.");
728                }
729
730                let client = new_client(driver.clone(), &vm)?;
731                'connect: loop {
732                    if reconnect {
733                        client.wait_for_server().await?;
734                    }
735                    let mut file_stream = client.kmsg(follow).await?;
736                    if verbose {
737                        eprintln!("Connected.");
738                    }
739
740                    while let Some(data) = file_stream.next().await {
741                        match data {
742                            Ok(data) => match kmsg::KmsgParsedEntry::new(&data) {
743                                Ok(message) => println!("{}", message.display(is_terminal)),
744                                Err(e) => println!("Invalid kmsg entry: {e:?}"),
745                            },
746                            Err(err) if reconnect && err.kind() == ErrorKind::ConnectionReset => {
747                                if verbose {
748                                    eprintln!(
749                                        "Connection reset to the diagnostics server. Reconnecting."
750                                    );
751                                }
752                                continue 'connect;
753                            }
754                            Err(err) => Err(err).context("failed to read kmsg")?,
755                        }
756                    }
757
758                    if reconnect {
759                        if verbose {
760                            eprintln!("Lost connection to the diagnostics server. Reconnecting.");
761                        }
762                        continue 'connect;
763                    }
764
765                    break;
766                }
767            }
768            Command::File { follow, file_path } => {
769                let client = new_client(driver.clone(), &vm)?;
770                let stream = client.read_file(follow, file_path).await?;
771                futures::io::copy(stream, &mut AllowStdIo::new(term::raw_stdout()))
772                    .await
773                    .context("failed to copy trace file")?;
774            }
775            Command::Gdbserver { multi, pid } => {
776                let client = new_client(driver.clone(), &vm)?;
777                // Pass the --once flag so that gdbserver exits after the stdio
778                // pipes are closed. Otherwise, gdbserver spins in a tight loop
779                // and never exits.
780                let gdbserver = "gdbserver --once";
781                let command = if multi {
782                    format!("{gdbserver} --multi -")
783                } else if let Some(pid) = pid {
784                    format!("{gdbserver} --attach - {pid}")
785                } else {
786                    format!("{gdbserver} --attach - \"$(cat /run/underhill.pid)\"")
787                };
788
789                run(&client, "/bin/sh", &["-c", &command]).await?;
790            }
791            Command::Gdbstub { port } => {
792                let vsock = match vm.id {
793                    VmId::HybridVsock(path) => {
794                        diag_client::connect_hybrid_vsock(&driver, &path, port).await?
795                    }
796                    #[cfg(windows)]
797                    VmId::HyperV(name) => {
798                        let vm_id = diag_client::hyperv::vm_id_from_name(&name)?;
799                        let stream =
800                            diag_client::hyperv::connect_vsock(&driver, vm_id, port).await?;
801                        PolledSocket::new(&driver, socket2::Socket::from(stream))?
802                    }
803                    #[cfg(windows)]
804                    VmId::HyperVId(vm_id) => {
805                        let stream =
806                            diag_client::hyperv::connect_vsock(&driver, vm_id, port).await?;
807                        PolledSocket::new(&driver, socket2::Socket::from(stream))?
808                    }
809                };
810
811                let vsock = Arc::new(vsock.into_inner());
812                // Spawn a thread to read stdin synchronously since pal_async
813                // does not offer a way to read it asynchronously.
814                let thread = std::thread::spawn({
815                    let vsock = vsock.clone();
816                    move || {
817                        let _ = std::io::copy(&mut std::io::stdin(), &mut vsock.as_ref());
818                    }
819                });
820
821                std::io::copy(&mut vsock.as_ref(), &mut term::raw_stdout())
822                    .context("failed stdout copy")?;
823                thread.join().unwrap();
824            }
825            Command::Crash {
826                crash_type,
827                pid,
828                name,
829            } => {
830                let client = new_client(driver.clone(), &vm)?;
831                let pid = if let Some(name) = name {
832                    client.get_pid(&name).await?
833                } else {
834                    pid.unwrap()
835                };
836                println!("Crashing PID: {pid}");
837                match crash_type {
838                    CrashType::UhPanic => {
839                        _ = client.crash(pid).await;
840                    }
841                }
842            }
843            Command::PacketCapture {
844                output,
845                seconds,
846                snaplen,
847            } => {
848                let client = new_client(driver.clone(), &vm)?;
849                println!(
850                    "Starting network packet capture. Wait for timeout or Ctrl-C to quit anytime."
851                );
852                let (_, num_streams) = client
853                    .packet_capture(PacketCaptureOperation::Query, 0, 0)
854                    .await?;
855                let file_stem = &output.file_stem().unwrap().to_string_lossy();
856                let extension = &output.extension().unwrap_or(OsStr::new("pcap"));
857                let mut new_output = PathBuf::from(&output);
858                let streams = client
859                    .packet_capture(PacketCaptureOperation::Start, num_streams, snaplen)
860                    .await?
861                    .0
862                    .into_iter()
863                    .enumerate()
864                    .map(|(i, i_stream)| {
865                        new_output.set_file_name(format!("{}-{}", file_stem, i));
866                        new_output.set_extension(extension);
867                        let mut out = AllowStdIo::new(fs_err::File::create(&new_output)?);
868                        Ok(async move { futures::io::copy(i_stream, &mut out).await })
869                    })
870                    .collect::<Result<Vec<_>, std::io::Error>>()?;
871                capture_packets(client, streams, seconds).await;
872            }
873            Command::CoreDump {
874                verbose,
875                pid,
876                name,
877                dst,
878            } => {
879                ensure_not_terminal(&dst)?;
880                let client = new_client(driver.clone(), &vm)?;
881                let pid = if let Some(name) = name {
882                    client.get_pid(&name).await?
883                } else {
884                    pid.unwrap()
885                };
886                println!("Dumping PID: {pid}");
887                let file = create_or_stderr(&dst)?;
888                client
889                    .core_dump(
890                        pid,
891                        AllowStdIo::new(file),
892                        AllowStdIo::new(std::io::stderr()),
893                        verbose,
894                    )
895                    .await?;
896            }
897            Command::Restart => {
898                let client = new_client(driver.clone(), &vm)?;
899                client.restart().await?;
900            }
901            Command::PerfTrace { output } => {
902                ensure_not_terminal(&output)?;
903
904                let client = new_client(driver.clone(), &vm)?;
905
906                // Flush the perf trace.
907                client
908                    .update("trace/perf/flush".to_owned(), "true".to_owned())
909                    .await
910                    .context("failed to flush perf")?;
911
912                let file = create_or_stderr(&output)?;
913                let stream = client
914                    .read_file(false, "underhill.perfetto".to_owned())
915                    .await
916                    .context("failed to read trace file")?;
917
918                futures::io::copy(stream, &mut AllowStdIo::new(file))
919                    .await
920                    .context("failed to copy trace file")?;
921            }
922            Command::VsockTcpRelay {
923                vsock_port,
924                tcp_port,
925                allow_remote,
926                reconnect,
927            } => {
928                let addr = if allow_remote { "0.0.0.0" } else { "127.0.0.1" };
929                let listener = TcpListener::bind((addr, tcp_port))
930                    .with_context(|| format!("binding to port {}", tcp_port))?;
931                println!("TCP listening on {}:{}", addr, tcp_port);
932                'connect: loop {
933                    let (tcp_socket, tcp_addr) = listener.accept()?;
934                    let tcp_socket = PolledSocket::new(&driver, tcp_socket)?;
935                    println!("TCP accept on {:?}", tcp_addr);
936
937                    // TODO: support reconnect attempt for vsock like kmsg
938                    let vsock = match vm.id {
939                        VmId::HybridVsock(ref path) => {
940                            // TODO: reconnection attempt logic like kmsg is
941                            // broken for hybrid_vsock with end of file error,
942                            // if this is started before the vm is started
943                            diag_client::connect_hybrid_vsock(&driver, path, vsock_port).await?
944                        }
945                        #[cfg(windows)]
946                        VmId::HyperV(ref name) => {
947                            let vm_id = diag_client::hyperv::vm_id_from_name(name)?;
948                            let stream =
949                                diag_client::hyperv::connect_vsock(&driver, vm_id, vsock_port)
950                                    .await?;
951                            PolledSocket::new(&driver, socket2::Socket::from(stream))?
952                        }
953                        #[cfg(windows)]
954                        VmId::HyperVId(ref vm_id) => {
955                            let stream =
956                                diag_client::hyperv::connect_vsock(&driver, *vm_id, vsock_port)
957                                    .await?;
958                            PolledSocket::new(&driver, socket2::Socket::from(stream))?
959                        }
960                    };
961                    println!("VSOCK connect to port {:?}", vsock_port);
962
963                    let (tcp_read, mut tcp_write) = tcp_socket.split();
964                    let (vsock_read, mut vsock_write) = vsock.split();
965                    let tx = futures::io::copy(tcp_read, &mut vsock_write);
966                    let rx = futures::io::copy(vsock_read, &mut tcp_write);
967                    let result = futures::future::try_join(tx, rx).await;
968                    match result {
969                        Ok(_) => {}
970                        Err(e) => match e.kind() {
971                            ErrorKind::ConnectionReset => {}
972                            _ => return Err(anyhow::Error::from(e)),
973                        },
974                    }
975                    println!("Connection closed");
976
977                    if reconnect {
978                        println!("Reconnecting...");
979                        continue 'connect;
980                    }
981
982                    break;
983                }
984            }
985            Command::Pause => {
986                let client = new_client(driver.clone(), &vm)?;
987                client.pause().await?;
988            }
989            Command::Resume => {
990                let client = new_client(driver.clone(), &vm)?;
991                client.resume().await?;
992            }
993            Command::DumpSavedState { output } => {
994                ensure_not_terminal(&output)?;
995                let client = new_client(driver.clone(), &vm)?;
996                let mut file = create_or_stderr(&output)?;
997                file.write_all(&client.dump_saved_state().await?)?;
998            }
999            Command::MemoryProfileTrace { pid, name, output } => {
1000                let client = new_client(driver.clone(), &vm)?;
1001                let pid = if let Some(name) = name {
1002                    client.get_pid(&name).await?
1003                } else if let Some(pid) = pid {
1004                    pid
1005                } else {
1006                    anyhow::bail!("either --pid or --name must be specified");
1007                };
1008                // Do not write anything on the stdout in case the output
1009                // is set to stdout, to avoid breaking the output format
1010                // of the trace.
1011                let mut file = create_or_stderr(&output)?;
1012                file.write_all(&client.memory_profile_trace(pid).await?)?;
1013            }
1014            Command::EfiDiagnostics { log_level, output } => {
1015                let client = new_client(driver.clone(), &vm)?;
1016                let arg = format!(
1017                    "{},{}",
1018                    log_level.as_inspect_value(),
1019                    output.as_inspect_value()
1020                );
1021                let value = client
1022                    .update("vm/uefi/process_diagnostics", &arg)
1023                    .await
1024                    .context("failed to process EFI diagnostics")?;
1025                match value.kind {
1026                    inspect::ValueKind::String(s) => print!("{s}"),
1027                    _ => print!("{value}"),
1028                }
1029            }
1030        }
1031        Ok(())
1032    })
1033}
1034
1035fn ensure_not_terminal(path: &Option<PathBuf>) -> anyhow::Result<()> {
1036    if path.is_none() && std::io::stdout().is_terminal() {
1037        anyhow::bail!("cannot write to terminal");
1038    }
1039    Ok(())
1040}
1041
1042fn create_or_stderr(path: &Option<PathBuf>) -> std::io::Result<fs_err::File> {
1043    let file = match path {
1044        Some(path) => fs_err::File::create(path)?,
1045        None => fs_err::File::from_parts(term::raw_stdout(), "stdout"),
1046    };
1047    Ok(file)
1048}
1049
1050async fn capture_packets(
1051    client: DiagClient,
1052    streams: Vec<impl Future<Output = Result<u64, std::io::Error>>>,
1053    capture_duration: Duration,
1054) {
1055    let mut capture_streams = FuturesUnordered::from_iter(streams);
1056    let (user_input_tx, mut user_input_rx) = mesh::channel();
1057    ctrlc::set_handler(move || user_input_tx.send(())).expect("Error setting Ctrl-C handler");
1058
1059    let mut ctx = mesh::CancelContext::new().with_timeout(capture_duration);
1060    let mut stop_signaled = std::pin::pin!(ctx.until_cancelled(user_input_rx.recv()));
1061
1062    let mut stop_streams = std::pin::pin!(async {
1063        if let Err(err) = client
1064            .packet_capture(PacketCaptureOperation::Stop, 0, 0)
1065            .await
1066        {
1067            eprintln!("Failed stop: {err}");
1068        }
1069    });
1070
1071    #[derive(PartialEq)]
1072    enum State {
1073        Running,
1074        Stopping,
1075        StoppingStreamsDone,
1076        Stopped,
1077    }
1078    let mut state = State::Running;
1079    loop {
1080        enum Event {
1081            Continue,
1082            StopSignaled,
1083            StopComplete,
1084            StreamsDone,
1085        }
1086        let stop = async {
1087            match state {
1088                State::Running => {
1089                    (&mut stop_signaled).await.ok();
1090                    Event::StopSignaled
1091                }
1092                State::Stopping | State::StoppingStreamsDone => {
1093                    (&mut stop_streams).await;
1094                    Event::StopComplete
1095                }
1096                State::Stopped => std::future::pending::<Event>().await,
1097            }
1098        };
1099        let process_streams = async {
1100            if state == State::StoppingStreamsDone {
1101                std::future::pending::<()>().await;
1102            }
1103            match capture_streams.next().await {
1104                Some(_) => Event::Continue,
1105                None => Event::StreamsDone,
1106            }
1107        };
1108        let event = (stop, process_streams).race();
1109
1110        // N.B Wait for all the copy tasks to complete to make sure the data is flushed to
1111        //     ensure compatibility with the packet capture protocol.
1112        match event.await {
1113            Event::Continue => continue,
1114            Event::StopSignaled => {
1115                println!("Stopping packet capture...");
1116                state = State::Stopping;
1117            }
1118            Event::StopComplete => {
1119                println!("Waiting for data to be flushed...");
1120                if state == State::Stopping {
1121                    state = State::Stopped;
1122                } else {
1123                    break;
1124                }
1125            }
1126            Event::StreamsDone if state == State::Stopping => {
1127                state = State::StoppingStreamsDone;
1128            }
1129            Event::StreamsDone => {
1130                if state != State::Stopped {
1131                    println!("Lost connection with network.");
1132                }
1133                break;
1134            }
1135        }
1136    }
1137    println!("All done.");
1138}
1139
1140#[cfg(all(test, windows))]
1141mod tests {
1142    use super::VmId;
1143
1144    #[test]
1145    fn bare_guid_is_treated_as_name() {
1146        // Fleet VM names are GUID-shaped; a bare value must be a name, not an ID.
1147        let id: VmId = "1965676a-8dd3-4b46-9439-40c6a30e5b1a".parse().unwrap();
1148        assert!(matches!(id, VmId::HyperV(name) if name == "1965676a-8dd3-4b46-9439-40c6a30e5b1a"));
1149    }
1150
1151    #[test]
1152    fn hyperv_prefix_is_a_name() {
1153        let id: VmId = "hyperv:my-vm".parse().unwrap();
1154        assert!(matches!(id, VmId::HyperV(name) if name == "my-vm"));
1155    }
1156
1157    #[test]
1158    fn hyperv_prefix_with_guid_is_still_a_name() {
1159        let id: VmId = "hyperv:1965676a-8dd3-4b46-9439-40c6a30e5b1a"
1160            .parse()
1161            .unwrap();
1162        assert!(matches!(id, VmId::HyperV(name) if name == "1965676a-8dd3-4b46-9439-40c6a30e5b1a"));
1163    }
1164
1165    #[test]
1166    fn hyperv_id_prefix_parses_guid() {
1167        let id: VmId = "hyperv-id:1965676a-8dd3-4b46-9439-40c6a30e5b1a"
1168            .parse()
1169            .unwrap();
1170        assert!(matches!(id, VmId::HyperVId(_)));
1171    }
1172
1173    #[test]
1174    fn hyperv_id_prefix_parses_braced_guid() {
1175        let id: VmId = "hyperv-id:{1965676a-8dd3-4b46-9439-40c6a30e5b1a}"
1176            .parse()
1177            .unwrap();
1178        assert!(matches!(id, VmId::HyperVId(_)));
1179    }
1180
1181    #[test]
1182    fn hyperv_id_prefix_with_invalid_guid_errors() {
1183        assert!("hyperv-id:not-a-guid".parse::<VmId>().is_err());
1184    }
1185
1186    #[test]
1187    fn vsock_prefix_is_hybrid_vsock() {
1188        let id: VmId = "vsock:/tmp/vm.sock".parse().unwrap();
1189        assert!(matches!(id, VmId::HybridVsock(_)));
1190    }
1191}