mirror of
https://github.com/hansjone/netx.git
synced 2026-10-09 04:20:45 +08:00
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:
parent
6ed08e24de
commit
6aedd2c522
3 changed files with 43 additions and 48 deletions
|
|
@ -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":
|
||||||
|
|
|
||||||
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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:"),
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue