mirror of
https://github.com/rcourtman/Pulse.git
synced 2026-09-10 02:25:56 +00:00
feat(audit): add comprehensive audit logging system
- Add SQLite-backed audit logger for persistent audit trails - Implement cryptographic signing for tamper detection - Add audit log export functionality - Add webhook notifications for audit events
This commit is contained in:
@@ -0,0 +1,251 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/csv"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ExportFormat defines the export file format.
|
||||
type ExportFormat string
|
||||
|
||||
const (
|
||||
ExportFormatCSV ExportFormat = "csv"
|
||||
ExportFormatJSON ExportFormat = "json"
|
||||
)
|
||||
|
||||
// ExportResult contains export data and metadata.
|
||||
type ExportResult struct {
|
||||
Data []byte
|
||||
ContentType string
|
||||
Filename string
|
||||
EventCount int
|
||||
}
|
||||
|
||||
// ExportEvent extends Event with verification status for exports.
|
||||
type ExportEvent struct {
|
||||
ID string `json:"id"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
EventType string `json:"event_type"`
|
||||
User string `json:"user,omitempty"`
|
||||
IP string `json:"ip,omitempty"`
|
||||
Path string `json:"path,omitempty"`
|
||||
Success bool `json:"success"`
|
||||
Details string `json:"details,omitempty"`
|
||||
Signature string `json:"signature,omitempty"`
|
||||
SignatureValid *bool `json:"signature_valid,omitempty"`
|
||||
}
|
||||
|
||||
// Exporter provides export functionality for audit logs.
|
||||
type Exporter struct {
|
||||
logger *SQLiteLogger
|
||||
}
|
||||
|
||||
// NewExporter creates a new exporter for the given logger.
|
||||
func NewExporter(logger *SQLiteLogger) *Exporter {
|
||||
return &Exporter{logger: logger}
|
||||
}
|
||||
|
||||
// Export generates an export in the specified format.
|
||||
func (e *Exporter) Export(filter QueryFilter, format ExportFormat, includeVerification bool) (*ExportResult, error) {
|
||||
// Remove limit for export (get all matching events)
|
||||
filter.Limit = 0
|
||||
filter.Offset = 0
|
||||
|
||||
events, err := e.logger.Query(filter)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query events for export: %w", err)
|
||||
}
|
||||
|
||||
// Convert to export events
|
||||
exportEvents := make([]ExportEvent, len(events))
|
||||
for i, event := range events {
|
||||
exportEvents[i] = ExportEvent{
|
||||
ID: event.ID,
|
||||
Timestamp: event.Timestamp,
|
||||
EventType: event.EventType,
|
||||
User: event.User,
|
||||
IP: event.IP,
|
||||
Path: event.Path,
|
||||
Success: event.Success,
|
||||
Details: event.Details,
|
||||
Signature: event.Signature,
|
||||
}
|
||||
|
||||
if includeVerification && event.Signature != "" {
|
||||
valid := e.logger.VerifySignature(event)
|
||||
exportEvents[i].SignatureValid = &valid
|
||||
}
|
||||
}
|
||||
|
||||
// Generate timestamp for filename
|
||||
timestamp := time.Now().Format("20060102-150405")
|
||||
|
||||
switch format {
|
||||
case ExportFormatCSV:
|
||||
return e.exportCSV(exportEvents, timestamp, includeVerification)
|
||||
case ExportFormatJSON:
|
||||
return e.exportJSON(exportEvents, timestamp)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported export format: %s", format)
|
||||
}
|
||||
}
|
||||
|
||||
// exportCSV generates a CSV export.
|
||||
func (e *Exporter) exportCSV(events []ExportEvent, timestamp string, includeVerification bool) (*ExportResult, error) {
|
||||
var buf bytes.Buffer
|
||||
writer := csv.NewWriter(&buf)
|
||||
|
||||
// Write header
|
||||
header := []string{"ID", "Timestamp", "Event Type", "User", "IP", "Path", "Success", "Details", "Signature"}
|
||||
if includeVerification {
|
||||
header = append(header, "Signature Valid")
|
||||
}
|
||||
if err := writer.Write(header); err != nil {
|
||||
return nil, fmt.Errorf("failed to write CSV header: %w", err)
|
||||
}
|
||||
|
||||
// Write rows
|
||||
for _, event := range events {
|
||||
success := "false"
|
||||
if event.Success {
|
||||
success = "true"
|
||||
}
|
||||
|
||||
row := []string{
|
||||
event.ID,
|
||||
event.Timestamp.Format(time.RFC3339),
|
||||
event.EventType,
|
||||
event.User,
|
||||
event.IP,
|
||||
event.Path,
|
||||
success,
|
||||
event.Details,
|
||||
event.Signature,
|
||||
}
|
||||
|
||||
if includeVerification {
|
||||
sigValid := ""
|
||||
if event.SignatureValid != nil {
|
||||
if *event.SignatureValid {
|
||||
sigValid = "true"
|
||||
} else {
|
||||
sigValid = "false"
|
||||
}
|
||||
}
|
||||
row = append(row, sigValid)
|
||||
}
|
||||
|
||||
if err := writer.Write(row); err != nil {
|
||||
return nil, fmt.Errorf("failed to write CSV row: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
writer.Flush()
|
||||
if err := writer.Error(); err != nil {
|
||||
return nil, fmt.Errorf("CSV writer error: %w", err)
|
||||
}
|
||||
|
||||
return &ExportResult{
|
||||
Data: buf.Bytes(),
|
||||
ContentType: "text/csv; charset=utf-8",
|
||||
Filename: fmt.Sprintf("audit-log-%s.csv", timestamp),
|
||||
EventCount: len(events),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// exportJSON generates a JSON export.
|
||||
func (e *Exporter) exportJSON(events []ExportEvent, timestamp string) (*ExportResult, error) {
|
||||
// Wrap in an object for better structure
|
||||
export := struct {
|
||||
ExportedAt time.Time `json:"exported_at"`
|
||||
EventCount int `json:"event_count"`
|
||||
Events []ExportEvent `json:"events"`
|
||||
}{
|
||||
ExportedAt: time.Now(),
|
||||
EventCount: len(events),
|
||||
Events: events,
|
||||
}
|
||||
|
||||
data, err := json.MarshalIndent(export, "", " ")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal JSON export: %w", err)
|
||||
}
|
||||
|
||||
return &ExportResult{
|
||||
Data: data,
|
||||
ContentType: "application/json; charset=utf-8",
|
||||
Filename: fmt.Sprintf("audit-log-%s.json", timestamp),
|
||||
EventCount: len(events),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// ExportSummary generates a summary of audit activity.
|
||||
type ExportSummary struct {
|
||||
TotalEvents int `json:"total_events"`
|
||||
SuccessCount int `json:"success_count"`
|
||||
FailureCount int `json:"failure_count"`
|
||||
EventsByType map[string]int `json:"events_by_type"`
|
||||
EventsByUser map[string]int `json:"events_by_user"`
|
||||
StartTime *time.Time `json:"start_time,omitempty"`
|
||||
EndTime *time.Time `json:"end_time,omitempty"`
|
||||
InvalidSigCount int `json:"invalid_signature_count,omitempty"`
|
||||
}
|
||||
|
||||
// GenerateSummary creates a summary of audit events matching the filter.
|
||||
func (e *Exporter) GenerateSummary(filter QueryFilter, verifySignatures bool) (*ExportSummary, error) {
|
||||
// Remove limit for summary
|
||||
filter.Limit = 0
|
||||
filter.Offset = 0
|
||||
|
||||
events, err := e.logger.Query(filter)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query events for summary: %w", err)
|
||||
}
|
||||
|
||||
summary := &ExportSummary{
|
||||
TotalEvents: len(events),
|
||||
EventsByType: make(map[string]int),
|
||||
EventsByUser: make(map[string]int),
|
||||
}
|
||||
|
||||
var minTime, maxTime *time.Time
|
||||
|
||||
for _, event := range events {
|
||||
if event.Success {
|
||||
summary.SuccessCount++
|
||||
} else {
|
||||
summary.FailureCount++
|
||||
}
|
||||
|
||||
summary.EventsByType[event.EventType]++
|
||||
|
||||
if event.User != "" {
|
||||
summary.EventsByUser[event.User]++
|
||||
}
|
||||
|
||||
// Track time range
|
||||
if minTime == nil || event.Timestamp.Before(*minTime) {
|
||||
t := event.Timestamp
|
||||
minTime = &t
|
||||
}
|
||||
if maxTime == nil || event.Timestamp.After(*maxTime) {
|
||||
t := event.Timestamp
|
||||
maxTime = &t
|
||||
}
|
||||
|
||||
// Verify signatures if requested
|
||||
if verifySignatures && event.Signature != "" {
|
||||
if !e.logger.VerifySignature(event) {
|
||||
summary.InvalidSigCount++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
summary.StartTime = minTime
|
||||
summary.EndTime = maxTime
|
||||
|
||||
return summary, nil
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
|
||||
"github.com/rs/zerolog/log"
|
||||
)
|
||||
|
||||
// Signer handles HMAC-SHA256 signing and verification for audit events.
|
||||
// The signing key is stored encrypted at rest using the provided crypto manager.
|
||||
type Signer struct {
|
||||
key []byte // 32-byte HMAC signing key
|
||||
}
|
||||
|
||||
// CryptoEncryptor interface for encrypting/decrypting the signing key.
|
||||
// This matches the methods from internal/crypto.CryptoManager.
|
||||
type CryptoEncryptor interface {
|
||||
Encrypt(plaintext []byte) ([]byte, error)
|
||||
Decrypt(ciphertext []byte) ([]byte, error)
|
||||
}
|
||||
|
||||
// NewSigner creates a new signer, loading or generating the HMAC key.
|
||||
// The key is stored encrypted in the data directory.
|
||||
// If cryptoMgr is nil, signing will be disabled (returns empty signatures).
|
||||
func NewSigner(dataDir string, cryptoMgr CryptoEncryptor) (*Signer, error) {
|
||||
if cryptoMgr == nil {
|
||||
log.Warn().Msg("Crypto manager not provided, audit signing disabled")
|
||||
return &Signer{key: nil}, nil
|
||||
}
|
||||
|
||||
keyPath := filepath.Join(dataDir, ".audit-signing.key")
|
||||
|
||||
// Try to load existing key
|
||||
if encryptedKey, err := os.ReadFile(keyPath); err == nil {
|
||||
key, err := cryptoMgr.Decrypt(encryptedKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to decrypt audit signing key: %w", err)
|
||||
}
|
||||
if len(key) != 32 {
|
||||
return nil, fmt.Errorf("invalid audit signing key length: got %d, want 32", len(key))
|
||||
}
|
||||
log.Debug().Msg("Loaded existing audit signing key")
|
||||
return &Signer{key: key}, nil
|
||||
}
|
||||
|
||||
// Generate new key
|
||||
key := make([]byte, 32)
|
||||
if _, err := io.ReadFull(rand.Reader, key); err != nil {
|
||||
return nil, fmt.Errorf("failed to generate audit signing key: %w", err)
|
||||
}
|
||||
|
||||
// Encrypt and save
|
||||
encryptedKey, err := cryptoMgr.Encrypt(key)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to encrypt audit signing key: %w", err)
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(keyPath), 0700); err != nil {
|
||||
return nil, fmt.Errorf("failed to create directory for audit signing key: %w", err)
|
||||
}
|
||||
|
||||
if err := os.WriteFile(keyPath, encryptedKey, 0600); err != nil {
|
||||
return nil, fmt.Errorf("failed to save audit signing key: %w", err)
|
||||
}
|
||||
|
||||
log.Info().Msg("Generated new audit signing key")
|
||||
return &Signer{key: key}, nil
|
||||
}
|
||||
|
||||
// Sign computes an HMAC-SHA256 signature over the event's canonical form.
|
||||
// Returns hex-encoded signature, or empty string if signing is disabled.
|
||||
func (s *Signer) Sign(event Event) string {
|
||||
if s.key == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
canonical := s.canonicalForm(event)
|
||||
mac := hmac.New(sha256.New, s.key)
|
||||
mac.Write([]byte(canonical))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
// Verify checks if the event's signature matches its content.
|
||||
// Returns true if the signature is valid, false if invalid or signing is disabled.
|
||||
func (s *Signer) Verify(event Event) bool {
|
||||
if s.key == nil || event.Signature == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
expected := s.Sign(event)
|
||||
return hmac.Equal([]byte(expected), []byte(event.Signature))
|
||||
}
|
||||
|
||||
// canonicalForm creates a deterministic string representation of an event for signing.
|
||||
// Format: ID|Timestamp(Unix)|EventType|User|IP|Path|Success(0/1)|Details
|
||||
func (s *Signer) canonicalForm(event Event) string {
|
||||
success := "0"
|
||||
if event.Success {
|
||||
success = "1"
|
||||
}
|
||||
|
||||
return event.ID + "|" +
|
||||
strconv.FormatInt(event.Timestamp.Unix(), 10) + "|" +
|
||||
event.EventType + "|" +
|
||||
event.User + "|" +
|
||||
event.IP + "|" +
|
||||
event.Path + "|" +
|
||||
success + "|" +
|
||||
event.Details
|
||||
}
|
||||
|
||||
// SigningEnabled returns true if the signer has a valid key.
|
||||
func (s *Signer) SigningEnabled() bool {
|
||||
return s.key != nil
|
||||
}
|
||||
|
||||
// ExportKey exports the signing key as base64 for backup purposes.
|
||||
// Returns empty string if signing is disabled.
|
||||
func (s *Signer) ExportKey() string {
|
||||
if s.key == nil {
|
||||
return ""
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(s.key)
|
||||
}
|
||||
@@ -0,0 +1,266 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// mockCryptoManager implements CryptoEncryptor for testing.
|
||||
type mockCryptoManager struct {
|
||||
key []byte
|
||||
}
|
||||
|
||||
func newMockCryptoManager() *mockCryptoManager {
|
||||
return &mockCryptoManager{
|
||||
key: []byte("0123456789abcdef0123456789abcdef"), // 32 bytes
|
||||
}
|
||||
}
|
||||
|
||||
func (m *mockCryptoManager) Encrypt(plaintext []byte) ([]byte, error) {
|
||||
// Simple XOR for testing (not secure, just for tests)
|
||||
result := make([]byte, len(plaintext))
|
||||
for i := range plaintext {
|
||||
result[i] = plaintext[i] ^ m.key[i%len(m.key)]
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (m *mockCryptoManager) Decrypt(ciphertext []byte) ([]byte, error) {
|
||||
// XOR is symmetric
|
||||
return m.Encrypt(ciphertext)
|
||||
}
|
||||
|
||||
func TestNewSigner(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
crypto := newMockCryptoManager()
|
||||
|
||||
// Create new signer
|
||||
signer, err := NewSigner(tempDir, crypto)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
if !signer.SigningEnabled() {
|
||||
t.Error("Expected signing to be enabled")
|
||||
}
|
||||
|
||||
// Verify key file was created
|
||||
keyPath := filepath.Join(tempDir, ".audit-signing.key")
|
||||
if _, err := os.Stat(keyPath); os.IsNotExist(err) {
|
||||
t.Error("Key file was not created")
|
||||
}
|
||||
|
||||
// Create another signer - should load existing key
|
||||
signer2, err := NewSigner(tempDir, crypto)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner (reload) failed: %v", err)
|
||||
}
|
||||
|
||||
// Both signers should produce the same signatures
|
||||
event := Event{
|
||||
ID: "test-123",
|
||||
Timestamp: time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC),
|
||||
EventType: "login",
|
||||
User: "admin",
|
||||
IP: "192.168.1.1",
|
||||
Path: "/api/auth",
|
||||
Success: true,
|
||||
Details: "test details",
|
||||
}
|
||||
|
||||
sig1 := signer.Sign(event)
|
||||
sig2 := signer2.Sign(event)
|
||||
|
||||
if sig1 != sig2 {
|
||||
t.Errorf("Signatures should match: got %s and %s", sig1, sig2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewSignerWithoutCrypto(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
// Create signer without crypto manager
|
||||
signer, err := NewSigner(tempDir, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
if signer.SigningEnabled() {
|
||||
t.Error("Expected signing to be disabled without crypto manager")
|
||||
}
|
||||
|
||||
event := Event{
|
||||
ID: "test-123",
|
||||
Timestamp: time.Now(),
|
||||
EventType: "test",
|
||||
}
|
||||
|
||||
sig := signer.Sign(event)
|
||||
if sig != "" {
|
||||
t.Error("Expected empty signature when signing is disabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignerSign(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
crypto := newMockCryptoManager()
|
||||
|
||||
signer, err := NewSigner(tempDir, crypto)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
event := Event{
|
||||
ID: "evt-001",
|
||||
Timestamp: time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC),
|
||||
EventType: "login",
|
||||
User: "testuser",
|
||||
IP: "10.0.0.1",
|
||||
Path: "/api/login",
|
||||
Success: true,
|
||||
Details: "successful login",
|
||||
}
|
||||
|
||||
sig := signer.Sign(event)
|
||||
|
||||
// Signature should be hex-encoded (64 characters for SHA256)
|
||||
if len(sig) != 64 {
|
||||
t.Errorf("Expected signature length 64, got %d", len(sig))
|
||||
}
|
||||
|
||||
// Same event should produce same signature
|
||||
sig2 := signer.Sign(event)
|
||||
if sig != sig2 {
|
||||
t.Error("Same event should produce same signature")
|
||||
}
|
||||
|
||||
// Different event should produce different signature
|
||||
event2 := event
|
||||
event2.User = "different"
|
||||
sig3 := signer.Sign(event2)
|
||||
if sig == sig3 {
|
||||
t.Error("Different events should produce different signatures")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignerVerify(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
crypto := newMockCryptoManager()
|
||||
|
||||
signer, err := NewSigner(tempDir, crypto)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
event := Event{
|
||||
ID: "evt-002",
|
||||
Timestamp: time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC),
|
||||
EventType: "config_change",
|
||||
User: "admin",
|
||||
IP: "192.168.1.100",
|
||||
Path: "/api/settings",
|
||||
Success: true,
|
||||
Details: "changed setting X",
|
||||
}
|
||||
|
||||
// Sign the event
|
||||
event.Signature = signer.Sign(event)
|
||||
|
||||
// Verify should succeed
|
||||
if !signer.Verify(event) {
|
||||
t.Error("Verify should return true for valid signature")
|
||||
}
|
||||
|
||||
// Tamper with event
|
||||
tamperedEvent := event
|
||||
tamperedEvent.User = "hacker"
|
||||
if signer.Verify(tamperedEvent) {
|
||||
t.Error("Verify should return false for tampered event")
|
||||
}
|
||||
|
||||
// Wrong signature
|
||||
wrongSigEvent := event
|
||||
wrongSigEvent.Signature = "0000000000000000000000000000000000000000000000000000000000000000"
|
||||
if signer.Verify(wrongSigEvent) {
|
||||
t.Error("Verify should return false for wrong signature")
|
||||
}
|
||||
|
||||
// Empty signature
|
||||
noSigEvent := event
|
||||
noSigEvent.Signature = ""
|
||||
if signer.Verify(noSigEvent) {
|
||||
t.Error("Verify should return false for empty signature")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignerCanonicalForm(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
crypto := newMockCryptoManager()
|
||||
|
||||
signer, err := NewSigner(tempDir, crypto)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
// Test that canonical form is deterministic
|
||||
event := Event{
|
||||
ID: "id123",
|
||||
Timestamp: time.Unix(1705315800, 0), // Fixed Unix timestamp
|
||||
EventType: "test",
|
||||
User: "user",
|
||||
IP: "1.2.3.4",
|
||||
Path: "/path",
|
||||
Success: true,
|
||||
Details: "details",
|
||||
}
|
||||
|
||||
sig1 := signer.Sign(event)
|
||||
sig2 := signer.Sign(event)
|
||||
|
||||
if sig1 != sig2 {
|
||||
t.Error("Canonical form should be deterministic")
|
||||
}
|
||||
|
||||
// Success=false should produce different signature
|
||||
event.Success = false
|
||||
sig3 := signer.Sign(event)
|
||||
if sig1 == sig3 {
|
||||
t.Error("Different success value should produce different signature")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignerExportKey(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
crypto := newMockCryptoManager()
|
||||
|
||||
signer, err := NewSigner(tempDir, crypto)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
key := signer.ExportKey()
|
||||
if key == "" {
|
||||
t.Error("ExportKey should return non-empty string")
|
||||
}
|
||||
|
||||
// Key should be base64 encoded (44 characters for 32 bytes)
|
||||
if len(key) != 44 {
|
||||
t.Errorf("Expected base64 key length 44, got %d", len(key))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignerExportKeyDisabled(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
signer, err := NewSigner(tempDir, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("NewSigner failed: %v", err)
|
||||
}
|
||||
|
||||
key := signer.ExportKey()
|
||||
if key != "" {
|
||||
t.Error("ExportKey should return empty string when signing is disabled")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,490 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
"github.com/rs/zerolog/log"
|
||||
)
|
||||
|
||||
// SQLiteLoggerConfig configures the SQLite audit logger.
|
||||
type SQLiteLoggerConfig struct {
|
||||
DataDir string // Directory for audit.db
|
||||
CryptoMgr CryptoEncryptor // For encrypting the signing key (optional)
|
||||
RetentionDays int // Days to keep events (default: 90, 0 = forever)
|
||||
}
|
||||
|
||||
// SQLiteLogger implements Logger with persistent SQLite storage and HMAC signing.
|
||||
type SQLiteLogger struct {
|
||||
mu sync.RWMutex
|
||||
db *sql.DB
|
||||
dbPath string
|
||||
signer *Signer
|
||||
webhookDelivery *WebhookDelivery
|
||||
retentionDays int
|
||||
stopChan chan struct{}
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewSQLiteLogger creates a new SQLite-backed audit logger.
|
||||
func NewSQLiteLogger(cfg SQLiteLoggerConfig) (*SQLiteLogger, error) {
|
||||
if cfg.DataDir == "" {
|
||||
return nil, fmt.Errorf("data directory is required")
|
||||
}
|
||||
|
||||
// Ensure directory exists
|
||||
auditDir := filepath.Join(cfg.DataDir, "audit")
|
||||
if err := os.MkdirAll(auditDir, 0700); err != nil {
|
||||
return nil, fmt.Errorf("failed to create audit directory: %w", err)
|
||||
}
|
||||
|
||||
dbPath := filepath.Join(auditDir, "audit.db")
|
||||
|
||||
db, err := sql.Open("sqlite3", dbPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open audit database: %w", err)
|
||||
}
|
||||
|
||||
// Configure SQLite for better concurrency and durability
|
||||
pragmas := []string{
|
||||
"PRAGMA journal_mode=WAL",
|
||||
"PRAGMA synchronous=NORMAL",
|
||||
"PRAGMA busy_timeout=5000",
|
||||
"PRAGMA foreign_keys=ON",
|
||||
"PRAGMA cache_size=-64000", // 64MB cache
|
||||
}
|
||||
|
||||
for _, pragma := range pragmas {
|
||||
if _, err := db.Exec(pragma); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("failed to set pragma %s: %w", pragma, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Initialize signer
|
||||
signer, err := NewSigner(auditDir, cfg.CryptoMgr)
|
||||
if err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("failed to initialize audit signer: %w", err)
|
||||
}
|
||||
|
||||
retentionDays := cfg.RetentionDays
|
||||
if retentionDays == 0 {
|
||||
retentionDays = 90 // Default
|
||||
}
|
||||
|
||||
l := &SQLiteLogger{
|
||||
db: db,
|
||||
dbPath: dbPath,
|
||||
signer: signer,
|
||||
retentionDays: retentionDays,
|
||||
stopChan: make(chan struct{}),
|
||||
}
|
||||
|
||||
// Initialize schema
|
||||
if err := l.initSchema(); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("failed to initialize schema: %w", err)
|
||||
}
|
||||
|
||||
// Load webhook URLs from config table
|
||||
urls := l.loadWebhookURLs()
|
||||
if len(urls) > 0 {
|
||||
l.webhookDelivery = NewWebhookDelivery(urls)
|
||||
l.webhookDelivery.Start()
|
||||
}
|
||||
|
||||
// Start retention worker if retention is enabled
|
||||
if retentionDays > 0 {
|
||||
l.wg.Add(1)
|
||||
go l.retentionWorker()
|
||||
}
|
||||
|
||||
log.Info().
|
||||
Str("dbPath", dbPath).
|
||||
Int("retentionDays", retentionDays).
|
||||
Bool("signingEnabled", signer.SigningEnabled()).
|
||||
Msg("SQLite audit logger initialized")
|
||||
|
||||
return l, nil
|
||||
}
|
||||
|
||||
// initSchema creates the database tables and runs migrations.
|
||||
func (l *SQLiteLogger) initSchema() error {
|
||||
schema := `
|
||||
CREATE TABLE IF NOT EXISTS audit_events (
|
||||
id TEXT PRIMARY KEY,
|
||||
timestamp INTEGER NOT NULL,
|
||||
event_type TEXT NOT NULL,
|
||||
user TEXT,
|
||||
ip TEXT,
|
||||
path TEXT,
|
||||
success INTEGER NOT NULL,
|
||||
details TEXT,
|
||||
signature TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_audit_timestamp ON audit_events(timestamp);
|
||||
CREATE INDEX IF NOT EXISTS idx_audit_event_type ON audit_events(event_type);
|
||||
CREATE INDEX IF NOT EXISTS idx_audit_user ON audit_events(user) WHERE user != '';
|
||||
CREATE INDEX IF NOT EXISTS idx_audit_success ON audit_events(success);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS audit_config (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL,
|
||||
updated_at INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS schema_version (
|
||||
version INTEGER PRIMARY KEY,
|
||||
applied_at INTEGER NOT NULL
|
||||
);
|
||||
`
|
||||
|
||||
_, err := l.db.Exec(schema)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create schema: %w", err)
|
||||
}
|
||||
|
||||
// Record schema version
|
||||
_, err = l.db.Exec(`INSERT OR IGNORE INTO schema_version (version, applied_at) VALUES (1, ?)`,
|
||||
time.Now().Unix())
|
||||
return err
|
||||
}
|
||||
|
||||
// Log records an audit event with HMAC signature.
|
||||
func (l *SQLiteLogger) Log(event Event) error {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
// Sign the event
|
||||
event.Signature = l.signer.Sign(event)
|
||||
|
||||
// Insert into database
|
||||
success := 0
|
||||
if event.Success {
|
||||
success = 1
|
||||
}
|
||||
|
||||
_, err := l.db.Exec(`
|
||||
INSERT INTO audit_events (id, timestamp, event_type, user, ip, path, success, details, signature)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
event.ID,
|
||||
event.Timestamp.Unix(),
|
||||
event.EventType,
|
||||
event.User,
|
||||
event.IP,
|
||||
event.Path,
|
||||
success,
|
||||
event.Details,
|
||||
event.Signature,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to insert audit event: %w", err)
|
||||
}
|
||||
|
||||
// Also log to zerolog for real-time visibility
|
||||
logEvent := log.With().
|
||||
Str("audit_id", event.ID).
|
||||
Str("event", event.EventType).
|
||||
Str("user", event.User).
|
||||
Str("ip", event.IP).
|
||||
Str("path", event.Path).
|
||||
Time("timestamp", event.Timestamp).
|
||||
Str("details", event.Details).
|
||||
Logger()
|
||||
|
||||
if event.Success {
|
||||
logEvent.Info().Msg("Audit event")
|
||||
} else {
|
||||
logEvent.Warn().Msg("Audit event - FAILED")
|
||||
}
|
||||
|
||||
// Send to webhooks if configured
|
||||
if l.webhookDelivery != nil {
|
||||
l.webhookDelivery.Enqueue(event)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Query retrieves audit events matching the filter.
|
||||
func (l *SQLiteLogger) Query(filter QueryFilter) ([]Event, error) {
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
|
||||
query := "SELECT id, timestamp, event_type, user, ip, path, success, details, signature FROM audit_events WHERE 1=1"
|
||||
args := []interface{}{}
|
||||
|
||||
if filter.ID != "" {
|
||||
query += " AND id = ?"
|
||||
args = append(args, filter.ID)
|
||||
}
|
||||
if filter.StartTime != nil {
|
||||
query += " AND timestamp >= ?"
|
||||
args = append(args, filter.StartTime.Unix())
|
||||
}
|
||||
if filter.EndTime != nil {
|
||||
query += " AND timestamp <= ?"
|
||||
args = append(args, filter.EndTime.Unix())
|
||||
}
|
||||
if filter.EventType != "" {
|
||||
query += " AND event_type = ?"
|
||||
args = append(args, filter.EventType)
|
||||
}
|
||||
if filter.User != "" {
|
||||
query += " AND user = ?"
|
||||
args = append(args, filter.User)
|
||||
}
|
||||
if filter.Success != nil {
|
||||
success := 0
|
||||
if *filter.Success {
|
||||
success = 1
|
||||
}
|
||||
query += " AND success = ?"
|
||||
args = append(args, success)
|
||||
}
|
||||
|
||||
query += " ORDER BY timestamp DESC"
|
||||
|
||||
if filter.Limit > 0 {
|
||||
query += " LIMIT ?"
|
||||
args = append(args, filter.Limit)
|
||||
}
|
||||
if filter.Offset > 0 {
|
||||
query += " OFFSET ?"
|
||||
args = append(args, filter.Offset)
|
||||
}
|
||||
|
||||
rows, err := l.db.Query(query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query audit events: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var events []Event
|
||||
for rows.Next() {
|
||||
var e Event
|
||||
var timestamp int64
|
||||
var success int
|
||||
var user, ip, path, details, signature sql.NullString
|
||||
|
||||
err := rows.Scan(&e.ID, ×tamp, &e.EventType, &user, &ip, &path, &success, &details, &signature)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to scan audit event: %w", err)
|
||||
}
|
||||
|
||||
e.Timestamp = time.Unix(timestamp, 0)
|
||||
e.Success = success == 1
|
||||
e.User = user.String
|
||||
e.IP = ip.String
|
||||
e.Path = path.String
|
||||
e.Details = details.String
|
||||
e.Signature = signature.String
|
||||
|
||||
events = append(events, e)
|
||||
}
|
||||
|
||||
return events, rows.Err()
|
||||
}
|
||||
|
||||
// Count returns the number of events matching the filter.
|
||||
func (l *SQLiteLogger) Count(filter QueryFilter) (int, error) {
|
||||
l.mu.RLock()
|
||||
defer l.mu.RUnlock()
|
||||
|
||||
query := "SELECT COUNT(*) FROM audit_events WHERE 1=1"
|
||||
args := []interface{}{}
|
||||
|
||||
if filter.ID != "" {
|
||||
query += " AND id = ?"
|
||||
args = append(args, filter.ID)
|
||||
}
|
||||
if filter.StartTime != nil {
|
||||
query += " AND timestamp >= ?"
|
||||
args = append(args, filter.StartTime.Unix())
|
||||
}
|
||||
if filter.EndTime != nil {
|
||||
query += " AND timestamp <= ?"
|
||||
args = append(args, filter.EndTime.Unix())
|
||||
}
|
||||
if filter.EventType != "" {
|
||||
query += " AND event_type = ?"
|
||||
args = append(args, filter.EventType)
|
||||
}
|
||||
if filter.User != "" {
|
||||
query += " AND user = ?"
|
||||
args = append(args, filter.User)
|
||||
}
|
||||
if filter.Success != nil {
|
||||
success := 0
|
||||
if *filter.Success {
|
||||
success = 1
|
||||
}
|
||||
query += " AND success = ?"
|
||||
args = append(args, success)
|
||||
}
|
||||
|
||||
var count int
|
||||
err := l.db.QueryRow(query, args...).Scan(&count)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to count audit events: %w", err)
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// GetWebhookURLs returns the configured webhook URLs.
|
||||
func (l *SQLiteLogger) GetWebhookURLs() []string {
|
||||
if l.webhookDelivery != nil {
|
||||
return l.webhookDelivery.GetURLs()
|
||||
}
|
||||
return l.loadWebhookURLs()
|
||||
}
|
||||
|
||||
// UpdateWebhookURLs updates the webhook configuration.
|
||||
func (l *SQLiteLogger) UpdateWebhookURLs(urls []string) error {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
// Save to config table
|
||||
value := strings.Join(urls, ",")
|
||||
_, err := l.db.Exec(`
|
||||
INSERT INTO audit_config (key, value, updated_at) VALUES ('webhook_urls', ?, ?)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
|
||||
value, time.Now().Unix())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to save webhook URLs: %w", err)
|
||||
}
|
||||
|
||||
// Update delivery worker
|
||||
if len(urls) > 0 {
|
||||
if l.webhookDelivery == nil {
|
||||
l.webhookDelivery = NewWebhookDelivery(urls)
|
||||
l.webhookDelivery.Start()
|
||||
} else {
|
||||
l.webhookDelivery.UpdateURLs(urls)
|
||||
}
|
||||
} else if l.webhookDelivery != nil {
|
||||
l.webhookDelivery.Stop()
|
||||
l.webhookDelivery = nil
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// VerifySignature checks if an event's signature is valid.
|
||||
func (l *SQLiteLogger) VerifySignature(event Event) bool {
|
||||
return l.signer.Verify(event)
|
||||
}
|
||||
|
||||
// Close gracefully shuts down the logger.
|
||||
func (l *SQLiteLogger) Close() error {
|
||||
close(l.stopChan)
|
||||
|
||||
if l.webhookDelivery != nil {
|
||||
l.webhookDelivery.Stop()
|
||||
}
|
||||
|
||||
l.wg.Wait()
|
||||
|
||||
if err := l.db.Close(); err != nil {
|
||||
return fmt.Errorf("failed to close audit database: %w", err)
|
||||
}
|
||||
|
||||
log.Info().Msg("SQLite audit logger closed")
|
||||
return nil
|
||||
}
|
||||
|
||||
// loadWebhookURLs loads webhook URLs from the config table.
|
||||
func (l *SQLiteLogger) loadWebhookURLs() []string {
|
||||
var value string
|
||||
err := l.db.QueryRow(`SELECT value FROM audit_config WHERE key = 'webhook_urls'`).Scan(&value)
|
||||
if err != nil || value == "" {
|
||||
return nil
|
||||
}
|
||||
return strings.Split(value, ",")
|
||||
}
|
||||
|
||||
// retentionWorker runs periodically to clean up old events.
|
||||
func (l *SQLiteLogger) retentionWorker() {
|
||||
defer l.wg.Done()
|
||||
|
||||
// Run at 3 AM daily
|
||||
ticker := time.NewTicker(24 * time.Hour)
|
||||
defer ticker.Stop()
|
||||
|
||||
// Also run once at startup after a short delay
|
||||
time.AfterFunc(5*time.Minute, func() {
|
||||
l.cleanupOldEvents()
|
||||
})
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-l.stopChan:
|
||||
return
|
||||
case <-ticker.C:
|
||||
l.cleanupOldEvents()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// cleanupOldEvents deletes events older than the retention period.
|
||||
func (l *SQLiteLogger) cleanupOldEvents() {
|
||||
if l.retentionDays <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
cutoff := time.Now().AddDate(0, 0, -l.retentionDays).Unix()
|
||||
|
||||
result, err := l.db.Exec(`DELETE FROM audit_events WHERE timestamp < ?`, cutoff)
|
||||
if err != nil {
|
||||
log.Error().Err(err).Msg("Failed to cleanup old audit events")
|
||||
return
|
||||
}
|
||||
|
||||
deleted, _ := result.RowsAffected()
|
||||
if deleted > 0 {
|
||||
log.Info().
|
||||
Int64("deleted", deleted).
|
||||
Int("retentionDays", l.retentionDays).
|
||||
Msg("Cleaned up old audit events")
|
||||
|
||||
// Log the cleanup as an audit event (without recursion - direct insert)
|
||||
_, _ = l.db.Exec(`
|
||||
INSERT INTO audit_events (id, timestamp, event_type, user, ip, path, success, details, signature)
|
||||
VALUES (?, ?, 'audit_cleanup', 'system', '', '', 1, ?, '')`,
|
||||
fmt.Sprintf("cleanup-%d", time.Now().Unix()),
|
||||
time.Now().Unix(),
|
||||
fmt.Sprintf("Deleted %d events older than %d days", deleted, l.retentionDays),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// GetRetentionDays returns the current retention period.
|
||||
func (l *SQLiteLogger) GetRetentionDays() int {
|
||||
return l.retentionDays
|
||||
}
|
||||
|
||||
// SetRetentionDays updates the retention period.
|
||||
func (l *SQLiteLogger) SetRetentionDays(days int) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
l.retentionDays = days
|
||||
|
||||
// Save to config
|
||||
_, _ = l.db.Exec(`
|
||||
INSERT INTO audit_config (key, value, updated_at) VALUES ('retention_days', ?, ?)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
|
||||
fmt.Sprintf("%d", days), time.Now().Unix())
|
||||
}
|
||||
@@ -0,0 +1,504 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
func TestNewSQLiteLogger(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
if logger.GetRetentionDays() != 30 {
|
||||
t.Errorf("Expected retention days 30, got %d", logger.GetRetentionDays())
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewSQLiteLoggerDefaultRetention(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
// RetentionDays not set
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
if logger.GetRetentionDays() != 90 {
|
||||
t.Errorf("Expected default retention days 90, got %d", logger.GetRetentionDays())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerLog(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
event := Event{
|
||||
ID: uuid.NewString(),
|
||||
Timestamp: time.Now(),
|
||||
EventType: "test_event",
|
||||
User: "testuser",
|
||||
IP: "192.168.1.1",
|
||||
Path: "/api/test",
|
||||
Success: true,
|
||||
Details: "test details",
|
||||
}
|
||||
|
||||
err = logger.Log(event)
|
||||
if err != nil {
|
||||
t.Fatalf("Log failed: %v", err)
|
||||
}
|
||||
|
||||
// Query the event back
|
||||
events, err := logger.Query(QueryFilter{ID: event.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
|
||||
if len(events) != 1 {
|
||||
t.Fatalf("Expected 1 event, got %d", len(events))
|
||||
}
|
||||
|
||||
retrieved := events[0]
|
||||
if retrieved.ID != event.ID {
|
||||
t.Errorf("ID mismatch: expected %s, got %s", event.ID, retrieved.ID)
|
||||
}
|
||||
if retrieved.EventType != event.EventType {
|
||||
t.Errorf("EventType mismatch: expected %s, got %s", event.EventType, retrieved.EventType)
|
||||
}
|
||||
if retrieved.User != event.User {
|
||||
t.Errorf("User mismatch: expected %s, got %s", event.User, retrieved.User)
|
||||
}
|
||||
if retrieved.Success != event.Success {
|
||||
t.Errorf("Success mismatch: expected %v, got %v", event.Success, retrieved.Success)
|
||||
}
|
||||
|
||||
// Event should have a signature
|
||||
if retrieved.Signature == "" {
|
||||
t.Error("Expected event to have a signature")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerQuery(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
// Log several events
|
||||
now := time.Now()
|
||||
events := []Event{
|
||||
{ID: "e1", Timestamp: now.Add(-2 * time.Hour), EventType: "login", User: "alice", Success: true},
|
||||
{ID: "e2", Timestamp: now.Add(-1 * time.Hour), EventType: "login", User: "bob", Success: true},
|
||||
{ID: "e3", Timestamp: now, EventType: "logout", User: "alice", Success: true},
|
||||
{ID: "e4", Timestamp: now, EventType: "login", User: "charlie", Success: false},
|
||||
}
|
||||
|
||||
for _, e := range events {
|
||||
if err := logger.Log(e); err != nil {
|
||||
t.Fatalf("Log failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Test event type filter
|
||||
t.Run("FilterByEventType", func(t *testing.T) {
|
||||
results, err := logger.Query(QueryFilter{EventType: "login"})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(results) != 3 {
|
||||
t.Errorf("Expected 3 login events, got %d", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
// Test user filter
|
||||
t.Run("FilterByUser", func(t *testing.T) {
|
||||
results, err := logger.Query(QueryFilter{User: "alice"})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(results) != 2 {
|
||||
t.Errorf("Expected 2 events for alice, got %d", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
// Test success filter
|
||||
t.Run("FilterBySuccess", func(t *testing.T) {
|
||||
success := false
|
||||
results, err := logger.Query(QueryFilter{Success: &success})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Errorf("Expected 1 failed event, got %d", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
// Test time range filter
|
||||
t.Run("FilterByTimeRange", func(t *testing.T) {
|
||||
start := now.Add(-90 * time.Minute)
|
||||
end := now.Add(-30 * time.Minute)
|
||||
results, err := logger.Query(QueryFilter{StartTime: &start, EndTime: &end})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(results) != 1 {
|
||||
t.Errorf("Expected 1 event in time range, got %d", len(results))
|
||||
}
|
||||
})
|
||||
|
||||
// Test limit and offset
|
||||
t.Run("LimitAndOffset", func(t *testing.T) {
|
||||
results, err := logger.Query(QueryFilter{Limit: 2, Offset: 1})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(results) != 2 {
|
||||
t.Errorf("Expected 2 events with limit, got %d", len(results))
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerCount(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
// Log several events
|
||||
for i := 0; i < 5; i++ {
|
||||
event := Event{
|
||||
ID: uuid.NewString(),
|
||||
Timestamp: time.Now(),
|
||||
EventType: "test",
|
||||
Success: i%2 == 0,
|
||||
}
|
||||
if err := logger.Log(event); err != nil {
|
||||
t.Fatalf("Log failed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Count all
|
||||
count, err := logger.Count(QueryFilter{})
|
||||
if err != nil {
|
||||
t.Fatalf("Count failed: %v", err)
|
||||
}
|
||||
if count != 5 {
|
||||
t.Errorf("Expected count 5, got %d", count)
|
||||
}
|
||||
|
||||
// Count successful
|
||||
success := true
|
||||
count, err = logger.Count(QueryFilter{Success: &success})
|
||||
if err != nil {
|
||||
t.Fatalf("Count failed: %v", err)
|
||||
}
|
||||
if count != 3 {
|
||||
t.Errorf("Expected 3 successful events, got %d", count)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerVerifySignature(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
event := Event{
|
||||
ID: uuid.NewString(),
|
||||
Timestamp: time.Now(),
|
||||
EventType: "verify_test",
|
||||
User: "testuser",
|
||||
Success: true,
|
||||
}
|
||||
|
||||
if err := logger.Log(event); err != nil {
|
||||
t.Fatalf("Log failed: %v", err)
|
||||
}
|
||||
|
||||
// Query the event back
|
||||
events, err := logger.Query(QueryFilter{ID: event.ID})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
|
||||
if len(events) != 1 {
|
||||
t.Fatalf("Expected 1 event, got %d", len(events))
|
||||
}
|
||||
|
||||
// Verify should succeed
|
||||
if !logger.VerifySignature(events[0]) {
|
||||
t.Error("VerifySignature should return true for valid event")
|
||||
}
|
||||
|
||||
// Tamper with event and verify should fail
|
||||
tamperedEvent := events[0]
|
||||
tamperedEvent.Details = "tampered"
|
||||
if logger.VerifySignature(tamperedEvent) {
|
||||
t.Error("VerifySignature should return false for tampered event")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerWebhooks(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
// Initially no webhooks
|
||||
urls := logger.GetWebhookURLs()
|
||||
if len(urls) != 0 {
|
||||
t.Errorf("Expected no webhooks initially, got %d", len(urls))
|
||||
}
|
||||
|
||||
// Add webhooks
|
||||
testURLs := []string{"https://example.com/webhook1", "https://example.com/webhook2"}
|
||||
err = logger.UpdateWebhookURLs(testURLs)
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateWebhookURLs failed: %v", err)
|
||||
}
|
||||
|
||||
urls = logger.GetWebhookURLs()
|
||||
if len(urls) != 2 {
|
||||
t.Errorf("Expected 2 webhooks, got %d", len(urls))
|
||||
}
|
||||
|
||||
// Clear webhooks
|
||||
err = logger.UpdateWebhookURLs([]string{})
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateWebhookURLs (clear) failed: %v", err)
|
||||
}
|
||||
|
||||
urls = logger.GetWebhookURLs()
|
||||
if len(urls) != 0 {
|
||||
t.Errorf("Expected 0 webhooks after clear, got %d", len(urls))
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerRetention(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 1, // 1 day retention for testing
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
// Log an old event (2 days ago)
|
||||
oldEvent := Event{
|
||||
ID: "old-event",
|
||||
Timestamp: time.Now().Add(-48 * time.Hour),
|
||||
EventType: "old",
|
||||
Success: true,
|
||||
}
|
||||
if err := logger.Log(oldEvent); err != nil {
|
||||
t.Fatalf("Log old event failed: %v", err)
|
||||
}
|
||||
|
||||
// Log a recent event
|
||||
newEvent := Event{
|
||||
ID: "new-event",
|
||||
Timestamp: time.Now(),
|
||||
EventType: "new",
|
||||
Success: true,
|
||||
}
|
||||
if err := logger.Log(newEvent); err != nil {
|
||||
t.Fatalf("Log new event failed: %v", err)
|
||||
}
|
||||
|
||||
// Run cleanup
|
||||
logger.cleanupOldEvents()
|
||||
|
||||
// Old event should be deleted
|
||||
events, err := logger.Query(QueryFilter{ID: "old-event"})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(events) != 0 {
|
||||
t.Error("Old event should have been deleted")
|
||||
}
|
||||
|
||||
// New event should still exist
|
||||
events, err = logger.Query(QueryFilter{ID: "new-event"})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(events) != 1 {
|
||||
t.Error("New event should still exist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerSetRetentionDays(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
logger.SetRetentionDays(60)
|
||||
if logger.GetRetentionDays() != 60 {
|
||||
t.Errorf("Expected retention days 60, got %d", logger.GetRetentionDays())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerPersistence(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
// Create logger and log an event
|
||||
logger1, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
|
||||
event := Event{
|
||||
ID: "persist-test",
|
||||
Timestamp: time.Now(),
|
||||
EventType: "persistence_test",
|
||||
User: "testuser",
|
||||
Success: true,
|
||||
}
|
||||
if err := logger1.Log(event); err != nil {
|
||||
t.Fatalf("Log failed: %v", err)
|
||||
}
|
||||
|
||||
// Close the logger
|
||||
logger1.Close()
|
||||
|
||||
// Create a new logger with same data dir
|
||||
logger2, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger (reload) failed: %v", err)
|
||||
}
|
||||
defer logger2.Close()
|
||||
|
||||
// Query the event - should still exist
|
||||
events, err := logger2.Query(QueryFilter{ID: "persist-test"})
|
||||
if err != nil {
|
||||
t.Fatalf("Query failed: %v", err)
|
||||
}
|
||||
if len(events) != 1 {
|
||||
t.Error("Event should persist across logger restarts")
|
||||
}
|
||||
|
||||
// Signature should still verify
|
||||
if len(events) > 0 && !logger2.VerifySignature(events[0]) {
|
||||
t.Error("Signature should still verify after restart")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSQLiteLoggerConcurrentAccess(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
|
||||
logger, err := NewSQLiteLogger(SQLiteLoggerConfig{
|
||||
DataDir: tempDir,
|
||||
CryptoMgr: newMockCryptoManager(),
|
||||
RetentionDays: 30,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewSQLiteLogger failed: %v", err)
|
||||
}
|
||||
defer logger.Close()
|
||||
|
||||
// Concurrent writes
|
||||
done := make(chan bool)
|
||||
for i := 0; i < 10; i++ {
|
||||
go func(n int) {
|
||||
for j := 0; j < 10; j++ {
|
||||
event := Event{
|
||||
ID: uuid.NewString(),
|
||||
Timestamp: time.Now(),
|
||||
EventType: "concurrent_test",
|
||||
User: "user",
|
||||
Success: true,
|
||||
}
|
||||
if err := logger.Log(event); err != nil {
|
||||
t.Errorf("Concurrent log failed: %v", err)
|
||||
}
|
||||
}
|
||||
done <- true
|
||||
}(i)
|
||||
}
|
||||
|
||||
// Wait for all goroutines
|
||||
for i := 0; i < 10; i++ {
|
||||
<-done
|
||||
}
|
||||
|
||||
// Verify count
|
||||
count, err := logger.Count(QueryFilter{EventType: "concurrent_test"})
|
||||
if err != nil {
|
||||
t.Fatalf("Count failed: %v", err)
|
||||
}
|
||||
if count != 100 {
|
||||
t.Errorf("Expected 100 events, got %d", count)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
package audit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/rs/zerolog/log"
|
||||
)
|
||||
|
||||
const (
|
||||
webhookQueueSize = 1000
|
||||
webhookMaxRetries = 3
|
||||
webhookTimeout = 30 * time.Second
|
||||
webhookWorkerCount = 3
|
||||
)
|
||||
|
||||
var webhookBackoff = []time.Duration{1 * time.Second, 5 * time.Second, 30 * time.Second}
|
||||
|
||||
// WebhookDelivery handles async webhook delivery with retries.
|
||||
type WebhookDelivery struct {
|
||||
mu sync.RWMutex
|
||||
urls []string
|
||||
client *http.Client
|
||||
queue chan Event
|
||||
stopChan chan struct{}
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
// WebhookPayload is the JSON payload sent to webhooks.
|
||||
type WebhookPayload struct {
|
||||
Event string `json:"event"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
Data Event `json:"data"`
|
||||
}
|
||||
|
||||
// NewWebhookDelivery creates a new webhook delivery worker.
|
||||
func NewWebhookDelivery(urls []string) *WebhookDelivery {
|
||||
return &WebhookDelivery{
|
||||
urls: urls,
|
||||
client: &http.Client{
|
||||
Timeout: webhookTimeout,
|
||||
Transport: &http.Transport{
|
||||
MaxIdleConns: 10,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
DisableCompression: true,
|
||||
MaxConnsPerHost: 5,
|
||||
MaxIdleConnsPerHost: 2,
|
||||
},
|
||||
},
|
||||
queue: make(chan Event, webhookQueueSize),
|
||||
stopChan: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Start begins the delivery worker goroutines.
|
||||
func (w *WebhookDelivery) Start() {
|
||||
for i := 0; i < webhookWorkerCount; i++ {
|
||||
w.wg.Add(1)
|
||||
go w.worker(i)
|
||||
}
|
||||
log.Debug().Int("workers", webhookWorkerCount).Msg("Audit webhook delivery started")
|
||||
}
|
||||
|
||||
// Stop gracefully stops the delivery workers.
|
||||
func (w *WebhookDelivery) Stop() {
|
||||
close(w.stopChan)
|
||||
w.wg.Wait()
|
||||
log.Debug().Msg("Audit webhook delivery stopped")
|
||||
}
|
||||
|
||||
// Enqueue adds an event to the delivery queue.
|
||||
// Non-blocking - drops events if queue is full.
|
||||
func (w *WebhookDelivery) Enqueue(event Event) {
|
||||
select {
|
||||
case w.queue <- event:
|
||||
// Enqueued successfully
|
||||
default:
|
||||
log.Warn().
|
||||
Str("event_id", event.ID).
|
||||
Str("event_type", event.EventType).
|
||||
Msg("Audit webhook queue full, dropping event")
|
||||
}
|
||||
}
|
||||
|
||||
// UpdateURLs updates the webhook URLs.
|
||||
func (w *WebhookDelivery) UpdateURLs(urls []string) {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
w.urls = urls
|
||||
}
|
||||
|
||||
// GetURLs returns the current webhook URLs.
|
||||
func (w *WebhookDelivery) GetURLs() []string {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
result := make([]string, len(w.urls))
|
||||
copy(result, w.urls)
|
||||
return result
|
||||
}
|
||||
|
||||
// worker processes events from the queue.
|
||||
func (w *WebhookDelivery) worker(id int) {
|
||||
defer w.wg.Done()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-w.stopChan:
|
||||
// Drain remaining events on shutdown (with timeout)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
w.drainQueue(ctx)
|
||||
cancel()
|
||||
return
|
||||
case event := <-w.queue:
|
||||
w.deliverToAll(event)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// drainQueue processes remaining events during shutdown.
|
||||
func (w *WebhookDelivery) drainQueue(ctx context.Context) {
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
remaining := len(w.queue)
|
||||
if remaining > 0 {
|
||||
log.Warn().Int("remaining", remaining).Msg("Audit webhook shutdown timeout, dropping remaining events")
|
||||
}
|
||||
return
|
||||
case event := <-w.queue:
|
||||
w.deliverToAll(event)
|
||||
default:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// deliverToAll sends an event to all configured webhooks.
|
||||
func (w *WebhookDelivery) deliverToAll(event Event) {
|
||||
w.mu.RLock()
|
||||
urls := make([]string, len(w.urls))
|
||||
copy(urls, w.urls)
|
||||
w.mu.RUnlock()
|
||||
|
||||
for _, url := range urls {
|
||||
if err := w.deliverWithRetry(url, event); err != nil {
|
||||
log.Error().
|
||||
Err(err).
|
||||
Str("url", url).
|
||||
Str("event_id", event.ID).
|
||||
Msg("Failed to deliver audit webhook after retries")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// deliverWithRetry attempts to deliver an event with exponential backoff.
|
||||
func (w *WebhookDelivery) deliverWithRetry(url string, event Event) error {
|
||||
var lastErr error
|
||||
|
||||
for attempt := 0; attempt <= webhookMaxRetries; attempt++ {
|
||||
if attempt > 0 {
|
||||
// Wait before retry
|
||||
backoffIdx := attempt - 1
|
||||
if backoffIdx >= len(webhookBackoff) {
|
||||
backoffIdx = len(webhookBackoff) - 1
|
||||
}
|
||||
time.Sleep(webhookBackoff[backoffIdx])
|
||||
}
|
||||
|
||||
err := w.deliver(url, event)
|
||||
if err == nil {
|
||||
if attempt > 0 {
|
||||
log.Debug().
|
||||
Str("url", url).
|
||||
Str("event_id", event.ID).
|
||||
Int("attempt", attempt+1).
|
||||
Msg("Audit webhook delivered after retry")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
lastErr = err
|
||||
log.Debug().
|
||||
Err(err).
|
||||
Str("url", url).
|
||||
Str("event_id", event.ID).
|
||||
Int("attempt", attempt+1).
|
||||
Int("maxAttempts", webhookMaxRetries+1).
|
||||
Msg("Audit webhook delivery attempt failed")
|
||||
}
|
||||
|
||||
return lastErr
|
||||
}
|
||||
|
||||
// deliver sends a single event to a webhook URL.
|
||||
func (w *WebhookDelivery) deliver(url string, event Event) error {
|
||||
payload := WebhookPayload{
|
||||
Event: "audit." + event.EventType,
|
||||
Timestamp: event.Timestamp,
|
||||
Data: event,
|
||||
}
|
||||
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal webhook payload: %w", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), webhookTimeout)
|
||||
defer cancel()
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create webhook request: %w", err)
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("User-Agent", "Pulse-Audit-Webhook/1.0")
|
||||
req.Header.Set("X-Pulse-Event", event.EventType)
|
||||
req.Header.Set("X-Pulse-Event-ID", event.ID)
|
||||
|
||||
resp, err := w.client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("webhook request failed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// Consider 2xx as success
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("webhook returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// QueueLength returns the current number of events in the queue.
|
||||
func (w *WebhookDelivery) QueueLength() int {
|
||||
return len(w.queue)
|
||||
}
|
||||
Reference in New Issue
Block a user