Retry ZTE prompt wait with RETURN after login banner timeout.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-08-11 19:22:56 +08:00
parent b4514c13aa
commit 54ca265430
3 changed files with 88 additions and 6 deletions

View file

@ -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()