package users

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

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

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

type Repository struct {
	db *pgxpool.Pool
}

type CreateParams struct {
	Email        string
	FullName     string
	PasswordHash string
	Role         string
}

type MFASettings struct {
	UserID               string     `json:"user_id"`
	TOTPEnabled          bool       `json:"totp_enabled"`
	EnabledAt            *time.Time `json:"enabled_at,omitempty"`
	TOTPSecretCiphertext string     `json:"-"`
}

type PasskeyCredential struct {
	ID             string     `json:"id"`
	UserID         string     `json:"user_id"`
	CredentialID   string     `json:"credential_id"`
	Nickname       string     `json:"nickname,omitempty"`
	Transports     []string   `json:"transports,omitempty"`
	AAGUID         string     `json:"aaguid,omitempty"`
	SignCount      int64      `json:"sign_count"`
	BackupEligible bool       `json:"backup_eligible"`
	BackupState    bool       `json:"backup_state"`
	Status         string     `json:"status"`
	LastUsedAt     *time.Time `json:"last_used_at,omitempty"`
	RevokedAt      *time.Time `json:"revoked_at,omitempty"`
	CreatedAt      time.Time  `json:"created_at"`
	UpdatedAt      time.Time  `json:"updated_at"`
}

type CreatePasskeyParams struct {
	UserID         string
	CredentialID   string
	PublicKey      string
	Nickname       string
	Transports     []string
	AAGUID         string
	SignCount      int64
	BackupEligible bool
	BackupState    bool
}

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

func (r *Repository) Create(ctx context.Context, params CreateParams) (domain.User, error) {
	role := params.Role
	if role == "" {
		role = "customer"
	}

	row := r.db.QueryRow(ctx, `
		INSERT INTO users (email, full_name, password_hash, role)
		VALUES ($1, $2, $3, $4)
		RETURNING id::text, email, full_name, password_hash, role, created_at
	`, strings.ToLower(params.Email), params.FullName, params.PasswordHash, role)

	user, err := scanUser(row)
	if err != nil {
		if isUniqueViolation(err) {
			return domain.User{}, domain.ErrConflict
		}
		return domain.User{}, err
	}

	return user, nil
}

func (r *Repository) FindByEmail(ctx context.Context, email string) (domain.User, error) {
	row := r.db.QueryRow(ctx, `
		SELECT id::text, email, full_name, password_hash, role, created_at
		FROM users
		WHERE lower(email) = lower($1)
	`, strings.TrimSpace(email))

	user, err := scanUser(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.User{}, domain.ErrNotFound
	}
	return user, err
}

func (r *Repository) FindByID(ctx context.Context, id string) (domain.User, error) {
	row := r.db.QueryRow(ctx, `
		SELECT id::text, email, full_name, password_hash, role, created_at
		FROM users
		WHERE id = $1
	`, id)

	user, err := scanUser(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.User{}, domain.ErrNotFound
	}
	return user, err
}

func (r *Repository) EffectiveScopes(ctx context.Context, user domain.User) ([]string, error) {
	base := security.ScopesForRole(user.Role)
	if user.ID == "" {
		return base, nil
	}

	rows, err := r.db.Query(ctx, `
		SELECT scope
		FROM admin_scope_assignments
		WHERE user_id = $1
			AND status = 'active'
			AND (expires_at IS NULL OR expires_at > now())
		UNION
		SELECT unnest(scopes)
		FROM admin_jit_elevation_requests
		WHERE admin_user_id = $1
			AND status = 'approved'
			AND starts_at <= now()
			AND expires_at > now()
	`, user.ID)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	extra := []string{}
	for rows.Next() {
		var scope string
		if err := rows.Scan(&scope); err != nil {
			return nil, err
		}
		extra = append(extra, scope)
	}
	if err := rows.Err(); err != nil {
		return nil, err
	}
	scopes, err := security.NormalizeScopes(security.MergeScopes(base, extra))
	if err != nil {
		return nil, err
	}
	return scopes, nil
}

