Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 90 additions & 0 deletions pr-triage/tests/test_e2e.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ def __init__(self):
self.calls = [] # (method, path, body_dict_or_None)
self.model_reply = None # str content, or Exception subclass to raise, or dict for raw payload
self.model_status = 200
self.model_mode = "json" # json | redirect | html
self.files = [{"filename": "docs/architecture.md", "status": "modified"}]
self.diff = "--- a/docs/architecture.md\n+++ b/docs/architecture.md\n@@ -1 +1 @@\n-old\n+new\n"
self.permission = "admin"
Expand Down Expand Up @@ -72,12 +73,25 @@ def _route(self, method):

if path == "/v1/chat/completions":
s.model_headers = {k.lower(): v for k, v in self.headers.items()}
if s.model_mode == "redirect":
self.send_response(302)
self.send_header("Location", f"http://127.0.0.1:{self.server.server_address[1]}/cdn-cgi/access/login/llm.ionite.io?kid=abc")
self.send_header("Content-Length", "0")
self.end_headers()
return
if s.model_mode == "html":
page = "<!DOCTYPE html><html><head><title>Sign in · Cloudflare Access</title></head><body>" + "login " * 40 + "</body></html>"
assert len(page) > 200 # the 120-char cap must be measurable
return self._send(200, page, "text/html; charset=utf-8")
if s.model_status != 200:
return self._send(s.model_status, {"error": "nope"})
if isinstance(s.model_reply, dict):
return self._send(200, s.model_reply)
return self._send(200, {"choices": [{"message": {"role": "assistant", "content": s.model_reply}}]})

if path.startswith("/cdn-cgi/access/login"):
return self._send(200, "<html>followed the redirect</html>", "text/html")

prefix = "/github/repos/LykosAI/Test"
if not path.startswith(prefix):
return self._send(404, {"message": "unknown"})
Expand Down Expand Up @@ -317,6 +331,82 @@ def test_empty_content_with_only_reasoning_is_the_offline_path(self):
self.assertEqual(self.calls("POST", "/pulls/7/reviews"), [])
self.assertEqual(len(self.calls("POST", "/issues/7/comments")), 1)

def test_access_redirect_is_reported_as_a_302_and_never_followed(self):
self.state.model_mode = "redirect"
result = self.run_script()
self.assertEqual(result.returncode, 0, result.stderr)
self.assertEqual(self.calls("POST", "/pulls/7/reviews"), [])
self.assertIn("could not reach my model", self.calls("POST", "/issues/7/comments")[0][2]["body"])
self.assertEqual(len(self.calls("POST", "/v1/chat/completions")), 1)
self.assertEqual([c for c in self.state.calls if c[1].startswith("/cdn-cgi/")], [])
warning = [l for l in result.stdout.splitlines() if "::warning::Model unreachable" in l][0]
self.assertIn("ModelReplyError", warning)
self.assertIn("HTTP 302 redirect to 127.0.0.1", warning)
self.assertNotIn("kid=abc", warning)
for secret in ("cf-id-value", "cf-secret-value", "llm-key-value"):
self.assertNotIn(secret, result.stdout + result.stderr)

def test_html_200_is_reported_with_status_url_content_type_and_capped_body(self):
self.state.model_mode = "html"
result = self.run_script()
self.assertEqual(result.returncode, 0, result.stderr)
self.assertEqual(self.calls("POST", "/pulls/7/reviews"), [])
self.assertIn("could not reach my model", self.calls("POST", "/issues/7/comments")[0][2]["body"])
warning = [l for l in result.stdout.splitlines() if "::warning::Model unreachable" in l][0]
self.assertIn("non-JSON reply: HTTP 200 from http://127.0.0.1:", warning)
self.assertIn("/v1/chat/completions", warning)
self.assertIn("Content-Type text/html; charset=utf-8", warning)
self.assertIn("<!DOCTYPE html><html><head><title>Sign in", warning)
self.assertNotIn("</html>", warning) # capped well before the end of the body
for secret in ("cf-id-value", "cf-secret-value", "llm-key-value"):
self.assertNotIn(secret, result.stdout + result.stderr)

def test_credential_presence_line_carries_booleans_and_lengths_only(self):
self.state.model_reply = '{"verdict": "human", "reason": "hm"}'
result = self.run_script()
line = [l for l in result.stdout.splitlines() if l.startswith("Model credentials in the environment:")][0]
self.assertIn("CF_ACCESS_CLIENT_ID present=True len=11", line)
self.assertIn("CF_ACCESS_CLIENT_SECRET present=True len=15", line)
self.assertIn("LLM_API_KEY present=True len=13", line)
for secret in ("cf-id-value", "cf-secret-value", "llm-key-value"):
self.assertNotIn(secret, line)

def test_missing_credentials_show_as_absent(self):
self.state.model_reply = '{"verdict": "human", "reason": "hm"}'
env = dict(
os.environ,
GITHUB_EVENT_PATH=self.event_path,
GITHUB_REPOSITORY="LykosAI/Test",
GITHUB_API_URL=f"http://127.0.0.1:{self.port}/github",
GITHUB_TOKEN="ghs_fake",
TRIAGE_MODEL="fake-model",
TRIAGE_ENDPOINT=f"http://127.0.0.1:{self.port}/v1",
CF_ACCESS_CLIENT_ID="",
PYTHONIOENCODING="utf-8",
)
env.pop("CF_ACCESS_CLIENT_SECRET", None)
env.pop("LLM_API_KEY", None)
result = subprocess.run([sys.executable, SCRIPT], env=env, capture_output=True, text=True, encoding="utf-8", timeout=60)
line = [l for l in result.stdout.splitlines() if l.startswith("Model credentials in the environment:")][0]
self.assertIn("CF_ACCESS_CLIENT_ID present=False len=0", line)
self.assertIn("CF_ACCESS_CLIENT_SECRET present=False len=0", line)
self.assertIn("LLM_API_KEY present=False len=0", line)

def test_ask_model_raises_a_diagnosable_error_on_redirect_and_html(self):
sys.path.insert(0, os.path.join(HERE, ".."))
import triage

endpoint = f"http://127.0.0.1:{self.port}/v1"
self.state.model_mode = "redirect"
with self.assertRaises(triage.ModelReplyError) as ctx:
triage.ask_model(endpoint, "m", "k", "i", "s", "sys", "user")
self.assertIn("HTTP 302 redirect to 127.0.0.1", str(ctx.exception))
self.state.model_mode = "html"
with self.assertRaises(triage.ModelReplyError) as ctx:
triage.ask_model(endpoint, "m", "k", "i", "s", "sys", "user")
self.assertIn("non-JSON reply: HTTP 200", str(ctx.exception))
self.assertIn("text/html", str(ctx.exception))

