mirror of
https://github.com/hansjone/netx.git
synced 2026-10-09 00:43:17 +08:00
Retry ZTE prompt wait with RETURN after login banner timeout.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
b4514c13aa
commit
54ca265430
3 changed files with 88 additions and 6 deletions
|
|
@ -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).
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue