From ca34ff1dc678ae69e3a5570ace38091aa1c06d55 Mon Sep 17 00:00:00 2001 From: Dennis Juhler Aagaard Date: Thu, 24 Sep 2026 17:00:07 +0200 Subject: [PATCH] fix: harden stelloauth oauth worker --- stelloauth/patches/cloakserve-loopback.patch | 14 + stelloauth/patches/stelloauth-security.patch | 341 +++++++++++++++++++ tests/test_security_patches.py | 56 +++ 3 files changed, 411 insertions(+) create mode 100644 stelloauth/patches/cloakserve-loopback.patch create mode 100644 stelloauth/patches/stelloauth-security.patch create mode 100644 tests/test_security_patches.py diff --git a/stelloauth/patches/cloakserve-loopback.patch b/stelloauth/patches/cloakserve-loopback.patch new file mode 100644 index 0000000..78c3c2b --- /dev/null +++ b/stelloauth/patches/cloakserve-loopback.patch @@ -0,0 +1,14 @@ +diff --git a/bin/cloakserve b/bin/cloakserve +index 0e545b2..04354aa 100755 +--- a/bin/cloakserve ++++ b/bin/cloakserve +@@ -903,8 +903,7 @@ def main() -> None: + port, + ) + +- in_container = os.path.exists("/.dockerenv") or os.path.exists("/run/.containerenv") +- host = "0.0.0.0" if in_container else "127.0.0.1" ++ host = "127.0.0.1" + web.run_app(app, host=host, port=port, print=None) + + diff --git a/stelloauth/patches/stelloauth-security.patch b/stelloauth/patches/stelloauth-security.patch new file mode 100644 index 0000000..7f97723 --- /dev/null +++ b/stelloauth/patches/stelloauth-security.patch @@ -0,0 +1,341 @@ +diff --git a/internal/app/oauth.go b/internal/app/oauth.go +index f0b9548..82f356f 100644 +--- a/internal/app/oauth.go ++++ b/internal/app/oauth.go +@@ -162,12 +162,11 @@ func performChromedpOAuth( + reqURL := e.Request.URL + // Capture OAuth redirect + if strings.HasPrefix(reqURL, redirectPrefix) { +- log.Printf("[%s] Redirect URL: %s", requestID, reqURL) ++ log.Printf("[%s] Captured OAuth redirect", requestID) + parsed, err := url.Parse(reqURL) + if err == nil { + if code := parsed.Query().Get("code"); code != "" { + oauthCode = code +- log.Printf("[%s] Captured OAuth code from redirect request", requestID) + } + } + } else if strings.Contains(reqURL, "OPErrorPage.php") { +@@ -175,7 +174,7 @@ func performChromedpOAuth( + // contextId when the login took too long). + if parsed, perr := url.Parse(reqURL); perr == nil { + flowError = friendlyOPError(parsed.Query().Get("code"), parsed.Query().Get("message")) +- log.Printf("[%s] Stellantis error page: %s", requestID, reqURL) ++ log.Printf("[%s] Stellantis error page received", requestID) + } + } else if debug != nil && isRelevantURL(reqURL) { + // Only show relevant OAuth flow URLs in debug output +@@ -452,7 +451,7 @@ func friendlyOPError(code, message string) string { + func codeFromLocation(browserCtx context.Context, redirectPrefix string) string { + var currentURL string + _ = chromedp.Run(browserCtx, chromedp.Location(¤tURL)) +- log.Printf("Current URL: %s", currentURL) ++ log.Printf("Current browser location checked") + if !strings.HasPrefix(currentURL, redirectPrefix) { + return "" + } +diff --git a/internal/app/server.go b/internal/app/server.go +index 48a487e..98accf9 100644 +--- a/internal/app/server.go ++++ b/internal/app/server.go +@@ -121,6 +121,13 @@ func getClientIP(r *http.Request) string { + return r.RemoteAddr + } + ++func remoteClientIP(r *http.Request) string { ++ if host, _, err := net.SplitHostPort(r.RemoteAddr); err == nil { ++ return host ++ } ++ return strings.TrimSpace(r.RemoteAddr) ++} ++ + // refundIfExpired gives back the rate-limit charge when the OAuth attempt + // failed with a transient "session expired" error (bounded by the limiter). + func refundIfExpired(clientIP, requestID string, err error) { +@@ -136,7 +143,7 @@ func handleOAuth(w http.ResponseWriter, r *http.Request) { + } + + // Get client IP early for rate limiting +- clientIP := getClientIP(r) ++ clientIP := remoteClientIP(r) + + // Check rate limit + if !rateLimiter.isAllowed(clientIP) { +@@ -160,7 +167,7 @@ func handleOAuth(w http.ResponseWriter, r *http.Request) { + // Generate request ID + requestID := uuid.New().String() + +- log.Printf("[%s] OAuth request from %s for user %s (%s/%s)", requestID, clientIP, req.Email, req.Brand, req.Country) ++ log.Printf("[%s] OAuth request from %s (%s/%s)", requestID, clientIP, req.Brand, req.Country) + + // Check if client accepts SSE + if r.Header.Get("Accept") == "text/event-stream" { +@@ -171,7 +178,7 @@ func handleOAuth(w http.ResponseWriter, r *http.Request) { + code, err := performOAuth(req, requestID, nil, nil) + if err != nil { + refundIfExpired(clientIP, requestID, err) +- log.Printf("[%s] OAuth failed: %s", requestID, err.Error()) ++ log.Printf("[%s] OAuth failed", requestID) + sendError(w, err.Error(), http.StatusBadRequest) + return + } +@@ -206,7 +213,7 @@ func handleOAuthSSE(w http.ResponseWriter, req OAuthRequest, requestID, clientIP + code, err := performOAuth(req, requestID, progress, debug) + if err != nil { + refundIfExpired(clientIP, requestID, err) +- log.Printf("[%s] OAuth failed: %s", requestID, err.Error()) ++ log.Printf("[%s] OAuth failed", requestID) + _, _ = fmt.Fprintf(w, "data: {\"type\":\"error\",\"message\":\"%s\"}\n\n", err.Error()) + flusher.Flush() + return +diff --git a/internal/app/server_test.go b/internal/app/server_test.go +index b56bb27..df05e37 100644 +--- a/internal/app/server_test.go ++++ b/internal/app/server_test.go +@@ -1,7 +1,9 @@ + package app + + import ( ++ "bytes" + "encoding/json" ++ "log" + "net/http" + "net/http/httptest" + "os" +@@ -188,3 +190,56 @@ func TestParseClientIP(t *testing.T) { + } + } + } ++ ++func TestRemoteClientIPIgnoresForwardedHeaders(t *testing.T) { ++ req := httptest.NewRequest(http.MethodPost, "/oauth", nil) ++ req.RemoteAddr = "192.0.2.10:1234" ++ req.Header.Set("Forwarded", "for=203.0.113.1") ++ req.Header.Set("X-Forwarded-For", "203.0.113.2") ++ req.Header.Set("X-Real-IP", "203.0.113.3") ++ ++ if got := remoteClientIP(req); got != "192.0.2.10" { ++ t.Errorf("remoteClientIP() = %q, want %q", got, "192.0.2.10") ++ } ++} ++ ++func TestRequestLogsDoNotContainCredentialsOrURLs(t *testing.T) { ++ sentinels := []string{ ++ "sentinel-email@example.invalid", ++ "sentinel-password", ++ "https://idpcvs.peugeot.com/am/oauth2/authorize?sentinel=authorize-url", ++ "mypeugeot://oauth2redirect/dk?sentinel=redirect-url", ++ "sentinel-oauth-code", ++ } ++ ++ var logs bytes.Buffer ++ originalWriter := log.Writer() ++ originalFlags := log.Flags() ++ log.SetOutput(&logs) ++ log.SetFlags(0) ++ t.Cleanup(func() { ++ log.SetOutput(originalWriter) ++ log.SetFlags(originalFlags) ++ }) ++ ++ body, err := json.Marshal(OAuthRequest{ ++ Brand: "MyPeugeot", ++ Country: "NOT-A-COUNTRY", ++ Email: strings.Join(sentinels, "|"), ++ Password: "sentinel-password", ++ }) ++ if err != nil { ++ t.Fatalf("marshal request: %v", err) ++ } ++ req := httptest.NewRequest(http.MethodPost, "/oauth", bytes.NewReader(body)) ++ req.RemoteAddr = "192.0.2.10:1234" ++ w := httptest.NewRecorder() ++ ++ handleOAuth(w, req) ++ ++ for _, sentinel := range sentinels { ++ if strings.Contains(logs.String(), sentinel) { ++ t.Errorf("logs contain sensitive value %q: %s", sentinel, logs.String()) ++ } ++ } ++} +diff --git a/internal/app/worker.go b/internal/app/worker.go +index 00b9cf8..0b6e6fd 100644 +--- a/internal/app/worker.go ++++ b/internal/app/worker.go +@@ -7,9 +7,20 @@ import ( + "log" + "net/http" + "net/url" ++ "strings" + "uuid" + ) + ++const maxWorkerRequestBytes = 64 << 10 ++ ++var approvedWorkerHosts = map[string]struct{}{ ++ "idpcvs.citroen.com": {}, ++ "idpcvs.driveds.com": {}, ++ "idpcvs.opel.com": {}, ++ "idpcvs.peugeot.com": {}, ++ "idpcvs.vauxhall.co.uk": {}, ++} ++ + // workerRequest is the worker-v2-compatible OAuth request: the caller supplies a + // fully-built authorize URL instead of brand/country. This mirrors the API of + // github.com/andreadegiovine/homeassistant-stellantis-vehicles-worker-v2 so the +@@ -40,15 +51,21 @@ func handleWorker(w http.ResponseWriter, r *http.Request) { + return + } + +- clientIP := getClientIP(r) ++ clientIP := remoteClientIP(r) + if !rateLimiter.isAllowed(clientIP) { + log.Printf("Rate limit exceeded for %s (worker)", clientIP) + sendWorkerError(w, "Rate limit exceeded. Try again later.", http.StatusTooManyRequests) + return + } + ++ r.Body = http.MaxBytesReader(w, r.Body, maxWorkerRequestBytes) + var req workerRequest + if err := json.UnmarshalRead(r.Body, &req); err != nil { ++ var maxBytesErr *http.MaxBytesError ++ if errors.As(err, &maxBytesErr) { ++ sendWorkerError(w, "Request body too large", http.StatusRequestEntityTooLarge) ++ return ++ } + sendWorkerError(w, "Invalid request body", http.StatusBadRequest) + return + } +@@ -58,6 +75,11 @@ func handleWorker(w http.ResponseWriter, r *http.Request) { + return + } + ++ if err := validateWorkerURL(req.URL); err != nil { ++ sendWorkerError(w, err.Error(), http.StatusBadRequest) ++ return ++ } ++ + scheme, err := redirectScheme(req.URL) + if err != nil { + sendWorkerError(w, err.Error(), http.StatusBadRequest) +@@ -65,12 +87,12 @@ func handleWorker(w http.ResponseWriter, r *http.Request) { + } + + requestID := uuid.New().String() +- log.Printf("[%s] worker OAuth request from %s for user %s", requestID, clientIP, req.Email) ++ log.Printf("[%s] worker OAuth request from %s", requestID, clientIP) + + code, err := performChromedpOAuth(req.URL, req.Email, req.Password, scheme, requestID, nil, nil) + if err != nil { + refundIfExpired(clientIP, requestID, err) +- log.Printf("[%s] worker OAuth failed: %s", requestID, err.Error()) ++ log.Printf("[%s] worker OAuth failed", requestID) + sendWorkerError(w, err.Error(), http.StatusBadRequest) + return + } +@@ -80,6 +102,23 @@ func handleWorker(w http.ResponseWriter, r *http.Request) { + _ = json.MarshalWrite(w, workerCode{Code: code}) + } + ++func validateWorkerURL(rawURL string) error { ++ parsed, err := url.Parse(rawURL) ++ if err != nil { ++ return errors.New("invalid url") ++ } ++ if parsed.Scheme != "https" || parsed.User != nil || parsed.Port() != "" || parsed.Fragment != "" { ++ return errors.New("invalid url") ++ } ++ if _, ok := approvedWorkerHosts[strings.ToLower(parsed.Hostname())]; !ok { ++ return errors.New("invalid url") ++ } ++ if parsed.EscapedPath() != "/am/oauth2/authorize" { ++ return errors.New("invalid url") ++ } ++ return nil ++} ++ + // redirectScheme extracts the custom redirect scheme (e.g. "mymap") from the + // authorize URL's redirect_uri query parameter; the code-capture listener keys on + // "://". +diff --git a/internal/app/worker_test.go b/internal/app/worker_test.go +index e64964d..7870e47 100644 +--- a/internal/app/worker_test.go ++++ b/internal/app/worker_test.go +@@ -8,6 +8,41 @@ import ( + "testing" + ) + ++func TestValidateWorkerURL(t *testing.T) { ++ allowed := []string{ ++ "https://idpcvs.citroen.com/am/oauth2/authorize?redirect_uri=mycitroen%3A%2F%2Foauth2redirect%2Fdk", ++ "https://idpcvs.driveds.com/am/oauth2/authorize?redirect_uri=mymap%3A%2F%2Foauth2redirect%2Fdk", ++ "https://idpcvs.opel.com/am/oauth2/authorize?redirect_uri=myopel%3A%2F%2Foauth2redirect%2Fdk", ++ "https://idpcvs.peugeot.com/am/oauth2/authorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "https://idpcvs.vauxhall.co.uk/am/oauth2/authorize?redirect_uri=myvauxhall%3A%2F%2Foauth2redirect%2Fgb", ++ "https://IDPCVS.OPEL.COM/am/oauth2/authorize?redirect_uri=myopel%3A%2F%2Foauth2redirect%2Fdk", ++ } ++ for _, rawURL := range allowed { ++ t.Run(rawURL, func(t *testing.T) { ++ if err := validateWorkerURL(rawURL); err != nil { ++ t.Fatalf("validateWorkerURL() error = %v", err) ++ } ++ }) ++ } ++ ++ rejected := map[string]string{ ++ "http": "http://idpcvs.peugeot.com/am/oauth2/authorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "userinfo": "https://user@idpcvs.peugeot.com/am/oauth2/authorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "explicit port": "https://idpcvs.peugeot.com:443/am/oauth2/authorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "evil suffix": "https://idpcvs.peugeot.com.evil.example/am/oauth2/authorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "extra path component": "https://idpcvs.peugeot.com/am/oauth2/authorize/extra?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "escaped path": "https://idpcvs.peugeot.com/am/oauth2%2Fauthorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk", ++ "fragment": "https://idpcvs.peugeot.com/am/oauth2/authorize?redirect_uri=mypeugeot%3A%2F%2Foauth2redirect%2Fdk#fragment", ++ } ++ for name, rawURL := range rejected { ++ t.Run(name, func(t *testing.T) { ++ if err := validateWorkerURL(rawURL); err == nil { ++ t.Fatal("validateWorkerURL() error = nil, want rejection") ++ } ++ }) ++ } ++} ++ + func TestHandleWorker_MethodNotAllowed(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/worker", nil) + w := httptest.NewRecorder() +@@ -30,6 +65,25 @@ func TestHandleWorker_InvalidBody(t *testing.T) { + } + } + ++func TestHandleWorker_RequestTooLarge(t *testing.T) { ++ body := `{"url":"` + strings.Repeat("x", 65<<10) + `"}` ++ req := httptest.NewRequest(http.MethodPost, "/worker", strings.NewReader(body)) ++ w := httptest.NewRecorder() ++ ++ handleWorker(w, req) ++ ++ if w.Code != http.StatusRequestEntityTooLarge { ++ t.Fatalf("status = %d, want %d", w.Code, http.StatusRequestEntityTooLarge) ++ } ++ var resp workerError ++ if err := json.NewDecoder(w.Body).Decode(&resp); err != nil { ++ t.Fatalf("decode response: %v", err) ++ } ++ if resp.Message != "Request body too large" { ++ t.Errorf("message = %q, want %q", resp.Message, "Request body too large") ++ } ++} ++ + func TestHandleWorker_MissingParams(t *testing.T) { + body := `{"url":"","email":"","password":""}` + req := httptest.NewRequest(http.MethodPost, "/worker", strings.NewReader(body)) +@@ -83,7 +137,7 @@ func TestRedirectScheme(t *testing.T) { + }, + { + name: "opel custom scheme", +- authURL: "https://example.com/authorize?redirect_uri=myopel%3A%2F%2Foauth2redirect%2Fde", ++ authURL: "https://idpcvs.opel.com/am/oauth2/authorize?redirect_uri=myopel%3A%2F%2Foauth2redirect%2Fde", + want: "myopel", + }, + { diff --git a/tests/test_security_patches.py b/tests/test_security_patches.py new file mode 100644 index 0000000..16fe94d --- /dev/null +++ b/tests/test_security_patches.py @@ -0,0 +1,56 @@ +from __future__ import annotations + +import subprocess +from pathlib import Path + +import pytest + + +ROOT = Path(__file__).parents[1] +STELLOAUTH_COMMIT = "367d4f8c02a3b072c59142c49dffc129edc8548b" +CLOAK_COMMIT = "f04c23da285b3b3d3cf10c8f9d282e7adc1d52ce" + + +def run(*args: str, cwd: Path): + return subprocess.run(args, cwd=cwd, text=True, capture_output=True, check=True) + + +def fetch_exact(tmp_path: Path, name: str, url: str, commit: str) -> Path: + target = tmp_path / name + run("git", "init", str(target), cwd=tmp_path) + run("git", "remote", "add", "origin", url, cwd=target) + run("git", "fetch", "--depth=1", "origin", commit, cwd=target) + run("git", "checkout", "--detach", "FETCH_HEAD", cwd=target) + assert run("git", "rev-parse", "HEAD", cwd=target).stdout.strip() == commit + return target + + +@pytest.mark.upstream +def test_stelloauth_patch_applies_and_tests_pass(tmp_path: Path) -> None: + source = fetch_exact( + tmp_path, + "stelloauth", + "https://github.com/tamcore/stelloauth.git", + STELLOAUTH_COMMIT, + ) + patch = ROOT / "stelloauth/patches/stelloauth-security.patch" + run("git", "apply", "--check", str(patch), cwd=source) + run("git", "apply", str(patch), cwd=source) + run("go", "test", "./...", cwd=source) + + +@pytest.mark.upstream +def test_cloak_patch_binds_only_loopback(tmp_path: Path) -> None: + source = fetch_exact( + tmp_path, + "cloakbrowser", + "https://github.com/CloakHQ/CloakBrowser.git", + CLOAK_COMMIT, + ) + patch = ROOT / "stelloauth/patches/cloakserve-loopback.patch" + run("git", "apply", "--check", str(patch), cwd=source) + run("git", "apply", str(patch), cwd=source) + wrapper = (source / "bin/cloakserve").read_text(encoding="utf-8") + assert 'host = "127.0.0.1"' in wrapper + assert 'host = "0.0.0.0" if in_container' not in wrapper + run("python3", "-m", "py_compile", "bin/cloakserve", cwd=source)