package app import ( "bytes" "net/http" "net/http/httptest" "os" "strings" "testing" "time" "github.com/xuri/excelize/v2" ) func TestWriteAuditWorkbookCreatesSafeExcelFile(t *testing.T) { t.Parallel() location, err := time.LoadLocation("America/Chicago") if err != nil { t.Fatal(err) } createdAt := time.Date(2026, time.July, 28, 20, 10, 32, 0, location) records := []auditRecord{{ EventType: "code_exchange_failed", CustomerLabel: "Acme & Sons", User: "taiadmin", Hostname: "=HYPERLINK(\"https://invalid.example\")", PackageSlug: "sentinelone-linux", SourceIP: "203.0.113.10", Details: "reason=code_not_found code_hint=ABCDE", CreatedAt: createdAt, }} var output bytes.Buffer if err := writeAuditWorkbook( &output, records, auditFilters{TimeRange: "all"}, createdAt, location, ); err != nil { t.Fatal(err) } if !bytes.HasPrefix(output.Bytes(), []byte("PK")) { t.Fatal("workbook is not a ZIP-based Excel file") } workbook, err := excelize.OpenReader(bytes.NewReader(output.Bytes())) if err != nil { t.Fatal(err) } defer func() { if err := workbook.Close(); err != nil { t.Error(err) } }() if sheets := workbook.GetSheetList(); len(sheets) != 1 || sheets[0] != "Audit Log" { t.Fatalf("sheets = %#v, want Audit Log", sheets) } customer, err := workbook.GetCellValue("Audit Log", "C5") if err != nil { t.Fatal(err) } if customer != "Acme & Sons" { t.Fatalf("customer = %q, want Acme & Sons", customer) } hostname, err := workbook.GetCellValue("Audit Log", "E5") if err != nil { t.Fatal(err) } if hostname != `=HYPERLINK("https://invalid.example")` { t.Fatalf("hostname = %q, formula-looking text was not preserved", hostname) } formula, err := workbook.GetCellFormula("Audit Log", "E5") if err != nil { t.Fatal(err) } if formula != "" { t.Fatalf("untrusted audit text was written as formula %q", formula) } if samplePath := os.Getenv("TAPM_AUDIT_SAMPLE_PATH"); samplePath != "" { if err := os.WriteFile(samplePath, output.Bytes(), 0o600); err != nil { t.Fatal(err) } } } func TestAuditExportReturnsAllRowsAndRecordsDownload(t *testing.T) { server, sessionToken := newAuthorizationTestServer(t) server.cfg.DisplayTimeZone = time.UTC for _, eventType := range []string{"older_event", "newer_event"} { if _, err := server.db.Exec( `INSERT INTO audit_events (event_type, details) VALUES (?, '')`, eventType, ); err != nil { t.Fatal(err) } } request := httptest.NewRequest(http.MethodGet, "/portal/audit/export?audit_range=all", nil) request.AddCookie(&http.Cookie{Name: sessionCookieName, Value: sessionToken}) response := httptest.NewRecorder() server.handleAuditExport(response, request) if response.Code != http.StatusOK { t.Fatalf("status = %d, body = %s", response.Code, response.Body.String()) } if contentType := response.Header().Get("Content-Type"); contentType != "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet" { t.Fatalf("content type = %q", contentType) } if disposition := response.Header().Get("Content-Disposition"); !strings.Contains( disposition, `attachment; filename="tapm-audit-`, ) { t.Fatalf("content disposition = %q", disposition) } workbook, err := excelize.OpenReader(bytes.NewReader(response.Body.Bytes())) if err != nil { t.Fatal(err) } defer func() { if err := workbook.Close(); err != nil { t.Error(err) } }() rows, err := workbook.GetRows("Audit Log") if err != nil { t.Fatal(err) } var exportedEvents []string for _, row := range rows { if len(row) > 1 && (row[1] == "older_event" || row[1] == "newer_event") { exportedEvents = append(exportedEvents, row[1]) } } if len(exportedEvents) != 2 { t.Fatal("complete export did not contain all retained audit rows") } var exportEvents int if err := server.db.QueryRow( `SELECT COUNT(*) FROM audit_events WHERE event_type = 'audit_exported'`, ).Scan(&exportEvents); err != nil { t.Fatal(err) } if exportEvents != 1 { t.Fatalf("audit_exported events = %d, want 1", exportEvents) } } func TestAuditExportRequiresTechnicianSession(t *testing.T) { server, _ := newAuthorizationTestServer(t) request := httptest.NewRequest(http.MethodGet, "/portal/audit/export?audit_range=all", nil) response := httptest.NewRecorder() server.Routes().ServeHTTP(response, request) if response.Code != http.StatusSeeOther { t.Fatalf("status = %d, want redirect", response.Code) } }