Fix WebCRT connect crash from invalid Netmiko session_log wrapper.

Use a BytesIO subclass for live connect echo so Netmiko accepts the log object again.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-08-28 21:55:20 +08:00
parent 6ed08e24de
commit 6aedd2c522
3 changed files with 43 additions and 48 deletions

View file

@ -184,32 +184,41 @@ def _emit_progress(progress_cb: Any, text: str) -> None:
_log.debug("connect progress_cb failed", exc_info=True) _log.debug("connect progress_cb failed", exc_info=True)
class _ProgressSessionLog: class _ProgressBytesIO(io.BytesIO):
"""Tee Netmiko ``session_log`` writes into a progress callback + BytesIO.""" """BytesIO session_log that also tees reads into a connect-progress callback.
Netmiko only accepts ``io.BufferedIOBase`` (or path / SessionLog). A plain
custom log object raises ``ValueError`` and breaks every WebCRT connect.
"""
def __init__(self, progress_cb: Any = None) -> None: def __init__(self, progress_cb: Any = None) -> None:
self._buf = io.BytesIO() super().__init__()
self._progress_cb = progress_cb self._progress_cb = progress_cb
self._lock = threading.Lock()
def write(self, data: Any) -> int: def write(self, b: Any) -> int: # noqa: ANN401
if isinstance(data, bytes): raw = b if isinstance(b, (bytes, bytearray)) else str(b or "").encode("utf-8", errors="replace")
raw = data if raw:
text = data.decode("utf-8", errors="replace") try:
else: _emit_progress(self._progress_cb, bytes(raw).decode("utf-8", errors="replace"))
text = str(data or "") except Exception:
raw = text.encode("utf-8", errors="replace") pass
with self._lock: return super().write(raw)
self._buf.write(raw)
_emit_progress(self._progress_cb, text)
return len(raw) class _ProgressSessionLog:
"""Deprecated alias kept for imports; prefer ``_ProgressBytesIO``."""
def __init__(self, progress_cb: Any = None) -> None:
self._buf = _ProgressBytesIO(progress_cb)
def write(self, data: Any) -> int: # noqa: ANN401
return self._buf.write(data)
def flush(self) -> None: def flush(self) -> None:
return None self._buf.flush()
def getvalue(self) -> bytes: def getvalue(self) -> bytes:
with self._lock: return self._buf.getvalue()
return self._buf.getvalue()
def _cisco_ios_collection_driver_class(base_cls: type) -> type: def _cisco_ios_collection_driver_class(base_cls: type) -> type:
@ -979,35 +988,9 @@ def open_netmiko_connection(
ka = keepalive ka = keepalive
if ka is None and interactive: if ka is None and interactive:
ka = int(getattr(settings, "webcrt_keepalive_sec", 0) or 0) or None ka = int(getattr(settings, "webcrt_keepalive_sec", 0) or 0) or None
# Never wrap session_log in a non-BufferedIOBase object — Netmiko rejects it
# with ValueError and every WebCRT session fails to open.
log = session_log log = session_log
if progress_cb is not None:
class _TeeLog(_ProgressSessionLog):
def write(self, data: Any) -> int: # noqa: ANN401
n = super().write(data)
if session_log is None:
return n
try:
if isinstance(data, (bytes, bytearray)):
session_log.write(data)
else:
session_log.write(str(data).encode("utf-8", errors="replace"))
except Exception:
try:
session_log.write(data)
except Exception:
pass
return n
def getvalue(self) -> bytes:
if session_log is not None and hasattr(session_log, "getvalue"):
try:
raw = session_log.getvalue()
return raw if isinstance(raw, (bytes, bytearray)) else str(raw).encode()
except Exception:
pass
return super().getvalue()
log = _TeeLog(progress_cb)
if creds.get("hop_enabled"): if creds.get("hop_enabled"):
vendor = _hop_vendor(creds) vendor = _hop_vendor(creds)
if vendor == "linux": if vendor == "linux":

View file

@ -19,6 +19,7 @@ from .ne_session_factory import (
get_cli_hop_guard, get_cli_hop_guard,
open_netmiko_connection, open_netmiko_connection,
) )
from .ne_session_connect import _ProgressBytesIO
from .webcrt_channel import ( from .webcrt_channel import (
_audit, _audit,
_capture_raw_channel, _capture_raw_channel,
@ -222,11 +223,12 @@ def _finish_connect(
connect_timeout: int, connect_timeout: int,
client: str, client: str,
) -> None: ) -> None:
log_buf = io.BytesIO()
def _progress(text: str) -> None: def _progress(text: str) -> None:
sess.push_connect_echo(text) sess.push_connect_echo(text)
# Must be BufferedIOBase subclass — Netmiko rejects custom log wrappers.
log_buf = _ProgressBytesIO(_progress)
try: try:
conn = open_netmiko_connection( conn = open_netmiko_connection(
creds, creds,

View file

@ -14,6 +14,16 @@ from netx_api.ne_session_connect import (
class PromptDetectTests(unittest.TestCase): class PromptDetectTests(unittest.TestCase):
def test_progress_bytesio_is_netmiko_compatible(self) -> None:
import io
from netx_api.ne_session_connect import _ProgressBytesIO
buf = _ProgressBytesIO(lambda _t: None)
self.assertIsInstance(buf, io.BufferedIOBase)
# Plain custom log objects are what broke every WebCRT connect (ValueError).
self.assertFalse(isinstance(object(), io.BufferedIOBase))
def test_huawei_username_and_password(self) -> None: def test_huawei_username_and_password(self) -> None:
self.assertEqual( self.assertEqual(
_prompt_needs_auth("Please input the username:"), _prompt_needs_auth("Please input the username:"),