package audit

import (
	"context"
	"crypto/sha256"
	"encoding/hex"
	"encoding/json"
	"time"

	"github.com/jackc/pgx/v5/pgxpool"

	"github.com/niels/banking-app/backend/internal/domain"
)

type Repository struct {
	db *pgxpool.Pool
}

type Event struct {
	ActorUserID *string
	EventType   string
	TargetType  string
	TargetID    string
	Metadata    any
	RemoteIP    string
	UserAgent   string
}

type IntegrityReport struct {
	CheckedEvents int       `json:"checked_events"`
	MissingHashes int       `json:"missing_hashes"`
	BrokenHashes  int       `json:"broken_hashes"`
	Valid         bool      `json:"valid"`
	CheckedAt     time.Time `json:"checked_at"`
}

func NewRepository(db *pgxpool.Pool) *Repository {
	return &Repository{db: db}
}

func (r *Repository) Record(ctx context.Context, event Event) error {
	metadata := []byte(`{}`)
	if event.Metadata != nil {
		encoded, err := json.Marshal(event.Metadata)
		if err != nil {
			return err
		}
		metadata = encoded
	}

	var actor any
	if event.ActorUserID != nil && *event.ActorUserID != "" {
		actor = *event.ActorUserID
	}

	tx, err := r.db.Begin(ctx)
	if err != nil {
		return err
	}
	defer tx.Rollback(ctx)

	var previousHash string
	if err := tx.QueryRow(ctx, `
		SELECT event_hash
		FROM audit_events
		WHERE event_hash <> ''
		ORDER BY created_at DESC, id DESC
		LIMIT 1
	`).Scan(&previousHash); err != nil {
		previousHash = ""
	}

	var id string
	var createdAt time.Time
	err = tx.QueryRow(ctx, `
		INSERT INTO audit_events (actor_user_id, event_type, target_type, target_id, metadata, remote_ip, user_agent)
		VALUES ($1, $2, $3, $4, $5, $6, $7)
		RETURNING id::text, created_at
	`, actor, event.EventType, event.TargetType, event.TargetID, metadata, event.RemoteIP, event.UserAgent).Scan(&id, &createdAt)
	if err != nil {
		return err
	}

	actorString := ""
	if event.ActorUserID != nil {
		actorString = *event.ActorUserID
	}
	eventHash := hashAuditEvent(id, previousHash, actorString, event.EventType, event.TargetType, event.TargetID, metadata, event.RemoteIP, event.UserAgent, createdAt)
	if _, err := tx.Exec(ctx, `
		UPDATE audit_events
		SET previous_hash = $2, event_hash = $3, hash_algorithm = 'sha256'
		WHERE id = $1
	`, id, previousHash, eventHash); err != nil {
		return err
	}

	return tx.Commit(ctx)
}

func (r *Repository) List(ctx context.Context, limit int) ([]domain.AuditEvent, error) {
	if limit <= 0 || limit > 100 {
		limit = 50
	}

	rows, err := r.db.Query(ctx, `
		SELECT id::text, COALESCE(actor_user_id::text, ''), event_type, target_type, target_id, metadata, remote_ip, user_agent, created_at
		FROM audit_events
		ORDER BY created_at DESC
		LIMIT $1
	`, limit)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	events := make([]domain.AuditEvent, 0, limit)
	for rows.Next() {
		var event domain.AuditEvent
		if err := rows.Scan(&event.ID, &event.ActorUserID, &event.EventType, &event.TargetType, &event.TargetID, &event.Metadata, &event.RemoteIP, &event.UserAgent, &event.CreatedAt); err != nil {
			return nil, err
		}
		events = append(events, event)
	}

	return events, rows.Err()
}

func (r *Repository) VerifyIntegrity(ctx context.Context, limit int) (IntegrityReport, error) {
	if limit <= 0 || limit > 10000 {
		limit = 10000
	}
	rows, err := r.db.Query(ctx, `
		SELECT id::text, COALESCE(previous_hash, ''), COALESCE(event_hash, ''), COALESCE(actor_user_id::text, ''),
			event_type, target_type, target_id, metadata, remote_ip, user_agent, created_at
		FROM audit_events
		ORDER BY created_at ASC, id ASC
		LIMIT $1
	`, limit)
	if err != nil {
		return IntegrityReport{}, err
	}
	defer rows.Close()

	report := IntegrityReport{Valid: true, CheckedAt: time.Now().UTC()}
	for rows.Next() {
		var id, previousHash, eventHash, actorID, eventType, targetType, targetID, remoteIP, userAgent string
		var metadata json.RawMessage
		var createdAt time.Time
		if err := rows.Scan(&id, &previousHash, &eventHash, &actorID, &eventType, &targetType, &targetID, &metadata, &remoteIP, &userAgent, &createdAt); err != nil {
			return IntegrityReport{}, err
		}
		report.CheckedEvents++
		if eventHash == "" {
			report.MissingHashes++
			report.Valid = false
			continue
		}
		expected := hashAuditEvent(id, previousHash, actorID, eventType, targetType, targetID, metadata, remoteIP, userAgent, createdAt)
		if expected != eventHash {
			report.BrokenHashes++
			report.Valid = false
		}
	}
	return report, rows.Err()
}

func hashAuditEvent(id, previousHash, actorID, eventType, targetType, targetID string, metadata []byte, remoteIP, userAgent string, createdAt time.Time) string {
	payload := id + "\x1f" +
		previousHash + "\x1f" +
		actorID + "\x1f" +
		eventType + "\x1f" +
		targetType + "\x1f" +
		targetID + "\x1f" +
		string(metadata) + "\x1f" +
		remoteIP + "\x1f" +
		userAgent + "\x1f" +
		createdAt.UTC().Format(time.RFC3339Nano)
	sum := sha256.Sum256([]byte(payload))
	return hex.EncodeToString(sum[:])
}
