From 79ae5ff31ca588c06c7606f53134eaa0e6591cdc Mon Sep 17 00:00:00 2001 From: oliver Date: Tue, 23 Jun 2026 21:31:46 +0800 Subject: [PATCH] fix(bastion): use keyboard-interactive auth for protocol-proxy bastions ZTE-TSM and similar bastions reject standard SSH password auth and require Vault password via keyboard-interactive; connect over an authenticated Paramiko session instead of re-handshaking with Netmiko. Co-authored-by: Cursor --- netx_api/ne_session_factory.py | 149 ++++++++++++++++++++++++++++++--- tests/test_bastion_hop.py | 61 +++++++++++--- 2 files changed, 190 insertions(+), 20 deletions(-) diff --git a/netx_api/ne_session_factory.py b/netx_api/ne_session_factory.py index 03cb683..65d8463 100644 --- a/netx_api/ne_session_factory.py +++ b/netx_api/ne_session_factory.py @@ -99,6 +99,116 @@ def render_hop_command(template: str, creds: dict[str, Any]) -> str: return out +def _bastion_interactive_handler(password: str): + """Reply to bastion keyboard-interactive prompts (Vault password, OTP, etc.).""" + + def handler(title: str, instructions: str, prompt_list: list[tuple[str, bool]]) -> list[str]: + if not prompt_list: + return [] + responses: list[str] = [] + for prompt, _echo in prompt_list: + pl = str(prompt or "").lower() + if re.search(r"password|vault|口令|密码|passcode|otp|token|verification|verify", pl): + responses.append(password) + else: + responses.append("") + if not any(responses): + responses = [password] * len(prompt_list) + return responses + + return handler + + +def _bastion_ssh_connect( + *, + host: str, + port: int, + username: str, + password: str, + timeout: int, +) -> paramiko.SSHClient: + """SSH to protocol-proxy bastion; interactive auth first (JumpServer/CBH/ZTE-TSM).""" + client = paramiko.SSHClient() + client.set_missing_host_key_policy(paramiko.AutoAddPolicy()) + transport = paramiko.Transport((host, int(port or 22))) + transport.banner_timeout = timeout + transport.auth_timeout = timeout + try: + transport.start_client(timeout=timeout) + except Exception: + try: + transport.close() + except Exception: + pass + raise + + handler = _bastion_interactive_handler(password) + auth_errors: list[Exception] = [] + for use_interactive in (True, False): + try: + if use_interactive: + transport.auth_interactive(username, handler) + else: + transport.auth_password(username, password, fallback=False) + if transport.is_authenticated(): + client._transport = transport # noqa: SLF001 + return client + except paramiko.BadAuthenticationType as exc: + auth_errors.append(exc) + if use_interactive and "keyboard-interactive" not in getattr(exc, "allowed_types", ()): + continue + except paramiko.AuthenticationException as exc: + auth_errors.append(exc) + continue + + try: + transport.close() + except Exception: + pass + if auth_errors: + raise auth_errors[-1] + raise paramiko.AuthenticationException("bastion_auth_failed") + + +def _netmiko_over_ssh_client( + ssh_client: paramiko.SSHClient, + *, + device_type: str, + host: str, + port: int, + username: str, + password: str, + enable_secret: str, + session_timeout: int | None, +) -> ConnectHandler: + """Netmiko session over an already-authenticated SSH client (bastion protocol proxy).""" + + class _PreauthConnectHandler(ConnectHandler): + def establish_connection(self, width: int = 511, height: int = 1000) -> None: + from netmiko.channel import SSHChannel + + self.remote_conn_pre = ssh_client + self.remote_conn = ssh_client.invoke_shell(term="vt100", width=width, height=height) + self.remote_conn.settimeout(self.blocking_timeout) + if self.keepalive: + chan_transport = self.remote_conn.transport + if chan_transport is not None: + chan_transport.set_keepalive(self.keepalive) + self.channel = SSHChannel(conn=self.remote_conn, encoding=self.encoding) + self.special_login_handler() + + dev = _base_connect_kwargs( + device_type=device_type, + host=host, + port=port, + username=username, + password=password, + enable_secret=enable_secret, + session_timeout=session_timeout, + ) + return _PreauthConnectHandler(**dev) + + def _base_connect_kwargs( *, device_type: str, @@ -252,16 +362,35 @@ def _connect_via_bastion(creds: dict[str, Any], *, session_timeout: int | None = composite_user = render_hop_command(str(creds.get("hop_command_template") or ""), creds) device_type = normalize_netmiko_device_type(creds["device_type"], creds["protocol"]) - hop_dev = _base_connect_kwargs( - device_type=device_type, - host=hop_host, - port=int(creds.get("hop_port") or 22), - username=composite_user, - password=hop_pass, - enable_secret=str(creds.get("enable_secret") or ""), - session_timeout=session_timeout or 180, - ) - conn = ConnectHandler(**hop_dev) + hop_port = int(creds.get("hop_port") or 22) + timeout = int(settings.ne_connect_timeout_sec or 30) + ssh_client = None + try: + ssh_client = _bastion_ssh_connect( + host=hop_host, + port=hop_port, + username=composite_user, + password=hop_pass, + timeout=timeout, + ) + conn = _netmiko_over_ssh_client( + ssh_client, + device_type=device_type, + host=hop_host, + port=hop_port, + username=composite_user, + password=hop_pass, + enable_secret=str(creds.get("enable_secret") or ""), + session_timeout=session_timeout or 180, + ) + except Exception: + if ssh_client is not None: + try: + ssh_client.close() + except Exception: + pass + raise + auth_mode = str(creds.get("hop_target_auth_mode") or "bastion_managed").strip().lower() if auth_mode == "manual": target_pass = str(creds.get("password") or "") diff --git a/tests/test_bastion_hop.py b/tests/test_bastion_hop.py index 51b0312..6ed1c82 100644 --- a/tests/test_bastion_hop.py +++ b/tests/test_bastion_hop.py @@ -66,12 +66,15 @@ class BastionConnectRoutingTests(unittest.TestCase): class BastionConnectImplTests(unittest.TestCase): - @patch("netx_api.ne_session_factory.ConnectHandler") - def test_bastion_managed_skips_secondary_auth(self, connect_handler) -> None: + @patch("netx_api.ne_session_factory._netmiko_over_ssh_client") + @patch("netx_api.ne_session_factory._bastion_ssh_connect") + def test_bastion_managed_skips_secondary_auth(self, bastion_ssh, netmiko_wrap) -> None: from netx_api.ne_session_factory import _connect_via_bastion + ssh_client = MagicMock() + bastion_ssh.return_value = ssh_client conn = MagicMock() - connect_handler.return_value = conn + netmiko_wrap.return_value = conn creds = { "hop_host": "1.1.1.1", "hop_username": "bastion-user", @@ -88,20 +91,32 @@ class BastionConnectImplTests(unittest.TestCase): "hop_vrf": "", } _connect_via_bastion(creds) - kwargs = connect_handler.call_args.kwargs - self.assertEqual(kwargs["host"], "1.1.1.1") - self.assertEqual(kwargs["username"], "bastion-user@target-user@2.2.2.2@1.1.1.1") - self.assertEqual(kwargs["password"], "vault-pass") + bastion_ssh.assert_called_once_with( + host="1.1.1.1", + port=22, + username="bastion-user@target-user@2.2.2.2@1.1.1.1", + password="vault-pass", + timeout=unittest.mock.ANY, + ) + netmiko_wrap.assert_called_once() + wrap_kwargs = netmiko_wrap.call_args.kwargs + self.assertEqual(wrap_kwargs["host"], "1.1.1.1") + self.assertEqual(wrap_kwargs["username"], "bastion-user@target-user@2.2.2.2@1.1.1.1") + self.assertEqual(wrap_kwargs["password"], "vault-pass") conn.disconnect.assert_not_called() @patch("netx_api.ne_session_factory._interactive_target_auth") @patch("netx_api.ne_session_factory._read_channel") - @patch("netx_api.ne_session_factory.ConnectHandler") - def test_bastion_manual_invokes_secondary_auth(self, connect_handler, _read, interact) -> None: + @patch("netx_api.ne_session_factory._netmiko_over_ssh_client") + @patch("netx_api.ne_session_factory._bastion_ssh_connect") + def test_bastion_manual_invokes_secondary_auth( + self, bastion_ssh, netmiko_wrap, _read, interact + ) -> None: from netx_api.ne_session_factory import _connect_via_bastion + bastion_ssh.return_value = MagicMock() conn = MagicMock() - connect_handler.return_value = conn + netmiko_wrap.return_value = conn creds = { "hop_host": "10.0.0.1", "hop_username": "bastion-user", @@ -121,5 +136,31 @@ class BastionConnectImplTests(unittest.TestCase): interact.assert_called_once_with(conn, "target-user", "target-pass") +class BastionInteractiveHandlerTests(unittest.TestCase): + def test_replies_to_vault_password_prompt(self) -> None: + from netx_api.ne_session_factory import _bastion_interactive_handler + + handler = _bastion_interactive_handler("vault-secret") + out = handler( + "Login", + "", + [("(ZTE-TSM@user@1.1.1.1@2.2.2.2) Vault Password:", False)], + ) + self.assertEqual(out, ["vault-secret"]) + + def test_replies_to_all_fields_when_prompt_unknown(self) -> None: + from netx_api.ne_session_factory import _bastion_interactive_handler + + handler = _bastion_interactive_handler("vault-secret") + out = handler("Login", "", [("Enter code:", False), ("Confirm:", False)]) + self.assertEqual(out, ["vault-secret", "vault-secret"]) + + def test_empty_prompt_list_returns_empty(self) -> None: + from netx_api.ne_session_factory import _bastion_interactive_handler + + handler = _bastion_interactive_handler("vault-secret") + self.assertEqual(handler("Login", "", []), []) + + if __name__ == "__main__": unittest.main()