Fix WebCRT command audit loss on merged stdin flush and early input.

Audit every completed stdin line per flush, keep a list of audit_line hints from the router, restore prompt enrichment and stdout fallback, and record actual bytes when hints disagree with what was sent.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-09-02 15:58:55 +08:00
parent 302ebfac23
commit fcef3b59ac
4 changed files with 260 additions and 35 deletions

View file

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

View file

@ -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:

View file

@ -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 = []