func (r *Repository) GetMFASettings(ctx context.Context, userID string) (MFASettings, error) {
	var settings MFASettings
	var enabledAt sql.NullTime
	err := r.db.QueryRow(ctx, `
		SELECT user_id::text, totp_enabled, enabled_at, COALESCE(totp_secret_ciphertext, '')
		FROM user_mfa_settings
		WHERE user_id = $1
	`, userID).Scan(
		&settings.UserID,
		&settings.TOTPEnabled,
		&enabledAt,
		&settings.TOTPSecretCiphertext,
	)
	if errors.Is(err, pgx.ErrNoRows) {
		return MFASettings{UserID: userID}, nil
	}
	if err != nil {
		return MFASettings{}, err
	}
	if enabledAt.Valid {
		settings.EnabledAt = &enabledAt.Time
	}
	return settings, nil
}

func (r *Repository) ActiveRecoveryCodeCount(ctx context.Context, userID string) (int64, error) {
	var count int64
	err := r.db.QueryRow(ctx, `
		SELECT COUNT(*)::bigint
		FROM user_mfa_recovery_codes
		WHERE user_id = $1
			AND used_at IS NULL
			AND replaced_at IS NULL
	`, userID).Scan(&count)
	return count, err
}

func (r *Repository) ReplaceRecoveryCodes(ctx context.Context, userID string, codeHashes []string) error {
	if len(codeHashes) == 0 {
		return fmt.Errorf("%w: at least one recovery code is required", domain.ErrValidation)
	}
	tx, err := r.db.BeginTx(ctx, pgx.TxOptions{})
	if err != nil {
		return err
	}
	defer tx.Rollback(ctx)

	if _, err := tx.Exec(ctx, `
		UPDATE user_mfa_recovery_codes
		SET replaced_at = now()
		WHERE user_id = $1
			AND used_at IS NULL
			AND replaced_at IS NULL
	`, userID); err != nil {
		return err
	}
	for _, hash := range codeHashes {
		if strings.TrimSpace(hash) == "" {
			return fmt.Errorf("%w: recovery code hash cannot be empty", domain.ErrValidation)
		}
		if _, err := tx.Exec(ctx, `
			INSERT INTO user_mfa_recovery_codes (user_id, code_hash)
			VALUES ($1, $2)
		`, userID, hash); err != nil {
			return err
		}
	}
	return tx.Commit(ctx)
}

func (r *Repository) UseRecoveryCode(ctx context.Context, userID, codeHash, remoteIP, userAgent string) error {
	remoteIP = strings.TrimSpace(remoteIP)
	if remoteIP == "" {
		remoteIP = "unknown"
	}
	tag, err := r.db.Exec(ctx, `
		UPDATE user_mfa_recovery_codes
		SET used_at = now(),
			used_remote_ip = NULLIF($3, ''),
			used_user_agent = NULLIF($4, '')
		WHERE user_id = $1
			AND code_hash = $2
			AND used_at IS NULL
			AND replaced_at IS NULL
	`, userID, codeHash, remoteIP, userAgent)
	if err != nil {
		return err
	}
	if tag.RowsAffected() == 0 {
		return domain.ErrInvalidCredentials
	}
	return nil
}

func (r *Repository) CreateWebAuthnChallenge(ctx context.Context, userID, challengeHash, challengeType string, expiresAt time.Time) error {
	challengeType = strings.ToLower(strings.TrimSpace(challengeType))
	if challengeType != "registration" && challengeType != "authentication" {
		return fmt.Errorf("%w: challenge_type must be registration or authentication", domain.ErrValidation)
	}
	_, err := r.db.Exec(ctx, `
		INSERT INTO webauthn_challenges (user_id, challenge_hash, challenge_type, expires_at)
		VALUES ($1, $2, $3, $4)
	`, userID, challengeHash, challengeType, expiresAt)
	return err
}

