From 54ca2654304a1f0e3fe9c887c36390fa816d3ddb Mon Sep 17 00:00:00 2001 From: oliver Date: Tue, 11 Aug 2026 19:22:56 +0800 Subject: [PATCH] Retry ZTE prompt wait with RETURN after login banner timeout. Co-authored-by: Cursor --- netx_api/ne_netmiko.py | 4 ++ netx_api/ne_session_connect.py | 55 +++++++++++++++++++++++++--- tests/test_ne_session_interactive.py | 35 ++++++++++++++++++ 3 files changed, 88 insertions(+), 6 deletions(-) diff --git a/netx_api/ne_netmiko.py b/netx_api/ne_netmiko.py index be2b61d..4c51c2a 100644 --- a/netx_api/ne_netmiko.py +++ b/netx_api/ne_netmiko.py @@ -33,6 +33,10 @@ def is_cisco_ios_device_type(device_type: str) -> bool: return "cisco_ios" in dt or dt in {"cisco_ios", "cisco_ios_ssh", "cisco_ios_telnet"} +def is_zte_device_type(device_type: str) -> bool: + return "zte" in str(device_type or "").strip().lower() + + def send_show_command(conn: Any, command: str, *, read_timeout: int = 120) -> str: """Send a show/display command via ``send_command`` (wait for device prompt). diff --git a/netx_api/ne_session_connect.py b/netx_api/ne_session_connect.py index eebda88..540e0ec 100644 --- a/netx_api/ne_session_connect.py +++ b/netx_api/ne_session_connect.py @@ -192,18 +192,59 @@ def _cisco_ios_collection_driver_class(base_cls: type) -> type: return _CiscoIosCollectionSession +def _zte_collection_driver_class(base_cls: type) -> type: + """ZTE ZXROS collection: after login banner, RETURN once if prompt wait times out.""" + + class _ZteCollectionSession(base_cls): # type: ignore[misc,valid-type] + def session_preparation(self) -> None: + prompt_pat = r"[>#]" + try: + self._test_channel_read(pattern=prompt_pat) + except Exception: + # Long MOTD / idle after bastion proxy — nudge then wait again. + try: + self.write_channel("\n") + except Exception: + pass + self._test_channel_read(pattern=prompt_pat) + self.set_base_prompt() + try: + self.disable_paging(command="terminal length 0", cmd_verify=False, pattern=prompt_pat) + except Exception: + try: + self.write_channel("\n") + self.disable_paging(command="terminal length 0", cmd_verify=False, pattern=prompt_pat) + except Exception: + pass + time.sleep(0.3 * self.global_delay_factor) + self.clear_buffer() + + _ZteCollectionSession.__name__ = f"ZteCollection{getattr(base_cls, '__name__', 'Netmiko')}" + return _ZteCollectionSession + + +def _collection_driver_class(device_type: str, base_cls: type) -> type: + """Vendor-specific collection session prep (non-interactive CLI / LLDP / sync).""" + from .ne_netmiko import is_cisco_ios_device_type, is_zte_device_type + + if is_cisco_ios_device_type(device_type): + return _cisco_ios_collection_driver_class(base_cls) + if is_zte_device_type(device_type): + return _zte_collection_driver_class(base_cls) + return base_cls + + def _build_netmiko_connection(dev: dict[str, Any], *, interactive: bool = False) -> ConnectHandler: """Instantiate Netmiko from connect kwargs; optional interactive skips paging cmds.""" - from .ne_netmiko import is_cisco_ios_device_type - device_type = str(dev.get("device_type") or "").strip() if interactive: base_cls = _netmiko_driver_class(device_type) return _interactive_driver_class(base_cls)(**dev) - if is_cisco_ios_device_type(device_type): - base_cls = _netmiko_driver_class(device_type) - return _cisco_ios_collection_driver_class(base_cls)(**dev) - return ConnectHandler(**dev) + raw_cls = _netmiko_driver_class(device_type) + base_cls = _collection_driver_class(device_type, raw_cls) + if base_cls is raw_cls: + return ConnectHandler(**dev) + return base_cls(**dev) def _netmiko_over_ssh_client( @@ -224,6 +265,8 @@ def _netmiko_over_ssh_client( base_cls = _netmiko_driver_class(device_type) if interactive: base_cls = _interactive_driver_class(base_cls) + else: + base_cls = _collection_driver_class(device_type, base_cls) class _PreauthSession(base_cls): # type: ignore[misc,valid-type] def establish_connection(self, width: int = 511, height: int = 1000) -> None: diff --git a/tests/test_ne_session_interactive.py b/tests/test_ne_session_interactive.py index 75c1e44..0638338 100644 --- a/tests/test_ne_session_interactive.py +++ b/tests/test_ne_session_interactive.py @@ -3,6 +3,7 @@ from __future__ import annotations import unittest from unittest.mock import MagicMock, patch +from netx_api.ne_session_connect import _collection_driver_class, _zte_collection_driver_class from netx_api.ne_session_factory import ( _build_netmiko_connection, _interactive_driver_class, @@ -49,5 +50,39 @@ class InteractiveNetmikoTests(unittest.TestCase): self.assertTrue(direct.call_args.kwargs.get("interactive")) +class ZteCollectionPrepTests(unittest.TestCase): + def test_collection_driver_selects_zte_wrapper(self) -> None: + base = _netmiko_driver_class("zte_zxros_ssh") + wrapped = _collection_driver_class("zte_zxros_ssh", base) + self.assertIsNot(wrapped, base) + self.assertTrue(wrapped.__name__.startswith("ZteCollection")) + + def test_prompt_timeout_sends_return_then_retries(self) -> None: + base = _netmiko_driver_class("zte_zxros_ssh") + cls = _zte_collection_driver_class(base) + sess = cls.__new__(cls) + calls: list[str] = [] + + def _test_channel_read(*, pattern: str): + calls.append(f"read:{pattern}") + if calls.count(f"read:{pattern}") == 1: + raise TimeoutError("ReadTimeout: Pattern not detected: '[>#]'") + + sess._test_channel_read = _test_channel_read # type: ignore[method-assign] + sess.write_channel = MagicMock(side_effect=lambda data: calls.append(f"write:{data!r}")) # type: ignore[method-assign] + sess.set_base_prompt = MagicMock() # type: ignore[method-assign] + sess.disable_paging = MagicMock(return_value="") # type: ignore[method-assign] + sess.clear_buffer = MagicMock() # type: ignore[method-assign] + sess.global_delay_factor = 0.01 + + sess.session_preparation() + + self.assertEqual(calls[0], "read:[>#]") + self.assertEqual(calls[1], "write:'\\n'") + self.assertEqual(calls[2], "read:[>#]") + sess.disable_paging.assert_called_once() + sess.set_base_prompt.assert_called_once() + + if __name__ == "__main__": unittest.main()