diff --git a/headroom/proxy/interceptors/astgrep.py b/headroom/proxy/interceptors/astgrep.py index 460192892..35ff5836e 100644 --- a/headroom/proxy/interceptors/astgrep.py +++ b/headroom/proxy/interceptors/astgrep.py @@ -13,15 +13,19 @@ from __future__ import annotations import json import logging import os +import re import shutil +import stat import subprocess import tempfile +from enum import Enum from pathlib import Path from typing import Any from headroom import binaries from headroom._subprocess import run from headroom.proxy import runtime_env +from headroom.proxy.project_context import get_current_cwd, is_current_request_trusted from . import base @@ -88,6 +92,177 @@ _PATTERNS: dict[str, list[str]] = { OUTLINE_MARKER = " # ... (body elided by Headroom; Read a specific line range to see it)\n" +# Per-client banner signatures -- add an entry only once a client's exact +# banner text is confirmed, never a generic keyword match. +_TRUNCATION_SIGNATURES: tuple[re.Pattern[str], ...] = ( + # Claude Code: "[Truncated: PARTIAL view -- : showing lines A-B of + # T total (...). Call Read with offset=N to see more.]" + re.compile( + r"\[\s*truncated\s*:\s*partial\s+view\b" + r"[^\[\]]*?" + r"showing\s+lines?\s+(?P\d+)\s*[-–]\s*(?P\d+)" + r"\s+of\s+(?P\d+)\s+total" + r"[^\[\]]*\]", + re.IGNORECASE, + ), +) + + +def _is_plausible_truncation_range( + start_line: int, end_line: int, total_lines: int, source_line_count: int +) -> bool: + # end_line == total_lines means the whole file was shown, not truncated. + if start_line < 1 or end_line < start_line or end_line >= total_lines: + return False + return end_line <= source_line_count + + +def _detect_truncation(source: str) -> tuple[int, int] | None: + """Return (end_line, total_lines) if `source` carries a recognized, + internally-consistent upstream truncation banner, else None.""" + source_line_count = len(source.splitlines()) + for pattern in _TRUNCATION_SIGNATURES: + for m in pattern.finditer(source): + start_line = int(m.group("start_line")) + end_line = int(m.group("end_line")) + total_lines = int(m.group("total_lines")) + if _is_plausible_truncation_range(start_line, end_line, total_lines, source_line_count): + return end_line, total_lines + return None + + +class ReadVerificationResult(Enum): + """Client-independent fallback for `_detect_truncation`'s banner regex: + compares tool_output against the real file on disk instead of parsing + client-specific prose. Only used when the banner regex finds nothing.""" + + COMPLETE = "complete" + TRUNCATED = "truncated" + # Unresolvable path, unreadable file, or mismatched content — never guess. + UNKNOWN = "unknown" + + +def _verify_truncation_on_disk_enabled() -> bool: + # Live read (not a module constant), matching _min_chars_to_rewrite()'s + # hot-reload behavior. + return runtime_env.getenv("HEADROOM_VERIFY_TRUNCATION_ON_DISK", "").lower() in ( + "1", + "true", + "yes", + ) + + +def _max_disk_verify_bytes() -> int: + # Live read, same hot-reload pattern as _min_chars_to_rewrite(). 5 MB + # comfortably covers real source files while bounding worst-case read + # time/memory for the disk-verify fallback. + try: + return int(runtime_env.getenv("HEADROOM_VERIFY_TRUNCATION_MAX_BYTES", "5000000")) + except (TypeError, ValueError): + return 5_000_000 + + +def _resolve_read_path_in_workspace(file_path: str, resolved_root: Path) -> Path | None: + """Confine `file_path` (relative or absolute) strictly under `resolved_root`. + + Mirrors memory_handler._resolve_native_path's join/resolve/relative_to + pattern. `resolved_root` must already be canonicalized by the caller (a + single source of truth for "inside the workspace" across a call) -- + this does not re-resolve it. `.resolve()` collapses symlinks, including + in intermediate path components, before the containment check runs, so + a symlink that points outside the workspace is rejected the same way a + `..` traversal is -- this is a path-segment containment check via + `relative_to()`, not a string-prefix check, so a sibling directory that + merely shares a string prefix with the root is correctly rejected too. + """ + candidate = Path(file_path) if Path(file_path).is_absolute() else resolved_root / file_path + try: + resolved = candidate.resolve(strict=True) + except OSError: + return None + try: + resolved.relative_to(resolved_root) + except ValueError: + return None + return resolved + + +def _read_disk_content_bounded(resolved: Path, max_bytes: int) -> str | None: + """Read `resolved` iff it's a regular file no larger than `max_bytes`. + + Opens with O_NONBLOCK so a FIFO/special file with no writer returns + immediately instead of blocking the calling thread indefinitely -- + O_NONBLOCK has no effect on reads once fstat confirms a regular file. + fstat runs on the already-open fd (not the path) so the type/size + check and the read happen on the same underlying file object, closing + the TOCTOU gap a separate path-based stat-then-open would leave. + """ + try: + fd = os.open(resolved, os.O_RDONLY | os.O_NONBLOCK) + except OSError: + return None + try: + st = os.fstat(fd) + if not stat.S_ISREG(st.st_mode) or st.st_size > max_bytes: + return None + with os.fdopen(fd, "r", encoding="utf-8") as f: + fd = -1 # ownership transferred to the file object + return f.read() + except (OSError, UnicodeDecodeError): + return None + finally: + if fd >= 0: + os.close(fd) + + +def _verify_read_against_disk( + file_path: str | None, + received_content: str, + cwd: str | None, + *, + trusted: bool, +) -> tuple[ReadVerificationResult, tuple[int, int] | None]: + """Compare `received_content` against the real file at `file_path`. + + `cwd` (the `x-headroom-cwd` header) is never authority on its own -- + `trusted` (whether the request's peer is loopback, a server-observed + fact, not a header) must hold before `cwd` or `file_path` are even + inspected, let alone touch disk. Once trusted, every read must be a + regular file whose fully resolved path is beneath the canonically + resolved `cwd`; anything else is UNKNOWN, never a guess. + + TRUNCATED requires an exact-prefix match with strictly more on disk — + a weaker match means the file diverged since the client read it, not + a provable truncation, so it's UNKNOWN. Returns `(visible_lines, + total_lines)` alongside TRUNCATED so the header can cite real numbers + without a second, potentially racy, read. + """ + if not trusted: + return ReadVerificationResult.UNKNOWN, None + if not file_path: + return ReadVerificationResult.UNKNOWN, None + if not cwd or not os.path.isabs(cwd): + return ReadVerificationResult.UNKNOWN, None + try: + resolved_root = Path(cwd).resolve(strict=True) + except OSError: + return ReadVerificationResult.UNKNOWN, None + if not resolved_root.is_dir(): + return ReadVerificationResult.UNKNOWN, None + resolved = _resolve_read_path_in_workspace(file_path, resolved_root) + if resolved is None: + return ReadVerificationResult.UNKNOWN, None + disk_content = _read_disk_content_bounded(resolved, _max_disk_verify_bytes()) + if disk_content is None: + return ReadVerificationResult.UNKNOWN, None + if disk_content == received_content: + return ReadVerificationResult.COMPLETE, None + if disk_content.startswith(received_content) and len(disk_content) > len(received_content): + visible_lines = len(received_content.splitlines()) + total_lines = len(disk_content.splitlines()) + return ReadVerificationResult.TRUNCATED, (visible_lines, total_lines) + return ReadVerificationResult.UNKNOWN, None + class AstGrepReadOutline: """Interceptor that outlines verbose code-file Read outputs.""" @@ -131,7 +306,18 @@ class AstGrepReadOutline: if not matches: return None - outline = _build_outline(matches, tool_output) + # Banner (cheap, no I/O) wins; disk verification is the opt-in fallback. + truncation = _detect_truncation(tool_output) + if truncation is None and _verify_truncation_on_disk_enabled(): + verdict, disk_truncation = _verify_read_against_disk( + _path_from_input(tool_input), + tool_output, + get_current_cwd(), + trusted=is_current_request_trusted(), + ) + if verdict is ReadVerificationResult.TRUNCATED: + truncation = disk_truncation + outline = _build_outline(matches, tool_output, truncation) return outline if outline else None def progressive_disclosure_key( @@ -258,12 +444,21 @@ def _run_ast_grep( return all_matches -def _build_outline(matches: list[dict[str, Any]], source: str) -> str | None: +def _build_outline( + matches: list[dict[str, Any]], + source: str, + truncation: tuple[int, int] | None = None, +) -> str | None: """Build a compact outline from ast-grep matches. Emits each definition's signature line + docstring (if next line is a string literal) + an elision marker. Matches are sorted by byte offset so the outline tracks the original file order. + + `truncation`, if given, is (end_line, total_lines) from an upstream + truncation banner already present in `source` (e.g. a client's own Read + token-cap notice). When set, the header states that the input was a + partial view instead of implying `source` is the whole file. """ lines = source.splitlines(keepends=True) outline_chunks: list[str] = [] @@ -292,11 +487,21 @@ def _build_outline(matches: list[dict[str, Any]], source: str) -> str | None: if not outline_chunks: return None - header = ( - "[headroom: outlined by ast-grep — " - f"{len(seen_starts)} definition(s); " - "bodies elided. Re-read the file with a line range to see a specific body.]\n" - ) + + if truncation: + end_line, total_lines = truncation + header = ( + "[headroom: outlined by ast-grep — " + f"{len(seen_starts)} definition(s) in the visible portion; " + f"input was truncated upstream (showing through line {end_line} of {total_lines} total). " + "Bodies elided. Re-read remaining lines to see more.]\n" + ) + else: + header = ( + "[headroom: outlined by ast-grep — " + f"{len(seen_starts)} definition(s); " + "bodies elided. Re-read the file with a line range to see a specific body.]\n" + ) return header + "".join(outline_chunks) diff --git a/headroom/proxy/project_context.py b/headroom/proxy/project_context.py index d6a446c68..a1d003834 100644 --- a/headroom/proxy/project_context.py +++ b/headroom/proxy/project_context.py @@ -11,6 +11,22 @@ savings to a project without threading a parameter through every handler. The value is sanitized (printable characters only, length-capped) before it is stored; an absent or unusable header simply leaves attribution off for that request, matching pre-feature behavior. + +The HTTP middleware also binds the raw ``x-headroom-cwd`` header into a +second, unsanitized contextvar for consumers that need the literal +filesystem path (e.g. verifying a Read tool_result against disk). Not +(yet) bound at the WebSocket accept paths — an absent cwd there is already +treated as "can't resolve, don't guess." + +A third contextvar, ``_current_request_trusted``, records whether the +active request's peer is loopback — a fact the HTTP middleware observes +from the TCP connection itself, not something a caller can assert via a +header. Consumers that turn ``_current_cwd`` into a filesystem read (again, +disk verification) must gate on this first: the cwd header alone is never +sufficient authority for touching disk, only a signal to interpret once the +peer is already known to be trusted. Defaults to ``False`` so any request +path that never calls the setter (e.g. the WebSocket accept path) is safe +by construction. """ from __future__ import annotations @@ -31,6 +47,15 @@ from headroom.proxy.savings_tracker import sanitize_project_name _current_project: ContextVar[str | None] = ContextVar("headroom_current_project", default=None) +# Unsanitized, unlike _current_project — consumers join this against a +# tool's file_path and read from disk, so it must stay the literal path. +_current_cwd: ContextVar[str | None] = ContextVar("headroom_current_cwd", default=None) + +# Server-observed (not header-derived) loopback signal — see module docstring. +_current_request_trusted: ContextVar[bool] = ContextVar( + "headroom_current_request_trusted", default=False +) + def set_current_project(project: str | None) -> None: """Bind the active request's project for downstream outcome recording.""" @@ -42,6 +67,26 @@ def get_current_project() -> str | None: return _current_project.get() +def set_current_cwd(cwd: str | None) -> None: + """Bind the active request's ``x-headroom-cwd`` header value, unmodified.""" + _current_cwd.set(cwd.strip() if isinstance(cwd, str) and cwd.strip() else None) + + +def get_current_cwd() -> str | None: + """Raw cwd header bound to the current request context, or ``None``.""" + return _current_cwd.get() + + +def set_current_request_trusted(trusted: bool) -> None: + """Bind whether the active request's peer is loopback.""" + _current_request_trusted.set(trusted) + + +def is_current_request_trusted() -> bool: + """Whether the active request's peer is loopback, or ``False`` if unset.""" + return _current_request_trusted.get() + + def strip_project_path_prefix(scope: MutableMapping[str, Any]) -> str | None: """Strip a ``/p/`` prefix from an ASGI scope, returning the name. @@ -58,8 +103,12 @@ __all__ = [ "PROJECT_HEADER", "PROJECT_PATH_PREFIX", "classify_project", + "get_current_cwd", "get_current_project", + "is_current_request_trusted", + "set_current_cwd", "set_current_project", + "set_current_request_trusted", "split_project_path", "strip_project_path_prefix", "with_project_prefix", diff --git a/headroom/proxy/runtime_env.py b/headroom/proxy/runtime_env.py index bc4c03bc2..b9d255f34 100644 --- a/headroom/proxy/runtime_env.py +++ b/headroom/proxy/runtime_env.py @@ -75,6 +75,16 @@ RUNTIME_ENV_KNOBS: tuple[Knob, ...] = ( "int", "Min tool-output chars before the ast-grep read rewrite.", ), + Knob( + "HEADROOM_VERIFY_TRUNCATION_ON_DISK", + "bool", + "Verify ast-grep Read truncation against disk when no client banner is found.", + ), + Knob( + "HEADROOM_VERIFY_TRUNCATION_MAX_BYTES", + "int", + "Max on-disk file size (bytes) the disk-verify fallback will read.", + ), ) _KNOBS_BY_ENV: dict[str, Knob] = {k.env: k for k in RUNTIME_ENV_KNOBS} diff --git a/headroom/proxy/server.py b/headroom/proxy/server.py index 233a0367c..9f1db4b8d 100644 --- a/headroom/proxy/server.py +++ b/headroom/proxy/server.py @@ -163,7 +163,9 @@ from headroom.proxy.modes import ( from headroom.proxy.probe_recorder import probe_recorder_from_env from headroom.proxy.project_context import ( classify_project, + set_current_cwd, set_current_project, + set_current_request_trusted, strip_project_path_prefix, ) from headroom.proxy.prometheus_metrics import PrometheusMetrics # noqa: F401 @@ -2618,6 +2620,7 @@ class WebSocketProjectPrefixMiddleware: name.decode("latin-1"): value.decode("latin-1") for name, value in scope["headers"] } set_current_project(classify_project(headers) or prefix_project) + # No set_current_cwd() here -- HTTP-only for now, see project_context.py. await self.app(scope, receive, send) @@ -3360,7 +3363,19 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: method = request.method query = request.url.query headers = dict(request.headers.items()) + client = getattr(request, "client", None) + client_addr = "" + client_host = None + if client is not None: + client_host = getattr(client, "host", None) + client_port = getattr(client, "port", None) + client_addr = f"{client_host}:{client_port}" if client_port else str(client_host) set_current_project(classify_project(headers) or prefix_project) + set_current_cwd(headers.get("x-headroom-cwd")) + # Server-observed (not header-derived) trust signal for consumers that + # turn x-headroom-cwd into a filesystem read (e.g. astgrep disk + # verification) -- the header alone is never sufficient authority. + set_current_request_trusted(is_loopback_host(client_host)) # Path-based Codex identification: stamp X-Client: codex on the # Responses endpoint for callers that don't otherwise classify (e.g. # Codex Desktop, whose User-Agent isn't a known codex UA). Without it @@ -3370,12 +3385,6 @@ def create_app(config: ProxyConfig | None = None) -> FastAPI: # makes every downstream classify_client(headers) read "codex". if should_stamp_codex_client(path, headers): request.scope["headers"].append((b"x-client", b"codex")) - client = getattr(request, "client", None) - client_addr = "" - if client is not None: - client_host = getattr(client, "host", None) - client_port = getattr(client, "port", None) - client_addr = f"{client_host}:{client_port}" if client_port else str(client_host) try: proxy.metrics.record_inbound_request(method=method, path=path) except Exception: diff --git a/tests/test_proxy_loopback_gating.py b/tests/test_proxy_loopback_gating.py index 49381d1ef..33dd8c7a5 100644 --- a/tests/test_proxy_loopback_gating.py +++ b/tests/test_proxy_loopback_gating.py @@ -745,3 +745,44 @@ def test_dashboard_client_cidr_does_not_expand_other_management_endpoints( assert client.get("/admin/upstream").status_code == 404 assert client.get("/debug/tasks").status_code == 404 assert client.post("/stats/reset").status_code == 404 + + +def _app_with_trust_probe() -> FastAPI: + """`_record_headroom_stack` binds `is_current_request_trusted()` from the + peer's loopback-ness before any route runs; this test-only route reads + it back so the binding itself (not just a downstream consumer like + astgrep's disk-verify fallback) can be asserted directly.""" + from headroom.proxy.project_context import is_current_request_trusted + + app = _make_app() + + @app.get("/__test/trusted") + def _trusted_probe() -> dict[str, bool]: + return {"trusted": is_current_request_trusted()} + + # A catch-all passthrough route registered by create_app() would + # otherwise shadow this path -- Starlette matches routes in + # registration order, not by specificity. + app.router.routes.insert(0, app.router.routes.pop()) + return app + + +def test_record_headroom_stack_binds_trusted_true_for_loopback_client() -> None: + client = TestClient( + _app_with_trust_probe(), base_url="http://127.0.0.1", client=("127.0.0.1", 12345) + ) + resp = client.get("/__test/trusted", headers={"x-headroom-cwd": "/some/spoofed/path"}) + assert resp.json() == {"trusted": True} + + +def test_record_headroom_stack_binds_trusted_false_for_non_loopback_client() -> None: + """A spoofed x-headroom-cwd header from a non-loopback peer must not + make the request trusted -- the header carries no authority on its + own. Direct regression test for the disk-verify oracle finding.""" + client = TestClient( + _app_with_trust_probe(), + base_url="http://100.82.0.2:8787", + client=("100.90.0.5", 12345), + ) + resp = client.get("/__test/trusted", headers={"x-headroom-cwd": "/some/spoofed/path"}) + assert resp.json() == {"trusted": False} diff --git a/tests/test_tool_result_interceptors.py b/tests/test_tool_result_interceptors.py index 432c2c51b..bc61e216a 100644 --- a/tests/test_tool_result_interceptors.py +++ b/tests/test_tool_result_interceptors.py @@ -2,6 +2,8 @@ from __future__ import annotations +import os +import sys import textwrap import pytest @@ -14,8 +16,17 @@ from headroom.proxy.interceptors import ( interceptor_failure_counts, register, ) -from headroom.proxy.interceptors.astgrep import AstGrepReadOutline +from headroom.proxy.interceptors.astgrep import ( + AstGrepReadOutline, + ReadVerificationResult, + _verify_read_against_disk, +) from headroom.proxy.interceptors.base import reset_interceptor_failure_counts +from headroom.proxy.project_context import ( + get_current_cwd, + set_current_cwd, + set_current_request_trusted, +) from headroom.tokenizer import Tokenizer @@ -211,6 +222,462 @@ def test_astgrep_outlines_large_python_read(tokenizer): assert "def apply_promo" in new_content # Bodies should NOT leak through unchanged. assert "total += item.price * item.qty" not in new_content + # Complete-file control: no truncation banner in the input -> no truncation marker. + assert "truncated upstream" not in new_content + + +def test_astgrep_flags_truncated_read(tokenizer): + truncated_source = ( + _PY_FIXTURE + "\n\n[Truncated: PARTIAL view — /repo/payments.py: " + "showing lines 1-42 of 90 total (26031 tokens, cap 25000). " + "Call Read with offset=43 to see more.]\n" + ) + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "abc", + "name": "Read", + "input": {"file_path": "/repo/payments.py"}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "abc", "content": truncated_source}], + }, + ] + result = apply_to_messages(messages, tokenizer) + assert len(result.spans) == 1 + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" in new_content + assert "showing through line 42 of 90 total" in new_content + # Still lists the definitions actually present in the visible portion. + assert "def process_payment" in new_content + assert "def apply_promo" in new_content + + +def _read_result_messages(content: str, file_path: str = "/repo/payments.py"): + return [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "abc", + "name": "Read", + "input": {"file_path": file_path}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "abc", "content": content}], + }, + ] + + +def test_astgrep_flags_truncated_read_wording_variants(tokenizer): + """The signature tolerates wording/casing/dash variation, but only inside + the recognized envelope — not as a synonym match over arbitrary prose.""" + truncated_source = ( + _PY_FIXTURE + "\n\n[TRUNCATED: PARTIAL VIEW — /repo/payments.py: " + "SHOWING LINES 1–42 of 90 TOTAL. Call Read with offset=43 to see more.]\n" + ) + messages = _read_result_messages(truncated_source) + result = apply_to_messages(messages, tokenizer) + assert len(result.spans) == 1 + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" in new_content + assert "showing through line 42 of 90 total" in new_content + + +def test_astgrep_ignores_truncation_phrase_in_comment(tokenizer): + """A count-shaped phrase in a plain comment, with no bracketed envelope, + must not be read as an upstream truncation claim.""" + source_with_comment = ( + _PY_FIXTURE + "\n\n# API pagination showing lines 10-20 of 30 total records\n" + ) + messages = _read_result_messages(source_with_comment) + result = apply_to_messages(messages, tokenizer) + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" not in new_content + + +def test_astgrep_ignores_bracketed_phrase_without_signature(tokenizer): + """Brackets plus a count-shaped phrase aren't enough on their own — the + exact recognized signature phrase must also be present.""" + source_with_bracket = ( + _PY_FIXTURE + "\n\n[Truncation happened; showing lines 1-42 of 90 total]\n" + ) + messages = _read_result_messages(source_with_bracket) + result = apply_to_messages(messages, tokenizer) + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" not in new_content + + +@pytest.mark.parametrize( + "banner_numbers", + [ + pytest.param("50-90 of 90", id="end_equals_total"), + pytest.param("42-10 of 90", id="end_less_than_start"), + pytest.param("0-42 of 90", id="start_is_zero"), + pytest.param("1-5000 of 9000", id="end_exceeds_visible_payload"), + ], +) +def test_astgrep_ignores_malformed_truncation_counts(tokenizer, banner_numbers): + truncated_source = ( + _PY_FIXTURE + "\n\n[Truncated: PARTIAL view — /repo/payments.py: " + f"showing lines {banner_numbers} total (26031 tokens, cap 25000). " + "Call Read with offset=43 to see more.]\n" + ) + messages = _read_result_messages(truncated_source) + result = apply_to_messages(messages, tokenizer) + new_content = result.messages[1]["content"][0]["content"] + assert "outlined by ast-grep" in new_content + assert "truncated upstream" not in new_content + + +def test_astgrep_accepts_truncation_at_start_equals_end(tokenizer): + """`end >= start` is inclusive — a single-line visible window is valid.""" + truncated_source = ( + _PY_FIXTURE + "\n\n[Truncated: PARTIAL view — /repo/payments.py: " + "showing lines 42-42 of 90 total (26031 tokens, cap 25000). " + "Call Read with offset=43 to see more.]\n" + ) + messages = _read_result_messages(truncated_source) + result = apply_to_messages(messages, tokenizer) + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" in new_content + + +# -------- Disk verification: client-independent truncation fallback ----- # + + +def test_set_get_current_cwd_round_trips(): + set_current_cwd(" /repo/project ") + try: + assert get_current_cwd() == "/repo/project" + finally: + set_current_cwd(None) + + +def test_set_current_cwd_none_and_blank_both_clear(): + set_current_cwd("/repo") + try: + assert get_current_cwd() == "/repo" + set_current_cwd(" ") + assert get_current_cwd() is None + finally: + set_current_cwd(None) + + +class TestVerifyReadAgainstDisk: + def test_missing_file_path_is_unknown(self): + verdict, info = _verify_read_against_disk(None, "abc", "/repo", trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_relative_path_without_cwd_is_unknown(self): + verdict, info = _verify_read_against_disk("payments.py", "abc", None, trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_untrusted_request_never_touches_disk(self, tmp_path, monkeypatch): + """A fully valid cwd/file_path/content combo that would resolve + COMPLETE if trusted must short-circuit to UNKNOWN -- and never call + os.open -- when the request isn't trusted. The regression guarantee + is "untrusted -> UNKNOWN without any filesystem operation," not just + "untrusted -> UNKNOWN.\"""" + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + opened: list[object] = [] + real_open = os.open + + def _tracking_open(*args, **kwargs): + opened.append(args) + return real_open(*args, **kwargs) + + monkeypatch.setattr(os, "open", _tracking_open) + verdict, info = _verify_read_against_disk(str(f), _PY_FIXTURE, str(tmp_path), trusted=False) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + assert opened == [] + + def test_missing_file_under_workspace_root_is_unknown(self, tmp_path): + verdict, info = _verify_read_against_disk( + str(tmp_path / "missing.py"), "abc", str(tmp_path), trusted=True + ) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_exact_match_is_complete(self, tmp_path): + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + verdict, info = _verify_read_against_disk(str(f), _PY_FIXTURE, str(tmp_path), trusted=True) + assert verdict is ReadVerificationResult.COMPLETE + assert info is None + + def test_strict_prefix_is_truncated(self, tmp_path): + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + partial = _PY_FIXTURE[:200] + verdict, info = _verify_read_against_disk(str(f), partial, str(tmp_path), trusted=True) + assert verdict is ReadVerificationResult.TRUNCATED + assert info == (len(partial.splitlines()), len(_PY_FIXTURE.splitlines())) + + def test_relative_path_resolves_against_cwd(self, tmp_path): + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + partial = _PY_FIXTURE[:200] + verdict, info = _verify_read_against_disk( + "payments.py", partial, str(tmp_path), trusted=True + ) + assert verdict is ReadVerificationResult.TRUNCATED + assert info is not None + + def test_content_mismatch_is_unknown_not_truncated(self, tmp_path): + # File diverged since the client read it -- not a clean prefix. + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE.replace("compute_subtotal", "compute_total"), encoding="utf-8") + verdict, info = _verify_read_against_disk(str(f), _PY_FIXTURE, str(tmp_path), trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_absolute_path_outside_workspace_root_is_unknown(self, tmp_path): + root = tmp_path / "project" + root.mkdir() + sibling = tmp_path / "other" + sibling.mkdir() + f = sibling / "secret.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + verdict, info = _verify_read_against_disk(str(f), _PY_FIXTURE, str(root), trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_relative_traversal_escapes_workspace_is_unknown(self, tmp_path): + root = tmp_path / "project" + root.mkdir() + f = tmp_path / "secret.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + verdict, info = _verify_read_against_disk( + "../secret.py", _PY_FIXTURE, str(root), trusted=True + ) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_sibling_directory_sharing_string_prefix_is_not_inside_workspace(self, tmp_path): + """`relative_to()` is a path-segment containment check, not a string + prefix check -- a target that merely starts with the root's string + must not be treated as inside it.""" + root = tmp_path / "project" + root.mkdir() + other = tmp_path / "project-other" + other.mkdir() + f = other / "secret.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + verdict, info = _verify_read_against_disk(str(f), _PY_FIXTURE, str(root), trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + @pytest.mark.skipif( + sys.platform == "win32", reason="symlinks need elevated privilege on Windows" + ) + def test_symlinked_file_escapes_workspace_is_unknown(self, tmp_path): + root = tmp_path / "project" + root.mkdir() + outside = tmp_path / "outside.py" + outside.write_text(_PY_FIXTURE, encoding="utf-8") + link = root / "link.py" + link.symlink_to(outside) + verdict, info = _verify_read_against_disk("link.py", _PY_FIXTURE, str(root), trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_directory_passed_as_file_path_is_unknown(self, tmp_path): + # Also proves O_NONBLOCK doesn't need a FIFO fixture to matter: this + # exercises the same os.open(..., O_NONBLOCK) path, just rejected by + # the S_ISREG check rather than by not blocking on a missing writer. + root = tmp_path / "project" + subdir = root / "subdir" + subdir.mkdir(parents=True) + verdict, info = _verify_read_against_disk("subdir", "abc", str(root), trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + def test_oversized_file_exceeding_byte_cap_is_unknown(self, tmp_path, monkeypatch): + monkeypatch.setenv("HEADROOM_VERIFY_TRUNCATION_MAX_BYTES", "10") + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + assert len(_PY_FIXTURE.encode("utf-8")) > 10 + # Content matches exactly -- would be COMPLETE without the cap. + verdict, info = _verify_read_against_disk(str(f), _PY_FIXTURE, str(tmp_path), trusted=True) + assert verdict is ReadVerificationResult.UNKNOWN + assert info is None + + +def test_astgrep_disk_verification_flags_truncation_when_no_banner( + tokenizer, tmp_path, monkeypatch +): + """Opted in, cwd bound, file on disk is strictly longer than the tool_result, + and no banner is present -- disk verification alone should qualify the header.""" + monkeypatch.setenv("HEADROOM_VERIFY_TRUNCATION_ON_DISK", "1") + monkeypatch.setenv("HEADROOM_INTERCEPT_READ_MIN_CHARS", "50") + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + # Split after most functions -- a too-short partial's elision markers + # can outweigh tiny bodies and trip the "refuse to enlarge" guard. + marker = "\n\ndef format_receipt" + partial = _PY_FIXTURE[: _PY_FIXTURE.index(marker)] # no banner text anywhere + assert "truncated" not in partial.lower() + + set_current_cwd(str(tmp_path)) + set_current_request_trusted(True) + try: + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "abc", + "name": "Read", + "input": {"file_path": "payments.py"}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "abc", "content": partial}], + }, + ] + result = apply_to_messages(messages, tokenizer) + finally: + set_current_cwd(None) + set_current_request_trusted(False) + + assert len(result.spans) == 1 + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" in new_content + visible_lines = len(partial.splitlines()) + total_lines = len(_PY_FIXTURE.splitlines()) + assert f"showing through line {visible_lines} of {total_lines} total" in new_content + assert "def compute_subtotal" in new_content + assert "def apply_promo" in new_content + # Beyond the truncation point -- never reached ast-grep, can't appear. + assert "def format_receipt" not in new_content + + +def test_astgrep_disk_verification_disabled_by_default(tokenizer, tmp_path, monkeypatch): + """Same truncated-on-disk scenario as above, but without the opt-in env + var -- must behave exactly like the no-signal case (no header change).""" + monkeypatch.delenv("HEADROOM_VERIFY_TRUNCATION_ON_DISK", raising=False) + monkeypatch.setenv("HEADROOM_INTERCEPT_READ_MIN_CHARS", "50") + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + marker = "\n\ndef format_receipt" + partial = _PY_FIXTURE[: _PY_FIXTURE.index(marker)] + + set_current_cwd(str(tmp_path)) + set_current_request_trusted(True) + try: + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "abc", + "name": "Read", + "input": {"file_path": "payments.py"}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "abc", "content": partial}], + }, + ] + result = apply_to_messages(messages, tokenizer) + finally: + set_current_cwd(None) + set_current_request_trusted(False) + + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" not in new_content + + +def test_astgrep_disk_verification_skips_when_request_untrusted(tokenizer, tmp_path, monkeypatch): + """Same truncated-on-disk scenario, opted in and cwd bound, but the + request's peer was never marked loopback-trusted (the contextvar's + False default) -- must behave exactly like the no-signal case.""" + monkeypatch.setenv("HEADROOM_VERIFY_TRUNCATION_ON_DISK", "1") + monkeypatch.setenv("HEADROOM_INTERCEPT_READ_MIN_CHARS", "50") + f = tmp_path / "payments.py" + f.write_text(_PY_FIXTURE, encoding="utf-8") + marker = "\n\ndef format_receipt" + partial = _PY_FIXTURE[: _PY_FIXTURE.index(marker)] + + set_current_cwd(str(tmp_path)) + # Deliberately not calling set_current_request_trusted(True). + try: + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "abc", + "name": "Read", + "input": {"file_path": "payments.py"}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "abc", "content": partial}], + }, + ] + result = apply_to_messages(messages, tokenizer) + finally: + set_current_cwd(None) + + new_content = result.messages[1]["content"][0]["content"] + assert "truncated upstream" not in new_content + + +def test_astgrep_banner_detection_takes_priority_over_disk_verification(tokenizer, monkeypatch): + """When a banner is already present, disk verification must not run at + all (no cwd bound here -- if it ran, resolution would fail anyway), and + the banner's own numbers must be what the header reports.""" + monkeypatch.setenv("HEADROOM_VERIFY_TRUNCATION_ON_DISK", "1") + truncated_source = ( + _PY_FIXTURE + "\n\n[Truncated: PARTIAL view — /repo/payments.py: " + "showing lines 1-42 of 90 total (26031 tokens, cap 25000). " + "Call Read with offset=43 to see more.]\n" + ) + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "abc", + "name": "Read", + "input": {"file_path": "/repo/payments.py"}, + } + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "abc", "content": truncated_source}], + }, + ] + result = apply_to_messages(messages, tokenizer) + new_content = result.messages[1]["content"][0]["content"] + assert "showing through line 42 of 90 total" in new_content def test_astgrep_skips_small_files(tokenizer):