1#![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 Shell {
81 #[clap(default_value = "/bin/sh")]
83 shell: String,
84 args: Vec<String>,
86 },
87 Run {
89 command: String,
91 args: Vec<String>,
93 },
94 #[clap(visible_alias = "i")]
96 Inspect {
97 #[clap(short)]
99 recursive: bool,
100 #[clap(short, long, requires("recursive"))]
102 limit: Option<usize>,
103 #[clap(short, long)]
105 json: bool,
106 #[clap(short)]
108 poll: bool,
109 #[clap(long, default_value = "1", requires("poll"))]
111 period: f64,
112 #[clap(long, requires("poll"))]
114 count: Option<usize>,
115 path: Option<String>,
117 #[clap(short, long, conflicts_with("recursive"))]
119 update: Option<String>,
120 #[clap(short, default_value = "1", conflicts_with("update"))]
122 timeout: u64,
123 },
124 #[clap(hide = true)]
126 Update {
127 path: String,
129 value: String,
131 },
132 Start {
137 #[clap(short, long)]
139 env: Vec<EnvString>,
140 #[clap(short, long)]
142 unset: Vec<String>,
143 args: Vec<String>,
145 },
146 Kmsg {
148 #[clap(short, long)]
150 follow: bool,
151 #[clap(short, long)]
153 reconnect: bool,
154 #[clap(short, long)]
156 verbose: bool,
157 #[cfg(windows)]
161 #[clap(long, conflicts_with = "reconnect")]
162 serial: bool,
163 #[cfg(windows)]
167 #[clap(long, requires = "serial")]
168 pipe_path: Option<String>,
169 },
170 File {
172 #[clap(short, long)]
174 follow: bool,
175 #[clap(short('p'), long)]
176 file_path: String,
177 },
178 Gdbserver {
188 #[clap(long)]
190 pid: Option<i32>,
191 #[clap(long, conflicts_with("pid"))]
193 multi: bool,
194 },
195 Gdbstub {
202 #[clap(short, long, default_value = "4")]
204 port: u32,
205 },
206 #[clap(group(
210 ArgGroup::new("process")
211 .required(true)
212 .args(&["pid", "name"]),
213 ))]
214 Crash {
215 crash_type: CrashType,
219 #[clap(short, long)]
221 pid: Option<i32>,
222 #[clap(short, long)]
224 name: Option<String>,
225 },
226 #[clap(group(
231 ArgGroup::new("process")
232 .required(true)
233 .args(&["pid", "name"]),
234 ))]
235 CoreDump {
236 #[clap(short, long)]
238 verbose: bool,
239 #[clap(short, long)]
241 pid: Option<i32>,
242 #[clap(short, long)]
244 name: Option<String>,
245 dst: Option<PathBuf>,
248 },
249 Restart,
251 PerfTrace {
254 #[clap(short)]
256 output: Option<PathBuf>,
257 },
258 VsockTcpRelay {
260 vsock_port: u32,
261 tcp_port: u16,
262 #[clap(long)]
263 allow_remote: bool,
264 #[clap(short, long)]
270 reconnect: bool,
271 },
272 Pause,
274 Resume,
276 DumpSavedState {
278 #[clap(short)]
280 output: Option<PathBuf>,
281 },
282 PacketCapture {
284 #[clap(short('w'), default_value = "nic")]
286 output: PathBuf,
287 #[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 #[clap(short('s'), long, default_value = "65535", value_parser = clap::value_parser!(u16).range(1..))]
292 snaplen: u16,
293 },
294 MemoryProfileTrace {
296 #[clap(short, long)]
298 pid: Option<i32>,
299 #[clap(short, long)]
301 name: Option<String>,
302 #[clap(short)]
304 output: Option<PathBuf>,
305 },
306 EfiDiagnostics {
312 log_level: EfiDiagnosticsLogLevel,
317 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 Ok(Self::HybridVsock(Path::new(s).to_owned()))
392 }
393 }
394}
395
396#[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 Default,
420 Info,
422 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 Stdout,
440 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
469async fn run(
471 client: &DiagClient,
472 command: impl AsRef<str>,
473 args: impl IntoIterator<Item = impl AsRef<str>>,
474) -> anyhow::Result<()> {
475 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 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 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 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 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 let vsock = match vm.id {
939 VmId::HybridVsock(ref path) => {
940 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 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 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 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}