diff --git a/scripts/nvx.py b/scripts/nvx.py index 2c777da2..c4076d8c 100644 --- a/scripts/nvx.py +++ b/scripts/nvx.py @@ -263,6 +263,25 @@ def _format_command(command: list[str]) -> str: return subprocess.list2cmdline(command) if os.name == "nt" else shlex.join(command) +def _extend_network_arguments(command: list[str], args: argparse.Namespace) -> None: + if args.net is not None: + command.extend(["--net", args.net, "--network-profile", args.network_profile]) + if args.network_egress is not None: + command.extend(["--network-egress", args.network_egress]) + if args.network_ingress is not None: + command.extend(["--network-ingress", args.network_ingress]) + for rule in args.network_egress_allow: + command.extend(["--network-egress-allow", rule]) + for rule in args.network_egress_deny: + command.extend(["--network-egress-deny", rule]) + if args.host_loopback is not None: + command.extend(["--host-loopback", args.host_loopback]) + if args.network_proxy is not None: + command.extend(["--network-proxy", args.network_proxy]) + for forward in args.host_loopback_forward: + command.extend(["--host-loopback-forward", forward]) + + def command_run(args: argparse.Namespace) -> None: if (args.net is None) != (args.network_profile is None): raise ScriptError("--net and --network-profile must be specified together") @@ -336,22 +355,7 @@ def command_run(args: argparse.Namespace) -> None: command.extend(["--mount", args.mount]) for denied_path in args.mount_deny: command.extend(["--mount-deny", str(denied_path)]) - if args.net is not None: - command.extend(["--net", args.net, "--network-profile", args.network_profile]) - if args.network_egress is not None: - command.extend(["--network-egress", args.network_egress]) - if args.network_ingress is not None: - command.extend(["--network-ingress", args.network_ingress]) - for rule in args.network_egress_allow: - command.extend(["--network-egress-allow", rule]) - for rule in args.network_egress_deny: - command.extend(["--network-egress-deny", rule]) - if args.host_loopback is not None: - command.extend(["--host-loopback", args.host_loopback]) - if args.network_proxy is not None: - command.extend(["--network-proxy", args.network_proxy]) - for forward in args.host_loopback_forward: - command.extend(["--host-loopback-forward", forward]) + _extend_network_arguments(command, args) if args.outcome_report is not None: command.extend(["--microvm-report", str(args.outcome_report)]) if args.cmdline: @@ -477,22 +481,7 @@ def command_sandbox(args: argparse.Namespace) -> None: "--cmdline", launch.kernel_command_line(args.cmdline), ] - if args.net is not None: - command.extend(["--net", args.net, "--network-profile", args.network_profile]) - if args.network_egress is not None: - command.extend(["--network-egress", args.network_egress]) - if args.network_ingress is not None: - command.extend(["--network-ingress", args.network_ingress]) - for rule in args.network_egress_allow: - command.extend(["--network-egress-allow", rule]) - for rule in args.network_egress_deny: - command.extend(["--network-egress-deny", rule]) - if args.host_loopback is not None: - command.extend(["--host-loopback", args.host_loopback]) - if args.network_proxy is not None: - command.extend(["--network-proxy", args.network_proxy]) - for forward in args.host_loopback_forward: - command.extend(["--host-loopback-forward", forward]) + _extend_network_arguments(command, args) if args.outcome_report is not None: command.extend(["--microvm-report", str(args.outcome_report)]) print(f">> {_format_command(command)}") diff --git a/scripts/test_nvx_tools.py b/scripts/test_nvx_tools.py index 9474c6a6..c681a3fc 100644 --- a/scripts/test_nvx_tools.py +++ b/scripts/test_nvx_tools.py @@ -710,6 +710,48 @@ def test_network_requires_explicit_portable_profile(self): with self.assertRaisesRegex(common.ScriptError, "--net and --network-profile"): nvx.command_run(missing_network) + def test_run_and_sandbox_forward_network_arguments(self): + network_arguments = ( + "--net 10.0.0.2/24 --network-profile portable " + "--network-egress deny --network-ingress deny " + "--network-egress-allow 140.82.112.0/20:tcp:443 " + "--network-egress-deny 10.0.0.0/8 --host-loopback allow " + "--network-proxy 10.0.0.1:3128 --host-loopback-forward tcp:8080:80" + ).split() + commands = [ + ["run", "--dry-run", *network_arguments], + [ + "sandbox", + "--dry-run", + "--layer", + "distro,distro.erofs,11111111-1111-1111-1111-111111111111", + "--scratch", + "scratch.ext4", + *network_arguments, + ], + ] + + def require(path: Path, _description: str) -> Path: + return path + + for arguments in commands: + with self.subTest(command=arguments[0]): + args = nvx.parse_args(arguments) + with ( + patch.object(nvx, "require_file", side_effect=require), + patch.object(sandbox, "require_file", side_effect=require), + patch.object( + nvx, "_format_command", return_value="formatted" + ) as format_command, + ): + args.handler(args) + + command = format_command.call_args.args[0] + self.assertEqual( + command[command.index("--net") :], + network_arguments, + ) + def test_run_parses_denied_filesystem_paths(self): args = nvx.parse_args( [