func (r *Repository) ConsumeWebAuthnChallenge(ctx context.Context, userID, challengeHash, challengeType string) error {
	tag, err := r.db.Exec(ctx, `
		UPDATE webauthn_challenges
		SET consumed_at = now()
		WHERE user_id = $1
			AND challenge_hash = $2
			AND challenge_type = $3
			AND consumed_at IS NULL
			AND expires_at > now()
	`, userID, challengeHash, strings.ToLower(strings.TrimSpace(challengeType)))
	if err != nil {
		return err
	}
	if tag.RowsAffected() == 0 {
		return fmt.Errorf("%w: passkey challenge is invalid or expired", domain.ErrValidation)
	}
	return nil
}

func (r *Repository) CreatePasskey(ctx context.Context, params CreatePasskeyParams) (PasskeyCredential, error) {
	if err := normalizePasskeyParams(&params); err != nil {
		return PasskeyCredential{}, err
	}
	row := r.db.QueryRow(ctx, `
		INSERT INTO webauthn_credentials (
			user_id, credential_id, public_key, nickname, transports, aaguid,
			sign_count, backup_eligible, backup_state
		)
		VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9)
		RETURNING id::text, user_id::text, credential_id, nickname, transports, aaguid,
			sign_count, backup_eligible, backup_state, status, last_used_at, revoked_at, created_at, updated_at
	`, params.UserID, params.CredentialID, params.PublicKey, params.Nickname, params.Transports,
		params.AAGUID, params.SignCount, params.BackupEligible, params.BackupState)
	credential, err := scanPasskey(row)
	if err != nil {
		if isUniqueViolation(err) {
			return PasskeyCredential{}, domain.ErrConflict
		}
		return PasskeyCredential{}, err
	}
	return credential, nil
}

func (r *Repository) ListPasskeys(ctx context.Context, userID string) ([]PasskeyCredential, error) {
	rows, err := r.db.Query(ctx, `
		SELECT id::text, user_id::text, credential_id, nickname, transports, aaguid,
			sign_count, backup_eligible, backup_state, status, last_used_at, revoked_at, created_at, updated_at
		FROM webauthn_credentials
		WHERE user_id = $1
		ORDER BY created_at DESC
	`, userID)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	credentials := []PasskeyCredential{}
	for rows.Next() {
		credential, err := scanPasskey(rows)
		if err != nil {
			return nil, err
		}
		credentials = append(credentials, credential)
	}
	return credentials, rows.Err()
}

func (r *Repository) RevokePasskey(ctx context.Context, userID, credentialID string) error {
	tag, err := r.db.Exec(ctx, `
		UPDATE webauthn_credentials
		SET status = 'revoked',
			revoked_at = now()
		WHERE id = $1
			AND user_id = $2
			AND status = 'active'
	`, credentialID, userID)
	if err != nil {
		return err
	}
	if tag.RowsAffected() == 0 {
		return domain.ErrNotFound
	}
	return nil
}

func (r *Repository) UpsertTOTPSetup(ctx context.Context, userID, ciphertext string) (MFASettings, error) {
	row := r.db.QueryRow(ctx, `
		INSERT INTO user_mfa_settings (user_id, totp_secret_ciphertext, totp_enabled, enabled_at)
		VALUES ($1, $2, false, NULL)
		ON CONFLICT (user_id) DO UPDATE
		SET totp_secret_ciphertext = EXCLUDED.totp_secret_ciphertext,
			totp_enabled = false,
			enabled_at = NULL,
			updated_at = now()
		RETURNING user_id::text, totp_enabled, enabled_at, COALESCE(totp_secret_ciphertext, '')
	`, userID, ciphertext)

	var settings MFASettings
	var enabledAt sql.NullTime
	if err := row.Scan(&settings.UserID, &settings.TOTPEnabled, &enabledAt, &settings.TOTPSecretCiphertext); err != nil {
		return MFASettings{}, err
	}
	if enabledAt.Valid {
		settings.EnabledAt = &enabledAt.Time
	}
	return settings, nil
}

func (r *Repository) EnableTOTP(ctx context.Context, userID string) (MFASettings, error) {
	row := r.db.QueryRow(ctx, `
		UPDATE user_mfa_settings
		SET totp_enabled = true,
			enabled_at = now(),
			updated_at = now()
		WHERE user_id = $1 AND totp_secret_ciphertext IS NOT NULL
		RETURNING user_id::text, totp_enabled, enabled_at, COALESCE(totp_secret_ciphertext, '')
	`, userID)

	var settings MFASettings
	var enabledAt sql.NullTime
	err := row.Scan(&settings.UserID, &settings.TOTPEnabled, &enabledAt, &settings.TOTPSecretCiphertext)
	if errors.Is(err, pgx.ErrNoRows) {
		return MFASettings{}, domain.ErrNotFound
	}
	if err != nil {
		return MFASettings{}, err
	}
	if enabledAt.Valid {
		settings.EnabledAt = &enabledAt.Time
	}
	return settings, nil
}

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

func scanUser(row scanner) (domain.User, error) {
	var user domain.User
	err := row.Scan(&user.ID, &user.Email, &user.FullName, &user.PasswordHash, &user.Role, &user.CreatedAt)
	return user, err
}

func scanPasskey(row scanner) (PasskeyCredential, error) {
	var credential PasskeyCredential
	var lastUsedAt sql.NullTime
	var revokedAt sql.NullTime
	err := row.Scan(
		&credential.ID,
		&credential.UserID,
		&credential.CredentialID,
		&credential.Nickname,
		&credential.Transports,
		&credential.AAGUID,
		&credential.SignCount,
		&credential.BackupEligible,
		&credential.BackupState,
		&credential.Status,
		&lastUsedAt,
		&revokedAt,
		&credential.CreatedAt,
		&credential.UpdatedAt,
	)
	if lastUsedAt.Valid {
		credential.LastUsedAt = &lastUsedAt.Time
	}
	if revokedAt.Valid {
		credential.RevokedAt = &revokedAt.Time
	}
	return credential, err
}

func normalizePasskeyParams(params *CreatePasskeyParams) error {
	params.UserID = strings.TrimSpace(params.UserID)
	params.CredentialID = strings.TrimSpace(params.CredentialID)
	params.PublicKey = strings.TrimSpace(params.PublicKey)
	params.Nickname = strings.TrimSpace(params.Nickname)
	params.AAGUID = strings.TrimSpace(params.AAGUID)
	if err := domain.ValidateUUID("user_id", params.UserID); err != nil {
		return err
	}
	if params.CredentialID == "" {
		return fmt.Errorf("%w: credential_id is required", domain.ErrValidation)
	}
	if params.PublicKey == "" {
		return fmt.Errorf("%w: public_key is required", domain.ErrValidation)
	}
	if params.SignCount < 0 {
		return fmt.Errorf("%w: sign_count cannot be negative", domain.ErrValidation)
	}
	cleanTransports := []string{}
	seen := map[string]bool{}
	for _, raw := range params.Transports {
		transport := strings.ToLower(strings.TrimSpace(raw))
		if transport == "" || seen[transport] {
			continue
		}
		switch transport {
		case "usb", "nfc", "ble", "internal", "hybrid":
			seen[transport] = true
			cleanTransports = append(cleanTransports, transport)
		default:
			return fmt.Errorf("%w: unsupported passkey transport %q", domain.ErrValidation, transport)
		}
	}
	params.Transports = cleanTransports
	return nil
}

func isUniqueViolation(err error) bool {
	var pgErr *pgconn.PgError
	return errors.As(err, &pgErr) && pgErr.Code == "23505"
}
