From 60d15504c4df8db250fc8ce2fcf6fdc327f54d82 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:57:49 -0700 Subject: [PATCH 01/24] feat: optional Forge.list_paths with PathListing, replay delegation and GitHub trees (#17 T8) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Repository context (#17) needs to know which files exist at the reviewed commit before it can match contract globs or resolve a path. This adds the optional listing surface and its first adapter. - forges/base.py: the frozen PathListing(paths, complete) dataclass, the MAX_LISTING_PAGES = 20 page cap for paged adapters, and the optional Protocol method list_paths(ref, *, sha) -> PathListing | None after get_pr_history. Callers resolve it with getattr, it never raises, and it returns None when no listing can be fetched. - forges/replay.py: ReplayForge.list_paths delegates to the inner forge like get_file_content, passing the caller's sha through, and gives None when the inner forge has no listing or raises. LocalDiffForge stays without one. - forges/github.py: one Git Trees API request (git/trees/{sha}?recursive=1) keeping only blob entries, sorted and deduplicated, with complete = not truncated. An empty sha, a transport error, a non-2xx status, a non-JSON body or a body without a tree list gives None, logged at DEBUG. - tests/test_forge_github.py: the timeout sweep pins the adapter's public method set, so list_paths is driven there and the routed session answers the trees route. - tests/test_issue_17_list_paths.py: new, 37 tests. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/forges/base.py | 32 +++ src/prxref/forges/github.py | 47 +++++ src/prxref/forges/replay.py | 16 +- tests/test_forge_github.py | 8 + tests/test_issue_17_list_paths.py | 338 ++++++++++++++++++++++++++++++ 5 files changed, 438 insertions(+), 3 deletions(-) create mode 100644 tests/test_issue_17_list_paths.py diff --git a/src/prxref/forges/base.py b/src/prxref/forges/base.py index 8f9450b..3d329d5 100644 --- a/src/prxref/forges/base.py +++ b/src/prxref/forges/base.py @@ -155,6 +155,24 @@ def __post_init__(self) -> None: raise ValueError(f"PRHistory.title_renames must hold TitleRename, got {type(rename).__name__}") +MAX_LISTING_PAGES = 20 + + +@dataclass(frozen=True) +class PathListing: + """Every file path in a repository at one commit, from ``list_paths``. + + ``paths`` holds repo-relative FILE paths, sorted and deduplicated, with + no directories and no submodules. ``complete`` is ``False`` when the + forge truncated the listing or the walk stopped at ``MAX_LISTING_PAGES`` + pages, so a path missing from an incomplete listing may still exist. + A listing is immutable and hashable. + """ + + paths: tuple[str, ...] + complete: bool + + SUMMARY_MARKER = "" # The attribution prefix every posted comment carries (the full line is @@ -269,6 +287,20 @@ def get_pr_history(self, ref: PRRef, *, head_sha: str | None = None) -> PRHistor """ ... + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """Return every file path in the repository at commit ``sha``. + + Optional: callers resolve it with ``getattr(forge, "list_paths", None)``, + so a Forge without it is still valid, and not every forge implements it + (repository context then works without a listing). The paths are + repo-relative file paths, sorted and deduplicated, with no directories + and no submodules. A listing the forge truncated, or whose walk stopped + at ``MAX_LISTING_PAGES`` pages, is returned with ``complete=False``. + Returns ``None`` when no listing can be fetched at all; this method + never raises. + """ + ... + def detect_forge(url: str) -> PRRef | None: """Try each registered forge's URL parser in order. diff --git a/src/prxref/forges/github.py b/src/prxref/forges/github.py index 112fdd7..85c5e11 100644 --- a/src/prxref/forges/github.py +++ b/src/prxref/forges/github.py @@ -22,6 +22,7 @@ DescriptionVersion, FeedReadError, InlineComment, + PathListing, PRData, PRHistory, PRRef, @@ -697,6 +698,52 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: return None return content.decode("utf-8", errors="replace") + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """Return every file path in the repository at commit ``sha``, best-effort. + + One request to the Git Trees API, ``git/trees/{sha}?recursive=1``, + with no paging. Only ``blob`` entries are kept, so directories + (``tree``) and submodules (``commit``) are dropped, and the paths are + sorted and deduplicated. ``complete`` is ``False`` when GitHub reports + the tree ``truncated``. An empty ``sha`` (no request is made), a + transport failure, a non-2xx status, a body that is not JSON, or one + with no ``tree`` list gives ``None``. Never raises. + """ + if not sha: + return None + url = ( + f"{self._api_base(ref)}/repos/{ref.owner}/{ref.repo}/git/trees/" + f"{quote(sha, safe='')}" + ) + try: + resp = self.session.get( + url, headers=self._headers(ref.host), params={"recursive": "1"}, + timeout=_REQUEST_TIMEOUT, + ) + except requests.RequestException as e: + logger.debug("list_paths failed for %s/%s@%s: %s", ref.owner, ref.repo, sha, e) + return None + if not resp.ok: + logger.debug( + "list_paths got HTTP %s for %s/%s@%s", resp.status_code, ref.owner, ref.repo, sha + ) + return None + try: + body = resp.json() + except ValueError as e: + logger.debug("list_paths got a non-JSON body for %s/%s@%s: %s", ref.owner, ref.repo, sha, e) + return None + tree = body.get("tree") if isinstance(body, dict) else None + if not isinstance(tree, list): + logger.debug("list_paths got no tree list for %s/%s@%s", ref.owner, ref.repo, sha) + return None + paths = { + entry["path"] for entry in tree + if isinstance(entry, dict) and entry.get("type") == "blob" + and isinstance(entry.get("path"), str) and entry["path"] + } + return PathListing(paths=tuple(sorted(paths)), complete=not bool(body.get("truncated"))) + def prune_inline_comments(self, ref: PRRef) -> int: """Delete prxref-attributed inline comments; returns the count removed. diff --git a/src/prxref/forges/replay.py b/src/prxref/forges/replay.py index 514706d..2bf54d2 100644 --- a/src/prxref/forges/replay.py +++ b/src/prxref/forges/replay.py @@ -7,8 +7,8 @@ ``--pr-url``): no network, no threads, no file reads. - :class:`ReplayForge` wraps a real forge and pins what the orchestrator sees: the diff of a commit range (``base_sha``/``head_sha``, through the inner - forge's optional ``get_compare_diff``) or a diff text, file reads at the - pinned head, optionally no existing threads, and the PR's title and + forge's optional ``get_compare_diff``) or a diff text, file reads and path + listings at the pinned head, optionally no existing threads, and the PR's title and description as they stood at a cutoff (issue #16). Both raise on every write method, as defence in depth: the CLI already forces @@ -29,7 +29,7 @@ from pathlib import Path from typing import Literal -from .base import Forge, InlineComment, PRData, PRHistory, PRRef, Thread, _require_aware +from .base import Forge, InlineComment, PathListing, PRData, PRHistory, PRRef, Thread, _require_aware NEVER_POSTS = "replay runs never write to a forge" @@ -288,6 +288,16 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: except Exception: # noqa: BLE001 - the Protocol says this never raises return None + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """The inner forge's path listing at ``sha``; ``None`` when it has none or it raises.""" + lister = getattr(self._inner, "list_paths", None) + if lister is None: + return None + try: + return lister(ref, sha=sha) + except Exception: # noqa: BLE001 - the Protocol says this never raises + return None + def post_summary(self, ref: PRRef, body: str) -> None: """Always raises ``RuntimeError``: a replay never writes to a forge.""" raise RuntimeError(NEVER_POSTS) diff --git a/tests/test_forge_github.py b/tests/test_forge_github.py index ff59e32..af57749 100644 --- a/tests/test_forge_github.py +++ b/tests/test_forge_github.py @@ -738,6 +738,11 @@ def get(url, headers=None, params=None, **kwargs): text="x = 1\n", headers={"Content-Type": "application/vnd.github.raw+json; charset=utf-8"}, ) + if "/git/trees/" in url: + return _mock_response(json_data={ + "sha": HEAD_SHA, "truncated": False, + "tree": [{"path": "src/app.py", "type": "blob"}], + }) if url.endswith("/issues/42/comments"): return _mock_response(json_data=summary_feed) if url.endswith("/pulls/42/comments"): @@ -791,6 +796,9 @@ def test_every_request_the_adapter_sends_carries_the_timeout(monkeypatch): ) == "x = 1\n", "prune_inline_comments": lambda: forge.prune_inline_comments(ref) == 1, "get_pr_history": lambda: forge.get_pr_history(ref).complete, + "list_paths": lambda: getattr( + forge.list_paths(ref, sha=HEAD_SHA), "paths", None + ) == ("src/app.py",), # Both branches: no summary yet (POST), and one to update (PATCH). "post_summary": lambda: ( forge.post_summary(ref, "first") is None diff --git a/tests/test_issue_17_list_paths.py b/tests/test_issue_17_list_paths.py new file mode 100644 index 0000000..9ece10d --- /dev/null +++ b/tests/test_issue_17_list_paths.py @@ -0,0 +1,338 @@ +"""Tests for the optional ``Forge.list_paths`` listing (issue #17, task T8). + +``PathListing`` and the Protocol method live in ``forges/base.py``; +``ReplayForge`` delegates the listing to the forge it wraps, ``LocalDiffForge`` +has none, and GitHub answers it from one Git Trees API request. +""" +from __future__ import annotations + +import dataclasses +import inspect +import json +import logging +from unittest.mock import MagicMock + +import pytest +import requests + +from prxref.forges import base +from prxref.forges.base import Forge, PathListing, PRRef +from prxref.forges.github import ForgeImpl +from prxref.forges.replay import LocalDiffForge, ReplayForge + +SHA = "e13a2c97d926386a950ae1a559d2ee50b113ad2e" +OTHER_SHA = "b" * 40 +REQUEST_TIMEOUT = (10.0, 30.0) +TREE = [ + {"mode": "100644", "path": "src/b.py", "sha": "1" * 40, "size": 10, "type": "blob"}, + {"mode": "040000", "path": "src", "sha": "2" * 40, "type": "tree"}, + {"mode": "160000", "path": "vendor/lib", "sha": "3" * 40, "type": "commit"}, + {"mode": "100644", "path": "README.md", "sha": "4" * 40, "size": 5, "type": "blob"}, + {"mode": "100644", "path": "src/a.py", "sha": "5" * 40, "size": 7, "type": "blob"}, + {"mode": "100644", "path": "src/a.py", "sha": "5" * 40, "size": 7, "type": "blob"}, +] +BLOB_PATHS = ("README.md", "src/a.py", "src/b.py") + + +def _mock_response(status_code=200, json_data=None, text="", content=None, headers=None): + resp = MagicMock(spec=requests.Response) + resp.status_code = status_code + resp.ok = 200 <= status_code < 300 + resp.headers = headers or {} + if json_data is not None: + resp.json.return_value = json_data + resp.text = json.dumps(json_data) + else: + resp.text = text + resp.json.side_effect = ValueError("No JSON") + resp.content = content if content is not None else resp.text.encode("utf-8") + resp.raise_for_status.side_effect = ( + None if resp.ok else requests.HTTPError(response=resp) + ) + return resp + + +def _ref(url="https://github.com/acme/api/pull/42"): + ref = ForgeImpl.parse_pr_url(url) + assert ref is not None + return ref + + +def _tree_body(tree=TREE, truncated=False): + return {"sha": SHA, "url": "https://api.example.com/tree", "tree": tree, "truncated": truncated} + + +def _session(response): + session = MagicMock(spec=requests.Session) + session.get.return_value = response + return session + + +# --- PathListing and the Protocol -------------------------------------------- + + +class TestPathListing: + def test_it_is_frozen(self): + listing = PathListing(paths=("a.py",), complete=True) + with pytest.raises(dataclasses.FrozenInstanceError): + listing.complete = False # type: ignore[misc] + + def test_it_is_hashable_and_compares_by_value(self): + one = PathListing(paths=("a.py", "b.py"), complete=True) + two = PathListing(paths=("a.py", "b.py"), complete=True) + assert hash(one) == hash(two) + assert one == two + assert len({one, two, PathListing(paths=("a.py", "b.py"), complete=False)}) == 2 + + def test_its_fields_are_exactly_paths_and_complete(self): + assert [field.name for field in dataclasses.fields(PathListing)] == ["paths", "complete"] + + +class TestProtocolDeclaration: + def test_the_page_cap_is_twenty(self): + assert base.MAX_LISTING_PAGES == 20 + + def test_the_protocol_declares_it_with_a_keyword_only_sha(self): + signature = inspect.signature(Forge.list_paths) + params = signature.parameters + assert list(params) == ["self", "ref", "sha"] + assert params["sha"].kind is inspect.Parameter.KEYWORD_ONLY + assert params["sha"].default is inspect.Parameter.empty + assert signature.return_annotation == "PathListing | None" + + def test_it_follows_get_pr_history_in_the_protocol(self): + names = [name for name in vars(Forge) if not name.startswith("_")] + assert names.index("list_paths") == names.index("get_pr_history") + 1 + + def test_the_docstring_says_it_is_optional_and_never_raises(self): + doc = inspect.getdoc(Forge.list_paths) + assert 'getattr(forge, "list_paths", None)' in doc + assert "complete=False" in doc + assert "MAX_LISTING_PAGES" in doc + assert "never raises" in doc + + +# --- GitHub ------------------------------------------------------------------- + + +class TestGitHubListPaths: + def test_it_keeps_only_blob_paths_sorted_and_deduplicated(self, monkeypatch): + monkeypatch.setenv("PRXREF_GITHUB_TOKEN", "t0ken") + session = _session(_mock_response(json_data=_tree_body())) + + listing = ForgeImpl(session=session).list_paths(_ref(), sha=SHA) + + assert listing == PathListing(paths=BLOB_PATHS, complete=True) + session.get.assert_called_once() + call = session.get.call_args + assert call.args[0] == f"https://api.github.com/repos/acme/api/git/trees/{SHA}" + assert call.kwargs["params"] == {"recursive": "1"} + assert call.kwargs["timeout"] == REQUEST_TIMEOUT + assert call.kwargs["headers"]["Accept"] == "application/vnd.github+json" + assert call.kwargs["headers"]["Authorization"] == "Bearer t0ken" + + def test_the_live_probe_entry_shape_is_read(self): + entry = { + "mode": "100644", "path": ".dockerignore", + "sha": "b1b58bb3cf967b160f341ebb00e6e8c99bee2955", "size": 116, "type": "blob", + "url": "https://api.github.com/repos/acme/api/git/blobs/b1b58bb3cf967b160f341ebb00e6e8c99bee2955", + } + session = _session(_mock_response(json_data=_tree_body(tree=[entry]))) + + listing = ForgeImpl(session=session).list_paths(_ref(), sha=SHA) + + assert listing == PathListing(paths=(".dockerignore",), complete=True) + + @pytest.mark.parametrize(("truncated", "complete"), [(False, True), (True, False)]) + def test_complete_mirrors_truncated(self, truncated, complete): + session = _session(_mock_response(json_data=_tree_body(truncated=truncated))) + + listing = ForgeImpl(session=session).list_paths(_ref(), sha=SHA) + + assert listing is not None + assert listing.complete is complete + assert listing.paths == BLOB_PATHS + + def test_malformed_entries_are_skipped(self): + tree = [ + "not-an-entry", + {"type": "blob"}, + {"type": "blob", "path": ""}, + {"type": "blob", "path": 7}, + {"type": "blob", "path": "kept.py"}, + ] + session = _session(_mock_response(json_data=_tree_body(tree=tree))) + + listing = ForgeImpl(session=session).list_paths(_ref(), sha=SHA) + + assert listing == PathListing(paths=("kept.py",), complete=True) + + def test_an_empty_tree_is_an_empty_complete_listing(self): + session = _session(_mock_response(json_data=_tree_body(tree=[]))) + + assert ForgeImpl(session=session).list_paths(_ref(), sha=SHA) == PathListing( + paths=(), complete=True + ) + + def test_the_sha_is_quoted_into_the_url(self): + session = _session(_mock_response(json_data=_tree_body())) + + ForgeImpl(session=session).list_paths(_ref(), sha="feat/x y") + + assert session.get.call_args.args[0] == ( + "https://api.github.com/repos/acme/api/git/trees/feat%2Fx%20y" + ) + + def test_an_enterprise_host_uses_its_own_api_base(self, monkeypatch): + monkeypatch.setenv("PRXREF_GITHUB_ENTERPRISE_TOKEN", "ghes-t0ken") + session = _session(_mock_response(json_data=_tree_body())) + ref = _ref("https://git.corp.example/acme/api/pull/7") + + listing = ForgeImpl(session=session).list_paths(ref, sha=SHA) + + assert listing == PathListing(paths=BLOB_PATHS, complete=True) + call = session.get.call_args + assert call.args[0] == f"https://git.corp.example/api/v3/repos/acme/api/git/trees/{SHA}" + assert call.kwargs["params"] == {"recursive": "1"} + assert call.kwargs["headers"]["Authorization"] == "Bearer ghes-t0ken" + + @pytest.mark.parametrize("status", [404, 409, 500]) + def test_a_non_2xx_status_gives_none(self, status): + session = _session(_mock_response(status, json_data={"message": "nope"})) + + assert ForgeImpl(session=session).list_paths(_ref(), sha=SHA) is None + + @pytest.mark.parametrize( + "error", [requests.ConnectionError("down"), requests.Timeout("slow"), requests.RequestException("x")] + ) + def test_a_request_exception_gives_none(self, error): + session = MagicMock(spec=requests.Session) + session.get.side_effect = error + + assert ForgeImpl(session=session).list_paths(_ref(), sha=SHA) is None + + def test_a_non_json_body_gives_none(self): + session = _session(_mock_response(text="gateway")) + + assert ForgeImpl(session=session).list_paths(_ref(), sha=SHA) is None + + @pytest.mark.parametrize( + "body", + [ + {"sha": SHA, "truncated": False}, + {"sha": SHA, "truncated": False, "tree": None}, + {"sha": SHA, "truncated": False, "tree": {"path": "a.py", "type": "blob"}}, + [{"path": "a.py", "type": "blob"}], + ], + ids=["missing-tree", "null-tree", "tree-not-a-list", "body-not-an-object"], + ) + def test_a_body_without_a_tree_list_gives_none(self, body): + session = _session(_mock_response(json_data=body)) + + assert ForgeImpl(session=session).list_paths(_ref(), sha=SHA) is None + + def test_an_empty_sha_gives_none_without_a_request(self): + session = MagicMock(spec=requests.Session) + + assert ForgeImpl(session=session).list_paths(_ref(), sha="") is None + session.get.assert_not_called() + + def test_failures_never_log_above_debug(self, caplog): + session = MagicMock(spec=requests.Session) + session.get.side_effect = [ + requests.ConnectionError("down"), + _mock_response(500, json_data={"message": "boom"}), + _mock_response(text="not json"), + _mock_response(json_data={"sha": SHA}), + ] + forge = ForgeImpl(session=session) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.github"): + results = [forge.list_paths(_ref(), sha=SHA) for _ in range(4)] + + assert results == [None, None, None, None] + assert len(caplog.records) == 4 + assert all(record.levelno <= logging.DEBUG for record in caplog.records) + + +# --- ReplayForge and LocalDiffForge ------------------------------------------- + + +class _ListingForge: + """An inner forge with a listing that records every call.""" + + name = "fake" + + def __init__(self, result=None, error=None): + self.calls: list[tuple[PRRef, str]] = [] + self._result = result + self._error = error + + def list_paths(self, ref, *, sha): + self.calls.append((ref, sha)) + if self._error is not None: + raise self._error + return self._result + + +class _NoListingForge: + """An inner forge with no ``list_paths`` at all.""" + + name = "bare" + + +LISTING = PathListing(paths=BLOB_PATHS, complete=True) + + +class TestReplayForgeListPaths: + def test_it_delegates_and_returns_the_inner_listing(self): + inner = _ListingForge(result=LISTING) + ref = _ref() + + result = ReplayForge(inner).list_paths(ref, sha=SHA) + + assert result is LISTING + assert inner.calls == [(ref, SHA)] + + def test_it_passes_the_callers_sha_through_unchanged_under_a_pin(self): + inner = _ListingForge(result=LISTING) + ref = _ref() + replay = ReplayForge(inner, base_sha="a" * 40, head_sha=SHA, diff_text="diff --git a/x b/x\n") + + replay.list_paths(ref, sha=OTHER_SHA) + + assert inner.calls == [(ref, OTHER_SHA)] + + def test_an_incomplete_or_absent_inner_listing_is_returned_as_is(self): + partial = PathListing(paths=("a.py",), complete=False) + assert ReplayForge(_ListingForge(result=partial)).list_paths(_ref(), sha=SHA) is partial + assert ReplayForge(_ListingForge(result=None)).list_paths(_ref(), sha=SHA) is None + + def test_an_inner_forge_without_list_paths_gives_none(self): + replay = ReplayForge(_NoListingForge()) + + assert getattr(replay, "list_paths", None) is not None + assert replay.list_paths(_ref(), sha=SHA) is None + + @pytest.mark.parametrize("error", [RuntimeError("boom"), requests.ConnectionError("down"), ValueError("bad")]) + def test_an_inner_forge_that_raises_gives_none(self, error): + inner = _ListingForge(error=error) + + assert ReplayForge(inner).list_paths(_ref(), sha=SHA) is None + assert len(inner.calls) == 1 + + def test_it_delegates_to_the_github_adapter_end_to_end(self): + session = _session(_mock_response(json_data=_tree_body())) + + listing = ReplayForge(ForgeImpl(session=session)).list_paths(_ref(), sha=SHA) + + assert listing == PathListing(paths=BLOB_PATHS, complete=True) + assert session.get.call_args.args[0].endswith(f"/git/trees/{SHA}") + + +class TestLocalDiffForgeHasNoListing: + def test_it_has_no_list_paths_attribute(self): + forge = LocalDiffForge("diff --git a/x b/x\n", path="change.diff") + + assert getattr(forge, "list_paths", None) is None + assert not hasattr(LocalDiffForge, "list_paths") From b106b2d8ac5f1b32827c42ef349df7ad0928be32 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 15:59:03 -0700 Subject: [PATCH 02/24] test: add the issue #17 repository-context acceptance fixture (T16) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds tests/fixtures/issue17/, a small Java repo fixture that reproduces the two repository-context misses from issue #17: an OpenAPI idempotency-key contract that sits outside the diff and disagrees with a new migration's unique index, and a TransportConfig mutual-exclusion constraint that changes in one diff chunk while a service in another chunk violates it. repo/ holds the working tree at the PR head (Java sources under src/main/java/, an OpenAPI spec under api/openapi/, and three Liquibase changesets under db/changelog/). pr.diff adds the 003 migration and modifies TransportConfig.java and ConnectorService.java; every added line matches the repo/ head content, verified by a new head-consistency test. cases.json carries the two "error" labels a resolver should surface, each anchored on an added diff line with a must_match regex. tests/test_issue_17_fixture.py checks the fixture structurally: cases.json loads and its labels anchor, build_chunks(max_files_per_chunk=1) puts TransportConfig.java and ConnectorService.java in different chunks, the OpenAPI spec and the earlier two migrations exist under repo/ but are absent from the diff, every added diff line matches repo/ at its new line number, and the fixture text uses only placeholder identities (acme, example.com). 🤖 Authored with Claude Code — Claude Sonnet 5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- tests/fixtures/issue17/cases.json | 27 +++++ tests/fixtures/issue17/pr.diff | 56 +++++++++ .../issue17/repo/api/openapi/connectors.yaml | 62 ++++++++++ .../db/changelog/001-create-connectors.sql | 8 ++ .../changelog/002-create-idempotency-keys.sql | 9 ++ .../db/changelog/003-idempotency-unique.sql | 4 + .../com/acme/connectors/ConnectorService.java | 27 +++++ .../com/acme/connectors/TransportConfig.java | 23 ++++ tests/test_issue_17_fixture.py | 114 ++++++++++++++++++ 9 files changed, 330 insertions(+) create mode 100644 tests/fixtures/issue17/cases.json create mode 100644 tests/fixtures/issue17/pr.diff create mode 100644 tests/fixtures/issue17/repo/api/openapi/connectors.yaml create mode 100644 tests/fixtures/issue17/repo/db/changelog/001-create-connectors.sql create mode 100644 tests/fixtures/issue17/repo/db/changelog/002-create-idempotency-keys.sql create mode 100644 tests/fixtures/issue17/repo/db/changelog/003-idempotency-unique.sql create mode 100644 tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/ConnectorService.java create mode 100644 tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/TransportConfig.java create mode 100644 tests/test_issue_17_fixture.py diff --git a/tests/fixtures/issue17/cases.json b/tests/fixtures/issue17/cases.json new file mode 100644 index 0000000..91ca9aa --- /dev/null +++ b/tests/fixtures/issue17/cases.json @@ -0,0 +1,27 @@ +{ + "version": 1, + "cases": [ + { + "id": "issue17-repo-context", + "diff_file": "pr.diff", + "expected": [ + { + "id": "L1", + "file": "db/changelog/003-idempotency-unique.sql", + "line": 4, + "severity": "error", + "text": "The unique index on idempotency_keys omits connector_id, so a key is only unique per (tenant, key) even though the OpenAPI IdempotencyKey contract (api/openapi/connectors.yaml) requires uniqueness per (tenant, connector, key).", + "must_match": "re:(?i)connector_id" + }, + { + "id": "L2", + "file": "src/main/java/com/acme/connectors/ConnectorService.java", + "line": 23, + "severity": "error", + "text": "createTransport passes both request.url() and request.legacyUrl() into the TransportConfig constructor, but TransportConfig's compact constructor (and the OpenAPI TransportConfig schema) require the two fields to be mutually exclusive, so this always throws IllegalArgumentException when a client sends both.", + "must_match": "re:(?i)(mutually exclusive|exactly one of)" + } + ] + } + ] +} diff --git a/tests/fixtures/issue17/pr.diff b/tests/fixtures/issue17/pr.diff new file mode 100644 index 0000000..de6b683 --- /dev/null +++ b/tests/fixtures/issue17/pr.diff @@ -0,0 +1,56 @@ +diff --git a/db/changelog/003-idempotency-unique.sql b/db/changelog/003-idempotency-unique.sql +new file mode 100644 +--- /dev/null ++++ b/db/changelog/003-idempotency-unique.sql +@@ -0,0 +1,4 @@ ++--liquibase formatted sql ++ ++--changeset acme:3 ++CREATE UNIQUE INDEX ux_idempotency_keys ON idempotency_keys (tenant_id, key); +diff --git a/src/main/java/com/acme/connectors/TransportConfig.java b/src/main/java/com/acme/connectors/TransportConfig.java +--- a/src/main/java/com/acme/connectors/TransportConfig.java ++++ b/src/main/java/com/acme/connectors/TransportConfig.java +@@ -8,6 +8,15 @@ + */ + public record TransportConfig(String url, String legacyUrl) { + ++ public TransportConfig { ++ boolean hasUrl = url != null; ++ boolean hasLegacyUrl = legacyUrl != null; ++ if (hasUrl == hasLegacyUrl) { ++ throw new IllegalArgumentException( ++ "exactly one of url or legacyUrl must be set"); ++ } ++ } ++ + public String effectiveUrl() { + return url != null ? url : legacyUrl; + } +diff --git a/src/main/java/com/acme/connectors/ConnectorService.java b/src/main/java/com/acme/connectors/ConnectorService.java +--- a/src/main/java/com/acme/connectors/ConnectorService.java ++++ b/src/main/java/com/acme/connectors/ConnectorService.java +@@ -1,5 +1,8 @@ + package com.acme.connectors; + ++import org.springframework.web.bind.annotation.PathVariable; ++import org.springframework.web.bind.annotation.PostMapping; ++import org.springframework.web.bind.annotation.RequestBody; + import org.springframework.web.bind.annotation.RequestHeader; + import org.springframework.web.bind.annotation.RestController; + +@@ -10,4 +13,15 @@ + public class ConnectorService { + + private final Map transportsByKey = new ConcurrentHashMap<>(); ++ ++ @PostMapping("/connectors/{connectorId}/transports") ++ public TransportConfig createTransport( ++ @RequestHeader("X-Tenant-Id") String tenantId, ++ @PathVariable String connectorId, ++ @RequestBody CreateTransportRequest request) { ++ String idempotencyKey = tenantId + ":" + request.idempotencyKey(); ++ TransportConfig config = new TransportConfig(request.url(), request.legacyUrl()); ++ transportsByKey.put(idempotencyKey, config); ++ return config; ++ } + } diff --git a/tests/fixtures/issue17/repo/api/openapi/connectors.yaml b/tests/fixtures/issue17/repo/api/openapi/connectors.yaml new file mode 100644 index 0000000..791eeda --- /dev/null +++ b/tests/fixtures/issue17/repo/api/openapi/connectors.yaml @@ -0,0 +1,62 @@ +openapi: 3.0.3 +info: + title: Acme Connectors API + version: "1.0.0" +paths: + /connectors/{connectorId}/transports: + post: + operationId: createTransport + summary: Create a transport configuration for a connector. + parameters: + - name: connectorId + in: path + required: true + schema: + type: string + requestBody: + required: true + content: + application/json: + schema: + $ref: '#/components/schemas/CreateTransportRequest' + responses: + '201': + description: The created transport configuration. + content: + application/json: + schema: + $ref: '#/components/schemas/TransportConfig' +components: + schemas: + TransportConfig: + type: object + description: >- + url and legacyUrl are mutually exclusive: exactly one of the two + fields must be set. + properties: + url: + type: string + legacyUrl: + type: string + CreateTransportRequest: + type: object + properties: + idempotencyKey: + type: string + url: + type: string + legacyUrl: + type: string + IdempotencyKey: + type: object + description: >- + An idempotency key is unique per (tenant, connector, key); replaying + the same key for a different connector in the same tenant must be + rejected, not silently merged. + properties: + tenantId: + type: string + connectorId: + type: string + key: + type: string diff --git a/tests/fixtures/issue17/repo/db/changelog/001-create-connectors.sql b/tests/fixtures/issue17/repo/db/changelog/001-create-connectors.sql new file mode 100644 index 0000000..9e3a429 --- /dev/null +++ b/tests/fixtures/issue17/repo/db/changelog/001-create-connectors.sql @@ -0,0 +1,8 @@ +--liquibase formatted sql + +--changeset acme:1 +CREATE TABLE connectors ( + id VARCHAR(36) NOT NULL PRIMARY KEY, + tenant_id VARCHAR(36) NOT NULL, + name VARCHAR(255) NOT NULL +); diff --git a/tests/fixtures/issue17/repo/db/changelog/002-create-idempotency-keys.sql b/tests/fixtures/issue17/repo/db/changelog/002-create-idempotency-keys.sql new file mode 100644 index 0000000..f62b489 --- /dev/null +++ b/tests/fixtures/issue17/repo/db/changelog/002-create-idempotency-keys.sql @@ -0,0 +1,9 @@ +--liquibase formatted sql + +--changeset acme:2 +CREATE TABLE idempotency_keys ( + tenant_id VARCHAR(36) NOT NULL, + connector_id VARCHAR(36) NOT NULL, + key VARCHAR(255) NOT NULL, + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); diff --git a/tests/fixtures/issue17/repo/db/changelog/003-idempotency-unique.sql b/tests/fixtures/issue17/repo/db/changelog/003-idempotency-unique.sql new file mode 100644 index 0000000..c7349bf --- /dev/null +++ b/tests/fixtures/issue17/repo/db/changelog/003-idempotency-unique.sql @@ -0,0 +1,4 @@ +--liquibase formatted sql + +--changeset acme:3 +CREATE UNIQUE INDEX ux_idempotency_keys ON idempotency_keys (tenant_id, key); diff --git a/tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/ConnectorService.java b/tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/ConnectorService.java new file mode 100644 index 0000000..a2d911f --- /dev/null +++ b/tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/ConnectorService.java @@ -0,0 +1,27 @@ +package com.acme.connectors; + +import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; +import org.springframework.web.bind.annotation.RequestHeader; +import org.springframework.web.bind.annotation.RestController; + +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; + +@RestController +public class ConnectorService { + + private final Map transportsByKey = new ConcurrentHashMap<>(); + + @PostMapping("/connectors/{connectorId}/transports") + public TransportConfig createTransport( + @RequestHeader("X-Tenant-Id") String tenantId, + @PathVariable String connectorId, + @RequestBody CreateTransportRequest request) { + String idempotencyKey = tenantId + ":" + request.idempotencyKey(); + TransportConfig config = new TransportConfig(request.url(), request.legacyUrl()); + transportsByKey.put(idempotencyKey, config); + return config; + } +} diff --git a/tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/TransportConfig.java b/tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/TransportConfig.java new file mode 100644 index 0000000..764ad7f --- /dev/null +++ b/tests/fixtures/issue17/repo/src/main/java/com/acme/connectors/TransportConfig.java @@ -0,0 +1,23 @@ +package com.acme.connectors; + +/** + * Transport configuration for a connector. Exactly one of {@code url} or + * {@code legacyUrl} must be set; the API contract at + * api/openapi/connectors.yaml#/components/schemas/TransportConfig documents + * the two fields as mutually exclusive. + */ +public record TransportConfig(String url, String legacyUrl) { + + public TransportConfig { + boolean hasUrl = url != null; + boolean hasLegacyUrl = legacyUrl != null; + if (hasUrl == hasLegacyUrl) { + throw new IllegalArgumentException( + "exactly one of url or legacyUrl must be set"); + } + } + + public String effectiveUrl() { + return url != null ? url : legacyUrl; + } +} diff --git a/tests/test_issue_17_fixture.py b/tests/test_issue_17_fixture.py new file mode 100644 index 0000000..34e9a5a --- /dev/null +++ b/tests/test_issue_17_fixture.py @@ -0,0 +1,114 @@ +"""Structural checks for the issue #17 acceptance fixture (T16). + +``tests/fixtures/issue17/`` reproduces the two repository-context misses from +issue #17: a contract file outside the diff (the OpenAPI idempotency-key +uniqueness contract) and a type changed in another diff chunk (the +``TransportConfig`` mutual-exclusion constructor). These tests validate the +fixture itself, not any resolver: ``cases.json`` loads and its labels anchor +on lines the diff adds, the diff chunks the two Java files apart, the OpenAPI +spec and the earlier migrations sit outside the diff, every added diff line +matches the head content committed under ``repo/``, and the fixture text +carries only placeholder identities. +""" +from __future__ import annotations + +import re +from pathlib import Path + +from prxref import eval_cases +from prxref.triage import build_chunks, parse_unified_diff + +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" +REPO = FIXTURE / "repo" + +TRANSPORT_CONFIG_PATH = "src/main/java/com/acme/connectors/TransportConfig.java" +CONNECTOR_SERVICE_PATH = "src/main/java/com/acme/connectors/ConnectorService.java" +MIGRATION_003_PATH = "db/changelog/003-idempotency-unique.sql" + +_EMAIL_RE = re.compile(r"[A-Za-z0-9._%+-]+@([A-Za-z0-9.-]+)") +_DOMAIN_RE = re.compile(r"\b[a-zA-Z0-9-]+\.(?:com|org|net|io|dev|co)\b") + + +def _diff_files(): + diff_text = (FIXTURE / "pr.diff").read_text(encoding="utf-8") + return parse_unified_diff(diff_text) + + +def test_cases_json_loads_and_holds_two_error_labels(): + cases = eval_cases.load_cases(str(FIXTURE / "cases.json")) + assert len(cases) == 1 + case = cases[0] + assert case.id == "issue17-repo-context" + assert len(case.expected) == 2 + assert {finding.severity for finding in case.expected} == {"error"} + assert all(finding.must_match and finding.must_match.startswith(eval_cases.MUST_MATCH_REGEX_PREFIX) + for finding in case.expected) + + +def test_labels_anchor_on_the_migration_and_the_transport_config_call(): + cases = eval_cases.load_cases(str(FIXTURE / "cases.json")) + anchors = {(finding.file, finding.line) for finding in cases[0].expected} + assert (MIGRATION_003_PATH, 4) in anchors + assert (CONNECTOR_SERVICE_PATH, 23) in anchors + + +def test_diff_touches_exactly_the_three_changed_files(): + files = _diff_files() + assert {f.path for f in files} == { + MIGRATION_003_PATH, + TRANSPORT_CONFIG_PATH, + CONNECTOR_SERVICE_PATH, + } + + +def test_chunking_puts_transport_config_and_connector_service_in_different_chunks(): + files = _diff_files() + chunks = build_chunks(files, max_files_per_chunk=1) + assert len(chunks) >= 2 + chunk_of = {f.path: index for index, chunk in enumerate(chunks) for f in chunk} + assert chunk_of[TRANSPORT_CONFIG_PATH] != chunk_of[CONNECTOR_SERVICE_PATH] + + +def test_openapi_spec_and_earlier_migrations_exist_but_sit_outside_the_diff(): + diff_paths = {f.path for f in _diff_files()} + outside_the_diff = { + "api/openapi/connectors.yaml", + "db/changelog/001-create-connectors.sql", + "db/changelog/002-create-idempotency-keys.sql", + } + assert diff_paths.isdisjoint(outside_the_diff) + for relative_path in outside_the_diff: + assert (REPO / relative_path).is_file() + + +def test_every_added_line_matches_the_repo_head_content(): + for file_diff in _diff_files(): + repo_lines = (REPO / file_diff.path).read_text(encoding="utf-8").splitlines() + for hunk in file_diff.hunks: + for line in hunk.lines: + if line.kind != "+" or line.new_line is None: + continue + index = line.new_line - 1 + assert 0 <= index < len(repo_lines), ( + f"{file_diff.path}:{line.new_line} is past the end of repo/{file_diff.path}" + ) + assert repo_lines[index] == line.text, ( + f"{file_diff.path}:{line.new_line} disagrees with repo/{file_diff.path}" + ) + + +def test_fixture_text_uses_only_placeholder_identities(): + paths = [FIXTURE / "pr.diff", FIXTURE / "cases.json"] + paths.extend(sorted(p for p in REPO.rglob("*") if p.is_file())) + assert len(paths) >= 8 + for path in paths: + text = path.read_text(encoding="utf-8") + for match in _EMAIL_RE.finditer(text): + assert match.group(1) == "example.com", ( + f"{path}: email domain {match.group(1)!r} is not the example.com placeholder" + ) + for match in _DOMAIN_RE.finditer(text): + domain = match.group().lower() + assert domain == "example.com", ( + f"{path}: domain {domain!r} is not the example.com placeholder" + ) From f78610daf2f6578d11eac058c11f74ae9709e8cf Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:00:22 -0700 Subject: [PATCH 03/24] feat: add RepoDir local filesystem reader for --repo-dir (T11) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds src/prxref/forges/repo_dir.py holding RepoDir, a confined local directory reader that stands in for a PR head with no network access, for the upcoming --repo-dir flag (wired by a later seat). RepoDir.read(path) matches the forge adapters' get_file_content contract byte-for-byte (missing/directory/unsafe-path/.git/oversized/binary all read as None, otherwise utf-8 with errors="replace", 512 KiB ceiling, size checked from a stat before the file is read). RepoDir.list_files() walks the tree, never follows a symlinked directory, skips .git at any depth, omits symlinks that escape the root, and caps at 100_000 files, returning (paths, complete). The confinement check resolves real paths and compares against root plus a path separator, not a string prefix, so a sibling directory that extends the root's name (".../r2" next to ".../r") is never mistaken for being inside it; a dedicated regression test and a falsifiability run against a naive startswith(root) check both confirm this. list_files() returns a plain tuple rather than the PathListing dataclass because that type is landing in forges/base.py under a parallel seat (T8); the wiring seat (T15) wraps this return value. tests/test_repo_dir_reader.py covers normal reads, missing/directory targets, unsafe paths (absolute, "..", backslash, empty, NUL), symlinked files and directories both inside and outside the root, .git exclusion, the 512 KiB boundary, NUL-byte and invalid-utf8 handling, sorted POSIX listing output, the file-count cap at and past the boundary, and the ValueError raised for a non-existent root. 🤖 Authored with Claude Code — Claude Sonnet 5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/forges/repo_dir.py | 121 +++++++++++++++++++++++ tests/test_repo_dir_reader.py | 175 ++++++++++++++++++++++++++++++++++ 2 files changed, 296 insertions(+) create mode 100644 src/prxref/forges/repo_dir.py create mode 100644 tests/test_repo_dir_reader.py diff --git a/src/prxref/forges/repo_dir.py b/src/prxref/forges/repo_dir.py new file mode 100644 index 0000000..66f42f0 --- /dev/null +++ b/src/prxref/forges/repo_dir.py @@ -0,0 +1,121 @@ +"""Local filesystem reader that stands in for a PR head via ``--repo-dir``. + +``RepoDir`` gives a review or an eval case repository context read from a +local working tree, with no network access. ``read`` matches the shape of +the forge adapters' ``get_file_content`` contract (see +``forges/github.py``); ``list_files`` gives a bounded, confined directory +listing. The CLI flag and its wiring into the review pipeline belong to a +later seat; this module only reads a confined filesystem tree. +""" + +from __future__ import annotations + +import os +import stat +from pathlib import PurePosixPath + +_MAX_FILE_CONTENT_BYTES = 512 * 1024 +_MAX_LISTED_FILES = 100_000 + + +class RepoDir: + """Reads a confined local directory tree in place of a forge working copy.""" + + def __init__(self, root: str | os.PathLike[str]) -> None: + """Resolve ``root`` with ``os.path.realpath``. + + Raises ``ValueError`` naming ``root`` when it is not an existing + directory. + """ + resolved = os.path.realpath(root) + if not os.path.isdir(resolved): + raise ValueError(f"repo-dir root is not an existing directory: {root}") + self.root = resolved + + def _confined_real_path(self, path: str) -> str | None: + """Return the confined, symlink-resolved filesystem path for ``path``. + + Returns None for an unsafe ``path`` (absolute, a ``..`` segment, a + backslash, a NUL byte, or empty), for one whose resolved real path + leaves the root, and for one that lands inside ``.git`` at any + depth. The confinement check compares resolved real paths, never + string prefixes. + """ + if not path or "\x00" in path or "\\" in path or os.path.isabs(path): + return None + if ".." in PurePosixPath(path).parts: + return None + candidate = os.path.join(self.root, path) + real = os.path.realpath(candidate) + root_prefix = self.root + os.sep + if real != self.root and not real.startswith(root_prefix): + return None + rel_parts = PurePosixPath(os.path.relpath(real, self.root)).parts + if ".git" in rel_parts: + return None + return real + + def read(self, path: str) -> str | None: + """Return the utf-8 text of ``path`` under the root, or None. + + Matches the forge adapters' ``get_file_content`` contract: never + raises, and returns None when the path is unsafe, missing, a + directory, escapes the root (directly or via a symlink), falls + inside ``.git``, is over 512 KiB, or holds a NUL byte. The size is + checked from a stat before the file is read. Otherwise the content + is decoded as utf-8 with ``errors="replace"``. + """ + real = self._confined_real_path(path) + if real is None: + return None + try: + st = os.stat(real) + except OSError: + return None + if not stat.S_ISREG(st.st_mode): + return None + if st.st_size > _MAX_FILE_CONTENT_BYTES: + return None + try: + with open(real, "rb") as fh: + content = fh.read() + except OSError: + return None + if b"\x00" in content: + return None + return content.decode("utf-8", errors="replace") + + def list_files(self) -> tuple[tuple[str, ...], bool]: + """Return sorted repo-relative POSIX file paths, and whether the walk was complete. + + The walk never follows a symlinked directory, skips any ``.git`` + directory at any depth, and omits symlinks whose target escapes the + root. It stops past ``_MAX_LISTED_FILES`` files and reports + ``complete=False`` only when a file beyond the cap was actually + dropped; a tree with exactly the cap's worth of files is complete. + """ + found: list[str] = [] + root_prefix = self.root + os.sep + complete = True + for dirpath, dirnames, filenames in os.walk(self.root, followlinks=False): + dirnames[:] = sorted(d for d in dirnames if d != ".git") + for name in sorted(filenames): + full = os.path.join(dirpath, name) + if os.path.islink(full): + real = os.path.realpath(full) + if real != self.root and not real.startswith(root_prefix): + continue + if not os.path.isfile(full): + continue + rel = PurePosixPath(os.path.relpath(full, self.root)) + if ".git" in rel.parts: + continue + if len(found) < _MAX_LISTED_FILES: + found.append(rel.as_posix()) + else: + complete = False + break + if not complete: + break + found.sort() + return tuple(found), complete diff --git a/tests/test_repo_dir_reader.py b/tests/test_repo_dir_reader.py new file mode 100644 index 0000000..f11cae6 --- /dev/null +++ b/tests/test_repo_dir_reader.py @@ -0,0 +1,175 @@ +"""``RepoDir`` reads a local working tree in place of a forge, for ``--repo-dir``. + +``read`` mirrors the forge adapters' ``get_file_content`` contract (never +raises, None for missing/binary/oversized/unsafe paths) and ``list_files`` +gives a bounded, confined listing. The confinement check must use +``os.path.realpath`` and not a string prefix: a sibling directory whose name +extends the root's name (``.../r2`` next to ``.../r``) must not be mistaken +for being inside it. +""" +from __future__ import annotations + +import os +import re + +import pytest + +from prxref.forges import repo_dir as repo_dir_module +from prxref.forges.repo_dir import RepoDir + + +def test_read_normal_file(tmp_path): + (tmp_path / "a.py").write_text("print('hi')\n") + rd = RepoDir(tmp_path) + assert rd.read("a.py") == "print('hi')\n" + + +def test_read_missing_file(tmp_path): + rd = RepoDir(tmp_path) + assert rd.read("nope.py") is None + + +def test_read_directory(tmp_path): + (tmp_path / "sub").mkdir() + rd = RepoDir(tmp_path) + assert rd.read("sub") is None + + +@pytest.mark.parametrize( + "unsafe", + ["../x", "/etc/hosts", "a/../../x", "a\\b", ""], +) +def test_unsafe_paths_read_none(tmp_path, unsafe): + rd = RepoDir(tmp_path) + assert rd.read(unsafe) is None + + +def test_symlink_to_file_outside_root_reads_none(tmp_path): + root = tmp_path / "root" + root.mkdir() + outside = tmp_path / "outside" + outside.mkdir() + secret = outside / "secret.txt" + secret.write_text("nope\n") + link = root / "link.txt" + link.symlink_to(secret) + rd = RepoDir(root) + assert rd.read("link.txt") is None + + +def test_symlink_to_file_inside_root_reads_fine(tmp_path): + (tmp_path / "real.txt").write_text("hello\n") + link = tmp_path / "link.txt" + link.symlink_to(tmp_path / "real.txt") + rd = RepoDir(tmp_path) + assert rd.read("link.txt") == "hello\n" + + +def test_symlinked_directory_outside_root_not_walked(tmp_path): + root = tmp_path / "root" + root.mkdir() + (root / "kept.txt").write_text("keep\n") + outside = tmp_path / "outside" + outside.mkdir() + (outside / "secret.txt").write_text("nope\n") + (root / "escape").symlink_to(outside) + rd = RepoDir(root) + paths, complete = rd.list_files() + assert paths == ("kept.txt",) + assert complete is True + + +def test_git_config_read_none_and_never_listed(tmp_path): + git_dir = tmp_path / ".git" + git_dir.mkdir() + (git_dir / "config").write_text("[core]\n") + (tmp_path / "kept.txt").write_text("keep\n") + rd = RepoDir(tmp_path) + assert rd.read(".git/config") is None + paths, complete = rd.list_files() + assert paths == ("kept.txt",) + assert complete is True + + +def test_file_exactly_512kib_reads(tmp_path): + size = 512 * 1024 + (tmp_path / "big.txt").write_bytes(b"a" * size) + rd = RepoDir(tmp_path) + content = rd.read("big.txt") + assert content is not None + assert len(content) == size + + +def test_file_512kib_plus_one_byte_reads_none(tmp_path): + size = 512 * 1024 + 1 + (tmp_path / "toobig.txt").write_bytes(b"a" * size) + rd = RepoDir(tmp_path) + assert rd.read("toobig.txt") is None + + +def test_nul_byte_reads_none(tmp_path): + (tmp_path / "bin.dat").write_bytes(b"abc\x00def") + rd = RepoDir(tmp_path) + assert rd.read("bin.dat") is None + + +def test_invalid_utf8_is_replaced_not_raised(tmp_path): + (tmp_path / "bad.txt").write_bytes(b"caf\xe9 au lait") + rd = RepoDir(tmp_path) + content = rd.read("bad.txt") + assert content is not None + assert "\N{REPLACEMENT CHARACTER}" in content + + +def test_listing_is_sorted_and_posix(tmp_path): + (tmp_path / "b").mkdir() + (tmp_path / "a.txt").write_text("1\n") + (tmp_path / "b" / "z.txt").write_text("2\n") + (tmp_path / "b" / "a.txt").write_text("3\n") + rd = RepoDir(tmp_path) + paths, complete = rd.list_files() + assert paths == ("a.txt", "b/a.txt", "b/z.txt") + assert complete is True + assert all("\\" not in p for p in paths) + + +def test_listing_cap_stops_and_reports_incomplete(tmp_path, monkeypatch): + for i in range(5): + (tmp_path / f"f{i}.txt").write_text("x\n") + monkeypatch.setattr(repo_dir_module, "_MAX_LISTED_FILES", 3) + rd = RepoDir(tmp_path) + paths, complete = rd.list_files() + assert len(paths) == 3 + assert complete is False + + +def test_listing_exactly_at_cap_is_complete(tmp_path, monkeypatch): + for i in range(3): + (tmp_path / f"f{i}.txt").write_text("x\n") + monkeypatch.setattr(repo_dir_module, "_MAX_LISTED_FILES", 3) + rd = RepoDir(tmp_path) + paths, complete = rd.list_files() + assert len(paths) == 3 + assert complete is True + + +def test_root_that_does_not_exist_raises_value_error_naming_it(tmp_path): + missing = tmp_path / "nope" + with pytest.raises(ValueError, match=re.escape(os.fspath(missing))): + RepoDir(missing) + + +def test_confinement_uses_realpath_not_string_prefix(tmp_path): + root = tmp_path / "r" + root.mkdir() + sibling = tmp_path / "r2" + sibling.mkdir() + secret = sibling / "secret" + secret.write_text("nope\n") + link = root / "link.txt" + link.symlink_to(secret) + rd = RepoDir(root) + assert rd.read("link.txt") is None + paths, complete = rd.list_files() + assert paths == () + assert complete is True From 2685c1454207acceb613e92fb353979a01fc107c Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:00:24 -0700 Subject: [PATCH 04/24] feat: add the four 0.16.0 repo-context config keys (#17 seat T1) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add PRXREF_REPO_CONTEXT (choice off|diff|repo, default off), PRXREF_REPO_CONTEXT_MAX_CHARS (int, default 12000, _Range(0)), PRXREF_CONTEXT_CONTRACT_GLOBS (list key defaulting to the OQ5 built-in contract-glob set; a set value replaces the default, empty reads as unset) and PRXREF_CONTEXT_EXCLUDE_GLOBS (list key, default empty, added to a floor a later seat owns) to config._DEFAULTS and its _INT_KEYS/_LIST_KEYS/_CHOICE_KEYS/_RANGES tables, per decisions.md "Config keys" (D2: all four surfaces land in this one seat). Document all four on every surface: the config.py module docstring (key table and the list-key sentence), .env.example, and docs/env-vars.md, and recompute the two stated key-count totals in docs/env-vars.md (63->67 configuration keys, 64->68 accepted variable names with the legacy alias) with code rather than by hand. Add _KEYS_0_16 and TestKeys016AreDocumented to tests/test_config.py, mirroring the 0.15 substring-trap guard, and a new tests/test_issue_17_config.py pinning load_config behaviour: defaults, exact case-sensitive matching for PRXREF_REPO_CONTEXT (like PRXREF_FAIL_ON), the open-zero _Range(0) rejection of 0 and negative values for the max-chars key, and the replace/empty-is-unset contract for the contract-globs default. Nothing reads these keys yet; wiring lands in later tasks (T13/T14). 🤖 Authored with Claude Code — Claude Sonnet 5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- .env.example | 33 +++++ docs/env-vars.md | 10 +- src/prxref/config.py | 69 ++++++++++- tests/test_config.py | 26 ++++ tests/test_issue_17_config.py | 219 ++++++++++++++++++++++++++++++++++ 5 files changed, 351 insertions(+), 6 deletions(-) create mode 100644 tests/test_issue_17_config.py diff --git a/.env.example b/.env.example index 0285181..420139b 100644 --- a/.env.example +++ b/.env.example @@ -279,6 +279,39 @@ PRXREF_MAX_CHUNKS=8 # a visible marker. Must be > 0. # PRXREF_TICKET_CONTEXT_MAX_CHARS=6000 +# Repository context (0.16.0): off (default) | diff | repo. off leaves every +# prompt, post, trace and default-verbosity log byte-identical to 0.15.0. +# diff adds cross-chunk definitions from other files already in the diff, +# plus diff-file entries, all read from the diff itself -- no repository +# reader is needed. repo also reads files outside the diff (import, +# path-convention and name-search definitions, plus contract excerpts) +# through the forge's repository reader when one is available, or +# --repo-dir. Matching is exact and case-sensitive, like PRXREF_FAIL_ON; any +# other value is a configuration error (exit 2). +# PRXREF_REPO_CONTEXT=off + +# Repository context (0.16.0): per-chunk character budget shared by the +# cross-chunk, contract and other repository-context entries. Must be > 0. +# PRXREF_REPO_CONTEXT_MAX_CHARS=12000 + +# Repository context (0.16.0): globs (matched like PRXREF_SIZE_IGNORE_GLOBS) +# selecting the contract files -- OpenAPI, JSON Schema, Liquibase/SQL +# migrations -- excerpted under "repo". A set value REPLACES the built-in +# set below rather than adding to it; an empty value reads as unset, so the +# built-in set stays -- there is no way to turn contract excerpts off on +# their own in 0.16.0 short of setting PRXREF_REPO_CONTEXT to off or diff. +# Built-in set: **/openapi*.y*ml, **/openapi*.json, **/openapi/**, +# **/swagger*, **/*.schema.json, **/db/changelog/**, **/db/migration/**, +# **/migrations/** +# PRXREF_CONTEXT_CONTRACT_GLOBS= + +# Repository context (0.16.0): globs (matched like PRXREF_SIZE_IGNORE_GLOBS) +# whose paths are never read for repository context, not even a diff file. +# ADDED to a floor that is always on: **/expected.json, **/cases.json, +# **/case.json, **/prxref-eval/**, **/.env*, **/*.pem, **/*.key. Empty +# (the default) adds nothing. +# PRXREF_CONTEXT_EXCLUDE_GLOBS= + # Jira base URL (scheme://host plus any context path) that ticket fetches are # looked up on, overriding a ticket URL's own base (a self-hosted board often # sits behind a different REST host than its browse URL). Jira credentials diff --git a/docs/env-vars.md b/docs/env-vars.md index 4da761f..f7c855c 100644 --- a/docs/env-vars.md +++ b/docs/env-vars.md @@ -55,6 +55,10 @@ Configuration is loaded from built-in defaults, overridden by environment variab | `PRXREF_PROMPTS_DIR` | *(empty — off)* | **Prompt template overrides (0.15.0).** Directory holding replacement `worker.md`, `systemic.md` and `summary.md` prompt templates; a file that is absent keeps the packaged one. Every template is validated before any network call: `worker.md` and `systemic.md` must keep the `## Review Context` marker and every packaged placeholder below it (the feature slots `{scope_example}` and `{rule_example}` are optional), and `summary.md` needs only `{findings}`. A missing directory, a failed check, or a template over 256 KiB raises `ConfigError` and `prxref review` exits `2`; an empty directory or an unknown placeholder only warns. The run record's `prompt_templates` stamps the directory and, for each template file present in it, edited or not, its path, SHA-256 and length; a template absent from the directory has no entry. `--prompts-dir DIR` wins for one run, and `prxref prompts export DIR [--force]` writes the packaged templates to start from. The judge prompt of `prxref eval` cannot be overridden. Read it from a trusted checkout: a PR that commits a template rewrites its own review. Unset uses the packaged templates, byte for byte. | | `PRXREF_TICKET_CONTEXT_FILE` | *(empty — off)* | Plain-text or Markdown file holding the ticket this PR implements. When set, every finding is marked in, out of, or of unknown ticket scope. An empty or whitespace-only file means "this PR has no ticket". A missing or non-UTF-8 file raises `ConfigError` and `prxref review` exits `2`. The webhook daemon ignores it (and says so once). `--context-file PATH` wins for one run, and `--context-file ""` turns it off. | | `PRXREF_TICKET_CONTEXT_MAX_CHARS` | `6000` | Characters of ticket text kept in the prompt; longer text is truncated with a visible marker. Must be **greater than 0**. | +| `PRXREF_REPO_CONTEXT` | `off` | **Repository context (0.16.0).** `off` (the default) leaves every prompt, post, trace and default-verbosity log byte-identical to 0.15.0. `diff` adds cross-chunk definitions from other files already in the diff, plus diff-file entries, all read from the diff itself, so no repository reader is needed. `repo` additionally reads files outside the diff — import, path-convention and name-search definitions, plus contract excerpts — through the forge's repository reader when one is available, or `--repo-dir`. Matching is exact and case-sensitive, like `PRXREF_FAIL_ON`: `Repo` or `REPO` is rejected. Any value outside `off`\|`diff`\|`repo` raises `ConfigError` and `prxref review` exits `2`. | +| `PRXREF_REPO_CONTEXT_MAX_CHARS` | `12000` | **Repository context (0.16.0).** Per-chunk character budget shared by the cross-chunk, contract and other repository-context entries admitted under `PRXREF_REPO_CONTEXT`. Must be **greater than 0**. | +| `PRXREF_CONTEXT_CONTRACT_GLOBS` | the built-in set below | **Repository context (0.16.0).** Globs (matched like `PRXREF_SIZE_IGNORE_GLOBS`: case-sensitive `fnmatch`, `*` crosses `/`) selecting the contract files — OpenAPI, JSON Schema, Liquibase/SQL migrations — excerpted under `repo`. A set value **REPLACES** the built-in set below, rather than adding to it; an empty value reads as unset, so the built-in set stays. There is no way to turn contract excerpts off on their own in 0.16.0 — set `PRXREF_REPO_CONTEXT` to `off` or `diff` instead. Built-in set: `**/openapi*.y*ml`, `**/openapi*.json`, `**/openapi/**`, `**/swagger*`, `**/*.schema.json`, `**/db/changelog/**`, `**/db/migration/**`, `**/migrations/**`. | +| `PRXREF_CONTEXT_EXCLUDE_GLOBS` | *(empty)* | **Repository context (0.16.0).** Globs (matched like `PRXREF_SIZE_IGNORE_GLOBS`) whose paths are never read for repository context, not even a diff file. **Added** to a floor that is always on: `**/expected.json`, `**/cases.json`, `**/case.json`, `**/prxref-eval/**`, `**/.env*`, `**/*.pem`, `**/*.key`. Empty (the default) adds nothing. | The replay flags of `prxref review` (`--base-sha`, `--head-sha`, `--no-threads`, `--diff-file`, `--as-of`, `--description-file`, `--no-description`) deliberately have no environment variable: set in the environment, a replay pin would silently pin every run, the webhook daemon's included. @@ -162,11 +166,11 @@ The two knobs above, plus the four opt-in levers added in 0.15.0 (`PRXREF_MAX_WA ## Environment Cross-Check & Defaults -The tables above define all **63** configuration keys in `src/prxref/config.py` (`_DEFAULTS`), and every one of them appears in `.env.example`: +The tables above define all **67** configuration keys in `src/prxref/config.py` (`_DEFAULTS`), and every one of them appears in `.env.example`: -- **LLM / Pipeline (45):** `PRXREF_LLM_BACKEND`, `PRXREF_LLM_BASE_URL`, `PRXREF_LLM_API_KEY`, `PRXREF_LLM_MODELS`, `PRXREF_LLM_REASONING_EFFORT`, `PRXREF_LLM_MAX_TOKENS`, `PRXREF_LLM_TIMEOUT`, `PRXREF_LLM_TEMPERATURE`, `PRXREF_LLM_SEED`, `PRXREF_LLM_CLI_PATH`, `PRXREF_LLM_CLI_CONCURRENCY`, `PRXREF_CONFIDENCE_FLOOR`, `PRXREF_MAX_ERROR_FINDINGS`, `PRXREF_MAX_WARNING_FINDINGS`, `PRXREF_MAX_OUTOFSCOPE_FINDINGS`, `PRXREF_MAX_FINDINGS_PER_RULE`, `PRXREF_GROUP_FINDINGS`, `PRXREF_DEDUP_SIMILARITY`, `PRXREF_MAX_CHUNKS`, `PRXREF_CHUNK_TOKEN_BUDGET`, `PRXREF_CHUNK_MAX_FILES`, `PRXREF_CHUNK_CONTEXT_LINES`, `PRXREF_MAX_WORKERS`, `PRXREF_MAX_INLINE_COMMENTS`, `PRXREF_FAIL_ON`, `PRXREF_DRY_RUN`, `PRXREF_TRACE_FILE`, `PRXREF_TRACE_DIR`, `PRXREF_POST_MODE`, `PRXREF_POST_VERDICT`, `PRXREF_PRICE_TABLE`, `PRXREF_POST_COST`, `PRXREF_SIZE_WARN_LINES`, `PRXREF_SIZE_WARN_FILES`, `PRXREF_SIZE_IGNORE_GLOBS`, `PRXREF_SPEC_SOURCES`, `PRXREF_SPEC_MAX_CHARS`, `PRXREF_SPEC_DIGEST_TOKENS`, `PRXREF_REVIEW_RULES`, `PRXREF_REVIEW_RULES_MAX_CHARS`, `PRXREF_SCOPED_RULES`, `PRXREF_SCOPED_RULES_MAX_CHARS`, `PRXREF_PROMPTS_DIR`, `PRXREF_TICKET_CONTEXT_FILE`, `PRXREF_TICKET_CONTEXT_MAX_CHARS` +- **LLM / Pipeline (49):** `PRXREF_LLM_BACKEND`, `PRXREF_LLM_BASE_URL`, `PRXREF_LLM_API_KEY`, `PRXREF_LLM_MODELS`, `PRXREF_LLM_REASONING_EFFORT`, `PRXREF_LLM_MAX_TOKENS`, `PRXREF_LLM_TIMEOUT`, `PRXREF_LLM_TEMPERATURE`, `PRXREF_LLM_SEED`, `PRXREF_LLM_CLI_PATH`, `PRXREF_LLM_CLI_CONCURRENCY`, `PRXREF_CONFIDENCE_FLOOR`, `PRXREF_MAX_ERROR_FINDINGS`, `PRXREF_MAX_WARNING_FINDINGS`, `PRXREF_MAX_OUTOFSCOPE_FINDINGS`, `PRXREF_MAX_FINDINGS_PER_RULE`, `PRXREF_GROUP_FINDINGS`, `PRXREF_DEDUP_SIMILARITY`, `PRXREF_MAX_CHUNKS`, `PRXREF_CHUNK_TOKEN_BUDGET`, `PRXREF_CHUNK_MAX_FILES`, `PRXREF_CHUNK_CONTEXT_LINES`, `PRXREF_MAX_WORKERS`, `PRXREF_MAX_INLINE_COMMENTS`, `PRXREF_FAIL_ON`, `PRXREF_DRY_RUN`, `PRXREF_TRACE_FILE`, `PRXREF_TRACE_DIR`, `PRXREF_POST_MODE`, `PRXREF_POST_VERDICT`, `PRXREF_PRICE_TABLE`, `PRXREF_POST_COST`, `PRXREF_SIZE_WARN_LINES`, `PRXREF_SIZE_WARN_FILES`, `PRXREF_SIZE_IGNORE_GLOBS`, `PRXREF_SPEC_SOURCES`, `PRXREF_SPEC_MAX_CHARS`, `PRXREF_SPEC_DIGEST_TOKENS`, `PRXREF_REVIEW_RULES`, `PRXREF_REVIEW_RULES_MAX_CHARS`, `PRXREF_SCOPED_RULES`, `PRXREF_SCOPED_RULES_MAX_CHARS`, `PRXREF_PROMPTS_DIR`, `PRXREF_TICKET_CONTEXT_FILE`, `PRXREF_TICKET_CONTEXT_MAX_CHARS`, `PRXREF_REPO_CONTEXT`, `PRXREF_REPO_CONTEXT_MAX_CHARS`, `PRXREF_CONTEXT_CONTRACT_GLOBS`, `PRXREF_CONTEXT_EXCLUDE_GLOBS` - **Per-Forge Auth (10):** `PRXREF_BITBUCKET_TOKEN`, `PRXREF_BITBUCKET_USER`, `PRXREF_BITBUCKET_APP_PASSWORD`, `PRXREF_BITBUCKET_SERVER_TOKEN`, `PRXREF_BITBUCKET_SERVER_USER`, `PRXREF_BITBUCKET_SERVER_PASSWORD`, `PRXREF_GITHUB_TOKEN`, `PRXREF_GITHUB_ENTERPRISE_TOKEN`, `PRXREF_GITLAB_TOKEN`, `PRXREF_AZURE_DEVOPS_TOKEN` - **Spec Sources / Jira (3):** `PRXREF_JIRA_BASE_URL`, `PRXREF_JIRA_EMAIL`, `PRXREF_JIRA_API_TOKEN` - **Webhooks (5):** `PRXREF_BITBUCKET_WEBHOOK_SECRET`, `PRXREF_GITHUB_WEBHOOK_SECRET`, `PRXREF_GITLAB_WEBHOOK_SECRET`, `PRXREF_AZURE_DEVOPS_WEBHOOK_SECRET`, `PRXREF_ALLOW_UNSIGNED` -*(63 configuration keys, plus one deprecated alias — `PRXREF_MAX_ERRORS` for `PRXREF_MAX_ERROR_FINDINGS` — for 64 accepted variable names.)* +*(67 configuration keys, plus one deprecated alias — `PRXREF_MAX_ERRORS` for `PRXREF_MAX_ERROR_FINDINGS` — for 68 accepted variable names.)* diff --git a/src/prxref/config.py b/src/prxref/config.py index 193b281..6108cb6 100644 --- a/src/prxref/config.py +++ b/src/prxref/config.py @@ -248,6 +248,45 @@ Characters of ticket text kept in the prompt; longer is truncated with a visible marker; positive int (default 6000) + PRXREF_REPO_CONTEXT Repository context (0.16.0): "off" (default) | + "diff" | "repo". "off" leaves every prompt, + post, trace and default-verbosity log + byte-identical to 0.15.0. "diff" adds + cross-chunk definitions from other files + already in the diff, plus diff-file entries, + all read from the diff itself; no repository + reader is needed. "repo" also reads files + outside the diff — import, path-convention and + name-search definitions, plus contract excerpts + — through the forge's repository reader when + one is available, or --repo-dir. Matching is + exact and case-sensitive, like PRXREF_FAIL_ON; + any other value is a configuration error + PRXREF_REPO_CONTEXT_MAX_CHARS Repository context (0.16.0): per-chunk + character budget shared by the cross-chunk, + contract and other repository-context entries; + positive int (default 12000) + PRXREF_CONTEXT_CONTRACT_GLOBS Repository context (0.16.0): globs (matched + like PRXREF_SIZE_IGNORE_GLOBS) selecting the + contract files — OpenAPI, JSON Schema, + Liquibase/SQL migrations — excerpted under + "repo". A set value REPLACES the built-in set + below rather than adding to it, and an empty + value reads as unset (the built-in set stays); + there is no way to turn contract excerpts off + on their own in 0.16.0 short of setting + PRXREF_REPO_CONTEXT to "off" or "diff". + Built-in set: **/openapi*.y*ml, + **/openapi*.json, **/openapi/**, **/swagger*, + **/*.schema.json, **/db/changelog/**, + **/db/migration/**, **/migrations/** + PRXREF_CONTEXT_EXCLUDE_GLOBS Repository context (0.16.0): globs (matched + like PRXREF_SIZE_IGNORE_GLOBS) whose paths are + never read for repository context, not even a + diff file. ADDED to a floor that is always on: + **/expected.json, **/cases.json, **/case.json, + **/prxref-eval/**, **/.env*, **/*.pem, + **/*.key. Empty (the default) adds nothing Spec sources / Jira: PRXREF_JIRA_BASE_URL Jira base URL (scheme://host plus any @@ -298,8 +337,9 @@ webhooks (default off; insecure) List-valued keys (PRXREF_LLM_MODELS, PRXREF_SPEC_SOURCES, -PRXREF_SIZE_IGNORE_GLOBS and PRXREF_SCOPED_RULES) split on any run of commas -and/or whitespace, so no item can contain either; a glob that must match a +PRXREF_SIZE_IGNORE_GLOBS, PRXREF_SCOPED_RULES, PRXREF_CONTEXT_CONTRACT_GLOBS +and PRXREF_CONTEXT_EXCLUDE_GLOBS) split on any run of commas and/or +whitespace, so no item can contain either; a glob that must match a literal space writes it as ``?``. Precedence: built-in defaults < environment < ``overrides`` kwargs. @@ -396,6 +436,26 @@ "prompts_dir": None, "ticket_context_file": "", "ticket_context_max_chars": 6000, + "repo_context": "off", + "repo_context_max_chars": 12000, + # The OQ5 built-in contract-glob set (map-17, decisions.md). Unlike the + # other _LIST_KEYS defaults, this one is non-empty: an env value REPLACES + # it rather than adding to it, and an empty value reads as unset (the + # normal "empty or whitespace-only reads as unset" rule), so this set + # stays in place. A ``set`` literal would work too -- the _LIST_KEYS + # coercion branch always returns a ``list`` -- but the default is typed + # as a ``list`` from the start so both paths give callers the same type. + "context_contract_globs": [ + "**/openapi*.y*ml", + "**/openapi*.json", + "**/openapi/**", + "**/swagger*", + "**/*.schema.json", + "**/db/changelog/**", + "**/db/migration/**", + "**/migrations/**", + ], + "context_exclude_globs": [], "jira_base_url": "", "jira_email": "", "jira_api_token": "", @@ -424,7 +484,7 @@ "llm_cli_concurrency", "review_rules_max_chars", "ticket_context_max_chars", "size_warn_lines", "size_warn_files", "max_warning_findings", "max_outofscope_findings", "scoped_rules_max_chars", - "max_findings_per_rule", + "max_findings_per_rule", "repo_context_max_chars", }) _FLOAT_KEYS = frozenset({"confidence_floor", "llm_timeout", "dedup_similarity"}) _BOOL_KEYS = frozenset({ @@ -432,6 +492,7 @@ }) _LIST_KEYS = frozenset({ "llm_models", "spec_sources", "size_ignore_globs", "scoped_rules", + "context_contract_globs", "context_exclude_globs", }) # An enum-valued key has no numeric interval to check, so its legal vocabulary @@ -441,6 +502,7 @@ # PRXREF_FAIL_ON=eror into an undetected "never". _CHOICE_KEYS: dict[str, frozenset[str]] = { "fail_on": frozenset({"never", "error", "any"}), + "repo_context": frozenset({"off", "diff", "repo"}), } # The posting-behaviour vocabulary, validated rather than trusted. Restated in @@ -522,6 +584,7 @@ def describe(self) -> str: "max_outofscope_findings": _Range(0, low_inclusive=True), "max_findings_per_rule": _Range(0, low_inclusive=True), "scoped_rules_max_chars": _Range(0), + "repo_context_max_chars": _Range(0), "confidence_floor": _Range(0.0, 1.0, low_inclusive=True), "dedup_similarity": _Range(0.0, 1.0), } diff --git a/tests/test_config.py b/tests/test_config.py index 1c0201f..b853551 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -1857,6 +1857,32 @@ def test_the_env_name_is_derived_for_the_suite_wide_clear(self, key, env): def test_each_surface_has_one_entry_that_names_the_release(self, key, env, surface): assert "0.15.0" in _doc_entry(surface, env) + +_KEYS_0_16 = [ + ("repo_context", "PRXREF_REPO_CONTEXT"), + ("repo_context_max_chars", "PRXREF_REPO_CONTEXT_MAX_CHARS"), + ("context_contract_globs", "PRXREF_CONTEXT_CONTRACT_GLOBS"), + ("context_exclude_globs", "PRXREF_CONTEXT_EXCLUDE_GLOBS"), +] + + +class TestKeys016AreDocumented: + """The four-surface rule, tightened for the 0.16 keys. + + ``test_docs_consistency`` checks a SUBSTRING, so ``PRXREF_REPO_CONTEXT`` + would pass on the strength of ``PRXREF_REPO_CONTEXT_MAX_CHARS`` alone. + """ + + @pytest.mark.parametrize("key,env", _KEYS_0_16) + def test_the_env_name_is_derived_for_the_suite_wide_clear(self, key, env): + assert config._ENV_PREFIX + key.upper() == env + assert env in prxref_env_names() + + @pytest.mark.parametrize("surface", sorted(SURFACES)) + @pytest.mark.parametrize("key,env", _KEYS_0_16) + def test_each_surface_has_one_entry_that_names_the_release(self, key, env, surface): + assert "0.16.0" in _doc_entry(surface, env) + @pytest.mark.parametrize("surface", sorted(SURFACES)) def test_the_outofscope_cap_says_it_is_not_ticket_scope_out(self, surface): entry = _doc_entry(surface, "PRXREF_MAX_OUTOFSCOPE_FINDINGS") diff --git a/tests/test_issue_17_config.py b/tests/test_issue_17_config.py new file mode 100644 index 0000000..c90ad7b --- /dev/null +++ b/tests/test_issue_17_config.py @@ -0,0 +1,219 @@ +"""PRXREF_REPO_CONTEXT / _MAX_CHARS / CONTEXT_CONTRACT_GLOBS / _EXCLUDE_GLOBS. + +The four 0.16.0 config keys (#17, seat T1). Nothing reads them yet -- wiring +is a later task -- so this module pins ``load_config`` behaviour only: +defaults, coercion, range/choice enforcement, and the OQ5 "a set value +REPLACES the built-in default, empty reads as unset" contract for +``PRXREF_CONTEXT_CONTRACT_GLOBS``. +""" +from __future__ import annotations + +import pytest + +from prxref import config +from prxref.config import load_config +from prxref.llm import ConfigError + +CONTRACT_GLOBS_DEFAULT = [ + "**/openapi*.y*ml", + "**/openapi*.json", + "**/openapi/**", + "**/swagger*", + "**/*.schema.json", + "**/db/changelog/**", + "**/db/migration/**", + "**/migrations/**", +] + + +class TestDefaults: + def test_repo_context_defaults_to_off(self): + assert config._DEFAULTS["repo_context"] == "off" + assert load_config()["repo_context"] == "off" + + def test_repo_context_max_chars_defaults_to_12000(self): + assert config._DEFAULTS["repo_context_max_chars"] == 12000 + value = load_config()["repo_context_max_chars"] + assert value == 12000 + assert isinstance(value, int) + + def test_context_contract_globs_defaults_to_the_oq5_set(self): + assert config._DEFAULTS["context_contract_globs"] == CONTRACT_GLOBS_DEFAULT + assert load_config()["context_contract_globs"] == CONTRACT_GLOBS_DEFAULT + + def test_context_exclude_globs_defaults_to_empty(self): + assert config._DEFAULTS["context_exclude_globs"] == [] + assert load_config()["context_exclude_globs"] == [] + + def test_defaults_are_independent_objects_across_calls(self): + """A caller mutating its own config dict must not corrupt the + built-in default or a later ``load_config()`` call.""" + cfg = load_config() + cfg["context_contract_globs"].append("**/mutated/**") + assert load_config()["context_contract_globs"] == CONTRACT_GLOBS_DEFAULT + assert config._DEFAULTS["context_contract_globs"] == CONTRACT_GLOBS_DEFAULT + + +class TestKeyDeclarations: + def test_repo_context_is_a_choice_key(self): + assert config._CHOICE_KEYS["repo_context"] == frozenset({"off", "diff", "repo"}) + assert "repo_context" not in config._INT_KEYS | config._FLOAT_KEYS + assert "repo_context" not in config._RANGES + + def test_repo_context_max_chars_is_an_int_with_an_open_zero_bound(self): + assert "repo_context_max_chars" in config._INT_KEYS + assert config._RANGES["repo_context_max_chars"] == config._Range(0) + assert config._RANGES["repo_context_max_chars"].low_inclusive is False + + def test_the_two_glob_keys_are_list_keys(self): + assert "context_contract_globs" in config._LIST_KEYS + assert "context_exclude_globs" in config._LIST_KEYS + for key in ("context_contract_globs", "context_exclude_globs"): + assert key not in config._RANGES + assert key not in config._CHOICE_KEYS + + +class TestRepoContextChoice: + @pytest.mark.parametrize("raw", ["off", "diff", "repo"]) + def test_each_legal_value_loads(self, monkeypatch, raw): + monkeypatch.setenv("PRXREF_REPO_CONTEXT", raw) + assert load_config()["repo_context"] == raw + + def test_empty_reads_as_unset(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT", "") + assert load_config()["repo_context"] == "off" + + def test_whitespace_only_reads_as_unset(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT", " ") + assert load_config()["repo_context"] == "off" + + @pytest.mark.parametrize("raw", ["Repo", "REPO", "Off", "DIFF", "bogus"]) + def test_a_value_outside_the_vocabulary_is_a_config_error(self, monkeypatch, raw): + """Matching is exact and case-sensitive, mirroring PRXREF_FAIL_ON: + ``Repo``/``REPO`` are rejected, not silently folded to ``repo``.""" + monkeypatch.setenv("PRXREF_REPO_CONTEXT", raw) + with pytest.raises(ConfigError, match=r"^PRXREF_REPO_CONTEXT: "): + load_config() + + def test_the_error_names_the_legal_values(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT", "bogus") + with pytest.raises(ConfigError) as exc: + load_config() + for word in ("off", "diff", "repo"): + assert word in str(exc.value) + + def test_an_override_wins_over_the_environment(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT", "off") + assert load_config(repo_context="repo")["repo_context"] == "repo" + + def test_an_override_cannot_smuggle_a_value_outside_the_vocabulary(self): + with pytest.raises(ConfigError, match=r"^repo_context: ") as exc: + load_config(repo_context="bogus") + assert "PRXREF_REPO_CONTEXT" not in str(exc.value) + + +class TestRepoContextMaxChars: + @pytest.mark.parametrize("raw,expected", [("1", 1), (" 500 ", 500), ("100000", 100_000)]) + def test_env_coerces_to_an_int(self, monkeypatch, raw, expected): + monkeypatch.setenv("PRXREF_REPO_CONTEXT_MAX_CHARS", raw) + value = load_config()["repo_context_max_chars"] + assert value == expected + assert isinstance(value, int) + + def test_empty_or_whitespace_reads_as_unset(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT_MAX_CHARS", " ") + assert load_config()["repo_context_max_chars"] == 12000 + + def test_negative_value_is_a_config_error_naming_the_variable(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT_MAX_CHARS", "-1") + with pytest.raises(ConfigError, match=r"^PRXREF_REPO_CONTEXT_MAX_CHARS: ") as exc: + load_config() + assert "greater than 0" in str(exc.value) + + def test_zero_is_rejected_per_range_zero_semantics(self, monkeypatch): + """``_Range(0)`` defaults ``low_inclusive`` to False, so the interval + is open at zero -- the same contract as PRXREF_MAX_CHUNKS and + PRXREF_SCOPED_RULES_MAX_CHARS. Zero is not a legal budget: it would + ask for repository context and admit none of it.""" + assert config._RANGES["repo_context_max_chars"].accepts(0) is False + monkeypatch.setenv("PRXREF_REPO_CONTEXT_MAX_CHARS", "0") + with pytest.raises(ConfigError, match=r"^PRXREF_REPO_CONTEXT_MAX_CHARS: "): + load_config() + + def test_malformed_value_names_the_variable(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT_MAX_CHARS", "lots") + with pytest.raises(ConfigError, match=r"^PRXREF_REPO_CONTEXT_MAX_CHARS: "): + load_config() + + def test_an_override_is_range_checked_and_named_as_itself(self): + with pytest.raises(ConfigError, match=r"^repo_context_max_chars: ") as exc: + load_config(repo_context_max_chars=0) + assert "PRXREF_REPO_CONTEXT_MAX_CHARS" not in str(exc.value) + + def test_an_override_wins_over_the_environment(self, monkeypatch): + monkeypatch.setenv("PRXREF_REPO_CONTEXT_MAX_CHARS", "3") + assert load_config(repo_context_max_chars=5)["repo_context_max_chars"] == 5 + + +class TestContextContractGlobs: + def test_a_set_value_replaces_the_default(self, monkeypatch): + monkeypatch.setenv("PRXREF_CONTEXT_CONTRACT_GLOBS", "custom/*.yaml") + assert load_config()["context_contract_globs"] == ["custom/*.yaml"] + + def test_splits_on_commas_and_whitespace_like_the_other_list_keys(self, monkeypatch): + monkeypatch.setenv( + "PRXREF_CONTEXT_CONTRACT_GLOBS", "a/*.yaml,b/*.json c/*.sql" + ) + assert load_config()["context_contract_globs"] == [ + "a/*.yaml", "b/*.json", "c/*.sql", + ] + + def test_an_empty_value_gives_the_default(self, monkeypatch): + monkeypatch.setenv("PRXREF_CONTEXT_CONTRACT_GLOBS", "") + assert load_config()["context_contract_globs"] == CONTRACT_GLOBS_DEFAULT + + def test_a_whitespace_only_value_gives_the_default(self, monkeypatch): + monkeypatch.setenv("PRXREF_CONTEXT_CONTRACT_GLOBS", " ") + assert load_config()["context_contract_globs"] == CONTRACT_GLOBS_DEFAULT + + def test_an_override_replaces_the_default_as_given(self): + assert load_config(context_contract_globs=["one/*.yaml"])[ + "context_contract_globs" + ] == ["one/*.yaml"] + + def test_an_override_of_the_empty_list_is_taken_literally(self): + """Unlike the empty-string environment path (unset -> default), an + override is used exactly as given: a caller that really wants zero + contract globs passes ``[]`` as an override, not through the + environment, where an empty string cannot spell it.""" + assert load_config(context_contract_globs=[])["context_contract_globs"] == [] + + def test_an_override_wins_over_the_environment(self, monkeypatch): + monkeypatch.setenv("PRXREF_CONTEXT_CONTRACT_GLOBS", "env/*.yaml") + assert load_config(context_contract_globs=["flag/*.yaml"])[ + "context_contract_globs" + ] == ["flag/*.yaml"] + + +class TestContextExcludeGlobs: + def test_splits_on_commas_and_whitespace_like_the_other_list_keys(self, monkeypatch): + monkeypatch.setenv( + "PRXREF_CONTEXT_EXCLUDE_GLOBS", "**/secrets/**,**/*.pem\t**/vendor/**" + ) + assert load_config()["context_exclude_globs"] == [ + "**/secrets/**", "**/*.pem", "**/vendor/**", + ] + + def test_empty_or_whitespace_reads_as_unset(self, monkeypatch): + monkeypatch.setenv("PRXREF_CONTEXT_EXCLUDE_GLOBS", " \t ") + assert load_config()["context_exclude_globs"] == [] + + def test_an_override_replaces_the_environment(self, monkeypatch): + monkeypatch.setenv("PRXREF_CONTEXT_EXCLUDE_GLOBS", "a.md,b.md") + assert load_config(context_exclude_globs=["c.md"])["context_exclude_globs"] == ["c.md"] + + def test_a_literal_space_is_written_as_a_question_mark(self, monkeypatch): + """Same doctrine as PRXREF_SIZE_IGNORE_GLOBS: a comma-or-whitespace + splitter cannot carry an item containing either character.""" + monkeypatch.setenv("PRXREF_CONTEXT_EXCLUDE_GLOBS", "path/with?space/**") + assert load_config()["context_exclude_globs"] == ["path/with?space/**"] From d986acedf223c8f885542216b5620525b24b2f85 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:01:07 -0700 Subject: [PATCH 05/24] feat: repo_context definitions core with Java type declarations (#17 T2) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add src/prxref/repo_context.py, the pure core the 0.16.0 repository-context feature builds on: the ContextEntry record (rendered() in the referenced_definitions "path:line: text" shape, record() for the run record), the REASONS admission ranks and KINDS, language_of (chunk_context's map plus .java), definition_regexes (a Java type-declaration regex for class, interface, record, enum and @interface after modifiers and same-line annotations; chunk_context's regexes otherwise), referenced_names (first-seen identifiers minus keywords, names defined on the added lines, and for Java the JDK names and non-type-like names), and find_definitions (the referenced_definitions scan loop generalized to any file's text, first definition per name, skip_lines, max_lines, the 512 KiB byte cap). chunk_context.py is untouched: Java stays out of its _definition_regexes because referenced_definitions runs with the feature off (D1). The module reuses chunk_context's underscore helpers so both scan and render alike. Tests: tests/test_repo_context_defs.py (65 tests). 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/repo_context.py | 232 ++++++++++++++++++++++ tests/test_repo_context_defs.py | 337 ++++++++++++++++++++++++++++++++ 2 files changed, 569 insertions(+) create mode 100644 src/prxref/repo_context.py create mode 100644 tests/test_repo_context_defs.py diff --git a/src/prxref/repo_context.py b/src/prxref/repo_context.py new file mode 100644 index 0000000..90ace55 --- /dev/null +++ b/src/prxref/repo_context.py @@ -0,0 +1,232 @@ +"""Repository context outside the diff: shared types and the definitions core. + +``chunk_context`` answers what a referenced identifier IS only when its +definition sits in the SAME changed file, and it knows no Java. This module +holds the pieces the repository-context feature (``PRXREF_REPO_CONTEXT``) +builds on: the :class:`ContextEntry` record every context source emits, the +admission ranks in :data:`REASONS`, a language map and definition regexes that +add Java, the names an added line references, and a definition scan over the +text of ANY file. + +The module is pure. It is stdlib plus :mod:`prxref.chunk_context`, performs no +I/O and no network, and callers hand it file text they already read. It +imports ``chunk_context``'s underscore helpers (``_language``, +``_definition_regexes``, ``_keywords``, ``_entry_text``, ``_IDENT_RE``) on +purpose, so both modules scan and render a definition the same way instead of +drifting apart. Java lives here rather than in +``chunk_context._definition_regexes`` because ``referenced_definitions`` runs +with the feature off, and feature-off output must stay byte-identical. +""" +from __future__ import annotations + +import re +from collections.abc import Iterable +from dataclasses import dataclass + +from . import chunk_context + +REASONS = ("cross-chunk", "contract", "diff-file", "import", "path-convention", "name-search") +KINDS = ("definition", "contract") + +_JAVA_DEF_RE = re.compile( + r"^\s*(?:(?:@(?!interface\b)[A-Za-z_$][\w$.]*(?:\((?:[^()]|\([^()]*\))*\))?" + r"|public|protected|private|static|final|abstract|sealed|non-sealed|strictfp)\s+)*" + r"(?:class|interface|record|enum|@interface)\s+([A-Za-z_$][A-Za-z0-9_$]*)" +) + +_JAVA_IMPORT_RE = re.compile(r"^\s*import\s+(?:static\s+)?((?:java|javax)\.[\w$.]*[\w$])(?:\.\*)?\s*;") + +_JAVA_KEYWORDS = frozenset({ + "abstract", "assert", "boolean", "break", "byte", "case", "catch", "char", + "class", "const", "continue", "default", "do", "double", "else", "enum", + "exports", "extends", "false", "final", "finally", "float", "for", "goto", + "if", "implements", "import", "instanceof", "int", "interface", "long", + "module", "native", "new", "non", "null", "open", "opens", "package", + "permits", "private", "protected", "provides", "public", "record", + "requires", "return", "sealed", "short", "static", "strictfp", "super", + "switch", "synchronized", "this", "throw", "throws", "to", "transient", + "transitive", "true", "try", "uses", "var", "void", "volatile", "when", + "while", "with", "yield", +}) + +_JDK_NAMES = frozenset({ + "AbstractMap", "ArithmeticException", "ArrayDeque", "ArrayIndexOutOfBoundsException", + "ArrayList", "Arrays", "AssertionError", "AtomicBoolean", "AtomicInteger", + "AtomicLong", "AtomicReference", "AutoCloseable", "Base64", "BigDecimal", + "BigInteger", "BiConsumer", "BiFunction", "BinaryOperator", "BiPredicate", + "Boolean", "Byte", "Callable", "CharSequence", "Character", "Charset", + "ChronoUnit", "Class", "ClassCastException", "Clock", "Cloneable", + "Collection", "Collections", "Collectors", "Comparable", "Comparator", + "CompletableFuture", "CompletionStage", "ConcurrentHashMap", "ConcurrentMap", + "Consumer", "CountDownLatch", "Date", "DateTimeFormatter", "Deprecated", + "Deque", "Double", "Duration", "Enum", "Error", "Exception", "ExecutorService", + "Executors", "File", "Files", "Float", "Function", "FunctionalInterface", + "Future", "HashMap", "HashSet", "IllegalArgumentException", + "IllegalStateException", "IndexOutOfBoundsException", "InputStream", "Instant", + "IntStream", "Integer", "InterruptedException", "IOException", "Iterable", + "Iterator", "LinkedHashMap", "LinkedHashSet", "LinkedList", "List", + "LocalDate", "LocalDateTime", "LocalTime", "Locale", "Long", "LongStream", + "Map", "Matcher", "Math", "NoSuchElementException", "NullPointerException", + "Number", "NumberFormatException", "Object", "Objects", "OffsetDateTime", + "Optional", "OptionalDouble", "OptionalInt", "OptionalLong", "OutputStream", + "Override", "Path", "Paths", "Pattern", "Period", "Predicate", "PriorityQueue", + "Queue", "Random", "Reader", "Record", "ReentrantLock", "Runnable", "Runtime", + "RuntimeException", "SafeVarargs", "Set", "Short", "SortedMap", "SortedSet", + "StandardCharsets", "Stream", "String", "StringBuilder", "StringJoiner", + "Supplier", "SuppressWarnings", "System", "Thread", "ThreadLocal", "Throwable", + "TimeUnit", "TreeMap", "TreeSet", "UUID", "UnaryOperator", + "UncheckedIOException", "UnsupportedOperationException", "URI", "URL", "Void", + "Writer", "ZoneId", "ZoneOffset", "ZonedDateTime", +}) + + +@dataclass(frozen=True) +class ContextEntry: + """One repository-context entry admitted into a worker chunk's prompt. + + ``path`` is repository-relative, ``line`` the 1-based start line in that + file, ``symbol`` the type or definition name (or, for a contract, the + operationId, schema name, route template or table name). ``kind`` is a + member of :data:`KINDS`, ``reason`` a member of :data:`REASONS`, and + ``text`` the entry body with per-entry caps already applied and no + ``path:line:`` prefix. + """ + + path: str + line: int + symbol: str + kind: str + reason: str + text: str + + def rendered(self) -> str: + """The prompt line, ``path:line: text``, the shape ``referenced_definitions`` emits.""" + return f"{self.path}:{self.line}: {self.text}" + + def record(self) -> dict: + """The run-record row: path, line, symbol, kind, reason, and the rendered length as ``chars``.""" + return { + "path": self.path, + "line": self.line, + "symbol": self.symbol, + "kind": self.kind, + "reason": self.reason, + "chars": len(self.rendered()), + } + + +def language_of(path: str) -> str: + """The definition language of ``path``: ``chunk_context``'s map plus ``"java"`` for ``.java``. + + Returns ``""`` for a path no language claims. + """ + if path.lower().endswith(".java"): + return "java" + return chunk_context._language(path) + + +def definition_regexes(language: str) -> tuple[re.Pattern[str], ...]: + """The definition regexes for ``language``; group 1 of a match is the defined name. + + ``"java"`` gets one regex for a type declaration (``class``, ``interface``, + ``record``, ``enum`` or ``@interface``, after optional modifiers and + same-line annotations); it never matches a method, a field or a local + variable. Every other language delegates to ``chunk_context``, which covers + js and python and returns ``()`` otherwise. + """ + if language == "java": + return (_JAVA_DEF_RE,) + return chunk_context._definition_regexes(language) + + +def _jdk_imports(added: Iterable[str]) -> set[str]: + names: set[str] = set() + for text in added: + match = _JAVA_IMPORT_RE.match(text) + if match: + names.update(match.group(1).split(".")) + return names + + +def _type_like(name: str) -> bool: + return name[:1].isupper() and any(c.islower() for c in name) + + +def referenced_names(added: Iterable[str], language: str) -> list[str]: + """Identifiers the added lines reference, in first-appearance order, deduplicated. + + Language keywords are removed, and so is every name a + :func:`definition_regexes` match defines on the added lines themselves, as + ``chunk_context.referenced_definitions`` does. For Java, JDK names are also + removed (every segment of a ``java.*`` or ``javax.*`` import on the added + lines, plus a built-in set of common ``java.lang``, ``java.util``, + ``java.time`` and related types), and only type-like names are kept: an + uppercase first letter and at least one lowercase letter, since a Java + definition regex can only ever find a type. That drops constants, locals, + methods and single-letter type parameters. + """ + lines = list(added) + java = language == "java" + keywords = _JAVA_KEYWORDS if java else chunk_context._keywords(language) + defined: set[str] = set() + for text in lines: + for regex in definition_regexes(language): + match = regex.match(text) + if match: + defined.add(match.group(1)) + dropped = keywords | defined + if java: + dropped = dropped | _JDK_NAMES | _jdk_imports(lines) + out: dict[str, None] = {} + for text in lines: + for name in chunk_context._IDENT_RE.findall(text): + if name in dropped or name in out: + continue + if java and not _type_like(name): + continue + out[name] = None + return list(out) + + +def find_definitions( + text: str, + names: Iterable[str], + *, + language: str, + skip_lines: frozenset[int] = frozenset(), + max_lines: int = chunk_context.MAX_LINES_PER_DEFINITION, +) -> list[tuple[str, int, str]]: + """``(symbol, line, entry_text)`` for the first definition of each wanted name. + + Scans ``text`` line by line with :func:`definition_regexes`, the loop + ``chunk_context.referenced_definitions`` runs, generalized to any file: + 1-based lines in ``skip_lines`` are passed over, only the first regex that + matches a line is consulted, and each name yields at most its FIRST + definition. ``entry_text`` is the defining line plus continuation lines up + to a balanced bracket or ``max_lines``. Results are in line order. Returns + ``[]`` when ``text`` is empty or larger than ``chunk_context.MAX_FILE_BYTES`` + in UTF-8, when no name is wanted, or when the language has no regexes. + """ + regexes = definition_regexes(language) + wanted = set(names) + if not regexes or not wanted or not text: + return [] + if len(text.encode("utf-8", "ignore")) > chunk_context.MAX_FILE_BYTES: + return [] + lines = text.splitlines() + seen: set[str] = set() + out: list[tuple[str, int, str]] = [] + for idx, line in enumerate(lines): + number = idx + 1 + if number in skip_lines: + continue + for regex in regexes: + match = regex.match(line) + if not match: + continue + name = match.group(1) + if name in wanted and name not in seen: + seen.add(name) + out.append((name, number, chunk_context._entry_text(lines, idx, max_lines))) + break + return out diff --git a/tests/test_repo_context_defs.py b/tests/test_repo_context_defs.py new file mode 100644 index 0000000..0261a44 --- /dev/null +++ b/tests/test_repo_context_defs.py @@ -0,0 +1,337 @@ +"""Unit tests for the definitions core of :mod:`prxref.repo_context`. + +The module is pure, so these pin the Java declaration regex, the names an added +line references, the definition scan over any file's text, and the shared +``ContextEntry`` shape without a forge, an LLM or a filesystem in the loop. +""" +from __future__ import annotations + +import dataclasses + +import pytest + +from prxref import chunk_context +from prxref.chunk_context import ChunkFile, referenced_definitions +from prxref.repo_context import ( + KINDS, + REASONS, + ContextEntry, + definition_regexes, + find_definitions, + language_of, + referenced_names, +) + + +def reader(files: dict[str, str]): + def _read(path: str) -> str | None: + return files.get(path) + return _read + + +def _java_match(line: str) -> str | None: + for regex in definition_regexes("java"): + match = regex.match(line) + if match: + return match.group(1) + return None + + +class TestVocabulary: + def test_reasons_order_is_the_admission_rank(self): + assert REASONS == ( + "cross-chunk", "contract", "diff-file", "import", "path-convention", "name-search", + ) + + def test_kinds(self): + assert KINDS == ("definition", "contract") + + +class TestLanguageOf: + @pytest.mark.parametrize("path", [ + "src/main/java/com/acme/ConnectorService.java", + "Legacy.JAVA", + "a/b/Mixed.Java", + ]) + def test_java_suffix_is_case_insensitive(self, path): + assert language_of(path) == "java" + + @pytest.mark.parametrize("path, expected", [ + ("web/app.ts", "js"), + ("web/view.tsx", "js"), + ("pkg/mod.py", "python"), + ("cmd/main.go", "go"), + ("src/lib.rs", "rust"), + ("README.md", ""), + ("build.gradle", ""), + ("Service.javax", ""), + ]) + def test_other_paths_delegate_to_chunk_context(self, path, expected): + assert language_of(path) == expected + assert language_of(path) == chunk_context._language(path) + + +class TestDefinitionRegexes: + def test_js_and_python_delegate_to_chunk_context(self): + assert definition_regexes("js") == chunk_context._definition_regexes("js") + assert definition_regexes("python") == chunk_context._definition_regexes("python") + + @pytest.mark.parametrize("language", ["go", "rust", "", "cobol"]) + def test_languages_without_regexes_get_none(self, language): + assert definition_regexes(language) == () + + def test_java_is_not_added_to_chunk_context(self): + assert chunk_context._definition_regexes("java") == () + assert len(definition_regexes("java")) == 1 + + +class TestJavaRegex: + @pytest.mark.parametrize("line, name", [ + ("public final class ConnectorService {", "ConnectorService"), + ("public sealed interface Transport permits A, B {", "Transport"), + ("enum Mode { A, B }", "Mode"), + ("@interface Audited {", "Audited"), + ("public @interface Audited {", "Audited"), + ("@Entity public class Foo extends Bar implements Q {", "Foo"), + ("@Deprecated @SuppressWarnings(\"unused\") abstract class Old {", "Old"), + ('@Table(name = "t", indexes = @Index(columnList = "a")) public class Tbl {', "Tbl"), + ("public non-sealed class Open extends Transport {", "Open"), + (" static final class Inner> {", "Inner"), + ("protected strictfp class Calc {", "Calc"), + ]) + def test_type_declarations_match(self, line, name): + assert _java_match(line) == name + + @pytest.mark.parametrize("line, name", [ + ("record TransportConfig(String url, String legacyUrl) {", "TransportConfig"), + ("public record Pair(A a, B b) implements Q {", "Pair"), + ]) + def test_record_declarations_match(self, line, name): + assert _java_match(line) == name + + @pytest.mark.parametrize("line", [ + "public void send(TransportConfig c) {", + "private final TransportConfig config;", + "TransportConfig cfg = new TransportConfig(a, b);", + "// class Foo in a comment", + 'String s = "class Foo";', + ' * class Foo in a javadoc line', + "var record = repository.find(id);", + " record.save();", + "public static Record record(Foo x) {", + "classic Foo bar;", + "return new Foo() {", + ]) + def test_methods_fields_locals_comments_and_strings_do_not_match(self, line): + assert _java_match(line) is None + + +class TestReferencedNames: + def test_java_drops_keywords_jdk_names_locals_and_own_declarations(self): + added = [ + "import com.acme.connectors.TransportConfig;", + "import java.util.concurrent.ConcurrentSkipListMap;", + "import static java.util.Objects.requireNonNull;", + "public record RetryPolicy(int max) {", + "private static final int MAX_RETRIES = 3;", + "public Optional send(TransportConfig config, Payload payload) {", + " var cfg = new ConcurrentSkipListMap();", + " List items = RetryPolicy.of(MAX_RETRIES);", + " return Optional.of(Encoder.encode(payload, config.legacyUrl()));", + ] + assert referenced_names(added, "java") == ["TransportConfig", "Payload", "Encoder"] + + def test_java_import_filter_is_what_drops_an_unlisted_jdk_type(self): + use = " var index = new ConcurrentSkipListMap();" + assert referenced_names([use], "java") == ["ConcurrentSkipListMap", "Payload"] + assert referenced_names( + ["import java.util.concurrent.ConcurrentSkipListMap;", use], "java", + ) == ["Payload"] + + def test_javax_imports_are_jdk_too(self): + added = ["import javax.annotation.processing.Generated;", "@Generated Widget w;"] + assert referenced_names(added, "java") == ["Widget"] + + def test_java_keeps_only_type_like_names(self): + added = ["ID = URL + Widget.SIZE_MAX + widgetCount + _Hidden + $Proxy + X;"] + assert referenced_names(added, "java") == ["Widget"] + + def test_python_matches_referenced_definitions_filtering(self): + added = [ + "from app.models import Helper", + "def local_fn(x):", + " return Helper(x) + other_value", + ] + assert referenced_names(added, "python") == ["app", "models", "Helper", "x", "other_value"] + + def test_js_matches_referenced_definitions_filtering(self): + added = [ + "import { Positive } from './checks';", + "export const total = (count: number) => Positive(count) + helper;", + ] + assert referenced_names(added, "js") == ["Positive", "checks", "count", "helper"] + + def test_accepts_a_one_shot_iterable(self): + added = iter(["public class Local {", " Remote r = Local.of();"]) + assert referenced_names(added, "java") == ["Remote"] + + def test_empty_added(self): + assert referenced_names([], "java") == [] + + +_JAVA_FILE = "\n".join([ + "package com.acme.connectors;", + "", + "import java.util.List;", + "", + "@Entity public class ConnectorService {", + " public void send(TransportConfig c) {", + " TransportConfig cfg = new TransportConfig(c.url(), null);", + " }", + "}", + "", + "public record TransportConfig(", + " String url,", + " String legacyUrl", + ") {}", + "", + "enum Mode { A, B }", +]) + +_PY_FILE = "\n".join([ + "class Foo:", + " pass", + "", + "def build(", + " a,", + " b,", + "):", + " return Foo()", + "", + "Foo = build(1, 2)", +]) + +_TS_FILE = "\n".join([ + "export interface TransportConfig {", + " url: string;", + "}", + "", + "export const send = (c: TransportConfig) => c.url;", + "function helper() { return 1; }", +]) + + +class TestFindDefinitions: + def test_java_types_in_line_order_with_continuation(self): + found = find_definitions( + _JAVA_FILE, ["Mode", "TransportConfig", "ConnectorService"], language="java", + ) + assert found == [ + ("ConnectorService", 5, "@Entity public class ConnectorService {\n" + " public void send(TransportConfig c) {\n" + " TransportConfig cfg = new TransportConfig(c.url(), null);\n" + " }\n" + "}"), + ("TransportConfig", 11, "public record TransportConfig(\n" + " String url,\n" + " String legacyUrl\n" + ") {}"), + ("Mode", 16, "enum Mode { A, B }"), + ] + + def test_java_ignores_names_nobody_wants(self): + assert find_definitions(_JAVA_FILE, ["Mode"], language="java") == [ + ("Mode", 16, "enum Mode { A, B }"), + ] + + def test_first_definition_only(self): + found = find_definitions(_PY_FILE, ["Foo"], language="python") + assert found == [("Foo", 1, "class Foo:")] + + def test_skip_lines_passes_over_a_definition(self): + found = find_definitions(_PY_FILE, ["Foo"], language="python", skip_lines=frozenset({1})) + assert found == [("Foo", 10, "Foo = build(1, 2)")] + + def test_max_lines_caps_the_continuation(self): + default = find_definitions(_PY_FILE, ["build"], language="python") + capped = find_definitions(_PY_FILE, ["build"], language="python", max_lines=2) + assert default == [("build", 4, "def build(\n a,\n b,\n):")] + assert capped == [("build", 4, "def build(\n a,")] + + def test_typescript(self): + found = find_definitions(_TS_FILE, ["helper", "TransportConfig", "send"], language="js") + assert found == [ + ("TransportConfig", 1, "export interface TransportConfig {\n url: string;\n}"), + ("send", 5, "export const send = (c: TransportConfig) => c.url;"), + ("helper", 6, "function helper() { return 1; }"), + ] + + def test_unknown_language_finds_nothing(self): + assert find_definitions(_PY_FILE, ["Foo"], language="") == [] + assert find_definitions(_PY_FILE, ["Foo"], language="go") == [] + + def test_java_text_under_a_js_language_finds_nothing(self): + assert find_definitions(_JAVA_FILE, ["Mode"], language="python") == [] + + def test_no_names_or_no_text(self): + assert find_definitions(_PY_FILE, [], language="python") == [] + assert find_definitions("", ["Foo"], language="python") == [] + + def test_byte_cap_is_inclusive_at_the_limit(self): + head = "class Foo:\n pass\n" + at_cap = head + "#" * (chunk_context.MAX_FILE_BYTES - len(head.encode("utf-8"))) + assert len(at_cap.encode("utf-8")) == chunk_context.MAX_FILE_BYTES + assert find_definitions(at_cap, ["Foo"], language="python") == [("Foo", 1, "class Foo:")] + assert find_definitions(at_cap + "#", ["Foo"], language="python") == [] + + def test_byte_cap_counts_utf8_bytes_not_characters(self): + head = "class Foo:\n pass\n# " + wide = head + "\N{SNOWMAN}" * (chunk_context.MAX_FILE_BYTES // 3) + assert len(wide) <= chunk_context.MAX_FILE_BYTES + assert len(wide.encode("utf-8")) > chunk_context.MAX_FILE_BYTES + assert find_definitions(wide, ["Foo"], language="python") == [] + + +class TestContextEntry: + def test_rendered_matches_the_referenced_definitions_line(self): + content = "import x\n\nclass Helper:\n pass\n\ndef use():\n return Helper()\n" + chunk = ChunkFile(path="pkg/a.py", added=(" return Helper()",), hunk_lines=frozenset({7})) + emitted = referenced_definitions([chunk], reader({"pkg/a.py": content})) + found = find_definitions( + content, + referenced_names(chunk.added, "python"), + language="python", + skip_lines=chunk.hunk_lines, + ) + entries = [ + ContextEntry(path="pkg/a.py", line=line, symbol=symbol, + kind="definition", reason="diff-file", text=text) + for symbol, line, text in found + ] + assert emitted == ["pkg/a.py:3: class Helper:"] + assert [entry.rendered() for entry in entries] == emitted + + def test_record_carries_the_rendered_length(self): + entry = ContextEntry( + path="com/acme/TransportConfig.java", line=11, symbol="TransportConfig", + kind="definition", reason="cross-chunk", + text="public record TransportConfig(\n String url\n) {}", + ) + assert entry.rendered() == ( + "com/acme/TransportConfig.java:11: public record TransportConfig(\n String url\n) {}" + ) + assert entry.record() == { + "path": "com/acme/TransportConfig.java", + "line": 11, + "symbol": "TransportConfig", + "kind": "definition", + "reason": "cross-chunk", + "chars": len(entry.rendered()), + } + assert entry.record()["chars"] == 84 + + def test_entries_are_frozen(self): + entry = ContextEntry(path="a.py", line=1, symbol="A", kind="definition", + reason="import", text="class A:") + with pytest.raises(dataclasses.FrozenInstanceError): + entry.line = 2 # type: ignore[misc] From b522b9e5dc869a009b8eb3eb6ba23e28422aa267 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:09:20 -0700 Subject: [PATCH 06/24] feat: list_paths for GitLab and Bitbucket Server / Data Center (#17 T9) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add the optional Forge.list_paths to the two paged adapters, behaving as T8's GitHub method does: file paths only, sorted and deduplicated, never raising, None on an empty sha with no request. GitLab walks repository/tree?recursive=true with per_page=_PAGE_SIZE and page=1, 2, 3 ..., keeps type == "blob" entries, and ends the walk when the X-Next-Page header is absent or empty (never on a short page), as the 2026-09-24 keyless probe of gitlab.com showed. Bitbucket Server walks the documented /files?at= listing with start/limit paging and ends on isLastPage (or its absence) or a missing nextPageStart, as the activity walk does. The endpoint shape comes from the REST documentation and was not probed live (no public instance), so it is UNVERIFIED. Both cap the walk at MAX_LISTING_PAGES (imported from forges.base). A walk the cap stops, or a failure on a later page, returns the paths read so far with complete=False; a failure on the first page returns None. Failures log at DEBUG, as get_file_content does. Tests: tests/test_issue_17_list_paths_paged.py (52 tests). 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/forges/bitbucket_server.py | 91 +++++ src/prxref/forges/gitlab.py | 81 ++++ tests/test_issue_17_list_paths_paged.py | 480 ++++++++++++++++++++++++ 3 files changed, 652 insertions(+) create mode 100644 tests/test_issue_17_list_paths_paged.py diff --git a/src/prxref/forges/bitbucket_server.py b/src/prxref/forges/bitbucket_server.py index e64599b..3798029 100644 --- a/src/prxref/forges/bitbucket_server.py +++ b/src/prxref/forges/bitbucket_server.py @@ -12,9 +12,11 @@ from prxref.forges.base import ( ATTRIBUTION_MARKER, + MAX_LISTING_PAGES, SUMMARY_MARKER, FeedReadError, InlineComment, + PathListing, PRData, PRRef, Thread, @@ -603,6 +605,95 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: return None return content.decode("utf-8", errors="replace") + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """Return every file path in the repository at commit ``sha``, best-effort. + + Walks the repository-level ``/files?at=`` listing with + ``start``/``limit`` paging, ``_PAGE_LIMIT`` entries a page. The + endpoint shape comes from the Bitbucket Server REST documentation + and was not probed live: a page is ``{"values": [...], "isLastPage", + "nextPageStart", ...}`` whose values are plain path strings, which + name files only. The non-empty strings are kept, sorted and + deduplicated. The walk ends when ``isLastPage`` is true (or absent) + or when ``nextPageStart`` is missing, as the activity walk does; a + page that says it is not the last but names no next start ends the + walk with ``complete=False``. It reads at most ``MAX_LISTING_PAGES`` + pages; when the last page it reads is not the last page, the paths + read so far come back with ``complete=False``. A failure on the first + page (a transport failure, a non-2xx status, a body that is not + JSON, or one with no ``values`` list) gives ``None``; the same + failure on a later page gives the paths read so far with + ``complete=False``. An empty ``sha`` gives ``None`` with no request. + ``ref.owner`` already carries the ``~slug`` form for a personal + repository. Never raises. + """ + if not sha: + return None + headers, auth = self._get_auth() + url = self._repo_url(ref, "/files") + where = f"{ref.owner}/{ref.repo}@{sha}" + paths: set[str] = set() + start = 0 + for page_number in range(1, MAX_LISTING_PAGES + 1): + try: + resp = self._session.get( + url, + params={"at": sha, "start": start, "limit": _PAGE_LIMIT}, + headers=headers, + auth=auth, + timeout=_REQUEST_TIMEOUT, + ) + except requests.RequestException as e: + return self._listing_stopped(paths, page_number, where, f"a transport failure ({e})") + if not resp.ok: + return self._listing_stopped(paths, page_number, where, f"HTTP {resp.status_code}") + try: + page = resp.json() + except ValueError as e: + return self._listing_stopped(paths, page_number, where, f"a non-JSON body ({e})") + values = page.get("values") if isinstance(page, dict) else None + if not isinstance(values, list): + return self._listing_stopped( + paths, page_number, where, f"a {type(page).__name__} body with no values list" + ) + paths.update(value for value in values if isinstance(value, str) and value) + if page.get("isLastPage", True): + return PathListing(paths=tuple(sorted(paths)), complete=True) + next_start = page.get("nextPageStart") + if next_start is None: + logger.debug( + "list_paths got a page that is not the last but has no " + "nextPageStart at page %d for %s; keeping the %d paths read so far", + page_number, where, len(paths), + ) + return PathListing(paths=tuple(sorted(paths)), complete=False) + start = next_start + logger.debug( + "list_paths stopped at the %d-page cap for %s with %d paths", + MAX_LISTING_PAGES, where, len(paths), + ) + return PathListing(paths=tuple(sorted(paths)), complete=False) + + @staticmethod + def _listing_stopped( + paths: set[str], page_number: int, where: str, reason: str + ) -> PathListing | None: + """Log why a ``list_paths`` walk stopped early and return what it has. + + A failure on the first page means there is no listing at all, so the + result is ``None``. A failure on a later page keeps the paths already + read, marked ``complete=False``, because a partial listing still + helps the name search. + """ + if page_number == 1: + logger.debug("list_paths got %s for %s", reason, where) + return None + logger.debug( + "list_paths got %s at page %d for %s; keeping the %d paths read so far", + reason, page_number, where, len(paths), + ) + return PathListing(paths=tuple(sorted(paths)), complete=False) + def prune_inline_comments(self, ref: PRRef) -> int: """Delete prxref-attributed inline comments; returns the count removed. diff --git a/src/prxref/forges/gitlab.py b/src/prxref/forges/gitlab.py index 08e6620..13043fe 100644 --- a/src/prxref/forges/gitlab.py +++ b/src/prxref/forges/gitlab.py @@ -13,9 +13,11 @@ from prxref.forges._diff_render import render_diff_entries as _render_diff_entries from prxref.forges.base import ( ATTRIBUTION_MARKER, + MAX_LISTING_PAGES, SUMMARY_MARKER, FeedReadError, InlineComment, + PathListing, PRData, PRRef, Thread, @@ -568,6 +570,85 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: return None return content.decode("utf-8", errors="replace") + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """Return every file path in the repository at commit ``sha``, best-effort. + + Walks the paged ``repository/tree?recursive=true`` listing, + ``_PAGE_SIZE`` entries a page, requesting pages 1, 2, 3 and so on. + Only ``blob`` entries are kept, so directories (``tree``) and + submodules (``commit``) are dropped, and the paths are sorted and + deduplicated. The walk ends when a page's ``X-Next-Page`` header is + absent or empty, never on a short page. It reads at most + ``MAX_LISTING_PAGES`` pages; when the last page it reads still names + a next page, the paths read so far come back with ``complete=False``. + A failure on the first page (a transport failure, a non-2xx status, a + body that is not JSON, or one that is not a list) gives ``None``; the + same failure on a later page gives the paths read so far with + ``complete=False``. An empty ``sha`` gives ``None`` with no request. + Never raises. + """ + if not sha: + return None + headers = self._get_auth_headers() + url = f"{self._api_base(ref)}/repository/tree" + where = f"{self._project_path(ref)}@{sha}" + paths: set[str] = set() + for page_number in range(1, MAX_LISTING_PAGES + 1): + params: dict[str, int | str] = { + "recursive": "true", + "per_page": _PAGE_SIZE, + "page": page_number, + "ref": sha, + } + try: + resp = self._session.get( + url, headers=headers, params=params, timeout=_REQUEST_TIMEOUT + ) + except requests.RequestException as e: + return self._listing_stopped(paths, page_number, where, f"a transport failure ({e})") + if not resp.ok: + return self._listing_stopped(paths, page_number, where, f"HTTP {resp.status_code}") + try: + entries = resp.json() + except ValueError as e: + return self._listing_stopped(paths, page_number, where, f"a non-JSON body ({e})") + if not isinstance(entries, list): + return self._listing_stopped( + paths, page_number, where, f"a {type(entries).__name__} body, not a list" + ) + paths.update( + entry["path"] for entry in entries + if isinstance(entry, dict) and entry.get("type") == "blob" + and isinstance(entry.get("path"), str) and entry["path"] + ) + if not (resp.headers.get("X-Next-Page") or "").strip(): + return PathListing(paths=tuple(sorted(paths)), complete=True) + logger.debug( + "list_paths stopped at the %d-page cap for %s with %d paths", + MAX_LISTING_PAGES, where, len(paths), + ) + return PathListing(paths=tuple(sorted(paths)), complete=False) + + @staticmethod + def _listing_stopped( + paths: set[str], page_number: int, where: str, reason: str + ) -> PathListing | None: + """Log why a ``list_paths`` walk stopped early and return what it has. + + A failure on the first page means there is no listing at all, so the + result is ``None``. A failure on a later page keeps the paths already + read, marked ``complete=False``, because a partial listing still + helps the name search. + """ + if page_number == 1: + logger.debug("list_paths got %s for %s", reason, where) + return None + logger.debug( + "list_paths got %s at page %d for %s; keeping the %d paths read so far", + reason, page_number, where, len(paths), + ) + return PathListing(paths=tuple(sorted(paths)), complete=False) + def prune_inline_comments(self, ref: PRRef) -> int: """Delete prxref-attributed inline comments; returns the count removed. diff --git a/tests/test_issue_17_list_paths_paged.py b/tests/test_issue_17_list_paths_paged.py new file mode 100644 index 0000000..1a3463e --- /dev/null +++ b/tests/test_issue_17_list_paths_paged.py @@ -0,0 +1,480 @@ +"""Tests for the paged ``list_paths`` listings (issue #17, task T9). + +GitLab walks ``repository/tree?recursive=true`` page by page and stops when +``X-Next-Page`` is absent or empty. Bitbucket Server / Data Center walks the +documented ``/files?at=`` listing with ``start``/``limit`` and stops on +``isLastPage`` or a missing ``nextPageStart``; its shape is UNVERIFIED (no +public instance). Both cap the walk at ``MAX_LISTING_PAGES`` pages. +""" +from __future__ import annotations + +import inspect +import json +import logging +from unittest.mock import MagicMock + +import pytest +import requests +from requests.structures import CaseInsensitiveDict + +from prxref.forges import base, bitbucket_server, gitlab +from prxref.forges.base import PathListing + +SHA = "1234567890abcdef1234567890abcdef12345678" +REQUEST_TIMEOUT = (10.0, 30.0) + +GL_PR_URL = "https://gitlab.example.com/acme/platform/api/-/merge_requests/7" +GL_TREE_URL = "https://gitlab.example.com/api/v4/projects/acme%2Fplatform%2Fapi/repository/tree" + +BBS_PR_URL = "https://bitbucket.example.com/projects/PLAT/repos/api/pull-requests/42" +BBS_FILES_URL = "https://bitbucket.example.com/rest/api/1.0/projects/PLAT/repos/api/files" +BBS_PERSONAL_PR_URL = "https://bitbucket.example.com/users/jdoe/repos/scratch/pull-requests/7" +BBS_PERSONAL_FILES_URL = "https://bitbucket.example.com/rest/api/1.0/projects/~jdoe/repos/scratch/files" + +TOKEN_VARS = ( + "PRXREF_GITLAB_TOKEN", + "PRXREF_BITBUCKET_SERVER_TOKEN", + "PRXREF_BITBUCKET_TOKEN", + "PRXREF_BITBUCKET_SERVER_USER", + "PRXREF_BITBUCKET_SERVER_PASSWORD", +) + + +@pytest.fixture(autouse=True) +def _no_ambient_tokens(monkeypatch): + for name in TOKEN_VARS: + monkeypatch.delenv(name, raising=False) + + +def _mock_response(status_code=200, json_data=None, text="", content=None, headers=None): + resp = MagicMock(spec=requests.Response) + resp.status_code = status_code + resp.ok = 200 <= status_code < 300 + resp.headers = headers or {} + if json_data is not None: + resp.json.return_value = json_data + resp.text = json.dumps(json_data) + else: + resp.text = text + resp.json.side_effect = ValueError("No JSON") + resp.content = content if content is not None else resp.text.encode("utf-8") + resp.raise_for_status.side_effect = ( + None if resp.ok else requests.HTTPError(response=resp) + ) + return resp + + +def _session(*responses): + session = MagicMock(spec=requests.Session) + session.get.side_effect = list(responses) + return session + + +def _gl_ref(url=GL_PR_URL): + ref = gitlab.ForgeImpl.parse_pr_url(url) + assert ref is not None + return ref + + +def _bbs_ref(url=BBS_PR_URL): + ref = bitbucket_server.ForgeImpl.parse_pr_url(url) + assert ref is not None + return ref + + +def _entry(path, kind="blob"): + mode = {"blob": "100644", "tree": "040000", "commit": "160000"}[kind] + return {"id": "0" * 40, "name": path.rsplit("/", 1)[-1], "type": kind, "path": path, "mode": mode} + + +def _gl_page(entries, next_page=""): + return _mock_response(200, json_data=entries, headers=CaseInsensitiveDict({"x-next-page": next_page})) + + +def _bbs_page(values, *, start=0, last=True, next_start=None): + body = {"values": values, "size": len(values), "isLastPage": last, "start": start, "limit": 100} + if next_start is not None: + body["nextPageStart"] = next_start + return _mock_response(200, json_data=body) + + +def _sent(session, key): + return [c.kwargs["params"][key] for c in session.get.call_args_list] + + +def _failure(kind): + if kind == "http-404": + return _mock_response(404, json_data={"message": "404 Tree Not Found"}) + if kind == "http-500": + return _mock_response(500, text="internal error") + if kind == "transport": + return requests.ConnectionError("connection refused") + return _mock_response(200, text="not json") + + +FAILURE_KINDS = ["http-404", "http-500", "transport", "non-json"] + + +# --- shared --------------------------------------------------------------------- + + +class TestSharedCap: + @pytest.mark.parametrize("module", [gitlab, bitbucket_server]) + def test_each_adapter_reads_the_cap_from_base(self, module): + assert module.MAX_LISTING_PAGES is base.MAX_LISTING_PAGES + assert not [name for name in vars(module) if "LISTING" in name and name != "MAX_LISTING_PAGES"] + + @pytest.mark.parametrize("forge_cls", [gitlab.ForgeImpl, bitbucket_server.ForgeImpl]) + def test_the_method_has_the_protocol_signature_and_a_docstring(self, forge_cls): + signature = inspect.signature(forge_cls.list_paths) + params = signature.parameters + assert list(params) == ["self", "ref", "sha"] + assert params["sha"].kind is inspect.Parameter.KEYWORD_ONLY + assert signature.return_annotation == "PathListing | None" + doc = inspect.getdoc(forge_cls.list_paths) + assert "Never raises" in doc + assert "MAX_LISTING_PAGES" in doc + + +# --- GitLab ------------------------------------------------------------------- + + +class TestGitLabListPaths: + def test_one_page_keeps_file_paths_sorted_and_deduplicated(self, monkeypatch): + monkeypatch.setenv("PRXREF_GITLAB_TOKEN", "t0ken") + session = _session(_gl_page([ + _entry("src/b.py"), + _entry("README.md"), + _entry("src/a.py"), + _entry("src/a.py"), + ])) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert listing == PathListing(paths=("README.md", "src/a.py", "src/b.py"), complete=True) + assert session.get.call_count == 1 + call = session.get.call_args + assert call.args[0] == GL_TREE_URL + assert call.kwargs["params"] == {"recursive": "true", "per_page": 100, "page": 1, "ref": SHA} + assert call.kwargs["headers"] == {"PRIVATE-TOKEN": "t0ken"} + assert call.kwargs["timeout"] == REQUEST_TIMEOUT + + def test_no_token_sends_no_auth_header(self): + session = _session(_gl_page([_entry("a.py")])) + + gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert session.get.call_args.kwargs["headers"] == {} + + def test_the_page_size_is_the_module_page_size(self): + session = _session(_gl_page([_entry("a.py")])) + + gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert session.get.call_args.kwargs["params"]["per_page"] == gitlab._PAGE_SIZE + + def test_three_pages_are_walked_in_order(self): + session = _session( + _gl_page([_entry("z/last.py"), _entry("a.py")], next_page="2"), + _gl_page([_entry("m/mid.py"), _entry("a.py")], next_page="3"), + _gl_page([_entry("b.py")], next_page=""), + ) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert listing == PathListing(paths=("a.py", "b.py", "m/mid.py", "z/last.py"), complete=True) + assert _sent(session, "page") == [1, 2, 3] + assert {c.args[0] for c in session.get.call_args_list} == {GL_TREE_URL} + assert {c.kwargs["params"]["ref"] for c in session.get.call_args_list} == {SHA} + assert {c.kwargs["timeout"] for c in session.get.call_args_list} == {REQUEST_TIMEOUT} + + def test_a_short_page_that_names_a_next_page_is_still_followed(self): + short = [_entry("a.py")] + assert len(short) < gitlab._PAGE_SIZE + session = _session( + _gl_page([_entry("first.py")], next_page="2"), + _gl_page(short, next_page="3"), + _gl_page([_entry("third.py")], next_page=""), + ) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert session.get.call_count == 3 + assert listing == PathListing(paths=("a.py", "first.py", "third.py"), complete=True) + + def test_the_page_cap_stops_the_walk_and_marks_the_listing_partial(self, monkeypatch): + monkeypatch.setattr(gitlab, "MAX_LISTING_PAGES", 2) + session = _session( + _gl_page([_entry("one.py")], next_page="2"), + _gl_page([_entry("two.py")], next_page="3"), + _gl_page([_entry("three.py")], next_page=""), + ) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert session.get.call_count == 2 + assert listing == PathListing(paths=("one.py", "two.py"), complete=False) + + def test_a_walk_that_ends_exactly_at_the_cap_is_whole(self, monkeypatch): + monkeypatch.setattr(gitlab, "MAX_LISTING_PAGES", 2) + session = _session( + _gl_page([_entry("one.py")], next_page="2"), + _gl_page([_entry("two.py")], next_page=""), + ) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert listing == PathListing(paths=("one.py", "two.py"), complete=True) + + @pytest.mark.parametrize("kind", FAILURE_KINDS) + def test_a_first_page_failure_gives_none(self, kind, caplog): + session = _session(_failure(kind)) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.gitlab"): + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert listing is None + records = [r for r in caplog.records if "list_paths" in r.getMessage()] + assert records and {r.levelno for r in records} == {logging.DEBUG} + + def test_a_wrong_shape_first_page_gives_none(self): + session = _session(_mock_response(200, json_data={"message": "not a list"})) + + assert gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) is None + + @pytest.mark.parametrize("kind", [*FAILURE_KINDS, "wrong-shape"]) + def test_a_second_page_failure_keeps_the_first_page_as_partial(self, kind, caplog): + second = ( + _mock_response(200, json_data={"message": "not a list"}) + if kind == "wrong-shape" else _failure(kind) + ) + session = _session(_gl_page([_entry("b.py"), _entry("a.py")], next_page="2"), second) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.gitlab"): + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert listing == PathListing(paths=("a.py", "b.py"), complete=False) + assert session.get.call_count == 2 + records = [r for r in caplog.records if "list_paths" in r.getMessage()] + assert records and {r.levelno for r in records} == {logging.DEBUG} + + def test_an_empty_sha_makes_no_request(self): + session = _session() + + assert gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha="") is None + session.get.assert_not_called() + + def test_directories_submodules_and_malformed_entries_are_dropped(self): + session = _session(_gl_page([ + _entry("src", kind="tree"), + _entry("vendor/lib", kind="commit"), + _entry("src/app.py"), + {"type": "blob", "path": ""}, + {"type": "blob", "path": None}, + {"type": "blob"}, + "not-a-dict", + ])) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert listing == PathListing(paths=("src/app.py",), complete=True) + + @pytest.mark.parametrize( + "headers", + [None, CaseInsensitiveDict({"x-next-page": ""}), CaseInsensitiveDict({"X-Next-Page": " "})], + ids=["header-absent", "header-empty", "header-blank"], + ) + def test_an_absent_or_empty_next_page_header_ends_the_walk(self, headers): + session = _session( + _mock_response(200, json_data=[_entry("a.py")] * 100, headers=headers), + _gl_page([_entry("never.py")]), + ) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert session.get.call_count == 1 + assert listing == PathListing(paths=("a.py",), complete=True) + + def test_the_next_page_header_is_read_case_insensitively(self): + session = _session( + _mock_response(200, json_data=[_entry("a.py")], headers=CaseInsensitiveDict({"X-NEXT-PAGE": "2"})), + _gl_page([_entry("b.py")]), + ) + + listing = gitlab.ForgeImpl(session=session).list_paths(_gl_ref(), sha=SHA) + + assert session.get.call_count == 2 + assert listing == PathListing(paths=("a.py", "b.py"), complete=True) + + +# --- Bitbucket Server / Data Center --------------------------------------------- + + +class TestBitbucketServerListPaths: + def test_one_page_keeps_file_paths_sorted_and_deduplicated(self, monkeypatch): + monkeypatch.setenv("PRXREF_BITBUCKET_SERVER_TOKEN", "t0ken") + session = _session(_bbs_page(["src/b.java", "README.md", "src/a.java", "src/a.java"])) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing == PathListing(paths=("README.md", "src/a.java", "src/b.java"), complete=True) + assert session.get.call_count == 1 + call = session.get.call_args + assert call.args[0] == BBS_FILES_URL + assert call.kwargs["params"] == {"at": SHA, "start": 0, "limit": 100} + assert call.kwargs["headers"] == {"Authorization": "Bearer t0ken"} + assert call.kwargs["auth"] is None + assert call.kwargs["timeout"] == REQUEST_TIMEOUT + + def test_basic_auth_is_sent_as_the_auth_pair(self, monkeypatch): + monkeypatch.setenv("PRXREF_BITBUCKET_SERVER_USER", "svc") + monkeypatch.setenv("PRXREF_BITBUCKET_SERVER_PASSWORD", "pw") + session = _session(_bbs_page(["a.py"])) + + bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert session.get.call_args.kwargs["headers"] == {} + assert session.get.call_args.kwargs["auth"] == ("svc", "pw") + + def test_the_page_limit_is_the_module_page_limit(self): + session = _session(_bbs_page(["a.py"])) + + bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert session.get.call_args.kwargs["params"]["limit"] == bitbucket_server._PAGE_LIMIT + + def test_three_pages_are_walked_in_order(self): + session = _session( + _bbs_page(["z/last.py", "a.py"], start=0, last=False, next_start=100), + _bbs_page(["m/mid.py", "a.py"], start=100, last=False, next_start=200), + _bbs_page(["b.py"], start=200, last=True), + ) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing == PathListing(paths=("a.py", "b.py", "m/mid.py", "z/last.py"), complete=True) + assert _sent(session, "start") == [0, 100, 200] + assert {c.args[0] for c in session.get.call_args_list} == {BBS_FILES_URL} + assert {c.kwargs["params"]["at"] for c in session.get.call_args_list} == {SHA} + assert {c.kwargs["timeout"] for c in session.get.call_args_list} == {REQUEST_TIMEOUT} + + def test_the_next_start_comes_from_the_page_not_a_fixed_stride(self): + session = _session( + _bbs_page(["a.py"], start=0, last=False, next_start=37), + _bbs_page(["b.py"], start=37, last=True), + ) + + bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert _sent(session, "start") == [0, 37] + + def test_the_page_cap_stops_the_walk_and_marks_the_listing_partial(self, monkeypatch): + monkeypatch.setattr(bitbucket_server, "MAX_LISTING_PAGES", 2) + session = _session( + _bbs_page(["one.py"], start=0, last=False, next_start=100), + _bbs_page(["two.py"], start=100, last=False, next_start=200), + _bbs_page(["three.py"], start=200, last=True), + ) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert session.get.call_count == 2 + assert listing == PathListing(paths=("one.py", "two.py"), complete=False) + + def test_a_walk_that_ends_exactly_at_the_cap_is_whole(self, monkeypatch): + monkeypatch.setattr(bitbucket_server, "MAX_LISTING_PAGES", 2) + session = _session( + _bbs_page(["one.py"], start=0, last=False, next_start=100), + _bbs_page(["two.py"], start=100, last=True), + ) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing == PathListing(paths=("one.py", "two.py"), complete=True) + + @pytest.mark.parametrize("kind", FAILURE_KINDS) + def test_a_first_page_failure_gives_none(self, kind, caplog): + session = _session(_failure(kind)) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.bitbucket_server"): + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing is None + records = [r for r in caplog.records if "list_paths" in r.getMessage()] + assert records and {r.levelno for r in records} == {logging.DEBUG} + + @pytest.mark.parametrize( + "body", + [["a.py"], {"size": 0, "isLastPage": True}, {"values": "a.py", "isLastPage": True}], + ids=["list-body", "no-values", "values-not-a-list"], + ) + def test_a_wrong_shape_first_page_gives_none(self, body): + session = _session(_mock_response(200, json_data=body)) + + assert bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) is None + + @pytest.mark.parametrize("kind", [*FAILURE_KINDS, "wrong-shape"]) + def test_a_second_page_failure_keeps_the_first_page_as_partial(self, kind, caplog): + second = ( + _mock_response(200, json_data={"errors": [{"message": "nope"}]}) + if kind == "wrong-shape" else _failure(kind) + ) + session = _session(_bbs_page(["b.py", "a.py"], start=0, last=False, next_start=100), second) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.bitbucket_server"): + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing == PathListing(paths=("a.py", "b.py"), complete=False) + assert session.get.call_count == 2 + records = [r for r in caplog.records if "list_paths" in r.getMessage()] + assert records and {r.levelno for r in records} == {logging.DEBUG} + + def test_an_empty_sha_makes_no_request(self): + session = _session() + + assert bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha="") is None + session.get.assert_not_called() + + def test_a_personal_repository_hits_the_tilde_url(self): + session = _session(_bbs_page(["notes.md"])) + ref = _bbs_ref(BBS_PERSONAL_PR_URL) + assert ref.owner == "~jdoe" + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(ref, sha=SHA) + + assert listing == PathListing(paths=("notes.md",), complete=True) + assert session.get.call_args.args[0] == BBS_PERSONAL_FILES_URL + + def test_empty_and_non_string_values_are_dropped(self): + session = _session(_bbs_page(["src/app.py", "", None, 7, {"path": "x.py"}])) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing == PathListing(paths=("src/app.py",), complete=True) + + def test_a_page_without_is_last_page_ends_the_walk_whole(self): + body = {"values": ["a.py"], "size": 1, "start": 0, "limit": 100, "nextPageStart": 100} + session = _session(_mock_response(200, json_data=body), _bbs_page(["never.py"])) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert session.get.call_count == 1 + assert listing == PathListing(paths=("a.py",), complete=True) + + def test_a_non_last_page_with_no_next_start_ends_the_walk_partial(self): + session = _session( + _bbs_page(["a.py"], start=0, last=False, next_start=None), + _bbs_page(["never.py"]), + ) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert session.get.call_count == 1 + assert listing == PathListing(paths=("a.py",), complete=False) + + def test_an_empty_listing_is_whole_and_empty(self): + session = _session(_bbs_page([])) + + listing = bitbucket_server.ForgeImpl(session=session).list_paths(_bbs_ref(), sha=SHA) + + assert listing == PathListing(paths=(), complete=True) From cd792cb1afbeff950018a704b1b2436c0f3a6053 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:09:24 -0700 Subject: [PATCH 07/24] feat: list_paths for Bitbucket Cloud and Azure DevOps (#17, T10) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Bitbucket Cloud walks GET /2.0/repositories/{owner}/{repo}/src/{sha}/ with max_depth=64 and pagelen=100, then follows each page's next URL verbatim with no params of its own, capped at MAX_LISTING_PAGES. Only commit_file entries are kept. The listing is complete=False when the cap stops a walk with a next still pending, when a later page cannot be read (the paths read so far are returned), or when a commit_directory sits at the depth limit. That last rule is inferred from the orchestrator's max_depth probe, not documented by Bitbucket, and the docstring says so. Azure DevOps makes one items?recursionLevel=Full request at the commit through the adapter's _get and _json helpers, keeps blob entries that are not folders, and strips the leading slash Azure puts on every path. An x-ms-continuationtoken header marks the listing complete=False and is not followed, because continuation on this endpoint is undocumented and was never observed. Both return None for an empty sha (no request) and for a failed first request, log only at DEBUG, and never raise. Tests are in tests/test_issue_17_list_paths_bbc_ado.py. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/forges/azure_devops.py | 54 +++ src/prxref/forges/bitbucket.py | 109 +++++ tests/test_issue_17_list_paths_bbc_ado.py | 506 ++++++++++++++++++++++ 3 files changed, 669 insertions(+) create mode 100644 tests/test_issue_17_list_paths_bbc_ado.py diff --git a/src/prxref/forges/azure_devops.py b/src/prxref/forges/azure_devops.py index 626dd09..0f60fa0 100644 --- a/src/prxref/forges/azure_devops.py +++ b/src/prxref/forges/azure_devops.py @@ -35,6 +35,7 @@ SUMMARY_MARKER, FeedReadError, InlineComment, + PathListing, PRData, PRRef, Thread, @@ -642,6 +643,59 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: return None return bytes(buf).decode("utf-8", errors="replace") + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """Return every file path in the repository at commit ``sha``, best-effort. + + One request to the ``items`` endpoint with ``recursionLevel=Full`` at + the commit, starting at the repository root, with no paging. Only + ``blob`` entries that are not flagged ``isFolder`` are kept, so the + root (``/``), directories (``tree``) and every other object type, + submodules included, are dropped. Azure DevOps returns each path with + a leading slash, which is stripped; the paths are sorted and + deduplicated. + + ``complete`` is ``False`` when the response carries an + ``x-ms-continuationtoken`` header. Continuation on this endpoint is + undocumented and was never observed, so the token is not followed: the + paths already returned are kept and the listing is marked incomplete. + An empty ``sha`` (no request is made), a transport failure, a non-2xx + status, a 203 sign-in page or other non-JSON body, or a body with no + ``value`` list gives ``None``. Never raises. + """ + if not sha: + return None + where = f"{ref.owner}/{ref.repo}@{sha}" + try: + resp = self._get( + f"{self._api_base(ref)}/items", + { + "recursionLevel": "Full", + "versionDescriptor.version": sha, + "versionDescriptor.versionType": "commit", + }, + ) + body = self._json(resp, "item listing") + except (requests.RequestException, ValueError) as e: + logger.debug("list_paths failed for %s: %s", where, e) + return None + values = body.get("value") + if not isinstance(values, list): + logger.debug("list_paths got no value list for %s", where) + return None + paths = { + entry["path"].lstrip("/") for entry in values + if isinstance(entry, dict) and entry.get("gitObjectType") == "blob" + and not entry.get("isFolder") + and isinstance(entry.get("path"), str) and entry["path"].lstrip("/") + } + complete = not resp.headers.get("x-ms-continuationtoken") + if not complete: + logger.debug( + "list_paths got an x-ms-continuationtoken for %s and did not follow it; " + "the listing is incomplete", where, + ) + return PathListing(paths=tuple(sorted(paths)), complete=complete) + def _read_threads(self, ref: PRRef) -> list[dict]: """Read every thread on the PR (one response; the API does not page them). diff --git a/src/prxref/forges/bitbucket.py b/src/prxref/forges/bitbucket.py index 1e7ab8d..7c91dac 100644 --- a/src/prxref/forges/bitbucket.py +++ b/src/prxref/forges/bitbucket.py @@ -13,10 +13,12 @@ from prxref.forges.base import ( ATTRIBUTION_MARKER, + MAX_LISTING_PAGES, SUMMARY_MARKER, DescriptionVersion, FeedReadError, InlineComment, + PathListing, PRData, PRHistory, PRRef, @@ -36,6 +38,9 @@ # 500 comments. 50 puts the ceiling far past any real PR, and running out of # budget is now a refusal to post rather than an invisible short read. _MAX_PAGES = 50 +# list_paths asks /src for this many directory levels; the orchestrator's +# probe saw 64 accepted and the whole tree returned. +_LISTING_MAX_DEPTH = 64 # get_file_content is best-effort context, not the review itself: a body past # this size (or one that looks binary) is worth skipping rather than shipping # hundreds of KB into a worker prompt. @@ -467,6 +472,110 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: return None return content.decode("utf-8", errors="replace") + def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: + """Return every file path in the repository at commit ``sha``, best-effort. + + Reads the repository-level ``/src/{sha}/`` listing with ``max_depth`` + set to ``_LISTING_MAX_DEPTH`` and ``pagelen`` to ``_PAGE_SIZE``, then + follows each page's ``next`` URL verbatim, adding no parameters of its + own, until a page carries no ``next``. Only ``commit_file`` entries are + kept, so directories (``commit_directory``) and every other type, + submodules included, are dropped; the paths are sorted and + deduplicated. + + ``complete`` is ``False`` when the walk stops at ``MAX_LISTING_PAGES`` + pages with a ``next`` still to follow, when a page after the first + cannot be read (the paths read so far are returned), or when a + returned ``commit_directory`` sits at the depth limit, its path + holding ``_LISTING_MAX_DEPTH - 1`` slashes. That depth rule is + inferred from probing, not documented by Bitbucket: ``max_depth=N`` + was observed to list N directory levels, so a directory on level N is + listed without its contents. An empty ``sha`` (no request is made) + gives ``None``, as does a first request that fails in transport, + returns a non-2xx status, or returns a body that is not JSON or holds + no ``values`` list. Never raises. + """ + if not sha: + return None + headers, auth = self._get_auth() + url = f"{_API_BASE}/repositories/{ref.owner}/{ref.repo}/src/{quote(sha, safe='')}/" + params: dict[str, int] | None = {"max_depth": _LISTING_MAX_DEPTH, "pagelen": _PAGE_SIZE} + where = f"{ref.owner}/{ref.repo}@{sha}" + files: set[str] = set() + at_depth_limit = False + for page in range(MAX_LISTING_PAGES): + body = self._read_listing_page(url, params, headers, auth, where) + if body is None: + if page == 0: + return None + return PathListing(paths=tuple(sorted(files)), complete=False) + for entry in body["values"]: + if not isinstance(entry, dict): + continue + path = entry.get("path") + if not isinstance(path, str) or not path: + continue + if entry.get("type") == "commit_file": + files.add(path) + elif ( + entry.get("type") == "commit_directory" + and path.rstrip("/").count("/") >= _LISTING_MAX_DEPTH - 1 + ): + at_depth_limit = True + next_url = body.get("next") + if not next_url: + if at_depth_limit: + logger.debug( + "list_paths reached max_depth=%d for %s; the listing is incomplete", + _LISTING_MAX_DEPTH, where, + ) + return PathListing(paths=tuple(sorted(files)), complete=not at_depth_limit) + url = next_url + params = None + logger.debug( + "list_paths stopped at the %d-page cap for %s; the listing is incomplete", + MAX_LISTING_PAGES, where, + ) + return PathListing(paths=tuple(sorted(files)), complete=False) + + def _read_listing_page( + self, + url: str, + params: dict[str, int] | None, + headers: dict[str, str], + auth: tuple[str, str] | None, + where: str, + ) -> dict | None: + """Return one ``/src`` listing page, or ``None`` (logged at DEBUG) when it cannot be read. + + A page is readable when the request succeeds with a 2xx status and a + JSON object holding a ``values`` list and, if it has one, a string + ``next``. + """ + try: + resp = self._session.get( + url, params=params, headers=headers, auth=auth, timeout=_REQUEST_TIMEOUT + ) + except requests.RequestException as e: + logger.debug("list_paths failed for %s: %s", where, e) + return None + if not resp.ok: + logger.debug("list_paths got HTTP %s for %s", resp.status_code, where) + return None + try: + body = resp.json() + except ValueError as e: + logger.debug("list_paths got a non-JSON body for %s: %s", where, e) + return None + if ( + not isinstance(body, dict) + or not isinstance(body.get("values"), list) + or not isinstance(body.get("next") or "", str) + ): + logger.debug("list_paths got no values list for %s", where) + return None + return body + def prune_inline_comments(self, ref: PRRef) -> int: """Delete prxref-attributed inline comments; returns the count removed. diff --git a/tests/test_issue_17_list_paths_bbc_ado.py b/tests/test_issue_17_list_paths_bbc_ado.py new file mode 100644 index 0000000..7ee17c2 --- /dev/null +++ b/tests/test_issue_17_list_paths_bbc_ado.py @@ -0,0 +1,506 @@ +"""Tests for ``list_paths`` on Bitbucket Cloud and Azure DevOps (issue #17, task T10). + +Bitbucket Cloud walks the paged ``/src/{sha}/`` listing, following ``next`` +verbatim; Azure DevOps answers from one ``items?recursionLevel=Full`` request. +The entry shapes are the ones the orchestrator's read-only probes observed. +""" +from __future__ import annotations + +import json +import logging +from unittest.mock import MagicMock + +import pytest +import requests + +from prxref.forges import base, bitbucket +from prxref.forges.azure_devops import ForgeImpl as AzureForge +from prxref.forges.base import PathListing, PRRef +from prxref.forges.bitbucket import ForgeImpl as BitbucketForge + +SHA = "0123456789abcdef0123456789abcdef01234567" +REQUEST_TIMEOUT = (10.0, 30.0) + +BB_PR_URL = "https://bitbucket.org/acme/api/pull-requests/42" +BB_SRC = f"https://api.bitbucket.org/2.0/repositories/acme/api/src/{SHA}/" + +ADO_PR_URL = "https://dev.azure.com/acme/AcmeWeb/_git/AcmeWeb/pullrequest/551" +ADO_ITEMS = "https://dev.azure.com/acme/AcmeWeb/_apis/git/repositories/AcmeWeb/items" +ADO_JSON = "application/json; charset=utf-8; api-version=7.1" + + +@pytest.fixture(autouse=True) +def _no_forge_credentials(monkeypatch): + """conftest clears PRXREF_* only; the Azure pipeline token is not one of them.""" + monkeypatch.delenv("SYSTEM_ACCESSTOKEN", raising=False) + monkeypatch.delenv("PRXREF_AZURE_DEVOPS_TOKEN", raising=False) + monkeypatch.delenv("PRXREF_BITBUCKET_TOKEN", raising=False) + monkeypatch.delenv("PRXREF_BITBUCKET_USER", raising=False) + monkeypatch.delenv("PRXREF_BITBUCKET_APP_PASSWORD", raising=False) + + +def _mock_response(status_code=200, json_data=None, text="", content=None, headers=None): + resp = MagicMock(spec=requests.Response) + resp.status_code = status_code + resp.ok = 200 <= status_code < 300 + resp.headers = headers or {} + if json_data is not None: + resp.json.return_value = json_data + resp.text = json.dumps(json_data) + else: + resp.text = text + resp.json.side_effect = ValueError("No JSON") + resp.content = content if content is not None else resp.text.encode("utf-8") + resp.raise_for_status.side_effect = ( + None if resp.ok else requests.HTTPError(response=resp) + ) + return resp + + +def _session(*responses): + session = MagicMock(spec=requests.Session) + session.get.side_effect = list(responses) + return session + + +# --- Bitbucket Cloud ----------------------------------------------------------- + + +def _bb_ref(url=BB_PR_URL): + ref = BitbucketForge.parse_pr_url(url) + assert ref is not None + return ref + + +def _bb_file(path): + return { + "path": path, + "commit": {"hash": SHA, "type": "commit"}, + "type": "commit_file", + "attributes": [], + "escaped_path": path, + "size": 10, + "mimetype": None, + "links": {"self": {"href": f"{BB_SRC}{path}"}}, + } + + +def _bb_dir(path): + return { + "path": path, + "commit": {"hash": SHA, "type": "commit"}, + "type": "commit_directory", + "links": {"self": {"href": f"{BB_SRC}{path}/"}}, + } + + +def _bb_page(values, next_url=None, page=1): + body = {"page": page, "pagelen": 100, "values": values} + if next_url: + body["next"] = next_url + return _mock_response(json_data=body) + + +def _bb_next(token): + return f"{BB_SRC}?max_depth=64&pagelen=100&page={token}" + + +class TestBitbucketCloudListPaths: + def test_one_page_keeps_only_commit_file_paths_sorted_and_deduplicated(self, monkeypatch): + monkeypatch.setenv("PRXREF_BITBUCKET_TOKEN", "t0ken") + values = [ + _bb_file("src/b.py"), + _bb_dir("src"), + {"path": "vendor/lib", "type": "commit_submodule"}, + _bb_file("README.md"), + _bb_file("src/a.py"), + _bb_file("src/a.py"), + ] + session = _session(_bb_page(values)) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert listing == PathListing(paths=("README.md", "src/a.py", "src/b.py"), complete=True) + session.get.assert_called_once() + call = session.get.call_args + assert call.args[0] == BB_SRC + assert call.kwargs["params"] == {"max_depth": 64, "pagelen": 100} + assert call.kwargs["timeout"] == REQUEST_TIMEOUT + assert call.kwargs["headers"] == {"Authorization": "Bearer t0ken"} + assert call.kwargs["auth"] is None + + def test_app_password_credentials_go_out_as_basic_auth(self, monkeypatch): + monkeypatch.setenv("PRXREF_BITBUCKET_USER", "bot") + monkeypatch.setenv("PRXREF_BITBUCKET_APP_PASSWORD", "app-pw") + session = _session(_bb_page([_bb_file("a.py")])) + + BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + call = session.get.call_args + assert call.kwargs["headers"] == {} + assert call.kwargs["auth"] == ("bot", "app-pw") + + def test_three_pages_follow_next_verbatim_and_union_the_paths(self, monkeypatch): + monkeypatch.setenv("PRXREF_BITBUCKET_TOKEN", "t0ken") + second, third = _bb_next("rXtr"), _bb_next("sYus") + session = _session( + _bb_page([_bb_file("c.py"), _bb_file("a.py")], next_url=second, page=1), + _bb_page([_bb_dir("pkg"), _bb_file("pkg/b.py")], next_url=third, page=2), + _bb_page([_bb_file("a.py"), _bb_file("pkg/d.py")], page=3), + ) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert listing == PathListing(paths=("a.py", "c.py", "pkg/b.py", "pkg/d.py"), complete=True) + calls = session.get.call_args_list + assert [c.args[0] for c in calls] == [BB_SRC, second, third] + assert calls[0].kwargs["params"] == {"max_depth": 64, "pagelen": 100} + for follow_up in calls[1:]: + assert not follow_up.kwargs.get("params") + assert follow_up.kwargs["headers"] == {"Authorization": "Bearer t0ken"} + assert follow_up.kwargs["timeout"] == REQUEST_TIMEOUT + + def test_the_depth_limit_is_sixty_four(self): + assert bitbucket._LISTING_MAX_DEPTH == 64 + + @pytest.mark.parametrize( + ("directory", "complete"), + [("a/b", False), ("a/b/", False), ("a", True)], + ids=["at-the-limit", "at-the-limit-trailing-slash", "one-below-the-limit"], + ) + def test_a_directory_at_the_depth_limit_marks_the_listing_incomplete( + self, monkeypatch, directory, complete + ): + monkeypatch.setattr(bitbucket, "_LISTING_MAX_DEPTH", 2) + session = _session(_bb_page([_bb_file("a/x.py"), _bb_dir(directory)])) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert listing == PathListing(paths=("a/x.py",), complete=complete) + assert session.get.call_args.kwargs["params"]["max_depth"] == 2 + + def test_a_directory_at_the_default_depth_limit_marks_the_listing_incomplete(self): + deep = "/".join(f"d{i}" for i in range(64)) + shallow = "/".join(f"d{i}" for i in range(63)) + assert deep.count("/") == 63 + session = _session(_bb_page([_bb_dir(shallow)]), _bb_page([_bb_dir(deep)])) + forge = BitbucketForge(session=session) + + assert forge.list_paths(_bb_ref(), sha=SHA) == PathListing(paths=(), complete=True) + assert forge.list_paths(_bb_ref(), sha=SHA) == PathListing(paths=(), complete=False) + + def test_a_directory_at_the_limit_on_an_earlier_page_still_counts(self, monkeypatch): + monkeypatch.setattr(bitbucket, "_LISTING_MAX_DEPTH", 2) + session = _session( + _bb_page([_bb_dir("a/b"), _bb_file("a/x.py")], next_url=_bb_next("p2")), + _bb_page([_bb_file("z.py")], page=2), + ) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert listing == PathListing(paths=("a/x.py", "z.py"), complete=False) + + def test_the_page_cap_is_the_shared_constant(self): + assert bitbucket.MAX_LISTING_PAGES is base.MAX_LISTING_PAGES + + def test_the_page_cap_stops_the_walk_and_marks_it_incomplete(self, monkeypatch): + monkeypatch.setattr(bitbucket, "MAX_LISTING_PAGES", 2) + session = _session( + _bb_page([_bb_file("a.py")], next_url=_bb_next("p2"), page=1), + _bb_page([_bb_file("b.py")], next_url=_bb_next("p3"), page=2), + _bb_page([_bb_file("c.py")], page=3), + ) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert session.get.call_count == 2 + assert listing == PathListing(paths=("a.py", "b.py"), complete=False) + + def test_a_walk_that_ends_on_the_last_allowed_page_is_complete(self, monkeypatch): + monkeypatch.setattr(bitbucket, "MAX_LISTING_PAGES", 2) + session = _session( + _bb_page([_bb_file("a.py")], next_url=_bb_next("p2"), page=1), + _bb_page([_bb_file("b.py")], page=2), + ) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert session.get.call_count == 2 + assert listing == PathListing(paths=("a.py", "b.py"), complete=True) + + def test_malformed_entries_are_skipped(self): + values = [ + "not-an-entry", + {"type": "commit_file"}, + {"type": "commit_file", "path": ""}, + {"type": "commit_file", "path": 7}, + {"type": "commit_directory"}, + _bb_file("kept.py"), + ] + session = _session(_bb_page(values)) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert listing == PathListing(paths=("kept.py",), complete=True) + + def test_an_empty_listing_is_an_empty_complete_listing(self): + session = _session(_bb_page([])) + + assert BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) == PathListing( + paths=(), complete=True + ) + + def test_the_sha_is_quoted_into_the_url(self): + session = _session(_bb_page([])) + + BitbucketForge(session=session).list_paths(_bb_ref(), sha="feat/x y") + + assert session.get.call_args.args[0] == ( + "https://api.bitbucket.org/2.0/repositories/acme/api/src/feat%2Fx%20y/" + ) + + @pytest.mark.parametrize( + "failure", + [ + _mock_response(404, json_data={"type": "error", "error": {"message": "Commit not found"}}), + _mock_response(500, text="boom"), + requests.ConnectionError("down"), + requests.Timeout("slow"), + _mock_response(text="gateway"), + _mock_response(json_data={"page": 1, "pagelen": 100}), + _mock_response(json_data={"page": 1, "pagelen": 100, "values": None}), + _mock_response(json_data={"page": 1, "pagelen": 100, "values": {"path": "a.py"}}), + _mock_response(json_data=[_bb_file("a.py")]), + _mock_response(json_data={"values": [_bb_file("a.py")], "next": 7}), + ], + ids=[ + "http-404", "http-500", "connection-error", "timeout", "non-json", + "missing-values", "null-values", "values-not-a-list", "body-not-an-object", "next-not-a-string", + ], + ) + def test_a_first_page_failure_gives_none(self, failure): + session = _session(failure) + + assert BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) is None + session.get.assert_called_once() + + @pytest.mark.parametrize( + "failure", + [ + _mock_response(500, text="boom"), + requests.ConnectionError("down"), + _mock_response(text="gateway"), + _mock_response(json_data={"page": 2, "pagelen": 100}), + ], + ids=["http-500", "connection-error", "non-json", "missing-values"], + ) + def test_a_later_page_failure_keeps_the_paths_so_far_as_incomplete(self, failure): + session = _session( + _bb_page([_bb_file("b.py"), _bb_file("a.py")], next_url=_bb_next("p2")), + failure, + ) + + listing = BitbucketForge(session=session).list_paths(_bb_ref(), sha=SHA) + + assert listing == PathListing(paths=("a.py", "b.py"), complete=False) + assert session.get.call_count == 2 + + def test_an_empty_sha_gives_none_without_a_request(self): + session = MagicMock(spec=requests.Session) + + assert BitbucketForge(session=session).list_paths(_bb_ref(), sha="") is None + session.get.assert_not_called() + + def test_failures_never_log_above_debug(self, caplog): + session = _session( + requests.ConnectionError("down"), + _mock_response(500, text="boom"), + _mock_response(text="not json"), + _mock_response(json_data={"page": 1}), + _bb_page([_bb_file("a.py")], next_url=_bb_next("p2")), + requests.ConnectionError("down"), + ) + forge = BitbucketForge(session=session) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.bitbucket"): + results = [forge.list_paths(_bb_ref(), sha=SHA) for _ in range(5)] + + assert results == [None, None, None, None, PathListing(paths=("a.py",), complete=False)] + assert len(caplog.records) == 5 + assert all(record.levelno <= logging.DEBUG for record in caplog.records) + + +# --- Azure DevOps -------------------------------------------------------------- + + +def _ado_ref(url=ADO_PR_URL): + ref = AzureForge.parse_pr_url(url) + assert ref is not None + return ref + + +def _ado_entry(path, git_object_type="blob", is_folder=None): + entry = { + "objectId": "1" * 40, + "gitObjectType": git_object_type, + "commitId": SHA, + "path": path, + "url": f"{ADO_ITEMS}?path={path}", + } + if is_folder is not None: + entry["isFolder"] = is_folder + return entry + + +def _ado_listing(values, headers=None): + return _mock_response( + json_data={"count": len(values), "value": values}, + headers={"Content-Type": ADO_JSON, **(headers or {})}, + ) + + +ADO_VALUES = [ + _ado_entry("/", "tree", is_folder=True), + _ado_entry("/src/b.py"), + _ado_entry("/src", "tree", is_folder=True), + _ado_entry("/.order"), + _ado_entry("/vendor/lib", "commit"), + _ado_entry("/src/a.py"), + _ado_entry("/src/a.py"), + _ado_entry("/odd", "blob", is_folder=True), +] +ADO_PATHS = (".order", "src/a.py", "src/b.py") + + +class TestAzureDevOpsListPaths: + def test_it_keeps_blob_paths_without_the_leading_slash_sorted_and_deduplicated(self): + session = _session(_ado_listing(ADO_VALUES)) + + listing = AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) + + assert listing == PathListing(paths=ADO_PATHS, complete=True) + session.get.assert_called_once() + call = session.get.call_args + assert call.args[0] == ADO_ITEMS + assert call.kwargs["params"] == { + "api-version": "7.1", + "recursionLevel": "Full", + "versionDescriptor.version": SHA, + "versionDescriptor.versionType": "commit", + } + assert call.kwargs["timeout"] == REQUEST_TIMEOUT + assert call.kwargs["headers"]["Accept"] == "application/json" + assert call.kwargs["headers"]["X-TFS-FedAuthRedirect"] == "Suppress" + assert "Authorization" not in call.kwargs["headers"] + + def test_the_personal_access_token_goes_out_as_basic_auth(self, monkeypatch): + monkeypatch.setenv("PRXREF_AZURE_DEVOPS_TOKEN", "p4t") + session = _session(_ado_listing(ADO_VALUES)) + + AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) + + assert session.get.call_args.kwargs["headers"]["Authorization"] == "Basic OnA0dA==" + + def test_the_live_probe_entry_shapes_are_read(self): + values = [ + {"objectId": "2" * 40, "gitObjectType": "tree", "commitId": SHA, "path": "/", + "isFolder": True, "url": ADO_ITEMS}, + {"objectId": "3" * 40, "gitObjectType": "blob", "commitId": SHA, "path": "/.order", + "url": ADO_ITEMS}, + ] + session = _session(_ado_listing(values)) + + listing = AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) + + assert listing == PathListing(paths=(".order",), complete=True) + + def test_a_continuation_token_marks_the_listing_incomplete_without_following_it(self): + session = _session( + _ado_listing(ADO_VALUES, headers={"x-ms-continuationtoken": "next-page"}), + _ado_listing([_ado_entry("/late.py")]), + ) + + listing = AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) + + assert listing == PathListing(paths=ADO_PATHS, complete=False) + session.get.assert_called_once() + + def test_malformed_entries_are_skipped(self): + values = [ + "not-an-entry", + {"gitObjectType": "blob"}, + {"gitObjectType": "blob", "path": ""}, + {"gitObjectType": "blob", "path": "/"}, + {"gitObjectType": "blob", "path": 7}, + _ado_entry("/kept.py"), + ] + session = _session(_ado_listing(values)) + + listing = AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) + + assert listing == PathListing(paths=("kept.py",), complete=True) + + def test_an_empty_value_list_is_an_empty_complete_listing(self): + session = _session(_ado_listing([])) + + assert AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) == PathListing( + paths=(), complete=True + ) + + @pytest.mark.parametrize( + "failure", + [ + _mock_response(404, json_data={"message": "TF401175"}, headers={"Content-Type": ADO_JSON}), + _mock_response(401, text="", headers={"Content-Type": ADO_JSON}), + _mock_response(500, text="boom", headers={"Content-Type": "text/plain"}), + _mock_response(203, text="sign in", headers={"Content-Type": "text/html"}), + _mock_response(200, text="sign in", headers={"Content-Type": "text/html; charset=utf-8"}), + requests.ConnectionError("reset"), + requests.Timeout("slow"), + _mock_response(200, text="{not json", headers={"Content-Type": ADO_JSON}), + _mock_response(json_data={"count": 0}, headers={"Content-Type": ADO_JSON}), + _mock_response(json_data={"count": 0, "value": None}, headers={"Content-Type": ADO_JSON}), + _mock_response(json_data={"count": 1, "value": {"path": "/a.py"}}, headers={"Content-Type": ADO_JSON}), + _mock_response(json_data=[_ado_entry("/a.py")], headers={"Content-Type": ADO_JSON}), + ], + ids=[ + "http-404", "http-401", "http-500", "http-203-sign-in", "http-200-html", "connection-error", + "timeout", "non-json", "missing-value", "null-value", "value-not-a-list", "body-not-an-object", + ], + ) + def test_a_failure_gives_none(self, failure): + session = _session(failure) + + assert AzureForge(session=session).list_paths(_ado_ref(), sha=SHA) is None + session.get.assert_called_once() + + def test_an_empty_sha_gives_none_without_a_request(self): + session = MagicMock(spec=requests.Session) + + assert AzureForge(session=session).list_paths(_ado_ref(), sha="") is None + session.get.assert_not_called() + + def test_a_foreign_ref_gives_none_without_a_request(self): + ref = PRRef(forge="azure-devops", host="github.com", owner="o", repo="r", number=1, + url="https://github.com/o/r/pull/1") + session = MagicMock(spec=requests.Session) + + assert AzureForge(session=session).list_paths(ref, sha=SHA) is None + session.get.assert_not_called() + + def test_failures_never_log_above_debug(self, caplog): + session = _session( + requests.ConnectionError("reset"), + _mock_response(404, json_data={"message": "TF401175"}, headers={"Content-Type": ADO_JSON}), + _mock_response(203, text="sign in", headers={"Content-Type": "text/html"}), + _mock_response(json_data={"count": 0}, headers={"Content-Type": ADO_JSON}), + _ado_listing(ADO_VALUES, headers={"x-ms-continuationtoken": "next-page"}), + ) + forge = AzureForge(session=session) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.azure_devops"): + results = [forge.list_paths(_ado_ref(), sha=SHA) for _ in range(5)] + + assert results == [None, None, None, None, PathListing(paths=ADO_PATHS, complete=False)] + assert len(caplog.records) == 5 + assert all(record.levelno <= logging.DEBUG for record in caplog.records) From 6eb22b6b4474260497c050899b122d679e63ea3f Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:11:34 -0700 Subject: [PATCH 08/24] feat: contract excerpters for repository context (#17, T4) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit New pure module src/prxref/repo_contracts.py (stdlib json and re only, no YAML library, no I/O) that cuts the slice of a contract file a changed route, operation, schema or table points at: - normalize_name folds idempotency_keys, idempotency-key and IdempotencyKey to one key; route_key makes {id}, :id, and routes compare equal to a spec path. - openapi_yaml_excerpts slices by indentation under paths:, the enclosing HTTP-method block of an operationId, and components.schemas (plus the Swagger 2.0 definitions:), skipping block-scalar text. - openapi_json_excerpts and json_schema_excerpts do the same through json.loads, with a best-effort line found by following the key path. - sql_excerpts finds CREATE TABLE, ALTER TABLE and CREATE [UNIQUE] INDEX ON statements; comments, quoted strings and $$ bodies do not end a statement, and a Liquibase --changeset line or a GO line does. - liquibase_excerpts slices XML, YAML and JSON changeSets holding tableName, baseTableName or referencedTableName. Every excerpt is capped at 40 lines and 2000 chars, and a cut block ends with an "N more lines" marker line that counts toward both caps. An excerpt that another one already shows in full is dropped. The regexes stay linear on long whitespace runs. Tests: tests/test_repo_context_contracts.py (63), with every fixture inline. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/repo_contracts.py | 681 +++++++++++++++++++++++++++ tests/test_repo_context_contracts.py | 639 +++++++++++++++++++++++++ 2 files changed, 1320 insertions(+) create mode 100644 src/prxref/repo_contracts.py create mode 100644 tests/test_repo_context_contracts.py diff --git a/src/prxref/repo_contracts.py b/src/prxref/repo_contracts.py new file mode 100644 index 0000000..8aa46cc --- /dev/null +++ b/src/prxref/repo_contracts.py @@ -0,0 +1,681 @@ +"""Contract excerpts for repository context: the OpenAPI, JSON Schema, SQL and +Liquibase slices that a changed route, operation, schema or table points at. + +A diff that adds a route handler, a migration or a DTO often changes behavior +that a contract file elsewhere in the repository pins down: the OpenAPI +operation for the route, the schema a payload must satisfy, the unique index on +a table. The excerpters here cut the matching slice out of such a file so a +worker can check the change against it. + +The module is pure and stdlib only. It performs no I/O: callers pass the file's +text. There is deliberately no YAML parser, because the core ships none, so +YAML is sliced by indentation, JSON goes through :mod:`json`, and SQL and XML +are scanned with regular expressions. Every excerpter degrades to ``[]`` on +text it cannot read rather than raising, because a review must never fail over +missing context. Each excerpt is capped at :data:`MAX_CONTRACT_LINES` lines and +:data:`MAX_CONTRACT_CHARS` characters. + +Names compare through :func:`normalize_name`, so a table called +``idempotency_keys`` matches a schema or class called ``IdempotencyKey``. +Routes compare through :func:`route_key`, so the code route +``/connectors/:id`` matches the spec path ``/connectors/{connectorId}``. +""" +from __future__ import annotations + +import json +import re +from collections.abc import Callable, Iterable, Iterator +from dataclasses import dataclass + +MAX_CONTRACT_LINES = 40 +MAX_CONTRACT_CHARS = 2000 + +_BOM = chr(0xFEFF) +_HTTP_METHODS = frozenset({"get", "put", "post", "delete", "options", "head", "patch", "trace"}) +_TABLE_KEYS = frozenset({"tableName", "baseTableName", "referencedTableName"}) + +_BRACE_PARAM = re.compile(r"\{[^{}/]*\}") +_ANGLE_PARAM = re.compile(r"<[^<>/]*>") + +_YAML_KEY = re.compile( + r"""(?P\ *)(?P(?:-[ \t]+)*)""" + r"""(?P"(?:[^"\\]|\\.)*"[ \t]*|'(?:[^']|'')*'[ \t]*""" + r"""|(?:[^\s#"'\-\[\]{},?&*!|>%@`]|-(?=\S))(?:[^#:]|:(?![ \t]|$))*):(?:[ \t]|$)""" +) +_YAML_QUOTED = re.compile(r""""(?:[^"\\]|\\.)*"|'(?:[^']|'')*'""") +_YAML_BLOCK_SCALAR = re.compile(r"(?:^[ \t]*(?:-[ \t]+)+|:[ \t]+)[|>][-+1-9]{0,2}[ \t]*(?:#.*)?$") + +_SQL_TOKEN = re.compile( + r"(?P--[ \t]*changeset\b[^\n]*)" + r"|(?P--[^\n]*|/\*.*?(?:\*/|\Z))" + r"|(?P'(?:[^']|'')*'|\"(?:[^\"]|\"\")*\"|`[^`]*`)" + r"|(?P\$(?P[A-Za-z_]\w*|)\$.*?\$(?P=tag)\$)" + r"|(?P^[ \t]*go[ \t\r]*$)" + r"|(?P;)", + re.S | re.M | re.I, +) +_SQL_PART = r"""(?:"[^"]+"|`[^`]+`|\[[^\]]+\]|[\w$]+)""" +_SQL_NAME = rf"(?P{_SQL_PART}(?:\s*\.\s*{_SQL_PART})*)" +_SQL_HEADS = ( + re.compile( + r"create\s+(?:or\s+replace\s+)?(?:(?:global|local)\s+)?(?:(?:temp|temporary|unlogged)\s+)?" + rf"table\s+(?:if\s+not\s+exists\s+)?{_SQL_NAME}", + re.I, + ), + re.compile(rf"alter\s+table\s+(?:if\s+exists\s+)?(?:only\s+)?{_SQL_NAME}", re.I), + re.compile( + r"create\s+(?:unique\s+)?(?:(?:clustered|nonclustered|bitmap|fulltext|spatial)\s+)?" + rf"index\b[^;]*?\bon\s+(?:only\s+)?{_SQL_NAME}", + re.I, + ), +) +_SQL_NAME_PART = re.compile(_SQL_PART) +_NON_SPACE = re.compile(r"\S") + +_XML_COMMENT = re.compile(r"|\Z)", re.S) +_XML_CHANGESET = re.compile( + r"""<(?:[\w.-]+:)?changeSet\b(?:"[^"]*"|'[^']*'|[^'">/]|/(?!>))*(?:/>|>.*?)""", + re.S, +) +_XML_TABLE_ATTR = re.compile(r"""\b(?:tableName|baseTableName|referencedTableName)\s*=\s*(?:"([^"]*)"|'([^']*)')""") + + +@dataclass(frozen=True) +class Excerpt: + """One slice of a contract file. + + ``line`` is the 1-based line the slice starts on (best effort for JSON, and + 1 when unknown), ``symbol`` is what matched, written as the file writes it, + and ``text`` is the slice, already capped to :data:`MAX_CONTRACT_LINES` and + :data:`MAX_CONTRACT_CHARS`. + """ + + line: int + symbol: str + text: str + + +def normalize_name(name: str) -> str: + """Fold a type, schema or table name so its spellings compare equal. + + The name is lowercased, every ``_`` and ``-`` is dropped, and ONE trailing + ``s`` is stripped, so ``idempotency_keys``, ``idempotency-key`` and + ``IdempotencyKey`` all give ``idempotencykey``. Excerpters treat a name + that folds to the empty string as matching nothing. + """ + folded = name.lower().replace("_", "").replace("-", "") + return folded[:-1] if folded.endswith("s") else folded + + +def route_key(route: str) -> str: + """Normalize an HTTP route template so a code route and a spec path compare equal. + + Every parameter becomes ``{}``: an OpenAPI or Spring ``{id}`` or + ``{connectorId}``, a Flask ```` or ````, and an Express + segment that starts with ``:`` (``:id``, ``:id?``). A parameter inside a + segment is replaced in place (``/files/{name}.json`` gives + ``/files/{}.json``). Empty segments are dropped, which collapses duplicate + slashes and strips a trailing slash, and the result always starts with + ``/``. Case is kept. ``/`` gives ``/``; a blank route gives ``""``, which + excerpters treat as matching nothing. + """ + segments = [segment for segment in route.strip().split("/") if segment] + if not segments: + return "/" if route.strip() else "" + keyed = [ + "{}" if segment.startswith(":") else _ANGLE_PARAM.sub("{}", _BRACE_PARAM.sub("{}", segment)) + for segment in segments + ] + return "/" + "/".join(keyed) + + +def openapi_yaml_excerpts( + text: str, + *, + routes: Iterable[str] = (), + operation_ids: Iterable[str] = (), + schemas: Iterable[str] = (), +) -> list[Excerpt]: + """Slice the matching parts out of an OpenAPI or Swagger document written in YAML. + + There is no YAML parser: each match is an indentation slice, the key line + plus every following line indented deeper than it, with blank and comment + lines inside kept and trailing ones trimmed. The slice is dedented by the + key's indentation. Lines inside a ``|`` or ``>`` block scalar are text, + never keys. + + - A route matches a key directly under the top-level ``paths:`` when the + two :func:`route_key` values are equal, and yields that path item. + - An operation id matches an ``operationId: `` line, quoted or not, + whose enclosing key is an HTTP method (``get:``, ``post:`` and so on), + and yields that operation block. An ``operationId`` under ``links:`` is + a reference, not an operation, and is ignored. + - A schema name matches a key directly under ``components:`` then + ``schemas:`` (or under the top-level ``definitions:`` of Swagger 2.0) + when the :func:`normalize_name` values are equal, and yields that schema. + + ``line`` is the key line's 1-based number. ``symbol`` is the spec's own + path key, operation id or schema key, unquoted. Results come in line order + with no duplicates: an excerpt is dropped when another one already shows + every line of it, so a path item found by route absorbs an operation found + by id, unless the path item's caps cut that operation off. A text whose + first character is ``{`` is JSON, which is also YAML, and goes through + :func:`openapi_json_excerpts`. + """ + text = text.removeprefix(_BOM) + if text.lstrip().startswith("{"): + return openapi_json_excerpts(text, routes=routes, operation_ids=operation_ids, schemas=schemas) + route_keys = _keys(routes, route_key) + op_ids = _keys(operation_ids, str.strip) + schema_keys = _keys(schemas, normalize_name) + if not (route_keys or op_ids or schema_keys): + return [] + doc = _Yaml(text) + top = doc.top_level() + found: list[tuple[int, str]] = [] + if route_keys and "paths" in top: + found += [(i, doc.key[i]) for i in doc.children(top["paths"]) if route_key(doc.key[i]) in route_keys] + if op_ids: + for i, key in enumerate(doc.key): + if key != "operationId" or doc.value[i] not in op_ids: + continue + parent = doc.parent(i) + if parent is not None and doc.key[parent].lower() in _HTTP_METHODS: + found.append((parent, doc.value[i])) + if schema_keys: + containers: list[int] = [] + if "components" in top: + containers += [i for i in doc.children(top["components"]) if doc.key[i] == "schemas"] + if "definitions" in top: + containers.append(top["definitions"]) + for container in containers: + found += [(i, doc.key[i]) for i in doc.children(container) if normalize_name(doc.key[i]) in schema_keys] + return _line_excerpts(doc.spans(found)) + + +def openapi_json_excerpts( + text: str, + *, + routes: Iterable[str] = (), + operation_ids: Iterable[str] = (), + schemas: Iterable[str] = (), +) -> list[Excerpt]: + """Slice the matching parts out of an OpenAPI or Swagger document written in JSON. + + The semantics are :func:`openapi_yaml_excerpts`'s, through + :func:`json.loads`: a route matches a key of ``paths``, an operation id an + HTTP-method object whose ``operationId`` equals it, and a schema name a key + of ``components.schemas`` or of the Swagger 2.0 ``definitions``. + + ``text`` is the matched member rendered as ``"": `` followed by + ``json.dumps(node, indent=2)``, capped the same way. The key is kept + because a rendered context entry shows only ``path:line: text``, and a bare + ``{`` would not say which path, operation or schema it is. ``line`` is best + effort: the line of the matched key found by following its key path + through the source (``paths``, then the path, then the method), else the + line of the deepest ancestor found, else 1. Results come in line order + with no duplicates; here a container absorbs what it contains only when it + is shown whole. Invalid JSON, or a document that is not an object, gives + ``[]``. + """ + text = text.removeprefix(_BOM) + route_keys = _keys(routes, route_key) + op_ids = _keys(operation_ids, str.strip) + schema_keys = _keys(schemas, normalize_name) + if not (route_keys or op_ids or schema_keys): + return [] + doc = _load_json(text) + if not isinstance(doc, dict): + return [] + hits: list[_Hit] = [] + paths = doc.get("paths") + if route_keys and isinstance(paths, dict): + hits += [ + _json_hit(text, ("paths", key), key, key, item) + for key, item in paths.items() + if route_key(key) in route_keys + ] + if op_ids: + for pointer, node in _walk(doc): + key = pointer[-1] if pointer else None + if not (isinstance(key, str) and key.lower() in _HTTP_METHODS and isinstance(node, dict)): + continue + op_id = node.get("operationId") + if isinstance(op_id, str) and op_id in op_ids: + hits.append(_json_hit(text, pointer, op_id, key, node)) + if schema_keys: + containers: list[tuple[tuple[str, ...], dict]] = [] + components = doc.get("components") + if isinstance(components, dict) and isinstance(components.get("schemas"), dict): + containers.append((("components", "schemas"), components["schemas"])) + if isinstance(doc.get("definitions"), dict): + containers.append((("definitions",), doc["definitions"])) + for base, members in containers: + hits += [ + _json_hit(text, (*base, name), name, name, node) + for name, node in members.items() + if normalize_name(name) in schema_keys + ] + return _json_excerpts(hits) + + +def json_schema_excerpts(text: str, *, names: Iterable[str] = ()) -> list[Excerpt]: + """Slice the matching schemas out of a JSON Schema document. + + A name matches the document itself when it equals the root ``title`` by + :func:`normalize_name`, with whitespace in the title dropped first, so + ``transport_requests`` matches ``"title": "Transport Request"``; that + excerpt is the whole document at line 1. A name also matches any member + of a ``$defs`` or ``definitions`` object at any depth (not one that is a + property called ``definitions``), rendered and located as in + :func:`openapi_json_excerpts`. Results come in line order with no + duplicates; the root, when shown whole, absorbs its own definitions. + Invalid JSON, or a document that is not an object, gives ``[]``. + """ + text = text.removeprefix(_BOM) + keys = _keys(names, normalize_name) + if not keys: + return [] + doc = _load_json(text) + if not isinstance(doc, dict): + return [] + hits: list[_Hit] = [] + title = doc.get("title") + if isinstance(title, str) and normalize_name("".join(title.split())) in keys: + hits.append(_Hit((), title, None, doc, 1)) + for pointer, node in _walk(doc): + if not pointer or pointer[-1] not in ("$defs", "definitions") or not isinstance(node, dict): + continue + if len(pointer) > 1 and pointer[-2] in ("properties", "patternProperties"): + continue + hits += [ + _json_hit(text, (*pointer, name), name, name, member) + for name, member in node.items() + if normalize_name(name) in keys + ] + return _json_excerpts(hits) + + +def sql_excerpts(text: str, *, tables: Iterable[str] = ()) -> list[Excerpt]: + """Slice the statements that touch a table out of a SQL script or migration. + + Three statement kinds count, case-insensitively: ``CREATE TABLE`` + (``OR REPLACE``, ``TEMP``, ``UNLOGGED`` and ``IF NOT EXISTS`` allowed), + ``ALTER TABLE`` (``IF EXISTS`` and ``ONLY`` allowed), and + ``CREATE [UNIQUE] INDEX … ON ``. The table matches when its bare + name (the last dot-separated part, with ``"``, backticks and ``[]`` + removed) equals a requested table by :func:`normalize_name`. + + A statement runs from its keyword to its terminating ``;``. A ``;`` inside + a ``--`` or ``/* */`` comment, a quoted string or identifier, or a + ``$$`` body does not end it. A Liquibase formatted-SQL ``--changeset`` line + and a batch-separator line ``GO`` also end a statement, so a changeset's + last statement may omit its ``;``; trailing comments such as + ``--rollback`` stay outside it. Comment lines, including the + ``--liquibase formatted sql`` header, never start one. ``line`` is the + keyword's 1-based line and ``symbol`` the bare table name as written. + Results come in file order. + """ + text = text.removeprefix(_BOM) + keys = _keys(tables, normalize_name) + if not keys: + return [] + out: list[Excerpt] = [] + for start, end, name in _sql_statements(text): + bare = _bare_table(name) + if normalize_name(bare) in keys: + capped, _ = _cap(_slice_lines(text, start, end)) + out.append(Excerpt(text.count("\n", 0, start) + 1, bare, capped)) + return out + + +def liquibase_excerpts(text: str, *, tables: Iterable[str] = ()) -> list[Excerpt]: + """Slice the changeSets that touch a table out of a Liquibase changelog. + + The format follows the first non-blank character: ``<`` is XML, ``{`` or + ``[`` is JSON, anything else is YAML. A changeSet matches when it holds a + ``tableName``, ``baseTableName`` or ``referencedTableName`` (an XML + attribute, a YAML key or a JSON member) equal to a requested table by + :func:`normalize_name`. + + - XML yields the ```` element through ````; + changeSets inside ```` comments are ignored. + - YAML yields the indentation slice of the ``changeSet:`` line (usually + ``- changeSet:``), as :func:`openapi_yaml_excerpts` slices a key. + - JSON yields ``"changeSet": `` plus the member's ``json.dumps(indent=2)``. + + ``line`` is the changeSet's 1-based line (for JSON, the line of its own + ``"changeSet"`` key) and ``symbol`` the first matching table name as + written. Formatted-SQL changelogs go through :func:`sql_excerpts`. + Results come in file order; invalid JSON gives ``[]``. + """ + text = text.removeprefix(_BOM) + keys = _keys(tables, normalize_name) + if not keys: + return [] + head = text.lstrip()[:1] + if head == "<": + return _liquibase_xml(text, keys) + if head in ("{", "["): + return _liquibase_json(text, keys) + doc = _Yaml(text) + found: list[tuple[int, str]] = [] + for i, key in enumerate(doc.key): + if key != "changeSet": + continue + for j in range(i + 1, doc.block_end(i) + 1): + if doc.key[j] in _TABLE_KEYS and normalize_name(doc.value[j]) in keys: + found.append((i, doc.value[j])) + break + return _line_excerpts(doc.spans(found)) + + +def _keys(values: Iterable[str], fold: Callable[[str], str]) -> frozenset[str]: + return frozenset(key for key in map(fold, values) if key) + + +def _cap(lines: list[str]) -> tuple[str, int]: + """``(text, shown)``: the capped text and how many source lines it shows whole. + + A cut block keeps as many leading lines as fit and ends with an + ``… N more lines`` line, which counts toward both caps. A first line too + long to fit even alone is itself cut and ends with ``…``. + """ + text = "\n".join(lines) + if len(lines) <= MAX_CONTRACT_LINES and len(text) <= MAX_CONTRACT_CHARS: + return text, len(lines) + for keep in range(min(len(lines), MAX_CONTRACT_LINES) - 1, 0, -1): + capped = "\n".join([*lines[:keep], f"… {len(lines) - keep} more lines"]) + if len(capped) <= MAX_CONTRACT_CHARS: + return capped, keep + if len(lines) == 1: + return lines[0][: MAX_CONTRACT_CHARS - 1] + "…", 0 + marker = f"… {len(lines) - 1} more lines" + return f"{lines[0][: MAX_CONTRACT_CHARS - len(marker) - 2]}…\n{marker}", 0 + + +def _dedent(line: str, width: int) -> str: + return line[width:] if not line[:width].strip() else line.lstrip() + + +def _slice_lines(text: str, start: int, end: int) -> list[str]: + """The source between two offsets as lines, dedented by the start's column.""" + column = start - (text.rfind("\n", 0, start) + 1) + lines = [line.rstrip() for line in text[start:end].split("\n")] + return [lines[0], *(_dedent(line, column) for line in lines[1:])] + + +@dataclass(frozen=True) +class _Span: + start: int + end: int + shown: int + symbol: str + text: str + + +def _covers(outer: _Span, inner: _Span) -> bool: + """True when ``outer`` already shows every line ``inner`` shows.""" + if (outer.start, outer.end) == (inner.start, inner.end): + return True + return ( + outer.start <= inner.start + and inner.end <= outer.end + and inner.start + max(inner.shown, 1) <= outer.start + outer.shown + ) + + +def _line_excerpts(spans: list[_Span]) -> list[Excerpt]: + kept: list[_Span] = [] + for span in sorted(spans, key=lambda s: (s.start, -s.end)): + if not any(_covers(other, span) for other in kept): + kept.append(span) + return [Excerpt(span.start + 1, span.symbol, span.text) for span in kept] + + +def _unquote(token: str) -> str: + if len(token) >= 2 and token[0] == token[-1] == '"': + try: + value = json.loads(token) + except ValueError: + return token[1:-1] + return value if isinstance(value, str) else token[1:-1] + if len(token) >= 2 and token[0] == token[-1] == "'": + return token[1:-1].replace("''", "'") + return token + + +def _yaml_scalar(rest: str) -> str: + rest = rest.strip() + quoted = _YAML_QUOTED.match(rest) + if quoted: + return _unquote(quoted.group()) + return rest.split(" #", 1)[0].rstrip() + + +class _Yaml: + """Per-line indentation facts about a YAML text, enough for indentation slices.""" + + def __init__(self, text: str) -> None: + self.lines = [line.rstrip() for line in text.split("\n")] + self.indent: list[int] = [] + self.key: list[str | None] = [] + self.value: list[str] = [] + self.dash: list[bool] = [] + self.filler: list[bool] = [] + self.scalar: list[bool] = [] + scalar_floor: int | None = None + for line in self.lines: + stripped = line.lstrip(" ") + indent = len(line) - len(stripped) + in_scalar = scalar_floor is not None and (not stripped or indent > scalar_floor) + if not in_scalar: + scalar_floor = None + filler = not stripped or (not in_scalar and stripped.startswith("#")) + match = None if in_scalar or filler else _YAML_KEY.match(line) + self.indent.append(indent) + self.key.append(_unquote(match.group("key").rstrip()) if match else None) + self.value.append(_yaml_scalar(line[match.end():]) if match else "") + self.dash.append(bool(match and match.group("dash"))) + self.filler.append(filler) + self.scalar.append(in_scalar) + if not in_scalar and not filler and _YAML_BLOCK_SCALAR.search(line): + scalar_floor = indent + + def _content(self, index: int) -> bool: + return not self.filler[index] and not self.scalar[index] + + def top_level(self) -> dict[str, int]: + """Each column-0 mapping key's first line.""" + top: dict[str, int] = {} + for i, key in enumerate(self.key): + if key is not None and self.indent[i] == 0 and not self.dash[i] and self._content(i): + top.setdefault(key, i) + return top + + def block_end(self, start: int) -> int: + """The last content line of the block the key on ``start`` opens.""" + last = start + for j in range(start + 1, len(self.lines)): + if self.filler[j]: + continue + if not self.scalar[j] and self.indent[j] <= self.indent[start]: + break + last = j + return last + + def children(self, parent: int) -> list[int]: + """The mapping keys directly under the key on ``parent``.""" + out: list[int] = [] + child_indent: int | None = None + for j in range(parent + 1, self.block_end(parent) + 1): + if not self._content(j): + continue + if child_indent is None: + child_indent = self.indent[j] + if self.indent[j] == child_indent and self.key[j] is not None and not self.dash[j]: + out.append(j) + return out + + def parent(self, index: int) -> int | None: + """The key line one level above ``index``, or None when that line holds no key.""" + for j in range(index - 1, -1, -1): + if self._content(j) and self.indent[j] < self.indent[index]: + return j if self.key[j] is not None else None + return None + + def spans(self, found: Iterable[tuple[int, str]]) -> list[_Span]: + out: list[_Span] = [] + for start, symbol in found: + end = self.block_end(start) + width = self.indent[start] + text, shown = _cap([_dedent(line, width) for line in self.lines[start : end + 1]]) + out.append(_Span(start, end, shown, symbol, text)) + return out + + +def _load_json(text: str) -> object: + try: + return json.loads(text) + except (ValueError, RecursionError): + return None + + +def _walk(root: object) -> Iterator[tuple[tuple[str | int, ...], object]]: + """Every node under ``root`` with its key path, in document order.""" + stack: list[tuple[tuple[str | int, ...], object]] = [((), root)] + while stack: + pointer, node = stack.pop() + yield pointer, node + if isinstance(node, dict): + stack.extend(((*pointer, key), value) for key, value in reversed(list(node.items()))) + elif isinstance(node, list): + stack.extend(((*pointer, i), value) for i, value in reversed(list(enumerate(node)))) + + +def _json_key_re(key: str) -> re.Pattern[str]: + forms = dict.fromkeys((json.dumps(key, ensure_ascii=False), json.dumps(key))) + return re.compile("(?:" + "|".join(map(re.escape, forms)) + r")\s*:") + + +def _json_line(text: str, pointer: Iterable[str | int]) -> int: + """Best-effort line of a key path: each string key is searched after its parent.""" + position, found = 0, -1 + for key in pointer: + if not isinstance(key, str): + continue + match = _json_key_re(key).search(text, position) + if match is None: + break + found, position = match.start(), match.end() + return text.count("\n", 0, found) + 1 if found >= 0 else 1 + + +@dataclass(frozen=True) +class _Hit: + pointer: tuple[str | int, ...] + symbol: str + label: str | None + node: object + line: int + + +def _json_hit(text: str, pointer: tuple[str | int, ...], symbol: str, label: str | None, node: object) -> _Hit: + return _Hit(pointer, symbol, label, node, _json_line(text, pointer)) + + +def _json_excerpts(hits: list[_Hit]) -> list[Excerpt]: + rendered: list[tuple[_Hit, int, str, bool]] = [] + for order, hit in enumerate(hits): + body = json.dumps(hit.node, indent=2, ensure_ascii=False) + if hit.label is not None: + body = f"{json.dumps(hit.label, ensure_ascii=False)}: {body}" + lines = body.split("\n") + text, shown = _cap(lines) + rendered.append((hit, order, text, shown < len(lines))) + kept: list[tuple[_Hit, int, str, bool]] = [] + for entry in sorted(rendered, key=lambda e: len(e[0].pointer)): + pointer = entry[0].pointer + if any( + other[0].pointer == pointer or (not other[3] and pointer[: len(other[0].pointer)] == other[0].pointer) + for other in kept + ): + continue + kept.append(entry) + kept.sort(key=lambda e: (e[0].line, e[1])) + return [Excerpt(hit.line, hit.symbol, text) for hit, _, text, _ in kept] + + +def _blank(fragment: str) -> str: + return re.sub(r"[^\n]", " ", fragment) + + +def _sql_statements(text: str) -> Iterator[tuple[int, int, str]]: + """``(start, end, table)`` for each table statement: keyword to terminator.""" + pieces: list[str] = [] + bounds: list[tuple[int, int]] = [] + copied = begin = 0 + for match in _SQL_TOKEN.finditer(text): + kind = match.lastgroup + if kind == "comment": + pieces += [text[copied : match.start()], _blank(match.group())] + copied = match.end() + elif kind == "end": + bounds.append((begin, match.end())) + begin = match.end() + elif kind in ("changeset", "go"): + bounds.append((begin, match.start())) + begin = match.end() + pieces.append(text[copied:]) + bounds.append((begin, len(text))) + cleaned = "".join(pieces) + for begin, end in bounds: + first = _NON_SPACE.search(cleaned, begin, end) + if first is None: + continue + for head in _SQL_HEADS: + match = head.match(cleaned, first.start(), end) + if match: + stop = begin + len(cleaned[begin:end].rstrip()) + yield first.start(), stop, match.group("name") + break + + +def _bare_table(name: str) -> str: + parts = _SQL_NAME_PART.findall(name) + return parts[-1].strip('"`[]') if parts else name + + +def _liquibase_xml(text: str, keys: frozenset[str]) -> list[Excerpt]: + blanked = _XML_COMMENT.sub(lambda m: _blank(m.group()), text) + out: list[Excerpt] = [] + for change_set in _XML_CHANGESET.finditer(blanked): + for attr in _XML_TABLE_ATTR.finditer(change_set.group()): + name = attr.group(1) if attr.group(1) is not None else attr.group(2) + if normalize_name(name) in keys: + capped, _ = _cap(_slice_lines(text, change_set.start(), change_set.end())) + out.append(Excerpt(text.count("\n", 0, change_set.start()) + 1, name, capped)) + break + return out + + +def _liquibase_json(text: str, keys: frozenset[str]) -> list[Excerpt]: + doc = _load_json(text) + if doc is None: + return [] + starts = [text.count("\n", 0, m.start()) + 1 for m in _json_key_re("changeSet").finditer(text)] + hits: list[_Hit] = [] + seen = 0 + for pointer, node in _walk(doc): + if not pointer or pointer[-1] != "changeSet": + continue + line = starts[seen] if seen < len(starts) else 1 + seen += 1 + if not isinstance(node, dict): + continue + for inner_pointer, value in _walk(node): + key = inner_pointer[-1] if inner_pointer else None + if key in _TABLE_KEYS and isinstance(value, str) and normalize_name(value) in keys: + hits.append(_Hit(pointer, value, "changeSet", node, line)) + break + return _json_excerpts(hits) diff --git a/tests/test_repo_context_contracts.py b/tests/test_repo_context_contracts.py new file mode 100644 index 0000000..fbd47bf --- /dev/null +++ b/tests/test_repo_context_contracts.py @@ -0,0 +1,639 @@ +"""Unit tests for :mod:`prxref.repo_contracts`, the contract excerpters. + +The module is pure, so every fixture is an inline string: OpenAPI in YAML and +JSON, a JSON Schema, a Liquibase formatted-SQL migration, and Liquibase XML, +YAML and JSON changelogs. The running example is the issue's: a transport +endpoint whose idempotency key is unique on (tenant, connector, key), where the +table is ``idempotency_keys`` and the schema is ``IdempotencyKey``. +""" +from __future__ import annotations + +import json +import re + +import pytest + +from prxref.repo_contracts import ( + MAX_CONTRACT_CHARS, + MAX_CONTRACT_LINES, + Excerpt, + json_schema_excerpts, + liquibase_excerpts, + normalize_name, + openapi_json_excerpts, + openapi_yaml_excerpts, + route_key, + sql_excerpts, +) + +ELLIPSIS = "\N{HORIZONTAL ELLIPSIS}" +MARKER = re.compile(ELLIPSIS + r" (\d+) more lines") + + +def line_of(text: str, prefix: str, after: int = 0) -> int: + """1-based number of the first line after line ``after`` that starts with ``prefix``.""" + return next( + number + for number, line in enumerate(text.split("\n"), 1) + if number > after and line.startswith(prefix) + ) + + +OPENAPI_YAML = """\ +openapi: 3.0.3 +info: + title: Acme connectors + version: 1.0.0 +paths: + /connectors/{connectorId}/transports: + parameters: + - name: connectorId + in: path + required: true + schema: + type: string + + post: + operationId: createTransport + summary: Create a transport for a connector + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/IdempotencyKey' + responses: + '201': + description: Created + get: + operationId: "listTransports" + responses: + '200': + description: OK + + "/tenants/{tenantId}": + get: + operationId: getTenant + description: | + An example kept as text, not as an operation: + post: + operationId: createTransport + responses: + '200': + description: OK + links: + CreateOne: + operationId: createTransport +components: + schemas: + IdempotencyKey: + type: object + description: Unique on (tenant, connector, key). + required: [tenant_id, connector_id, key] + properties: + key: + type: string + + Tenant: + type: object +""" + +TRANSPORTS = "/connectors/{connectorId}/transports" + + +class TestNormalizeName: + def test_plural_snake_case_and_pascal_case_fold_to_one_key(self): + assert normalize_name("idempotency_keys") == "idempotencykey" + assert normalize_name("IdempotencyKey") == "idempotencykey" + assert normalize_name("idempotency-keys") == "idempotencykey" + + def test_exactly_one_trailing_s_is_stripped(self): + assert normalize_name("Keys") == "key" + assert normalize_name("statuss") == "status" + assert normalize_name("s") == "" + + +class TestRouteKey: + @pytest.mark.parametrize("route", [ + "/connectors/{connectorId}/transports", + "/connectors/{id}/transports", + "/connectors/:id/transports", + "/connectors//transports", + "/connectors//transports", + "/connectors/{id:[0-9]+}/transports", + "/connectors/{id}/transports/", + "//connectors//{id}/transports", + "connectors/{id}/transports", + ]) + def test_every_parameter_style_gives_one_key(self, route): + assert route_key(route) == "/connectors/{}/transports" + + def test_case_is_kept(self): + assert route_key("/Connectors/{id}") != route_key("/connectors/{id}") + + def test_a_parameter_inside_a_segment_is_replaced_in_place(self): + assert route_key("/files/{name}.json") == "/files/{}.json" + + def test_root_and_blank(self): + assert route_key("/") == "/" + assert route_key("") == "" + assert route_key(" ") == "" + + +class TestOpenApiYaml: + @pytest.mark.parametrize("route", [ + "/connectors/{id}/transports", + "/connectors/:id/transports", + "/connectors//transports/", + ]) + def test_route_finds_the_path_item(self, route): + [excerpt] = openapi_yaml_excerpts(OPENAPI_YAML, routes=[route]) + assert excerpt.symbol == TRANSPORTS + assert excerpt.line == line_of(OPENAPI_YAML, f" {TRANSPORTS}:") + lines = excerpt.text.split("\n") + assert lines[0] == f"{TRANSPORTS}:" + assert lines[1] == " parameters:" + assert lines[-1] == " description: OK" + assert "\n\n post:" in excerpt.text + assert "operationId: createTransport" in excerpt.text + assert "listTransports" in excerpt.text + assert "tenants" not in excerpt.text + + def test_operation_id_finds_the_enclosing_operation_only(self): + [excerpt] = openapi_yaml_excerpts(OPENAPI_YAML, operation_ids=["createTransport"]) + assert excerpt == Excerpt( + line=line_of(OPENAPI_YAML, " post:"), + symbol="createTransport", + text=excerpt.text, + ) + lines = excerpt.text.split("\n") + assert lines[:2] == ["post:", " operationId: createTransport"] + assert lines[-1] == " description: Created" + assert "listTransports" not in excerpt.text + + def test_quoted_operation_id_and_quoted_path_key(self): + [operation] = openapi_yaml_excerpts(OPENAPI_YAML, operation_ids=["listTransports"]) + assert operation.symbol == "listTransports" + assert operation.text.startswith("get:\n") + [path_item] = openapi_yaml_excerpts(OPENAPI_YAML, routes=["/tenants/:tenantId"]) + assert path_item.symbol == "/tenants/{tenantId}" + assert path_item.line == line_of(OPENAPI_YAML, ' "/tenants/{tenantId}":') + assert path_item.text.split("\n")[-1] == " operationId: createTransport" + + @pytest.mark.parametrize("name", ["idempotency_keys", "IdempotencyKey", "idempotency-key"]) + def test_schema_matches_by_normalized_name(self, name): + [excerpt] = openapi_yaml_excerpts(OPENAPI_YAML, schemas=[name]) + assert excerpt.symbol == "IdempotencyKey" + assert excerpt.line == line_of(OPENAPI_YAML, " IdempotencyKey:") + assert excerpt.text.split("\n") == [ + "IdempotencyKey:", + " type: object", + " description: Unique on (tenant, connector, key).", + " required: [tenant_id, connector_id, key]", + " properties:", + " key:", + " type: string", + ] + + def test_a_route_and_its_operation_id_give_one_excerpt(self): + excerpts = openapi_yaml_excerpts( + OPENAPI_YAML, routes=["/connectors/:id/transports"], operation_ids=["createTransport"], + ) + assert [e.symbol for e in excerpts] == [TRANSPORTS] + + def test_results_come_in_line_order(self): + excerpts = openapi_yaml_excerpts( + OPENAPI_YAML, + schemas=["idempotency_keys"], + operation_ids=["getTenant"], + routes=["/connectors/:id/transports"], + ) + assert [e.symbol for e in excerpts] == [TRANSPORTS, "getTenant", "IdempotencyKey"] + assert [e.line for e in excerpts] == sorted(e.line for e in excerpts) + + def test_nothing_asked_or_nothing_matching_gives_nothing(self): + assert openapi_yaml_excerpts(OPENAPI_YAML) == [] + assert openapi_yaml_excerpts(OPENAPI_YAML, routes=["/connectors"], schemas=["widgets"]) == [] + assert openapi_yaml_excerpts("", routes=["/connectors"]) == [] + + def test_swagger_two_definitions_hold_schemas(self): + spec = "swagger: '2.0'\ndefinitions:\n IdempotencyKey:\n type: object\n Tenant:\n type: object\n" + [excerpt] = openapi_yaml_excerpts(spec, schemas=["idempotency_keys"]) + assert excerpt == Excerpt(3, "IdempotencyKey", "IdempotencyKey:\n type: object") + + def test_a_path_key_may_hold_a_colon_that_is_not_followed_by_a_space(self): + spec = "paths:\n /connectors/{id}:cancel:\n post:\n operationId: cancelConnector\n" + [excerpt] = openapi_yaml_excerpts(spec, routes=["/connectors/{connectorId}:cancel"]) + assert (excerpt.line, excerpt.symbol) == (2, "/connectors/{id}:cancel") + assert excerpt.text.split("\n")[-1] == " operationId: cancelConnector" + + def test_crlf_line_endings_and_a_byte_order_mark(self): + text = chr(0xFEFF) + OPENAPI_YAML.replace("\n", "\r\n") + assert openapi_yaml_excerpts(text, routes=["/connectors/:id/transports"], schemas=["IdempotencyKey"]) == ( + openapi_yaml_excerpts(OPENAPI_YAML, routes=["/connectors/:id/transports"], schemas=["IdempotencyKey"]) + ) + + def test_a_json_text_goes_through_the_json_excerpter(self): + text = json.dumps(OPENAPI_JSON, indent=2) + assert openapi_yaml_excerpts(text, routes=["/connectors/:id/transports"]) == openapi_json_excerpts( + text, routes=["/connectors/:id/transports"], + ) + + +def _schema_doc(body: list[str]) -> str: + return "\n".join([ + "components:", + " schemas:", + " Wide:", + *(f" {line}" for line in body), + " Other:", + " type: object", + ]) + + +def _operations(count: int, filler: int) -> str: + lines = ["paths:", f" {TRANSPORTS}:"] + for method, op_id in [("get", "listTransports"), ("put", "replaceTransport"), ("post", "createTransport")][:count]: + lines += [f" {method}:", f" operationId: {op_id}"] + lines += [f" x-note-{i:02d}: filler" for i in range(filler)] + return "\n".join(lines) + + +class TestCaps: + def test_a_sixty_line_schema_is_cut_to_forty_lines(self): + body = ["type: object", "properties:", *(f" field_{i:02d}: {{type: string}}" for i in range(57))] + [excerpt] = openapi_yaml_excerpts(_schema_doc(body), schemas=["wide"]) + source = ["Wide:", *(f" {line}" for line in body)] + assert len(source) == 60 + lines = excerpt.text.split("\n") + assert len(lines) == MAX_CONTRACT_LINES == 40 + assert lines[:39] == source[:39] + assert lines[-1] == f"{ELLIPSIS} 21 more lines" + assert len(excerpt.text) <= MAX_CONTRACT_CHARS + + def test_forty_lines_stay_whole_and_forty_one_are_cut(self): + whole = openapi_yaml_excerpts(_schema_doc([f"f{i}: x" for i in range(39)]), schemas=["wide"])[0] + assert len(whole.text.split("\n")) == 40 + assert not MARKER.search(whole.text) + cut = openapi_yaml_excerpts(_schema_doc([f"f{i}: x" for i in range(40)]), schemas=["wide"])[0] + lines = cut.text.split("\n") + assert len(lines) == 40 + assert lines[-1] == f"{ELLIPSIS} 2 more lines" + + def test_a_three_thousand_char_block_is_cut_under_the_char_cap(self): + body = [f"f{i}: {'x' * 300}" for i in range(10)] + source = ["Wide:", *(f" {line}" for line in body)] + assert len("\n".join(source)) > 3000 + [excerpt] = openapi_yaml_excerpts(_schema_doc(body), schemas=["wide"]) + assert len(excerpt.text) <= MAX_CONTRACT_CHARS + lines = excerpt.text.split("\n") + kept = len(lines) - 1 + assert lines[:kept] == source[:kept] + assert lines[-1] == f"{ELLIPSIS} {len(source) - kept} more lines" + one_more = "\n".join([*source[: kept + 1], f"{ELLIPSIS} {len(source) - kept - 1} more lines"]) + assert len(one_more) > MAX_CONTRACT_CHARS + + def test_a_single_overlong_line_is_itself_cut(self): + columns = ", ".join(f"column_{i:04d} BIGINT" for i in range(150)) + text = f"CREATE TABLE idempotency_keys ({columns});" + assert len(text) > 2500 + [excerpt] = sql_excerpts(text, tables=["IdempotencyKey"]) + assert len(excerpt.text) <= MAX_CONTRACT_CHARS + assert excerpt.text.startswith("CREATE TABLE idempotency_keys (") + assert excerpt.text.endswith(ELLIPSIS) + + def test_a_cut_path_item_does_not_absorb_an_operation_it_cuts_off(self): + spec = _operations(3, 18) + route = ["/connectors/:id/transports"] + shown = openapi_yaml_excerpts(spec, routes=route, operation_ids=["listTransports"]) + assert [e.symbol for e in shown] == [TRANSPORTS] + hidden = openapi_yaml_excerpts(spec, routes=route, operation_ids=["createTransport"]) + assert [e.symbol for e in hidden] == [TRANSPORTS, "createTransport"] + assert "createTransport" not in hidden[0].text + assert hidden[1].line == line_of(spec, " post:") + + +OPENAPI_JSON = { + "openapi": "3.0.3", + "info": {"title": "Acme connectors", "version": "1.0.0"}, + "paths": { + "/tenants/{tenantId}": { + "post": {"operationId": "updateTenant", "responses": {"200": {"description": "OK"}}}, + }, + TRANSPORTS: { + "post": { + "operationId": "createTransport", + "requestBody": { + "content": {"application/json": {"schema": {"$ref": "#/components/schemas/IdempotencyKey"}}}, + }, + "responses": {"201": {"description": "Created"}}, + }, + }, + }, + "components": { + "schemas": { + "Tenant": {"type": "object"}, + "IdempotencyKey": { + "type": "object", + "description": "Unique on (tenant, connector, key).", + "properties": {"key": {"type": "string"}}, + }, + }, + }, +} + + +class TestOpenApiJson: + TEXT = json.dumps(OPENAPI_JSON, indent=2) + + @pytest.mark.parametrize("route", ["/connectors/{id}/transports", "/connectors/:id/transports"]) + def test_route_finds_the_path_item(self, route): + [excerpt] = openapi_json_excerpts(self.TEXT, routes=[route]) + assert excerpt.symbol == TRANSPORTS + assert excerpt.line == line_of(self.TEXT, f' "{TRANSPORTS}": {{') + assert excerpt.text == f'"{TRANSPORTS}": ' + json.dumps(OPENAPI_JSON["paths"][TRANSPORTS], indent=2) + + def test_operation_id_finds_its_operation_and_its_own_line(self): + [excerpt] = openapi_json_excerpts(self.TEXT, operation_ids=["createTransport"]) + assert excerpt.symbol == "createTransport" + path_line = line_of(self.TEXT, f' "{TRANSPORTS}": {{') + assert excerpt.line == line_of(self.TEXT, ' "post": {', after=path_line) + assert excerpt.text.startswith('"post": {\n "operationId": "createTransport",') + + def test_schema_matches_by_normalized_name(self): + [excerpt] = openapi_json_excerpts(self.TEXT, schemas=["idempotency_keys"]) + assert excerpt.symbol == "IdempotencyKey" + assert excerpt.line == line_of(self.TEXT, ' "IdempotencyKey": {') + assert "Unique on (tenant, connector, key)." in excerpt.text + + def test_a_route_and_its_operation_id_give_one_excerpt(self): + excerpts = openapi_json_excerpts( + self.TEXT, routes=["/connectors/:id/transports"], operation_ids=["createTransport"], + ) + assert [e.symbol for e in excerpts] == [TRANSPORTS] + + def test_swagger_two_definitions_hold_schemas(self): + text = json.dumps({"swagger": "2.0", "definitions": {"IdempotencyKey": {"type": "object"}}}, indent=2) + [excerpt] = openapi_json_excerpts(text, schemas=["idempotency_keys"]) + assert excerpt.line == line_of(text, ' "IdempotencyKey": {') == 4 + + @pytest.mark.parametrize("text", ["{not json", "[]", "", '"paths"']) + def test_invalid_or_non_object_json_gives_nothing(self, text): + assert openapi_json_excerpts(text, routes=["/connectors/:id/transports"], schemas=["IdempotencyKey"]) == [] + + +JSON_SCHEMA = { + "$schema": "https://json-schema.org/draft/2020-12/schema", + "title": "Transport Request", + "type": "object", + "properties": {"idempotencyKey": {"$ref": "#/$defs/IdempotencyKey"}}, + "$defs": { + "Tenant": {"type": "object"}, + "IdempotencyKey": { + "type": "object", + "description": "Unique on (tenant, connector, key).", + "required": ["tenant_id", "connector_id", "key"], + }, + }, +} + + +class TestJsonSchema: + TEXT = json.dumps(JSON_SCHEMA, indent=2) + + def test_a_defs_member_matches_by_normalized_name(self): + [excerpt] = json_schema_excerpts(self.TEXT, names=["idempotency_keys"]) + assert excerpt.symbol == "IdempotencyKey" + assert excerpt.line == line_of(self.TEXT, ' "IdempotencyKey": {') + assert excerpt.text == '"IdempotencyKey": ' + json.dumps(JSON_SCHEMA["$defs"]["IdempotencyKey"], indent=2) + + def test_the_root_title_matches_and_absorbs_its_definitions(self): + [excerpt] = json_schema_excerpts(self.TEXT, names=["transport_requests", "idempotency_keys"]) + assert excerpt == Excerpt(1, "Transport Request", json.dumps(JSON_SCHEMA, indent=2)) + + def test_draft_seven_definitions_and_not_a_property_called_definitions(self): + schema = { + "definitions": {"IdempotencyKey": {"type": "object"}}, + "properties": {"definitions": {"type": "object", "properties": {"tenant": {"type": "string"}}}}, + } + text = json.dumps(schema, indent=2) + assert [e.symbol for e in json_schema_excerpts(text, names=["IdempotencyKeys"])] == ["IdempotencyKey"] + assert json_schema_excerpts(text, names=["type"]) == [] + + @pytest.mark.parametrize("text", ["{", "[1, 2]", ""]) + def test_invalid_or_non_object_json_gives_nothing(self, text): + assert json_schema_excerpts(text, names=["IdempotencyKey"]) == [] + + +MIGRATION = """\ +--liquibase formatted sql + +--changeset acme:1 +CREATE TABLE idempotency_keys ( + id BIGSERIAL PRIMARY KEY, + tenant_id BIGINT NOT NULL, -- the owning tenant; never null + connector_id BIGINT NOT NULL, + key VARCHAR(255) NOT NULL +); +--rollback DROP TABLE idempotency_keys; + +--changeset acme:2 +ALTER TABLE public."idempotency_keys" ADD CONSTRAINT uq_idempotency_keys_scope UNIQUE (tenant_id, connector_id, key); +--rollback ALTER TABLE idempotency_keys DROP CONSTRAINT uq_idempotency_keys_scope; + +--changeset acme:3 +create unique index ux on idempotency_keys (tenant_id, key); + +--changeset acme:4 +CREATE TABLE connectors ( + id BIGSERIAL PRIMARY KEY, + note TEXT DEFAULT 'a;b' +); +""" + + +class TestSql: + def test_create_alter_and_index_statements_match_the_schema_name(self): + excerpts = sql_excerpts(MIGRATION, tables=["IdempotencyKey"]) + assert [e.line for e in excerpts] == [ + line_of(MIGRATION, "CREATE TABLE idempotency_keys"), + line_of(MIGRATION, "ALTER TABLE"), + line_of(MIGRATION, "create unique index"), + ] + assert [e.symbol for e in excerpts] == ["idempotency_keys"] * 3 + create, alter, index = (e.text for e in excerpts) + assert create.split("\n") == [ + "CREATE TABLE idempotency_keys (", + " id BIGSERIAL PRIMARY KEY,", + " tenant_id BIGINT NOT NULL, -- the owning tenant; never null", + " connector_id BIGINT NOT NULL,", + " key VARCHAR(255) NOT NULL", + ");", + ] + assert alter == ( + 'ALTER TABLE public."idempotency_keys" ADD CONSTRAINT uq_idempotency_keys_scope ' + "UNIQUE (tenant_id, connector_id, key);" + ) + assert index == "create unique index ux on idempotency_keys (tenant_id, key);" + + def test_a_semicolon_in_a_string_does_not_end_the_statement(self): + [excerpt] = sql_excerpts(MIGRATION, tables=["connectors"]) + assert excerpt.text.split("\n")[-2:] == [" note TEXT DEFAULT 'a;b'", ");"] + + def test_an_unrelated_table_is_not_found(self): + assert sql_excerpts(MIGRATION, tables=["tenants"]) == [] + assert sql_excerpts(MIGRATION) == [] + + def test_a_changeset_line_ends_a_statement_that_has_no_semicolon(self): + text = ( + "--liquibase formatted sql\n" + "--changeset acme:5\n" + "CREATE INDEX ix_keys_tenant ON idempotency_keys (tenant_id)\n" + "--rollback DROP INDEX ix_keys_tenant;\n" + "--changeset acme:6\n" + "CREATE TABLE tenants (id BIGINT)\n" + ) + assert sql_excerpts(text, tables=["idempotency_keys"]) == [ + Excerpt(3, "idempotency_keys", "CREATE INDEX ix_keys_tenant ON idempotency_keys (tenant_id)"), + ] + assert sql_excerpts(text, tables=["tenant"]) == [Excerpt(6, "tenants", "CREATE TABLE tenants (id BIGINT)")] + + @pytest.mark.parametrize("newline", ["\n", "\r\n"]) + def test_quoted_and_schema_qualified_names_use_the_bare_table(self, newline): + text = newline.join([ + "ALTER TABLE `acme`.`idempotency_keys` ADD UNIQUE KEY uq (tenant_id, `key`);", + "GO", + "CREATE NONCLUSTERED INDEX ix ON [dbo].[IdempotencyKeys] ([key])", + "GO", + "CREATE TABLE IF NOT EXISTS acme.tenants (id BIGINT);", + ]) + excerpts = sql_excerpts(text, tables=["idempotency_key"]) + assert [(e.line, e.symbol) for e in excerpts] == [(1, "idempotency_keys"), (3, "IdempotencyKeys")] + assert excerpts[1].text == "CREATE NONCLUSTERED INDEX ix ON [dbo].[IdempotencyKeys] ([key])" + + def test_block_comments_and_dollar_bodies_hide_their_semicolons(self): + text = ( + "/* ALTER TABLE idempotency_keys DROP COLUMN key; */\n" + "CREATE FUNCTION touch() RETURNS trigger AS $body$\n" + "BEGIN\n" + " ALTER TABLE idempotency_keys ADD COLUMN touched BIGINT;\n" + "END;\n" + "$body$ LANGUAGE plpgsql;\n" + "CREATE INDEX ON ONLY idempotency_keys (key);\n" + ) + assert sql_excerpts(text, tables=["IdempotencyKey"]) == [ + Excerpt(7, "idempotency_keys", "CREATE INDEX ON ONLY idempotency_keys (key);"), + ] + + def test_crlf_line_endings_and_a_byte_order_mark(self): + text = chr(0xFEFF) + MIGRATION.replace("\n", "\r\n") + excerpts = sql_excerpts(text, tables=["IdempotencyKey"]) + assert [e.line for e in excerpts] == [e.line for e in sql_excerpts(MIGRATION, tables=["IdempotencyKey"])] + assert excerpts[0].text.split("\n")[-1] == ");" + assert "\r" not in "".join(e.text for e in excerpts) + + +LIQUIBASE_XML = """\ + + + + + + + + + + + + + + + + + + + +""" + +LIQUIBASE_YAML = """\ +databaseChangeLog: + - changeSet: + id: 1 + author: acme + changes: + - createTable: + tableName: idempotency_keys + columns: + - column: + name: id + type: BIGINT + - changeSet: + id: 2 + author: acme + changes: + - addUniqueConstraint: + tableName: "idempotency_keys" + columnNames: tenant_id, connector_id, key + - changeSet: + id: 3 + author: acme + changes: + - createTable: + tableName: connectors +""" + + +class TestLiquibase: + def test_xml_changesets_holding_the_table(self): + excerpts = liquibase_excerpts(LIQUIBASE_XML, tables=["IdempotencyKey"]) + assert [(e.line, e.symbol) for e in excerpts] == [ + (line_of(LIQUIBASE_XML, ' ', + ' ', + ' ', + " ", + "", + ] + assert all('id="0"' not in e.text for e in excerpts) + + def test_xml_changeset_for_another_table(self): + [excerpt] = liquibase_excerpts(LIQUIBASE_XML, tables=["connectors"]) + assert excerpt.text.split("\n") == [ + '', + ' ', + "", + ] + + def test_yaml_changesets_holding_the_table(self): + excerpts = liquibase_excerpts(LIQUIBASE_YAML, tables=["IdempotencyKey"]) + assert [(e.line, e.symbol) for e in excerpts] == [(2, "idempotency_keys"), (12, "idempotency_keys")] + first = excerpts[0].text.split("\n") + assert first[:3] == ["- changeSet:", " id: 1", " author: acme"] + assert first[-1] == " type: BIGINT" + assert excerpts[1].text.split("\n")[-1] == " columnNames: tenant_id, connector_id, key" + + def test_json_changesets_holding_the_table(self): + changelog = {"databaseChangeLog": [ + {"changeSet": {"id": "1", "author": "acme", "changes": [{"createTable": {"tableName": "connectors"}}]}}, + {"changeSet": {"id": "2", "author": "acme", "changes": [ + {"createIndex": {"tableName": "idempotency_keys", "indexName": "ux", "unique": True}}, + ]}}, + ]} + text = json.dumps(changelog, indent=2) + [excerpt] = liquibase_excerpts(text, tables=["IdempotencyKey"]) + second = line_of(text, ' "changeSet": {', after=line_of(text, ' "changeSet": {')) + assert (excerpt.line, excerpt.symbol) == (second, "idempotency_keys") + assert excerpt.text.startswith('"changeSet": {\n "id": "2",') + + def test_an_unrelated_table_or_invalid_json_gives_nothing(self): + assert liquibase_excerpts(LIQUIBASE_XML, tables=["tenants"]) == [] + assert liquibase_excerpts(LIQUIBASE_YAML, tables=["tenants"]) == [] + assert liquibase_excerpts("{broken", tables=["IdempotencyKey"]) == [] From d251e69a7d8859f0048b6fd492f2e9ff73049cfe Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:13:34 -0700 Subject: [PATCH 09/24] feat: run-scoped repository reader with single-flight reads and caps (#17 T12) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add src/prxref/repo_reader.py, the one reader a review run shares across its parallel chunk workers when repository context is on. RepoReader wraps a fetch(path) and an optional lister(): - read(path) is cached and single-flight per path: concurrent callers share one fetch, and a None result is cached, so a miss costs one fetch per run. It is not capped; its fetches count in stats()["reads"] only. - chunk_reader() returns a new callable per chunk that enforces MAX_CHUNK_READS (16) per chunk and MAX_RUN_READS (200) across all chunk readers. A cached or in-flight path is served without counting; a path refused at a cap returns None without a fetch, stays uncached, and sets read_cap_hit. - listing() calls the lister at most once per run, single-flight, and caches its result, with a raise or a non-PathListing read as None. - exclude(path) makes a path unreadable, uncounted and unlisted at the one boundary every read crosses; an exclude that raises fails closed. - stats() snapshots {"reads", "read_cap_hit", "listing"} for the run record. forge_reader(forge, ref, sha) mirrors _make_file_reader's gate (None without get_file_content or with an empty sha) and passes ref and sha to both get_file_content and list_paths; repo_dir_reader(repo_dir) wraps a RepoDir. orchestrator.py is untouched: the wiring seat builds the reader and maps stats() onto the run record. Tests: tests/test_orchestrator_repo_reader.py (49 tests). 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/repo_reader.py | 287 ++++++++++++ tests/test_orchestrator_repo_reader.py | 592 +++++++++++++++++++++++++ 2 files changed, 879 insertions(+) create mode 100644 src/prxref/repo_reader.py create mode 100644 tests/test_orchestrator_repo_reader.py diff --git a/src/prxref/repo_reader.py b/src/prxref/repo_reader.py new file mode 100644 index 0000000..78e6f45 --- /dev/null +++ b/src/prxref/repo_reader.py @@ -0,0 +1,287 @@ +"""The run-scoped repository reader behind repository context (#17). + +One ``RepoReader`` serves every file read and the one path listing a review +run makes when repository context is on. It sits in front of a forge's +optional ``get_file_content`` and ``list_paths`` (``forge_reader``) or a +local ``--repo-dir`` tree (``repo_dir_reader``) and gives the chunk workers, +which run in parallel, one shared view of the repository: + +- every path is fetched at most once per run, single-flight across threads, + and a miss is cached like a hit; +- the listing is fetched at most once per run, single-flight as well; +- each chunk's reader is capped at ``MAX_CHUNK_READS`` uncached reads and all + chunk readers together at ``MAX_RUN_READS``, so repository context bounds + both its network cost and its prompt growth; +- an excluded path is never read and never listed, at the one boundary every + read crosses; +- ``stats`` reports the counters the run record carries. + +The module does no prompt work and knows nothing about which paths are +worth reading; the caller decides that and passes plain callables in. +""" + +from __future__ import annotations + +import logging +import threading +from collections.abc import Callable + +from prxref.forges.base import Forge, PathListing, PRRef +from prxref.forges.repo_dir import RepoDir + +logger = logging.getLogger(__name__) + +MAX_RUN_READS = 200 +MAX_CHUNK_READS = 16 +READER_KINDS = ("forge", "repo-dir") + + +class _Slot: + """One single-flight result: ``done`` is set once ``value`` is final.""" + + __slots__ = ("done", "value") + + def __init__(self) -> None: + self.done = threading.Event() + self.value: str | PathListing | None = None + + +class RepoReader: + """A cached, single-flight, capped view of one repository for one review run. + + ``fetch(path)`` returns a file's text or None; ``lister()`` returns a + ``PathListing`` or None, and ``lister=None`` means the source cannot list. + Either callable may raise or return the wrong type: an exception is + logged at DEBUG and read as None, and a ``fetch`` result that is not a + ``str`` (or a ``lister`` result that is not a ``PathListing``) is read as + None. ``kind`` is ``"forge"`` or ``"repo-dir"``, the run record's + ``reader`` field. ``exclude(path)`` returning true (or raising) makes a + path unreadable and unlisted. + + The object is thread-safe and meant to be shared by every chunk worker + of a run. ``read`` is cached and single-flight but not capped: its + fetches count in ``stats()["reads"]`` and never spend either cap. + ``chunk_reader`` hands each chunk a reader that also enforces + ``chunk_cap`` for that chunk and ``run_cap`` across all chunk readers. + A path already cached or in flight is served without counting against + either cap, whichever reader fetched it. + + Determinism: below the caps, what every read returns is independent of + the order in which parallel workers ask, because each path is fetched + once and every caller sees that one result. Once the run cap is reached, + WHICH chunk meets it first depends on thread scheduling; that is + accepted only because ``stats()["read_cap_hit"]`` records that it + happened. + """ + + def __init__( + self, + fetch: Callable[[str], str | None], + lister: Callable[[], PathListing | None] | None, + *, + kind: str, + exclude: Callable[[str], bool] | None = None, + run_cap: int = MAX_RUN_READS, + chunk_cap: int = MAX_CHUNK_READS, + ) -> None: + """Wrap ``fetch`` and ``lister``; raises ``ValueError`` for a ``kind`` outside ``READER_KINDS``.""" + if kind not in READER_KINDS: + raise ValueError(f"RepoReader kind must be one of {READER_KINDS}, got {kind!r}") + self.kind = kind + self._fetch = fetch + self._lister = lister + self._exclude = exclude + self._run_cap = run_cap + self._chunk_cap = chunk_cap + self._lock = threading.Lock() + self._slots: dict[str, _Slot] = {} + self._reads = 0 + self._run_used = 0 + self._read_cap_hit = False + self._listing_slot: _Slot | None = None + + def _excluded(self, path: str) -> bool: + if self._exclude is None: + return False + try: + return bool(self._exclude(path)) + except Exception as e: # noqa: BLE001 - an exclude that cannot answer fails closed + logger.debug("repository context exclude(%s) failed: %s", path, e) + return True + + def _fetch_value(self, path: str) -> str | None: + try: + value = self._fetch(path) + except Exception as e: # noqa: BLE001 - context is never worth a failed review + logger.debug("%s read of %s failed: %s", self.kind, path, e) + return None + return value if isinstance(value, str) else None + + def _get(self, path: str, admit: Callable[[], bool] | None) -> str | None: + if self._excluded(path): + return None + with self._lock: + slot = self._slots.get(path) + owner = slot is None + if owner: + if admit is not None and not admit(): + self._read_cap_hit = True + return None + slot = _Slot() + self._slots[path] = slot + self._reads += 1 + if owner: + try: + slot.value = self._fetch_value(path) + finally: + slot.done.set() + else: + slot.done.wait() + value = slot.value + return value if isinstance(value, str) else None + + def read(self, path: str) -> str | None: + """Return the text of ``path``, or None; cached and single-flight, never capped. + + The first caller for an uncached path fetches it; a concurrent caller + for the same path waits for that one fetch and gets its result. A + None result is cached too, so a miss costs one fetch per run. An + excluded path returns None without fetching or counting. + """ + return self._get(path, None) + + def chunk_reader(self) -> Callable[[str], str | None]: + """Return a new ``read(path) -> str | None`` for one chunk, enforcing both caps. + + A path already cached or in flight is served as ``read`` serves it, + without counting. An uncached path counts one against this chunk's + ``chunk_cap`` and one against the run's ``run_cap``, which every + chunk reader of this ``RepoReader`` shares. When either cap is + already reached, it returns None without fetching, leaves the path + uncached, and marks ``read_cap_hit``. An excluded path returns None + without fetching or counting. + """ + used = 0 + + def admit() -> bool: + nonlocal used + if used >= self._chunk_cap or self._run_used >= self._run_cap: + return False + used += 1 + self._run_used += 1 + return True + + def read(path: str) -> str | None: + return self._get(path, admit) + + return read + + def _list_value(self) -> PathListing | None: + try: + listing = self._lister() + except Exception as e: # noqa: BLE001 - context is never worth a failed review + logger.debug("%s path listing failed: %s", self.kind, e) + return None + if not isinstance(listing, PathListing): + return None + if self._exclude is None: + return listing + kept = tuple(path for path in listing.paths if not self._excluded(path)) + return PathListing(paths=kept, complete=listing.complete) + + def listing(self) -> PathListing | None: + """Return the repository's path listing, or None; the lister runs at most once per run. + + Concurrent callers share one call to ``lister``, and its result is + cached, including None, a raise read as None, and a result that is + not a ``PathListing`` read as None. Excluded paths are filtered out, + and ``complete`` stays as the source reported it. With + ``lister=None`` this returns None and calls nothing. + """ + if self._lister is None: + return None + with self._lock: + slot = self._listing_slot + owner = slot is None + if owner: + slot = self._listing_slot = _Slot() + if owner: + try: + slot.value = self._list_value() + finally: + slot.done.set() + else: + slot.done.wait() + value = slot.value + return value if isinstance(value, PathListing) else None + + def stats(self) -> dict: + """Return a snapshot of the run's counters for the run record. + + ``{"reads": int, "read_cap_hit": bool, "listing": {"paths": int, + "complete": bool} | None}``. ``reads`` counts every call that reached + ``fetch``, from ``read`` and chunk readers alike. ``read_cap_hit`` is + true once any chunk reader refused a path at a cap. ``listing`` is + None until ``listing()`` has finished, and when it gave None; its + ``paths`` counts the listing after exclusion. Taking a snapshot never + calls ``fetch`` or ``lister``. + """ + with self._lock: + reads = self._reads + cap_hit = self._read_cap_hit + slot = self._listing_slot + listing = None + if slot is not None and slot.done.is_set() and isinstance(slot.value, PathListing): + listing = {"paths": len(slot.value.paths), "complete": slot.value.complete} + return {"reads": reads, "read_cap_hit": cap_hit, "listing": listing} + + +def forge_reader( + forge: Forge, + ref: PRRef, + sha: str | None, + *, + exclude: Callable[[str], bool] | None = None, +) -> RepoReader | None: + """Return a ``RepoReader`` over a forge at commit ``sha``, or None. + + None when the forge has no ``get_file_content`` or ``sha`` is empty, + exactly as the orchestrator's same-file reader decides. Reads call + ``get_file_content(ref, path, sha=sha)``. The listing calls + ``list_paths(ref, sha=sha)`` when the forge has that method; otherwise + the reader has no lister and ``listing()`` is None. ``kind`` is + ``"forge"``. + """ + getter = getattr(forge, "get_file_content", None) + if getter is None or not sha: + return None + list_paths = getattr(forge, "list_paths", None) + + def fetch(path: str) -> str | None: + return getter(ref, path, sha=sha) + + def list_at_sha() -> PathListing | None: + return list_paths(ref, sha=sha) + + lister = None if list_paths is None else list_at_sha + return RepoReader(fetch, lister, kind="forge", exclude=exclude) + + +def repo_dir_reader( + repo_dir: RepoDir, + *, + exclude: Callable[[str], bool] | None = None, +) -> RepoReader: + """Return a ``RepoReader`` over a local ``RepoDir`` tree. + + Reads call ``repo_dir.read``; the listing wraps ``repo_dir.list_files()`` + into a ``PathListing``. Both caps apply as they do to a forge, because + they bound prompt growth as well as network cost. ``kind`` is + ``"repo-dir"``. + """ + + def lister() -> PathListing: + paths, complete = repo_dir.list_files() + return PathListing(paths=tuple(paths), complete=complete) + + return RepoReader(repo_dir.read, lister, kind="repo-dir", exclude=exclude) diff --git a/tests/test_orchestrator_repo_reader.py b/tests/test_orchestrator_repo_reader.py new file mode 100644 index 0000000..e306061 --- /dev/null +++ b/tests/test_orchestrator_repo_reader.py @@ -0,0 +1,592 @@ +"""The run-scoped repository reader (#17 T12): ``prxref.repo_reader``. + +Plain fakes only: no forge adapter, no network. The single-flight tests hold +the first fetch (or listing) open on a ``threading.Event`` while seven more +threads ask for the same thing, and assert the call count BEFORE releasing +it, so a reader that lets a second caller fetch while the first is in +flight fails here rather than passing on timing luck. +""" +from __future__ import annotations + +import logging +import threading +import time + +import pytest + +from prxref.forges.base import PathListing +from prxref.forges.replay import ReplayForge +from prxref.forges.repo_dir import RepoDir +from prxref.repo_reader import ( + MAX_CHUNK_READS, + MAX_RUN_READS, + READER_KINDS, + RepoReader, + forge_reader, + repo_dir_reader, +) + +WAIT = 5.0 +SHA = "c" * 40 +REF = object() +THREADS = 8 + + +class Fetch: + """A recording ``fetch`` over a dict; paths it does not hold read as None.""" + + def __init__(self, files: dict[str, object] | None = None): + self.files = files if files is not None else {} + self.calls: list[str] = [] + self.lock = threading.Lock() + + def __call__(self, path: str): + with self.lock: + self.calls.append(path) + return self.files.get(path) + + +class BlockingFetch(Fetch): + """A ``Fetch`` whose FIRST call signals ``entered`` and then blocks on ``release``.""" + + def __init__(self, files: dict[str, object]): + super().__init__(files) + self.entered = threading.Event() + self.release = threading.Event() + + def __call__(self, path: str): + with self.lock: + self.calls.append(path) + first = len(self.calls) == 1 + if first: + self.entered.set() + self.release.wait(WAIT) + return self.files.get(path) + + +class Lister: + """A recording ``lister`` returning a fixed result, or raising it when it is an exception.""" + + def __init__(self, result: object): + self.result = result + self.calls = 0 + + def __call__(self): + self.calls += 1 + if isinstance(self.result, BaseException): + raise self.result + return self.result + + +def _reader(fetch=None, lister=None, **kw) -> RepoReader: + return RepoReader(fetch if fetch is not None else Fetch(), lister, kind=kw.pop("kind", "forge"), **kw) + + +def _run_threads(target, count: int = THREADS) -> tuple[list[threading.Thread], list[object]]: + """Start ``count`` threads that pass a barrier together, then call ``target``.""" + barrier = threading.Barrier(count) + results: list[object] = [] + results_lock = threading.Lock() + + def work(): + barrier.wait(WAIT) + value = target() + with results_lock: + results.append(value) + + threads = [threading.Thread(target=work, daemon=True) for _ in range(count)] + for thread in threads: + thread.start() + return threads, results + + +def _join(threads: list[threading.Thread]) -> None: + for thread in threads: + thread.join(WAIT) + assert not any(thread.is_alive() for thread in threads) + + +class TestModuleConstants: + def test_the_caps_match_the_design_budget(self): + assert MAX_RUN_READS == 200 + assert MAX_CHUNK_READS == 16 + + def test_the_reader_kinds_are_the_record_vocabulary(self): + assert READER_KINDS == ("forge", "repo-dir") + + def test_an_unknown_kind_is_refused(self): + with pytest.raises(ValueError, match="kind"): + RepoReader(Fetch(), None, kind="local") + + def test_the_default_caps_are_the_module_constants(self): + fetch = Fetch({f"f{i}": "x" for i in range(MAX_CHUNK_READS + 1)}) + chunk = _reader(fetch).chunk_reader() + for i in range(MAX_CHUNK_READS): + assert chunk(f"f{i}") == "x" + assert chunk(f"f{MAX_CHUNK_READS}") is None + assert len(fetch.calls) == MAX_CHUNK_READS + + +class TestSingleFlight: + def test_eight_threads_reading_one_path_fetch_it_once(self): + fetch = BlockingFetch({"src/A.java": "class A {}"}) + reader = _reader(fetch) + threads, results = _run_threads(lambda: reader.read("src/A.java")) + try: + assert fetch.entered.wait(WAIT) + time.sleep(0.2) + assert fetch.calls == ["src/A.java"] + finally: + fetch.release.set() + _join(threads) + assert fetch.calls == ["src/A.java"] + assert results == ["class A {}"] * THREADS + assert reader.stats()["reads"] == 1 + + def test_eight_chunk_readers_on_one_in_flight_path_fetch_once_and_count_once(self): + fetch = BlockingFetch({"src/A.java": "class A {}"}) + reader = _reader(fetch, run_cap=1, chunk_cap=1) + threads, results = _run_threads(lambda: reader.chunk_reader()("src/A.java")) + try: + assert fetch.entered.wait(WAIT) + time.sleep(0.2) + assert fetch.calls == ["src/A.java"] + finally: + fetch.release.set() + _join(threads) + assert results == ["class A {}"] * THREADS + assert reader.stats() == {"reads": 1, "read_cap_hit": False, "listing": None} + + def test_eight_threads_listing_call_the_lister_once(self): + entered = threading.Event() + release = threading.Event() + listing = PathListing(paths=("a.py", "b.py"), complete=True) + calls: list[int] = [] + + def lister(): + calls.append(1) + entered.set() + release.wait(WAIT) + return listing + + reader = _reader(lister=lister) + threads, results = _run_threads(reader.listing) + try: + assert entered.wait(WAIT) + time.sleep(0.2) + assert len(calls) == 1 + assert reader.stats()["listing"] is None + finally: + release.set() + _join(threads) + assert len(calls) == 1 + assert results == [listing] * THREADS + assert reader.stats()["listing"] == {"paths": 2, "complete": True} + + def test_a_raising_fetch_releases_its_waiters_with_none(self): + entered = threading.Event() + release = threading.Event() + calls: list[str] = [] + + def fetch(path): + calls.append(path) + entered.set() + release.wait(WAIT) + raise OSError("boom") + + reader = _reader(fetch) + threads, results = _run_threads(lambda: reader.read("x.py"), count=3) + assert entered.wait(WAIT) + release.set() + _join(threads) + assert results == [None, None, None] + assert calls == ["x.py"] + + +class TestCache: + def test_a_missing_path_is_fetched_once(self): + fetch = Fetch() + reader = _reader(fetch) + assert reader.read("missing.py") is None + assert reader.read("missing.py") is None + assert fetch.calls == ["missing.py"] + assert reader.stats()["reads"] == 1 + + def test_a_hit_is_fetched_once(self): + fetch = Fetch({"a.py": "A"}) + reader = _reader(fetch) + assert reader.read("a.py") == "A" + assert reader.chunk_reader()("a.py") == "A" + assert fetch.calls == ["a.py"] + + def test_a_raising_fetch_reads_as_none_logs_at_debug_and_is_cached(self, caplog): + def fetch(path): + raise RuntimeError("transport down") + + reader = _reader(fetch) + with caplog.at_level(logging.DEBUG, logger="prxref"): + assert reader.read("a.py") is None + assert reader.read("a.py") is None + debug = [r for r in caplog.records if "a.py" in r.getMessage()] + assert len(debug) == 1 + assert debug[0].levelno == logging.DEBUG + assert "transport down" in debug[0].getMessage() + assert reader.stats()["reads"] == 1 + + def test_a_non_str_result_reads_as_none(self): + fetch = Fetch({"a.py": b"bytes", "b.py": 42}) + reader = _reader(fetch) + assert reader.read("a.py") is None + assert reader.chunk_reader()("b.py") is None + + +class TestChunkCap: + def test_the_fourth_uncached_path_is_refused_without_a_fetch(self): + fetch = Fetch({p: p.upper() for p in "abcdefg"}) + reader = _reader(fetch, chunk_cap=3) + chunk = reader.chunk_reader() + assert [chunk(p) for p in "abc"] == ["A", "B", "C"] + assert reader.stats()["read_cap_hit"] is False + assert chunk("d") is None + assert fetch.calls == ["a", "b", "c"] + assert reader.stats()["read_cap_hit"] is True + + def test_a_cached_path_still_reads_after_the_cap(self): + fetch = Fetch({p: p.upper() for p in "abcd"}) + reader = _reader(fetch, chunk_cap=3) + chunk = reader.chunk_reader() + for p in "abc": + chunk(p) + assert chunk("d") is None + assert chunk("a") == "A" + assert fetch.calls == ["a", "b", "c"] + + def test_a_refused_path_is_not_cached(self): + fetch = Fetch({p: p.upper() for p in "abcd"}) + reader = _reader(fetch, chunk_cap=3) + chunk = reader.chunk_reader() + for p in "abc": + chunk(p) + assert chunk("d") is None + assert reader.chunk_reader()("d") == "D" + assert fetch.calls == ["a", "b", "c", "d"] + + def test_a_second_chunk_reader_gets_its_own_budget(self): + fetch = Fetch({p: p.upper() for p in "abcdefg"}) + reader = _reader(fetch, chunk_cap=3) + first = reader.chunk_reader() + for p in "abcd": + first(p) + second = reader.chunk_reader() + assert [second(p) for p in "def"] == ["D", "E", "F"] + assert second("g") is None + assert fetch.calls == ["a", "b", "c", "d", "e", "f"] + assert reader.stats()["reads"] == 6 + + def test_each_call_returns_a_new_reader(self): + reader = _reader() + assert reader.chunk_reader() is not reader.chunk_reader() + + +class TestRunCap: + def test_two_chunk_readers_share_the_run_cap(self): + fetch = Fetch({p: p.upper() for p in "abcdefgh"}) + reader = _reader(fetch, run_cap=4, chunk_cap=3) + first = reader.chunk_reader() + second = reader.chunk_reader() + assert [first(p) for p in "abc"] == ["A", "B", "C"] + assert second("d") == "D" + assert reader.stats()["read_cap_hit"] is False + assert second("e") is None + assert first("f") is None + assert reader.chunk_reader()("g") is None + assert fetch.calls == ["a", "b", "c", "d"] + assert reader.stats() == {"reads": 4, "read_cap_hit": True, "listing": None} + + def test_a_cached_path_reads_after_the_run_cap(self): + fetch = Fetch({p: p.upper() for p in "abcde"}) + reader = _reader(fetch, run_cap=2, chunk_cap=3) + first = reader.chunk_reader() + first("a") + first("b") + second = reader.chunk_reader() + assert second("c") is None + assert second("a") == "A" + assert fetch.calls == ["a", "b"] + + +class TestReadIgnoresCaps: + def test_read_fetches_past_both_caps_and_counts_in_reads(self): + fetch = Fetch({p: p.upper() for p in "abcde"}) + reader = _reader(fetch, run_cap=1, chunk_cap=1) + assert [reader.read(p) for p in "abcde"] == ["A", "B", "C", "D", "E"] + assert fetch.calls == list("abcde") + assert reader.stats() == {"reads": 5, "read_cap_hit": False, "listing": None} + + def test_read_does_not_spend_the_chunk_readers_budget(self): + fetch = Fetch({p: p.upper() for p in "abcdef"}) + reader = _reader(fetch, run_cap=1, chunk_cap=1) + for p in "abc": + reader.read(p) + chunk = reader.chunk_reader() + assert chunk("a") == "A" + assert chunk("d") == "D" + assert chunk("e") is None + assert fetch.calls == ["a", "b", "c", "d"] + assert reader.stats() == {"reads": 4, "read_cap_hit": True, "listing": None} + + def test_read_serves_a_path_a_chunk_reader_fetched(self): + fetch = Fetch({"a": "A"}) + reader = _reader(fetch) + reader.chunk_reader()("a") + assert reader.read("a") == "A" + assert fetch.calls == ["a"] + + +class TestExclude: + @staticmethod + def _exclude(path: str) -> bool: + return path.endswith("expected.json") or path.startswith(".env") + + def test_an_excluded_path_is_never_fetched_nor_counted(self): + fetch = Fetch({"cases/expected.json": "{}", ".env": "SECRET=x", "a.py": "A"}) + reader = _reader(fetch, exclude=self._exclude, chunk_cap=1, run_cap=1) + assert reader.read("cases/expected.json") is None + chunk = reader.chunk_reader() + assert chunk(".env") is None + assert chunk("cases/expected.json") is None + assert chunk("a.py") == "A" + assert fetch.calls == ["a.py"] + assert reader.stats() == {"reads": 1, "read_cap_hit": False, "listing": None} + + def test_excluded_paths_are_dropped_from_the_listing_and_complete_is_kept(self): + raw = PathListing(paths=(".env", "a.py", "cases/expected.json", "src/B.java"), complete=False) + reader = _reader(lister=Lister(raw), exclude=self._exclude) + assert reader.listing() == PathListing(paths=("a.py", "src/B.java"), complete=False) + assert reader.stats()["listing"] == {"paths": 2, "complete": False} + + def test_a_raising_exclude_fails_closed(self): + def exclude(path): + raise ValueError("bad glob") + + fetch = Fetch({"a.py": "A"}) + reader = _reader(fetch, lister=Lister(PathListing(paths=("a.py",), complete=True)), exclude=exclude) + assert reader.read("a.py") is None + assert reader.chunk_reader()("a.py") is None + assert fetch.calls == [] + assert reader.listing() == PathListing(paths=(), complete=True) + + def test_without_exclude_the_listing_is_returned_as_given(self): + raw = PathListing(paths=("a.py",), complete=True) + assert _reader(lister=Lister(raw)).listing() is raw + + +class TestListing: + def test_no_lister_gives_none(self): + reader = _reader(lister=None) + assert reader.listing() is None + assert reader.stats()["listing"] is None + + def test_the_lister_runs_once_and_its_result_is_cached(self): + lister = Lister(PathListing(paths=("a.py", "b.py", "c.py"), complete=True)) + reader = _reader(lister=lister) + assert reader.stats()["listing"] is None + first = reader.listing() + assert reader.listing() is first + assert lister.calls == 1 + assert reader.stats()["listing"] == {"paths": 3, "complete": True} + + def test_a_raising_lister_gives_none_and_is_not_retried(self, caplog): + lister = Lister(RuntimeError("tree API down")) + reader = _reader(lister=lister) + with caplog.at_level(logging.DEBUG, logger="prxref"): + assert reader.listing() is None + assert reader.listing() is None + assert lister.calls == 1 + assert reader.stats()["listing"] is None + assert any("tree API down" in r.getMessage() and r.levelno == logging.DEBUG for r in caplog.records) + + def test_a_none_listing_is_cached(self): + lister = Lister(None) + reader = _reader(lister=lister) + assert reader.listing() is None + assert reader.listing() is None + assert lister.calls == 1 + assert reader.stats()["listing"] is None + + @pytest.mark.parametrize("bad", [("a.py",), ["a.py"], {"paths": ("a.py",), "complete": True}, "a.py"]) + def test_a_non_path_listing_gives_none(self, bad): + lister = Lister(bad) + reader = _reader(lister=lister) + assert reader.listing() is None + assert reader.listing() is None + assert lister.calls == 1 + assert reader.stats()["listing"] is None + + def test_stats_never_calls_the_lister_or_fetch(self): + fetch = Fetch() + lister = Lister(PathListing(paths=(), complete=True)) + reader = _reader(fetch, lister=lister) + reader.stats() + assert lister.calls == 0 + assert fetch.calls == [] + + def test_stats_is_a_snapshot(self): + reader = _reader(Fetch({"a": "A"})) + before = reader.stats() + reader.read("a") + assert before["reads"] == 0 + assert reader.stats()["reads"] == 1 + + +class FakeForge: + """A forge fake with the optional reader and listing, every call recorded with its ref and sha.""" + + name = "fake" + + def __init__(self, files: dict[str, str] | None = None, listing: object = None): + self.files = files or {} + self.listing = listing + self.reads: list[tuple[object, str, str]] = [] + self.lists: list[tuple[object, str]] = [] + + def get_file_content(self, ref, path, *, sha): + self.reads.append((ref, path, sha)) + return self.files.get(path) + + def list_paths(self, ref, *, sha): + self.lists.append((ref, sha)) + return self.listing + + +class ReadOnlyForge: + """A forge fake with ``get_file_content`` and no ``list_paths``.""" + + name = "read-only" + + def __init__(self, files: dict[str, str] | None = None): + self.files = files or {} + self.reads: list[tuple[object, str, str]] = [] + + def get_file_content(self, ref, path, *, sha): + self.reads.append((ref, path, sha)) + return self.files.get(path) + + +class NoReaderForge: + """A forge fake with neither optional method.""" + + name = "no-reader" + + +class TestForgeReader: + def test_no_get_file_content_gives_none(self): + assert forge_reader(NoReaderForge(), REF, SHA) is None + + @pytest.mark.parametrize("sha", ["", None]) + def test_an_empty_sha_gives_none(self, sha): + forge = FakeForge({"a.py": "A"}) + assert forge_reader(forge, REF, sha) is None + assert forge.reads == [] + + def test_ref_and_sha_pass_through_to_both_calls(self): + listing = PathListing(paths=("a.py",), complete=True) + forge = FakeForge({"a.py": "A"}, listing) + reader = forge_reader(forge, REF, SHA) + assert reader is not None + assert reader.kind == "forge" + assert reader.read("a.py") == "A" + assert reader.chunk_reader()("b.py") is None + assert reader.listing() is listing + assert forge.reads == [(REF, "a.py", SHA), (REF, "b.py", SHA)] + assert forge.lists == [(REF, SHA)] + assert reader.stats() == {"reads": 2, "read_cap_hit": False, "listing": {"paths": 1, "complete": True}} + + def test_no_list_paths_gives_no_listing(self): + reader = forge_reader(ReadOnlyForge({"a.py": "A"}), REF, SHA) + assert reader is not None + assert reader.read("a.py") == "A" + assert reader.listing() is None + assert reader.stats()["listing"] is None + + def test_a_raising_get_file_content_reads_as_none_at_debug(self, caplog): + class Raising(ReadOnlyForge): + def get_file_content(self, ref, path, *, sha): + raise ConnectionError("reset by peer") + + reader = forge_reader(Raising(), REF, SHA) + with caplog.at_level(logging.DEBUG, logger="prxref"): + assert reader.read("a.py") is None + assert [r.levelno for r in caplog.records if "reset by peer" in r.getMessage()] == [logging.DEBUG] + + def test_exclude_is_passed_through(self): + forge = FakeForge({"a.py": "A", "k.pem": "KEY"}, PathListing(paths=("a.py", "k.pem"), complete=True)) + reader = forge_reader(forge, REF, SHA, exclude=lambda p: p.endswith(".pem")) + assert reader.read("k.pem") is None + assert reader.listing() == PathListing(paths=("a.py",), complete=True) + assert forge.reads == [] + + def test_a_replay_forge_reads_and_lists_through_its_inner_forge(self): + listing = PathListing(paths=("src/A.java", "src/B.java"), complete=False) + inner = FakeForge({"src/A.java": "class A {}"}, listing) + reader = forge_reader(ReplayForge(inner, hide_threads=True), REF, SHA) + assert reader is not None + assert reader.read("src/A.java") == "class A {}" + assert reader.chunk_reader()("src/B.java") is None + assert reader.listing() == listing + assert inner.reads == [(REF, "src/A.java", SHA), (REF, "src/B.java", SHA)] + assert inner.lists == [(REF, SHA)] + assert reader.stats() == {"reads": 2, "read_cap_hit": False, "listing": {"paths": 2, "complete": False}} + + def test_a_replay_forge_over_an_inner_without_list_paths_lists_none_once(self): + replay = ReplayForge(ReadOnlyForge({"a.py": "A"})) + reader = forge_reader(replay, REF, SHA) + assert reader is not None + assert reader.read("a.py") == "A" + assert reader.listing() is None + assert reader.stats()["listing"] is None + + +class TestRepoDirReader: + @staticmethod + def _tree(tmp_path): + (tmp_path / "src").mkdir() + (tmp_path / "src" / "A.java").write_text("class A {}\n") + (tmp_path / "src" / "B.java").write_text("class B {}\n") + (tmp_path / "cases").mkdir() + (tmp_path / "cases" / "expected.json").write_text("{}\n") + return RepoDir(tmp_path) + + def test_it_reads_lists_and_is_kind_repo_dir(self, tmp_path): + reader = repo_dir_reader(self._tree(tmp_path)) + assert reader.kind == "repo-dir" + assert reader.read("src/A.java") == "class A {}\n" + assert reader.read("src/missing.java") is None + assert reader.listing() == PathListing( + paths=("cases/expected.json", "src/A.java", "src/B.java"), complete=True + ) + assert reader.stats() == {"reads": 2, "read_cap_hit": False, "listing": {"paths": 3, "complete": True}} + + def test_exclude_applies(self, tmp_path): + reader = repo_dir_reader(self._tree(tmp_path), exclude=lambda p: p.endswith("expected.json")) + assert reader.read("cases/expected.json") is None + assert reader.listing() == PathListing(paths=("src/A.java", "src/B.java"), complete=True) + assert reader.stats()["reads"] == 0 + + def test_the_default_chunk_cap_applies(self, tmp_path): + for i in range(MAX_CHUNK_READS + 1): + (tmp_path / f"f{i:02d}.txt").write_text(f"{i}\n") + reader = repo_dir_reader(RepoDir(tmp_path)) + chunk = reader.chunk_reader() + assert [chunk(f"f{i:02d}.txt") for i in range(MAX_CHUNK_READS)] == [f"{i}\n" for i in range(MAX_CHUNK_READS)] + assert chunk(f"f{MAX_CHUNK_READS:02d}.txt") is None + assert reader.stats() == {"reads": MAX_CHUNK_READS, "read_cap_hit": True, "listing": None} + + def test_an_incomplete_walk_stays_incomplete(self, tmp_path, monkeypatch): + from prxref.forges import repo_dir as repo_dir_module + + monkeypatch.setattr(repo_dir_module, "_MAX_LISTED_FILES", 1) + listing = repo_dir_reader(self._tree(tmp_path)).listing() + assert listing is not None + assert listing.complete is False + assert len(listing.paths) == 1 From f8d50edbf5deb43ddc1b9ce86141ef7bf0d34d8c Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:17:28 -0700 Subject: [PATCH 10/24] feat: cross-chunk and diff-file definitions for repository context (#17 T5) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit New pure module src/prxref/repo_crosschunk.py with diff_definitions(chunk, all_files, read) -> list[ContextEntry]. It closes issue #17 miss (b): a type changed in one chunk (TransportConfig's new compact constructor) is now shown to the worker of another chunk that calls it (ConnectorService), together with the changed lines inside that type, not just its declaration. Names are the union of repo_context.referenced_names over the chunk files' added lines. Every diff file with definition regexes is searched with find_definitions, reading at the PR head or, with no reader or a None read, scanning each contiguous run of known hunk lines on its own so entries keep true new-file line numbers and never show a line the hunks lack. Skips: this chunk's own hunk lines, js/python files in the chunk when a reader is given (referenced_definitions already covers them, D1), and removed files. A hit in a file outside the chunk that has + lines anywhere in the PR is "cross-chunk" and also yields one change entry per contiguous + run after the definition line and before the next definition indented no deeper (capped at MAX_CHANGE_LINES = 12 plus a "... N more changed lines" line); every other hit is "diff-file". Output is ordered by (REASONS rank, path, line) and deduplicated on (path, line). No budget: T7 owns it. Tests: tests/test_repo_context_crosschunk.py, including the issue17 fixture end to end with a reader and from hunks only. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/repo_crosschunk.py | 234 ++++++++++++++++ tests/test_repo_context_crosschunk.py | 388 ++++++++++++++++++++++++++ 2 files changed, 622 insertions(+) create mode 100644 src/prxref/repo_crosschunk.py create mode 100644 tests/test_repo_context_crosschunk.py diff --git a/src/prxref/repo_crosschunk.py b/src/prxref/repo_crosschunk.py new file mode 100644 index 0000000..3bd1b21 --- /dev/null +++ b/src/prxref/repo_crosschunk.py @@ -0,0 +1,234 @@ +"""Cross-chunk and diff-file definitions for one worker chunk. + +``chunk_context.referenced_definitions`` shows a worker the definition of a +name its added lines reference only when that definition sits in the SAME +changed file, and it knows no Java. When a PR changes a type in one chunk and +another chunk calls it, the calling chunk's worker never sees the new +invariant. :func:`diff_definitions` closes that gap for the ``diff`` level of +``PRXREF_REPO_CONTEXT``: it searches every file of the PR, not just the +chunk's own, and for a type changed in another chunk it also shows the changed +lines inside that type. + +File text comes from the ``read(path) -> str | None`` callable at the PR head. +With no reader, or when a read returns ``None``, the only text known is the +file's hunks (its context and ``+`` lines), and entries are built from those +alone, so a forge without file reads still gets hunk-based entries. The module +is pure: stdlib plus :mod:`prxref.repo_context` and :mod:`prxref.chunk_context`, +with no I/O except through ``read``. It applies no budget; per-entry caps are +the only limits here. +""" +from __future__ import annotations + +from collections.abc import Callable, Sequence + +from . import chunk_context +from .repo_context import ( + REASONS, + ContextEntry, + definition_regexes, + find_definitions, + language_of, + referenced_names, +) + +MAX_CHANGE_LINES = 12 + +_SAME_FILE_LANGUAGES = frozenset({"js", "python"}) + + +def _diff_lines(f: object) -> list[object]: + return [ + line + for hunk in getattr(f, "hunks", None) or [] + for line in getattr(hunk, "lines", None) or [] + ] + + +def _hunk_text(f: object) -> dict[int, str]: + known: dict[int, str] = {} + for line in _diff_lines(f): + new_line = getattr(line, "new_line", None) + if getattr(line, "kind", " ") in (" ", "+") and isinstance(new_line, int): + known.setdefault(new_line, getattr(line, "text", "")) + return known + + +def _added_text(f: object) -> dict[int, str]: + added: dict[int, str] = {} + for line in _diff_lines(f): + new_line = getattr(line, "new_line", None) + if getattr(line, "kind", " ") == "+" and isinstance(new_line, int): + added.setdefault(new_line, getattr(line, "text", "")) + return added + + +def _runs(numbers: Sequence[int]) -> list[list[int]]: + runs: list[list[int]] = [] + for number in sorted(numbers): + if runs and number == runs[-1][-1] + 1: + runs[-1].append(number) + else: + runs.append([number]) + return runs + + +def _hunk_definitions( + known: dict[int, str], + names: Sequence[str], + language: str, + skip: frozenset[int], +) -> list[tuple[str, int, str]]: + found: list[tuple[str, int, str]] = [] + seen: set[str] = set() + for run in _runs(list(known)): + wanted = [name for name in names if name not in seen] + if not wanted: + break + offset = run[0] - 1 + local_skip = frozenset(number - offset for number in run if number in skip) + text = "\n".join(known[number] for number in run) + for symbol, line, body in find_definitions( + text, wanted, language=language, skip_lines=local_skip + ): + seen.add(symbol) + found.append((symbol, line + offset, body)) + return found + + +def _indent(text: str) -> int: + return len(text) - len(text.lstrip()) + + +def _window_end(known: dict[int, str], line: int, language: str) -> int | None: + regexes = definition_regexes(language) + depth = _indent(known.get(line, "")) + for number in sorted(known): + if number <= line: + continue + text = known[number] + if _indent(text) <= depth and any(regex.match(text) for regex in regexes): + return number + return None + + +def _change_text(lines: Sequence[str]) -> str: + shown = [text.rstrip() for text in lines[:MAX_CHANGE_LINES]] + hidden = len(lines) - len(shown) + if hidden > 0: + shown.append(f"… {hidden} more changed lines") + return "\n".join(shown) + + +def _change_entries( + path: str, + symbol: str, + line: int, + end: int | None, + added: dict[int, str], +) -> list[ContextEntry]: + inside = [n for n in added if n > line and (end is None or n < end)] + return [ + ContextEntry( + path=path, + line=run[0], + symbol=symbol, + kind="definition", + reason="cross-chunk", + text=_change_text([added[n] for n in run]), + ) + for run in _runs(inside) + ] + + +def diff_definitions( + chunk: Sequence[object], + all_files: Sequence[object], + read: Callable[[str], str | None] | None, +) -> list[ContextEntry]: + """Definitions from the PR's diff files for the names one chunk's added lines reference. + + ``chunk`` holds the worker's ``triage.FileDiff`` records and ``all_files`` + the whole PR's, both duck-typed on ``path``, ``status`` and ``hunks`` + (whose ``lines`` carry ``kind``, ``text`` and ``new_line``). The wanted + names are the union, in first-appearance order, of + :func:`~prxref.repo_context.referenced_names` over each chunk file's added + lines in that file's language. + + Every file in ``all_files`` with definition regexes is searched with + :func:`~prxref.repo_context.find_definitions`, reading its text with + ``read`` or, when ``read`` is ``None`` or returns ``None``, taking it from + the file's hunks: each contiguous run of known new-file lines is scanned + on its own, so an entry carries true line numbers and never shows a line + the hunks lack. A removed file is skipped, and so are the lines of this + chunk's own hunks. A js or python file in the chunk is skipped entirely + when ``read`` is given, because ``referenced_definitions`` already covers + it; a Java file in the chunk is searched outside its hunks. + + A hit in a file outside the chunk that has ``+`` lines anywhere in the PR + has reason ``"cross-chunk"``, and every other hit ``"diff-file"``. A + cross-chunk hit also yields one entry per contiguous run of ``+`` lines + after the definition line and before the next definition that is not + indented deeper than it (or the end of the known text): same symbol and + reason, ``line`` the run's first new-file line, and the run's text capped + at :data:`MAX_CHANGE_LINES` lines plus a ``… N more changed lines`` line. + Every entry's kind is ``"definition"``. + + Entries are ordered by ``(reason rank in REASONS, path, line)`` and + deduplicated on ``(path, line)``, the first kept. No budget is applied. + """ + own_files = chunk_context.chunk_files(chunk) + own = {entry.path: entry for entry in reversed(own_files)} + names: dict[str, None] = {} + for entry in own_files: + for name in referenced_names(entry.added, language_of(entry.path)): + names.setdefault(name, None) + if not names: + return [] + wanted = list(names) + changed = { + getattr(f, "path", "") + for f in all_files + if any(getattr(line, "kind", " ") == "+" for line in _diff_lines(f)) + } + + collected: list[ContextEntry] = [] + searched: set[str] = set() + for f in all_files: + path = getattr(f, "path", "") or "" + if not path or path in searched: + continue + searched.add(path) + if getattr(f, "status", "") == "removed": + continue + language = language_of(path) + if not definition_regexes(language): + continue + mine = own.get(path) + if mine is not None and read is not None and language in _SAME_FILE_LANGUAGES: + continue + skip = mine.hunk_lines if mine is not None else frozenset() + text = read(path) if read is not None else None + if isinstance(text, str): + known = dict(enumerate(text.splitlines(), start=1)) + found = find_definitions(text, wanted, language=language, skip_lines=skip) + else: + known = _hunk_text(f) + found = _hunk_definitions(known, wanted, language, skip) + reason = "cross-chunk" if mine is None and path in changed else "diff-file" + for symbol, line, body in found: + collected.append(ContextEntry(path, line, symbol, "definition", reason, body)) + if reason == "cross-chunk": + added = _added_text(f) + for symbol, line, _ in found: + end = _window_end(known, line, language) + collected.extend(_change_entries(path, symbol, line, end, added)) + + collected.sort(key=lambda e: (REASONS.index(e.reason), e.path, e.line)) + out: list[ContextEntry] = [] + kept: set[tuple[str, int]] = set() + for entry in collected: + key = (entry.path, entry.line) + if key not in kept: + kept.add(key) + out.append(entry) + return out diff --git a/tests/test_repo_context_crosschunk.py b/tests/test_repo_context_crosschunk.py new file mode 100644 index 0000000..63fbf9f --- /dev/null +++ b/tests/test_repo_context_crosschunk.py @@ -0,0 +1,388 @@ +"""Tests for ``prxref.repo_crosschunk.diff_definitions`` (issue #17, miss b). + +The fixture half runs the real parser and chunker over +``tests/fixtures/issue17``: ``TransportConfig`` gains a compact constructor +that makes ``url`` and ``legacyUrl`` mutually exclusive in one chunk, while +``ConnectorService`` in another chunk calls ``new TransportConfig(...)``. The +rest builds duck-typed stand-ins for ``triage.FileDiff`` so each rule (reasons, +skips, change windows, caps, order, dedup) is pinned on its own. +""" +from __future__ import annotations + +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from prxref import repo_crosschunk +from prxref.repo_context import REASONS, ContextEntry +from prxref.repo_crosschunk import MAX_CHANGE_LINES, diff_definitions +from prxref.triage import build_chunks, parse_unified_diff + +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" +REPO = FIXTURE / "repo" +TRANSPORT_CONFIG = "src/main/java/com/acme/connectors/TransportConfig.java" +CONNECTOR_SERVICE = "src/main/java/com/acme/connectors/ConnectorService.java" +MIGRATION = "db/changelog/003-idempotency-unique.sql" +EXCLUSIVITY = "exactly one of url or legacyUrl must be set" + + +class _Reader: + """A ``read(path)`` callable over a dict that records every path asked for.""" + + def __init__(self, texts: dict[str, str]): + self.texts = texts + self.calls: list[str] = [] + + def __call__(self, path: str) -> str | None: + self.calls.append(path) + return self.texts.get(path) + + +class _RepoReader: + """A ``read(path)`` callable over the fixture's PR-head working tree.""" + + def __init__(self): + self.calls: list[str] = [] + + def __call__(self, path: str) -> str | None: + self.calls.append(path) + target = REPO / path + return target.read_text(encoding="utf-8") if target.is_file() else None + + +def _file(path: str, *hunks: tuple[int, list[str]], status: str = "modified") -> SimpleNamespace: + """A FileDiff stand-in; each hunk is ``(new_start, lines)`` with a ``+``/``-``/`` `` prefix per line.""" + built = [] + for new_start, body in hunks: + number = new_start + lines = [] + for raw in body: + kind, text = raw[0], raw[1:] + if kind == "-": + lines.append(SimpleNamespace(kind="-", text=text, new_line=None)) + else: + lines.append(SimpleNamespace(kind=kind, text=text, new_line=number)) + number += 1 + built.append(SimpleNamespace(lines=lines)) + return SimpleNamespace(path=path, status=status, hunks=built) + + +def _keys(entries: list[ContextEntry]) -> list[tuple[str, int, str, str, str]]: + return [(e.path, e.line, e.symbol, e.kind, e.reason) for e in entries] + + +def _fixture(): + files = parse_unified_diff((FIXTURE / "pr.diff").read_text(encoding="utf-8")) + chunks = build_chunks(files, max_files_per_chunk=1) + return files, chunks + + +def _chunk_holding(chunks, path: str): + return next(chunk for chunk in chunks if any(f.path == path for f in chunk)) + + +def _hunk_lines(files, path: str) -> dict[int, str]: + diff = next(f for f in files if f.path == path) + return { + line.new_line: line.text + for hunk in diff.hunks + for line in hunk.lines + if line.kind in (" ", "+") and line.new_line is not None + } + + +def _assert_shows_only_known_lines(entries: list[ContextEntry], known_by_path: dict[str, dict[int, str]]): + for entry in entries: + known = known_by_path[entry.path] + for offset, text in enumerate(entry.text.split("\n")): + if text.startswith("\N{HORIZONTAL ELLIPSIS} "): + continue + assert entry.line + offset in known, (entry.path, entry.line + offset) + assert known[entry.line + offset].rstrip() == text + + +class TestFixture: + def test_connector_service_chunk_gets_the_transport_config_change(self): + files, chunks = _fixture() + read = _RepoReader() + entries = diff_definitions(_chunk_holding(chunks, CONNECTOR_SERVICE), files, read) + + cross = [e for e in entries if e.reason == "cross-chunk" and e.symbol == "TransportConfig"] + assert any(EXCLUSIVITY in e.text for e in cross) + assert _keys(entries) == [ + (TRANSPORT_CONFIG, 9, "TransportConfig", "definition", "cross-chunk"), + (TRANSPORT_CONFIG, 11, "TransportConfig", "definition", "cross-chunk"), + ] + declaration, change = entries + assert declaration.text.startswith("public record TransportConfig(String url, String legacyUrl) {") + assert EXCLUSIVITY not in declaration.text + assert EXCLUSIVITY in change.text + assert "if (hasUrl == hasLegacyUrl) {" in change.text + head = (REPO / TRANSPORT_CONFIG).read_text(encoding="utf-8").splitlines() + assert change.text == "\n".join(text.rstrip() for text in head[10:19]) + + def test_reads_only_the_files_a_definition_regex_covers(self): + files, chunks = _fixture() + read = _RepoReader() + diff_definitions(_chunk_holding(chunks, CONNECTOR_SERVICE), files, read) + assert read.calls == [TRANSPORT_CONFIG, CONNECTOR_SERVICE] + + @pytest.mark.parametrize("read", [None, lambda path: None], ids=["no-reader", "reader-returns-none"]) + def test_without_file_text_the_entries_come_from_hunks(self, read): + files, chunks = _fixture() + chunk = _chunk_holding(chunks, CONNECTOR_SERVICE) + entries = diff_definitions(chunk, files, read) + + assert [(e.path, e.line, e.reason) for e in entries] == [ + (TRANSPORT_CONFIG, 9, "cross-chunk"), + (TRANSPORT_CONFIG, 11, "cross-chunk"), + ] + assert EXCLUSIVITY in entries[1].text + _assert_shows_only_known_lines(entries, {TRANSPORT_CONFIG: _hunk_lines(files, TRANSPORT_CONFIG)}) + assert entries == diff_definitions(chunk, files, _RepoReader()) + + @pytest.mark.parametrize("path", [TRANSPORT_CONFIG, MIGRATION]) + def test_the_other_chunks_get_nothing(self, path): + files, chunks = _fixture() + assert diff_definitions(_chunk_holding(chunks, path), files, _RepoReader()) == [] + + +CALLER = _file( + "app/Caller.java", + (10, [" void run() {", "+ Limits limits = Limits.defaults();", " }"]), +) + + +class TestReasons: + def test_unchanged_definition_in_another_diff_file_is_diff_file(self): + limits = _file("app/Limits.java", (4, [" int max;", "- int min;", " }"])) + read = _Reader({"app/Limits.java": "package app;\n\npublic final class Limits {\n int max;\n}\n"}) + entries = diff_definitions([CALLER], [CALLER, limits], read) + assert _keys(entries) == [("app/Limits.java", 3, "Limits", "definition", "diff-file")] + assert entries[0].text == "public final class Limits {\n int max;\n}" + + def test_a_changed_file_outside_the_chunk_is_cross_chunk_even_when_the_definition_is_untouched(self): + limits = _file("app/Limits.java", (6, [" class Note {", "+ // note", " }"])) + read = _Reader({ + "app/Limits.java": "package app;\n\npublic final class Limits {\n}\n\nclass Note {\n // note\n}\n", + }) + entries = diff_definitions([CALLER], [CALLER, limits], read) + assert _keys(entries) == [("app/Limits.java", 3, "Limits", "definition", "cross-chunk")] + + def test_java_definition_in_the_chunks_own_file_outside_its_hunks_is_diff_file(self): + widget = _file( + "p/Widget.java", + (7, [" Widget copy() {", "+ return new Widget();", " }"]), + ) + text = ( + "package p;\n\npublic class Widget {\n\n private int size;\n\n" + " Widget copy() {\n return new Widget();\n }\n}\n" + ) + entries = diff_definitions([widget], [widget], _Reader({"p/Widget.java": text})) + assert _keys(entries) == [("p/Widget.java", 3, "Widget", "definition", "diff-file")] + + +class TestSkips: + def test_lines_inside_the_chunks_own_hunks_are_skipped(self): + widget = _file( + "p/Widget.java", + (3, [" public class Widget {", "+ Widget copy() { return new Widget(); }", " }"]), + ) + text = "package p;\n\npublic class Widget {\n Widget copy() { return new Widget(); }\n}\n" + read = _Reader({"p/Widget.java": text}) + assert diff_definitions([widget], [widget], read) == [] + assert read.calls == ["p/Widget.java"] + + def test_a_python_file_in_the_chunk_is_skipped_when_a_reader_is_given(self): + service = _file("app/service.py", (20, [" def handle(req):", "+ return build_reply(req)"])) + other = _file("app/other.py", (4, [" x = 1", "-y = 2"])) + read = _Reader({ + "app/service.py": "import os\ndef build_reply(req):\n return req\n", + "app/other.py": "def build_reply(req):\n return None\n\nx = 1\n", + }) + entries = diff_definitions([service], [service, other], read) + assert _keys(entries) == [("app/other.py", 1, "build_reply", "definition", "diff-file")] + assert read.calls == ["app/other.py"] + + def test_a_js_file_in_the_chunk_is_skipped_when_a_reader_is_given(self): + page = _file("web/page.ts", (5, [" export function page() {", "+ return renderCard(props);"])) + read = _Reader({"web/page.ts": "export function renderCard(props) {\n return props;\n}\n"}) + assert diff_definitions([page], [page], read) == [] + assert read.calls == [] + + def test_a_removed_file_is_skipped(self): + caller = _file("app/Caller.java", (4, ["+ Legacy old = null;"])) + legacy = _file("old/Legacy.java", (0, ["-public class Legacy {}"]), status="removed") + read = _Reader({"old/Legacy.java": "public class Legacy {}\n"}) + assert diff_definitions([caller], [caller, legacy], read) == [] + assert "old/Legacy.java" not in read.calls + + def test_no_referenced_names_reads_nothing(self): + deletion = _file("app/Gone.java", (3, ["- Legacy old = null;"])) + other = _file("app/Legacy.java", (1, [" public class Legacy {}"])) + read = _Reader({"app/Legacy.java": "public class Legacy {}\n"}) + assert diff_definitions([deletion], [deletion, other], read) == [] + assert read.calls == [] + + +class TestChangeEntries: + def test_the_window_ends_at_the_next_top_level_definition(self): + caller = _file("app/Caller.java", (1, ["+First first = new First();"])) + target = _file( + "p/First.java", + (1, [ + " package p;", " ", " public class First {", "+ int added;", " }", " ", + " class Second {", "+ int other;", " }", + ]), + ) + entries = diff_definitions([caller], [caller, target], None) + assert [(e.line, e.symbol, e.text) for e in entries] == [ + (3, "First", "public class First {\n int added;\n}"), + (4, "First", " int added;"), + ] + + def test_a_nested_type_does_not_end_the_window_and_dedup_keeps_the_first(self): + caller = _file("app/Caller.java", (1, ["+ Outer.Inner x = new Outer.Inner(1);"])) + outer = _file( + "p/Outer.java", + (1, [ + " package p;", " ", " public class Outer {", " public record Inner(int a) {", + " }", " void go() {", " run();", "+ check();", "+ log();", + " }", " }", + ]), + ) + entries = diff_definitions([caller], [caller, outer], None) + assert [(e.line, e.symbol, e.reason) for e in entries] == [ + (3, "Outer", "cross-chunk"), + (4, "Inner", "cross-chunk"), + (8, "Outer", "cross-chunk"), + ] + assert entries[2].text == " check();\n log();" + + def test_python_windows_follow_indentation(self): + view = _file("app/views.py", (10, [" def view(request):", "+ cfg = Settings.load()"])) + settings = _file( + "app/settings.py", + (3, [" ", " class Settings:", "+ debug = False", " "]), + (11, [" ", " def helper():", "+ return 1"]), + ) + text = ( + "import os\n\n\nclass Settings:\n debug = False\n\n @classmethod\n" + " def load(cls):\n return cls()\n\n\ndef helper():\n return 1\n" + ) + entries = diff_definitions([view], [view, settings], _Reader({"app/settings.py": text})) + assert [(e.line, e.symbol, e.text) for e in entries] == [ + (4, "Settings", "class Settings:"), + (5, "Settings", " debug = False"), + (8, "load", " def load(cls):"), + ] + + def test_a_wholly_added_type_keeps_its_change_entry(self): + caller = _file("app/Caller.java", (1, ["+NewType t = new NewType(null, \"b\");"])) + body = [ + "package p;", "", "public record NewType(String a, String b) {", "", + " static final int LIMIT = 3;", "", " public NewType {", + " if (a == null && b == null) {", + " throw new IllegalArgumentException(\"a or b is required\");", + " }", " }", "}", + ] + added = _file("p/NewType.java", (1, ["+" + text for text in body]), status="added") + entries = diff_definitions([caller], [caller, added], None) + assert [(e.line, e.reason) for e in entries] == [(3, "cross-chunk"), (4, "cross-chunk")] + assert "a or b is required" not in entries[0].text + assert entries[1].text == "\n".join(body[3:12]) + + @pytest.mark.parametrize( + ("run", "tail"), + [(20, ["\N{HORIZONTAL ELLIPSIS} 8 more changed lines"]), (MAX_CHANGE_LINES, [])], + ) + def test_a_long_run_is_capped(self, run, tail): + caller = _file("app/Caller.java", (1, ["+ Big big = new Big();"])) + fields = [f" int f{i};" for i in range(1, run + 1)] + big = _file("p/Big.java", (1, [" public class Big {", *("+" + f for f in fields), " }"])) + entries = diff_definitions([caller], [caller, big], None) + assert MAX_CHANGE_LINES == 12 + assert [e.line for e in entries] == [1, 2] + assert entries[1].text.split("\n") == fields[:12] + tail + + +class TestHunkText: + def test_a_definition_at_the_edge_of_a_hunk_shows_no_unknown_line(self): + caller = _file("app/Caller.java", (1, ["+ Spill s = new Spill();"])) + spill = _file( + "lib/Spill.java", + (1, [" package lib;", "+", " public class Spill extends Base {"]), + (10, ["+ int late;", " }"]), + ) + entries = diff_definitions([caller], [caller, spill], None) + assert [(e.line, e.text) for e in entries] == [ + (3, "public class Spill extends Base {"), + (10, " int late;"), + ] + _assert_shows_only_known_lines(entries, {"lib/Spill.java": { + 1: "package lib;", 2: "", 3: "public class Spill extends Base {", 10: " int late;", 11: "}", + }}) + fields = [f" int {name};" for name in ("a", "b", "c", "d", "e", "f")] + full = "\n".join(["package lib;", "", "public class Spill extends Base {", *fields, " int late;", "}"]) + read = diff_definitions([caller], [caller, spill], _Reader({"lib/Spill.java": full + "\n"})) + assert [(e.line, e.text) for e in read] == [ + (3, "\n".join(["public class Spill extends Base {", *fields[:5]])), + (10, " int late;"), + ] + + def test_hunk_runs_keep_their_new_file_line_numbers(self): + caller = _file("app/Caller.java", (1, ["+ Far far = Near.make();"])) + target = _file( + "p/Types.java", + (40, [" class Near {", "+ int n1;", " }"]), + (90, [" class Far {", " }"]), + ) + entries = diff_definitions([caller], [caller, target], None) + assert [(e.line, e.symbol) for e in entries] == [(40, "Near"), (41, "Near"), (90, "Far")] + _assert_shows_only_known_lines(entries, {"p/Types.java": { + 40: "class Near {", 41: " int n1;", 42: "}", 90: "class Far {", 91: "}", + }}) + + +class TestOrder: + def test_reason_rank_then_path_then_line(self): + main = _file( + "app/Main.java", + (5, [" class Main {", "+ Zeta z = new Zeta(); Alpha a = Alpha.of(); Beta b = new Beta();", " }"]), + ) + alpha = _file("z/Alpha.java", (1, [" public class Alpha {", "+ int a;", " }"])) + beta = _file("b/Beta.java", (1, [" public class Beta {", "- int old;", " }"])) + zeta = _file("a/Zeta.java", (1, [" public class Zeta {", "+ int z;", " }"])) + entries = diff_definitions([main], [main, alpha, beta, zeta], None) + assert [(e.path, e.line, e.reason) for e in entries] == [ + ("a/Zeta.java", 1, "cross-chunk"), + ("a/Zeta.java", 2, "cross-chunk"), + ("z/Alpha.java", 1, "cross-chunk"), + ("z/Alpha.java", 2, "cross-chunk"), + ("b/Beta.java", 1, "diff-file"), + ] + assert all(e.reason in REASONS and e.kind == "definition" for e in entries) + + def test_names_are_the_union_over_the_chunks_files(self): + java = _file("app/Main.java", (1, ["+ Alpha a = null;"])) + python = _file("app/run.py", (1, ["+beta_result = compute_beta()"])) + alpha = _file("z/Alpha.java", (1, [" public class Alpha {", " }"])) + helpers = _file("lib/helpers.py", (1, [" def compute_beta():", " return 2"])) + entries = diff_definitions([java, python], [java, python, alpha, helpers], None) + assert [(e.path, e.symbol, e.reason) for e in entries] == [ + ("lib/helpers.py", "compute_beta", "diff-file"), + ("z/Alpha.java", "Alpha", "diff-file"), + ] + + def test_the_module_is_pure(self): + source = Path(repo_crosschunk.__file__).read_text(encoding="utf-8") + imports = sorted( + line.strip() for line in source.splitlines() if line.startswith(("import ", "from ")) + ) + assert imports == [ + "from . import chunk_context", + "from .repo_context import (", + "from __future__ import annotations", + "from collections.abc import Callable, Sequence", + ] From a8cf9ed91896c2c4823cce27e9432052574e18a9 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:20:50 -0700 Subject: [PATCH 11/24] feat: repository-context resolver (candidate files per referenced name) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Issue #17, task T3. New pure module src/prxref/repo_resolve.py turns a referencing file, its text and the names referenced_names reports into an ordered list of Candidate(name, path, reason) files for the caller to read. It does no I/O: reading candidates and running find_definitions is T7's job. - Imports first: Java same-organization imports (first two package segments) under the source root derived from the package line, wildcard and nested-type imports included; Python from-imports at the root and under src/, relative dots counted from the file; TS/JS relative import/export specifiers probed as .ts, .tsx, .d.ts, .js, /index.ts. - Then the Java same-package convention /.java for names no single-name import binds. - Then a name search over a listing: same-language files whose stem matches (exact, case-insensitive, Python snake_case), deepest shared directory first, at most 3 per name. - Dead-candidate filters: a complete listing drops absent import and convention paths; third-party Java, stdlib Python and bare TS/JS imports give no candidate and are not name-searched; languages without definition regexes resolve nothing. Tests: tests/test_repo_context_resolve.py (81 tests, pure), including the issue17 fixture: TransportConfig resolves by path-convention and the Spring imports give no candidate. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/repo_resolve.py | 472 +++++++++++++++++++++++++++ tests/test_repo_context_resolve.py | 506 +++++++++++++++++++++++++++++ 2 files changed, 978 insertions(+) create mode 100644 src/prxref/repo_resolve.py create mode 100644 tests/test_repo_context_resolve.py diff --git a/src/prxref/repo_resolve.py b/src/prxref/repo_resolve.py new file mode 100644 index 0000000..5933e72 --- /dev/null +++ b/src/prxref/repo_resolve.py @@ -0,0 +1,472 @@ +"""Candidate files for the definitions a changed file references (repository context). + +:func:`prxref.repo_context.referenced_names` says WHICH identifiers a chunk's +added lines use. This module says WHERE their definitions probably live: it +turns a referencing file plus those names into an ordered list of +:class:`Candidate` files. Reading the candidates and running +:func:`prxref.repo_context.find_definitions` over their text is the caller's +job, not this module's. + +The module is pure. It is stdlib plus :mod:`prxref.repo_context`, and it +performs no I/O, no reads and no network. Every candidate costs the caller a +read, and a miss costs one too, so the rules stay narrow and the resolution is +best effort by design: no build file, ``tsconfig`` path alias, ``sys.path`` or +classpath is consulted. + +Candidates come in three groups, highest rank first, each carrying a member of +:data:`prxref.repo_context.REASONS` as its ``reason``: + +- ``"import"``: Java imports from the file's own organization under the source + root its ``package`` line implies, Python ``from ... import`` statements, + and TS/JS relative ``import ... from`` and ``export ... from`` specifiers. +- ``"path-convention"``: a Java type from the file's own package, + ``/.java``. +- ``"name-search"``: listing files of the same language whose stem equals the + name, when the caller has a repository listing. +""" +from __future__ import annotations + +import posixpath +import re +import sys +from collections.abc import Collection, Sequence +from dataclasses import dataclass + +from .repo_context import definition_regexes, language_of + +MAX_NAME_SEARCH_PER_NAME = 3 + +_JAVA_PACKAGE_RE = re.compile(r"^[ \t]*package[ \t]+([\w$.]+)[ \t]*;", re.M) +_JAVA_IMPORT_RE = re.compile(r"^[ \t]*import[ \t]+(static[ \t]+)?([\w$.]+?)(\.\*)?[ \t]*;", re.M) +_JAVA_EXTERNAL_ROOTS = frozenset({"java", "javax"}) + +_PY_FROM_RE = re.compile( + r"^[ \t]*from[ \t]+(\.*)[ \t]*([A-Za-z_][\w.]*)?[ \t]+import[ \t]+(\([^)]*\)|\([^\n]*|[^\n]*)", + re.M, +) +_PY_ITEM_RE = re.compile(r"^([A-Za-z_]\w*)(?:\s+as\s+([A-Za-z_]\w*))?$") +_PY_STDLIB = sys.stdlib_module_names + +_JS_IDENT = r"[A-Za-z_$][\w$]*" +_JS_FROM_RE = re.compile( + r"(? str | None: + parts: list[str] = [] + for part in path.split("/"): + if part in ("", "."): + continue + if part == "..": + if not parts: + return None + parts.pop() + continue + parts.append(part) + return "/".join(parts) or None + + +def _java_root(directory: str, package: str) -> str | None: + if not package: + return None + package_dir = package.replace(".", "/") + if directory == package_dir: + return "" + if directory.endswith("/" + package_dir): + return directory[: -len(package_dir)] + return None + + +def _java_type_index(parts: Sequence[str]) -> int | None: + for index, part in enumerate(parts): + if part[:1].isupper(): + return index + return None + + +def _java( + directory: str, text: str, names: Sequence[str] +) -> tuple[list[_Found], list[_Found], set[str]]: + wanted = set(names) + package_match = _JAVA_PACKAGE_RE.search(text) + package = package_match.group(1) if package_match else "" + root = _java_root(directory, package) + org = package.split(".")[:2] if package else [] + statements = [ + (bool(m.group(1)), m.group(2).split("."), bool(m.group(3))) for m in _JAVA_IMPORT_RE.finditer(text) + ] + claimed = {parts[-1] for _, parts, wildcard in statements if not wildcard} + external: set[str] = set() + imports: list[_Found] = [] + for static, parts, wildcard in statements: + third_party = parts[0] in _JAVA_EXTERNAL_ROOTS or bool(org and parts[: len(org)] != org) + if third_party and not wildcard: + external.add(parts[-1]) + if static or third_party or root is None: + continue + index = _java_type_index(parts) + if wildcard: + for name in names: + if name in claimed: + continue + file_parts = parts[: index + 1] if index is not None else [*parts, name] + imports.append((name, name, root + "/".join(file_parts) + ".java")) + continue + name = parts[-1] + if name in wanted: + file_parts = parts[: index + 1] if index is not None else parts + imports.append((name, name, root + "/".join(file_parts) + ".java")) + conventions = [ + (name, name, f"{directory}/{name}.java" if directory else f"{name}.java") + for name in names + if name not in claimed + ] + return imports, conventions, external + + +def _py_items(clause: str) -> list[tuple[str, str]]: + body = re.sub(r"#[^\n]*", "", clause).strip() + if body.startswith("("): + body = body[1:].split(")", 1)[0] + else: + body = body.split(";", 1)[0] + items: list[tuple[str, str]] = [] + for raw in body.split(","): + item = " ".join(raw.split()) + if item == "*": + items.append(("*", "*")) + continue + match = _PY_ITEM_RE.match(item) + if match: + items.append((match.group(1), match.group(2) or match.group(1))) + return items + + +def _py_module_files(directory: str, dots: str, module: str) -> list[str]: + relative = module.replace(".", "/") + if dots: + base = "/".join([directory, *[".."] * (len(dots) - 1)]) + if relative: + probes = [f"{base}/{relative}.py", f"{base}/{relative}/__init__.py"] + else: + probes = [f"{base}/__init__.py"] + else: + probes = [ + f"{relative}.py", + f"{relative}/__init__.py", + f"src/{relative}.py", + f"src/{relative}/__init__.py", + ] + return [p for p in map(_normalize, probes) if p] + + +def _python(directory: str, text: str, names: Sequence[str]) -> tuple[list[_Found], set[str]]: + wanted = set(names) + imports: list[_Found] = [] + external: set[str] = set() + bound: set[str] = set() + stars: list[list[str]] = [] + for match in _PY_FROM_RE.finditer(text.replace("\\\n", " ")): + dots, module, clause = match.group(1), match.group(2) or "", match.group(3) + if not dots and not module: + continue + items = _py_items(clause) + bound.update(local for _, local in items if local != "*") + if not dots and module.split(".")[0] in _PY_STDLIB: + external.update(local for _, local in items if local != "*") + continue + files = _py_module_files(directory, dots, module) + for original, local in items: + if original == "*": + stars.append(files) + elif local in wanted: + imports.extend((local, original, f) for f in files) + for files in stars: + for name in names: + if name not in bound: + imports.extend((name, name, f) for f in files) + return imports, external + + +def _js_module_files(directory: str, specifier: str) -> list[str]: + if not (specifier in (".", "..") or specifier.startswith(("./", "../"))): + return [] + base = f"{directory}/{specifier}" if directory else specifier + lower = specifier.lower() + if specifier in (".", "..") or specifier.endswith("/"): + probes = [base.rstrip("/") + "/index.ts"] + elif lower.endswith(_JS_ASSET_SUFFIXES): + return [] + elif lower.endswith(".js"): + probes = [base[:-3] + suffix for suffix in _JS_ESM_PROBES] + elif lower.endswith(_JS_LITERAL_SUFFIXES): + probes = [base] + else: + probes = [base + suffix for suffix in _JS_PROBES] + return [p for p in map(_normalize, probes) if p] + + +def _js_clause(clause: str) -> tuple[list[tuple[str, str]], list[str]]: + pairs: list[tuple[str, str]] = [] + namespaces: list[str] = [] + brace = clause.find("{") + head = clause if brace == -1 else clause[:brace] + body = "" if brace == -1 else clause[brace + 1 : clause.rfind("}")] + for raw in head.split(","): + part = " ".join(raw.split()) + namespace = _JS_NAMESPACE_RE.match(part) + if namespace: + namespaces.append(namespace.group(1)) + elif part and part != "*" and re.fullmatch(_JS_IDENT, part): + pairs.append(("default", part)) + for raw in _JS_COMMENT_RE.sub("", body).split(","): + match = _JS_ITEM_RE.match(" ".join(raw.split())) + if match: + pairs.append((match.group(1), match.group(2) or match.group(1))) + return pairs, namespaces + + +def _js(directory: str, text: str, names: Sequence[str]) -> tuple[list[_Found], set[str]]: + wanted = set(names) + imports: list[_Found] = [] + external: set[str] = set() + for match in _JS_FROM_RE.finditer(text): + keyword, clause, specifier = match.group(1), match.group(2), match.group(4) + pairs, namespaces = _js_clause(clause) + files = _js_module_files(directory, specifier) + if not files: + if keyword == "import" and not specifier.startswith((".", "/")): + external.update(local for _, local in pairs) + continue + for original, local in pairs: + keys = (local,) if keyword == "import" else (original, local) + key = next((k for k in keys if k in wanted), None) + if key is None: + continue + target = local if original == "default" else original + imports.extend((key, target, f) for f in files) + if keyword != "import": + continue + for namespace in namespaces: + member_re = re.compile(rf"(? str: + dot = base.rfind(".") + stem = base[:dot] if dot > 0 else base + return stem[:-2] if stem.endswith(".d") else stem + + +def _snake_case(name: str) -> str: + return _SNAKE_RE.sub("_", name).lower() + + +def _shared_depth(left: str, right: str) -> int: + depth = 0 + for a, b in zip(left.split("/") if left else [], right.split("/") if right else [], strict=False): + if a != b: + break + depth += 1 + return depth + + +def _name_search( + source: str, language: str, names: Sequence[str], listing: Collection[str], skip: set[str] +) -> list[_Found]: + python = language == "python" + searched = [n for n in names if n not in skip and not (n.startswith("__") and n.endswith("__"))] + targets = {n.lower() for n in searched} + if python: + targets |= {_snake_case(n) for n in searched} + by_lower: dict[str, list[tuple[str, str]]] = {} + for raw in listing: + stem = _stem(raw.rsplit("/", 1)[-1]) + folded = stem.lower() + if folded not in targets or language_of(raw) != language: + continue + path = _normalize(raw) + if path is None or path == source: + continue + by_lower.setdefault(folded, []).append((stem, path)) + directory = posixpath.dirname(source) + + def ordered(paths: list[str]) -> list[str]: + return sorted(paths, key=lambda p: (-_shared_depth(posixpath.dirname(p), directory), p)) + + found: list[_Found] = [] + for name in searched: + folded = name.lower() + hits = by_lower.get(folded, []) + tiers = [ + [p for stem, p in hits if stem == name], + [p for stem, p in hits if stem != name], + ] + snake = _snake_case(name) + if python and snake != folded: + tiers.append([p for _, p in by_lower.get(snake, [])]) + picked: list[str] = [] + for tier in tiers: + for path in ordered(tier): + if path not in picked: + picked.append(path) + found.extend((name, name, p) for p in picked[:MAX_NAME_SEARCH_PER_NAME]) + return found + + +def resolve_candidates( + path: str, + text: str | None, + names: Sequence[str], + *, + listing: Collection[str] | None, + listing_complete: bool = False, +) -> list[Candidate]: + """Ordered candidate files that may define ``names``, as referenced from ``path``. + + ``path`` is the referencing file, repository-relative. ``text`` is its + content at the head, or the chunk's added lines joined with newlines when + there is no reader, so no ``package`` line or import is assumed present; + ``None`` means no imports are known. ``names`` comes from + :func:`prxref.repo_context.referenced_names`. ``listing`` holds the + repository's file paths, or is ``None``; ``listing_complete`` says the + listing was not truncated. + + Rules, by :func:`prxref.repo_context.language_of`: + + - Java. The source root is ``path``'s directory minus the package + directory of its first ``package`` line (``package com.acme.connectors;`` + in ``src/main/java/com/acme/connectors/X.java`` gives ``src/main/java/``); + with no package line, or a directory that does not end with the package + path, the root is unknown and no import resolves. Only an import sharing + the first two package segments with the file's own package counts: + ``import com.acme.x.T;`` gives ``com/acme/x/T.java`` when ``T`` is + referenced, and a wildcard ``import com.acme.x.*;`` gives + ``com/acme/x/.java`` for each referenced name no + single-name import binds. ``import static``, ``java.*``, ``javax.*`` and + other organizations' imports give nothing. A nested type is best effort: + the file is the first segment with an uppercase initial, so + ``a.b.Outer.Inner`` maps to ``a/b/Outer.java``. Then, as the + ``"path-convention"`` group, ``/.java`` for every name no + single-name import binds; a single-name import is any non-wildcard + import, static or not, of any organization, resolved or not. + - Python. ``from a.b import C`` (also parenthesized, aliased or ``*``) + gives ``a/b.py``, ``a/b/__init__.py``, ``src/a/b.py`` and + ``src/a/b/__init__.py``, in that order. Relative dots count from + ``path``'s directory (``.`` is that directory, each further dot one level + up), give ``/x.py`` and ``/x/__init__.py``, and for a bare + ``from . import C`` just ``/__init__.py``. ``import a.b`` gives + nothing (its names are attributes), and neither does a standard-library + module. + - TS/JS (``language_of`` reports ``"js"``). A relative specifier of + ``import {C} from``, ``import C from``, ``import * as N from`` (for each + referenced ``N.``) or ``export {C} from`` tries ``x.ts``, + ``x.tsx``, ``x.d.ts``, ``x.js`` and ``x/index.ts`` against ``path``'s + directory. A ``.js`` specifier (TypeScript ESM) tries the same stem as + ``.ts``, ``.tsx``, ``.d.ts`` and ``.js``; another script extension is + tried literally; an asset (``.css``, ``.json``, ``.svg`` and similar) + gives nothing. Bare package specifiers give nothing. + - Name search, every language, only when ``listing`` is not ``None``: + listing files of the same language whose stem (the basename less its + last extension and a ``.d`` before it) equals the name case-sensitively, + then case-insensitively, then, for Python, equals the snake_case form of + a CamelCase name. Within each tier the deepest directory shared with + ``path`` comes first, then the path. At most + :data:`MAX_NAME_SEARCH_PER_NAME` per name. Names bound by an import + judged external (another organization's Java import, a + standard-library Python module, a bare TS/JS package) and dunder names + are not searched. + + Candidates are ordered by group (imports, path conventions, name search), + then by the name's position in ``names``; within one name, imports keep + statement order and each rule's probe order above. Duplicates by + ``(name, path)`` keep the first, highest-ranked reason. ``path`` itself is + never a candidate. With a listing and ``listing_complete``, an import or + path-convention candidate the listing lacks is dropped. Paths are + normalized POSIX, and a candidate that would escape the repository root is + dropped. A language without definition regexes + (:func:`prxref.repo_context.definition_regexes`) returns ``[]``, since no + candidate could yield a definition. There is no error path: an + unresolvable name is simply absent. + """ + source = _normalize(path) + if source is None: + return [] + language = language_of(source) + if not definition_regexes(language): + return [] + ordered_names = list(dict.fromkeys(names)) + rank = {name: index for index, name in enumerate(ordered_names)} + directory = posixpath.dirname(source) + body = text or "" + conventions: list[_Found] = [] + if language == "java": + imports, conventions, external = _java(directory, body, ordered_names) + elif language == "python": + imports, external = _python(directory, body, ordered_names) + elif language == "js": + imports, external = _js(directory, body, ordered_names) + else: + imports, external = [], set() + listed: Collection[str] | None = None + if listing is not None and listing_complete: + listed = listing if isinstance(listing, (set, frozenset)) else frozenset(listing) + out: list[Candidate] = [] + seen: set[tuple[str, str]] = set() + + def admit(reason: str, found: list[_Found], *, filtered: bool) -> None: + for _, name, raw in sorted(found, key=lambda entry: rank[entry[0]]): + candidate = _normalize(raw) + if candidate is None or candidate == source: + continue + if filtered and listed is not None and candidate not in listed: + continue + if (name, candidate) in seen: + continue + seen.add((name, candidate)) + out.append(Candidate(name=name, path=candidate, reason=reason)) + + admit("import", imports, filtered=True) + admit("path-convention", conventions, filtered=True) + if listing is not None: + admit("name-search", _name_search(source, language, ordered_names, listing, external), filtered=False) + return out diff --git a/tests/test_repo_context_resolve.py b/tests/test_repo_context_resolve.py new file mode 100644 index 0000000..c41e2da --- /dev/null +++ b/tests/test_repo_context_resolve.py @@ -0,0 +1,506 @@ +"""Unit tests for :mod:`prxref.repo_resolve`, the repository-context resolver. + +The resolver is pure: it turns a referencing file, its text and the names its +added lines reference into ordered candidate files, and never reads one. These +tests pin the per-language import rules, the Java same-package convention, the +name search over a listing, the two dead-candidate filters and the ordering, +against the issue #17 fixture and small synthetic sources. No reads of the +repository under review, no forge, no network. +""" +from __future__ import annotations + +import dataclasses +from pathlib import Path + +import pytest + +from prxref.chunk_context import chunk_files +from prxref.repo_context import REASONS, language_of, referenced_names +from prxref.repo_resolve import MAX_NAME_SEARCH_PER_NAME, Candidate, resolve_candidates +from prxref.triage import parse_unified_diff + +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" +REPO = FIXTURE / "repo" +SERVICE = "src/main/java/com/acme/connectors/ConnectorService.java" +CONFIG = "src/main/java/com/acme/connectors/TransportConfig.java" +LISTING = sorted(p.relative_to(REPO).as_posix() for p in REPO.rglob("*") if p.is_file()) +SPRING = ("PathVariable", "PostMapping", "RequestBody", "RequestHeader") + + +def resolve(path, text, names, *, listing=None, complete=False): + return [ + (c.name, c.path, c.reason) + for c in resolve_candidates(path, text, names, listing=listing, listing_complete=complete) + ] + + +def fixture_added(path: str) -> tuple[str, ...]: + files = parse_unified_diff((FIXTURE / "pr.diff").read_text(encoding="utf-8")) + for chunk_file in chunk_files(files): + if chunk_file.path == path: + return chunk_file.added + raise AssertionError(f"{path} is not in the fixture diff") + + +def fixture_names(path: str) -> list[str]: + return referenced_names(fixture_added(path), language_of(path)) + + +def fixture_text(path: str) -> str: + return (REPO / path).read_text(encoding="utf-8") + + +def probes(stem: str, suffixes) -> list[str]: + return [stem + suffix for suffix in suffixes] + + +TS_PROBES = (".ts", ".tsx", ".d.ts", ".js", "/index.ts") +PY_PROBES = (".py", "/__init__.py") + + +class TestCandidateShape: + def test_candidate_is_frozen(self): + candidate = Candidate(name="A", path="a/A.java", reason="import") + with pytest.raises(dataclasses.FrozenInstanceError): + candidate.path = "b/A.java" + + def test_every_reason_is_an_admission_rank(self): + assert {"import", "path-convention", "name-search"} <= set(REASONS) + text = "package com.acme.billing;\nimport com.acme.shared.Money;\n" + found = resolve( + "src/main/java/com/acme/billing/Invoice.java", + text, + ["Money", "Tax", "Ledger"], + listing=["lib/Ledger.java"], + ) + assert {reason for _, _, reason in found} == {"import", "path-convention", "name-search"} + + def test_cap_constant(self): + assert MAX_NAME_SEARCH_PER_NAME == 3 + + +class TestFixture: + def test_fixture_names_include_the_record_and_the_spring_annotations(self): + names = fixture_names(SERVICE) + assert "TransportConfig" in names + assert set(SPRING) <= set(names) + + def test_transport_config_resolves_by_same_package_convention(self): + found = resolve(SERVICE, fixture_text(SERVICE), fixture_names(SERVICE), listing=LISTING, complete=True) + assert found == [("TransportConfig", CONFIG, "path-convention")] + + def test_spring_imports_give_no_candidate_even_without_a_listing(self): + found = resolve(SERVICE, fixture_text(SERVICE), fixture_names(SERVICE)) + assert not {name for name, _, _ in found} & set(SPRING) + assert found == [ + ("TransportConfig", CONFIG, "path-convention"), + ("Tenant", "src/main/java/com/acme/connectors/Tenant.java", "path-convention"), + ("Id", "src/main/java/com/acme/connectors/Id.java", "path-convention"), + ( + "CreateTransportRequest", + "src/main/java/com/acme/connectors/CreateTransportRequest.java", + "path-convention", + ), + ] + + def test_added_lines_as_text_lose_the_context_import(self): + text = "\n".join(fixture_added(SERVICE)) + found = resolve(SERVICE, text, fixture_names(SERVICE)) + names = [name for name, _, _ in found] + assert "RequestHeader" in names + assert not {"PathVariable", "PostMapping", "RequestBody"} & set(names) + + def test_the_changed_record_never_resolves_to_itself(self): + names = fixture_names(CONFIG) + assert names == ["TransportConfig"] + assert resolve(CONFIG, fixture_text(CONFIG), names, listing=LISTING, complete=True) == [] + assert resolve(CONFIG, fixture_text(CONFIG), names, listing=LISTING) == [] + + def test_in_org_import_resolves_under_the_derived_source_root(self): + text = fixture_text(SERVICE).replace( + "package com.acme.connectors;\n", "package com.acme.connectors;\n\nimport com.acme.shared.Tenant;\n", 1 + ) + tenant = "src/main/java/com/acme/shared/Tenant.java" + found = resolve(SERVICE, text, fixture_names(SERVICE), listing=[*LISTING, tenant], complete=True) + assert found == [ + ("Tenant", tenant, "import"), + ("TransportConfig", CONFIG, "path-convention"), + ] + + def test_sql_file_resolves_nothing(self): + path = "db/changelog/003-idempotency-unique.sql" + assert resolve(path, fixture_text(path), fixture_names(path), listing=LISTING) == [] + + +BILLING = "src/main/java/com/acme/billing/Invoice.java" +HEADER = "package com.acme.billing;\n\n" + + +class TestJava: + def test_explicit_import(self): + found = resolve(BILLING, HEADER + "import com.acme.shared.Money;\n", ["Money"]) + assert found == [("Money", "src/main/java/com/acme/shared/Money.java", "import")] + + def test_explicit_import_of_an_unreferenced_name_gives_nothing(self): + assert resolve(BILLING, HEADER + "import com.acme.shared.Money;\n", ["Tax"]) == [ + ("Tax", "src/main/java/com/acme/billing/Tax.java", "path-convention") + ] + + def test_wildcard_import_probes_each_name_then_conventions_follow(self): + found = resolve(BILLING, HEADER + "import com.acme.shared.*;\n", ["Money", "Tax"]) + assert found == [ + ("Money", "src/main/java/com/acme/shared/Money.java", "import"), + ("Tax", "src/main/java/com/acme/shared/Tax.java", "import"), + ("Money", "src/main/java/com/acme/billing/Money.java", "path-convention"), + ("Tax", "src/main/java/com/acme/billing/Tax.java", "path-convention"), + ] + + def test_wildcard_skips_a_name_a_single_type_import_binds(self): + text = HEADER + "import com.acme.shared.*;\nimport com.acme.other.Money;\n" + found = resolve(BILLING, text, ["Money"]) + assert found == [("Money", "src/main/java/com/acme/other/Money.java", "import")] + + def test_static_import_gives_no_candidate(self): + text = HEADER + "import static com.acme.shared.Money.Zero;\nimport static com.acme.shared.Money.*;\n" + assert resolve(BILLING, text, ["Zero"]) == [] + assert [r for _, _, r in resolve(BILLING, text, ["Money"])] == ["path-convention"] + + def test_jdk_import_is_skipped_and_not_name_searched(self): + text = HEADER + "import java.util.List;\nimport javax.inject.Named;\n" + listing = ["lib/List.java", "lib/Named.java"] + assert resolve(BILLING, text, ["List", "Named"], listing=listing) == [] + + def test_other_organization_is_third_party(self): + text = HEADER + "import org.acme.Widget;\nimport com.other.Gadget;\nimport com.acme.shared.Money;\n" + listing = ["lib/Widget.java", "lib/Gadget.java"] + found = resolve(BILLING, text, ["Widget", "Gadget", "Money"], listing=listing) + assert found == [("Money", "src/main/java/com/acme/shared/Money.java", "import")] + + def test_one_segment_package_is_its_own_organization(self): + text = "package acme;\nimport acme.shared.Money;\nimport other.Tax;\n" + assert resolve("src/acme/Invoice.java", text, ["Money", "Tax"]) == [ + ("Money", "src/acme/shared/Money.java", "import") + ] + + def test_unknown_root_skips_imports_but_not_conventions(self): + path = "src/main/java/com/acme/wrong/Invoice.java" + found = resolve(path, HEADER + "import com.acme.shared.Money;\n", ["Money", "Tax"]) + assert found == [("Tax", "src/main/java/com/acme/wrong/Tax.java", "path-convention")] + + def test_no_package_line_resolves_no_import(self): + found = resolve(BILLING, "import com.acme.shared.Money;\nimport com.acme.shared.*;\n", ["Money", "Tax"]) + assert found == [("Tax", "src/main/java/com/acme/billing/Tax.java", "path-convention")] + + def test_nested_type_maps_to_the_outer_file(self): + found = resolve(BILLING, HEADER + "import com.acme.shared.Money.Currency;\n", ["Currency"]) + assert found == [("Currency", "src/main/java/com/acme/shared/Money.java", "import")] + + def test_nested_wildcard_maps_to_the_outer_file(self): + found = resolve(BILLING, HEADER + "import com.acme.shared.Money.*;\n", ["Currency"]) + assert found == [ + ("Currency", "src/main/java/com/acme/shared/Money.java", "import"), + ("Currency", "src/main/java/com/acme/billing/Currency.java", "path-convention"), + ] + + def test_text_none_still_gives_conventions(self): + assert resolve(BILLING, None, ["Money"]) == [ + ("Money", "src/main/java/com/acme/billing/Money.java", "path-convention") + ] + + def test_package_at_the_repository_root(self): + found = resolve("com/acme/billing/Invoice.java", HEADER + "import com.acme.shared.Money;\n", ["Money"]) + assert found == [("Money", "com/acme/shared/Money.java", "import")] + + def test_top_level_file_convention(self): + assert resolve("Invoice.java", None, ["Money"]) == [("Money", "Money.java", "path-convention")] + + +class TestPython: + def test_absolute_import_tries_root_then_src(self): + found = resolve("app/service.py", "from acme.models import User\n", ["User"]) + assert found == [ + ("User", "acme/models.py", "import"), + ("User", "acme/models/__init__.py", "import"), + ("User", "src/acme/models.py", "import"), + ("User", "src/acme/models/__init__.py", "import"), + ] + + def test_relative_single_dot(self): + found = resolve("pkg/m.py", "from .x import C\n", ["C"]) + assert found == [("C", "pkg/x.py", "import"), ("C", "pkg/x/__init__.py", "import")] + + def test_relative_double_dot(self): + found = resolve("pkg/sub/m.py", "from ..y.z import C\n", ["C"]) + assert found == [("C", "pkg/y/z.py", "import"), ("C", "pkg/y/z/__init__.py", "import")] + + def test_relative_to_the_root_is_kept(self): + found = resolve("pkg/m.py", "from ..z import C\n", ["C"]) + assert found == [("C", "z.py", "import"), ("C", "z/__init__.py", "import")] + + def test_relative_beyond_the_root_is_dropped(self): + assert resolve("pkg/m.py", "from ...z import C\n", ["C"]) == [] + assert resolve("m.py", "from ..z import C\n", ["C"]) == [] + + def test_bare_relative_package(self): + assert resolve("pkg/sub/m.py", "from . import C\n", ["C"]) == [("C", "pkg/sub/__init__.py", "import")] + assert resolve("pkg/sub/m.py", "from .. import C\n", ["C"]) == [("C", "pkg/__init__.py", "import")] + + def test_plain_import_is_skipped(self): + assert resolve("app/m.py", "import acme.models\nimport acme.models as am\n", ["acme", "models", "am"]) == [] + + def test_alias_looks_up_the_imported_name(self): + found = resolve("app/m.py", "from acme.models import User as Account\n", ["Account"]) + assert [(name, reason) for name, _, reason in found] == [("User", "import")] * 4 + + def test_parenthesized_import_orders_by_names(self): + text = "from .models import (\n User, # the user\n Group,\n)\n" + found = resolve("app/m.py", text, ["Group", "User"]) + assert found == [ + ("Group", "app/models.py", "import"), + ("Group", "app/models/__init__.py", "import"), + ("User", "app/models.py", "import"), + ("User", "app/models/__init__.py", "import"), + ] + + def test_indented_and_continued_import(self): + text = "if TYPE_CHECKING:\n from .models import User, \\\n Group\n" + found = resolve("app/m.py", text, ["Group"]) + assert found == [("Group", "app/models.py", "import"), ("Group", "app/models/__init__.py", "import")] + + def test_star_import_probes_every_unbound_name(self): + text = "from .models import *\nfrom .other import Group\n" + found = resolve("app/m.py", text, ["User", "Group"]) + assert found == [ + ("User", "app/models.py", "import"), + ("User", "app/models/__init__.py", "import"), + ("Group", "app/other.py", "import"), + ("Group", "app/other/__init__.py", "import"), + ] + + def test_standard_library_is_skipped_and_not_name_searched(self): + text = "from __future__ import annotations\nfrom typing import Any\nfrom collections import OrderedDict\n" + listing = ["lib/any.py", "lib/ordered_dict.py", "lib/annotations.py"] + assert resolve("app/m.py", text, ["annotations", "Any", "OrderedDict"], listing=listing) == [] + + def test_pyi_is_python(self): + found = resolve("app/m.pyi", "from .x import C\n", ["C"]) + assert found == [("C", "app/x.py", "import"), ("C", "app/x/__init__.py", "import")] + + +APP = "web/src/app.ts" + + +class TestTypeScript: + def test_named_import(self): + found = resolve(APP, "import {Button} from './ui/button';\n", ["Button"]) + assert found == [("Button", p, "import") for p in probes("web/src/ui/button", TS_PROBES)] + + def test_default_import(self): + found = resolve(APP, 'import Card from "../card";\n', ["Card"]) + assert found == [("Card", p, "import") for p in probes("web/card", TS_PROBES)] + + def test_namespace_import_probes_referenced_members(self): + text = "import * as Api from './api';\nconst user = Api.fetchUser(1);\n" + found = resolve(APP, text, ["Api", "fetchUser", "user"]) + assert found == [("fetchUser", p, "import") for p in probes("web/src/api", TS_PROBES)] + + def test_export_from(self): + found = resolve(APP, "export {Modal} from './modal';\n", ["Modal"]) + assert found == [("Modal", p, "import") for p in probes("web/src/modal", TS_PROBES)] + + def test_export_default_as(self): + found = resolve(APP, "export { default as Panel } from './panel';\n", ["Panel"]) + assert found == [("Panel", p, "import") for p in probes("web/src/panel", TS_PROBES)] + + def test_import_alias_looks_up_the_exported_name(self): + found = resolve(APP, "import {Button as Btn} from './button';\n", ["Btn"]) + assert {name for name, _, _ in found} == {"Button"} + + def test_default_and_named_together(self): + text = "import React, { useThing } from './react-lite';\n" + found = resolve(APP, text, ["useThing", "React"]) + assert [name for name, _, _ in found] == ["useThing"] * 5 + ["React"] * 5 + + def test_multiline_type_only_import_with_comments(self): + text = "import type {\n User,\n Group, // the group\n} from './types';\n" + found = resolve(APP, text, ["Group"]) + assert found == [("Group", p, "import") for p in probes("web/src/types", TS_PROBES)] + + def test_candidate_order_follows_names_not_statements(self): + text = "import {B} from './b';\nimport {A} from './a';\n" + found = resolve(APP, text, ["A", "B"]) + assert [path for _, path, _ in found] == probes("web/src/a", TS_PROBES) + probes("web/src/b", TS_PROBES) + + def test_bare_package_is_skipped_and_not_name_searched(self): + text = "import {useState} from 'react';\nimport {Thing} from '@scope/pkg';\n" + listing = ["web/src/useState.ts", "web/src/Thing.ts"] + assert resolve(APP, text, ["useState", "Thing"], listing=listing) == [] + + def test_esm_js_specifier_tries_the_typescript_source(self): + found = resolve(APP, "import {slug} from './util.js';\n", ["slug"]) + assert found == [("slug", p, "import") for p in probes("web/src/util", (".ts", ".tsx", ".d.ts", ".js"))] + + def test_explicit_script_extension_is_literal(self): + found = resolve(APP, "import {slug} from './util.mjs';\n", ["slug"]) + assert found == [("slug", "web/src/util.mjs", "import")] + + def test_asset_import_gives_nothing(self): + text = "import styles from './app.module.css';\nimport logo from './logo.svg';\n" + assert resolve(APP, text, ["styles", "logo"]) == [] + + def test_directory_specifier_tries_index(self): + assert resolve(APP, "import {A} from '.';\n", ["A"]) == [("A", "web/src/index.ts", "import")] + assert resolve(APP, "import {A} from '../';\n", ["A"]) == [("A", "web/index.ts", "import")] + + def test_escaping_specifier_is_dropped(self): + assert resolve("app.ts", "import X from '../x';\n", ["X"]) == [] + + def test_plain_javascript_uses_the_same_rules(self): + found = resolve("lib/a.js", "import {B} from './b';\n", ["B"]) + assert found == [("B", p, "import") for p in probes("lib/b", TS_PROBES)] + + def test_side_effect_import_and_export_star_give_nothing(self): + text = "import './polyfill';\nexport * from './all';\n" + assert resolve(APP, text, ["polyfill", "all"]) == [] + + +class TestFilters: + def test_complete_listing_drops_absent_import_and_convention(self): + text = HEADER + "import com.acme.shared.Money;\n" + listing = ["src/main/java/com/acme/billing/Tax.java"] + found = resolve(BILLING, text, ["Money", "Tax", "Fee"], listing=listing, complete=True) + assert found == [("Tax", "src/main/java/com/acme/billing/Tax.java", "path-convention")] + + def test_incomplete_listing_keeps_them(self): + text = HEADER + "import com.acme.shared.Money;\n" + listing = ["src/main/java/com/acme/billing/Tax.java"] + found = resolve(BILLING, text, ["Money", "Tax", "Fee"], listing=listing, complete=False) + assert found == [ + ("Money", "src/main/java/com/acme/shared/Money.java", "import"), + ("Tax", "src/main/java/com/acme/billing/Tax.java", "path-convention"), + ("Fee", "src/main/java/com/acme/billing/Fee.java", "path-convention"), + ] + + def test_complete_listing_filters_python_probes(self): + listing = ["src/acme/models.py"] + found = resolve("app/m.py", "from acme.models import User\n", ["User"], listing=listing, complete=True) + assert found == [("User", "src/acme/models.py", "import")] + + def test_complete_flag_without_a_listing_filters_nothing(self): + found = resolve("app/m.py", "from acme.models import User\n", ["User"], listing=None, complete=True) + assert len(found) == 4 + + def test_listing_as_a_set(self): + listing = {"src/acme/models.py"} + found = resolve("app/m.py", "from acme.models import User\n", ["User"], listing=listing, complete=True) + assert found == [("User", "src/acme/models.py", "import")] + + def test_spring_import_in_the_fixture_gives_no_candidate(self): + found = resolve(SERVICE, fixture_text(SERVICE), list(SPRING), listing=LISTING, complete=False) + assert found == [] + + +class TestNameSearch: + def test_no_listing_no_search(self): + assert resolve("app/m.py", None, ["Widget"]) == [] + + def test_case_sensitive_beats_case_insensitive(self): + listing = ["a/widget.ts", "z/Widget.ts"] + assert resolve(APP, None, ["Widget"], listing=listing) == [ + ("Widget", "z/Widget.ts", "name-search"), + ("Widget", "a/widget.ts", "name-search"), + ] + + def test_python_snake_case_form(self): + listing = ["lib/transport_config.py", "lib/TransportConfig.ts"] + assert resolve("app/m.py", None, ["TransportConfig"], listing=listing) == [ + ("TransportConfig", "lib/transport_config.py", "name-search") + ] + + def test_snake_case_is_python_only(self): + assert resolve(APP, None, ["TransportConfig"], listing=["lib/transport_config.ts"]) == [] + + def test_snake_case_ranks_after_case_insensitive(self): + listing = ["a/transport_config.py", "z/transportconfig.py"] + assert [p for _, p, _ in resolve("app/m.py", None, ["TransportConfig"], listing=listing)] == [ + "z/transportconfig.py", + "a/transport_config.py", + ] + + def test_cap_of_three_per_name(self): + listing = [f"m{i}/Widget.py" for i in range(5)] + found = resolve("app/m.py", None, ["Widget"], listing=listing) + assert [p for _, p, _ in found] == ["m0/Widget.py", "m1/Widget.py", "m2/Widget.py"] + + def test_deepest_shared_prefix_first_then_path(self): + listing = ["z/Foo.py", "a/Foo.py", "a/b/x/Foo.py", "a/b/Foo.py"] + found = resolve("a/b/c/m.py", None, ["Foo"], listing=listing) + assert [p for _, p, _ in found] == ["a/b/Foo.py", "a/b/x/Foo.py", "a/Foo.py"] + + def test_shared_prefix_counts_whole_directories(self): + listing = ["ab/Foo.py", "a/Foo.py"] + found = resolve("a/m.py", None, ["Foo"], listing=listing) + assert [p for _, p, _ in found] == ["a/Foo.py", "ab/Foo.py"] + + def test_other_language_excluded(self): + listing = ["web/Widget.ts", "lib/Widget.py", "src/Widget.java", "docs/Widget.md"] + assert resolve("app/m.py", None, ["Widget"], listing=listing) == [("Widget", "lib/Widget.py", "name-search")] + + def test_declaration_file_stem(self): + assert resolve(APP, None, ["Widget"], listing=["types/Widget.d.ts"]) == [ + ("Widget", "types/Widget.d.ts", "name-search") + ] + + def test_dunder_names_are_not_searched(self): + assert resolve("app/m.py", None, ["__init__"], listing=["pkg/__init__.py"]) == [] + + def test_language_without_definition_regexes_resolves_nothing(self): + listing = ["x/idempotency_keys.sql", "cmd/Server.go"] + assert resolve("db/001.sql", None, ["idempotency_keys"], listing=listing) == [] + assert resolve("cmd/main.go", None, ["Server"], listing=listing) == [] + + def test_name_search_follows_imports_and_conventions(self): + text = HEADER + "import com.acme.shared.Money;\n" + listing = ["lib/Money.java", "lib/Tax.java"] + assert resolve(BILLING, text, ["Money", "Tax"], listing=listing) == [ + ("Money", "src/main/java/com/acme/shared/Money.java", "import"), + ("Tax", "src/main/java/com/acme/billing/Tax.java", "path-convention"), + ("Money", "lib/Money.java", "name-search"), + ("Tax", "lib/Tax.java", "name-search"), + ] + + +class TestDedupAndPaths: + def test_duplicate_keeps_the_first_reason(self): + text = HEADER + "import com.acme.billing.*;\n" + listing = ["src/main/java/com/acme/billing/Money.java"] + assert resolve(BILLING, text, ["Money"], listing=listing) == [ + ("Money", "src/main/java/com/acme/billing/Money.java", "import") + ] + + def test_convention_and_name_search_on_one_path_keep_the_convention(self): + found = resolve(SERVICE, None, ["TransportConfig"], listing=LISTING) + assert found == [("TransportConfig", CONFIG, "path-convention")] + + def test_referencing_path_is_never_a_candidate(self): + path = "src/main/java/com/acme/billing/Money.java" + text = HEADER + "import com.acme.billing.*;\n" + assert resolve(path, text, ["Money"], listing=[path], complete=True) == [] + + def test_duplicate_names_resolve_once(self): + assert resolve(BILLING, None, ["Money", "Money"]) == [ + ("Money", "src/main/java/com/acme/billing/Money.java", "path-convention") + ] + + def test_paths_are_normalized(self): + found = resolve("./web//src/app.ts", "import {A} from './x/../y';\n", ["A"]) + assert [p for _, p, _ in found] == probes("web/src/y", TS_PROBES) + for _, path, _ in found: + assert not path.startswith(("/", "./")) and "/./" not in path and ".." not in path.split("/") + + def test_leading_slash_is_stripped_and_self_still_excluded(self): + listing = ["lib/Widget.py"] + assert resolve("/lib/Widget.py", None, ["Widget"], listing=listing) == [] + + def test_unresolvable_names_are_absent(self): + assert resolve("app/m.py", "from .x import C\n", ["Nope"], listing=[]) == [] From 4cfa8e3fd1fbb0fc79c276aa45b6e6c95fbe90b5 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:22:24 -0700 Subject: [PATCH 12/24] test: exercise the (path, line) dedup in the cross-chunk nested-type test MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The nested-type test named dedup but never produced two entries on one (path, line), so removing the dedup loop left the suite green. Making the nested record's declaration an added line puts Outer's change run and Inner's definition on the same line; the test now pins that the definition entry wins that tie, and a no-dedup mutation goes red. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- tests/test_repo_context_crosschunk.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_repo_context_crosschunk.py b/tests/test_repo_context_crosschunk.py index 63fbf9f..fdb7459 100644 --- a/tests/test_repo_context_crosschunk.py +++ b/tests/test_repo_context_crosschunk.py @@ -247,7 +247,7 @@ def test_a_nested_type_does_not_end_the_window_and_dedup_keeps_the_first(self): outer = _file( "p/Outer.java", (1, [ - " package p;", " ", " public class Outer {", " public record Inner(int a) {", + " package p;", " ", " public class Outer {", "+ public record Inner(int a) {", " }", " void go() {", " run();", "+ check();", "+ log();", " }", " }", ]), @@ -258,6 +258,7 @@ def test_a_nested_type_does_not_end_the_window_and_dedup_keeps_the_first(self): (4, "Inner", "cross-chunk"), (8, "Outer", "cross-chunk"), ] + assert entries[1].text == " public record Inner(int a) {\n }" assert entries[2].text == " check();\n log();" def test_python_windows_follow_indentation(self): From cd09953739a94a5b1ba901ca3fd590eb063209b5 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:37:24 -0700 Subject: [PATCH 13/24] feat: add repo_dir eval case field (issue #17, T15a) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Repository context (PRXREF_REPO_CONTEXT=repo) needs a zero-network way to run against a fixture or a checked-out PR head, so eval cases gain an optional repo_dir field: a directory holding the repository at the PR head, read and listed by repository context in place of a forge. - EvalCase gains repo_dir: str | None = None as its LAST field, so no positional construction shifts. - _CASE_KEYS gains "repo_dir" directly after "context_file", which fixes the allowed: list in the unknown-field message and the key order case_to_json writes; the cases.json form joins a relative repo_dir onto the dataset directory as context_file is joined, and requires an existing directory; the case-*/ directory form picks up a repo/ subdirectory as ticket.md becomes context_file. - case_to_json/case_from_json_record round-trip repo_dir, both set and None; a 0.15-shaped case.json record with no repo_dir key still loads. - The issue #17 acceptance fixture gains "repo_dir": "repo" and L1's must_match loosens from re:(?i)connector_id to re:(?i)connector, so it also matches a correct finding written as "connectorId" or "the connector" (the fixture-seat correction carried into this wave). - docs/evals.md documents the field for both dataset forms. This seat builds only the eval-case half. The --repo-dir CLI flag, the threading into evals._run_case and the orchestrator, the run-record fields and the README CLI Flags entry belong to seat T15b, which reads case.repo_dir once it lands. 🤖 Authored with Claude Code — Claude Sonnet 5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- docs/evals.md | 17 ++- src/prxref/eval_cases.py | 66 ++++---- tests/fixtures/issue17/cases.json | 3 +- tests/test_eval_run.py | 3 +- tests/test_issue_17_eval_repo_dir.py | 217 +++++++++++++++++++++++++++ 5 files changed, 272 insertions(+), 34 deletions(-) create mode 100644 tests/test_issue_17_eval_repo_dir.py diff --git a/docs/evals.md b/docs/evals.md index 505751b..b9c6d79 100644 --- a/docs/evals.md +++ b/docs/evals.md @@ -86,6 +86,7 @@ no case runs. "id": "local-1", "diff_file": "diffs/local-1.diff", "context_file": "tickets/local-1.md", + "repo_dir": "repos/local-1", "spec": ["docs/specs"], "expected": [] } @@ -105,15 +106,17 @@ is refused, and a field set to `null` counts as absent. | `diff_file` | one of these two | A unified diff (`git diff` or `git format-patch` output) holding at least one file diff. | | `base_sha`, `head_sha` | with `pr_url` | The pinned range, as the replay flags take it: a pair, full 40- or 64-character hex, two different commits. Stored lowercased. | | `context_file` | no | The ticket the PR implements, given to the review as `--context-file`. | +| `repo_dir` | no | A directory holding the repository at the PR head. Repository context (`PRXREF_REPO_CONTEXT=repo`) reads and lists files there instead of calling a forge; the field has no effect when repository context is off. | | `spec` | no | Spec sources, as `--spec` takes them: one string or an array of strings, each a local path or an `http(s)` URL. | A case replays either a local diff (`diff_file`), or a pull request pinned to a range (`pr_url` with both SHAs). It may also give `pr_url` beside `diff_file`, with or without the SHAs, as `prxref review` allows. SHAs need `pr_url`, and `pr_url` without `diff_file` needs both SHAs. A relative -`diff_file`, `context_file` or local `spec` path is read relative to the -directory holding `cases.json`. `context_file` and every local `spec` path -must exist, and a `spec` URL is kept as given. +`diff_file`, `context_file`, `repo_dir` or local `spec` path is read +relative to the directory holding `cases.json`. `context_file` and every +local `spec` path must exist, `repo_dir` must be an existing directory, and +a `spec` URL is kept as given. Each entry of `expected` is one label, a finding a human reviewer left: @@ -144,6 +147,7 @@ case, read in name order, and its id is the directory name. | `diff.patch` | yes | `diff_file` | | `expected.json` | yes | `expected`: a JSON array of labels | | `ticket.md` | no | `context_file` | +| `repo/` | no | `repo_dir` | | `docs/` | no | the case's one `spec` source | `meta.json` is not read. `expected.json` spells two label fields @@ -165,7 +169,8 @@ the path. case does not have. - Neither `pr_url` nor `diff_file`; a `pr_url` no forge recognises; SHAs that break the rules above; a `diff_file` that cannot be read or holds no - file diff; a `context_file` or local `spec` path that does not exist. + file diff; a `context_file` or local `spec` path that does not exist; a + `repo_dir` that is not an existing directory. - `expected` missing or not an array. - A label with a required field missing, a `line` below `1`, a `severity` outside the five, a `category`, `text` or `must_match` that is not a @@ -291,8 +296,8 @@ prxref-eval/ --out ``` - **`cases//case.json`** holds the case as the harness read it: `id`, - `pr_url`, `base_sha`, `head_sha`, `diff_file`, `context_file`, `spec` (a - list) and `expected`, each label with all eight fields, in the + `pr_url`, `base_sha`, `head_sha`, `diff_file`, `context_file`, `repo_dir`, + `spec` (a list) and `expected`, each label with all eight fields, in the `cases.json` spelling. `eval score` grades against this copy, so it needs no `--cases` and never opens the files the case names. - **`cases//record.json`** is the review's record, exactly what diff --git a/src/prxref/eval_cases.py b/src/prxref/eval_cases.py index f68aa60..31d5104 100644 --- a/src/prxref/eval_cases.py +++ b/src/prxref/eval_cases.py @@ -7,21 +7,21 @@ A ``cases.json`` file is ``{"version": 1, "cases": [...]}``. Each case is an object with a required ``id`` and ``expected`` and the optional ``pr_url``, -``base_sha``, ``head_sha``, ``diff_file``, ``context_file`` and ``spec`` (one -string or a list of them). ``expected`` is a list, possibly empty, of human -findings, each with ``id``, ``file``, ``line`` and ``severity`` and the -optional ``category``, ``accepted``, ``text`` and ``must_match``. An optional -field set to ``null`` counts as absent. A relative ``diff_file``, -``context_file`` or ``spec`` path is read relative to the directory holding -the ``cases.json`` file; a ``spec`` entry that is an ``http(s)`` URL is kept -as given. +``base_sha``, ``head_sha``, ``diff_file``, ``context_file``, ``repo_dir`` and +``spec`` (one string or a list of them). ``expected`` is a list, possibly +empty, of human findings, each with ``id``, ``file``, ``line`` and +``severity`` and the optional ``category``, ``accepted``, ``text`` and +``must_match``. An optional field set to ``null`` counts as absent. A +relative ``diff_file``, ``context_file``, ``repo_dir`` or ``spec`` path is +read relative to the directory holding the ``cases.json`` file; a ``spec`` +entry that is an ``http(s)`` URL is kept as given. A ``case-*/`` directory supplies ``diff.patch`` (required) as its ``diff_file``, ``expected.json`` (required: the ``expected`` list above) and, -when present, ``ticket.md`` as its ``context_file`` and ``docs/`` as its one -``spec`` source. Its id is the directory name. ``expected.json`` spells -``line`` as ``line_hint`` and ``category`` as ``source``, and every message -about it uses those spellings. +when present, ``ticket.md`` as its ``context_file``, ``repo/`` as its +``repo_dir`` and ``docs/`` as its one ``spec`` source. Its id is the +directory name. ``expected.json`` spells ``line`` as ``line_hint`` and +``category`` as ``source``, and every message about it uses those spellings. A case replays either a local diff (``diff_file``) or a pinned commit range of a pull request (``pr_url`` with both ``base_sha`` and ``head_sha``), and may @@ -72,7 +72,8 @@ _FULL_SHA_RE = re.compile(r"[0-9a-fA-F]{40}(?:[0-9a-fA-F]{24})?") _TOP_KEYS = ("version", "cases") _CASE_KEYS = ( - "id", "pr_url", "base_sha", "head_sha", "diff_file", "context_file", "spec", "expected", + "id", "pr_url", "base_sha", "head_sha", "diff_file", "context_file", "repo_dir", "spec", + "expected", ) _EXPECTED_KEYS = ( "id", "file", "line", "severity", "category", "accepted", "text", "must_match", @@ -109,12 +110,15 @@ class EvalCase: """One validated eval case: what to replay and the findings it should yield. ``id`` is a single safe path segment, unique within its dataset. - ``diff_file`` and ``context_file`` are paths ready to open, already - joined onto the ``cases.json`` directory or the case directory; ``spec`` - holds the spec sources in order, each such a path or an ``http(s)`` URL, - and is empty when the case sets none. ``base_sha`` and ``head_sha`` are - both full lowercased SHAs, and ``pr_url`` is set, or both are ``None``. - ``expected`` keeps the labels in file order and may be empty. + ``diff_file``, ``context_file`` and ``repo_dir`` are paths ready to open, + already joined onto the ``cases.json`` directory or the case directory; + ``repo_dir`` is a directory holding the repository at the PR head, read + by repository context (``PRXREF_REPO_CONTEXT=repo``) in place of a forge. + ``spec`` holds the spec sources in order, each such a path or an + ``http(s)`` URL, and is empty when the case sets none. ``base_sha`` and + ``head_sha`` are both full lowercased SHAs, and ``pr_url`` is set, or + both are ``None``. ``expected`` keeps the labels in file order and may + be empty. """ id: str @@ -125,6 +129,7 @@ class EvalCase: diff_file: str | None = None context_file: str | None = None spec: tuple[str, ...] = () + repo_dir: str | None = None def load_cases(path: str | os.PathLike[str], *, source: str = "--cases") -> list[EvalCase]: @@ -171,11 +176,12 @@ def case_to_json(case: EvalCase) -> dict[str, Any]: The keys follow the ``cases.json`` spelling in a fixed order: ``id``, ``pr_url``, ``base_sha``, ``head_sha``, ``diff_file``, ``context_file``, - ``spec`` (a list, empty when the case sets none) and ``expected``, a list - of labels each keyed ``id``, ``file``, ``line``, ``severity``, - ``category``, ``accepted``, ``text`` and ``must_match``. Every field is - written, ``null`` when unset. Paths are written exactly as the case holds - them, already joined onto the dataset directory when it was loaded. + ``repo_dir``, ``spec`` (a list, empty when the case sets none) and + ``expected``, a list of labels each keyed ``id``, ``file``, ``line``, + ``severity``, ``category``, ``accepted``, ``text`` and ``must_match``. + Every field is written, ``null`` when unset. Paths are written exactly + as the case holds them, already joined onto the dataset directory when + it was loaded. """ return { "id": case.id, @@ -184,6 +190,7 @@ def case_to_json(case: EvalCase) -> dict[str, Any]: "head_sha": case.head_sha, "diff_file": case.diff_file, "context_file": case.context_file, + "repo_dir": case.repo_dir, "spec": list(case.spec), "expected": [ {key: getattr(finding, key) for key in _EXPECTED_KEYS} for finding in case.expected @@ -219,7 +226,7 @@ def case_from_json_record(obj: Any, *, source: str = "case.json") -> EvalCase: ) texts = { key: _optional_text(obj, key, where, source) - for key in ("pr_url", "base_sha", "head_sha", "diff_file", "context_file") + for key in ("pr_url", "base_sha", "head_sha", "diff_file", "context_file", "repo_dir") } spec = obj.get("spec") if spec is None: @@ -292,6 +299,7 @@ def _case_from_json(entry: Any, index: int, base: Path, source: str) -> EvalCase head_sha = _optional_text(entry, "head_sha", where, source) diff_file = _optional_text(entry, "diff_file", where, source) context_file = _optional_text(entry, "context_file", where, source) + repo_dir = _optional_text(entry, "repo_dir", where, source) base_sha, head_sha = _check_replay(pr_url, base_sha, head_sha, diff_file, where, source) if pr_url is not None and detect_forge(pr_url) is None: raise ConfigError( @@ -302,6 +310,10 @@ def _case_from_json(entry: Any, index: int, base: Path, source: str) -> EvalCase context_file = _join(base, context_file) if not Path(context_file).is_file(): raise ConfigError(f"{source}: {where}: context_file: no such file {context_file!r}") + if repo_dir is not None: + repo_dir = _join(base, repo_dir) + if not Path(repo_dir).is_dir(): + raise ConfigError(f"{source}: {where}: repo_dir: no such directory {repo_dir!r}") spec = _spec_sources(entry.get("spec"), base, where, source) if "expected" not in entry: raise ConfigError(f"{source}: {where}: expected: required (an array of findings, possibly empty)") @@ -312,7 +324,7 @@ def _case_from_json(entry: Any, index: int, base: Path, source: str) -> EvalCase _check_anchors(expected, files, where, "expected", _JSON_SPELLING, source) return EvalCase( id=case_id, expected=expected, pr_url=pr_url, base_sha=base_sha, head_sha=head_sha, - diff_file=diff_file, context_file=context_file, spec=spec, + diff_file=diff_file, context_file=context_file, repo_dir=repo_dir, spec=spec, ) @@ -413,11 +425,13 @@ def _case_from_directory(case_dir: Path, source: str) -> EvalCase: _check_anchors(expected, files, where, "expected.json", _DIRECTORY_SPELLING, source) ticket = case_dir / "ticket.md" docs = case_dir / "docs" + repo = case_dir / "repo" return EvalCase( id=case_id, expected=expected, diff_file=diff_file, context_file=str(ticket) if ticket.is_file() else None, + repo_dir=str(repo) if repo.is_dir() else None, spec=(str(docs),) if docs.is_dir() else (), ) diff --git a/tests/fixtures/issue17/cases.json b/tests/fixtures/issue17/cases.json index 91ca9aa..4b08b00 100644 --- a/tests/fixtures/issue17/cases.json +++ b/tests/fixtures/issue17/cases.json @@ -4,6 +4,7 @@ { "id": "issue17-repo-context", "diff_file": "pr.diff", + "repo_dir": "repo", "expected": [ { "id": "L1", @@ -11,7 +12,7 @@ "line": 4, "severity": "error", "text": "The unique index on idempotency_keys omits connector_id, so a key is only unique per (tenant, key) even though the OpenAPI IdempotencyKey contract (api/openapi/connectors.yaml) requires uniqueness per (tenant, connector, key).", - "must_match": "re:(?i)connector_id" + "must_match": "re:(?i)connector" }, { "id": "L2", diff --git a/tests/test_eval_run.py b/tests/test_eval_run.py index 6dc3bb6..dad9eb8 100644 --- a/tests/test_eval_run.py +++ b/tests/test_eval_run.py @@ -628,7 +628,8 @@ def test_a_minimal_case_round_trips(self): def test_the_keys_come_in_the_cases_json_spelling_and_a_fixed_order(self): record = case_to_json(_full_case()) assert list(record) == [ - "id", "pr_url", "base_sha", "head_sha", "diff_file", "context_file", "spec", "expected", + "id", "pr_url", "base_sha", "head_sha", "diff_file", "context_file", "repo_dir", "spec", + "expected", ] assert list(record["expected"][0]) == [ "id", "file", "line", "severity", "category", "accepted", "text", "must_match", diff --git a/tests/test_issue_17_eval_repo_dir.py b/tests/test_issue_17_eval_repo_dir.py new file mode 100644 index 0000000..aee3667 --- /dev/null +++ b/tests/test_issue_17_eval_repo_dir.py @@ -0,0 +1,217 @@ +"""The eval-case half of ``--repo-dir`` (issue #17, T15a, OQ4). + +``repo_dir`` is an optional :class:`~prxref.eval_cases.EvalCase` field: a +directory holding the repository at the PR head. Repository context +(``PRXREF_REPO_CONTEXT=repo``) reads and lists files there instead of +calling a forge. This module only loads, validates, joins and round-trips +the field; the CLI flag, the threading into ``evals._run_case`` and the +orchestrator are T15b's job, and neither is exercised here. +""" +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +from prxref import cli +from prxref.eval_cases import ( + EvalCase, + ExpectedFinding, + case_from_json_record, + case_to_json, + load_cases, +) +from prxref.llm import ConfigError + +DIFF = ( + "diff --git a/src/app.py b/src/app.py\n" + "--- a/src/app.py\n" + "+++ b/src/app.py\n" + "@@ -1,2 +1,3 @@\n" + " import os\n" + "+import sys\n" + " print(os.name)\n" +) +DROP = object() +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" + + +def _finding(**over) -> dict: + entry = {"id": "H1", "file": "src/app.py", "line": 2, "severity": "error"} + entry.update(over) + return {key: value for key, value in entry.items() if value is not DROP} + + +def _dir_finding(**over) -> dict: + entry = {"id": "H1", "file": "src/app.py", "line_hint": 2, "severity": "error"} + entry.update(over) + return {key: value for key, value in entry.items() if value is not DROP} + + +def _case(**over) -> dict: + entry = {"id": "case-a", "diff_file": "change.diff", "expected": [_finding()]} + entry.update(over) + return {key: value for key, value in entry.items() if value is not DROP} + + +def _dataset(tmp_path: Path, cases) -> Path: + (tmp_path / "change.diff").write_text(DIFF, encoding="utf-8") + path = tmp_path / "cases.json" + path.write_text(json.dumps({"version": 1, "cases": cases}), encoding="utf-8") + return path + + +def _refusal(path) -> str: + with pytest.raises(ConfigError) as exc: + load_cases(path) + return str(exc.value) + + +class TestCasesJsonRepoDir: + def test_a_relative_repo_dir_is_joined_onto_the_dataset_directory(self, tmp_path): + (tmp_path / "repo").mkdir() + path = _dataset(tmp_path, [_case(repo_dir="repo")]) + + (case,) = load_cases(path) + + assert case.repo_dir == str(tmp_path / "repo") + + def test_an_absolute_repo_dir_is_kept_as_given(self, tmp_path): + repo = tmp_path / "elsewhere" / "repo" + repo.mkdir(parents=True) + path = _dataset(tmp_path, [_case(repo_dir=str(repo))]) + + (case,) = load_cases(path) + + assert case.repo_dir == str(repo) + + def test_a_missing_repo_dir_is_refused_naming_the_case_and_field(self, tmp_path): + message = _refusal(_dataset(tmp_path, [_case(repo_dir="missing")])) + + assert message.startswith("--cases: case 'case-a': repo_dir: no such directory "), message + + def test_a_repo_dir_that_is_a_file_is_refused(self, tmp_path): + (tmp_path / "repo-file").write_text("not a directory\n", encoding="utf-8") + + message = _refusal(_dataset(tmp_path, [_case(repo_dir="repo-file")])) + + assert message.startswith("--cases: case 'case-a': repo_dir: no such directory "), message + + def test_a_non_string_repo_dir_is_rejected_by_optional_text(self, tmp_path): + message = _refusal(_dataset(tmp_path, [_case(repo_dir=3)])) + + assert message == ( + "--cases: case 'case-a': repo_dir: must be a non-empty string or null, " + "got the number 3" + ) + + def test_a_null_repo_dir_counts_as_absent(self, tmp_path): + path = _dataset(tmp_path, [_case(repo_dir=None)]) + + (case,) = load_cases(path) + + assert case.repo_dir is None + + def test_repo_dir_may_sit_beside_context_file_and_diff_file(self, tmp_path): + (tmp_path / "repo").mkdir() + (tmp_path / "ticket.md").write_text("placeholder\n", encoding="utf-8") + path = _dataset(tmp_path, [_case(repo_dir="repo", context_file="ticket.md")]) + + (case,) = load_cases(path) + + assert case.repo_dir == str(tmp_path / "repo") + assert case.context_file == str(tmp_path / "ticket.md") + assert case.diff_file == str(tmp_path / "change.diff") + + def test_the_unknown_field_message_lists_repo_dir_in_its_position(self, tmp_path): + message = _refusal(_dataset(tmp_path, [_case(bogus_field=1)])) + + assert ( + "allowed: id, pr_url, base_sha, head_sha, diff_file, context_file, repo_dir, " + "spec, expected" in message + ) + + +class TestDirectoryFormRepoDir: + def test_a_repo_subdirectory_becomes_repo_dir(self, tmp_path): + case_dir = tmp_path / "case-a" + case_dir.mkdir() + (case_dir / "diff.patch").write_text(DIFF, encoding="utf-8") + (case_dir / "expected.json").write_text(json.dumps([_dir_finding()]), encoding="utf-8") + (case_dir / "repo").mkdir() + + (case,) = load_cases(tmp_path) + + assert case.repo_dir == str(case_dir / "repo") + + def test_a_case_directory_without_repo_gives_none(self, tmp_path): + case_dir = tmp_path / "case-a" + case_dir.mkdir() + (case_dir / "diff.patch").write_text(DIFF, encoding="utf-8") + (case_dir / "expected.json").write_text(json.dumps([_dir_finding()]), encoding="utf-8") + + (case,) = load_cases(tmp_path) + + assert case.repo_dir is None + + +class TestRepoDirRoundTrip: + def _case_with_repo_dir(self) -> EvalCase: + return EvalCase( + id="case-a", + expected=(ExpectedFinding(id="H1", file="src/app.py", line=2, severity="error"),), + diff_file="/data/cases/change.diff", + context_file="/data/cases/ticket.md", + spec=("/data/cases/docs",), + repo_dir="/data/cases/repo", + ) + + def test_repo_dir_round_trips_when_set(self): + case = self._case_with_repo_dir() + + assert case_from_json_record(case_to_json(case)) == case + + def test_repo_dir_round_trips_when_unset(self): + case = EvalCase(id="case-a", expected=(), diff_file="/data/cases/change.diff") + assert case.repo_dir is None + + assert case_from_json_record(case_to_json(case)) == case + + def test_a_0_15_shaped_record_with_no_repo_dir_key_still_loads(self): + record = { + "id": "case-a", + "pr_url": None, + "base_sha": None, + "head_sha": None, + "diff_file": "/data/cases/change.diff", + "context_file": None, + "spec": [], + "expected": [], + } + + case = case_from_json_record(record) + + assert case.repo_dir is None + assert case.diff_file == "/data/cases/change.diff" + + +class TestFixtureRepoDir: + def test_the_issue17_fixture_case_gets_a_repo_dir_pointing_at_the_committed_repo(self): + (case,) = load_cases(str(FIXTURE / "cases.json")) + + assert case.repo_dir is not None + assert case.repo_dir.endswith(str(Path("tests") / "fixtures" / "issue17" / "repo")) + assert Path(case.repo_dir).is_dir() + + +class TestCliExit2ThroughRealEntryPoint: + def test_a_bad_repo_dir_exits_2_naming_the_field_without_monkeypatching(self, tmp_path, capsys): + cases = _dataset(tmp_path, [_case(repo_dir="missing")]) + + code = cli.main(["eval", "run", "--cases", str(cases), "--label", "L"]) + + assert code == 2 + err = capsys.readouterr().err + assert "configuration error:" in err + assert "repo_dir" in err From f7919b58db4d3681ef6cefc58b27276871809127 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:46:38 -0700 Subject: [PATCH 14/24] feat: contract triggers, contract-file selection and per-chunk contract entries (#17) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds the T6 surface to src/prxref/repo_contracts.py for issue #17's misses (a) and (c): a migration whose unique index disagrees with an OpenAPI contract outside the diff, and a route whose operation and schemas the worker never sees. - ContractTriggers and contract_triggers(): route literals (Spring, JAX-RS, FastAPI/Flask, Express; class-level @RequestMapping/@Path, Blueprint url_prefix and APIRouter prefix joined when the head text is given), tables (T4's _SQL_HEADS and _bare_table, REFERENCES targets, Liquibase tableName/baseTableName/referencedTableName in XML, YAML and JSON), referenced names, and operation ids. Every scan is linear on a 512 KiB line; the INDEX ... ON head is bounded at the next CREATE/ALTER. - literal_contract_paths() and select_contract_files(): the run's contract files from the globs, the listing and the diff paths. - earlier_migrations(): up to MAX_EARLIER_MIGRATIONS = 4 same-directory predecessors under a natural sort. - contract_excerpts(): extension dispatch to T4's excerpters, plus whole-file fragments matched by stem. - contract_entries(): the per-chunk entry point for T7, reading at most MAX_SPEC_FILES = 6 ranked spec files plus earlier migrations, through the reader it is given, and returning ContextEntry(kind="contract", reason="contract") ordered and deduplicated by (path, line). Tests: tests/test_repo_context_triggers.py (82), including the issue17 fixture end to end. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/repo_contracts.py | 505 +++++++++++++++++++- tests/test_repo_context_triggers.py | 708 ++++++++++++++++++++++++++++ 2 files changed, 1206 insertions(+), 7 deletions(-) create mode 100644 tests/test_repo_context_triggers.py diff --git a/src/prxref/repo_contracts.py b/src/prxref/repo_contracts.py index 8aa46cc..d14e24a 100644 --- a/src/prxref/repo_contracts.py +++ b/src/prxref/repo_contracts.py @@ -7,28 +7,49 @@ a table. The excerpters here cut the matching slice out of such a file so a worker can check the change against it. -The module is pure and stdlib only. It performs no I/O: callers pass the file's -text. There is deliberately no YAML parser, because the core ships none, so -YAML is sliced by indentation, JSON goes through :mod:`json`, and SQL and XML -are scanned with regular expressions. Every excerpter degrades to ``[]`` on -text it cannot read rather than raising, because a review must never fail over -missing context. Each excerpt is capped at :data:`MAX_CONTRACT_LINES` lines and +The module is pure. It is stdlib plus the pure :mod:`prxref.repo_context`, +:mod:`prxref.chunk_context` and :func:`prxref.rules.match_globs`, and it +performs no I/O: callers pass the file's text, and :func:`contract_entries` +reads only through the ``read`` callable it is handed. There is deliberately +no YAML parser, because the core ships none, so YAML is sliced by indentation, +JSON goes through :mod:`json`, and SQL and XML are scanned with regular +expressions. Every excerpter degrades to ``[]`` on text it cannot read rather +than raising, because a review must never fail over missing context. Each +excerpt is capped at :data:`MAX_CONTRACT_LINES` lines and :data:`MAX_CONTRACT_CHARS` characters. Names compare through :func:`normalize_name`, so a table called ``idempotency_keys`` matches a schema or class called ``IdempotencyKey``. Routes compare through :func:`route_key`, so the code route ``/connectors/:id`` matches the spec path ``/connectors/{connectorId}``. + +The second half of the module turns a worker chunk into contract entries. +:func:`contract_triggers` reads a changed file's added lines for the routes, +tables, names and operation ids a contract can match. +:func:`select_contract_files` picks the repository's contract files from the +configured globs once per run, and :func:`literal_contract_paths` names the +globs that are plain paths. :func:`earlier_migrations` finds the migrations a +changed one builds on, :func:`contract_excerpts` sends one contract file to +the excerpter its format needs, and :func:`contract_entries` ties them +together for one chunk as :class:`prxref.repo_context.ContextEntry` records. +Every regular expression here runs in linear time on a 512 KiB line. """ from __future__ import annotations import json +import posixpath import re -from collections.abc import Callable, Iterable, Iterator +from collections.abc import Callable, Collection, Iterable, Iterator, Sequence from dataclasses import dataclass +from . import chunk_context +from .repo_context import _JAVA_KEYWORDS, ContextEntry, language_of, referenced_names +from .rules import match_globs + MAX_CONTRACT_LINES = 40 MAX_CONTRACT_CHARS = 2000 +MAX_EARLIER_MIGRATIONS = 4 +MAX_SPEC_FILES = 6 _BOM = chr(0xFEFF) _HTTP_METHODS = frozenset({"get", "put", "post", "delete", "options", "head", "patch", "trace"}) @@ -679,3 +700,473 @@ def _liquibase_json(text: str, keys: frozenset[str]) -> list[Excerpt]: hits.append(_Hit(pointer, value, "changeSet", node, line)) break return _json_excerpts(hits) + + +_SPEC_SUFFIXES = (".yaml", ".yml", ".json") +_MIGRATION_SUFFIXES = (".sql", ".xml", ".yaml", ".yml", ".json") +_PREFIX_LOOKBACK_LINES = 40 + +_OPEN_CALL = re.compile(r"(?class|interface)[ \t]+[A-Za-z_$]" +) +_FLASK_PREFIX = re.compile( + r"""(? ContractTriggers: + """The contract triggers on one changed file's added lines. + + ``added`` holds the file's ``+`` lines; they are scanned joined by + newlines, so a statement or annotation that wraps across added lines + still counts. ``text`` is the file's full content at the head, or None. + + - ``routes``: the path literals of route declarations. Spring + ``@GetMapping``, ``@PostMapping``, ``@PutMapping``, ``@DeleteMapping``, + ``@PatchMapping`` and ``@RequestMapping`` give their positional string, + their ``value =`` or ``path =`` string, or every string of an array + form ``{"/a", "/b"}``, never a ``produces``, ``consumes``, ``name``, + ``headers`` or ``params`` string. JAX-RS ``@Path("...")`` counts, as do + FastAPI and Flask ``@.get|post|put|delete|patch|route|api_route("...")`` + and Express ``.get|post|put|delete|patch|all|use('/...')`` when + the literal starts with ``/``. Empty literals are dropped. When + ``text`` is given and a route was found, class-level prefixes are read + from it: a Spring ``@RequestMapping`` or JAX-RS ``@Path`` in the + annotations directly above a class or interface declaration, a Flask + ``Blueprint(..., url_prefix="...")`` and a FastAPI + ``APIRouter(prefix="...")``, at most the first of each kind in the + file. After the bare routes come, for each prefix, the prefix joined to + every route with exactly one ``/`` between them, because Spring and + Flask accept a segment written without its leading slash. + - ``tables``: every ``CREATE TABLE``, ``ALTER TABLE`` and + ``CREATE [UNIQUE] INDEX ... ON
`` head this module's + :func:`sql_excerpts` recognizes (the table, never the index name), each + ``REFERENCES
`` foreign-key target, and every Liquibase + ``tableName``, ``baseTableName`` and ``referencedTableName`` value, in + XML (``tableName="t"``), YAML (``tableName: t``) or JSON + (``"tableName": "t"``). Each is reduced to its bare name. This applies + to every language, because SQL also shows up in Python and Java + migrations. An index head is read only up to the next ``CREATE`` or + ``ALTER`` keyword, which keeps the scan linear. + - ``names``: :func:`prxref.repo_context.referenced_names` of the added + lines in the path's language. + - ``operation_ids``: the identifiers written directly before a ``(`` + (declarations and calls alike), minus the language's keywords. + """ + body = "\n".join(added) + language = language_of(path) + routes = _routes(body) + if routes and text is not None: + prefixes = _route_prefixes(text.removeprefix(_BOM)) + routes += [_join_route(prefix, route) for prefix in prefixes for route in routes] + keywords = _JAVA_KEYWORDS if language == "java" else chunk_context._keywords(language) + return ContractTriggers( + routes=_unique(routes), + tables=_unique(_tables(body)), + names=_unique(referenced_names(added, language)), + operation_ids=_unique(name for name in _OPEN_CALL.findall(body) if name not in keywords), + ) + + +def literal_contract_paths(globs: Sequence[str]) -> list[str]: + """The contract globs that are plain paths, in glob order, deduplicated. + + A glob is literal when it holds no ``*``, ``?`` or ``[`` and does not + start with ``!``. It names one repository-relative path, which is read + directly even when no listing shows it: a miss costs one read. A literal + that a ``!`` negation in the same list vetoes is left out, as + :func:`prxref.rules.match_globs` would leave it out. Blank globs are + ignored. + """ + out: dict[str, None] = {} + for glob in globs: + if not glob.strip() or glob.startswith("!") or any(c in glob for c in "*?["): + continue + if match_globs(glob, globs): + out.setdefault(glob) + return list(out) + + +def select_contract_files( + globs: Sequence[str], + *, + listing: Collection[str] | None, + diff_paths: Sequence[str], +) -> list[str]: + """The run's contract files: sorted, deduplicated repository-relative paths. + + A path from ``listing`` (the repository's file listing at the head, or + None when there is none) or from ``diff_paths`` counts when + :func:`prxref.rules.match_globs` selects it with ``globs``, and every + :func:`literal_contract_paths` path counts even when absent from both. + The caller passes the diff paths that were not removed; nothing is + filtered by status here. The result depends on no chunk, so it is + computed once per run. + """ + selected = {path for path in (*(listing or ()), *diff_paths) if path and match_globs(path, globs)} + selected.update(literal_contract_paths(globs)) + return sorted(selected) + + +def earlier_migrations(path: str, contract_paths: Sequence[str]) -> list[str]: + """The contract files a migration at ``path`` builds on, in ascending order. + + A candidate sits in the same directory as ``path``, has an extension + :func:`contract_excerpts` handles as a migration (``.sql``, ``.xml``, + ``.yaml``, ``.yml`` or ``.json``, case-insensitive), and sorts before + ``path`` under a natural sort of the basenames, where digit runs compare + as numbers, so ``V9__a.sql`` sorts before ``V10__b.sql``. At most + :data:`MAX_EARLIER_MIGRATIONS` are returned: the ones nearest to + ``path``. ``path`` itself is never returned. + """ + directory = posixpath.dirname(path) + own = _natural_key(posixpath.basename(path)) + earlier = sorted( + { + candidate + for candidate in contract_paths + if candidate != path + and posixpath.dirname(candidate) == directory + and _suffix(candidate) in _MIGRATION_SUFFIXES + and _natural_key(posixpath.basename(candidate)) < own + }, + key=lambda candidate: _natural_key(posixpath.basename(candidate)), + ) + return earlier[-MAX_EARLIER_MIGRATIONS:] + + +def contract_excerpts(path: str, text: str, triggers: ContractTriggers) -> list[Excerpt]: + """The excerpts one contract file gives for ``triggers``, chosen by extension. + + The extension is compared case-insensitively. ``names + tables`` below is + the two tuples concatenated and deduplicated. + + - ``.sql`` goes to :func:`sql_excerpts` with the tables. + - ``.xml`` goes to :func:`liquibase_excerpts` with the tables. + - ``.yaml`` and ``.yml``: a top-level ``openapi:`` or ``swagger:`` key (a + line that starts in column 0, the key optionally quoted) goes to + :func:`openapi_yaml_excerpts` with the routes, the operation ids and + ``names + tables`` as schemas. Otherwise a top-level + ``databaseChangeLog:`` goes to :func:`liquibase_excerpts`. Otherwise + the file is a fragment, such as one schema of a split spec: when the + :func:`normalize_name` of its stem equals that of a name or table, the + whole file is one excerpt at line 1 whose symbol is the stem, capped as + every excerpt is; otherwise it gives nothing. A text whose first + non-blank character is ``{`` is JSON and takes the ``.json`` branch. + - ``.json`` is parsed once to sniff its top-level keys: ``openapi`` or + ``swagger`` goes to :func:`openapi_json_excerpts`, + ``databaseChangeLog`` to :func:`liquibase_excerpts`, and anything else, + invalid JSON included, to :func:`json_schema_excerpts` with + ``names + tables``. + - Any other extension gives ``[]``. + + The stem is the basename up to its first dot, so + ``transport-config.schema.json`` has the stem ``transport-config``. + """ + suffix = _suffix(path) + text = text.removeprefix(_BOM) + schemas = _unique((*triggers.names, *triggers.tables)) + if suffix == ".sql": + return sql_excerpts(text, tables=triggers.tables) + if suffix == ".xml": + return liquibase_excerpts(text, tables=triggers.tables) + if suffix not in _SPEC_SUFFIXES: + return [] + if suffix != ".json" and not text.lstrip().startswith("{"): + if _YAML_OPENAPI.search(text): + return openapi_yaml_excerpts( + text, routes=triggers.routes, operation_ids=triggers.operation_ids, schemas=schemas + ) + if _YAML_CHANGELOG.search(text): + return liquibase_excerpts(text, tables=triggers.tables) + return _fragment_excerpts(path, text, schemas) + doc = _load_json(text) + keys = doc if isinstance(doc, dict) else {} + if "openapi" in keys or "swagger" in keys: + return openapi_json_excerpts( + text, routes=triggers.routes, operation_ids=triggers.operation_ids, schemas=schemas + ) + if "databaseChangeLog" in keys: + return liquibase_excerpts(text, tables=triggers.tables) + return json_schema_excerpts(text, names=schemas) + + +def contract_entries( + chunk: Sequence[object], + *, + contract_paths: Sequence[str], + read: Callable[[str], str | None] | None, + priority: Sequence[str] = (), +) -> list[ContextEntry]: + """The contract entries for one worker chunk, ordered by ``(path, line)``. + + ``chunk`` is the chunk's file diffs, duck-typed on ``path``, ``hunks`` and + ``lines`` (``kind``, ``text``, ``new_line``) as + :func:`prxref.chunk_context.chunk_files` reads them. ``contract_paths`` is + :func:`select_contract_files`'s result. ``read`` is the chunk's capped + reader: each call may cost a read, and a path it excludes or cannot read + gives None, which is skipped silently. ``read`` None gives ``[]`` with no + work. ``priority`` lists paths that rank first among the spec files, + normally :func:`literal_contract_paths` of the globs. + + Each chunk file's triggers come from :func:`contract_triggers`. The file's + own text is read only when its added lines hold a route, so its + class-level prefix can be found. The triggers of all the chunk's files + are merged in first-appearance order, and when every field is empty the + result is ``[]`` with no read at all. + + The files then read, in this order and once each, are: + + 1. when there are tables, the :func:`earlier_migrations` of each chunk + file that is itself a contract path, at most + :data:`MAX_EARLIER_MIGRATIONS` per file; + 2. at most :data:`MAX_SPEC_FILES` spec candidates, the contract paths + ending in ``.yaml``, ``.yml`` or ``.json``, ranked: a path in + ``priority``, then a path whose stem matches a name or table by + :func:`normalize_name`, then a basename that starts with ``openapi`` + or ``swagger`` or ends with ``.schema.json`` (case-insensitive), then + the rest, ties broken by path. The ranking is deterministic, so every + chunk picks the same root spec. + + A contract path that is itself a file of this chunk is never read, and + ``.sql`` and ``.xml`` contract files are read only as earlier + migrations. Each file read goes through :func:`contract_excerpts`, and + each :class:`Excerpt` becomes a ``ContextEntry`` of kind and reason + ``"contract"``. Entries that share a ``(path, line)`` keep the first. + No character budget applies here: the only caps are the per-excerpt + ones, :data:`MAX_EARLIER_MIGRATIONS` and :data:`MAX_SPEC_FILES`. + """ + if read is None: + return [] + files = chunk_context.chunk_files(chunk) + per_file: list[ContractTriggers] = [] + for changed in files: + triggers = contract_triggers(changed.path, changed.added) + if triggers.routes: + text = read(changed.path) + if text is not None: + triggers = contract_triggers(changed.path, changed.added, text=text) + per_file.append(triggers) + merged = ContractTriggers( + routes=_unique(route for t in per_file for route in t.routes), + tables=_unique(table for t in per_file for table in t.tables), + names=_unique(name for t in per_file for name in t.names), + operation_ids=_unique(op for t in per_file for op in t.operation_ids), + ) + if not (merged.routes or merged.tables or merged.names or merged.operation_ids): + return [] + in_chunk = {changed.path for changed in files} + queue: dict[str, None] = {} + if merged.tables: + contract_set = set(contract_paths) + for changed in files: + if changed.path in contract_set: + earlier = earlier_migrations(changed.path, contract_paths) + queue.update(dict.fromkeys(p for p in earlier if p not in in_chunk)) + wanted = _keys((*merged.names, *merged.tables), normalize_name) + ranks = frozenset(priority) + specs = sorted( + {p for p in contract_paths if _suffix(p) in _SPEC_SUFFIXES and p not in in_chunk and p not in queue}, + key=lambda p: (_spec_rank(p, ranks, wanted), p), + ) + queue.update(dict.fromkeys(specs[:MAX_SPEC_FILES])) + entries: dict[tuple[str, int], ContextEntry] = {} + for path in queue: + text = read(path) + if text is None: + continue + for excerpt in contract_excerpts(path, text, merged): + entries.setdefault( + (path, excerpt.line), + ContextEntry( + path=path, line=excerpt.line, symbol=excerpt.symbol, kind="contract", reason="contract", + text=excerpt.text, + ), + ) + return [entries[key] for key in sorted(entries)] + + +def _unique(values: Iterable[str]) -> tuple[str, ...]: + return tuple(dict.fromkeys(value for value in values if value)) + + +def _suffix(path: str) -> str: + return posixpath.splitext(path)[1].lower() + + +def _stem(path: str) -> str: + return posixpath.basename(path).split(".", 1)[0] + + +def _natural_key(name: str) -> tuple[tuple[object, ...], str]: + """A natural sort key: digit runs compare by value, without converting them to ``int``.""" + parts = _DIGITS.split(name) + return ( + tuple((len(part.lstrip("0")), part.lstrip("0")) if i % 2 else part for i, part in enumerate(parts)), + name, + ) + + +def _spec_rank(path: str, priority: frozenset[str], wanted: frozenset[str]) -> int: + if path in priority: + return 0 + if normalize_name(_stem(path)) in wanted: + return 1 + base = posixpath.basename(path).lower() + if base.startswith(("openapi", "swagger")) or base.endswith(".schema.json"): + return 2 + return 3 + + +def _fragment_excerpts(path: str, text: str, schemas: Iterable[str]) -> list[Excerpt]: + stem = _stem(path) + if normalize_name(stem) not in _keys(schemas, normalize_name): + return [] + lines = [line.rstrip() for line in text.split("\n")] + while lines and not lines[-1]: + lines.pop() + if not lines: + return [] + capped, _ = _cap(lines) + return [Excerpt(1, stem, capped)] + + +def _first_group(match: re.Match[str]) -> str: + return next(group for group in match.groups() if group is not None) + + +def _tables(body: str) -> list[str]: + """Bare table names in ``body``, in order of appearance.""" + found: list[tuple[int, str]] = [] + verbs = [match.start() for match in _SQL_VERB.finditer(body)] + for start, end in zip(verbs, [*verbs[1:], len(body)], strict=False): + for head in _SQL_HEADS: + match = head.match(body, start, end) + if match: + found.append((start, match.group("name"))) + break + found += [(match.start(), match.group("name")) for match in _SQL_REFERENCES.finditer(body)] + found += [(match.start(), _first_group(match)) for match in _LIQUIBASE_TABLE.finditer(body)] + bare = (_bare_table(name.strip()) for _, name in sorted(found)) + return [name for name in bare if _WORD.search(name)] + + +def _annotation_paths(args: str) -> list[str]: + """The route strings of a Java annotation's arguments: positional, ``value =`` or ``path =``.""" + out: list[str] = [] + for match in _ANNOTATION_ARG.finditer(args): + key, value = match.groups() + if key is not None and key not in ("value", "path"): + continue + out += _JAVA_STRING.findall(value) if value.startswith("{") else [value[1:-1]] + return out + + +def _routes(body: str) -> list[str]: + found: list[tuple[int, int, str]] = [] + for match in _SPRING_MAPPING.finditer(body): + found += [(match.start(), i, route) for i, route in enumerate(_annotation_paths(match.group(2)))] + for match in _JAXRS_PATH.finditer(body): + found += [(match.start(), i, route) for i, route in enumerate(_annotation_paths(match.group(1)))] + for regex in (_PY_ROUTE, _EXPRESS_ROUTE): + found += [(match.start(), 0, _first_group(match)) for match in regex.finditer(body)] + return [route for _, _, route in sorted(found) if route] + + +def _annotation_block(lines: list[str], index: int) -> list[str]: + """The annotation lines directly above line ``index``, a wrapped annotation's arguments included.""" + block: list[str] = [] + depth = 0 + for j in range(index - 1, max(index - 1 - _PREFIX_LOOKBACK_LINES, -1), -1): + line = lines[j] + depth += line.count(")") - line.count("(") + if depth <= 0 and not line.lstrip().startswith("@"): + break + block.append(line) + return block[::-1] + + +def _class_prefixes(text: str) -> tuple[str | None, str | None]: + """The first class-level Spring ``@RequestMapping`` and JAX-RS ``@Path`` paths, or None.""" + lines = text.split("\n") + spring: str | None = None + jaxrs: str | None = None + for index, line in enumerate(lines): + if "class" not in line and "interface" not in line: + continue + decl = _JAVA_TYPE_DECL.match(line) + if decl is None: + continue + head = "\n".join([*_annotation_block(lines, index), line[: decl.start("kw")]]) + if spring is None: + mappings = (m for m in _SPRING_MAPPING.finditer(head) if m.group(1) == "Request") + spring = next((p for m in mappings for p in _annotation_paths(m.group(2))), None) + if jaxrs is None: + jaxrs = next((p for m in _JAXRS_PATH.finditer(head) for p in _annotation_paths(m.group(1))), None) + if spring is not None and jaxrs is not None: + break + return spring, jaxrs + + +def _route_prefixes(text: str) -> list[str]: + prefixes = list(_class_prefixes(text)) + for regex in (_FLASK_PREFIX, _FASTAPI_PREFIX): + match = regex.search(text) + prefixes.append(_first_group(match) if match else None) + return [prefix for prefix in prefixes if prefix] + + +def _join_route(prefix: str, route: str) -> str: + return prefix.rstrip("/") + "/" + route.lstrip("/") diff --git a/tests/test_repo_context_triggers.py b/tests/test_repo_context_triggers.py new file mode 100644 index 0000000..49b2b02 --- /dev/null +++ b/tests/test_repo_context_triggers.py @@ -0,0 +1,708 @@ +"""Unit tests for the contract-trigger half of :mod:`prxref.repo_contracts`. + +Covers :func:`contract_triggers` (routes, tables, names and operation ids on +added lines), :func:`select_contract_files` and :func:`literal_contract_paths`, +:func:`earlier_migrations`, the :func:`contract_excerpts` dispatch, and the +per-chunk :func:`contract_entries`. Everything is pure: the only file reads are +of ``tests/fixtures/issue17/repo/`` through a tiny reader defined here. +""" +from __future__ import annotations + +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from prxref import config +from prxref.repo_context import ContextEntry, language_of, referenced_names +from prxref.repo_contracts import ( + MAX_CONTRACT_LINES, + MAX_EARLIER_MIGRATIONS, + MAX_SPEC_FILES, + ContractTriggers, + Excerpt, + contract_entries, + contract_excerpts, + contract_triggers, + earlier_migrations, + json_schema_excerpts, + liquibase_excerpts, + literal_contract_paths, + openapi_json_excerpts, + openapi_yaml_excerpts, + select_contract_files, + sql_excerpts, +) +from prxref.triage import build_chunks, parse_unified_diff + +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" +REPO = FIXTURE / "repo" +SPEC = "api/openapi/connectors.yaml" +MIGRATION_001 = "db/changelog/001-create-connectors.sql" +MIGRATION_002 = "db/changelog/002-create-idempotency-keys.sql" +MIGRATION_003 = "db/changelog/003-idempotency-unique.sql" +CONNECTOR_SERVICE = "src/main/java/com/acme/connectors/ConnectorService.java" +TRANSPORT_CONFIG = "src/main/java/com/acme/connectors/TransportConfig.java" + +ELLIPSIS = "\N{HORIZONTAL ELLIPSIS}" + + +def fixture_reader(log: list[str]): + """A reader over the fixture repository that records every path asked for.""" + + def read(path: str) -> str | None: + log.append(path) + target = REPO / path + return target.read_text(encoding="utf-8") if target.is_file() else None + + return read + + +def dict_reader(files: dict[str, str | None], log: list[str]): + """A reader over an in-memory repository that records every path asked for.""" + + def read(path: str) -> str | None: + log.append(path) + return files.get(path) + + return read + + +def stand_in(path: str, added: list[str]) -> SimpleNamespace: + """A FileDiff stand-in whose one hunk adds ``added``.""" + lines = [SimpleNamespace(kind="+", text=text, new_line=i + 1) for i, text in enumerate(added)] + return SimpleNamespace(path=path, hunks=[SimpleNamespace(lines=lines)]) + + +def routes(added: list[str], path: str = "src/Api.java", text: str | None = None) -> tuple[str, ...]: + return contract_triggers(path, added, text=text).routes + + +def tables(added: list[str], path: str = "db/V1__x.sql") -> tuple[str, ...]: + return contract_triggers(path, added).tables + + +def summary(entries: list[ContextEntry]) -> list[tuple[str, int, str]]: + return [(entry.path, entry.line, entry.symbol) for entry in entries] + + +@pytest.fixture(scope="module") +def world(): + """The fixture diff parsed and chunked one file per chunk, plus its contract paths.""" + files = parse_unified_diff((FIXTURE / "pr.diff").read_text(encoding="utf-8")) + chunks = build_chunks(files, max_files_per_chunk=1) + listing = sorted(p.relative_to(REPO).as_posix() for p in REPO.rglob("*") if p.is_file()) + contract_paths = select_contract_files( + config._DEFAULTS["context_contract_globs"], + listing=listing, + diff_paths=[f.path for f in files if f.status != "removed"], + ) + by_path = {chunk[0].path: chunk for chunk in chunks} + return SimpleNamespace(chunks=chunks, by_path=by_path, contract_paths=contract_paths) + + +class TestIssue17Fixture: + """The fixture end to end: parse, chunk one file per chunk, select, read, excerpt.""" + + def run(self, world, path: str) -> tuple[list[ContextEntry], list[str]]: + log: list[str] = [] + entries = contract_entries(world.by_path[path], contract_paths=world.contract_paths, read=fixture_reader(log)) + return entries, log + + def test_three_single_file_chunks(self, world): + assert [[f.path for f in chunk] for chunk in world.chunks] == [ + [CONNECTOR_SERVICE], + [TRANSPORT_CONFIG], + [MIGRATION_003], + ] + + def test_default_globs_select_the_spec_and_every_migration(self, world): + assert world.contract_paths == [SPEC, MIGRATION_001, MIGRATION_002, MIGRATION_003] + + def test_migration_chunk_triggers_and_earlier_migrations(self, world): + chunk_file = world.by_path[MIGRATION_003][0] + added = [line.text for hunk in chunk_file.hunks for line in hunk.lines if line.kind == "+"] + triggers = contract_triggers(MIGRATION_003, added) + assert triggers.tables == ("idempotency_keys",) + assert triggers.routes == () + assert "ux_idempotency_keys" not in triggers.tables + assert earlier_migrations(MIGRATION_003, world.contract_paths) == [MIGRATION_001, MIGRATION_002] + + def test_migration_chunk_gets_the_create_table_and_the_idempotency_schema(self, world): + entries, log = self.run(world, MIGRATION_003) + assert summary(entries) == [ + (SPEC, 50, "IdempotencyKey"), + (MIGRATION_002, 4, "idempotency_keys"), + ] + table = entries[1] + assert table.text.startswith("CREATE TABLE idempotency_keys (") + assert "connector_id" in table.text + assert "(tenant, connector, key)" in entries[0].text + assert log == [MIGRATION_001, MIGRATION_002, SPEC] + assert all(entry.path not in (MIGRATION_001, MIGRATION_003) for entry in entries) + + def test_connector_service_chunk_gets_the_path_item_and_both_schemas(self, world): + entries, log = self.run(world, CONNECTOR_SERVICE) + assert summary(entries) == [ + (SPEC, 6, "/connectors/{connectorId}/transports"), + (SPEC, 31, "TransportConfig"), + (SPEC, 41, "CreateTransportRequest"), + ] + path_item = entries[0].text.split("\n") + assert len(path_item) == 23 + assert " operationId: createTransport" in path_item + assert "mutually exclusive" in entries[1].text + assert log == [CONNECTOR_SERVICE, SPEC] + + def test_connector_service_triggers(self, world): + chunk_file = world.by_path[CONNECTOR_SERVICE][0] + added = [line.text for hunk in chunk_file.hunks for line in hunk.lines if line.kind == "+"] + text = (REPO / CONNECTOR_SERVICE).read_text(encoding="utf-8") + triggers = contract_triggers(CONNECTOR_SERVICE, added, text=text) + assert triggers.routes == ("/connectors/{connectorId}/transports",) + assert triggers.tables == () + assert "createTransport" in triggers.operation_ids + assert {"TransportConfig", "CreateTransportRequest"} <= set(triggers.names) + + def test_transport_config_chunk_gets_its_schema_only(self, world): + entries, log = self.run(world, TRANSPORT_CONFIG) + assert summary(entries) == [(SPEC, 31, "TransportConfig")] + assert "mutually exclusive" in entries[0].text + assert log == [SPEC] + chunk_file = world.by_path[TRANSPORT_CONFIG][0] + added = [line.text for hunk in chunk_file.hunks for line in hunk.lines if line.kind == "+"] + assert contract_triggers(TRANSPORT_CONFIG, added) == ContractTriggers( + routes=(), tables=(), names=("TransportConfig",), operation_ids=("IllegalArgumentException",) + ) + + def test_every_entry_is_a_contract(self, world): + for chunk in world.chunks: + entries, _ = self.run(world, chunk[0].path) + assert entries + assert {(entry.kind, entry.reason) for entry in entries} == {("contract", "contract")} + + +class TestRoutes: + def test_spring_positional(self): + assert routes([' @GetMapping("/orders/{id}")']) == ("/orders/{id}",) + + def test_spring_value(self): + assert routes(['@PostMapping(value = "/orders", produces = "application/json")']) == ("/orders",) + + def test_spring_path(self): + assert routes(['@PutMapping(consumes = "text/plain", path = "/orders/{id}")']) == ("/orders/{id}",) + + def test_spring_array_forms(self): + assert routes(['@DeleteMapping({"/a", "/b"})']) == ("/a", "/b") + assert routes(['@RequestMapping(path = {"/c", "/d"}, method = RequestMethod.GET)']) == ("/c", "/d") + + def test_spring_ignores_produces_consumes_name_headers_and_params(self): + line = ( + '@PatchMapping(produces = {"application/json"}, consumes = "text/plain", name = "patch",' + ' headers = "X-A=1", params = "p=1", value = "/orders")' + ) + assert routes([line]) == ("/orders",) + assert routes(['@GetMapping(produces = "application/json")']) == () + + def test_spring_annotation_wrapped_across_added_lines(self): + assert routes(["@GetMapping(", ' value = "/wrapped",', ' produces = "application/json")']) == ( + "/wrapped", + ) + + def test_jax_rs_path(self): + assert routes([' @Path("/items/{id}")', ' public Item get(@PathParam("id") String id) {']) == ( + "/items/{id}", + ) + + def test_fastapi(self): + added = ['@router.get("/items/{item_id}", response_model=Item)', "@app.api_route('/health', methods=['GET'])"] + assert routes(added, path="app/api.py") == ("/items/{item_id}", "/health") + + def test_flask(self): + added = ['@bp.route("/users/", methods=["POST"])', "@app.post('/login')"] + assert routes(added, path="app/views.py") == ("/users/", "/login") + + def test_express_needs_a_leading_slash(self): + added = [ + "router.post('/orders/:id', handler);", + 'app.get("orders", handler);', + "app.use(`/static`, serve);", + "const value = cache.get(key);", + ] + assert routes(added, path="web/server.js") == ("/orders/:id", "/static") + + def test_empty_literals_are_dropped(self): + assert routes(['@router.get("")'], path="app/api.py") == () + + def test_first_appearance_order_and_dedup(self): + added = ['@GetMapping("/b")', '@GetMapping("/a")', '@PostMapping("/b")'] + assert routes(added) == ("/b", "/a") + + +SPRING_CONTROLLER = """package com.acme.orders; + +import org.springframework.web.bind.annotation.*; + +@RestController +@RequestMapping("/api/v1") +public class OrderController { + + @RequestMapping("/not-a-prefix") + @GetMapping("/orders") + public List list() { + return List.of(); + } +} +""" + +WRAPPED_SPRING_CONTROLLER = """package com.acme.orders; + +@RestController +@RequestMapping( + value = {"/api", "/legacy"}, + produces = "application/json") +public class OrderController { +} +""" + + +class TestRoutePrefixes: + def test_class_level_spring_prefix_is_joined(self): + assert routes([' @GetMapping("/orders")'], text=SPRING_CONTROLLER) == ("/orders", "/api/v1/orders") + + def test_no_text_means_no_prefix(self): + assert routes([' @GetMapping("/orders")']) == ("/orders",) + + def test_wrapped_class_annotation_gives_its_first_path(self): + assert routes(['@GetMapping("/orders")'], text=WRAPPED_SPRING_CONTROLLER) == ("/orders", "/api/orders") + + def test_segments_without_a_slash_are_joined_with_one(self): + text = '@RestController\n@RequestMapping("api/")\npublic class C {\n}\n' + assert routes(['@GetMapping("orders")'], text=text) == ("orders", "api/orders") + + def test_a_method_level_request_mapping_is_not_a_prefix(self): + text = 'public class C {\n @RequestMapping("/method")\n public void m() {}\n}\n' + assert routes(['@GetMapping("/x")'], text=text) == ("/x",) + + def test_class_level_jax_rs_path(self): + text = '@Path("/connectors")\n@Produces(MediaType.APPLICATION_JSON)\npublic class ConnectorResource {\n}\n' + assert routes([' @Path("{id}/transports")'], text=text) == ( + "{id}/transports", + "/connectors/{id}/transports", + ) + + def test_same_line_annotation_counts(self): + text = '@RequestMapping("/api") public interface OrdersApi {\n}\n' + assert routes(['@GetMapping("/orders")'], text=text) == ("/orders", "/api/orders") + + def test_flask_url_prefix(self): + text = 'bp = Blueprint("users", __name__, url_prefix="/users")\n\n\n@bp.route("/")\ndef show(uid):\n' + assert routes(['@bp.route("/")'], path="app/users.py", text=text) == ( + "/", + "/users/", + ) + + def test_fastapi_prefix(self): + text = 'router = APIRouter(prefix="/items", tags=["items"])\n\n@router.get("/{item_id}")\n' + assert routes(['@router.get("/{item_id}")'], path="app/items.py", text=text) == ( + "/{item_id}", + "/items/{item_id}", + ) + + def test_only_the_first_prefix_of_each_kind(self): + text = ( + 'a = Blueprint("a", __name__, url_prefix="/first")\n' + 'b = Blueprint("b", __name__, url_prefix="/second")\n' + 'r = APIRouter(prefix="/router")\n' + ) + assert routes(['@a.get("/x")'], path="app/x.py", text=text) == ("/x", "/first/x", "/router/x") + + def test_text_without_a_route_on_the_added_lines_adds_nothing(self): + assert routes(["return List.of();"], text=SPRING_CONTROLLER) == () + + +class TestTables: + def test_create_table_with_schema_prefix_and_quoting(self): + assert tables(['CREATE TABLE IF NOT EXISTS "public"."Orders" (']) == ("Orders",) + assert tables(["create table `shop`.`line_items` ("]) == ("line_items",) + assert tables(["CREATE TABLE [dbo].[Invoices] ("]) == ("Invoices",) + + def test_alter_table(self): + assert tables(["ALTER TABLE ONLY public.accounts ADD COLUMN region text;"]) == ("accounts",) + + def test_index_on_gives_the_table_never_the_index_name(self): + found = tables(["CREATE UNIQUE INDEX ux_idempotency_keys ON idempotency_keys (tenant_id, key);"]) + assert found == ("idempotency_keys",) + + def test_index_wrapped_across_added_lines(self): + assert tables(["CREATE INDEX ix_orders_tenant", " ON orders (tenant_id);"]) == ("orders",) + + def test_references_is_a_foreign_key_target(self): + added = [ + " connector_id VARCHAR(36) REFERENCES connectors (id),", + " tenant_id VARCHAR(36) REFERENCES tenants ON DELETE CASCADE", + ] + assert tables(added) == ("connectors", "tenants") + + def test_prose_references_is_not_a_table(self): + assert tables(["# this references the old layout"], path="app/models.py") == () + + def test_liquibase_xml(self): + added = [ + '', + '', + ] + assert tables(added, path="db/changelog/004.xml") == ("idempotency_keys", "orders", "tenants") + + def test_liquibase_yaml(self): + added = [" - createIndex:", " tableName: idempotency_keys", " baseTableName: 'orders'"] + assert tables(added, path="db/changelog/004.yaml") == ("idempotency_keys", "orders") + + def test_liquibase_json(self): + added = [' "tableName": "idempotency_keys",', ' "referencedTableName" : "tenants"'] + assert tables(added, path="db/changelog/004.json") == ("idempotency_keys", "tenants") + + def test_sql_inside_a_java_string(self): + added = [' jdbc.execute("ALTER TABLE accounts ADD COLUMN region text");'] + assert tables(added, path="src/main/java/Migrate.java") == ("accounts",) + + def test_first_appearance_order_and_dedup(self): + added = ["ALTER TABLE b ADD x int;", "CREATE TABLE a (id int);", "ALTER TABLE b ADD y int;"] + assert tables(added) == ("b", "a") + + +class TestNamesAndOperationIds: + def test_java_line(self): + path = "src/main/java/com/acme/OrderService.java" + added = [ + " public Order createOrder(CreateOrderRequest request) {", + " for(int i = 0; i < n; i++) save(i);", + ] + triggers = contract_triggers(path, added) + assert triggers.names == tuple(referenced_names(added, language_of(path))) + assert triggers.names == ("Order", "CreateOrderRequest") + assert triggers.operation_ids == ("createOrder", "save") + + def test_python_line(self): + path = "app/orders.py" + added = [" result = create_order(payload); print(len(result)); obj.save (x)"] + triggers = contract_triggers(path, added) + assert triggers.names == tuple(referenced_names(added, language_of(path))) + assert triggers.operation_ids == ("create_order",) + + def test_nothing_on_blank_lines(self): + assert contract_triggers("x.txt", ["", "}"]) == ContractTriggers() + + +class TestSelection: + GLOBS = config._DEFAULTS["context_contract_globs"] + + def test_a_glob_match_from_the_listing(self): + listing = ["api/openapi.yaml", "src/Order.java", "db/migration/V1__init.sql"] + assert select_contract_files(self.GLOBS, listing=listing, diff_paths=[]) == [ + "api/openapi.yaml", + "db/migration/V1__init.sql", + ] + + def test_a_diff_only_path(self): + assert select_contract_files(self.GLOBS, listing=["README.md"], diff_paths=["db/migration/V2__x.sql"]) == [ + "db/migration/V2__x.sql" + ] + + def test_a_literal_glob_absent_from_the_listing_is_included(self): + globs = ["docs/api/contract.yaml", "**/openapi*.y*ml"] + assert select_contract_files(globs, listing=["a/openapi.yml"], diff_paths=[]) == [ + "a/openapi.yml", + "docs/api/contract.yaml", + ] + + def test_a_negation_vetoes_listing_paths_and_literals(self): + globs = ["**/migrations/**", "migrations/archive/old.sql", "!**/archive/**"] + listing = ["migrations/001.sql", "migrations/archive/000.sql"] + assert select_contract_files(globs, listing=listing, diff_paths=[]) == ["migrations/001.sql"] + + def test_listing_none_keeps_diff_paths_and_literals(self): + globs = ["api/spec.json", "**/db/changelog/**"] + assert select_contract_files(globs, listing=None, diff_paths=["db/changelog/2.sql", "src/A.java"]) == [ + "api/spec.json", + "db/changelog/2.sql", + ] + + def test_sorted_and_deduplicated(self): + globs = ["**/*.schema.json", "b.schema.json"] + listing = ["z/c.schema.json", "b.schema.json", "a.schema.json"] + assert select_contract_files(globs, listing=listing, diff_paths=["a.schema.json", "z/c.schema.json"]) == [ + "a.schema.json", + "b.schema.json", + "z/c.schema.json", + ] + + def test_literal_contract_paths(self): + globs = ["a.yaml", "b/*.yaml", "!c.yaml", "a.yaml", "", " ", "d?.yaml", "e[1].yaml", "f.json", "c.yaml"] + assert literal_contract_paths(globs) == ["a.yaml", "f.json"] + + def test_the_default_globs_hold_no_literal(self): + assert literal_contract_paths(self.GLOBS) == [] + + +class TestEarlierMigrations: + def test_natural_sort_puts_v9_before_v10(self): + paths = ["db/migration/V10__b.sql", "db/migration/V9__a.sql", "db/migration/V11__c.sql"] + assert earlier_migrations("db/migration/V11__c.sql", paths) == [ + "db/migration/V9__a.sql", + "db/migration/V10__b.sql", + ] + + def test_at_most_the_nearest_four(self): + paths = [f"db/migration/V{i}__step.sql" for i in range(1, 9)] + assert MAX_EARLIER_MIGRATIONS == 4 + assert earlier_migrations("db/migration/V8__step.sql", paths) == [ + f"db/migration/V{i}__step.sql" for i in (4, 5, 6, 7) + ] + + def test_other_directories_other_extensions_later_files_and_itself_are_excluded(self): + paths = [ + "db/changelog/001.sql", + "db/changelog/002.md", + "db/changelog/003.java", + "db/changelog/nested/001.sql", + "db/other/001.sql", + "db/changelog/004.sql", + "db/changelog/005.sql", + ] + assert earlier_migrations("db/changelog/004.sql", paths) == ["db/changelog/001.sql"] + + def test_every_migration_extension_counts(self): + paths = ["m/1.sql", "m/2.XML", "m/3.yaml", "m/4.yml", "m/5.json", "m/6.sql"] + assert earlier_migrations("m/6.sql", paths) == ["m/2.XML", "m/3.yaml", "m/4.yml", "m/5.json"] + + +OPENAPI_YAML = """openapi: 3.0.3 +paths: + /orders: + post: + operationId: createOrder +components: + schemas: + Order: + type: object +""" + +LIQUIBASE_YAML = """databaseChangeLog: + - changeSet: + id: 1 + changes: + - createTable: + tableName: orders +""" + +OPENAPI_JSON = ( + '{"openapi":"3.0.0","paths":{"/orders":{"post":{"operationId":"createOrder"}}},' + '"components":{"schemas":{"Order":{"type":"object"}}}}' +) + +LIQUIBASE_JSON = ( + '{"databaseChangeLog": [{"changeSet": {"id": "1", "changes": [{"createTable": {"tableName": "orders"}}]}}]}' +) + +SCHEMA_JSON = '{\n "title": "Order",\n "type": "object",\n "properties": {"openapi": {"type": "string"}}\n}\n' + +TRIGGERS = ContractTriggers( + routes=("/orders",), tables=("orders",), names=("Order", "orders"), operation_ids=("createOrder",) +) + + +class TestDispatch: + def test_sql(self): + text = "CREATE TABLE orders (id int);\nCREATE TABLE other (id int);\n" + assert contract_excerpts("db/V1__init.SQL", text, TRIGGERS) == sql_excerpts(text, tables=("orders",)) + assert [e.symbol for e in contract_excerpts("db/V1__init.sql", text, TRIGGERS)] == ["orders"] + + def test_xml(self): + text = '\n\n\n\n' + found = contract_excerpts("db/changelog/1.Xml", text, TRIGGERS) + assert found == liquibase_excerpts(text, tables=("orders",)) + assert [e.line for e in found] == [2] + + def test_openapi_yaml(self): + found = contract_excerpts("api/openapi.yaml", OPENAPI_YAML, TRIGGERS) + assert found == openapi_yaml_excerpts( + OPENAPI_YAML, routes=("/orders",), operation_ids=("createOrder",), schemas=("Order", "orders") + ) + assert [(e.line, e.symbol) for e in found] == [(3, "/orders"), (8, "Order")] + + def test_swagger_yml_with_a_bom(self): + text = '\N{ZERO WIDTH NO-BREAK SPACE}swagger: "2.0"\ndefinitions:\n Order:\n type: object\n' + assert [(e.line, e.symbol) for e in contract_excerpts("api/swagger.yml", text, TRIGGERS)] == [(3, "Order")] + + def test_liquibase_yaml(self): + found = contract_excerpts("db/changelog/1.yaml", LIQUIBASE_YAML, TRIGGERS) + assert found == liquibase_excerpts(LIQUIBASE_YAML, tables=("orders",)) + assert [e.line for e in found] == [2] + + def test_fragment_matched_by_stem(self): + text = "type: object\nproperties:\n id:\n type: string\n\n" + found = contract_excerpts("api/openapi/schemas/Order.yaml", text, TRIGGERS) + assert found == [Excerpt(1, "Order", "type: object\nproperties:\n id:\n type: string")] + + def test_fragment_matched_by_a_table_through_normalize_name(self): + triggers = ContractTriggers(tables=("idempotency_keys",)) + found = contract_excerpts("api/schemas/idempotency-key.v2.yml", "type: object\n", triggers) + assert found == [Excerpt(1, "idempotency-key", "type: object")] + + def test_fragment_not_matched(self): + assert contract_excerpts("api/openapi/schemas/Invoice.yaml", "type: object\n", TRIGGERS) == [] + + def test_fragment_is_capped_like_every_excerpt(self): + text = "\n".join(f"line{i}: x" for i in range(60)) + (excerpt,) = contract_excerpts("api/schemas/Order.yaml", text, TRIGGERS) + lines = excerpt.text.split("\n") + assert len(lines) == MAX_CONTRACT_LINES + assert lines[-1] == f"{ELLIPSIS} 21 more lines" + + def test_json_openapi(self): + found = contract_excerpts("api/openapi.json", OPENAPI_JSON, TRIGGERS) + assert found == openapi_json_excerpts( + OPENAPI_JSON, routes=("/orders",), operation_ids=("createOrder",), schemas=("Order", "orders") + ) + assert [e.symbol for e in found] == ["/orders", "Order"] + + def test_json_liquibase(self): + found = contract_excerpts("db/changelog/1.json", LIQUIBASE_JSON, TRIGGERS) + assert found == liquibase_excerpts(LIQUIBASE_JSON, tables=("orders",)) + assert [e.symbol for e in found] == ["orders"] + + def test_json_schema_even_with_a_nested_openapi_key(self): + found = contract_excerpts("schemas/order.schema.json", SCHEMA_JSON, TRIGGERS) + assert found == json_schema_excerpts(SCHEMA_JSON, names=("Order", "orders")) + assert [(e.line, e.symbol) for e in found] == [(1, "Order")] + + def test_json_text_in_a_yaml_file_takes_the_json_branch(self): + found = contract_excerpts("api/openapi.yaml", OPENAPI_JSON, TRIGGERS) + assert [e.symbol for e in found] == ["/orders", "Order"] + + def test_unknown_extensions(self): + assert contract_excerpts("api/openapi.md", OPENAPI_YAML, TRIGGERS) == [] + assert contract_excerpts("src/Order.java", "public class Order {}", TRIGGERS) == [] + + +ORDER_SERVICE = "src/main/java/com/acme/OrderService.java" +ORDER_LINE = " Order order = repository.find(id);" + + +class TestContractEntries: + def test_read_none_gives_nothing(self): + chunk = [stand_in(ORDER_SERVICE, [ORDER_LINE])] + assert contract_entries(chunk, contract_paths=["api/openapi.yaml"], read=None) == [] + + def test_empty_triggers_read_nothing(self): + log: list[str] = [] + chunk = [stand_in("notes.txt", ["", "}", " "])] + found = contract_entries(chunk, contract_paths=["api/openapi.yaml", "db/1.sql"], read=dict_reader({}, log)) + assert found == [] + assert log == [] + + def test_a_contract_file_in_the_chunk_is_not_read(self): + log: list[str] = [] + chunk = [stand_in("api/openapi.yaml", [" operationId: createOrder"]), stand_in(ORDER_SERVICE, [ORDER_LINE])] + files = {"api/openapi.yaml": OPENAPI_YAML, "api/schemas/Order.yaml": "type: object\n"} + found = contract_entries(chunk, contract_paths=sorted(files), read=dict_reader(files, log)) + assert log == ["api/schemas/Order.yaml"] + assert summary(found) == [("api/schemas/Order.yaml", 1, "Order")] + + def test_max_spec_files_keeps_the_six_best_ranked(self): + candidates = [ + "a/openapi.yaml", + "b/Order.yaml", + "c/swagger.json", + "d/thing.schema.json", + "e/aaa.yaml", + "f/bbb.yaml", + "g/ccc.yml", + "z/priority.yaml", + ] + log: list[str] = [] + found = contract_entries( + [stand_in(ORDER_SERVICE, [ORDER_LINE])], + contract_paths=[*candidates, "db/changelog/1.sql", "db/changelog/2.xml"], + read=dict_reader({}, log), + priority=["z/priority.yaml", "not/a/contract.yaml"], + ) + assert MAX_SPEC_FILES == 6 + assert found == [] + assert log == [ + "z/priority.yaml", + "b/Order.yaml", + "a/openapi.yaml", + "c/swagger.json", + "d/thing.schema.json", + "e/aaa.yaml", + ] + + def test_a_none_read_is_skipped(self): + log: list[str] = [] + files = {"api/schemas/Order.yaml": "type: object\n", "api/openapi.yaml": None} + found = contract_entries( + [stand_in(ORDER_SERVICE, [ORDER_LINE])], contract_paths=sorted(files), read=dict_reader(files, log) + ) + assert log == ["api/schemas/Order.yaml", "api/openapi.yaml"] + assert summary(found) == [("api/schemas/Order.yaml", 1, "Order")] + + def test_order_by_path_and_line_and_dedup_on_both(self): + files = {"b/openapi.json": OPENAPI_JSON, "a/openapi.yaml": OPENAPI_YAML} + chunk = [stand_in("src/orders.js", ["router.post('/orders', createOrder);", "const o = new Order();"])] + found = contract_entries(chunk, contract_paths=sorted(files), read=dict_reader(files, [])) + assert summary(found) == [ + ("a/openapi.yaml", 3, "/orders"), + ("a/openapi.yaml", 8, "Order"), + ("b/openapi.json", 1, "/orders"), + ] + + def test_triggers_merge_across_the_chunk_files(self): + spec = "openapi: 3.0.3\npaths:\n /orders:\n post: {}\ncomponents:\n schemas:\n Invoice:\n type: x" + files = {"a/openapi.yaml": spec} + chunk = [ + stand_in("src/a.js", ["router.post('/orders', h);"]), + stand_in("src/main/java/com/acme/Billing.java", [" Invoice invoice = billing.find(id);"]), + ] + found = contract_entries(chunk, contract_paths=sorted(files), read=dict_reader(files, [])) + assert summary(found) == [("a/openapi.yaml", 3, "/orders"), ("a/openapi.yaml", 7, "Invoice")] + + def test_only_a_route_bearing_file_is_read_for_its_prefix(self): + log: list[str] = [] + controller = "src/main/java/com/acme/OrderController.java" + spec = "openapi: 3.0.3\npaths:\n /api/v1/orders:\n get: {}\n" + files = {controller: SPRING_CONTROLLER, "api/openapi.yaml": spec} + chunk = [stand_in(controller, [' @GetMapping("/orders")']), stand_in(ORDER_SERVICE, [ORDER_LINE])] + found = contract_entries(chunk, contract_paths=["api/openapi.yaml"], read=dict_reader(files, log)) + assert log == [controller, "api/openapi.yaml"] + assert summary(found) == [("api/openapi.yaml", 3, "/api/v1/orders")] + + def test_migrations_are_read_only_when_there_are_tables(self): + paths = ["db/changelog/1.sql", "db/changelog/2.sql", "db/changelog/3.sql"] + files = {"db/changelog/1.sql": "CREATE TABLE orders (id int);\n", "db/changelog/2.sql": "SELECT 1;\n"} + log: list[str] = [] + comment_only = [stand_in("db/changelog/3.sql", ["-- a comment"])] + assert contract_entries(comment_only, contract_paths=paths, read=dict_reader(files, log)) == [] + assert log == [] + with_table = [stand_in("db/changelog/3.sql", ["ALTER TABLE orders ADD region text;"])] + found = contract_entries(with_table, contract_paths=paths, read=dict_reader(files, log)) + assert log == ["db/changelog/1.sql", "db/changelog/2.sql"] + assert summary(found) == [("db/changelog/1.sql", 1, "orders")] + + def test_an_earlier_migration_in_the_chunk_is_not_read(self): + paths = ["db/changelog/1.sql", "db/changelog/2.sql", "db/changelog/3.sql"] + log: list[str] = [] + chunk = [ + stand_in("db/changelog/2.sql", ["CREATE TABLE orders (id int);"]), + stand_in("db/changelog/3.sql", ["ALTER TABLE orders ADD region text;"]), + ] + contract_entries(chunk, contract_paths=paths, read=dict_reader({}, log)) + assert log == ["db/changelog/1.sql"] + + def test_sql_and_xml_contracts_are_never_read_as_specs(self): + log: list[str] = [] + paths = ["db/changelog/1.sql", "db/changelog/2.xml", "other/orders.sql"] + chunk = [stand_in("src/orders.py", ["save_order(order)"])] + assert contract_entries(chunk, contract_paths=paths, read=dict_reader({}, log)) == [] + assert log == [] From 57390d6c61427f5c1898268efce9d183f4263a0a Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:06:07 -0700 Subject: [PATCH 15/24] feat: per-chunk repository context unit (sources, order, budget, exclude floor) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Issue #17, task T7. build_unit_context (new src/prxref/repo_unit.py) combines the three repository-context sources for one worker chunk: diff_definitions (diff and repo), contract_entries and the resolver (repo with a reader only), called in admission-rank order so a capped reader spends its reads on the higher-ranked sources first. Every read goes through a guard that refuses an excluded path without calling the reader (an exclude that raises fails closed). Entries at excluded paths are dropped, the rest ordered by (REASONS rank, path, line), deduplicated on (path, line), and admitted against one character budget that stops at the first entry that does not fit; the omitted line goes to the block of the first entry left out and does not count toward the budget. "off" returns EMPTY_UNIT with no read, and "repo" with no reader equals "diff" with no reader (OQ2). repo_context.py gains EXCLUDE_FLOOR and exclude_predicate(extra_globs), which matches the floor and the extra globs as two separate match_globs lists, so a "!" in PRXREF_CONTEXT_EXCLUDE_GLOBS can never re-admit a label file. chunk_context.render_context_blocks gains extra_def_lines (under the one definitions header) and contract_lines (under the new CONTRACT_HEADER); with the defaults it is byte-identical to before (D1), pinned against a verbatim oracle copy. Nothing calls the new code yet; T13 wires it. Tests: tests/test_repo_context_unit.py (94), including the issue17 fixture end to end through the real capped reader, the observed read order, the chunk cap degrading, and the budget, order and exclusion rules on canned sources. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/chunk_context.py | 29 +- src/prxref/repo_context.py | 43 +- src/prxref/repo_unit.py | 266 +++++++++++ tests/test_repo_context_unit.py | 820 ++++++++++++++++++++++++++++++++ 4 files changed, 1146 insertions(+), 12 deletions(-) create mode 100644 src/prxref/repo_unit.py create mode 100644 tests/test_repo_context_unit.py diff --git a/src/prxref/chunk_context.py b/src/prxref/chunk_context.py index d597634..265ccc2 100644 --- a/src/prxref/chunk_context.py +++ b/src/prxref/chunk_context.py @@ -39,6 +39,7 @@ DEPENDENCY_HEADER = "### Dependency versions" DEFINITIONS_HEADER = "### Definitions referenced by this chunk" SIBLING_HEADER = "### Other files changed in this PR" +CONTRACT_HEADER = "### Contract excerpts" _JS_SUFFIXES = (".js", ".jsx", ".mjs", ".cjs", ".ts", ".tsx", ".mts", ".cts") @@ -444,17 +445,31 @@ def referenced_definitions( return out -def render_context_blocks(dep_lines: Sequence[str], def_lines: Sequence[str]) -> str: - """Render the two prompt blocks, omitting each when it has no lines. - - Returns the empty string when both are empty, so the prompt slot leaves no - stray header behind. +def render_context_blocks( + dep_lines: Sequence[str], + def_lines: Sequence[str], + extra_def_lines: Sequence[str] = (), + contract_lines: Sequence[str] = (), +) -> str: + """Render the prompt blocks, omitting each when it has no lines. + + The dependency block comes first. The definitions block holds + ``def_lines`` followed by ``extra_def_lines`` (repository-context + definitions) under one :data:`DEFINITIONS_HEADER`, and is present when + either is non-empty. The contracts block, ``contract_lines`` under + :data:`CONTRACT_HEADER`, comes last. With the two optional arguments + empty, the output is exactly the two-block rendering repository context + predates. Returns the empty string when every list is empty, so the + prompt slot leaves no stray header behind. """ blocks: list[str] = [] if dep_lines: blocks.append(DEPENDENCY_HEADER + "\n\n" + "\n".join(dep_lines)) - if def_lines: - blocks.append(DEFINITIONS_HEADER + "\n\n" + "\n".join(def_lines)) + definitions = [*def_lines, *extra_def_lines] + if definitions: + blocks.append(DEFINITIONS_HEADER + "\n\n" + "\n".join(definitions)) + if contract_lines: + blocks.append(CONTRACT_HEADER + "\n\n" + "\n".join(contract_lines)) return "\n\n".join(blocks) diff --git a/src/prxref/repo_context.py b/src/prxref/repo_context.py index 90ace55..909bb73 100644 --- a/src/prxref/repo_context.py +++ b/src/prxref/repo_context.py @@ -5,11 +5,13 @@ holds the pieces the repository-context feature (``PRXREF_REPO_CONTEXT``) builds on: the :class:`ContextEntry` record every context source emits, the admission ranks in :data:`REASONS`, a language map and definition regexes that -add Java, the names an added line references, and a definition scan over the -text of ANY file. +add Java, the names an added line references, a definition scan over the +text of ANY file, and the exclude floor (:data:`EXCLUDE_FLOOR`, +:func:`exclude_predicate`) that keeps label files and secrets out of every read. -The module is pure. It is stdlib plus :mod:`prxref.chunk_context`, performs no -I/O and no network, and callers hand it file text they already read. It +The module is pure. It is stdlib plus :mod:`prxref.chunk_context` and +:func:`prxref.rules.match_globs`, performs no I/O and no network, and +callers hand it file text they already read. It imports ``chunk_context``'s underscore helpers (``_language``, ``_definition_regexes``, ``_keywords``, ``_entry_text``, ``_IDENT_RE``) on purpose, so both modules scan and render a definition the same way instead of @@ -20,10 +22,11 @@ from __future__ import annotations import re -from collections.abc import Iterable +from collections.abc import Callable, Iterable, Sequence from dataclasses import dataclass from . import chunk_context +from .rules import match_globs REASONS = ("cross-chunk", "contract", "diff-file", "import", "path-convention", "name-search") KINDS = ("definition", "contract") @@ -230,3 +233,33 @@ def find_definitions( out.append((name, number, chunk_context._entry_text(lines, idx, max_lines))) break return out + + +EXCLUDE_FLOOR = ( + "**/expected.json", + "**/cases.json", + "**/case.json", + "**/prxref-eval/**", + "**/.env*", + "**/*.pem", + "**/*.key", +) + + +def exclude_predicate(extra_globs: Sequence[str] = ()) -> Callable[[str], bool]: + """A ``path -> bool`` that is true for a path repository context must never read. + + A path is excluded when :func:`prxref.rules.match_globs` selects it with + :data:`EXCLUDE_FLOOR` (eval labels, eval output, dotenv files, keys), or + with ``extra_globs`` (``PRXREF_CONTEXT_EXCLUDE_GLOBS``) when that list is + non-empty. The two lists are matched separately, so the extra globs only + ever ADD exclusions: a ``!`` negation in ``extra_globs`` vetoes only the + extra list's own positives and can never re-admit a floor path, which it + would if both were one ``match_globs`` list. + """ + extra = tuple(extra_globs) + + def excluded(path: str) -> bool: + return match_globs(path, EXCLUDE_FLOOR) or bool(extra and match_globs(path, extra)) + + return excluded diff --git a/src/prxref/repo_unit.py b/src/prxref/repo_unit.py new file mode 100644 index 0000000..374a91f --- /dev/null +++ b/src/prxref/repo_unit.py @@ -0,0 +1,266 @@ +"""One worker chunk's repository context: sources, order, budget and exclusion (#17). + +:func:`build_unit_context` combines the three repository-context sources for +one chunk into the lines its prompt carries: + +1. :func:`prxref.repo_crosschunk.diff_definitions`: definitions from the PR's + own diff files (``cross-chunk`` and ``diff-file``), at the ``diff`` and + ``repo`` levels; +2. :func:`prxref.repo_contracts.contract_entries`: contract excerpts + (``contract``), at the ``repo`` level only; +3. the resolver, :func:`prxref.repo_resolve.resolve_candidates` plus + :func:`prxref.repo_context.find_definitions` over the candidate files + (``import``, ``path-convention``, ``name-search``), at the ``repo`` level + only. + +The sources are called in that order, which is their admission rank, so under +a capped reader the higher-ranked sources get their reads first. Every read +goes through one guard that refuses an excluded path without calling the +reader. The entries are then filtered by the exclude predicate, ordered by +``(REASONS rank, path, line)``, deduplicated on ``(path, line)`` and admitted +against one character budget, which stops at the first entry that does not +fit. + +The module is pure: stdlib plus the pure repository-context modules and +:mod:`prxref.chunk_context`, with no I/O except through the ``read`` callable +the caller passes, nothing from :mod:`prxref.forges`, and no threads. The same +inputs and the same reader answers give the same result and the same read +order. +""" +from __future__ import annotations + +from collections.abc import Callable, Collection, Sequence +from dataclasses import dataclass + +from .chunk_context import chunk_files +from .repo_context import ( + REASONS, + ContextEntry, + definition_regexes, + find_definitions, + language_of, + referenced_names, +) +from .repo_contracts import contract_entries +from .repo_crosschunk import diff_definitions +from .repo_resolve import resolve_candidates + +MODES = ("off", "diff", "repo") + + +@dataclass(frozen=True) +class UnitContext: + """The repository context admitted into one worker chunk's prompt. + + ``definition_lines`` and ``contract_lines`` are the rendered admitted + entries of each kind, in admission order, ready for + :func:`prxref.chunk_context.render_context_blocks` as ``extra_def_lines`` + and ``contract_lines``; when entries were left out, one omitted line (a + horizontal ellipsis, then ``N more context entries omitted``) ends the + list of the kind of the first entry left out. ``entries`` holds the admitted entries in + admission order, and ``omitted`` counts the entries the budget left out. + """ + + definition_lines: tuple[str, ...] + contract_lines: tuple[str, ...] + entries: tuple[ContextEntry, ...] + omitted: int + + def record(self) -> dict: + """The run-record row for this chunk: ``{"entries": [...], "omitted": N}``.""" + return {"entries": [entry.record() for entry in self.entries], "omitted": self.omitted} + + +EMPTY_UNIT = UnitContext((), (), (), 0) + + +def _excluded(exclude: Callable[[str], bool] | None, path: str) -> bool: + if exclude is None: + return False + try: + return bool(exclude(path)) + except Exception: # noqa: BLE001 + return True + + +def _guard( + read: Callable[[str], str | None] | None, + exclude: Callable[[str], bool] | None, +) -> Callable[[str], str | None] | None: + if read is None: + return None + + def guarded(path: str) -> str | None: + if _excluded(exclude, path): + return None + return read(path) + + return guarded + + +def _resolver_entries( + chunk: Sequence[object], + all_files: Sequence[object], + read: Callable[[str], str | None], + *, + listing_paths: Collection[str] | None, + listing_complete: bool, + exclude: Callable[[str], bool] | None, + found: Collection[str], +) -> list[ContextEntry]: + known = set(found) + removed = {getattr(f, "path", "") for f in chunk if getattr(f, "status", "") == "removed"} + diff_paths = {getattr(f, "path", "") for f in all_files} + out: list[ContextEntry] = [] + for changed in chunk_files(chunk): + if changed.path in removed: + continue + language = language_of(changed.path) + if not definition_regexes(language): + continue + names = [name for name in referenced_names(changed.added, language) if name not in known] + if not names: + continue + text = read(changed.path) + if not isinstance(text, str): + text = "\n".join(changed.added) + candidates = resolve_candidates( + changed.path, text, names, listing=listing_paths, listing_complete=listing_complete + ) + by_path: dict[str, dict[str, str]] = {} + for candidate in candidates: + by_path.setdefault(candidate.path, {}).setdefault(candidate.name, candidate.reason) + for path, reasons in by_path.items(): + if path in diff_paths or _excluded(exclude, path): + continue + wanted = [name for name in reasons if name not in known] + if not wanted: + continue + target = read(path) + if not isinstance(target, str): + continue + for symbol, line, body in find_definitions(target, wanted, language=language_of(path)): + out.append(ContextEntry(path, line, symbol, "definition", reasons[symbol], body)) + known.add(symbol) + return out + + +def _merge(entries: list[ContextEntry], exclude: Callable[[str], bool] | None) -> list[ContextEntry]: + kept = [entry for entry in entries if not _excluded(exclude, entry.path)] + kept.sort(key=lambda entry: (REASONS.index(entry.reason), entry.path, entry.line)) + out: list[ContextEntry] = [] + seen: set[tuple[str, int]] = set() + for entry in kept: + key = (entry.path, entry.line) + if key not in seen: + seen.add(key) + out.append(entry) + return out + + +def _admit(entries: list[ContextEntry], max_chars: int) -> UnitContext: + admitted: list[ContextEntry] = [] + used = 0 + for entry in entries: + size = len(entry.rendered()) + if used + size > max_chars: + break + admitted.append(entry) + used += size + omitted = len(entries) - len(admitted) + definition_lines = [entry.rendered() for entry in admitted if entry.kind == "definition"] + contract_lines = [entry.rendered() for entry in admitted if entry.kind != "definition"] + if omitted: + marker = f"\N{HORIZONTAL ELLIPSIS} {omitted} more context entries omitted" + if entries[len(admitted)].kind == "definition": + definition_lines.append(marker) + else: + contract_lines.append(marker) + return UnitContext(tuple(definition_lines), tuple(contract_lines), tuple(admitted), omitted) + + +def build_unit_context( + chunk: Sequence[object], + all_files: Sequence[object], + *, + mode: str, + read: Callable[[str], str | None] | None, + max_chars: int, + listing_paths: Collection[str] | None = None, + listing_complete: bool = False, + contract_paths: Sequence[str] = (), + contract_priority: Sequence[str] = (), + exclude: Callable[[str], bool] | None = None, +) -> UnitContext: + """The repository context for one worker chunk at level ``mode``. + + ``chunk`` holds the chunk's ``triage.FileDiff`` records and ``all_files`` + the whole PR's, duck-typed as :func:`prxref.repo_crosschunk.diff_definitions` + reads them. ``mode`` is ``"off"``, ``"diff"`` or ``"repo"``; any other + value raises ``ValueError``, and ``"off"`` returns :data:`EMPTY_UNIT` + without calling anything. ``read`` is the chunk's capped reader, or None + when there is none. ``listing_paths`` is the run's path listing as a set + built once per run (or None), and ``listing_complete`` says it was not + truncated. ``contract_paths`` and ``contract_priority`` are the run's + :func:`prxref.repo_contracts.select_contract_files` and + :func:`prxref.repo_contracts.literal_contract_paths` results. + ``exclude(path)`` true marks a path repository context must never read or + show, normally :func:`prxref.repo_context.exclude_predicate`; an + ``exclude`` that raises counts as true. + + Every read goes through a guard that returns None for an excluded path + without calling ``read``. The sources, called in this order: + + 1. :func:`prxref.repo_crosschunk.diff_definitions` with the guarded + reader, at both levels; + 2. at ``"repo"`` with a reader, :func:`prxref.repo_contracts.contract_entries`; + 3. at ``"repo"`` with a reader, the resolver. For each chunk file that is + not removed and whose language has definition regexes, the names its + added lines reference, less every symbol a definition entry at a + non-excluded path has already found, go to + :func:`prxref.repo_resolve.resolve_candidates` with the file's text + (its added lines when the read gives None). The candidates' distinct + paths are walked in first-appearance order; a path that is a PR diff + file or excluded is skipped, and so is one whose candidate names are + all found already, with no read. Each other path is read once, and + :func:`prxref.repo_context.find_definitions` over it gives a + ``definition`` entry whose reason is that name's candidate reason + there; its symbol then counts as found. + + ``"repo"`` with ``read`` None is exactly ``"diff"`` with ``read`` None. + + Entries at excluded paths are dropped (this covers the hunk-derived + entries of an excluded diff file), the rest are ordered by ``(REASONS + rank, path, line)`` and deduplicated on ``(path, line)``, the first kept. + They are admitted in that order while the sum of their rendered lengths + stays within ``max_chars``; admission stops at the first entry that does + not fit, even when a later one would. The omitted line does not count + toward ``max_chars``. + """ + if mode not in MODES: + raise ValueError(f"repository context mode must be one of {MODES}, got {mode!r}") + if mode == "off": + return EMPTY_UNIT + guarded = _guard(read, exclude) + entries = list(diff_definitions(chunk, all_files, guarded)) + if mode == "repo" and guarded is not None: + found = { + entry.symbol + for entry in entries + if entry.kind == "definition" and not _excluded(exclude, entry.path) + } + entries.extend( + contract_entries(chunk, contract_paths=contract_paths, read=guarded, priority=contract_priority) + ) + entries.extend( + _resolver_entries( + chunk, + all_files, + guarded, + listing_paths=listing_paths, + listing_complete=listing_complete, + exclude=exclude, + found=found, + ) + ) + return _admit(_merge(entries, exclude), max_chars) diff --git a/tests/test_repo_context_unit.py b/tests/test_repo_context_unit.py new file mode 100644 index 0000000..59c2ee9 --- /dev/null +++ b/tests/test_repo_context_unit.py @@ -0,0 +1,820 @@ +"""Tests for the per-chunk repository-context unit (issue #17, task T7). + +Covers :func:`prxref.repo_unit.build_unit_context` (sources, order, budget, +the omitted line, exclusion and the guard), :func:`prxref.repo_context.exclude_predicate` +with :data:`prxref.repo_context.EXCLUDE_FLOOR`, and the two optional arguments +of :func:`prxref.chunk_context.render_context_blocks`. + +The fixture half runs the real parser, chunker, sources and capped reader over +``tests/fixtures/issue17``. The budget and order half replaces the three +sources bound in ``prxref.repo_unit`` with canned entries. The rest builds +duck-typed ``triage.FileDiff`` stand-ins so each rule is pinned on its own. +No test touches the network. +""" +from __future__ import annotations + +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from prxref import config, repo_unit +from prxref.chunk_context import ( + CONTRACT_HEADER, + DEFINITIONS_HEADER, + DEPENDENCY_HEADER, + render_context_blocks, +) +from prxref.forges.repo_dir import RepoDir +from prxref.repo_context import EXCLUDE_FLOOR, REASONS, ContextEntry, exclude_predicate +from prxref.repo_contracts import literal_contract_paths, select_contract_files +from prxref.repo_reader import RepoReader, repo_dir_reader +from prxref.repo_resolve import Candidate +from prxref.repo_unit import EMPTY_UNIT, MODES, UnitContext, build_unit_context +from prxref.rules import match_globs +from prxref.triage import build_chunks, parse_unified_diff + +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" +REPO = FIXTURE / "repo" +TRANSPORT_CONFIG = "src/main/java/com/acme/connectors/TransportConfig.java" +CONNECTOR_SERVICE = "src/main/java/com/acme/connectors/ConnectorService.java" +MIGRATION = "db/changelog/003-idempotency-unique.sql" +IDEMPOTENCY_TABLE = "db/changelog/002-create-idempotency-keys.sql" +SPEC = "api/openapi/connectors.yaml" +EXCLUSIVITY = "exactly one of url or legacyUrl must be set" +MAX_CHARS = 12000 + + +def _omitted(count: int) -> str: + return f"\N{HORIZONTAL ELLIPSIS} {count} more context entries omitted" + + +class _Recording: + """Wraps a ``read(path)`` callable and records every path asked for, in call order.""" + + def __init__(self, read): + self.read = read + self.calls: list[str] = [] + + def __call__(self, path: str) -> str | None: + self.calls.append(path) + return self.read(path) + + def distinct(self) -> list[str]: + return list(dict.fromkeys(self.calls)) + + +def _texts(texts: dict[str, str]) -> _Recording: + return _Recording(texts.get) + + +def _repo_text(path: str) -> str | None: + target = REPO / path + return target.read_text(encoding="utf-8") if target.is_file() else None + + +def _keys(entries) -> list[tuple[str, int, str, str, str]]: + return [(e.path, e.line, e.symbol, e.kind, e.reason) for e in entries] + + +def _file(path: str, *hunks: tuple[int, list[str]], status: str = "modified") -> SimpleNamespace: + """A FileDiff stand-in; each hunk is ``(new_start, lines)`` with a ``+``/``-``/`` `` prefix per line.""" + built = [] + for new_start, body in hunks: + number = new_start + lines = [] + for raw in body: + kind, text = raw[0], raw[1:] + if kind == "-": + lines.append(SimpleNamespace(kind="-", text=text, new_line=None)) + else: + lines.append(SimpleNamespace(kind=kind, text=text, new_line=number)) + number += 1 + built.append(SimpleNamespace(lines=lines)) + return SimpleNamespace(path=path, status=status, hunks=built) + + +def _run() -> SimpleNamespace: + """The once-per-run inputs, built the way the orchestrator builds them.""" + files = parse_unified_diff((FIXTURE / "pr.diff").read_text(encoding="utf-8")) + chunks = build_chunks(files, max_files_per_chunk=1) + exclude = exclude_predicate() + reader = repo_dir_reader(RepoDir(REPO), exclude=exclude) + listing = reader.listing() + globs = config._DEFAULTS["context_contract_globs"] + return SimpleNamespace( + files=files, + chunks=chunks, + reader=reader, + exclude=exclude, + listing=listing, + listing_paths=frozenset(listing.paths), + listing_complete=listing.complete, + contract_paths=select_contract_files( + globs, listing=listing.paths, diff_paths=[f.path for f in files if f.status != "removed"] + ), + contract_priority=literal_contract_paths(globs), + ) + + +def _chunk_holding(chunks, path: str): + return next(chunk for chunk in chunks if any(f.path == path for f in chunk)) + + +def _unit(run, path: str, *, mode: str, read, **overrides) -> UnitContext: + kwargs = dict( + mode=mode, + read=read, + max_chars=MAX_CHARS, + listing_paths=run.listing_paths, + listing_complete=run.listing_complete, + contract_paths=run.contract_paths, + contract_priority=run.contract_priority, + exclude=run.exclude, + ) + kwargs.update(overrides) + return build_unit_context(_chunk_holding(run.chunks, path), run.files, **kwargs) + + +def _rendered(unit: UnitContext) -> str: + return render_context_blocks( + [], [], extra_def_lines=unit.definition_lines, contract_lines=unit.contract_lines + ) + + +class TestFixtureRepoMode: + def test_run_inputs_match_what_t6_observed(self): + run = _run() + assert run.listing_complete is True + assert run.contract_paths == [ + SPEC, + "db/changelog/001-create-connectors.sql", + IDEMPOTENCY_TABLE, + MIGRATION, + ] + assert run.contract_priority == [] + assert [[f.path for f in chunk] for chunk in run.chunks] == [ + [CONNECTOR_SERVICE], + [TRANSPORT_CONFIG], + [MIGRATION], + ] + + def test_connector_service_chunk_gets_cross_chunk_then_contract_entries(self): + run = _run() + unit = _unit(run, CONNECTOR_SERVICE, mode="repo", read=run.reader.chunk_reader()) + + assert _keys(unit.entries) == [ + (TRANSPORT_CONFIG, 9, "TransportConfig", "definition", "cross-chunk"), + (TRANSPORT_CONFIG, 11, "TransportConfig", "definition", "cross-chunk"), + (SPEC, 6, "/connectors/{connectorId}/transports", "contract", "contract"), + (SPEC, 31, "TransportConfig", "contract", "contract"), + (SPEC, 41, "CreateTransportRequest", "contract", "contract"), + ] + assert unit.omitted == 0 + assert unit.definition_lines == tuple(e.rendered() for e in unit.entries[:2]) + assert unit.contract_lines == tuple(e.rendered() for e in unit.entries[2:]) + assert "mutually exclusive" in unit.entries[3].text + + def test_connector_service_prompt_holds_both_blocks(self): + run = _run() + unit = _unit(run, CONNECTOR_SERVICE, mode="repo", read=run.reader.chunk_reader()) + text = _rendered(unit) + + assert text.startswith(DEFINITIONS_HEADER + "\n\n") + assert text.count(CONTRACT_HEADER) == 1 + definitions, contracts = text.split(CONTRACT_HEADER) + assert EXCLUSIVITY in definitions + assert "if (hasUrl == hasLegacyUrl) {" in definitions + assert f"{SPEC}:31: TransportConfig:" in contracts + assert "mutually exclusive" in contracts + + def test_migration_chunk_gets_the_table_and_the_schema(self): + run = _run() + unit = _unit(run, MIGRATION, mode="repo", read=run.reader.chunk_reader()) + + assert _keys(unit.entries) == [ + (SPEC, 50, "IdempotencyKey", "contract", "contract"), + (IDEMPOTENCY_TABLE, 4, "idempotency_keys", "contract", "contract"), + ] + assert unit.entries[1].text.startswith("CREATE TABLE idempotency_keys (") + assert unit.definition_lines == () + assert unit.contract_lines == tuple(e.rendered() for e in unit.entries) + + def test_transport_config_chunk_gets_its_schema_and_nothing_from_the_diff(self): + run = _run() + repo = _unit(run, TRANSPORT_CONFIG, mode="repo", read=run.reader.chunk_reader()) + diff = _unit(run, TRANSPORT_CONFIG, mode="diff", read=run.reader.chunk_reader()) + + assert diff.entries == () + assert _keys(repo.entries) == [(SPEC, 31, "TransportConfig", "contract", "contract")] + + def test_connector_service_reads_and_zero_resolver_candidate_reads(self): + run = _run() + read = _Recording(run.reader.chunk_reader()) + _unit(run, CONNECTOR_SERVICE, mode="repo", read=read) + + assert read.distinct() == [TRANSPORT_CONFIG, CONNECTOR_SERVICE, SPEC] + assert read.calls == [ + TRANSPORT_CONFIG, + CONNECTOR_SERVICE, + CONNECTOR_SERVICE, + SPEC, + CONNECTOR_SERVICE, + ] + assert run.reader.stats()["reads"] == 3 + + def test_the_complete_listing_is_what_keeps_the_resolver_from_reading(self): + run = _run() + read = _Recording(run.reader.chunk_reader()) + unit = _unit(run, CONNECTOR_SERVICE, mode="repo", read=read, listing_complete=False) + + convention = "src/main/java/com/acme/connectors/" + assert read.distinct() == [ + TRANSPORT_CONFIG, + CONNECTOR_SERVICE, + SPEC, + convention + "Tenant.java", + convention + "Id.java", + convention + "CreateTransportRequest.java", + ] + assert [e.reason for e in unit.entries] == ["cross-chunk"] * 2 + ["contract"] * 3 + + +class TestFixtureDiffModeAndNoReader: + def test_diff_mode_gives_only_the_cross_chunk_entries_and_reads_no_contract(self): + run = _run() + read = _Recording(run.reader.chunk_reader()) + unit = _unit(run, CONNECTOR_SERVICE, mode="diff", read=read) + + assert _keys(unit.entries) == [ + (TRANSPORT_CONFIG, 9, "TransportConfig", "definition", "cross-chunk"), + (TRANSPORT_CONFIG, 11, "TransportConfig", "definition", "cross-chunk"), + ] + assert unit.contract_lines == () + assert read.calls == [TRANSPORT_CONFIG, CONNECTOR_SERVICE] + assert SPEC not in read.calls + + def test_diff_without_a_reader_gives_the_hunk_only_entries(self): + run = _run() + unit = _unit(run, CONNECTOR_SERVICE, mode="diff", read=None) + + assert [(e.path, e.line, e.reason) for e in unit.entries] == [ + (TRANSPORT_CONFIG, 9, "cross-chunk"), + (TRANSPORT_CONFIG, 11, "cross-chunk"), + ] + assert EXCLUSIVITY in unit.entries[1].text + assert run.reader.stats()["reads"] == 0 + + @pytest.mark.parametrize("path", [CONNECTOR_SERVICE, TRANSPORT_CONFIG, MIGRATION]) + def test_repo_without_a_reader_equals_diff_without_a_reader(self, path): + run = _run() + repo = _unit(run, path, mode="repo", read=None) + diff = _unit(run, path, mode="diff", read=None) + assert repo == diff + + def test_off_is_the_empty_unit_with_zero_reads(self): + run = _run() + read = _Recording(run.reader.chunk_reader()) + unit = _unit(run, CONNECTOR_SERVICE, mode="off", read=read) + + assert unit is EMPTY_UNIT + assert read.calls == [] + assert run.reader.stats()["reads"] == 0 + + def test_the_same_inputs_give_the_same_unit_and_the_same_reads(self): + first_run, second_run = _run(), _run() + first = _Recording(_repo_text) + second = _Recording(_repo_text) + a = _unit(first_run, CONNECTOR_SERVICE, mode="repo", read=first) + b = _unit(second_run, CONNECTOR_SERVICE, mode="repo", read=second) + assert a == b + assert first.calls == second.calls + + +class TestReadCap: + def test_the_real_chunk_cap_degrades_to_fewer_entries(self): + run = _run() + capped = RepoReader(RepoDir(REPO).read, None, kind="repo-dir", exclude=run.exclude, chunk_cap=2) + unit = _unit(run, CONNECTOR_SERVICE, mode="repo", read=capped.chunk_reader()) + + assert [e.reason for e in unit.entries] == ["cross-chunk", "cross-chunk"] + assert capped.stats()["read_cap_hit"] is True + assert capped.stats()["reads"] == 2 + + @pytest.mark.parametrize( + ("allowed", "count"), + [(0, 2), (1, 2), (2, 2), (3, 2), (4, 5), (5, 5)], + ) + def test_a_reader_that_goes_dry_raises_nothing(self, allowed, count): + run = _run() + calls: list[str] = [] + + def read(path: str) -> str | None: + calls.append(path) + return _repo_text(path) if len(calls) <= allowed else None + + unit = _unit(run, CONNECTOR_SERVICE, mode="repo", read=read) + assert len(unit.entries) == count + assert unit.omitted == 0 + + +@pytest.fixture +def canned(monkeypatch): + """Replace the three sources ``repo_unit`` binds with canned entry lists.""" + sources: dict[str, list[ContextEntry]] = {"diff": [], "contract": [], "resolver": []} + called: list[str] = [] + + def diff(chunk, all_files, read): + called.append("diff") + return list(sources["diff"]) + + def contracts(chunk, *, contract_paths, read, priority=()): + called.append("contract") + return list(sources["contract"]) + + def resolver(chunk, all_files, read, **kwargs): + called.append("resolver") + return list(sources["resolver"]) + + monkeypatch.setattr(repo_unit, "diff_definitions", diff) + monkeypatch.setattr(repo_unit, "contract_entries", contracts) + monkeypatch.setattr(repo_unit, "_resolver_entries", resolver) + return SimpleNamespace(sources=sources, called=called) + + +def _entry(path: str, line: int, reason: str, *, kind: str = "definition", text: str = "t") -> ContextEntry: + return ContextEntry(path, line, "S", kind, reason, text) + + +def _build(max_chars: int = MAX_CHARS, **kwargs) -> UnitContext: + kwargs.setdefault("mode", "repo") + kwargs.setdefault("read", lambda path: None) + return build_unit_context([], [], max_chars=max_chars, **kwargs) + + +class TestBudget: + def test_admits_exactly_the_prefix_that_fits_and_appends_the_omitted_line(self, canned): + a = _entry("a.java", 1, "cross-chunk", text="a" * 20) + b = _entry("b.java", 1, "cross-chunk", text="b" * 20) + c = _entry("c.java", 1, "cross-chunk", text="c" * 20) + canned.sources["diff"] = [a, b, c] + size = len(a.rendered()) + + unit = _build(max_chars=2 * size) + + assert unit.entries == (a, b) + assert unit.omitted == 1 + assert unit.definition_lines == (a.rendered(), b.rendered(), _omitted(1)) + assert unit.contract_lines == () + + def test_one_char_short_admits_one_fewer(self, canned): + a = _entry("a.java", 1, "cross-chunk", text="a" * 20) + b = _entry("b.java", 1, "cross-chunk", text="b" * 20) + canned.sources["diff"] = [a, b] + + unit = _build(max_chars=len(a.rendered()) + len(b.rendered()) - 1) + + assert unit.entries == (a,) + assert unit.definition_lines == (a.rendered(), _omitted(1)) + + def test_a_later_smaller_entry_is_not_admitted_after_a_non_fit(self, canned): + small = _entry("a.java", 1, "cross-chunk", text="s") + big = _entry("b.java", 1, "cross-chunk", text="b" * 200) + later = _entry("c.java", 1, "cross-chunk", text="s") + canned.sources["diff"] = [small, big, later] + + unit = _build(max_chars=len(small.rendered()) + len(later.rendered()) + 10) + + assert unit.entries == (small,) + assert unit.omitted == 2 + assert later.rendered() not in unit.definition_lines + assert unit.definition_lines == (small.rendered(), _omitted(2)) + + def test_the_omitted_line_goes_to_the_contract_block_when_a_contract_is_first_left_out(self, canned): + definition = _entry("a.java", 1, "cross-chunk", text="d") + contract = _entry("spec.yaml", 5, "contract", kind="contract", text="c" * 200) + tail = _entry("z.java", 1, "diff-file", text="d") + canned.sources["diff"] = [definition, tail] + canned.sources["contract"] = [contract] + + unit = _build(max_chars=len(definition.rendered()) + 5) + + assert unit.entries == (definition,) + assert unit.definition_lines == (definition.rendered(),) + assert unit.contract_lines == (_omitted(2),) + + def test_the_omitted_line_goes_to_the_definition_block_when_a_definition_is_first_left_out(self, canned): + contract = _entry("spec.yaml", 5, "contract", kind="contract", text="c") + definition = _entry("z.java", 1, "diff-file", text="d" * 200) + canned.sources["contract"] = [contract] + canned.sources["diff"] = [definition] + + unit = _build(max_chars=len(contract.rendered())) + + assert unit.contract_lines == (contract.rendered(),) + assert unit.definition_lines == (_omitted(1),) + + def test_zero_budget_admits_nothing(self, canned): + canned.sources["diff"] = [_entry("a.java", 1, "cross-chunk")] + unit = _build(max_chars=0) + assert unit.entries == () + assert unit.definition_lines == (_omitted(1),) + + def test_nothing_to_admit_is_an_empty_unit(self, canned): + assert _build() == EMPTY_UNIT + + def test_record_lists_the_admitted_entries_and_the_omitted_count(self, canned): + a = _entry("a.java", 1, "cross-chunk", text="a") + b = _entry("b.java", 1, "cross-chunk", text="b" * 50) + canned.sources["diff"] = [a, b] + + unit = _build(max_chars=len(a.rendered())) + + assert unit.record() == {"entries": [a.record()], "omitted": 1} + assert EMPTY_UNIT.record() == {"entries": [], "omitted": 0} + + +class TestOrder: + def test_reason_rank_beats_path(self, canned): + canned.sources["diff"] = [_entry("a/x.java", 1, "diff-file"), _entry("z/y.java", 1, "cross-chunk")] + canned.sources["contract"] = [_entry("y/spec.yaml", 1, "contract", kind="contract")] + canned.sources["resolver"] = [ + _entry("a/n.java", 1, "name-search"), + _entry("b/p.java", 1, "path-convention"), + _entry("c/i.java", 1, "import"), + ] + + unit = _build() + + assert [e.reason for e in unit.entries] == list(REASONS) + assert [e.path for e in unit.entries] == [ + "z/y.java", "y/spec.yaml", "a/x.java", "c/i.java", "b/p.java", "a/n.java", + ] + + def test_within_a_reason_path_then_line(self, canned): + canned.sources["resolver"] = [ + _entry("c/i.java", 9, "import"), + _entry("c/i.java", 2, "import"), + _entry("b/i.java", 5, "import"), + ] + unit = _build() + assert [(e.path, e.line) for e in unit.entries] == [("b/i.java", 5), ("c/i.java", 2), ("c/i.java", 9)] + + def test_path_line_dedup_keeps_the_higher_rank(self, canned): + lower = _entry("x/A.java", 5, "diff-file", text="from the diff") + higher = _entry("x/A.java", 5, "contract", kind="contract", text="from the contract") + canned.sources["diff"] = [lower] + canned.sources["contract"] = [higher] + + unit = _build() + + assert unit.entries == (higher,) + assert unit.definition_lines == () + + def test_sources_are_called_in_rank_order(self, canned): + _build() + assert canned.called == ["diff", "contract", "resolver"] + + def test_diff_mode_calls_only_the_diff_source(self, canned): + canned.sources["contract"] = [_entry("spec.yaml", 1, "contract", kind="contract")] + canned.sources["resolver"] = [_entry("a.java", 1, "import")] + unit = _build(mode="diff") + assert canned.called == ["diff"] + assert unit == EMPTY_UNIT + + def test_repo_without_a_reader_calls_only_the_diff_source(self, canned): + _build(read=None) + assert canned.called == ["diff"] + + def test_off_calls_no_source(self, canned): + canned.sources["diff"] = [_entry("a.java", 1, "cross-chunk")] + assert _build(mode="off") is EMPTY_UNIT + assert canned.called == [] + + @pytest.mark.parametrize("mode", ["", "on", "Repo", "full", None]) + def test_an_unknown_mode_raises(self, canned, mode): + with pytest.raises(ValueError, match="mode"): + _build(mode=mode) + + def test_modes(self): + assert MODES == ("off", "diff", "repo") + assert EMPTY_UNIT == UnitContext((), (), (), 0) + + +SERVICE = "src/main/java/com/acme/app/Service.java" +SERVICE_TEXT = "package com.acme.app;\n\npublic class Service {\n Widget w = new Widget();\n}\n" +WIDGET = "src/main/java/com/acme/app/Widget.java" +WIDGET_TEXT = "package com.acme.app;\n\npublic class Widget {\n int size;\n}\n" + + +def _service_chunk(status: str = "modified") -> SimpleNamespace: + return _file(SERVICE, (3, [" public class Service {", "+ Widget w = new Widget();", " }"]), status=status) + + +class TestResolver: + def test_a_convention_candidate_is_read_and_admitted(self): + chunk = [_service_chunk()] + read = _texts({SERVICE: SERVICE_TEXT, WIDGET: WIDGET_TEXT}) + + unit = build_unit_context( + chunk, chunk, mode="repo", read=read, max_chars=MAX_CHARS, + listing_paths=frozenset({SERVICE, WIDGET}), listing_complete=True, + ) + + assert _keys(unit.entries) == [(WIDGET, 3, "Widget", "definition", "path-convention")] + assert unit.entries[0].text.startswith("public class Widget {") + assert read.calls == [SERVICE, SERVICE, WIDGET] + + def test_a_name_the_diff_already_resolved_causes_no_resolver_read(self): + legacy = "src/main/java/com/acme/legacy/Widget.java" + other = _file(legacy, (7, [" public class Widget {", "+ int legacySize;", " }"])) + chunk = [_service_chunk()] + read = _texts({SERVICE: SERVICE_TEXT, WIDGET: WIDGET_TEXT}) + + unit = build_unit_context( + chunk, [*chunk, other], mode="repo", read=read, max_chars=MAX_CHARS, + listing_paths=frozenset({SERVICE, WIDGET, legacy}), listing_complete=True, + ) + + assert WIDGET not in read.calls + assert read.calls == [SERVICE, legacy] + assert [(e.path, e.reason) for e in unit.entries] == [(legacy, "cross-chunk"), (legacy, "cross-chunk")] + + def test_a_name_the_resolver_found_is_not_searched_again_for_the_next_file(self): + second = "src/main/java/com/acme/app/Other.java" + chunk = [ + _service_chunk(), + _file(second, (3, [" public class Other {", "+ Widget w;", " }"])), + ] + read = _texts({SERVICE: SERVICE_TEXT, WIDGET: WIDGET_TEXT}) + + unit = build_unit_context( + chunk, chunk, mode="repo", read=read, max_chars=MAX_CHARS, + listing_paths=frozenset({SERVICE, WIDGET, second}), listing_complete=True, + ) + + assert read.calls.count(WIDGET) == 1 + assert read.calls == [SERVICE, second, SERVICE, WIDGET] + assert [e.path for e in unit.entries] == [WIDGET] + + def test_a_removed_chunk_file_is_not_resolved(self): + read = _texts({SERVICE: SERVICE_TEXT, WIDGET: WIDGET_TEXT}) + removed = [_service_chunk(status="removed")] + kept = [_service_chunk()] + kwargs = dict( + mode="repo", max_chars=MAX_CHARS, listing_paths=frozenset({SERVICE, WIDGET}), listing_complete=True + ) + + assert build_unit_context(removed, removed, read=read, **kwargs).entries == () + assert WIDGET not in read.calls + assert len(build_unit_context(kept, kept, read=read, **kwargs).entries) == 1 + + def test_a_candidate_that_is_a_diff_file_is_left_to_the_diff_source(self): + widget_diff = _file(WIDGET, (4, ["+ int size;"])) + chunk = [_service_chunk()] + read = _texts({SERVICE: SERVICE_TEXT}) + + build_unit_context( + chunk, [*chunk, widget_diff], mode="repo", read=read, max_chars=MAX_CHARS, + listing_paths=frozenset({SERVICE, WIDGET}), listing_complete=True, + ) + + assert read.calls.count(WIDGET) == 1 + + def test_a_path_whose_names_are_all_found_is_not_read(self): + service = "app/service.py" + chunk = [_file(service, (1, ["+w = Widget()"]))] + first, second = "lib/Widget.py", "pkg/widget.py" + listing = frozenset({service, first, second}) + widget = "class Widget:\n pass\n" + hit = _texts({service: "w = Widget()\n", first: widget, second: widget}) + miss = _texts({service: "w = Widget()\n", first: "x = 1\n", second: widget}) + + found = build_unit_context(chunk, chunk, mode="repo", read=hit, max_chars=MAX_CHARS, listing_paths=listing) + control = build_unit_context(chunk, chunk, mode="repo", read=miss, max_chars=MAX_CHARS, listing_paths=listing) + + assert hit.calls == [service, first] + assert [(e.path, e.reason) for e in found.entries] == [(first, "name-search")] + assert miss.calls == [service, first, second] + assert [(e.path, e.reason) for e in control.entries] == [(second, "name-search")] + + def test_an_import_candidate_carries_the_import_reason(self): + service = "app/service.py" + chunk = [_file(service, (1, ["+from pkg.models import Widget", "+w = Widget()"]))] + read = _texts({"pkg/models.py": "class Widget:\n size = 0\n"}) + + unit = build_unit_context(chunk, chunk, mode="repo", read=read, max_chars=MAX_CHARS) + + assert _keys(unit.entries) == [("pkg/models.py", 1, "Widget", "definition", "import")] + assert read.calls == [service, "pkg/models.py"] + + +class TestExclusion: + def test_a_resolver_candidate_at_a_label_file_is_never_read(self, monkeypatch): + label = "evals/cases.json" + library = "lib/Widget.java" + monkeypatch.setattr( + repo_unit, + "resolve_candidates", + lambda path, text, names, **kwargs: [ + Candidate("Widget", label, "import"), + Candidate("Widget", library, "import"), + ], + ) + chunk = [_service_chunk()] + guarded = _texts({SERVICE: SERVICE_TEXT, library: "public class Widget {\n}\n"}) + unguarded = _texts({SERVICE: SERVICE_TEXT, library: "public class Widget {\n}\n"}) + + unit = build_unit_context( + chunk, chunk, mode="repo", read=guarded, max_chars=MAX_CHARS, exclude=exclude_predicate() + ) + build_unit_context(chunk, chunk, mode="repo", read=unguarded, max_chars=MAX_CHARS, exclude=None) + + assert label not in guarded.calls + assert library in guarded.calls + assert [(e.path, e.reason) for e in unit.entries] == [(library, "import")] + assert label in unguarded.calls + + def test_a_real_resolver_candidate_under_an_extra_glob_is_never_read(self): + service = "app/service.py" + chunk = [_file(service, (1, ["+from vault.config import Settings", "+s = Settings()"]))] + texts = { + service: "from vault.config import Settings\ns = Settings()\n", + "vault/config.py": "class Settings:\n pass\n", + } + guarded, control = _texts(texts), _texts(texts) + + excluded = build_unit_context( + chunk, chunk, mode="repo", read=guarded, max_chars=MAX_CHARS, exclude=exclude_predicate(["**/vault/**"]) + ) + admitted = build_unit_context( + chunk, chunk, mode="repo", read=control, max_chars=MAX_CHARS, exclude=exclude_predicate() + ) + + assert not any("vault/" in path for path in guarded.calls) + assert excluded.entries == () + assert "vault/config.py" in control.calls + assert [e.path for e in admitted.entries] == ["vault/config.py"] + + @pytest.mark.parametrize("with_reader", [True, False], ids=["reader", "no-reader"]) + @pytest.mark.parametrize("mode", ["diff", "repo"]) + def test_an_excluded_diff_files_hunk_entries_are_dropped(self, mode, with_reader): + generated = "generated/Widget.java" + other = _file(generated, (3, [" public class Widget {", "+ int size;", " }"])) + chunk = [_service_chunk()] + read = _texts({SERVICE: SERVICE_TEXT}) if with_reader else None + + dropped = build_unit_context( + chunk, [*chunk, other], mode=mode, read=read, max_chars=MAX_CHARS, + exclude=exclude_predicate(["generated/**"]), + ) + kept = build_unit_context( + chunk, [*chunk, other], mode=mode, read=None, max_chars=MAX_CHARS, exclude=exclude_predicate() + ) + + assert dropped.entries == () + assert read is None or generated not in read.calls + assert [e.path for e in kept.entries] == [generated, generated] + + def test_the_guard_refuses_an_excluded_path_without_calling_read(self, monkeypatch): + seen: dict[str, str | None] = {} + + def diff(chunk, all_files, read): + for path in (".env", "keys/deploy.pem", "src/A.java"): + seen[path] = read(path) + return [] + + monkeypatch.setattr(repo_unit, "diff_definitions", diff) + read = _texts({".env": "SECRET=1", "keys/deploy.pem": "-----", "src/A.java": "class A {}"}) + + build_unit_context([], [], mode="diff", read=read, max_chars=MAX_CHARS, exclude=exclude_predicate()) + + assert read.calls == ["src/A.java"] + assert seen == {".env": None, "keys/deploy.pem": None, "src/A.java": "class A {}"} + + def test_an_exclude_that_raises_fails_closed(self): + run = _run() + + def broken(path: str) -> bool: + raise RuntimeError("glob engine down") + + read = _Recording(_repo_text) + unit = _unit(run, CONNECTOR_SERVICE, mode="repo", read=read, exclude=broken) + + assert read.calls == [] + assert unit.entries == () + assert unit.omitted == 0 + + +class TestExcludePredicate: + def test_the_floor_is_exactly_the_decided_set(self): + assert EXCLUDE_FLOOR == ( + "**/expected.json", + "**/cases.json", + "**/case.json", + "**/prxref-eval/**", + "**/.env*", + "**/*.pem", + "**/*.key", + ) + + FLOOR_SAMPLES = { + "**/expected.json": "tests/evals/cases/rate-limit/expected.json", + "**/cases.json": "cases.json", + "**/case.json": "evals/one/case.json", + "**/prxref-eval/**": "out/prxref-eval/run-1/summary.json", + "**/.env*": "deploy/.env.production", + "**/*.pem": "certs/server.pem", + "**/*.key": "keys/private.key", + } + + def test_every_floor_glob_has_a_sample(self): + assert tuple(self.FLOOR_SAMPLES) == EXCLUDE_FLOOR + + @pytest.mark.parametrize(("glob", "sample"), list(FLOOR_SAMPLES.items())) + def test_every_floor_glob_excludes_its_sample(self, glob, sample): + assert match_globs(sample, [glob]) + assert exclude_predicate()(sample) is True + assert exclude_predicate(["docs/**"])(sample) is True + + @pytest.mark.parametrize("path", ["src/main/java/A.java", "api/openapi/connectors.yaml", "docs/cases.md"]) + def test_ordinary_paths_are_not_excluded(self, path): + assert exclude_predicate()(path) is False + + def test_a_negation_in_the_extra_globs_cannot_re_admit_a_floor_path(self): + assert exclude_predicate(["!**/cases.json"])("a/cases.json") is True + + def test_extra_globs_add_to_the_floor(self): + assert exclude_predicate(["docs/**"])("docs/x.md") is True + assert exclude_predicate()("docs/x.md") is False + + def test_a_negation_in_the_extra_globs_still_applies_to_the_extra_positives(self): + predicate = exclude_predicate(["docs/**", "!docs/keep.md"]) + assert predicate("docs/x.md") is True + assert predicate("docs/keep.md") is False + + def test_an_iterator_of_extra_globs_is_read_once(self): + predicate = exclude_predicate(iter(["docs/**"])) + assert predicate("docs/a.md") is True + assert predicate("docs/b.md") is True + + +def _oracle(dep_lines, def_lines): + """``render_context_blocks`` exactly as it was before repository context (the D1 oracle).""" + blocks: list[str] = [] + if dep_lines: + blocks.append(DEPENDENCY_HEADER + "\n\n" + "\n".join(dep_lines)) + if def_lines: + blocks.append(DEFINITIONS_HEADER + "\n\n" + "\n".join(def_lines)) + return "\n\n".join(blocks) + + +_DEPS = {"none": [], "one": ["effect@4.0.0"], "two": ["a@1.0.0", "b@2.0.0"], "tuple": ("c@3.0.0",)} +_DEFS = { + "none": [], + "one": ["a.ts:1: const X = 1;"], + "capped": ["a.ts:1: x", "b.ts:2: y", "\N{HORIZONTAL ELLIPSIS} 1 more definitions omitted"], + "tuple": ("m.py:3: def f():",), +} + + +class TestRenderContextBlocks: + @pytest.mark.parametrize("deps", list(_DEPS), ids=list(_DEPS)) + @pytest.mark.parametrize("defs", list(_DEFS), ids=list(_DEFS)) + def test_the_defaults_are_byte_identical_to_the_old_rendering(self, deps, defs): + dep_lines, def_lines = _DEPS[deps], _DEFS[defs] + expected = _oracle(dep_lines, def_lines) + assert render_context_blocks(dep_lines, def_lines) == expected + assert render_context_blocks(dep_lines, def_lines, (), ()) == expected + assert render_context_blocks(dep_lines, def_lines, extra_def_lines=[], contract_lines=[]) == expected + + def test_the_contract_header(self): + assert CONTRACT_HEADER == "### Contract excerpts" + + def test_extra_definitions_follow_under_one_header(self): + out = render_context_blocks(["effect@4.0.0"], ["a.ts:1: x"], extra_def_lines=["B.java:9: record B() {"]) + assert out == ( + DEPENDENCY_HEADER + "\n\neffect@4.0.0\n\n" + + DEFINITIONS_HEADER + "\n\na.ts:1: x\nB.java:9: record B() {" + ) + assert out.count(DEFINITIONS_HEADER) == 1 + + def test_extra_definitions_alone_bring_the_header(self): + out = render_context_blocks([], [], extra_def_lines=["B.java:9: record B() {"]) + assert out == DEFINITIONS_HEADER + "\n\nB.java:9: record B() {" + + def test_contracts_come_last(self): + out = render_context_blocks( + ["effect@4.0.0"], + ["a.ts:1: x"], + extra_def_lines=["B.java:9: y"], + contract_lines=["s.yaml:3: B:", "s.yaml:9: C:"], + ) + assert out == ( + DEPENDENCY_HEADER + "\n\neffect@4.0.0\n\n" + + DEFINITIONS_HEADER + "\n\na.ts:1: x\nB.java:9: y\n\n" + + CONTRACT_HEADER + "\n\ns.yaml:3: B:\ns.yaml:9: C:" + ) + + def test_contracts_alone(self): + assert render_context_blocks([], [], contract_lines=["s.yaml:3: B:"]) == CONTRACT_HEADER + "\n\ns.yaml:3: B:" From 311a5484c32dadee9341198153ab892101bd3cc3 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:36:04 -0700 Subject: [PATCH 16/24] feat: wire repository context into orchestrate_review MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit orchestrate_review takes five new keyword-only parameters after max_findings_per_rule: repo_context ("off" by default), repo_context_max_chars (12000), context_contract_globs, context_exclude_globs and repo_dir. An unknown level raises ValueError before any forge call. Off is byte-identical to a run without the kwargs; the run record gains a repo_context key that is None when off. On, the run builds one RepoReader (repo_dir first, else the forge at the head sha), takes the listing and the contract files once per run in repo mode, and builds each chunk's unit context inside its worker before the first attempt. Diff-file reads go through the shared uncapped read and every other path through the chunk's capped reader. A unit build that raises gives that chunk the empty unit and one WARNING. The timeout retry carries neither new block and marks the chunk's row retry_dropped. Repo mode with no reader or no listing logs one WARNING naming PRXREF_REPO_CONTEXT. The trace gains one chunk context event per chunk and one repo_context ok event per run. The CLI and JSON output are not wired yet (T14); tests/test_cli.py and tests/test_run_record.py carry NOT_WIRED_YET / NOT_IN_JSON_YET tripwires for that seat to remove. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/orchestrator.py | 325 +++++++++++- tests/test_cli.py | 12 +- tests/test_orchestrator.py | 8 +- tests/test_orchestrator_repo_context.py | 636 ++++++++++++++++++++++++ tests/test_orchestrator_rule_cap.py | 3 +- tests/test_run_record.py | 11 +- 6 files changed, 969 insertions(+), 26 deletions(-) create mode 100644 tests/test_orchestrator_repo_context.py diff --git a/src/prxref/orchestrator.py b/src/prxref/orchestrator.py index 2584825..3715562 100644 --- a/src/prxref/orchestrator.py +++ b/src/prxref/orchestrator.py @@ -125,7 +125,7 @@ keys that :func:`_run_record` stamps on every exit (``cost_usd``, ``cost_estimated``, ``review_rules``, ``ticket_context``, ``spec_grounding``, ``size_advisory``, ``prompt_templates``, - ``scoped_rules``, ``rule_counts``; ``replay`` on replays only, and + ``scoped_rules``, ``rule_counts``, ``repo_context``; ``replay`` on replays only, and ``cost_api_equivalent`` on claude-cli-priced runs only). 7. Verdict: ``"Error"`` when every CHUNK review failed (a sweep success on a dead worker pool cannot carry the run); ``"Request-Changes"`` @@ -166,13 +166,13 @@ import threading import time from collections import Counter -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from concurrent.futures import ThreadPoolExecutor -from dataclasses import replace +from dataclasses import dataclass, replace from typing import Any from urllib.parse import urlparse -from . import chunk_context, costs, heuristics, reviewer, specs, systemic +from . import chunk_context, costs, heuristics, repo_contracts, repo_reader, repo_unit, reviewer, specs, systemic from .forges.base import ( ATTRIBUTION_MARKER, Forge, @@ -181,6 +181,7 @@ PRRef, Thread, ) +from .forges.repo_dir import RepoDir from .llm import LLMClient from .markers import OUT_OF_TICKET_MARKER, SEVERITY_MARKERS, inline_header, marker_for from .prompt_templates import CONTEXT_MARKER, REVIEW_TEMPLATES, PromptTemplates, packaged_text, placeholders @@ -209,6 +210,7 @@ prompt_example_titles, rule_cap_counts, ) +from .repo_context import exclude_predicate from .reviewer import NO_PROMPT_CONTEXT, PromptContext, fill_template from .trace import Tracer, get_tracer from .triage import ( @@ -431,6 +433,11 @@ def orchestrate_review( max_warning_findings: int | None = None, max_outofscope_findings: int | None = None, max_findings_per_rule: int = 2, + repo_context: str = "off", + repo_context_max_chars: int = 12000, + context_contract_globs: Sequence[str] = (), + context_exclude_globs: Sequence[str] = (), + repo_dir: RepoDir | None = None, ) -> dict: """Run one full review pass over a PR and optionally post results. @@ -438,9 +445,9 @@ def orchestrate_review( chunks_reviewed, chunks_failed, elapsed_ms, input_tokens, output_tokens, posted, sampling, cost_usd, cost_estimated, review_rules, ticket_context, spec_grounding, size_advisory, prompt_templates, scoped_rules, - rule_counts}``, plus ``replay`` on a replay run only. + rule_counts, repo_context}``, plus ``replay`` on a replay run only. Every exit, error and empty-diff exits included, goes through - :func:`_run_record`, so the last nine keys are always present and are + :func:`_run_record`, so the last ten keys are always present and are ``None`` (``cost_usd``: ``0.0`` before any LLM request; ``cost_estimated``: ``False``) when their feature is off or the run never reached it. ``cost_usd`` is ``None`` when the cost is unknown, never ``0``. @@ -454,6 +461,8 @@ def orchestrate_review( chunking, or LLM — the run degrades to verdict ``"Error"`` with a posted notice when ``post`` is true. Degenerate arguments are part of that: a caller passing ``max_chunks=0`` gets an error run, not a ``ValueError``. + The one exception is an unknown ``repo_context`` level, which raises + ``ValueError`` before any forge call (see below). ``chunk_count`` counts the review units: ``len(chunks)`` plus one for the systemic sweep, which runs whenever at least one chunk exists (an @@ -655,7 +664,69 @@ def orchestrate_review( unlimited and reads no environment variable; ``0`` drops every finding of that severity. ``outofscope`` is the minor severity, not the ticket scope ``out``. + + ``repo_context`` is the repository-context level + (``PRXREF_REPO_CONTEXT``): ``"off"`` (the default), ``"diff"`` or + ``"repo"`` (:data:`prxref.repo_unit.MODES`). Any other value raises + ``ValueError`` before any forge call; config rejects it long before, so + this guards a library caller only. ``repo_context_max_chars`` + (``PRXREF_REPO_CONTEXT_MAX_CHARS``) is each chunk's budget for it, + ``context_contract_globs`` (``PRXREF_CONTEXT_CONTRACT_GLOBS``) selects + the contract files, where ``()`` means none (the built-in set is + config's default, which the CLI passes), and ``context_exclude_globs`` + (``PRXREF_CONTEXT_EXCLUDE_GLOBS``) adds to the exclude floor of + :func:`prxref.repo_context.exclude_predicate`: an excluded path is + never read, listed or shown. ``repo_dir`` is a + :class:`prxref.forges.repo_dir.RepoDir` to read the repository from in + place of the forge; this function does not validate it (``RepoDir`` + does, when it is built). Off, nothing new is built or called: the + prompts, posts, trace and logs are exactly a run without these + arguments, and the record's ``repo_context`` key is ``None``. + + On, the run has ONE reader (:class:`prxref.repo_reader.RepoReader`): + over ``repo_dir`` when given, else over the forge's optional + ``get_file_content`` and ``list_paths`` at ``pr.source_sha``, and none + when the forge cannot read or the PR has no head sha. At ``"repo"`` with + a reader, the path listing is taken once and the contract files are + selected once, before the chunk workers run. ``"repo"`` with no reader + (only diff-only entries from hunk lines are left), or with a reader but + no listing (no name search and no glob-matched contract files), logs one + WARNING naming ``PRXREF_REPO_CONTEXT``; ``"diff"`` with no reader + builds its entries from hunk lines and logs nothing. Each chunk's worker + builds its context once, before its first attempt, with + :func:`prxref.repo_unit.build_unit_context`. It reads a PR diff file + through the reader's shared, uncapped ``read`` and every other path + through a fresh ``chunk_reader()``, so the per-chunk read cap is spent + on the paths outside the diff alone, and the entries do not depend on + which chunk reads a shared diff file first. The context's definition + lines extend the definitions block and its contract lines form a + ``### Contract excerpts`` block, on the first attempt only: the timeout + retry carries neither. A build that raises gives that chunk no context + and one WARNING naming the chunk; the review goes on. The dependency and + same-file definition blocks keep their own reader in every mode, so a + diff file can be fetched once by each reader. + + The ``repo_context`` key of every exit, when on, is ``{"mode", + "max_chars", "contract_globs", "exclude_globs", "reader", "listing", + "reads", "read_cap_hit", "units"}``. Until the chunk workers finish, + ``reader`` and ``listing`` are ``None``, ``reads`` is 0, + ``read_cap_hit`` is false and ``units`` is ``None``. After them, + ``reader`` is the reader's ``kind`` (``"forge"`` or ``"repo-dir"``) or + ``None``; ``listing`` (``{"paths", "complete"}`` or ``None``), + ``reads`` and ``read_cap_hit`` come from one + :meth:`~prxref.repo_reader.RepoReader.stats` snapshot, so ``reads`` + counts repository-context fetches only; and ``units`` is ``{"chunks": + [{"entries", "omitted", "retry_dropped"}, ...]}``, one row per chunk in + chunk order, where ``retry_dropped`` is true when the timeout retry ran + (and so ran without the chunk's repository context). The sweep has no + row. One ``chunk context`` trace event per chunk (``index``, ``total``, + ``entries``, ``omitted``, ``chars``) and one ``repo_context ok`` event + per run carry the same figures, only when on. """ + if repo_context not in repo_unit.MODES: + raise ValueError( + f"repo_context must be one of {repo_unit.MODES}, got {repo_context!r}" + ) t0 = time.perf_counter() tracer = get_tracer(trace_file) sampling = _sampling(llm) @@ -675,6 +746,7 @@ def orchestrate_review( "prompt_templates": None, "scoped_rules": None, "rule_counts": None, + "repo_context": None, } # Resolved once, before the first exit, so every exit records them and the # empty-diff summary gets the ticket note. An inactive (empty) ticket is @@ -689,6 +761,18 @@ def orchestrate_review( if scoped_rules is not None: scoped_meta = {**scoped_rules.record(), "max_chars": scoped_rules_max_chars} run_inputs["scoped_rules"] = {**scoped_meta, "units": None} + if repo_context != "off": + run_inputs["repo_context"] = { + "mode": repo_context, + "max_chars": repo_context_max_chars, + "contract_globs": list(context_contract_globs), + "exclude_globs": list(context_exclude_globs), + "reader": None, + "listing": None, + "reads": 0, + "read_cap_hit": False, + "units": None, + } summary_template = prompts.override("summary") if prompts is not None else "" ticket_active = ticket is not None and bool(ticket.active) ticket_note = ticket.note() if ticket is not None else "" @@ -961,12 +1045,30 @@ def orchestrate_review( prompts, feature="finding grouping" if group_findings else "the per-rule cap", ) reader = _make_file_reader(forge, ref, pr) + repo_plan: _RepoPlan | None = None + unit_records: list[dict[str, Any] | None] | None = None + if repo_context != "off": + repo_plan = _plan_repo_context( + repo_context, forge, ref, pr, files, run_inputs["repo_context"], repo_dir=repo_dir, + ) + unit_records = [None] * len(chunks) results = _run_workers( llm, chunks, pr, max_tokens=max_tokens, max_workers=max_workers, context_lines=context_lines, tracer=tracer, reader=reader, all_files=files, trace_dir=trace_dir, prompt_context=prompt_context, scoped_blocks=scoped_blocks, + repo_plan=repo_plan, unit_records=unit_records, ) + if repo_plan is not None and unit_records is not None: + run_inputs["repo_context"] = _repo_context_record( + run_inputs["repo_context"], repo_plan, unit_records, + ) + record = run_inputs["repo_context"] + tracer.event( + "repo_context", "ok", mode=record["mode"], reader=record["reader"], + listing=record["listing"], reads=record["reads"], + read_cap_hit=record["read_cap_hit"], + ) # One more worker-style unit, not inside the pool: the sweep digests the # WHOLE diff, so it only has something to say once every chunk result — @@ -1773,23 +1875,191 @@ def read(path: str) -> str | None: return read -def _context_blocks(chunk, reader, *, include_definitions: bool) -> str: - """Render the chunk's dependency and definition blocks; never raises.""" - if reader is None: +def _context_blocks( + chunk, reader, *, include_definitions: bool, unit: repo_unit.UnitContext | None = None, +) -> str: + """Render the chunk's dependency, definition and contract blocks; never raises. + + ``unit`` is the chunk's repository context. Its definition lines follow + the same-file definitions under one header and its contract lines form + the contracts block; they render with no ``reader`` too, over empty + dependency and same-file lists. ``None``, or a unit with no lines, is + exactly the rendering without repository context. + """ + extra = unit.definition_lines if unit is not None else () + contracts = unit.contract_lines if unit is not None else () + if reader is None and not (extra or contracts): return "" + deps: list[str] = [] + defs: list[str] = [] + if reader is not None: + try: + files = chunk_context.chunk_files(chunk) + deps = chunk_context.dependency_versions(files, reader) + defs = ( + chunk_context.referenced_definitions(files, reader) + if include_definitions else [] + ) + except Exception as e: # noqa: BLE001 + logger.debug("chunk context unavailable: %s", e) + if not (extra or contracts): + return "" + deps, defs = [], [] try: - files = chunk_context.chunk_files(chunk) - deps = chunk_context.dependency_versions(files, reader) - defs = ( - chunk_context.referenced_definitions(files, reader) - if include_definitions else [] + return chunk_context.render_context_blocks( + deps, defs, extra_def_lines=extra, contract_lines=contracts, ) - return chunk_context.render_context_blocks(deps, defs) except Exception as e: # noqa: BLE001 logger.debug("chunk context unavailable: %s", e) return "" +@dataclass(frozen=True) +class _RepoPlan: + """One run's repository-context inputs, fixed before the chunk workers start. + + ``reader`` is the run's one :class:`prxref.repo_reader.RepoReader`, or + ``None``. ``diff_paths`` holds every PR diff file's path: a read of one + goes to the shared, uncapped ``reader.read``. The listing and contract + fields are the ``"repo"`` level's once-per-run inputs, and stay empty at + ``"diff"`` or without a reader. + """ + + mode: str + reader: repo_reader.RepoReader | None + max_chars: int + exclude: Callable[[str], bool] + diff_paths: frozenset[str] + listing_paths: frozenset[str] | None = None + listing_complete: bool = False + contract_paths: tuple[str, ...] = () + contract_priority: tuple[str, ...] = () + + +def _plan_repo_context( + mode: str, forge: Forge, ref: PRRef, pr: PRData, files: Sequence[Any], + initial: Mapping[str, Any], *, repo_dir: RepoDir | None, +) -> _RepoPlan: + """Build the run's reader and its once-per-run inputs, for a level other than ``"off"``. + + ``initial`` is the run's initial ``repo_context`` record, whose + ``max_chars`` and glob lists are the inputs. The reader reads + ``repo_dir`` when it is given, else the forge at the PR's head sha. At + ``"repo"`` with a reader, the listing is taken here, once, and the + contract files are selected here, once. At ``"repo"``, a missing reader + or a missing listing logs the run's one WARNING naming + ``PRXREF_REPO_CONTEXT``. + """ + exclude = exclude_predicate(initial["exclude_globs"]) + if repo_dir is not None: + reader = repo_reader.repo_dir_reader(repo_dir, exclude=exclude) + else: + reader = repo_reader.forge_reader( + forge, ref, getattr(pr, "source_sha", "") or "", exclude=exclude, + ) + diff_paths = frozenset(f.path for f in files) + if mode != "repo": + return _RepoPlan(mode, reader, initial["max_chars"], exclude, diff_paths) + if reader is None: + logger.warning( + "PRXREF_REPO_CONTEXT=repo, but there is no repository reader (the forge cannot " + "read files at the PR head and no repository directory was given); repository " + "context is limited to diff-only entries from hunk lines", + ) + return _RepoPlan(mode, None, initial["max_chars"], exclude, diff_paths) + listing = reader.listing() + if listing is None: + logger.warning( + "PRXREF_REPO_CONTEXT=repo, but the repository path listing is unavailable; " + "repository context runs with no name search and no glob-matched contract files " + "outside the PR's own files", + ) + globs = list(initial["contract_globs"]) + return _RepoPlan( + mode, reader, initial["max_chars"], exclude, diff_paths, + listing_paths=frozenset(listing.paths) if listing is not None else None, + listing_complete=listing.complete if listing is not None else False, + contract_paths=tuple(repo_contracts.select_contract_files( + globs, listing=listing.paths if listing is not None else None, + diff_paths=[f.path for f in files if f.status != "removed"], + )), + contract_priority=tuple(repo_contracts.literal_contract_paths(globs)), + ) + + +def _routed_read(reader: repo_reader.RepoReader, diff_paths: frozenset[str]) -> Callable[[str], str | None]: + """One chunk's ``read``: a PR diff file through the shared ``reader.read``, any other path capped. + + The capped half is a fresh :meth:`~prxref.repo_reader.RepoReader.chunk_reader`, + so the per-chunk cap is spent on paths outside the diff alone. Both + halves refuse an excluded path. + """ + capped = reader.chunk_reader() + shared = reader.read + + def read(path: str) -> str | None: + return shared(path) if path in diff_paths else capped(path) + + return read + + +def _chunk_unit( + plan: _RepoPlan, chunk, all_files, *, index: int, total: int, +) -> repo_unit.UnitContext: + """Build one chunk's repository context in its worker; never raises. + + A build that raises gives :data:`prxref.repo_unit.EMPTY_UNIT` and one + WARNING naming the chunk. + """ + read = _routed_read(plan.reader, plan.diff_paths) if plan.reader is not None else None + try: + return repo_unit.build_unit_context( + chunk, all_files if all_files is not None else chunk, + mode=plan.mode, read=read, max_chars=plan.max_chars, + listing_paths=plan.listing_paths, listing_complete=plan.listing_complete, + contract_paths=plan.contract_paths, contract_priority=plan.contract_priority, + exclude=plan.exclude, + ) + except Exception as e: # noqa: BLE001 - context is never worth a failed review + logger.warning( + "[chunk %d/%d] repository context failed (continuing without it): %s", + index, total, e, + ) + return repo_unit.EMPTY_UNIT + + +def _unit_row(unit: repo_unit.UnitContext, *, retry_dropped: bool = False) -> dict[str, Any]: + """One chunk's ``repo_context`` units row: the unit's record plus ``retry_dropped``.""" + return {**unit.record(), "retry_dropped": retry_dropped} + + +def _repo_context_record( + initial: Mapping[str, Any], plan: _RepoPlan, unit_records: Sequence[dict[str, Any] | None], +) -> dict[str, Any]: + """The ``repo_context`` record once the chunk workers are done. + + ``reader`` is the reader's ``kind`` or ``None``; ``listing``, ``reads`` + and ``read_cap_hit`` come from one ``stats()`` snapshot (``None``, 0 + and false without a reader); ``units`` lists the rows in chunk order, + where a chunk whose worker left no row gets the empty unit's. + """ + reader = plan.reader + stats = reader.stats() if reader is not None else {"reads": 0, "read_cap_hit": False, "listing": None} + return { + **initial, + "reader": reader.kind if reader is not None else None, + "listing": stats["listing"], + "reads": stats["reads"], + "read_cap_hit": stats["read_cap_hit"], + "units": { + "chunks": [ + row if row is not None else _unit_row(repo_unit.EMPTY_UNIT) + for row in unit_records + ], + }, + } + + def _scoped_unit_blocks( scoped_rules: Any, chunks, always_on, *, max_chars: int, ) -> tuple[list[Any], Any]: @@ -1856,6 +2126,8 @@ def _run_workers( trace_dir: str | None = None, prompt_context: PromptContext = NO_PROMPT_CONTEXT, scoped_blocks: Sequence[Any] | None = None, + repo_plan: _RepoPlan | None = None, + unit_records: list[dict[str, Any] | None] | None = None, ) -> list[dict]: # Never below 1: ThreadPoolExecutor rejects a zero-width pool, and a # library caller is not gated by config's range check. @@ -1893,6 +2165,7 @@ def _heartbeat() -> None: trace_label=f"chunk{i}", trace_dir=trace_dir, prompt_context=prompt_context, scoped_block=scoped_blocks[i] if scoped_blocks is not None else None, + repo_plan=repo_plan, unit_records=unit_records, ) for i, chunk in enumerate(chunks) ] @@ -1941,6 +2214,7 @@ def _invoke_chunk( reader=None, *, include_definitions: bool = True, all_files=None, trace_label: str = "", trace_dir: str | None = None, prompt_context: PromptContext = NO_PROMPT_CONTEXT, + unit: repo_unit.UnitContext | None = None, ) -> dict: """One normalized :func:`reviewer.review_chunk` call; never raises. @@ -1960,14 +2234,16 @@ def _invoke_chunk( unchanged on both attempts: it is intent, not bulk context, and a dict-shaped finding keeps its ``scope`` only when :attr:`reviewer.PromptContext.scope_active`, and its ``rule`` only when - :attr:`reviewer.PromptContext.rule_active`. + :attr:`reviewer.PromptContext.rule_active`. ``unit`` is the chunk's + repository context (:func:`_context_blocks`); :func:`_run_worker` passes + it on the first attempt only, and ``None`` renders exactly as before. The shape carries the reviewer's reported ``cost_usd`` and ``cost_source`` beside the token counts; a call that raised, or a stub whose meta lacks them, gives ``None`` and ``""``. Pricing is left to :func:`_stamp_run_cost`, over the whole run. """ - blocks = _context_blocks(chunk, reader, include_definitions=include_definitions) + blocks = _context_blocks(chunk, reader, include_definitions=include_definitions, unit=unit) try: res = reviewer.review_chunk( llm, chunk, pr_title=pr.title, pr_description=pr.description, @@ -2025,6 +2301,8 @@ def _run_worker( trace_label: str = "", trace_dir: str | None = None, *, prompt_context: PromptContext = NO_PROMPT_CONTEXT, scoped_block: Any = None, + repo_plan: _RepoPlan | None = None, + unit_records: list[dict[str, Any] | None] | None = None, ) -> dict: tracer = tracer if tracer is not None else get_tracer() t0 = time.perf_counter() @@ -2046,9 +2324,20 @@ def _run_worker( files=[f.path for f in chunk], **({"rules": _scoped_rows(scoped_block)} if scoped_block is not None else {}), ) + unit: repo_unit.UnitContext | None = None + if repo_plan is not None: + unit = _chunk_unit(repo_plan, chunk, all_files, index=index, total=total) + if unit_records is not None: + unit_records[index - 1] = _unit_row(unit) + tracer.event( + "chunk", "context", index=index, total=total, + entries=len(unit.entries), omitted=unit.omitted, + chars=sum(len(entry.rendered()) for entry in unit.entries), + ) res = _invoke_chunk( llm, chunk, pr, max_tokens, context_lines, reader, all_files=all_files, trace_label=trace_label, trace_dir=trace_dir, prompt_context=unit_context, + unit=unit, ) if ( res["error"] @@ -2070,6 +2359,8 @@ def _run_worker( # The dependency block is a handful of tokens and survives; the # definitions block is the bulky one and is dropped, because shrinking # the prompt is the entire point of this retry. + if unit is not None and unit_records is not None: + unit_records[index - 1] = _unit_row(unit, retry_dropped=True) res = _invoke_chunk( llm, chunk, pr, max_tokens, _TIMEOUT_RETRY_CONTEXT_LINES, reader, include_definitions=False, all_files=all_files, diff --git a/tests/test_cli.py b/tests/test_cli.py index 2742d48..a45f0c8 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -901,6 +901,14 @@ def test_run_review_passes_only_real_orchestrate_kwargs(fake_runtime, monkeypatc assert sorted(set(calls[0]) - set(params)) == [] +# ``orchestrate_review`` parameters that are config keys but that ``_run_review`` +# does not pass yet. The test below asserts they are still absent from the call, +# so passing one fails it until the name is removed from here. +NOT_WIRED_YET = frozenset({ + "repo_context", "repo_context_max_chars", "context_contract_globs", "context_exclude_globs", +}) + + def test_run_review_passes_every_configured_orchestrate_kwarg(fake_runtime, monkeypatch): """The reverse direction: every ``orchestrate_review`` parameter that is also a ``load_config`` key is handed over by ``_run_review``. @@ -914,8 +922,9 @@ def test_run_review_passes_every_configured_orchestrate_kwarg(fake_runtime, monk real = real_orchestrator.orchestrate_review assert sys.modules["prxref.orchestrator"].orchestrate_review is not real params = inspect.signature(real).parameters - expected = {name for name in params if name in config._DEFAULTS} + expected = {name for name in params if name in config._DEFAULTS} - NOT_WIRED_YET assert expected, "no orchestrate parameter is a config key, so the check is vacuous" + assert NOT_WIRED_YET <= set(params) ref = PRRef( forge="github", host="github.com", owner="org", repo="repo", number=7, url="https://github.com/org/repo/pull/7", @@ -927,6 +936,7 @@ def test_run_review_passes_every_configured_orchestrate_kwarg(fake_runtime, monk calls = fake_runtime["orchestrate_calls"] assert len(calls) == 1 assert sorted(expected - set(calls[0])) == [] + assert not NOT_WIRED_YET & set(calls[0]), "now passed: drop it from NOT_WIRED_YET" assert {"rules", "ticket", "replay"} <= set(calls[0]) diff --git a/tests/test_orchestrator.py b/tests/test_orchestrator.py index 248c7b4..bbf656f 100644 --- a/tests/test_orchestrator.py +++ b/tests/test_orchestrator.py @@ -278,7 +278,7 @@ def test_posts_summary_and_inline_comments(self): "sampling", "cost_usd", "cost_estimated", "review_rules", "ticket_context", "spec_grounding", "size_advisory", "prompt_templates", "scoped_rules", - "rule_counts", + "rule_counts", "repo_context", } assert res["verdict"] == "Request-Changes" assert len(res["findings_active"]) == 2 @@ -747,7 +747,7 @@ def test_result_key_set_is_unchanged_by_the_new_knob(self): "sampling", "cost_usd", "cost_estimated", "review_rules", "ticket_context", "spec_grounding", "size_advisory", "prompt_templates", "scoped_rules", - "rule_counts", + "rule_counts", "repo_context", } @@ -1062,7 +1062,7 @@ def test_result_key_set_is_unchanged_by_the_new_knobs(self): "sampling", "cost_usd", "cost_estimated", "review_rules", "ticket_context", "spec_grounding", "size_advisory", "prompt_templates", "scoped_rules", - "rule_counts", + "rule_counts", "repo_context", } @@ -1072,7 +1072,7 @@ def test_result_key_set_is_unchanged_by_the_new_knobs(self): "elapsed_ms", "input_tokens", "output_tokens", "posted", "sampling", "cost_usd", "cost_estimated", "review_rules", "ticket_context", "spec_grounding", "size_advisory", "prompt_templates", "scoped_rules", - "rule_counts", + "rule_counts", "repo_context", } diff --git a/tests/test_orchestrator_repo_context.py b/tests/test_orchestrator_repo_context.py new file mode 100644 index 0000000..b691316 --- /dev/null +++ b/tests/test_orchestrator_repo_context.py @@ -0,0 +1,636 @@ +"""Repository context in the orchestrator (#17 T13): ``orchestrate_review(repo_context=...)``. + +The orchestrator takes the level (``off``, ``diff`` or ``repo``), the per-chunk +budget, the contract and exclude globs and an optional ``RepoDir``, and wires +them: one ``RepoReader`` per run, the listing and the contract files once per +run, each chunk's ``build_unit_context`` inside its worker before the first +attempt, the new blocks on the first attempt only, the ``repo_context`` record +on every exit, one WARNING when ``repo`` degrades, and the trace events. + +These tests run the REAL reviewer end to end over ``tests/fixtures/issue17`` +and capture every prompt with a recording LLM; none uses the ``contract_stubs`` +fixture, which replaces the worker template. With ``max_files_per_chunk=1`` +the fixture chunks into [ConnectorService], [TransportConfig], [003], in that +order. No test touches the network. +""" +from __future__ import annotations + +import inspect +import json +import logging +import threading +from collections import Counter +from pathlib import Path + +import pytest + +from prxref import config, orchestrator, repo_unit +from prxref.chunk_context import CONTRACT_HEADER, DEFINITIONS_HEADER, render_context_blocks +from prxref.forges.base import PathListing +from prxref.forges.replay import LocalDiffForge +from prxref.forges.repo_dir import RepoDir +from prxref.llm import InvokeResult +from prxref.repo_context import exclude_predicate +from prxref.repo_contracts import literal_contract_paths, select_contract_files +from prxref.repo_reader import repo_dir_reader +from prxref.repo_unit import EMPTY_UNIT, build_unit_context +from prxref.triage import DEFAULT_TOKEN_BUDGET, build_chunks, parse_unified_diff +from tests.test_orchestrator import REF, FakeForge + +FIXTURE = Path(__file__).parent / "fixtures" / "issue17" +REPO = FIXTURE / "repo" +DIFF_FILE = FIXTURE / "pr.diff" +DIFF = DIFF_FILE.read_text(encoding="utf-8") +CONNECTOR_SERVICE = "src/main/java/com/acme/connectors/ConnectorService.java" +TRANSPORT_CONFIG = "src/main/java/com/acme/connectors/TransportConfig.java" +MIGRATION = "db/changelog/003-idempotency-unique.sql" +IDEMPOTENCY_TABLE = "db/changelog/002-create-idempotency-keys.sql" +SPEC = "api/openapi/connectors.yaml" +CHUNK_ORDER = [CONNECTOR_SERVICE, TRANSPORT_CONFIG, MIGRATION] +EXCLUSIVITY = "exactly one of url or legacyUrl must be set" +GLOBS = tuple(config._DEFAULTS["context_contract_globs"]) +MAX_CHARS = config._DEFAULTS["repo_context_max_chars"] +WARN_NAME = "PRXREF_REPO_CONTEXT" +NO_FINDINGS = '{"findings": []}' +INITIAL_KEYS = { + "mode", "max_chars", "contract_globs", "exclude_globs", "reader", "listing", + "reads", "read_cap_hit", "units", +} + +PY_HELPER = ( + "diff --git a/tools/sync.py b/tools/sync.py\n" + "new file mode 100644\n" + "--- /dev/null\n" + "+++ b/tools/sync.py\n" + "@@ -0,0 +1,3 @@\n" + "+from tools.helpers import load_config\n" + "+\n" + "+settings = load_config()\n" +) + + +class _ReadingForge(FakeForge): + """FakeForge over a diff, with ``get_file_content`` reading a tree and counting calls per path.""" + + def __init__(self, diff: str = DIFF, root: Path = REPO): + super().__init__(diff=diff) + self.root = Path(root) + self.content_calls: Counter[str] = Counter() + self.get_pr_calls = 0 + self._lock = threading.Lock() + + def get_pr(self, ref): + self.get_pr_calls += 1 + return super().get_pr(ref) + + def get_file_content(self, ref, path, *, sha): + with self._lock: + self.content_calls[path] += 1 + target = self.root / path + return target.read_text(encoding="utf-8") if target.is_file() else None + + +class _RepoForge(_ReadingForge): + """``_ReadingForge`` plus ``list_paths`` over the same tree, counting calls.""" + + def __init__(self, diff: str = DIFF, root: Path = REPO): + super().__init__(diff, root) + self.list_calls = 0 + + def list_paths(self, ref, *, sha): + with self._lock: + self.list_calls += 1 + paths, complete = RepoDir(self.root).list_files() + return PathListing(paths=paths, complete=complete) + + +class _RecordingLLM: + """Records every ``(system, user)`` prompt; the first worker prompt of ``timeout_path`` times out.""" + + def __init__(self, *, timeout_path: str | None = None): + self.calls: list[tuple[str, str]] = [] + self.timeout_path = timeout_path + self._timed_out = False + self._lock = threading.Lock() + + def invoke(self, system, user, *, max_tokens=4096, json_mode=False, timeout_s=60.0): + with self._lock: + self.calls.append((system, user)) + fire = ( + self.timeout_path is not None + and not self._timed_out + and f"diff --git a/{self.timeout_path} " in user + ) + if fire: + self._timed_out = True + if fire: + raise TimeoutError("request timeout after 60s") + return InvokeResult( + text=NO_FINDINGS, input_tokens=10, output_tokens=5, + model="fake-model", backend="fake", elapsed_ms=1, + ) + + +@pytest.fixture(autouse=True) +def _pinned_clock(monkeypatch): + monkeypatch.setattr(orchestrator, "_elapsed_ms", lambda t0: 0) + + +def _review(forge, llm: _RecordingLLM | None = None, *, ref=REF, **kwargs): + llm = llm if llm is not None else _RecordingLLM() + kwargs.setdefault("max_files_per_chunk", 1) + res = orchestrator.orchestrate_review(forge, ref, llm, post=False, **kwargs) + return res, llm + + +def _worker_prompts(llm: _RecordingLLM) -> dict[str, list[str]]: + """Every worker prompt (the sweep, which runs last, excluded), grouped by the chunk's file, in call order.""" + *workers, _sweep = llm.calls + grouped: dict[str, list[str]] = {path: [] for path in CHUNK_ORDER} + for _system, user in workers: + owners = [path for path in CHUNK_ORDER if f"diff --git a/{path} " in user] + assert len(owners) == 1, owners + grouped[owners[0]].append(user) + return grouped + + +def _events(path: Path) -> list[dict]: + return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()] + + +def _stable(events: list[dict]) -> list[tuple]: + """Every event but its wall-clock fields.""" + return [ + (e["node"], e["phase"], {k: v for k, v in (e.get("meta") or {}).items() if k != "elapsed_ms"}) + for e in events + ] + + +def _warnings(caplog) -> list[logging.LogRecord]: + return [r for r in caplog.records if r.levelno == logging.WARNING and WARN_NAME in r.getMessage()] + + +def _keys(row: dict) -> list[tuple[str, int, str]]: + return [(e["path"], e["line"], e["reason"]) for e in row["entries"]] + + +def _expected_rows( + mode: str, *, read: bool = True, listing: bool = True, globs=GLOBS, diff: str = DIFF, + root: Path = REPO, max_files_per_chunk: int = 1, token_budget: int = DEFAULT_TOKEN_BUDGET, +) -> list[dict]: + """The units rows ``build_unit_context`` gives each chunk when called directly with a fresh reader.""" + files = parse_unified_diff(diff) + chunks = build_chunks(files, token_budget=token_budget, max_files_per_chunk=max_files_per_chunk) + exclude = exclude_predicate(()) + reader = repo_dir_reader(RepoDir(root), exclude=exclude) if read else None + got = reader.listing() if reader is not None and listing and mode == "repo" else None + contract_paths: list[str] = [] + priority: list[str] = [] + if mode == "repo" and reader is not None: + contract_paths = select_contract_files( + list(globs), listing=got.paths if got is not None else None, + diff_paths=[f.path for f in files if f.status != "removed"], + ) + priority = literal_contract_paths(list(globs)) + rows = [] + for chunk in chunks: + unit = build_unit_context( + chunk, files, mode=mode, read=reader.chunk_reader() if reader is not None else None, + max_chars=MAX_CHARS, + listing_paths=frozenset(got.paths) if got is not None else None, + listing_complete=got.complete if got is not None else False, + contract_paths=contract_paths, contract_priority=priority, exclude=exclude, + ) + rows.append({**unit.record(), "retry_dropped": False}) + return rows + + +REPO_CONNECTOR_KEYS = [ + (TRANSPORT_CONFIG, 9, "cross-chunk"), + (TRANSPORT_CONFIG, 11, "cross-chunk"), + (SPEC, 6, "contract"), + (SPEC, 31, "contract"), + (SPEC, 41, "contract"), +] + + +class TestTheParameters: + NEW = ["repo_context", "repo_context_max_chars", "context_contract_globs", "context_exclude_globs", "repo_dir"] + + def test_the_five_kwargs_follow_max_findings_per_rule_keyword_only(self): + params = inspect.signature(orchestrator.orchestrate_review).parameters + names = list(params) + start = names.index("max_findings_per_rule") + 1 + assert names[start:start + 5] == self.NEW + for name in self.NEW: + assert params[name].kind is inspect.Parameter.KEYWORD_ONLY, name + + def test_the_defaults_are_off_and_restate_config(self): + params = inspect.signature(orchestrator.orchestrate_review).parameters + assert params["repo_context"].default == "off" == config._DEFAULTS["repo_context"] + assert params["repo_context_max_chars"].default == 12000 == MAX_CHARS + assert params["context_contract_globs"].default == () + assert params["context_exclude_globs"].default == () + assert params["repo_dir"].default is None + + @pytest.mark.parametrize("mode", ["bogus", "", "OFF", None]) + def test_an_unknown_level_raises_before_any_forge_call(self, mode): + forge = _RepoForge() + with pytest.raises(ValueError, match="repo_context"): + _review(forge, repo_context=mode) + assert forge.get_pr_calls == 0 + assert forge.list_calls == 0 + assert not forge.content_calls + + +class TestOffIsByteIdentical: + """D1: ``off`` with every other new kwarg set is exactly a run without them.""" + + def _run(self, tmp_path, caplog, name, **kwargs): + forge = _RepoForge(DIFF + PY_HELPER) + trace = tmp_path / f"{name}.jsonl" + caplog.clear() + with caplog.at_level(logging.DEBUG, logger="prxref"): + res, llm = _review( + forge, _RecordingLLM(timeout_path=CONNECTOR_SERVICE), + max_files_per_chunk=4, max_workers=1, trace_file=str(trace), **kwargs, + ) + logs = [ + (r.name, r.levelname, r.getMessage()) for r in caplog.records + if r.levelno >= logging.INFO + ] + return res, llm.calls, _stable(_events(trace)), logs, forge + + def test_off_with_every_kwarg_matches_a_run_without_them(self, tmp_path, caplog): + base = self._run(tmp_path, caplog, "base") + off = self._run( + tmp_path, caplog, "off", + repo_context="off", repo_context_max_chars=5, context_contract_globs=GLOBS, + context_exclude_globs=("**/*.yaml",), repo_dir=RepoDir(REPO), + ) + assert off[:4] == base[:4] + assert off[4].content_calls == base[4].content_calls + assert base[4].content_calls, "the old reader read nothing, so the call-count check is vacuous" + assert off[4].list_calls == base[4].list_calls == 0 + assert base[0]["repo_context"] is None and off[0]["repo_context"] is None + assert ("chunk", "retry") in [(n, p) for n, p, _ in base[2]] + for node, phase, _meta in off[2]: + assert (node, phase) != ("chunk", "context") + assert node != "repo_context" + + def test_the_oracle_sees_a_repo_run_differ(self, tmp_path, caplog): + """The control: the same oracle over ``repo`` must come out different on every axis it checks.""" + base = self._run(tmp_path, caplog, "base") + on = self._run(tmp_path, caplog, "on", repo_context="repo", context_contract_globs=GLOBS) + assert on[1] != base[1] + assert on[2] != base[2] + assert on[4].list_calls == 1 + assert on[4].content_calls != base[4].content_calls + assert on[0]["repo_context"] is not None + + +class TestDiffLevel: + def test_the_connector_chunk_gets_the_cross_chunk_definitions_and_no_contracts(self): + forge = _RepoForge() + res, llm = _review(forge, repo_context="diff", context_contract_globs=GLOBS, max_workers=1) + prompts = _worker_prompts(llm) + (connector,) = prompts[CONNECTOR_SERVICE] + block = connector[connector.index(DEFINITIONS_HEADER):] + assert f"{TRANSPORT_CONFIG}:9: " in block + assert f"{TRANSPORT_CONFIG}:11: " in block + assert EXCLUSIVITY in block + assert all(CONTRACT_HEADER not in p for group in prompts.values() for p in group) + assert forge.list_calls == 0 + record = res["repo_context"] + assert (record["mode"], record["reader"], record["listing"]) == ("diff", "forge", None) + assert record["reads"] == 2 + assert record["read_cap_hit"] is False + assert record["units"]["chunks"] == _expected_rows("diff") + assert _keys(record["units"]["chunks"][0]) == REPO_CONNECTOR_KEYS[:2] + + def test_an_off_run_has_no_definitions_block_on_that_chunk(self): + """The control for the test above: the block is the level's doing, not the fixture's. + + The exclusivity text itself is in every run's "Other files changed" digest, so the + control looks for the definition's ``path:line:`` rows instead. + """ + _, llm = _review(_RepoForge(), max_workers=1) + (connector,) = _worker_prompts(llm)[CONNECTOR_SERVICE] + assert DEFINITIONS_HEADER not in connector + assert f"{TRANSPORT_CONFIG}:9: " not in connector + assert f"{TRANSPORT_CONFIG}:11: " not in connector + + def test_diff_with_no_reader_builds_hunk_entries_and_logs_nothing(self, caplog): + with caplog.at_level(logging.WARNING, logger="prxref"): + res, llm = _review(FakeForge(diff=DIFF), repo_context="diff", max_workers=1) + assert _warnings(caplog) == [] + (connector,) = _worker_prompts(llm)[CONNECTOR_SERVICE] + assert EXCLUSIVITY in connector[connector.index(DEFINITIONS_HEADER):] + record = res["repo_context"] + assert (record["reader"], record["listing"], record["reads"]) == (None, None, 0) + assert record["units"]["chunks"] == _expected_rows("diff", read=False) + + +class TestRepoLevel: + def test_the_connector_chunk_gets_its_contract_excerpts(self, caplog): + forge = _RepoForge() + with caplog.at_level(logging.WARNING, logger="prxref"): + res, llm = _review(forge, repo_context="repo", context_contract_globs=GLOBS, max_workers=1) + assert _warnings(caplog) == [] + (connector,) = _worker_prompts(llm)[CONNECTOR_SERVICE] + contracts = connector[connector.index(CONTRACT_HEADER):] + for line in (6, 31, 41): + assert f"{SPEC}:{line}: " in contracts + assert "mutually exclusive" in contracts + assert forge.list_calls == 1 + record = res["repo_context"] + assert set(record) == INITIAL_KEYS + assert record["mode"] == "repo" and record["reader"] == "forge" + assert record["max_chars"] == MAX_CHARS + assert record["contract_globs"] == list(GLOBS) and record["exclude_globs"] == [] + assert record["listing"] == {"paths": 6, "complete": True} + assert record["reads"] > 0 + assert record["read_cap_hit"] is False + rows = record["units"]["chunks"] + assert rows == _expected_rows("repo") + assert [_keys(row) for row in rows] == [ + REPO_CONNECTOR_KEYS, + [(SPEC, 31, "contract")], + [(SPEC, 50, "contract"), (IDEMPOTENCY_TABLE, 4, "contract")], + ] + + def test_empty_contract_globs_mean_no_contract_files(self): + res, llm = _review(_RepoForge(), repo_context="repo", max_workers=1) + assert all(CONTRACT_HEADER not in p for group in _worker_prompts(llm).values() for p in group) + assert _keys(res["repo_context"]["units"]["chunks"][0]) == REPO_CONNECTOR_KEYS[:2] + + def test_an_excluded_spec_is_never_read(self): + forge = _RepoForge() + res, llm = _review( + forge, repo_context="repo", context_contract_globs=GLOBS, + context_exclude_globs=["api/**"], max_workers=1, + ) + assert forge.content_calls[SPEC] == 0 + assert res["repo_context"]["exclude_globs"] == ["api/**"] + assert res["repo_context"]["listing"] == {"paths": 5, "complete": True} + assert all(SPEC not in p for group in _worker_prompts(llm).values() for p in group) + + def test_repo_dir_on_a_local_diff_forge(self, caplog): + """No ``get_file_content`` and no head sha: the old blocks have no reader, and the new ones still render.""" + forge = LocalDiffForge(DIFF, path=str(DIFF_FILE)) + with caplog.at_level(logging.WARNING, logger="prxref"): + res, llm = _review( + forge, ref=LocalDiffForge.ref_for(str(DIFF_FILE)), repo_context="repo", + context_contract_globs=GLOBS, repo_dir=RepoDir(REPO), max_workers=1, + ) + assert _warnings(caplog) == [] + record = res["repo_context"] + assert record["reader"] == "repo-dir" + assert record["listing"] == {"paths": 6, "complete": True} + assert record["units"]["chunks"] == _expected_rows("repo") + (connector,) = _worker_prompts(llm)[CONNECTOR_SERVICE] + assert f"{SPEC}:6: " in connector[connector.index(CONTRACT_HEADER):] + assert EXCLUSIVITY in connector[connector.index(DEFINITIONS_HEADER):] + + def test_repo_with_no_reader_warns_once_and_keeps_the_hunk_entries(self, caplog): + with caplog.at_level(logging.WARNING, logger="prxref"): + res, llm = _review( + FakeForge(diff=DIFF), repo_context="repo", context_contract_globs=GLOBS, max_workers=4, + ) + (warning,) = _warnings(caplog) + assert "hunk lines" in warning.getMessage() + record = res["repo_context"] + assert (record["reader"], record["listing"], record["reads"], record["read_cap_hit"]) == ( + None, None, 0, False, + ) + rows = record["units"]["chunks"] + assert rows == _expected_rows("repo", read=False) + assert _keys(rows[0]) == REPO_CONNECTOR_KEYS[:2] + (connector,) = _worker_prompts(llm)[CONNECTOR_SERVICE] + assert f"{TRANSPORT_CONFIG}:9: " in connector[connector.index(DEFINITIONS_HEADER):] + assert CONTRACT_HEADER not in connector + + def test_repo_with_a_reader_but_no_listing_warns_once(self, caplog): + forge = _ReadingForge() + assert not hasattr(forge, "list_paths") + with caplog.at_level(logging.WARNING, logger="prxref"): + res, _ = _review(forge, repo_context="repo", context_contract_globs=GLOBS, max_workers=4) + (warning,) = _warnings(caplog) + assert "listing" in warning.getMessage() + record = res["repo_context"] + assert (record["reader"], record["listing"]) == ("forge", None) + assert record["units"]["chunks"] == _expected_rows("repo", listing=False) + + +class TestTheTimeoutRetry: + def test_the_retry_carries_neither_new_block(self, tmp_path): + trace = tmp_path / "trace.jsonl" + res, llm = _review( + _RepoForge(), _RecordingLLM(timeout_path=CONNECTOR_SERVICE), repo_context="repo", + context_contract_globs=GLOBS, max_workers=1, trace_file=str(trace), + ) + first, retry = _worker_prompts(llm)[CONNECTOR_SERVICE] + assert CONTRACT_HEADER in first and DEFINITIONS_HEADER in first + assert CONTRACT_HEADER not in retry + assert DEFINITIONS_HEADER not in retry + assert f"{SPEC}:" not in retry and f"{TRANSPORT_CONFIG}:9: " not in retry + assert res["chunks_failed"] == 0 + assert [(n, p) for n, p, _ in _stable(_events(trace))].count(("chunk", "retry")) == 1 + rows = res["repo_context"]["units"]["chunks"] + assert [row["retry_dropped"] for row in rows] == [True, False, False] + expected = _expected_rows("repo") + expected[0]["retry_dropped"] = True + assert rows == expected + + +class TestABuildThatRaises: + def test_the_chunk_gets_the_empty_unit_and_one_warning(self, monkeypatch, caplog): + real = repo_unit.build_unit_context + + def flaky(chunk, all_files, **kwargs): + if any(f.path == CONNECTOR_SERVICE for f in chunk): + raise RuntimeError("resolver exploded") + return real(chunk, all_files, **kwargs) + + monkeypatch.setattr(orchestrator.repo_unit, "build_unit_context", flaky) + with caplog.at_level(logging.WARNING, logger="prxref"): + res, llm = _review(_RepoForge(), repo_context="repo", context_contract_globs=GLOBS, max_workers=1) + failures = [r for r in caplog.records if "repository context failed" in r.getMessage()] + assert len(failures) == 1 + assert "[chunk 1/3]" in failures[0].getMessage() + assert failures[0].levelno == logging.WARNING + assert res["chunks_failed"] == 0 + assert res["verdict"] == "Approved" + rows = res["repo_context"]["units"]["chunks"] + assert rows[0] == {**EMPTY_UNIT.record(), "retry_dropped": False} + assert rows[1:] == _expected_rows("repo")[1:] + (connector,) = _worker_prompts(llm)[CONNECTOR_SERVICE] + assert CONTRACT_HEADER not in connector + + +class TestEarlyExits: + def _initial(self, mode: str) -> dict: + return { + "mode": mode, "max_chars": MAX_CHARS, "contract_globs": list(GLOBS), "exclude_globs": [], + "reader": None, "listing": None, "reads": 0, "read_cap_hit": False, "units": None, + } + + @pytest.mark.parametrize("mode", ["diff", "repo"]) + def test_an_empty_diff_carries_the_initial_record(self, mode): + forge = _RepoForge("") + res, _ = _review(forge, repo_context=mode, context_contract_globs=GLOBS) + assert res["repo_context"] == self._initial(mode) + assert forge.list_calls == 0 and not forge.content_calls + + def test_an_empty_diff_off_is_none(self): + res, _ = _review(_RepoForge("")) + assert res["repo_context"] is None + + def test_a_get_pr_failure_carries_the_initial_record(self): + forge = _RepoForge() + forge.fail.add("get_pr") + res, _ = _review(forge, repo_context="repo", context_contract_globs=GLOBS) + assert res["verdict"] == "Error" + assert res["repo_context"] == self._initial("repo") + + def test_a_total_failure_carries_the_filled_record(self): + class _DeadLLM(_RecordingLLM): + def invoke(self, system, user, **kwargs): + raise RuntimeError("no model") + + res, _ = _review(_RepoForge(), _DeadLLM(), repo_context="repo", context_contract_globs=GLOBS) + assert res["verdict"] == "Error" + assert res["repo_context"]["reader"] == "forge" + assert res["repo_context"]["units"]["chunks"] == _expected_rows("repo") + + +class TestDeterminism: + def test_one_and_four_workers_give_equal_records(self): + forge = _RepoForge() + records = [ + _review(forge, repo_context="repo", context_contract_globs=GLOBS, max_workers=workers)[0]["repo_context"] + for workers in (1, 4, 4) + ] + assert records[0]["units"]["chunks"] == _expected_rows("repo") + assert records[1] == records[0] + assert records[2] == records[0] + + +class TestTheTrace: + def test_one_context_event_per_chunk_and_one_run_event(self, tmp_path): + trace = tmp_path / "trace.jsonl" + res, _ = _review( + _RepoForge(), repo_context="repo", context_contract_globs=GLOBS, max_workers=4, + trace_file=str(trace), + ) + events = _events(trace) + record = res["repo_context"] + rows = record["units"]["chunks"] + context = sorted( + (e["meta"] for e in events if (e["node"], e["phase"]) == ("chunk", "context")), + key=lambda meta: meta["index"], + ) + assert context == [ + { + "index": i, "total": 3, "entries": len(row["entries"]), "omitted": row["omitted"], + "chars": sum(entry["chars"] for entry in row["entries"]), + } + for i, row in enumerate(rows, start=1) + ] + (run,) = [e for e in events if e["node"] == "repo_context"] + assert run["phase"] == "ok" + assert run["meta"] == {key: record[key] for key in ("mode", "reader", "listing", "reads", "read_cap_hit")} + sweep_start = next( + i for i, e in enumerate(events) if (e["node"], e["phase"]) == ("sweep", "start") + ) + assert events.index(run) < sweep_start + + +def _java_type(name: str) -> str: + return f"package com.acme.widgets;\n\npublic class {name} {{\n private final int size = 0;\n}}\n" + + +def _new_file_diff(path: str, text: str) -> str: + lines = text.splitlines() + body = "".join(f"+{line}\n" for line in lines) + return ( + f"diff --git a/{path} b/{path}\nnew file mode 100644\n--- /dev/null\n+++ b/{path}\n" + f"@@ -0,0 +1,{len(lines)} @@\n{body}" + ) + + +WIDGET_SPEC = ( + "openapi: 3.0.3\n" + "info:\n" + " title: Widgets\n" + ' version: "1"\n' + "paths:\n" + " /widgets:\n" + " post:\n" + " operationId: createWidget\n" + " responses:\n" + " '201':\n" + " description: The created widget.\n" +) + + +class TestDiffFileReadsAreUncapped: + """D-L: a chunk's diff-file reads go to the shared read, so its capped reads are left for everything else.""" + + PARTS = [f"Part{i:02d}" for i in range(1, 21)] + CONTROLLER = "src/main/java/com/acme/widgets/WidgetController.java" + + def _tree(self, root: Path) -> str: + texts = {f"src/main/java/com/acme/widgets/{name}.java": _java_type(name) for name in self.PARTS} + texts[self.CONTROLLER] = ( + "package com.acme.widgets;\n\n" + "public class WidgetController {\n\n" + ' @PostMapping("/widgets")\n' + f" public Widget create({', '.join(f'{n} p{i}' for i, n in enumerate(self.PARTS))}) {{\n" + " return null;\n" + " }\n" + "}\n" + ) + for path, text in {**texts, "api/openapi.yaml": WIDGET_SPEC}.items(): + target = root / path + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(text, encoding="utf-8") + return "".join(_new_file_diff(path, text) for path, text in texts.items()) + + def test_a_chunk_of_many_diff_files_still_gets_its_contract(self, tmp_path): + root = tmp_path / "repo" + diff = self._tree(root) + shape = {"max_files_per_chunk": 30, "token_budget": 200_000} + assert len(build_chunks(parse_unified_diff(diff), **shape)) == 1 + capped = _expected_rows("repo", diff=diff, root=root, **shape) + assert all(e["kind"] != "contract" for e in capped[0]["entries"]), ( + "the capped reader was not starved, so this test cannot see the routing" + ) + res, llm = _review( + FakeForge(diff=diff), repo_context="repo", context_contract_globs=GLOBS, + repo_dir=RepoDir(root), **shape, + ) + record = res["repo_context"] + (row,) = record["units"]["chunks"] + assert [(e["path"], e["line"], e["kind"]) for e in row["entries"]] == [("api/openapi.yaml", 6, "contract")] + assert record["read_cap_hit"] is False + worker = llm.calls[0][1] + assert "api/openapi.yaml:6: " in worker[worker.index(CONTRACT_HEADER):] + + +class TestContextBlocks: + """``_context_blocks`` with a unit and no old reader, and with no unit, called directly.""" + + def _chunk(self): + return parse_unified_diff(DIFF)[:1] + + def test_no_reader_and_no_unit_is_empty(self): + assert orchestrator._context_blocks(self._chunk(), None, include_definitions=True) == "" + assert orchestrator._context_blocks(self._chunk(), None, include_definitions=True, unit=EMPTY_UNIT) == "" + + def test_a_unit_renders_without_the_old_reader(self): + unit = repo_unit.UnitContext(("a.java:1: class A {}",), ("spec.yaml:2: /a:",), (), 0) + out = orchestrator._context_blocks(self._chunk(), None, include_definitions=True, unit=unit) + assert out == render_context_blocks([], [], extra_def_lines=unit.definition_lines, + contract_lines=unit.contract_lines) + assert out.startswith(DEFINITIONS_HEADER) and CONTRACT_HEADER in out diff --git a/tests/test_orchestrator_rule_cap.py b/tests/test_orchestrator_rule_cap.py index 9d0d9e7..7ad2b5c 100644 --- a/tests/test_orchestrator_rule_cap.py +++ b/tests/test_orchestrator_rule_cap.py @@ -605,8 +605,9 @@ def _pinned_clock(self, monkeypatch): def test_cap_zero_with_rules_is_byte_identical_to_base(self, name): text, res = rules_capture(name, max_findings_per_rule=0) assert _sha(text) == RULES_GOLDEN[name] - assert set(res) == set(A81_RECORD_KEYS) | {"rule_counts"} + assert set(res) == set(A81_RECORD_KEYS) | {"rule_counts", "repo_context"} assert res["rule_counts"] is None + assert res["repo_context"] is None payload = list(cli._build_json_result(res)) assert payload == [*A81_JSON_KEYS[:-1], "rule_counts", "sampling"] diff --git a/tests/test_run_record.py b/tests/test_run_record.py index c62eb22..4597c73 100644 --- a/tests/test_run_record.py +++ b/tests/test_run_record.py @@ -51,12 +51,16 @@ RECORD_KEYS = { "cost_usd", "cost_estimated", "review_rules", "ticket_context", "spec_grounding", "size_advisory", "prompt_templates", "scoped_rules", - "rule_counts", + "rule_counts", "repo_context", } NULL_WHEN_OFF = ( "review_rules", "ticket_context", "spec_grounding", "size_advisory", "prompt_templates", - "scoped_rules", "rule_counts", + "scoped_rules", "rule_counts", "repo_context", ) +# Record keys the orchestrator stamps that ``--format json`` does not carry +# yet. The JSON test below asserts they are still absent, so emitting one +# fails it until the key is removed from here. +NOT_IN_JSON_YET = frozenset({"repo_context"}) REPLAY = { "base_sha": "b" * 40, @@ -278,8 +282,9 @@ def test_json_payload_is_normal_keys_plus_replay(self, monkeypatch, tmp_path, pa assert "replay" not in normal_payload assert set(replay_payload) == set(normal_payload) | {"replay"} assert replay_payload["replay"] == REPLAY - for key in RECORD_KEYS: + for key in RECORD_KEYS - NOT_IN_JSON_YET: assert key in normal_payload, key + assert not NOT_IN_JSON_YET & set(normal_payload), "now emitted: drop it from NOT_IN_JSON_YET" assert normal_payload["size_advisory"] is None assert normal_payload["cost_estimated"] is False From 5d8cd80b088ab5a53fb31a18d1ce1f4c9e8f645d Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:36:46 -0700 Subject: [PATCH 17/24] fix: GitLab list_paths walks GraphQL tree.blobs before the REST tree MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The live check V1 found that GitLab's REST repository/tree?recursive=true listing returns every directory before any file across the whole recursive walk. Under the 20-page cap, a project with more than 2000 entries lost files, and gitlab-org/gitlab came back with zero paths and complete=False. list_paths now walks GraphQL's files-only tree.blobs connection first (POST https://{host}/api/graphql, variables p/ref/after, the same PRIVATE-TOKEN header and timeout as every other request), following endCursor while hasNextPage for at most MAX_LISTING_PAGES pages. The cap, or a later-page failure, keeps the paths read so far with complete=False. When the first GraphQL page is unusable, the unchanged REST walk answers. That covers a transport failure, a non-2xx status, a non-JSON body, a non-empty errors array, a null data.project, any other wrong shape, and an empty first page with hasNextPage false. GraphQL answers a sha it cannot resolve with that empty page, so REST decides between 404 -> None and an empty tree. The fallback reason is logged at DEBUG. Tests: tests/test_issue_17_gitlab_graphql.py (58 new, including the directories-first defect with a control). T9's REST tests pass unedited. 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/forges/gitlab.py | 135 ++++++- tests/test_issue_17_gitlab_graphql.py | 550 ++++++++++++++++++++++++++ 2 files changed, 682 insertions(+), 3 deletions(-) create mode 100644 tests/test_issue_17_gitlab_graphql.py diff --git a/src/prxref/forges/gitlab.py b/src/prxref/forges/gitlab.py index 13043fe..81e02aa 100644 --- a/src/prxref/forges/gitlab.py +++ b/src/prxref/forges/gitlab.py @@ -48,6 +48,12 @@ # it is still bounded: GitLab's validation errors run long enough to bury the # log line that carries them. _ERROR_DETAIL_CHARS = 400 +_TREE_BLOBS_QUERY = ( + "query($p: ID!, $ref: String!, $after: String) { project(fullPath: $p) { repository { " + "tree(ref: $ref, recursive: true) { blobs(first: 100, after: $after) { " + "pageInfo { hasNextPage endCursor } nodes { path type mode } } } } } }" +) +_TREE_BLOBS_PATH = ("data", "project", "repository", "tree", "blobs") def _response_detail(resp: requests.Response) -> str: @@ -573,8 +579,37 @@ def get_file_content(self, ref: PRRef, path: str, *, sha: str) -> str | None: def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: """Return every file path in the repository at commit ``sha``, best-effort. - Walks the paged ``repository/tree?recursive=true`` listing, - ``_PAGE_SIZE`` entries a page, requesting pages 1, 2, 3 and so on. + Two walks, GraphQL first. The REST ``repository/tree?recursive=true`` + listing returns every directory before any file, across the whole + recursive walk, so on a project with more entries than the page cap + holds, the capped REST walk loses files, and returns none at all once + the directories alone fill the cap. GraphQL's ``tree.blobs`` + connection lists files only, so its whole page budget goes to paths. + + The GraphQL walk POSTs ``_TREE_BLOBS_QUERY`` (a read: the text begins + with ``query``) to ``https://{host}/api/graphql``, the host the REST + API is on, with the variables ``p`` (the project's full path), ``ref`` + (``sha``) and ``after`` (the previous page's ``endCursor``, ``None`` + on the first page), the same auth header as every other request, and + the request timeout. It follows ``endCursor`` while ``hasNextPage`` is + true and reads at most ``MAX_LISTING_PAGES`` pages of 100 nodes. Node + paths that are non-empty strings are kept, sorted and deduplicated. A + page with ``hasNextPage`` false ends the walk with ``complete=True``. + When the cap stops the walk while ``hasNextPage`` is still true, or a + later page fails, the paths read so far come back with + ``complete=False``. + + The REST walk answers instead when the FIRST GraphQL page is unusable: + a transport failure, a non-2xx status, a body that is not JSON, a + non-empty ``errors`` array, a null ``data.project``, any other wrong + shape (``hasNextPage`` true with no ``endCursor`` included), or a first + page with no usable path and ``hasNextPage`` false. GraphQL answers a + sha it cannot resolve with exactly that empty page, which it does not + distinguish from an empty tree, so REST decides between them. The + fallback reason is logged at DEBUG. + + The REST walk requests pages 1, 2, 3 and so on of + ``repository/tree?recursive=true``, ``_PAGE_SIZE`` entries a page. Only ``blob`` entries are kept, so directories (``tree``) and submodules (``commit``) are dropped, and the paths are sorted and deduplicated. The walk ends when a page's ``X-Next-Page`` header is @@ -589,9 +624,103 @@ def list_paths(self, ref: PRRef, *, sha: str) -> PathListing | None: """ if not sha: return None + where = f"{self._project_path(ref)}@{sha}" + listing = self._list_paths_graphql(ref, sha=sha, where=where) + if listing is not None: + return listing + return self._list_paths_rest(ref, sha=sha, where=where) + + def _list_paths_graphql(self, ref: PRRef, *, sha: str, where: str) -> PathListing | None: + """Walk GraphQL's ``tree.blobs`` connection; ``None`` hands the listing to REST.""" + headers = self._get_auth_headers() + url = f"https://{ref.host}/api/graphql" + project_path = self._project_path(ref) + paths: set[str] = set() + after: str | None = None + for page_number in range(1, MAX_LISTING_PAGES + 1): + variables = {"p": project_path, "ref": sha, "after": after} + try: + resp = self._session.post( + url, + json={"query": _TREE_BLOBS_QUERY, "variables": variables}, + headers=headers, + timeout=_REQUEST_TIMEOUT, + ) + except requests.RequestException as e: + return self._graphql_stopped(paths, page_number, where, f"a transport failure ({e})") + try: + page_paths, has_next, after = self._read_blob_page(resp, first_page=page_number == 1) + except ValueError as e: + return self._graphql_stopped(paths, page_number, where, str(e)) + paths.update(page_paths) + if not has_next: + return PathListing(paths=tuple(sorted(paths)), complete=True) + logger.debug( + "list_paths stopped at the %d-page cap of the GraphQL listing for %s with %d paths", + MAX_LISTING_PAGES, where, len(paths), + ) + return PathListing(paths=tuple(sorted(paths)), complete=False) + + @staticmethod + def _read_blob_page( + resp: requests.Response, *, first_page: bool + ) -> tuple[list[str], bool, str | None]: + """Parse one ``tree.blobs`` page into (paths, hasNextPage, endCursor). + + Raises ``ValueError`` naming what makes the page unusable. A first + page with no usable path and ``hasNextPage`` false is unusable too, + because GraphQL gives that same page for a sha it cannot resolve. + """ + if not resp.ok: + raise ValueError(f"HTTP {resp.status_code}") + try: + body = resp.json() + except ValueError as e: + raise ValueError(f"a non-JSON body ({e})") from e + if not isinstance(body, dict): + raise ValueError(f"a {type(body).__name__} body, not an object") + if body.get("errors"): + raise ValueError(f"GraphQL errors ({_response_detail(resp)})") + node = body + for key in _TREE_BLOBS_PATH: + child = node.get(key) + if not isinstance(child, dict): + raise ValueError(f"no {key} object (a {type(child).__name__})") + node = child + page_info = node.get("pageInfo") + nodes = node.get("nodes") + if not isinstance(page_info, dict) or not isinstance(nodes, list): + raise ValueError("a blobs connection without a pageInfo object and a nodes list") + has_next = page_info.get("hasNextPage") + cursor = page_info.get("endCursor") + if not isinstance(has_next, bool): + raise ValueError(f"a {type(has_next).__name__} hasNextPage, not a bool") + if has_next and not (isinstance(cursor, str) and cursor): + raise ValueError("hasNextPage true with no endCursor") + paths = [ + entry["path"] for entry in nodes + if isinstance(entry, dict) and isinstance(entry.get("path"), str) and entry["path"] + ] + if first_page and not paths and not has_next: + raise ValueError("an empty first page, which is also the answer for a sha GitLab cannot resolve") + return paths, has_next, cursor if has_next else None + + def _graphql_stopped( + self, paths: set[str], page_number: int, where: str, reason: str + ) -> PathListing | None: + """Log why the GraphQL walk stopped; ``None`` on the first page falls back to REST.""" + if page_number == 1: + logger.debug( + "list_paths falls back to the REST tree walk for %s: the GraphQL listing gave %s", + where, reason, + ) + return None + return self._listing_stopped(paths, page_number, where, f"{reason} from the GraphQL listing") + + def _list_paths_rest(self, ref: PRRef, *, sha: str, where: str) -> PathListing | None: + """Walk the REST ``repository/tree?recursive=true`` listing page by page.""" headers = self._get_auth_headers() url = f"{self._api_base(ref)}/repository/tree" - where = f"{self._project_path(ref)}@{sha}" paths: set[str] = set() for page_number in range(1, MAX_LISTING_PAGES + 1): params: dict[str, int | str] = { diff --git a/tests/test_issue_17_gitlab_graphql.py b/tests/test_issue_17_gitlab_graphql.py new file mode 100644 index 0000000..1b89e83 --- /dev/null +++ b/tests/test_issue_17_gitlab_graphql.py @@ -0,0 +1,550 @@ +"""Tests for GitLab's GraphQL-first ``list_paths`` (issue #17, task T19). + +The REST ``repository/tree?recursive=true`` listing returns every directory +before any file, across the whole recursive walk, so under the page cap a large +project loses files: the live check V1 saw gitlab-org/gitlab's first 21 pages +come back 100% ``tree`` and ``list_paths`` return zero paths. The adapter now +walks GraphQL's files-only ``tree.blobs`` connection first and falls back to the +unchanged REST walk when the first GraphQL page is unusable. The response +shapes below are the ones V1 observed live on gitlab.com. +""" +from __future__ import annotations + +import copy +import json +import logging +import re + +import pytest +import requests +from requests.structures import CaseInsensitiveDict + +from prxref.forges import gitlab +from prxref.forges.base import PathListing + +SHA = "1234567890abcdef1234567890abcdef12345678" +REQUEST_TIMEOUT = (10.0, 30.0) + +GL_PR_URL = "https://gitlab.example.com/acme/platform/api/-/merge_requests/7" +PROJECT = "acme/platform/api" +GRAPHQL_URL = "https://gitlab.example.com/api/graphql" +TREE_URL = "https://gitlab.example.com/api/v4/projects/acme%2Fplatform%2Fapi/repository/tree" + +LIVE_QUERY = ( + "query($p: ID!, $ref: String!, $after: String) { project(fullPath: $p) { repository { " + "tree(ref: $ref, recursive: true) { blobs(first: 100, after: $after) { " + "pageInfo { hasNextPage endCursor } nodes { path type mode } } } } } }" +) +LIVE_EMPTY_PAGE = { + "data": {"project": {"repository": {"tree": {"blobs": { + "pageInfo": {"hasNextPage": False, "endCursor": None}, "nodes": [], + }}}}}, +} +LIVE_UNKNOWN_PROJECT = {"data": {"project": None}} + +_NOT_JSON = object() + + +@pytest.fixture(autouse=True) +def _no_ambient_token(monkeypatch): + monkeypatch.delenv("PRXREF_GITLAB_TOKEN", raising=False) + + +class FakeResponse: + """The slice of ``requests.Response`` the adapter reads.""" + + def __init__(self, status_code=200, *, json_data=_NOT_JSON, text="", headers=None): + self.status_code = status_code + self.ok = 200 <= status_code < 300 + self.headers = headers if headers is not None else CaseInsensitiveDict() + self._json = json_data + self.text = text if json_data is _NOT_JSON else json.dumps(json_data) + + def json(self): + if self._json is _NOT_JSON: + raise requests.exceptions.JSONDecodeError("Expecting value", self.text, 0) + return self._json + + +class FakeSession: + """Records every request; answers POSTs and GETs from their own queues, in order.""" + + def __init__(self, graphql=(), rest=()): + self._queues = {"POST": list(graphql), "GET": list(rest)} + self.calls: list[tuple[str, str, dict]] = [] + + def _answer(self, verb, url, kwargs): + self.calls.append((verb, url, copy.deepcopy(kwargs))) + queue = self._queues[verb] + if not queue: + raise AssertionError(f"unexpected {verb} {url}") + item = queue.pop(0) + if isinstance(item, BaseException): + raise item + return item + + def post(self, url, **kwargs): + return self._answer("POST", url, kwargs) + + def get(self, url, **kwargs): + return self._answer("GET", url, kwargs) + + def verbs(self): + return [verb for verb, _, _ in self.calls] + + def of(self, verb): + return [(url, kwargs) for v, url, kwargs in self.calls if v == verb] + + +def _ref(url=GL_PR_URL): + ref = gitlab.ForgeImpl.parse_pr_url(url) + assert ref is not None + return ref + + +def _list(session, ref=None, sha=SHA): + return gitlab.ForgeImpl(session=session).list_paths(ref or _ref(), sha=sha) + + +def _blobs_body(paths, *, has_next=False, cursor=None, nodes=None): + if nodes is None: + nodes = [{"path": p, "type": "blob", "mode": "100644"} for p in paths] + return {"data": {"project": {"repository": {"tree": {"blobs": { + "pageInfo": {"hasNextPage": has_next, "endCursor": cursor}, "nodes": nodes, + }}}}}} + + +def _gql(paths, *, has_next=False, cursor=None): + return FakeResponse(200, json_data=_blobs_body(paths, has_next=has_next, cursor=cursor)) + + +def _entry(path, kind="blob"): + mode = {"blob": "100644", "tree": "040000", "commit": "160000"}[kind] + return {"id": "0" * 40, "name": path.rsplit("/", 1)[-1], "type": kind, "path": path, "mode": mode} + + +def _rest_page(entries, next_page=""): + return FakeResponse(200, json_data=entries, headers=CaseInsensitiveDict({"x-next-page": next_page})) + + +def _rest_params(page): + return {"recursive": "true", "per_page": 100, "page": page, "ref": SHA} + + +def _debug_lines(caplog): + records = [r for r in caplog.records if "list_paths" in r.getMessage()] + assert records, "no list_paths log line" + assert {r.levelno for r in records} == {logging.DEBUG} + return [r.getMessage() for r in records] + + +# --- the GraphQL walk ------------------------------------------------------------ + + +class TestGraphQLWalk: + def test_one_page_is_one_post_of_the_query_and_no_rest_call(self, monkeypatch): + monkeypatch.setenv("PRXREF_GITLAB_TOKEN", "t0ken") + session = FakeSession(graphql=[ + _gql(["src/b.py", "README.md", "src/a.py", "src/a.py"], has_next=False, cursor="NDA"), + ]) + + listing = _list(session) + + assert listing == PathListing(paths=("README.md", "src/a.py", "src/b.py"), complete=True) + assert session.verbs() == ["POST"] + url, kwargs = session.of("POST")[0] + assert url == GRAPHQL_URL + assert set(kwargs) == {"json", "headers", "timeout"} + assert set(kwargs["json"]) == {"query", "variables"} + assert kwargs["json"]["query"] == LIVE_QUERY + assert kwargs["json"]["variables"] == {"p": PROJECT, "ref": SHA, "after": None} + assert kwargs["headers"] == {"PRIVATE-TOKEN": "t0ken"} + assert kwargs["timeout"] == REQUEST_TIMEOUT + + def test_the_query_is_a_read_whose_first_word_is_query(self): + session = FakeSession(graphql=[_gql(["a.py"])]) + + _list(session) + + sent = session.of("POST")[0][1]["json"]["query"] + assert re.match(r"query\b", sent) + assert sent.split("(", 1)[0] == "query" + assert "mutation" not in sent + assert gitlab._TREE_BLOBS_QUERY == LIVE_QUERY + + def test_no_token_sends_no_auth_header(self): + session = FakeSession(graphql=[_gql(["a.py"])]) + + _list(session) + + assert session.of("POST")[0][1]["headers"] == {} + + def test_pages_are_followed_by_threading_each_end_cursor_into_after(self): + session = FakeSession(graphql=[ + _gql(["z/last.py", "a.py"], has_next=True, cursor="MTAw"), + _gql(["m/mid.py", "a.py"], has_next=True, cursor="MjAw"), + _gql(["b.py"], has_next=False, cursor="MjUw"), + ]) + + listing = _list(session) + + assert listing == PathListing(paths=("a.py", "b.py", "m/mid.py", "z/last.py"), complete=True) + posts = session.of("POST") + assert session.verbs() == ["POST", "POST", "POST"] + assert [kwargs["json"]["variables"]["after"] for _, kwargs in posts] == [None, "MTAw", "MjAw"] + assert {url for url, _ in posts} == {GRAPHQL_URL} + assert {kwargs["json"]["variables"]["p"] for _, kwargs in posts} == {PROJECT} + assert {kwargs["json"]["variables"]["ref"] for _, kwargs in posts} == {SHA} + assert {kwargs["json"]["query"] for _, kwargs in posts} == {LIVE_QUERY} + assert {kwargs["timeout"] for _, kwargs in posts} == {REQUEST_TIMEOUT} + + def test_two_pages_via_the_cursor(self): + session = FakeSession(graphql=[ + _gql(["one.py"], has_next=True, cursor="MTAw"), + _gql(["two.py"], has_next=False, cursor=None), + ]) + + listing = _list(session) + + assert listing == PathListing(paths=("one.py", "two.py"), complete=True) + assert [kw["json"]["variables"] for _, kw in session.of("POST")] == [ + {"p": PROJECT, "ref": SHA, "after": None}, + {"p": PROJECT, "ref": SHA, "after": "MTAw"}, + ] + assert session.of("GET") == [] + + def test_nodes_whose_path_is_not_a_non_empty_string_are_dropped(self): + nodes = [ + {"path": ".gitattributes", "type": "blob", "mode": "100644"}, + {"path": "bin/run", "type": "blob", "mode": "100755"}, + {"path": "", "type": "blob", "mode": "100644"}, + {"path": None, "type": "blob", "mode": "100644"}, + {"type": "blob", "mode": "100644"}, + {"path": 7}, + "not-a-dict", + None, + ] + session = FakeSession(graphql=[FakeResponse(200, json_data=_blobs_body([], nodes=nodes))]) + + listing = _list(session) + + assert listing == PathListing(paths=(".gitattributes", "bin/run"), complete=True) + assert session.verbs() == ["POST"] + + def test_a_later_empty_page_ends_the_walk_whole(self): + session = FakeSession(graphql=[ + _gql(["a.py"], has_next=True, cursor="MQ"), + FakeResponse(200, json_data=LIVE_EMPTY_PAGE), + ]) + + listing = _list(session) + + assert listing == PathListing(paths=("a.py",), complete=True) + assert session.verbs() == ["POST", "POST"] + + def test_the_post_goes_to_the_pr_host(self): + ref = _ref("https://gitlab.com/acme/tools/sub/-/merge_requests/3") + session = FakeSession(graphql=[_gql(["a.py"])]) + + _list(session, ref=ref) + + url, kwargs = session.of("POST")[0] + assert url == "https://gitlab.com/api/graphql" + assert kwargs["json"]["variables"]["p"] == "acme/tools/sub" + + def test_an_empty_errors_array_is_not_a_failure(self): + body = {"errors": [], **_blobs_body(["a.py"])} + session = FakeSession(graphql=[FakeResponse(200, json_data=body)]) + + assert _list(session) == PathListing(paths=("a.py",), complete=True) + assert session.verbs() == ["POST"] + + def test_an_empty_sha_makes_no_request(self): + session = FakeSession() + + assert _list(session, sha="") is None + assert session.calls == [] + + +# --- the page cap ------------------------------------------------------------------ + + +class TestGraphQLCap: + def test_the_cap_stops_the_walk_partial_with_no_rest_call(self, caplog): + cap = gitlab.MAX_LISTING_PAGES + pages = [_gql([f"p{n:02d}.py"], has_next=True, cursor=f"c{n}") for n in range(cap + 1)] + session = FakeSession(graphql=pages) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.gitlab"): + listing = _list(session) + + assert listing == PathListing(paths=tuple(f"p{n:02d}.py" for n in range(cap)), complete=False) + assert session.verbs() == ["POST"] * cap + assert [kw["json"]["variables"]["after"] for _, kw in session.of("POST")] == [ + None, *(f"c{n}" for n in range(cap - 1)), + ] + assert any("cap" in line and "GraphQL" in line for line in _debug_lines(caplog)) + + def test_a_walk_that_ends_exactly_at_the_cap_is_whole(self): + cap = gitlab.MAX_LISTING_PAGES + pages = [_gql([f"p{n:02d}.py"], has_next=n < cap - 1, cursor=f"c{n}") for n in range(cap)] + session = FakeSession(graphql=pages) + + listing = _list(session) + + assert listing == PathListing(paths=tuple(f"p{n:02d}.py" for n in range(cap)), complete=True) + assert session.verbs() == ["POST"] * cap + + def test_the_cap_is_read_from_the_module(self, monkeypatch): + monkeypatch.setattr(gitlab, "MAX_LISTING_PAGES", 2) + session = FakeSession(graphql=[ + _gql(["one.py"], has_next=True, cursor="c1"), + _gql(["two.py"], has_next=True, cursor="c2"), + _gql(["three.py"]), + ]) + + assert _list(session) == PathListing(paths=("one.py", "two.py"), complete=False) + assert session.verbs() == ["POST", "POST"] + + +# --- a failure after the first page -------------------------------------------------- + + +def _unusable(kind): + if kind == "transport": + return requests.ConnectionError("connection refused") + if kind == "timeout": + return requests.Timeout("read timed out") + if kind == "http-401": + return FakeResponse(401, json_data={"message": "401 Unauthorized"}) + if kind == "http-404": + return FakeResponse(404, json_data={"message": "404 Not Found"}) + if kind == "http-500": + return FakeResponse(500, text="internal error") + if kind == "non-json": + return FakeResponse(200, text="not json") + bodies = { + "list-body": [], + "null-body": None, + "errors": {"errors": [{"message": "Field 'blobs' doesn't exist on type 'Tree'"}]}, + "errors-with-data": {"errors": [{"message": "partial"}], **_blobs_body(["x.py"])}, + "project-null": LIVE_UNKNOWN_PROJECT, + "data-null": {"data": None}, + "no-data": {}, + "repository-null": {"data": {"project": {"repository": None}}}, + "tree-null": {"data": {"project": {"repository": {"tree": None}}}}, + "blobs-missing": {"data": {"project": {"repository": {"tree": {}}}}}, + "nodes-not-a-list": {"data": {"project": {"repository": {"tree": {"blobs": { + "pageInfo": {"hasNextPage": False, "endCursor": None}, "nodes": {"path": "a.py"}, + }}}}}}, + "page-info-missing": {"data": {"project": {"repository": {"tree": {"blobs": { + "nodes": [{"path": "a.py"}], + }}}}}}, + "has-next-not-a-bool": _blobs_body(["a.py"], has_next="false", cursor=None), + "has-next-without-cursor": _blobs_body(["a.py"], has_next=True, cursor=None), + "has-next-with-empty-cursor": _blobs_body(["a.py"], has_next=True, cursor=""), + } + return FakeResponse(200, json_data=bodies[kind]) + + +LATER_PAGE_FAILURES = [ + "transport", "timeout", "http-401", "http-500", "non-json", "list-body", "errors", + "project-null", "no-data", "nodes-not-a-list", "has-next-not-a-bool", "has-next-without-cursor", +] + + +class TestGraphQLLaterPageFailure: + @pytest.mark.parametrize("kind", LATER_PAGE_FAILURES) + def test_a_later_page_failure_keeps_the_paths_read_so_far(self, kind, caplog): + session = FakeSession(graphql=[ + _gql(["b.py", "a.py"], has_next=True, cursor="MTAw"), + _unusable(kind), + ]) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.gitlab"): + listing = _list(session) + + assert listing == PathListing(paths=("a.py", "b.py"), complete=False) + assert session.verbs() == ["POST", "POST"] + lines = _debug_lines(caplog) + assert any("page 2" in line and "GraphQL" in line for line in lines) + + def test_a_third_page_failure_keeps_both_pages(self): + session = FakeSession(graphql=[ + _gql(["a.py"], has_next=True, cursor="c1"), + _gql(["b.py"], has_next=True, cursor="c2"), + requests.ConnectionError("reset"), + ]) + + assert _list(session) == PathListing(paths=("a.py", "b.py"), complete=False) + assert session.verbs() == ["POST", "POST", "POST"] + + +# --- falling back to the REST walk --------------------------------------------------- + + +FIRST_PAGE_FALLBACKS = [ + "transport", "timeout", "http-401", "http-404", "http-500", "non-json", + "list-body", "null-body", "errors", "errors-with-data", "project-null", "data-null", + "no-data", "repository-null", "tree-null", "blobs-missing", "nodes-not-a-list", + "page-info-missing", "has-next-not-a-bool", "has-next-without-cursor", + "has-next-with-empty-cursor", +] + + +class TestFallbackToRest: + @pytest.mark.parametrize("kind", FIRST_PAGE_FALLBACKS) + def test_an_unusable_first_page_runs_the_rest_walk(self, kind, caplog): + session = FakeSession( + graphql=[_unusable(kind)], + rest=[_rest_page([_entry("src", kind="tree"), _entry("src/app.py")])], + ) + + with caplog.at_level(logging.DEBUG, logger="prxref.forges.gitlab"): + listing = _list(session) + + assert listing == PathListing(paths=("src/app.py",), complete=True) + assert session.verbs() == ["POST", "GET"] + url, kwargs = session.of("GET")[0] + assert url == TREE_URL + assert kwargs["params"] == _rest_params(1) + assert any("falls back to the REST tree walk" in line for line in _debug_lines(caplog)) + + @pytest.mark.parametrize( + "first", + [ + pytest.param(LIVE_EMPTY_PAGE, id="live-empty-page"), + pytest.param(_blobs_body([], cursor="NDA"), id="empty-page-with-a-cursor"), + pytest.param( + _blobs_body([], nodes=[{"path": ""}, {"type": "blob"}, "x"]), + id="no-usable-path", + ), + ], + ) + def test_an_empty_first_page_runs_the_rest_walk(self, first): + session = FakeSession( + graphql=[FakeResponse(200, json_data=first)], + rest=[_rest_page([_entry("a.py")])], + ) + + assert _list(session) == PathListing(paths=("a.py",), complete=True) + assert session.verbs() == ["POST", "GET"] + + def test_a_bad_sha_empty_first_page_with_a_rest_404_gives_none(self): + session = FakeSession( + graphql=[FakeResponse(200, json_data=LIVE_EMPTY_PAGE)], + rest=[FakeResponse(404, json_data={"message": "404 Tree Not Found"})], + ) + + assert _list(session) is None + assert session.verbs() == ["POST", "GET"] + + def test_an_empty_tree_gives_an_empty_complete_listing(self): + session = FakeSession( + graphql=[FakeResponse(200, json_data=LIVE_EMPTY_PAGE)], + rest=[_rest_page([])], + ) + + assert _list(session) == PathListing(paths=(), complete=True) + assert session.verbs() == ["POST", "GET"] + + def test_an_unknown_project_with_a_rest_404_gives_none(self): + session = FakeSession( + graphql=[FakeResponse(200, json_data=LIVE_UNKNOWN_PROJECT)], + rest=[FakeResponse(404, json_data={"message": "404 Project Not Found"})], + ) + + assert _list(session) is None + assert session.verbs() == ["POST", "GET"] + + def test_the_fallback_walk_pages_and_authenticates_as_before(self, monkeypatch): + monkeypatch.setenv("PRXREF_GITLAB_TOKEN", "t0ken") + session = FakeSession( + graphql=[FakeResponse(401, json_data={"message": "401 Unauthorized"})], + rest=[ + _rest_page([_entry("b.py"), _entry("lib", kind="tree")], next_page="2"), + _rest_page([_entry("a.py"), _entry("vendor/x", kind="commit")], next_page=""), + ], + ) + + listing = _list(session) + + assert listing == PathListing(paths=("a.py", "b.py"), complete=True) + assert session.verbs() == ["POST", "GET", "GET"] + assert session.of("POST")[0][1]["headers"] == {"PRIVATE-TOKEN": "t0ken"} + gets = session.of("GET") + assert [kw["params"] for _, kw in gets] == [_rest_params(1), _rest_params(2)] + assert {kw["headers"]["PRIVATE-TOKEN"] for _, kw in gets} == {"t0ken"} + assert {kw["timeout"] for _, kw in gets} == {REQUEST_TIMEOUT} + + def test_a_rest_first_page_failure_after_the_fallback_gives_none(self): + session = FakeSession( + graphql=[requests.ConnectionError("refused")], + rest=[requests.ConnectionError("refused")], + ) + + assert _list(session) is None + assert session.verbs() == ["POST", "GET"] + + +# --- the defect V1 observed live --------------------------------------------------- + + +def _directories_first_rest_listing(): + """REST as V1 saw it: every page up to the cap is directories, the files come after.""" + cap = gitlab.MAX_LISTING_PAGES + pages = [ + _rest_page( + [_entry(f"dir{page:02d}/sub{n:02d}", kind="tree") for n in range(100)], + next_page=str(page + 2), + ) + for page in range(cap) + ] + pages.append(_rest_page([_entry("README.md"), _entry("src/app.py")], next_page="")) + return pages + + +BLOBS = [f"dir{n:02d}/file{m}.py" for n in range(30) for m in range(10)] + + +class TestDirectoriesBeforeFiles: + def test_graphql_lists_the_files_the_rest_walk_never_reaches(self): + session = FakeSession( + graphql=[ + _gql(BLOBS[:100], has_next=True, cursor="MTAw"), + _gql(BLOBS[100:200], has_next=True, cursor="MjAw"), + _gql(BLOBS[200:], has_next=False, cursor="MzAw"), + ], + rest=_directories_first_rest_listing(), + ) + + listing = _list(session) + + assert listing == PathListing(paths=tuple(sorted(BLOBS)), complete=True) + assert len(listing.paths) == 300 + assert session.verbs() == ["POST"] * 3 + + def test_control_without_graphql_the_rest_walk_finds_no_file(self): + session = FakeSession( + graphql=[requests.ConnectionError("refused")], + rest=_directories_first_rest_listing(), + ) + + listing = _list(session) + + assert listing == PathListing(paths=(), complete=False) + assert session.verbs() == ["POST"] + ["GET"] * gitlab.MAX_LISTING_PAGES + + +# --- the documented contract ------------------------------------------------------- + + +class TestDocstring: + def test_the_docstring_names_both_walks_and_why_graphql_comes_first(self): + doc = " ".join((gitlab.ForgeImpl.list_paths.__doc__ or "").split()) + assert "tree.blobs" in doc + assert "every directory before any file" in doc + assert "repository/tree?recursive=true" in doc + assert "REST walk answers instead when the FIRST GraphQL page is unusable" in doc + assert "Never raises" in doc + assert "MAX_LISTING_PAGES" in doc From 3e467cedf2ec0fe6c0c33b75c59ec17ce4707701 Mon Sep 17 00:00:00 2001 From: Stephen Blatt <5125883+sblattj@users.noreply.github.com> Date: Thu, 24 Sep 2026 17:50:53 -0700 Subject: [PATCH 18/24] fix: retry an empty model reply once in the reviewer MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A reply with no text (or whitespace only) and a non-truncation finish reason failed its unit at parse time ("no parseable content") and was never asked again, though the provider billed it. The review of 0.15.0's own pull request lost 1 of 8 chunks this way. _invoke_and_parse now repeats that call once, with the same prompt and the same budget, after one WARNING naming the unit label and the finish reason. It never retries a reply stopped at the budget (length or max_tokens, the truncation path), a non-empty reply that fails to parse, or a call that raised. Tokens and elapsed_ms cover both calls, model is the second call's, and cost folds through costs.combine_reported: the sum when both calls reported a figure, else None. A second empty reply fails the unit with the same error as before. The chunk and the systemic sweep share the path. Worst case per chunk is 4 invoke calls with the orchestrator's timeout retry, 2 for the sweep; the module, review_chunk and review_systemic docstrings are corrected to say so. Tests: tests/test_empty_reply_retry.py (23 cases, including one run through orchestrate_review). 🤖 Authored with Claude Code — Claude Opus 5.5 via Claude Code Claude-Session-Id: f915b0ff-b27c-46f3-bc7c-430dca61d459 --- src/prxref/reviewer.py | 133 ++++++++++++-- tests/test_empty_reply_retry.py | 311 ++++++++++++++++++++++++++++++++ 2 files changed, 426 insertions(+), 18 deletions(-) create mode 100644 tests/test_empty_reply_retry.py diff --git a/src/prxref/reviewer.py b/src/prxref/reviewer.py index 4898684..2b9763a 100644 --- a/src/prxref/reviewer.py +++ b/src/prxref/reviewer.py @@ -1,8 +1,9 @@ """Worker-review layer: one LLM call per diff chunk, plus the systemic sweep. The reviewer renders ``prompts/worker.md`` with the chunk's unified diff, -makes a single :meth:`LLMClient.invoke` call (no retries — the model -fallback chain handles transient failures), and maps the JSON response +makes a single :meth:`LLMClient.invoke` call (the model fallback chain +handles transient failures; the one retry here is for a reply that came +back empty, made once with the same prompt), and maps the JSON response onto :class:`prxref.triage.Finding` records. Unparseable or malformed responses degrade to ``([], [])`` with a logged warning; this layer never raises. :func:`review_systemic` is the same contract over the @@ -60,7 +61,7 @@ from typing import Any from .chunk_context import sibling_summary_block -from .costs import valid_usd +from .costs import combine_reported, valid_usd from .forges.base import Thread from .llm import LLMClient from .parser import loads_lenient @@ -156,6 +157,55 @@ def _budget_stop_reason(result: Any) -> str: return literal if literal.lower() in _TRUNCATION_FINISH_REASONS else "" +def _is_empty_reply(result: Any) -> bool: + """True when a completion came back with no text: empty or whitespace only. + + Only a string is judged. A backend or test double whose ``text`` is not a + string keeps the path it always had, where the parse names what it got. + """ + text = getattr(result, "text", None) + return isinstance(text, str) and not text.strip() + + +def _finish_reason_text(result: Any) -> str: + """The provider's stop reason for a log line, stripped; ``-`` when none was reported.""" + reason = getattr(result, "finish_reason", "") + if not isinstance(reason, str) or not reason.strip(): + return "-" + return reason.strip() + + +def _result_cost(result: Any) -> tuple[float | None, str]: + """``(cost_usd, cost_source)`` as the backend reported them for one call. + + The figure goes through :func:`prxref.costs.valid_usd`; the source is + ``""`` whenever the figure is ``None``. + """ + cost = valid_usd(getattr(result, "cost_usd", None)) + source = str(getattr(result, "cost_source", "") or "") if cost is not None else "" + return cost, source + + +def _fold_retry_usage(meta: dict, result: Any) -> None: + """Add the empty-reply retry's usage to ``meta``, which holds the first call's. + + Both calls were billed, so ``input_tokens`` and ``output_tokens`` are + summed. ``model`` becomes the retry's, the reply the unit went on to use. + The reported cost folds through :func:`prxref.costs.combine_reported`, + the rule a backend applies to the attempts inside one invoke: the sum + when both calls reported a figure, else ``None`` with source ``""``. A + ``None`` leaves the unit to :func:`prxref.costs.run_cost`, which prices it + from the summed tokens when a price table knows the model. + """ + meta["input_tokens"] = (meta["input_tokens"] or 0) + (result.input_tokens or 0) + meta["output_tokens"] = (meta["output_tokens"] or 0) + (result.output_tokens or 0) + meta["model"] = result.model + meta["cost_usd"], meta["cost_source"] = combine_reported([ + (meta["cost_usd"], meta["cost_source"]), + _result_cost(result), + ]) + + def load_prompt(name: str) -> str: """Load a prompt template from the packaged ``prxref/prompts`` directory.""" fname = f"{name}.md" if not name.endswith(".md") else name @@ -458,6 +508,10 @@ def _write_trace_files( Each file lands via a temp file plus :func:`os.replace`, so a concurrent reader never observes a half-written file, and a timeout retry simply overwrites: the trace ends up showing the attempt whose result was used. + The empty-reply retry in :func:`_invoke_and_parse` writes only once, + after its final call, so the files describe that call (the meta's token + counts cover both calls) and the empty first reply is visible only in its + WARNING log line. Empty ``trace_dir`` is the declared off switch and does nothing — no directory, no syscalls, no cost. Every failure (directory cannot be created, path unwritable, disk full) is a logged warning and nothing more: @@ -518,6 +572,24 @@ def _invoke_and_parse( ``accept_rule`` does the same for ``rule`` (through :func:`prxref.triage.normalize_rule`); false, the default, leaves every finding's ``rule`` at ``None``. + + A reply that came back EMPTY (no text, or whitespace only) is asked for + again, once, with the same prompt and the same budget, after one WARNING + ``