mirror of
https://github.com/hansjone/netx.git
synced 2026-10-09 06:40: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)
|
||||
|
||||
|
||||
class _ProgressSessionLog:
|
||||
"""Tee Netmiko ``session_log`` writes into a progress callback + BytesIO."""
|
||||
class _ProgressBytesIO(io.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:
|
||||
self._buf = io.BytesIO()
|
||||
super().__init__()
|
||||
self._progress_cb = progress_cb
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def write(self, data: Any) -> int:
|
||||
if isinstance(data, bytes):
|
||||
raw = data
|
||||
text = data.decode("utf-8", errors="replace")
|
||||
else:
|
||||
text = str(data or "")
|
||||
raw = text.encode("utf-8", errors="replace")
|
||||
with self._lock:
|
||||
self._buf.write(raw)
|
||||
_emit_progress(self._progress_cb, text)
|
||||
return len(raw)
|
||||
def write(self, b: Any) -> int: # noqa: ANN401
|
||||
raw = b if isinstance(b, (bytes, bytearray)) else str(b or "").encode("utf-8", errors="replace")
|
||||
if raw:
|
||||
try:
|
||||
_emit_progress(self._progress_cb, bytes(raw).decode("utf-8", errors="replace"))
|
||||
except Exception:
|
||||
pass
|
||||
return super().write(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:
|
||||
return None
|
||||
self._buf.flush()
|
||||
|
||||
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:
|
||||
|
|
@ -979,35 +988,9 @@ def open_netmiko_connection(
|
|||
ka = keepalive
|
||||
if ka is None and interactive:
|
||||
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
|
||||
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"):
|
||||
vendor = _hop_vendor(creds)
|
||||
if vendor == "linux":
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from .ne_session_factory import (
|
|||
get_cli_hop_guard,
|
||||
open_netmiko_connection,
|
||||
)
|
||||
from .ne_session_connect import _ProgressBytesIO
|
||||
from .webcrt_channel import (
|
||||
_audit,
|
||||
_capture_raw_channel,
|
||||
|
|
@ -222,11 +223,12 @@ def _finish_connect(
|
|||
connect_timeout: int,
|
||||
client: str,
|
||||
) -> None:
|
||||
log_buf = io.BytesIO()
|
||||
|
||||
def _progress(text: str) -> None:
|
||||
sess.push_connect_echo(text)
|
||||
|
||||
# Must be BufferedIOBase subclass — Netmiko rejects custom log wrappers.
|
||||
log_buf = _ProgressBytesIO(_progress)
|
||||
|
||||
try:
|
||||
conn = open_netmiko_connection(
|
||||
creds,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue