package authguard

import (
	"context"
	"database/sql"
	"encoding/json"
	"errors"
	"fmt"
	"strings"
	"time"

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

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

const (
	EventLoginFailedVelocity = "login_failed_velocity"
	EventLoginLockoutCreated = "login_lockout_created"
	EventLoginBlocked        = "login_blocked"
	EventLoginRateLimited    = "login_rate_limited"
	EventSuspiciousLogin     = "suspicious_login"
)

type Config struct {
	IdentifierWindow       time.Duration
	IdentifierMaxFailures  int
	IPWindow               time.Duration
	IPMaxFailures          int
	LockoutDuration        time.Duration
	SuspiciousLookback     time.Duration
	RecentFailureLookback  time.Duration
	RecentFailureThreshold int
	DistinctIPThreshold24h int
}

type Repository struct {
	db  *pgxpool.Pool
	cfg Config
}

type SecurityEvent struct {
	ID                    string          `json:"id"`
	UserID                string          `json:"user_id,omitempty"`
	UserEmail             string          `json:"user_email,omitempty"`
	UserFullName          string          `json:"user_full_name,omitempty"`
	EventType             string          `json:"event_type"`
	Severity              string          `json:"severity"`
	Status                string          `json:"status"`
	Identifier            string          `json:"identifier"`
	RemoteIP              string          `json:"remote_ip"`
	UserAgent             string          `json:"user_agent,omitempty"`
	Details               json.RawMessage `json:"details"`
	ResolvedByAdminUserID string          `json:"resolved_by_admin_user_id,omitempty"`
	ResolutionNote        string          `json:"resolution_note,omitempty"`
	ResolvedAt            *time.Time      `json:"resolved_at,omitempty"`
	CreatedAt             time.Time       `json:"created_at"`
	UpdatedAt             time.Time       `json:"updated_at"`
}

type Metrics struct {
	OpenEvents         int64 `json:"open_events"`
	HighOpenEvents     int64 `json:"high_open_events"`
	CriticalOpenEvents int64 `json:"critical_open_events"`
	ActiveLockouts     int64 `json:"active_lockouts"`
	Failed15m          int64 `json:"failed_15m"`
	Blocked24h         int64 `json:"blocked_24h"`
	Suspicious24h      int64 `json:"suspicious_24h"`
	Successful24h      int64 `json:"successful_24h"`
}

type Dashboard struct {
	Metrics Metrics         `json:"metrics"`
	Events  []SecurityEvent `json:"events"`
}

type LoginSecurityResult struct {
	Suspicious bool     `json:"suspicious"`
	Reasons    []string `json:"reasons"`
	Severity   string   `json:"severity"`
}

type EventDecisionParams struct {
	EventID        string
	AdminUserID    string
	Status         string
	ResolutionNote string
}

func DefaultConfig() Config {
	return Config{
		IdentifierWindow:       15 * time.Minute,
		IdentifierMaxFailures:  5,
		IPWindow:               15 * time.Minute,
		IPMaxFailures:          30,
		LockoutDuration:        15 * time.Minute,
		SuspiciousLookback:     90 * 24 * time.Hour,
		RecentFailureLookback:  time.Hour,
		RecentFailureThreshold: 3,
		DistinctIPThreshold24h: 3,
	}
}

func NewRepository(db *pgxpool.Pool, cfg Config) *Repository {
	cfg = normalizeConfig(cfg)
	return &Repository{db: db, cfg: cfg}
}

func (r *Repository) CheckLoginAllowed(ctx context.Context, identifier, remoteIP, userAgent string) error {
	identifier = normalizeIdentifier(identifier)
	remoteIP = normalizeRemoteIP(remoteIP)
	userAgent = strings.TrimSpace(userAgent)

	var lockedUntil time.Time
	err := r.db.QueryRow(ctx, `
		SELECT locked_until
		FROM auth_lockouts
		WHERE released_at IS NULL
			AND locked_until > now()
			AND (identifier = $1 OR remote_ip = $2)
		ORDER BY locked_until DESC
		LIMIT 1
	`, identifier, remoteIP).Scan(&lockedUntil)
	if errors.Is(err, pgx.ErrNoRows) {
		return nil
	}
	if err != nil {
		return err
	}

	_ = r.recordSecurityEvent(ctx, nil, EventLoginBlocked, "high", identifier, remoteIP, userAgent, map[string]any{
		"locked_until": lockedUntil,
	})
	return fmt.Errorf("%w: login is temporarily locked until %s", domain.ErrRateLimited, lockedUntil.UTC().Format(time.RFC3339))
}

func (r *Repository) RecordLoginFailure(ctx context.Context, userID, identifier, remoteIP, userAgent, reason string) error {
	identifier = normalizeIdentifier(identifier)
	remoteIP = normalizeRemoteIP(remoteIP)
	userAgent = strings.TrimSpace(userAgent)
	reason = strings.TrimSpace(reason)
	if reason == "" {
		reason = "invalid_credentials"
	}

	tx, err := r.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.Serializable})
	if err != nil {
		return err
	}
	defer tx.Rollback(ctx)

	if err := insertLoginAttempt(ctx, tx, userID, identifier, remoteIP, userAgent, false, reason); err != nil {
		return err
	}
	identifierFailures, err := countFailures(ctx, tx, "identifier", identifier, r.cfg.IdentifierWindow)
	if err != nil {
		return err
	}
	ipFailures, err := countFailures(ctx, tx, "remote_ip", remoteIP, r.cfg.IPWindow)
	if err != nil {
		return err
	}

	if identifierFailures >= int64(r.cfg.IdentifierMaxFailures) || ipFailures >= int64(r.cfg.IPMaxFailures) {
		lockReason := "identifier_failure_threshold"
		failedCount := identifierFailures
		if ipFailures >= int64(r.cfg.IPMaxFailures) {
			lockReason = "ip_failure_threshold"
			failedCount = ipFailures
		}
		if err := createLockoutIfAbsent(ctx, tx, userID, identifier, remoteIP, lockReason, int(failedCount), r.cfg.LockoutDuration); err != nil {
			return err
		}
		if err := insertSecurityEvent(ctx, tx, userID, EventLoginLockoutCreated, "high", identifier, remoteIP, userAgent, map[string]any{
			"reason":              lockReason,
			"identifier_failures": identifierFailures,
			"ip_failures":         ipFailures,
			"lockout_seconds":     int(r.cfg.LockoutDuration.Seconds()),
		}); err != nil {
			return err
		}
	} else if identifierFailures+1 >= int64(r.cfg.IdentifierMaxFailures) || ipFailures+3 >= int64(r.cfg.IPMaxFailures) {
		if err := insertSecurityEvent(ctx, tx, userID, EventLoginFailedVelocity, "medium", identifier, remoteIP, userAgent, map[string]any{
			"identifier_failures": identifierFailures,
			"ip_failures":         ipFailures,
		}); err != nil {
			return err
		}
	}

	return tx.Commit(ctx)
}

func (r *Repository) RecordLoginSuccessAndDetect(ctx context.Context, userID, identifier, remoteIP, userAgent string) (LoginSecurityResult, error) {
	identifier = normalizeIdentifier(identifier)
	remoteIP = normalizeRemoteIP(remoteIP)
	userAgent = strings.TrimSpace(userAgent)

	tx, err := r.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead})
	if err != nil {
		return LoginSecurityResult{}, err
	}
	defer tx.Rollback(ctx)

	if err := insertLoginAttempt(ctx, tx, userID, identifier, remoteIP, userAgent, true, ""); err != nil {
		return LoginSecurityResult{}, err
	}

	signals, severity, err := suspiciousSignals(ctx, tx, userID, identifier, remoteIP, userAgent, r.cfg)
	if err != nil {
		return LoginSecurityResult{}, err
	}
	result := LoginSecurityResult{
		Suspicious: len(signals) > 0,
		Reasons:    signals,
		Severity:   severity,
	}
	if result.Suspicious {
		if err := insertSecurityEvent(ctx, tx, userID, EventSuspiciousLogin, severity, identifier, remoteIP, userAgent, map[string]any{
			"signals": signals,
		}); err != nil {
			return LoginSecurityResult{}, err
		}
	}

	if err := tx.Commit(ctx); err != nil {
		return LoginSecurityResult{}, err
	}
	return result, nil
}

func (r *Repository) Dashboard(ctx context.Context, status string, limit int) (Dashboard, error) {
	metrics, err := r.Metrics(ctx)
	if err != nil {
		return Dashboard{}, err
	}
	events, err := r.ListEvents(ctx, status, limit)
	if err != nil {
		return Dashboard{}, err
	}
	return Dashboard{Metrics: metrics, Events: events}, nil
}

func (r *Repository) Metrics(ctx context.Context) (Metrics, error) {
	var metrics Metrics
	err := r.db.QueryRow(ctx, `
		SELECT
			(SELECT COUNT(*) FROM security_events WHERE status = 'open')::bigint,
			(SELECT COUNT(*) FROM security_events WHERE status = 'open' AND severity = 'high')::bigint,
			(SELECT COUNT(*) FROM security_events WHERE status = 'open' AND severity = 'critical')::bigint,
			(SELECT COUNT(*) FROM auth_lockouts WHERE released_at IS NULL AND locked_until > now())::bigint,
			(SELECT COUNT(*) FROM auth_login_attempts WHERE success = false AND created_at >= now() - interval '15 minutes')::bigint,
			(SELECT COUNT(*) FROM security_events WHERE event_type IN ('login_blocked', 'login_rate_limited') AND created_at >= now() - interval '24 hours')::bigint,
			(SELECT COUNT(*) FROM security_events WHERE event_type = 'suspicious_login' AND created_at >= now() - interval '24 hours')::bigint,
			(SELECT COUNT(*) FROM auth_login_attempts WHERE success = true AND created_at >= now() - interval '24 hours')::bigint
	`).Scan(
		&metrics.OpenEvents,
		&metrics.HighOpenEvents,
		&metrics.CriticalOpenEvents,
		&metrics.ActiveLockouts,
		&metrics.Failed15m,
		&metrics.Blocked24h,
		&metrics.Suspicious24h,
		&metrics.Successful24h,
	)
	return metrics, err
}

func (r *Repository) ListEvents(ctx context.Context, status string, limit int) ([]SecurityEvent, error) {
	limit = normalizeLimit(limit)
	status = strings.ToLower(strings.TrimSpace(status))
	if status == "" {
		status = "open"
	}
	if status != "all" && status != "open" && status != "reviewing" && status != "resolved" && status != "ignored" {
		return nil, fmt.Errorf("%w: invalid security event status", domain.ErrValidation)
	}

	rows, err := r.db.Query(ctx, `
		SELECT se.id::text, COALESCE(se.user_id::text, ''), COALESCE(u.email, ''), COALESCE(u.full_name, ''),
			se.event_type, se.severity, se.status, se.identifier, se.remote_ip, COALESCE(se.user_agent, ''),
			se.details, COALESCE(se.resolved_by_admin_user_id::text, ''), COALESCE(se.resolution_note, ''),
			se.resolved_at, se.created_at, se.updated_at
		FROM security_events se
		LEFT JOIN users u ON u.id = se.user_id
		WHERE ($1 = 'all' OR se.status = $1)
		ORDER BY
			CASE se.severity WHEN 'critical' THEN 1 WHEN 'high' THEN 2 WHEN 'medium' THEN 3 ELSE 4 END,
			se.created_at DESC
		LIMIT $2
	`, status, limit)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	events := []SecurityEvent{}
	for rows.Next() {
		event, err := scanSecurityEvent(rows)
		if err != nil {
			return nil, err
		}
		events = append(events, event)
	}
	return events, rows.Err()
}

func (r *Repository) DecideEvent(ctx context.Context, params EventDecisionParams) (SecurityEvent, error) {
	if err := normalizeEventDecision(&params); err != nil {
		return SecurityEvent{}, err
	}

	var resolvedBy any
	var resolvedAt any
	if params.Status == "resolved" || params.Status == "ignored" {
		resolvedBy = params.AdminUserID
		resolvedAt = time.Now().UTC()
	}
	row := r.db.QueryRow(ctx, `
		UPDATE security_events
		SET status = $2,
			resolution_note = NULLIF($3, ''),
			resolved_by_admin_user_id = $4,
			resolved_at = $5::timestamptz
		WHERE id = $1
		RETURNING id::text, COALESCE(user_id::text, ''), ''::text, ''::text,
			event_type, severity, status, identifier, remote_ip, COALESCE(user_agent, ''),
			details, COALESCE(resolved_by_admin_user_id::text, ''), COALESCE(resolution_note, ''),
			resolved_at, created_at, updated_at
	`, params.EventID, params.Status, params.ResolutionNote, resolvedBy, resolvedAt)

	event, err := scanSecurityEvent(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return SecurityEvent{}, domain.ErrNotFound
	}
	return event, err
}

func (r *Repository) recordSecurityEvent(ctx context.Context, userID *string, eventType, severity, identifier, remoteIP, userAgent string, details map[string]any) error {
	value := ""
	if userID != nil {
		value = *userID
	}
	tx, err := r.db.BeginTx(ctx, pgx.TxOptions{})
	if err != nil {
		return err
	}
	defer tx.Rollback(ctx)
	if err := insertSecurityEvent(ctx, tx, value, eventType, severity, identifier, remoteIP, userAgent, details); err != nil {
		return err
	}
	return tx.Commit(ctx)
}

func insertLoginAttempt(ctx context.Context, tx pgx.Tx, userID, identifier, remoteIP, userAgent string, success bool, failureReason string) error {
	var user any
	if strings.TrimSpace(userID) != "" {
		user = strings.TrimSpace(userID)
	}
	_, err := tx.Exec(ctx, `
		INSERT INTO auth_login_attempts (user_id, identifier, remote_ip, user_agent, success, failure_reason)
		VALUES ($1, $2, $3, NULLIF($4, ''), $5, NULLIF($6, ''))
	`, user, identifier, remoteIP, userAgent, success, failureReason)
	return err
}

func countFailures(ctx context.Context, tx pgx.Tx, column, value string, window time.Duration) (int64, error) {
	if column != "identifier" && column != "remote_ip" {
		return 0, fmt.Errorf("%w: invalid failure counter", domain.ErrValidation)
	}
	var count int64
	err := tx.QueryRow(ctx, fmt.Sprintf(`
		SELECT COUNT(*)::bigint
		FROM auth_login_attempts
		WHERE success = false
			AND %s = $1
			AND created_at >= now() - make_interval(secs => $2)
	`, column), value, int(window.Seconds())).Scan(&count)
	return count, err
}

func createLockoutIfAbsent(ctx context.Context, tx pgx.Tx, userID, identifier, remoteIP, reason string, failedCount int, duration time.Duration) error {
	var exists bool
	if err := tx.QueryRow(ctx, `
		SELECT EXISTS (
			SELECT 1
			FROM auth_lockouts
			WHERE released_at IS NULL
				AND locked_until > now()
				AND (identifier = $1 OR remote_ip = $2)
		)
	`, identifier, remoteIP).Scan(&exists); err != nil {
		return err
	}
	if exists {
		return nil
	}
	var user any
	if strings.TrimSpace(userID) != "" {
		user = strings.TrimSpace(userID)
	}
	_, err := tx.Exec(ctx, `
		INSERT INTO auth_lockouts (user_id, identifier, remote_ip, reason, failed_count, locked_until)
		VALUES ($1, $2, $3, $4, $5, now() + make_interval(secs => $6))
	`, user, identifier, remoteIP, reason, failedCount, int(duration.Seconds()))
	return err
}

func suspiciousSignals(ctx context.Context, tx pgx.Tx, userID, identifier, remoteIP, userAgent string, cfg Config) ([]string, string, error) {
	var previousSessions int64
	var seenIP bool
	var seenUserAgent bool
	var distinctIPs24h int64
	var recentFailures int64
	if err := tx.QueryRow(ctx, `
		SELECT
			(SELECT COUNT(*) FROM auth_sessions WHERE user_id = $1)::bigint,
			(SELECT EXISTS (
				SELECT 1 FROM auth_sessions
				WHERE user_id = $1 AND remote_ip = $2 AND created_at >= now() - make_interval(secs => $4)
			)),
			(SELECT EXISTS (
				SELECT 1 FROM auth_sessions
				WHERE user_id = $1 AND COALESCE(user_agent, '') = $3 AND created_at >= now() - make_interval(secs => $4)
			)),
			(SELECT COUNT(DISTINCT remote_ip)::bigint FROM auth_sessions
				WHERE user_id = $1 AND remote_ip IS NOT NULL AND created_at >= now() - interval '24 hours'),
			(SELECT COUNT(*)::bigint FROM auth_login_attempts
				WHERE success = false AND (identifier = $5 OR remote_ip = $2)
					AND created_at >= now() - make_interval(secs => $6))
	`, userID, remoteIP, userAgent, int(cfg.SuspiciousLookback.Seconds()), identifier, int(cfg.RecentFailureLookback.Seconds())).Scan(
		&previousSessions,
		&seenIP,
		&seenUserAgent,
		&distinctIPs24h,
		&recentFailures,
	); err != nil {
		return nil, "", err
	}

	signals := []string{}
	severity := "medium"
	if previousSessions > 0 && remoteIP != "" && !seenIP {
		signals = append(signals, "new_remote_ip")
		severity = "high"
	}
	if previousSessions > 0 && userAgent != "" && !seenUserAgent {
		signals = append(signals, "new_user_agent")
	}
	if distinctIPs24h+1 >= int64(cfg.DistinctIPThreshold24h) {
		signals = append(signals, "multiple_ips_24h")
		severity = "high"
	}
	if recentFailures >= int64(cfg.RecentFailureThreshold) {
		signals = append(signals, "recent_failed_attempts_before_success")
	}
	return signals, severity, nil
}

func insertSecurityEvent(ctx context.Context, tx pgx.Tx, userID, eventType, severity, identifier, remoteIP, userAgent string, details map[string]any) error {
	if details == nil {
		details = map[string]any{}
	}
	encoded, err := json.Marshal(details)
	if err != nil {
		return err
	}
	var user any
	if strings.TrimSpace(userID) != "" {
		user = strings.TrimSpace(userID)
	}
	_, err = tx.Exec(ctx, `
		INSERT INTO security_events (user_id, event_type, severity, identifier, remote_ip, user_agent, details)
		VALUES ($1, $2, $3, $4, $5, NULLIF($6, ''), $7)
	`, user, eventType, severity, identifier, remoteIP, userAgent, encoded)
	return err
}

func normalizeEventDecision(params *EventDecisionParams) error {
	params.EventID = strings.TrimSpace(params.EventID)
	params.AdminUserID = strings.TrimSpace(params.AdminUserID)
	params.Status = strings.ToLower(strings.TrimSpace(params.Status))
	params.ResolutionNote = strings.TrimSpace(params.ResolutionNote)
	if err := domain.ValidateUUID("id", params.EventID); err != nil {
		return err
	}
	if err := domain.ValidateUUID("admin_user_id", params.AdminUserID); err != nil {
		return err
	}
	switch params.Status {
	case "open", "reviewing":
		return nil
	case "resolved", "ignored":
		if len(params.ResolutionNote) < 8 || len(params.ResolutionNote) > 500 {
			return fmt.Errorf("%w: resolution_note must be between 8 and 500 characters", domain.ErrValidation)
		}
		return nil
	default:
		return fmt.Errorf("%w: invalid security event status", domain.ErrValidation)
	}
}

func normalizeConfig(cfg Config) Config {
	defaults := DefaultConfig()
	if cfg.IdentifierWindow <= 0 {
		cfg.IdentifierWindow = defaults.IdentifierWindow
	}
	if cfg.IdentifierMaxFailures <= 0 {
		cfg.IdentifierMaxFailures = defaults.IdentifierMaxFailures
	}
	if cfg.IPWindow <= 0 {
		cfg.IPWindow = defaults.IPWindow
	}
	if cfg.IPMaxFailures <= 0 {
		cfg.IPMaxFailures = defaults.IPMaxFailures
	}
	if cfg.LockoutDuration <= 0 {
		cfg.LockoutDuration = defaults.LockoutDuration
	}
	if cfg.SuspiciousLookback <= 0 {
		cfg.SuspiciousLookback = defaults.SuspiciousLookback
	}
	if cfg.RecentFailureLookback <= 0 {
		cfg.RecentFailureLookback = defaults.RecentFailureLookback
	}
	if cfg.RecentFailureThreshold <= 0 {
		cfg.RecentFailureThreshold = defaults.RecentFailureThreshold
	}
	if cfg.DistinctIPThreshold24h <= 0 {
		cfg.DistinctIPThreshold24h = defaults.DistinctIPThreshold24h
	}
	return cfg
}

func normalizeLimit(limit int) int {
	if limit <= 0 || limit > 100 {
		return 50
	}
	return limit
}

func normalizeIdentifier(identifier string) string {
	return strings.ToLower(strings.TrimSpace(identifier))
}

func normalizeRemoteIP(remoteIP string) string {
	remoteIP = strings.TrimSpace(remoteIP)
	if remoteIP == "" {
		return "unknown"
	}
	return remoteIP
}

type scanner interface {
	Scan(dest ...any) error
}

func scanSecurityEvent(row scanner) (SecurityEvent, error) {
	var event SecurityEvent
	var resolvedAt sql.NullTime
	err := row.Scan(
		&event.ID,
		&event.UserID,
		&event.UserEmail,
		&event.UserFullName,
		&event.EventType,
		&event.Severity,
		&event.Status,
		&event.Identifier,
		&event.RemoteIP,
		&event.UserAgent,
		&event.Details,
		&event.ResolvedByAdminUserID,
		&event.ResolutionNote,
		&resolvedAt,
		&event.CreatedAt,
		&event.UpdatedAt,
	)
	if resolvedAt.Valid {
		event.ResolvedAt = &resolvedAt.Time
	}
	return event, err
}
