diff --git a/netx_api/webcrt_channel.py b/netx_api/webcrt_channel.py index 1503997..29f307f 100644 --- a/netx_api/webcrt_channel.py +++ b/netx_api/webcrt_channel.py @@ -623,6 +623,139 @@ def extract_last_prompt_command(text: str) -> str | None: return None +def _command_tail(line: str) -> str: + s = normalize_audit_line(line) + for pat in ( + r"^[\w.-]+(?:\([^)]+\))*[#>]\s*(.*)$", + r"^<[^>]+>\s*(.*)$", + r"^\[[^\]]+\]\s*(.*)$", + ): + m = re.match(pat, s, flags=re.I) + if m: + return m.group(1).strip() + return s.strip() + + +def _attach_prompt_prefix(typed: str, hint: str) -> str | None: + cmd = str(typed or "").strip() + if not cmd: + return None + h = normalize_audit_line(hint) + if not _has_cli_prompt_prefix(h): + return None + m = re.match(r"^([\w.-]+(?:\([^)]+\))*[#>])\s*", h, flags=re.I) + if m: + return f"{m.group(1)}{cmd}" + m = re.match(r"^(<[^>]+>)\s*", h) + if m: + return f"{m.group(1)}{cmd}" + m = re.match(r"^(\[[^\]]+\])\s*", h) + if m: + return f"{m.group(1)}{cmd}" + return None + + +def pick_audit_command( + stdin_line: str, + audit_hint: str | None, + *, + prompt_hint: str = "", + stdout_tail: str = "", + source: str = "stdin", +) -> str | None: + """Pick auditable text for one completed stdin line (actual send + optional xterm hint).""" + typed = normalize_audit_line(stdin_line).strip() + hint = normalize_audit_line(audit_hint) if audit_hint else "" + src = str(source or "stdin") + + if typed and _is_device_output_line(typed): + return None + + if src == "post_login" and typed: + return typed + + if src == "early_stdin" and typed and not _is_prompt_only_line(typed): + return typed + + if hint and is_auditable_command_line(hint): + if not typed: + return hint + hint_cmd = _command_tail(hint) or hint + if typed == hint_cmd or hint_cmd.startswith(typed): + return hint + # Material disagreement — record bytes actually sent, not xterm hint. + return typed if typed else hint + + if typed and is_auditable_command_line(typed): + return typed + + for prefix_src in (hint, prompt_hint): + enriched = _attach_prompt_prefix(typed, prefix_src) + if enriched and is_auditable_command_line(enriched): + return enriched + + if stdout_tail: + ext = extract_last_prompt_command(stdout_tail) + if ext and is_auditable_command_line(ext): + ext_cmd = _command_tail(ext) or ext + if typed and (typed in ext_cmd or ext_cmd.endswith(typed)): + return ext + + if src == "stdin" and typed and not _is_prompt_only_line(typed): + return typed + + return None + + +def resolve_audit_commands( + buf_lines: list[str], + *, + audit_line: str | None = None, + audit_lines: list[str] | None = None, + prompt_hint: str = "", + stdout_tail: str = "", + source: str = "stdin", +) -> list[str]: + """Map all completed stdin lines in one flush to auditable command strings.""" + src = str(source or "stdin") + if src == "prompt_sync": + return [] + + if not buf_lines: + if audit_line and is_auditable_command_line(normalize_audit_line(audit_line)): + return [normalize_audit_line(audit_line)] + return [] + + hints: list[str | None] = [None] * len(buf_lines) + merged: list[str] = [str(x).strip() for x in (audit_lines or []) if str(x).strip()] + if audit_line and str(audit_line).strip(): + if not merged: + merged = [str(audit_line).strip()] + elif merged[-1] != str(audit_line).strip(): + merged.append(str(audit_line).strip()) + + if merged: + if len(merged) == len(buf_lines): + hints = list(merged) + else: + start = max(0, len(buf_lines) - len(merged)) + for j, h in enumerate(merged): + hints[start + j] = h + + out: list[str] = [] + for i, typed in enumerate(buf_lines): + cmd = pick_audit_command( + typed, + hints[i], + prompt_hint=prompt_hint, + stdout_tail=stdout_tail if i == len(buf_lines) - 1 else "", + source=src, + ) + if cmd: + out.append(cmd[:512]) + return out + + def feed_command_line_buffer(buf: str, data: str, *, max_line: int = 512) -> tuple[str, list[str]]: """Accumulate stdin into completed command lines (Enter / CR / LF). diff --git a/netx_api/webcrt_router.py b/netx_api/webcrt_router.py index 089053c..9308d1f 100644 --- a/netx_api/webcrt_router.py +++ b/netx_api/webcrt_router.py @@ -464,7 +464,8 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: loop = asyncio.get_running_loop() budget = max(30, int(settings.webcrt_connect_timeout_sec or 90)) + 15 deadline = time.time() + budget - early_stdin: list[str] = [] + early_stdin_parts: list[str] = [] + early_audit_lines: list[str] = [] async def _flush_connect_echo(cur_sess: Any) -> None: try: @@ -484,7 +485,7 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: async def _drain_client_during_connect() -> bool: """Handle ping/stdin/resize while connect runs. False = client gone.""" - nonlocal early_stdin + nonlocal early_stdin_parts, early_audit_lines while True: try: msg_raw = await asyncio.wait_for(websocket.receive(), timeout=0.01) @@ -499,7 +500,7 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: # Binary frames during connect: treat as stdin if decodeable. if "bytes" in msg_raw and msg_raw["bytes"] is not None: try: - early_stdin.append( + early_stdin_parts.append( _decode_bytes(bytes(msg_raw["bytes"]), sess.encoding) ) except Exception: @@ -508,7 +509,7 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: try: msg = json.loads(raw) except json.JSONDecodeError: - early_stdin.append(str(raw)) + early_stdin_parts.append(str(raw)) continue mtype = str(msg.get("type") or "").strip().lower() if mtype == "ping": @@ -519,7 +520,10 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: elif mtype == "stdin": data = msg.get("data") if data is not None: - early_stdin.append(str(data)) + early_stdin_parts.append(str(data)) + audit_raw = msg.get("audit_line") + if audit_raw is not None and str(audit_raw).strip(): + early_audit_lines.append(str(audit_raw).strip()[:512]) elif mtype == "resize": # Ignore until ready (PTY size already set from create). pass @@ -599,12 +603,15 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: return # Keystrokes typed while login was already on screen (before ready). - if early_stdin and sess is not None and not sess.closed and sess.conn is not None: + if early_stdin_parts and sess is not None and not sess.closed and sess.conn is not None: try: await loop.run_in_executor( webcrt_io_executor(), - sess.write_stdin, - "".join(early_stdin), + lambda: sess.write_stdin( + "".join(early_stdin_parts), + audit_source="early_stdin", + audit_lines=early_audit_lines or None, + ), ) except Exception: _log.debug( @@ -693,21 +700,25 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: stop = asyncio.Event() stdin_buf: list[str] = [] - stdin_audit_line: str | None = None + stdin_audit_lines: list[str] = [] stdin_flush_task: asyncio.Task[None] | None = None async def flush_stdin() -> None: - nonlocal stdin_buf, stdin_audit_line + nonlocal stdin_buf, stdin_audit_lines if not stdin_buf: return data = "".join(stdin_buf) - audit_line = stdin_audit_line + audit_lines = list(stdin_audit_lines) stdin_buf = [] - stdin_audit_line = None + stdin_audit_lines = [] try: await asyncio.get_running_loop().run_in_executor( webcrt_io_executor(), - lambda: sess.write_stdin(data, audit_line=audit_line), + lambda: sess.write_stdin( + data, + audit_lines=audit_lines or None, + audit_line=audit_lines[-1] if audit_lines else None, + ), ) except Exception as exc: await websocket.send_json( @@ -836,7 +847,7 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None: continue audit_raw = msg.get("audit_line") if audit_raw is not None and str(audit_raw).strip(): - stdin_audit_line = str(audit_raw).strip()[:512] + stdin_audit_lines.append(str(audit_raw).strip()[:512]) stdin_buf.append(str(data)) # Coalesce high-frequency keystrokes briefly. if len(stdin_buf) >= 8: diff --git a/netx_api/webcrt_session_model.py b/netx_api/webcrt_session_model.py index aeb9bd5..af5a25e 100644 --- a/netx_api/webcrt_session_model.py +++ b/netx_api/webcrt_session_model.py @@ -326,13 +326,21 @@ class WebcrtSession: return "stale" return chunk # bytes | None - def write_stdin(self, data: str, *, audit_source: str = "stdin", audit_line: str | None = None) -> None: + def write_stdin( + self, + data: str, + *, + audit_source: str = "stdin", + audit_line: str | None = None, + audit_lines: list[str] | None = None, + ) -> None: if self.closed or self.conn is None: raise RuntimeError("session_closed") text = str(data or "") if not text: return audit_override = str(audit_line).strip() if audit_line is not None else None + audit_list = [str(x).strip() for x in (audit_lines or []) if str(x).strip()] or None if self.cli_keymap: text = map_network_cli_keys( text, @@ -343,7 +351,12 @@ class WebcrtSession: text = map_network_cli_enter(text, self.conn) if not text: return - self._note_stdin_for_audit(text, source=audit_source, audit_line=audit_override) + self._note_stdin_for_audit( + text, + source=audit_source, + audit_line=audit_override, + audit_lines=audit_list, + ) with self._write_lock: # Prefer raw channel I/O for interactive typing (char echo / backspace). channel = getattr(self.conn, "remote_conn", None) @@ -377,34 +390,28 @@ class WebcrtSession: *, source: str = "stdin", audit_line: str | None = None, + audit_lines: list[str] | None = None, ) -> None: """Extract completed command lines from stdin and emit webcrt.command audits.""" with self._cmd_buf_lock: self._cmd_buf, buf_lines = feed_command_line_buffer(self._cmd_buf, text) redacted = bool(self._password_mode) - if redacted and (buf_lines or audit_line): + if redacted and (buf_lines or audit_line or audit_lines): self._password_mode = False if "\r" in text or "\n" in text: - from .webcrt_channel import is_auditable_command_line, normalize_audit_line + from .webcrt_channel import resolve_audit_commands - src = str(source or "stdin") - cmd: str | None = None - if redacted and (audit_line or buf_lines): - lines = ["***"] - elif src == "prompt_sync": - lines = [] - elif audit_line and str(audit_line).strip(): - candidate = normalize_audit_line(audit_line) - if is_auditable_command_line(candidate): - cmd = candidate - lines = [cmd] if cmd else [] - elif src == "post_login": - if buf_lines and str(buf_lines[-1]).strip(): - cmd = str(buf_lines[-1]).strip() - lines = [cmd] if cmd else [] + if redacted and (buf_lines or audit_line or audit_lines): + lines = ["***"] * max(1, len(buf_lines)) else: - # Interactive WebCRT: only trust frontend audit_line at Enter. - lines = [] + lines = resolve_audit_commands( + buf_lines, + audit_line=audit_line, + audit_lines=audit_lines, + prompt_hint=self._last_prompt_line, + stdout_tail=self._stdout_tail, + source=str(source or "stdin"), + ) self._last_prompt_line = "" else: lines = [] diff --git a/tests/test_webcrt_audit.py b/tests/test_webcrt_audit.py index ddc0ae8..9c2a866 100644 --- a/tests/test_webcrt_audit.py +++ b/tests/test_webcrt_audit.py @@ -16,6 +16,8 @@ from netx_api.webcrt_channel import ( is_auditable_command_line, _is_device_output_line, _is_prompt_only_line, + resolve_audit_commands, + pick_audit_command, ) from netx_api.webcrt_session_model import WebcrtSession @@ -141,6 +143,48 @@ class AuditableCommandLineTests(unittest.TestCase): def test_plain_stdin_without_prompt_not_auditable(self) -> None: self.assertFalse(is_auditable_command_line("display version")) + out = resolve_audit_commands( + ["display version"], + prompt_hint="6150#", + source="stdin", + ) + self.assertEqual(out, ["6150#display version"]) + + def test_plain_stdin_records_actual_without_hint(self) -> None: + out = resolve_audit_commands(["display version"], source="stdin") + self.assertEqual(out, ["display version"]) + + +class ResolveAuditCommandsTests(unittest.TestCase): + def test_merged_flush_audits_all_buf_lines(self) -> None: + """Interleave flush: multiple Enter in one write_stdin must not drop middle commands.""" + out = resolve_audit_commands( + ["show version", "show ll n b", "show intf"], + audit_lines=["AL5458#show intf"], + prompt_hint="AL5458#show version", + source="stdin", + ) + self.assertEqual(len(out), 3) + self.assertIn("show version", out[0]) + self.assertIn("show ll n b", out[1]) + self.assertIn("show intf", out[2]) + + def test_hint_disagreement_records_actual_stdin(self) -> None: + cmd = pick_audit_command( + "how ll n b", + "AL5458#show ll n b", + prompt_hint="AL5458#", + source="stdin", + ) + self.assertEqual(cmd, "how ll n b") + + def test_early_stdin_plain_command(self) -> None: + out = resolve_audit_commands( + ["show version"], + source="early_stdin", + prompt_hint="AL5458#", + ) + self.assertEqual(out, ["show version"]) class PasswordPromptTests(unittest.TestCase): @@ -339,6 +383,36 @@ class SessionCommandAuditTests(unittest.TestCase): sess.write_stdin("\r", audit_source="prompt_sync") mock_audit.assert_not_called() + @patch("netx_api.webcrt_session_model._audit") + def test_merged_write_stdin_audits_every_command(self, mock_audit: MagicMock) -> None: + conn = MagicMock() + conn.RETURN = "\n" + conn.remote_conn = MagicMock(spec=["recv_ready", "recv", "exit_status_ready", "resize_pty"]) + del conn.remote_conn.send + conn.write_channel = MagicMock() + + sess = WebcrtSession( + session_id="s-merge", + ne_id="ne1", + ne_name="lab", + ne_ip="1.2.3.4", + protocol="ssh", + cols=80, + rows=24, + cli_keymap=False, + conn=conn, + ) + sess._last_prompt_line = "AL5458#" + sess.write_stdin( + "show version\rshow ll n b\rshow intf\r", + audit_lines=["AL5458#show intf"], + ) + self.assertEqual(mock_audit.call_count, 3) + recorded = [c.kwargs["command"] for c in mock_audit.call_args_list] + self.assertTrue(any("show version" in x for x in recorded)) + self.assertTrue(any("show ll n b" in x for x in recorded)) + self.assertTrue(any("show intf" in x for x in recorded)) + @patch("netx_api.webcrt_session_model._audit") def test_device_error_audit_line_rejected(self, mock_audit: MagicMock) -> None: conn = MagicMock()