diff --git a/internal/app/download.go b/internal/app/download.go index 113bd31..4f90f00 100644 --- a/internal/app/download.go +++ b/internal/app/download.go @@ -61,11 +61,17 @@ func (s *Server) handleExchange(w http.ResponseWriter, r *http.Request) { request.RequestedAction = strings.TrimSpace(request.RequestedAction) request.RequestedPackage = strings.TrimSpace(request.RequestedPackage) request.InstallationID = strings.TrimSpace(request.InstallationID) - if request.Code == "" || len(request.HostFingerprint) < 16 || + if len(request.HostFingerprint) < 16 || request.Hostname == "" || len(request.Hostname) > 255 || (request.RequestedAction != "" && !validSlug(request.RequestedAction)) || (request.RequestedPackage != "" && !validSlug(request.RequestedPackage)) { - writeJSON(w, http.StatusBadRequest, map[string]string{"error": "code and host identity are required"}) + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "valid host identity is required"}) + return + } + if !validDeploymentCodeFormat(request.Code) { + _ = s.audit(r.Context(), "code_exchange_failed", "", nil, + request.Hostname, "", sourceIP, codeExchangeFailureDetails(request.Code, errCodeNotFound)) + writeJSON(w, http.StatusForbidden, map[string]string{"error": "authorization is invalid"}) return } diff --git a/internal/app/download_test.go b/internal/app/download_test.go index 0bf0138..4d12aed 100644 --- a/internal/app/download_test.go +++ b/internal/app/download_test.go @@ -61,6 +61,41 @@ func TestCodeExchangeFailureReasons(t *testing.T) { } } +func TestMalformedCodeAttemptIsAudited(t *testing.T) { + server, _ := newAuthorizationTestServer(t) + body, err := json.Marshal(exchangeRequest{ + Code: "not-a-code", + HostFingerprint: "0123456789abcdef", + Hostname: "pve01", + }) + if err != nil { + t.Fatal(err) + } + request := httptest.NewRequest(http.MethodPost, "/api/v1/exchange", bytes.NewReader(body)) + response := httptest.NewRecorder() + server.handleExchange(response, request) + if response.Code != http.StatusForbidden { + t.Fatalf("status = %d, want %d", response.Code, http.StatusForbidden) + } + + var hostname, details string + if err := server.db.QueryRow( + `SELECT hostname, details + FROM audit_events + WHERE event_type = 'code_exchange_failed'`, + ).Scan(&hostname, &details); err != nil { + t.Fatal(err) + } + if hostname != "pve01" || + !strings.Contains(details, "reason=code_not_found") || + !strings.Contains(details, "normalized_length=10") || + !strings.Contains(details, "format=unexpected") || + !strings.Contains(details, "code_hint=unavailable") || + strings.Contains(details, "not-a-code") { + t.Fatalf("unsafe or incomplete malformed-code audit: hostname=%q details=%q", hostname, details) + } +} + func TestFailedExpiredCodeAttemptIsLinkedToDeployment(t *testing.T) { server, _ := newAuthorizationTestServer(t) code := "TAPM-12345-ABCDE"