def test_github_outage_exits_zero(self):
self.state.model_reply = '{"verdict": "approve", "reason": "ok"}'
env_override = f"http://127.0.0.1:{self.port}/nowhere"
Expand Down
84 changes: 67 additions & 17 deletions pr-triage/triage.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
import re
import sys
import urllib.error
import urllib.parse
import urllib.request
from dataclasses import dataclass, field
from typing import Any
Expand Down Expand Up @@ -315,10 +316,32 @@ def render_comment_body(


class HttpError(Exception):
def __init__(self, status: int, body: str):
def __init__(self, status: int, body: str, headers: dict[str, str] | None = None):
super().__init__(f"HTTP {status}: {body[:300]}")
self.status = status
self.body = body
self.headers = headers or {}


class ModelReplyError(Exception):
"""The model endpoint answered with something other than a chat completion.

The message is safe to log: status, final URL, content type, redirect host and a
capped body prefix; never a request header value.
"""


class _NoRedirect(urllib.request.HTTPRedirectHandler):
def redirect_request(self, req, fp, code, msg, headers, newurl):
return None


@dataclass(frozen=True)
class HttpResponse:
status: int
body: str
headers: dict[str, str]
url: str


def http(
Expand All @@ -327,18 +350,20 @@ def http(
headers: dict[str, str],
body: Any = None,
timeout: float = GITHUB_TIMEOUT_SECONDS,
) -> tuple[int, str, dict[str, str]]:
follow_redirects: bool = True,
) -> HttpResponse:
data = None
req_headers = dict(headers)
if body is not None:
data = json.dumps(body).encode("utf-8")
req_headers["Content-Type"] = "application/json"
req = urllib.request.Request(url, data=data, method=method, headers=req_headers)
opener = urllib.request.build_opener() if follow_redirects else urllib.request.build_opener(_NoRedirect())
try:
with urllib.request.urlopen(req, timeout=timeout) as resp:
return resp.status, resp.read().decode("utf-8", "replace"), dict(resp.headers)
with opener.open(req, timeout=timeout) as resp:
return HttpResponse(resp.status, resp.read().decode("utf-8", "replace"), dict(resp.headers), resp.geturl())
except urllib.error.HTTPError as e:
raise HttpError(e.code, e.read().decode("utf-8", "replace")) from e
raise HttpError(e.code, e.read().decode("utf-8", "replace"), dict(e.headers)) from e


class GitHub:
Expand All @@ -356,8 +381,7 @@ def _url(self, path: str) -> str:
return f"{self.api_url}/repos/{self.repo}{path}"

def get_json(self, path: str) -> Any:
_, text, _ = http("GET", self._url(path), self.headers)
return json.loads(text)
return json.loads(http("GET", self._url(path), self.headers).body)

def get_paginated(self, path: str) -> list[Any]:
items: list[Any] = []
Expand All @@ -372,8 +396,7 @@ def get_paginated(self, path: str) -> list[Any]:

def get_diff(self, number: int) -> str:
headers = dict(self.headers, Accept="application/vnd.github.diff")
_, text, _ = http("GET", self._url(f"/pulls/{number}"), headers)
return text
return http("GET", self._url(f"/pulls/{number}"), headers).body

def author_permission(self, login: str) -> str | None:
try:
Expand Down Expand Up @@ -433,6 +456,17 @@ def ask_model(
system_prompt: str,
user_content: str,
) -> str | None:
log(
"Model credentials in the environment: "
+ ", ".join(
f"{name} present={bool(value)} len={len(value)}"
for name, value in (
("CF_ACCESS_CLIENT_ID", cf_id),
("CF_ACCESS_CLIENT_SECRET", cf_secret),
("LLM_API_KEY", api_key),
)
)
)
headers = {
"Authorization": f"Bearer {api_key}",
"CF-Access-Client-Id": cf_id,
Expand All @@ -449,14 +483,30 @@ def ask_model(
"max_tokens": 2500,
"reasoning_effort": "low",
}
_, text, _ = http(
"POST",
endpoint.rstrip("/") + "/chat/completions",
headers,
payload,
timeout=MODEL_TIMEOUT_SECONDS,
)
data = json.loads(text)
try:
resp = http(
"POST",
endpoint.rstrip("/") + "/chat/completions",
headers,
payload,
timeout=MODEL_TIMEOUT_SECONDS,
follow_redirects=False,
)
except HttpError as e:
if 300 <= e.status < 400:
location = e.headers.get("Location") or e.headers.get("location") or ""
host = urllib.parse.urlsplit(location).netloc or "(no Location header)"
raise ModelReplyError(f"HTTP {e.status} redirect to {host}; an Access login page means the service token was not accepted") from e
raise
content_type = resp.headers.get("Content-Type") or resp.headers.get("content-type") or "(none)"
try:
data = json.loads(resp.body)
except json.JSONDecodeError as e:
raise ModelReplyError(
f"non-JSON reply: HTTP {resp.status} from {resp.url}, Content-Type {content_type}, body starts {resp.body[:120]!r}"
) from e
if not isinstance(data, dict):
raise ModelReplyError(f"JSON reply is not an object: HTTP {resp.status} from {resp.url}, body starts {resp.body[:120]!r}")
choices = data.get("choices") or []
if not choices:
return None
Expand Down
Loading