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", }, {