feat(dashboard): drop token auth, trust oauth2-proxy X-Forwarded-User

This commit is contained in:
Claude
2026-08-22 18:03:33 +00:00
parent c9809bce9b
commit 99255bb167
4 changed files with 141 additions and 203 deletions
+35 -73
View File
@@ -9,13 +9,13 @@ ThreadingHTTPServer shape) that renders three views off the feedback SQLite:
GET /r/<owner>/<name>/<index> one PR's findings GET /r/<owner>/<name>/<index> one PR's findings
GET /r/<owner>/<name>/<index>/raw raw Markdown body (via Gitea) GET /r/<owner>/<name>/<index>/raw raw Markdown body (via Gitea)
GET /static/style.css CSS GET /static/style.css CSS
GET /login login form
POST /login check token, set cookie
POST /r/<owner>/<name>/edit mutate .pr-review.json (Tasks C+D) POST /r/<owner>/<name>/edit mutate .pr-review.json (Tasks C+D)
Auth: when `PRAGENT_DASHBOARD_TOKEN` is set, every route except `/login`, Auth: oauth2-proxy fronts this service in-cluster. Every route except
`/static/*`, and `POST /login` requires `Cookie: pragent_dash=<token>`. `/static/*` requires the `X-Forwarded-User` header (set by oauth2-proxy
Unset → tailnet-only mode, no auth. once the user has logged in via Logto). Missing header → 401 +
`WWW-Authenticate: Basic realm="pragent-dashboard"` so oauth2-proxy
intercepts the response.
DB: `PRAGENT_FEEDBACK_DB` points at the SQLite file the webhook server DB: `PRAGENT_FEEDBACK_DB` points at the SQLite file the webhook server
also writes. Per-request open (SQLite is cheap, no concurrency hazard, also writes. Per-request open (SQLite is cheap, no concurrency hazard,
@@ -33,6 +33,7 @@ import datetime
import html import html
import json import json
import os import os
import secrets
import string import string
import urllib.error import urllib.error
import urllib.parse import urllib.parse
@@ -47,13 +48,15 @@ from pilot import dashboard_data
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
FEEDBACK_DB = "" # legacy; readers should call _feedback_db() FEEDBACK_DB = "" # legacy; readers should call _feedback_db()
DASHBOARD_TOKEN = "" # legacy; readers should call _dashboard_token()
PORT = int(os.environ.get("DASHBOARD_PORT", "8081")) PORT = int(os.environ.get("DASHBOARD_PORT", "8081"))
GITEA_API = "" # legacy; readers should call _gitea_api() GITEA_API = "" # legacy; readers should call _gitea_api()
BOT_TOKEN = "" # legacy; readers should call _bot_token() BOT_TOKEN = "" # legacy; readers should call _bot_token()
COOKIE_NAME = "pragent_dash" # CSRF secret for the edit form. Regenerated per process (each Python
# interpreter launch). Behind oauth2-proxy this is enough — only an
# already-authenticated same-tab request can read this and echo it back.
_CSRF_SECRET: str = secrets.token_urlsafe(24)
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -67,10 +70,6 @@ def _feedback_db() -> str:
return os.environ.get("PRAGENT_FEEDBACK_DB", "") return os.environ.get("PRAGENT_FEEDBACK_DB", "")
def _dashboard_token() -> str:
return os.environ.get("PRAGENT_DASHBOARD_TOKEN", "")
def _bot_token() -> str: def _bot_token() -> str:
return os.environ.get("PRAGENT_BOT_TOKEN", "") return os.environ.get("PRAGENT_BOT_TOKEN", "")
@@ -231,18 +230,6 @@ ${findings_table}
""") """)
_LOGIN = string.Template("""<h2>Login</h2>
${flash}
<form method="post" action="/login">
<label for="token">Dashboard token</label>
<input type="text" id="token" name="token" autofocus>
<div class="row">
<button type="submit">Sign in</button>
</div>
</form>
""")
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Small helpers # Small helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -418,10 +405,6 @@ def _pr_body(data: dict, *, repo_url: str) -> str:
) )
def _login_body(flash: str = "") -> str:
return _LOGIN.substitute(flash=_esc(flash))
def _page(title: str, body: str, *, repos_first: str = "") -> str: def _page(title: str, body: str, *, repos_first: str = "") -> str:
db_status = _feedback_db() or "(no DB configured)" db_status = _feedback_db() or "(no DB configured)"
return _BASE.substitute( return _BASE.substitute(
@@ -467,16 +450,14 @@ def _http(method: str, url: str, *, token: str = "", body: dict | None = None,
def _is_authed(headers) -> bool: def _is_authed(headers) -> bool:
"""True when token matches; false when token unset OR cookie absent/wrong.""" """True when oauth2-proxy forwarded a verified user.
tok = _dashboard_token()
if not tok: oauth2-proxy sets `X-Forwarded-User` (and friends) only after a
return True # no token → tailnet-only mode successful Logto login + email allowlist check. Unauthenticated
cookie = headers.get("Cookie") or "" requests never see the header, so the dashboard never has to know
for part in cookie.split(";"): about cookies, secrets, or Logto's token shape.
k, _, v = part.strip().partition("=") """
if k == COOKIE_NAME and v == tok: return bool((headers.get("X-Forwarded-User") or "").strip())
return True
return False
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -505,11 +486,10 @@ def _route_repo(owner: str, name: str) -> bytes:
current_model = cfg.get("model", "") current_model = cfg.get("model", "")
elif err and err != "404": elif err and err != "404":
flash = f"could not read .pr-review.json: {err}" flash = f"could not read .pr-review.json: {err}"
csrf = _dashboard_token() or "noauth"
body = _repo_body( body = _repo_body(
data, data,
repo_url=repo_url, repo_url=repo_url,
csrf=csrf, csrf=_CSRF_SECRET,
current_model=current_model, current_model=current_model,
current_static_message=current_static_message, current_static_message=current_static_message,
flash=flash, flash=flash,
@@ -549,25 +529,11 @@ def _route_static_css() -> bytes:
return STYLE_CSS.encode() return STYLE_CSS.encode()
def _route_login_get() -> bytes:
return _page("login", _login_body()).encode()
def _route_login_post(token: str) -> tuple[int, dict, bytes]:
"""Compare token to env. Sets cookie on match, 401 otherwise."""
expected = _dashboard_token()
if not expected or token != expected:
return 401, {}, b"invalid token"
cookie = f"{COOKIE_NAME}={expected}; HttpOnly; SameSite=Strict; Path=/"
return 302, {"Set-Cookie": cookie, "Location": "/"}, b""
def _route_edit(owner: str, name: str, form: dict) -> tuple[int, dict, bytes]: def _route_edit(owner: str, name: str, form: dict) -> tuple[int, dict, bytes]:
"""Mutate .pr-review.json via the Gitea contents API (Tasks C+D).""" """Mutate .pr-review.json via the Gitea contents API (Tasks C+D)."""
repo_url = f"{owner}/{name}" repo_url = f"{owner}/{name}"
csrf = form.get("_csrf", "") csrf = form.get("_csrf", "")
expected_csrf = _dashboard_token() or "noauth" if csrf != _CSRF_SECRET:
if csrf != expected_csrf:
return 302, {"Location": f"/r/{repo_url}"}, b"" return 302, {"Location": f"/r/{repo_url}"}, b""
static_message = (form.get("static_message") or "").strip()[:400] static_message = (form.get("static_message") or "").strip()[:400]
model = (form.get("model") or "").strip() model = (form.get("model") or "").strip()
@@ -682,19 +648,27 @@ class Handler(BaseHTTPRequestHandler):
self.end_headers() self.end_headers()
self.wfile.write(body) self.wfile.write(body)
def _unauthorized(self) -> None:
"""401 + Basic challenge so oauth2-proxy intercepts and redirects to Logto."""
body = b"unauthorized\n"
self.send_response(401)
self.send_header("Content-Type", "text/plain; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.send_header("WWW-Authenticate", 'Basic realm="pragent-dashboard"')
self.end_headers()
self.wfile.write(body)
# --- GET ----------------------------------------------------------------- # --- GET -----------------------------------------------------------------
def do_GET(self): def do_GET(self):
path = self.path path = self.path
# Static is exempt from auth so the login page can load its CSS. # Static is exempt from auth (also unauthenticated browser fingerprinting
# noise, but it's the same CSS regardless of viewer).
if path == "/static/style.css": if path == "/static/style.css":
self._send(200, _route_static_css(), content_type="text/css; charset=utf-8") self._send(200, _route_static_css(), content_type="text/css; charset=utf-8")
return return
if path == "/login":
self._send(200, _route_login_get())
return
if not _is_authed(self.headers): if not _is_authed(self.headers):
self._redirect("/login") self._unauthorized()
return return
if path == "/" or path == "": if path == "/" or path == "":
@@ -724,17 +698,8 @@ class Handler(BaseHTTPRequestHandler):
def do_POST(self): def do_POST(self):
path = self.path path = self.path
# Login POST is exempt from auth (obviously).
if path == "/login":
length = int(self.headers.get("Content-Length", "0") or "0")
raw = self.rfile.read(length) if length else b""
form = urllib.parse.parse_qs(raw.decode("utf-8", errors="replace"))
token = (form.get("token") or [""])[0]
status, extra, body = _route_login_post(token)
self._send(status, body, content_type="text/plain", extra_headers=extra)
return
if not _is_authed(self.headers): if not _is_authed(self.headers):
self._redirect("/login") self._unauthorized()
return return
# /r/<owner>/<name>/edit # /r/<owner>/<name>/edit
m = _EDIT_RE.match(path) m = _EDIT_RE.match(path)
@@ -775,10 +740,7 @@ def main() -> int:
if not _feedback_db(): if not _feedback_db():
print("pragent-dashboard: WARNING: PRAGENT_FEEDBACK_DB not set; dashboard will be empty", print("pragent-dashboard: WARNING: PRAGENT_FEEDBACK_DB not set; dashboard will be empty",
flush=True) flush=True)
if _dashboard_token(): print("pragent-dashboard: auth via oauth2-proxy (X-Forwarded-User required)", flush=True)
print("pragent-dashboard: auth ON (PRAGENT_DASHBOARD_TOKEN set)", flush=True)
else:
print("pragent-dashboard: auth OFF (tailnet-only mode)", flush=True)
server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler) server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler)
print(f"pragent-dashboard: listening on :{PORT}", flush=True) print(f"pragent-dashboard: listening on :{PORT}", flush=True)
try: try:
+70 -98
View File
@@ -1,8 +1,9 @@
"""Tests for pilot/dashboard.py — stdlib HTTP server rendering dashboard HTML. """Tests for pilot/dashboard.py — stdlib HTTP server rendering dashboard HTML.
We spin up the server on an ephemeral port in setUp, drive it with We spin up the server on an ephemeral port in setUp, drive it with
http.client, and tear it down in tearDown. Tests are split between http.client, and tear it down in tearDown. Auth is now performed by
no-auth (default — env unset) and auth-on (PRAGENT_DASHBOARD_TOKEN set). oauth2-proxy: the dashboard trusts `X-Forwarded-User` set by the proxy
and returns 401 (with a Basic challenge) when the header is missing.
The dashboard reads `PRAGENT_FEEDBACK_DB` and renders views via The dashboard reads `PRAGENT_FEEDBACK_DB` and renders views via
`dashboard_data`. We seed an in-memory SQLite at `tmp_path` for each `dashboard_data`. We seed an in-memory SQLite at `tmp_path` for each
@@ -69,15 +70,72 @@ def _post(port: int, path: str, body: bytes, headers: dict | None = None) -> tup
return r.status, h, body_b return r.status, h, body_b
class TestDashboardNoAuth(unittest.TestCase): class TestDashboardAuth(unittest.TestCase):
"""PRAGENT_DASHBOARD_TOKEN unset → no cookie check.""" """Auth gate: require X-Forwarded-User (set by oauth2-proxy).
When the header is missing every non-static route returns 401 with a
Basic challenge, which lets oauth2-proxy redirect the browser to
Logto. Static is exempt so the unauthenticated probe traffic doesn't
loop the proxy through the auth flow.
"""
def setUp(self): def setUp(self):
self.tmp = tempfile.TemporaryDirectory() self.tmp = tempfile.TemporaryDirectory()
self.db = f"{self.tmp.name}/f.db" self.db = f"{self.tmp.name}/f.db"
os.environ["PRAGENT_FEEDBACK_DB"] = self.db os.environ["PRAGENT_FEEDBACK_DB"] = self.db
os.environ["PRAGENT_DASHBOARD_TOKEN"] = ""
os.environ["DASHBOARD_PORT"] = str(0) # we override below os.environ["DASHBOARD_PORT"] = str(0) # we override below
self.port = _free_port()
from http.server import ThreadingHTTPServer
self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler)
self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True)
self.thread.start()
def tearDown(self):
self.srv.shutdown()
self.srv.server_close()
self.thread.join(timeout=2)
for k in ("PRAGENT_FEEDBACK_DB", "DASHBOARD_PORT"):
os.environ.pop(k, None)
self.tmp.cleanup()
def test_anonymous_overview_returns_401_with_basic_challenge(self):
status, h, body = _get(self.port, "/")
self.assertEqual(status, 401)
self.assertEqual(h.get("WWW-Authenticate"), 'Basic realm="pragent-dashboard"')
self.assertEqual(body, b"unauthorized\n")
def test_anonymous_repo_returns_401(self):
status, _h, _body = _get(self.port, "/r/alpha/one")
self.assertEqual(status, 401)
def test_anonymous_post_returns_401(self):
status, _h, _body = _post(self.port, "/r/alpha/one/edit", b"x=1")
self.assertEqual(status, 401)
def test_authenticated_overview_succeeds(self):
status, h, body = _get(self.port, "/", headers={"X-Forwarded-User": "marcos@example.com"})
self.assertEqual(status, 200)
self.assertIn(b"Overview", body)
def test_static_does_not_require_auth(self):
status, h, body = _get(self.port, "/static/style.css")
self.assertEqual(status, 200)
self.assertIn("text/css", h.get("Content-Type", ""))
self.assertGreater(len(body), 50)
def test_empty_x_forwarded_user_treated_as_anonymous(self):
status, _h, _body = _get(self.port, "/", headers={"X-Forwarded-User": " "})
self.assertEqual(status, 401)
class TestDashboardRender(unittest.TestCase):
"""Render-only tests — X-Forwarded-User set, real seeded data."""
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.db = f"{self.tmp.name}/f.db"
os.environ["PRAGENT_FEEDBACK_DB"] = self.db
# Seed: 2 repos, a couple of reviews + findings each. # Seed: 2 repos, a couple of reviews + findings each.
conn = feedback.init(self.db) conn = feedback.init(self.db)
for repo, n_prs in (("alpha/one", 2), ("beta/two", 1)): for repo, n_prs in (("alpha/one", 2), ("beta/two", 1)):
@@ -99,22 +157,21 @@ class TestDashboardNoAuth(unittest.TestCase):
conn.close() conn.close()
self.port = _free_port() self.port = _free_port()
# Re-bind the handler's server-class addr by constructing a fresh server.
from http.server import ThreadingHTTPServer from http.server import ThreadingHTTPServer
self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler) self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler)
self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True) self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True)
self.thread.start() self.thread.start()
self.auth_hdr = {"X-Forwarded-User": "marcos@example.com"}
def tearDown(self): def tearDown(self):
self.srv.shutdown() self.srv.shutdown()
self.srv.server_close() self.srv.server_close()
self.thread.join(timeout=2) self.thread.join(timeout=2)
for k in ("PRAGENT_FEEDBACK_DB", "PRAGENT_DASHBOARD_TOKEN", "DASHBOARD_PORT"): os.environ.pop("PRAGENT_FEEDBACK_DB", None)
os.environ.pop(k, None)
self.tmp.cleanup() self.tmp.cleanup()
def test_overview_200_contains_repo_names(self): def test_overview_200_contains_repo_names(self):
status, _h, body = _get(self.port, "/") status, _h, body = _get(self.port, "/", headers=self.auth_hdr)
self.assertEqual(status, 200) self.assertEqual(status, 200)
text = body.decode() text = body.decode()
self.assertIn("Overview", text) self.assertIn("Overview", text)
@@ -122,7 +179,7 @@ class TestDashboardNoAuth(unittest.TestCase):
self.assertIn("beta/two", text) self.assertIn("beta/two", text)
def test_repo_page_200(self): def test_repo_page_200(self):
status, _h, body = _get(self.port, "/r/alpha/one") status, _h, body = _get(self.port, "/r/alpha/one", headers=self.auth_hdr)
self.assertEqual(status, 200) self.assertEqual(status, 200)
text = body.decode() text = body.decode()
self.assertIn("alpha/one", text) self.assertIn("alpha/one", text)
@@ -130,102 +187,17 @@ class TestDashboardNoAuth(unittest.TestCase):
self.assertIn("src/file_0.py", text) self.assertIn("src/file_0.py", text)
def test_pr_page_200(self): def test_pr_page_200(self):
status, _h, body = _get(self.port, "/r/alpha/one/1") status, _h, body = _get(self.port, "/r/alpha/one/1", headers=self.auth_hdr)
self.assertEqual(status, 200) self.assertEqual(status, 200)
text = body.decode() text = body.decode()
self.assertIn("alpha/one", text) self.assertIn("alpha/one", text)
self.assertIn("#1", text) self.assertIn("#1", text)
self.assertIn("src/file_0.py", text) self.assertIn("src/file_0.py", text)
def test_static_css_200(self):
status, h, body = _get(self.port, "/static/style.css")
self.assertEqual(status, 200)
self.assertIn("text/css", h.get("Content-Type", ""))
self.assertGreater(len(body), 50)
def test_login_form(self):
status, _h, body = _get(self.port, "/login")
self.assertEqual(status, 200)
self.assertIn(b'name="token"', body)
def test_login_post_wrong_token_401(self):
status, _h, body = _post(
self.port, "/login", b"token=wrong",
)
self.assertEqual(status, 401)
def test_unknown_route_404(self): def test_unknown_route_404(self):
status, _h, _body = _get(self.port, "/no/such/route") status, _h, _body = _get(self.port, "/no/such/route", headers=self.auth_hdr)
self.assertEqual(status, 404) self.assertEqual(status, 404)
class TestDashboardWithAuth(unittest.TestCase):
"""PRAGENT_DASHBOARD_TOKEN set → all routes (except /login) require cookie."""
def setUp(self):
self.tmp = tempfile.TemporaryDirectory()
self.db = f"{self.tmp.name}/f.db"
os.environ["PRAGENT_FEEDBACK_DB"] = self.db
os.environ["PRAGENT_DASHBOARD_TOKEN"] = "secret-token-xyz"
# Seed one repo with a single PR.
conn = feedback.init(self.db)
rid = feedback.record_review(
conn, repo="o/r", pr=1, head_sha="x",
review_id_gitea=1, body_comment_id=2,
)
feedback.record_inline_finding(
conn, review_id=rid, repo="o/r", pr=1,
path="a.py", line=1, severity="low", problem="x",
)
conn.close()
self.port = _free_port()
from http.server import ThreadingHTTPServer
self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler)
self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True)
self.thread.start()
def tearDown(self):
self.srv.shutdown()
self.srv.server_close()
self.thread.join(timeout=2)
for k in ("PRAGENT_FEEDBACK_DB", "PRAGENT_DASHBOARD_TOKEN"):
os.environ.pop(k, None)
self.tmp.cleanup()
def test_overview_without_cookie_redirects_to_login(self):
status, h, _b = _get(self.port, "/")
self.assertEqual(status, 302)
self.assertIn("/login", h.get("Location", ""))
def test_login_with_correct_token_sets_cookie_and_redirects(self):
status, h, _b = _post(
self.port, "/login", b"token=secret-token-xyz",
)
self.assertEqual(status, 302)
self.assertIn("/", h.get("Location", ""))
sc = h.get("Set-Cookie", "")
self.assertIn("pragent_dash=secret-token-xyz", sc)
self.assertIn("HttpOnly", sc)
self.assertIn("SameSite=Strict", sc)
def test_overview_with_correct_cookie_succeeds(self):
status, _h, body = _get(
self.port, "/",
headers={"Cookie": "pragent_dash=secret-token-xyz"},
)
self.assertEqual(status, 200)
self.assertIn(b"Overview", body)
def test_login_form_does_not_require_auth(self):
status, _h, _b = _get(self.port, "/login")
self.assertEqual(status, 200)
def test_static_does_not_require_auth(self):
status, _h, _b = _get(self.port, "/static/style.css")
self.assertEqual(status, 200)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+33 -29
View File
@@ -2,7 +2,7 @@
We mock the Gitea HTTP layer (urllib.request.urlopen) so the test never We mock the Gitea HTTP layer (urllib.request.urlopen) so the test never
touches the network. The dashboard handler is responsible for: touches the network. The dashboard handler is responsible for:
* auth + CSRF * auth (X-Forwarded-User set by oauth2-proxy) + CSRF
* read .pr-review.json via GET (404 → start from {}) * read .pr-review.json via GET (404 → start from {})
* validate model against cost_model.PRICES * validate model against cost_model.PRICES
* PUT the updated file back, with sha + base64 content * PUT the updated file back, with sha + base64 content
@@ -35,12 +35,12 @@ def _free_port() -> int:
return port return port
def _post(port: int, path: str, body: bytes, *, cookie: str | None = None) -> tuple[int, dict, bytes]: def _post(port: int, path: str, body: bytes, *, headers: dict | None = None) -> tuple[int, dict, bytes]:
conn = http.client.HTTPConnection("127.0.0.1", port, timeout=5) conn = http.client.HTTPConnection("127.0.0.1", port, timeout=5)
headers = {"Content-Type": "application/x-www-form-urlencoded"} hdrs = {"Content-Type": "application/x-www-form-urlencoded"}
if cookie: if headers:
headers["Cookie"] = cookie hdrs.update(headers)
conn.request("POST", path, body=body, headers=headers) conn.request("POST", path, body=body, headers=hdrs)
r = conn.getresponse() r = conn.getresponse()
body_b = r.read() body_b = r.read()
h = dict(r.getheaders()) h = dict(r.getheaders())
@@ -68,7 +68,6 @@ class TestDashboardEdit(unittest.TestCase):
self.tmp = tempfile.TemporaryDirectory() self.tmp = tempfile.TemporaryDirectory()
self.db = f"{self.tmp.name}/f.db" self.db = f"{self.tmp.name}/f.db"
os.environ["PRAGENT_FEEDBACK_DB"] = self.db os.environ["PRAGENT_FEEDBACK_DB"] = self.db
os.environ["PRAGENT_DASHBOARD_TOKEN"] = "secret-edit-token"
os.environ["PRAGENT_BOT_TOKEN"] = "bot-token" os.environ["PRAGENT_BOT_TOKEN"] = "bot-token"
# Seed a row so the repo page is meaningful. # Seed a row so the repo page is meaningful.
conn = feedback.init(self.db) conn = feedback.init(self.db)
@@ -82,7 +81,10 @@ class TestDashboardEdit(unittest.TestCase):
self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler) self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler)
self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True) self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True)
self.thread.start() self.thread.start()
self.cookie = "pragent_dash=secret-edit-token" # Pull the per-process CSRF secret from the rendered repo page — the
# edit form embeds the same token as a hidden input.
self.csrf = dash._CSRF_SECRET
self.auth_hdr = {"X-Forwarded-User": "marcos@example.com"}
# Records of HTTP calls made by the handler. # Records of HTTP calls made by the handler.
self.calls: list[tuple[str, str, dict | None, bytes | None]] = [] self.calls: list[tuple[str, str, dict | None, bytes | None]] = []
@@ -91,7 +93,7 @@ class TestDashboardEdit(unittest.TestCase):
self.srv.shutdown() self.srv.shutdown()
self.srv.server_close() self.srv.server_close()
self.thread.join(timeout=2) self.thread.join(timeout=2)
for k in ("PRAGENT_FEEDBACK_DB", "PRAGENT_DASHBOARD_TOKEN", "PRAGENT_BOT_TOKEN"): for k in ("PRAGENT_FEEDBACK_DB", "PRAGENT_BOT_TOKEN"):
os.environ.pop(k, None) os.environ.pop(k, None)
self.tmp.cleanup() self.tmp.cleanup()
@@ -114,12 +116,12 @@ class TestDashboardEdit(unittest.TestCase):
def test_edit_updates_static_message_and_model(self): def test_edit_updates_static_message_and_model(self):
form = ( form = (
b"_csrf=secret-edit-token" f"_csrf={self.csrf}"
b"&static_message=Hello%20world" f"&static_message=Hello%20world"
b"&model=claude-sonnet-5" f"&model=claude-sonnet-5"
) ).encode()
with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen): with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen):
status, h, _b = _post(self.port, "/r/o/r/edit", form, cookie=self.cookie) status, h, _b = _post(self.port, "/r/o/r/edit", form, headers=self.auth_hdr)
self.assertEqual(status, 302) self.assertEqual(status, 302)
self.assertEqual(h.get("Location"), "/r/o/r") self.assertEqual(h.get("Location"), "/r/o/r")
@@ -139,24 +141,24 @@ class TestDashboardEdit(unittest.TestCase):
def test_edit_strips_static_message_to_400(self): def test_edit_strips_static_message_to_400(self):
long_msg = "x" * 600 long_msg = "x" * 600
form = ( form = (
f"_csrf=secret-edit-token" f"_csrf={self.csrf}"
f"&static_message={long_msg}" f"&static_message={long_msg}"
f"&model=claude-haiku-4-5" f"&model=claude-haiku-4-5"
).encode() ).encode()
with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen): with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen):
_post(self.port, "/r/o/r/edit", form, cookie=self.cookie) _post(self.port, "/r/o/r/edit", form, headers=self.auth_hdr)
put = next(c for c in self.calls if c[0] == "PUT") put = next(c for c in self.calls if c[0] == "PUT")
cfg = json.loads(base64.b64decode(json.loads(put[3])["content"])) cfg = json.loads(base64.b64decode(json.loads(put[3])["content"]))
self.assertEqual(len(cfg["static_message"]), 400) self.assertEqual(len(cfg["static_message"]), 400)
def test_edit_rejects_unknown_model_with_flash(self): def test_edit_rejects_unknown_model_with_flash(self):
form = ( form = (
b"_csrf=secret-edit-token" f"_csrf={self.csrf}"
b"&static_message=hi" f"&static_message=hi"
b"&model=does-not-exist" f"&model=does-not-exist"
) ).encode()
with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen): with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen):
status, h, _b = _post(self.port, "/r/o/r/edit", form, cookie=self.cookie) status, h, _b = _post(self.port, "/r/o/r/edit", form, headers=self.auth_hdr)
self.assertEqual(status, 302) self.assertEqual(status, 302)
self.assertIn("flash=", h.get("Location", "")) self.assertIn("flash=", h.get("Location", ""))
# No PUT should have been issued. # No PUT should have been issued.
@@ -164,16 +166,18 @@ class TestDashboardEdit(unittest.TestCase):
self.assertEqual(put_calls, []) self.assertEqual(put_calls, [])
def test_edit_requires_auth(self): def test_edit_requires_auth(self):
form = b"_csrf=secret-edit-token&static_message=x&model=claude-haiku-4-5" form = f"_csrf={self.csrf}&static_message=x&model=claude-haiku-4-5".encode()
with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen): with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen):
status, h, _b = _post(self.port, "/r/o/r/edit", form) status, h, _b = _post(self.port, "/r/o/r/edit", form)
self.assertEqual(status, 302) self.assertEqual(status, 401)
self.assertEqual(h.get("Location"), "/login") self.assertEqual(h.get("WWW-Authenticate"), 'Basic realm="pragent-dashboard"')
# No Gitea calls at all — auth gate fires first.
self.assertEqual(self.calls, [])
def test_edit_csrf_mismatch_redirects_without_save(self): def test_edit_csrf_mismatch_redirects_without_save(self):
form = b"_csrf=wrong&static_message=x&model=claude-haiku-4-5" form = f"_csrf=wrong&static_message=x&model=claude-haiku-4-5".encode()
with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen): with patch.object(dash.urllib.request, "urlopen", side_effect=self._urlopen):
status, h, _b = _post(self.port, "/r/o/r/edit", form, cookie=self.cookie) status, h, _b = _post(self.port, "/r/o/r/edit", form, headers=self.auth_hdr)
self.assertEqual(status, 302) self.assertEqual(status, 302)
self.assertEqual(h.get("Location"), "/r/o/r") self.assertEqual(h.get("Location"), "/r/o/r")
put_calls = [c for c in self.calls if c[0] == "PUT"] put_calls = [c for c in self.calls if c[0] == "PUT"]
@@ -192,9 +196,9 @@ class TestDashboardEdit(unittest.TestCase):
return _FakeResp(201, b"{}") return _FakeResp(201, b"{}")
return _FakeResp(404, b"") return _FakeResp(404, b"")
form = b"_csrf=secret-edit-token&static_message=hi&model=claude-haiku-4-5" form = f"_csrf={self.csrf}&static_message=hi&model=claude-haiku-4-5".encode()
with patch.object(dash.urllib.request, "urlopen", side_effect=_route): with patch.object(dash.urllib.request, "urlopen", side_effect=_route):
status, h, _b = _post(self.port, "/r/o/r/edit", form, cookie=self.cookie) status, h, _b = _post(self.port, "/r/o/r/edit", form, headers=self.auth_hdr)
self.assertEqual(status, 302) self.assertEqual(status, 302)
put = next(c for c in self.calls if c[0] == "PUT") put = next(c for c in self.calls if c[0] == "PUT")
payload = json.loads(put[3]) payload = json.loads(put[3])
@@ -202,4 +206,4 @@ class TestDashboardEdit(unittest.TestCase):
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+3 -3
View File
@@ -38,7 +38,6 @@ class TestRepoEditSelect(unittest.TestCase):
self.tmp = tempfile.TemporaryDirectory() self.tmp = tempfile.TemporaryDirectory()
self.db = f"{self.tmp.name}/f.db" self.db = f"{self.tmp.name}/f.db"
os.environ["PRAGENT_FEEDBACK_DB"] = self.db os.environ["PRAGENT_FEEDBACK_DB"] = self.db
os.environ["PRAGENT_DASHBOARD_TOKEN"] = ""
os.environ["PRAGENT_BOT_TOKEN"] = "" os.environ["PRAGENT_BOT_TOKEN"] = ""
conn = feedback.init(self.db) conn = feedback.init(self.db)
feedback.record_review(conn, repo="o/r", pr=1, head_sha="x") feedback.record_review(conn, repo="o/r", pr=1, head_sha="x")
@@ -49,17 +48,18 @@ class TestRepoEditSelect(unittest.TestCase):
self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler) self.srv = ThreadingHTTPServer(("127.0.0.1", self.port), dash.Handler)
self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True) self.thread = threading.Thread(target=self.srv.serve_forever, daemon=True)
self.thread.start() self.thread.start()
self.auth_hdr = {"X-Forwarded-User": "marcos@example.com"}
def tearDown(self): def tearDown(self):
self.srv.shutdown() self.srv.shutdown()
self.srv.server_close() self.srv.server_close()
self.thread.join(timeout=2) self.thread.join(timeout=2)
for k in ("PRAGENT_FEEDBACK_DB", "PRAGENT_DASHBOARD_TOKEN", "PRAGENT_BOT_TOKEN"): for k in ("PRAGENT_FEEDBACK_DB", "PRAGENT_BOT_TOKEN"):
os.environ.pop(k, None) os.environ.pop(k, None)
self.tmp.cleanup() self.tmp.cleanup()
def test_repo_page_renders_select_with_one_option_per_price(self): def test_repo_page_renders_select_with_one_option_per_price(self):
status, _h, body = _get(self.port, "/r/o/r") status, _h, body = _get(self.port, "/r/o/r", headers=self.auth_hdr)
self.assertEqual(status, 200) self.assertEqual(status, 200)
text = body.decode() text = body.decode()
self.assertIn('<select id="model" name="model">', text) self.assertIn('<select id="model" name="model">', text)