package ledger

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

	"github.com/jackc/pgx/v5"

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

type TrialBalanceReport struct {
	AsOf       time.Time              `json:"as_of"`
	Currencies []TrialBalanceCurrency `json:"currencies"`
}

type TrialBalanceCurrency struct {
	Currency         string            `json:"currency"`
	TotalDebitCents  int64             `json:"total_debit_cents"`
	TotalCreditCents int64             `json:"total_credit_cents"`
	Balanced         bool              `json:"balanced"`
	Rows             []TrialBalanceRow `json:"rows"`
}

type TrialBalanceRow struct {
	LedgerAccountID    string `json:"ledger_account_id"`
	OwnerUserID        string `json:"owner_user_id,omitempty"`
	ReferenceType      string `json:"reference_type"`
	ReferenceID        string `json:"reference_id"`
	Currency           string `json:"currency"`
	NormalBalance      string `json:"normal_balance"`
	DebitCents         int64  `json:"debit_cents"`
	CreditCents        int64  `json:"credit_cents"`
	DebitBalanceCents  int64  `json:"debit_balance_cents"`
	CreditBalanceCents int64  `json:"credit_balance_cents"`
}

type ExportFilters struct {
	From       *time.Time
	To         *time.Time
	Currency   string
	SourceType string
	Limit      int
}

type ExportRow struct {
	JournalEntryID             string          `json:"journal_entry_id"`
	JournalLineID              string          `json:"journal_line_id"`
	EventType                  string          `json:"event_type"`
	SourceType                 string          `json:"source_type"`
	SourceID                   string          `json:"source_id"`
	IdempotencyKey             string          `json:"idempotency_key,omitempty"`
	Description                string          `json:"description,omitempty"`
	LedgerAccountID            string          `json:"ledger_account_id"`
	LedgerAccountOwnerUserID   string          `json:"ledger_account_owner_user_id,omitempty"`
	LedgerAccountReferenceType string          `json:"ledger_account_reference_type"`
	LedgerAccountReferenceID   string          `json:"ledger_account_reference_id"`
	LedgerAccountNormalBalance string          `json:"ledger_account_normal_balance"`
	Direction                  string          `json:"direction"`
	AmountCents                int64           `json:"amount_cents"`
	Currency                   string          `json:"currency"`
	Metadata                   json.RawMessage `json:"metadata"`
	JournalEntryCreatedAt      time.Time       `json:"journal_entry_created_at"`
	JournalLineCreatedAt       time.Time       `json:"journal_line_created_at"`
}

type JournalReversal struct {
	ID                     string                     `json:"id"`
	OriginalJournalEntryID string                     `json:"original_journal_entry_id"`
	ReversalJournalEntryID string                     `json:"reversal_journal_entry_id,omitempty"`
	RequestedByAdminUserID string                     `json:"requested_by_admin_user_id,omitempty"`
	SourceType             string                     `json:"source_type"`
	MovementType           string                     `json:"movement_type"`
	ReversalType           string                     `json:"reversal_type"`
	Status                 string                     `json:"status"`
	Reason                 string                     `json:"reason"`
	Metadata               json.RawMessage            `json:"metadata"`
	OriginalEntry          *domain.LedgerJournalEntry `json:"original_entry,omitempty"`
	ReversalEntry          *domain.LedgerJournalEntry `json:"reversal_entry,omitempty"`
	CreatedAt              time.Time                  `json:"created_at"`
	CompletedAt            *time.Time                 `json:"completed_at,omitempty"`
}

type ReversalParams struct {
	OriginalJournalEntryID string
	AdminUserID            string
	Reason                 string
	ReversalType           string
	Metadata               any
}

type CurrencyRoundingPolicy struct {
	Currency                        string    `json:"currency"`
	MinorUnit                       int       `json:"minor_unit"`
	FXRoundingMode                  string    `json:"fx_rounding_mode"`
	CashRoundingIncrementMinorUnits int64     `json:"cash_rounding_increment_minor_units"`
	PolicySource                    string    `json:"policy_source"`
	EffectiveFrom                   time.Time `json:"effective_from"`
	UpdatedAt                       time.Time `json:"updated_at"`
}

type reversalPolicy struct {
	SourceType           string
	MovementType         string
	ProductStateRequired bool
	AllowedTypes         []string
	Notes                string
}

func (r *Repository) TrialBalance(ctx context.Context, asOf *time.Time) (TrialBalanceReport, error) {
	reportAsOf := time.Now().UTC()
	var asOfArg any
	if asOf != nil {
		reportAsOf = asOf.UTC()
		asOfArg = reportAsOf
	}

	rows, err := r.db.Query(ctx, `
		WITH eligible_lines AS (
			SELECT jl.*
			FROM ledger_journal_lines jl
			JOIN ledger_journal_entries je ON je.id = jl.journal_entry_id
			WHERE ($1::timestamptz IS NULL OR je.created_at <= $1)
		),
		totals AS (
			SELECT
				la.id,
				COALESCE(SUM(CASE WHEN jl.direction = 'debit' THEN jl.amount_cents ELSE 0 END), 0)::bigint AS debit_cents,
				COALESCE(SUM(CASE WHEN jl.direction = 'credit' THEN jl.amount_cents ELSE 0 END), 0)::bigint AS credit_cents
			FROM ledger_accounts la
			LEFT JOIN eligible_lines jl ON jl.ledger_account_id = la.id
			GROUP BY la.id
		)
		SELECT la.id::text, COALESCE(la.owner_user_id::text, ''), la.reference_type, la.reference_id,
			TRIM(la.currency)::text, la.normal_balance, t.debit_cents, t.credit_cents,
			GREATEST(t.debit_cents - t.credit_cents, 0)::bigint AS debit_balance_cents,
			GREATEST(t.credit_cents - t.debit_cents, 0)::bigint AS credit_balance_cents
		FROM ledger_accounts la
		JOIN totals t ON t.id = la.id
		ORDER BY TRIM(la.currency)::text, la.reference_type, la.reference_id
	`, asOfArg)
	if err != nil {
		return TrialBalanceReport{}, err
	}
	defer rows.Close()

	byCurrency := map[string]*TrialBalanceCurrency{}
	order := []string{}
	for rows.Next() {
		var row TrialBalanceRow
		if err := rows.Scan(
			&row.LedgerAccountID,
			&row.OwnerUserID,
			&row.ReferenceType,
			&row.ReferenceID,
			&row.Currency,
			&row.NormalBalance,
			&row.DebitCents,
			&row.CreditCents,
			&row.DebitBalanceCents,
			&row.CreditBalanceCents,
		); err != nil {
			return TrialBalanceReport{}, err
		}
		bucket := byCurrency[row.Currency]
		if bucket == nil {
			bucket = &TrialBalanceCurrency{Currency: row.Currency}
			byCurrency[row.Currency] = bucket
			order = append(order, row.Currency)
		}
		bucket.TotalDebitCents += row.DebitBalanceCents
		bucket.TotalCreditCents += row.CreditBalanceCents
		bucket.Rows = append(bucket.Rows, row)
	}
	if err := rows.Err(); err != nil {
		return TrialBalanceReport{}, err
	}

	currencies := make([]TrialBalanceCurrency, 0, len(order))
	for _, currency := range order {
		bucket := byCurrency[currency]
		bucket.Balanced = bucket.TotalDebitCents == bucket.TotalCreditCents
		currencies = append(currencies, *bucket)
	}
	return TrialBalanceReport{AsOf: reportAsOf, Currencies: currencies}, nil
}

func (r *Repository) Export(ctx context.Context, filters ExportFilters) ([]ExportRow, error) {
	filters.Currency = domain.NormalizeCurrency(filters.Currency)
	if filters.Currency != "" {
		if err := domain.ValidateCurrency(filters.Currency); err != nil {
			return nil, err
		}
	}
	filters.SourceType = strings.TrimSpace(filters.SourceType)
	limit := filters.Limit
	if limit <= 0 || limit > 10_000 {
		limit = 1_000
	}
	var fromArg, toArg any
	if filters.From != nil {
		fromArg = filters.From.UTC()
	}
	if filters.To != nil {
		toArg = filters.To.UTC()
	}

	rows, err := r.db.Query(ctx, `
		SELECT
			je.id::text,
			jl.id::text,
			je.event_type,
			je.source_type,
			je.source_id,
			COALESCE(je.idempotency_key, ''),
			COALESCE(je.description, ''),
			la.id::text,
			COALESCE(la.owner_user_id::text, ''),
			la.reference_type,
			la.reference_id,
			la.normal_balance,
			jl.direction,
			jl.amount_cents,
			TRIM(jl.currency)::text,
			je.metadata,
			je.created_at,
			jl.created_at
		FROM ledger_journal_entries je
		JOIN ledger_journal_lines jl ON jl.journal_entry_id = je.id
		JOIN ledger_accounts la ON la.id = jl.ledger_account_id
		WHERE ($1::timestamptz IS NULL OR je.created_at >= $1)
			AND ($2::timestamptz IS NULL OR je.created_at <= $2)
			AND ($3 = '' OR TRIM(jl.currency)::text = $3)
			AND ($4 = '' OR je.source_type = $4)
		ORDER BY je.created_at DESC, jl.created_at DESC, jl.id
		LIMIT $5
	`, fromArg, toArg, filters.Currency, filters.SourceType, limit)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	exportRows := []ExportRow{}
	for rows.Next() {
		var row ExportRow
		if err := rows.Scan(
			&row.JournalEntryID,
			&row.JournalLineID,
			&row.EventType,
			&row.SourceType,
			&row.SourceID,
			&row.IdempotencyKey,
			&row.Description,
			&row.LedgerAccountID,
			&row.LedgerAccountOwnerUserID,
			&row.LedgerAccountReferenceType,
			&row.LedgerAccountReferenceID,
			&row.LedgerAccountNormalBalance,
			&row.Direction,
			&row.AmountCents,
			&row.Currency,
			&row.Metadata,
			&row.JournalEntryCreatedAt,
			&row.JournalLineCreatedAt,
		); err != nil {
			return nil, err
		}
		exportRows = append(exportRows, row)
	}
	return exportRows, rows.Err()
}

func (r *Repository) ListRoundingPolicies(ctx context.Context) ([]CurrencyRoundingPolicy, error) {
	rows, err := r.db.Query(ctx, `
		SELECT currency, minor_unit, fx_rounding_mode, cash_rounding_increment_minor_units,
			policy_source, effective_from, updated_at
		FROM currency_rounding_policies
		ORDER BY currency
	`)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	policies := []CurrencyRoundingPolicy{}
	for rows.Next() {
		var policy CurrencyRoundingPolicy
		if err := rows.Scan(
			&policy.Currency,
			&policy.MinorUnit,
			&policy.FXRoundingMode,
			&policy.CashRoundingIncrementMinorUnits,
			&policy.PolicySource,
			&policy.EffectiveFrom,
			&policy.UpdatedAt,
		); err != nil {
			return nil, err
		}
		policies = append(policies, policy)
	}
	return policies, rows.Err()
}

func (r *Repository) ListReversals(ctx context.Context, limit int) ([]JournalReversal, error) {
	if limit <= 0 || limit > 500 {
		limit = 100
	}
	rows, err := r.db.Query(ctx, reversalSelect()+`
		ORDER BY r.created_at DESC
		LIMIT $1
	`, limit)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	reversals := []JournalReversal{}
	for rows.Next() {
		item, err := scanReversal(rows)
		if err != nil {
			return nil, err
		}
		reversals = append(reversals, item)
	}
	return reversals, rows.Err()
}

func (r *Repository) CreateJournalReversal(ctx context.Context, params ReversalParams) (JournalReversal, error) {
	if err := normalizeReversalParams(&params); err != nil {
		return JournalReversal{}, err
	}
	metadata, err := marshalMetadata(params.Metadata)
	if err != nil {
		return JournalReversal{}, err
	}

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

	original, err := findJournalEntryForUpdate(ctx, tx, params.OriginalJournalEntryID)
	if err != nil {
		return JournalReversal{}, err
	}
	if original.SourceType == "ledger_reversal" {
		return JournalReversal{}, fmt.Errorf("%w: reversal entries cannot be reversed directly", domain.ErrValidation)
	}
	policy, err := findReversalPolicy(ctx, tx, original.SourceType)
	if err != nil {
		return JournalReversal{}, err
	}
	if !containsString(policy.AllowedTypes, params.ReversalType) {
		return JournalReversal{}, fmt.Errorf("%w: reversal_type is not allowed for source_type", domain.ErrValidation)
	}

	reversal, inserted, err := insertPendingReversal(ctx, tx, original.ID, params, policy, metadata)
	if err != nil {
		return JournalReversal{}, err
	}
	if !inserted {
		existing, err := findReversalByOriginal(ctx, tx, original.ID)
		if err != nil {
			return JournalReversal{}, err
		}
		if err := hydrateReversalEntries(ctx, tx, &existing); err != nil {
			return JournalReversal{}, err
		}
		if err := tx.Commit(ctx); err != nil {
			return JournalReversal{}, err
		}
		return existing, nil
	}

	lines := make([]LineParams, 0, len(original.Lines))
	for _, line := range original.Lines {
		lines = append(lines, LineParams{
			LedgerAccountID: line.LedgerAccountID,
			Direction:       reverseDirection(line.Direction),
			AmountCents:     line.AmountCents,
			Currency:        line.Currency,
		})
	}
	entry, err := r.Post(ctx, tx, PostParams{
		EventType:      original.EventType + ".reversed",
		SourceType:     "ledger_reversal",
		SourceID:       reversal.ID,
		IdempotencyKey: "ledger-reversal-" + reversal.ID,
		Description:    "Reversal for journal entry " + original.ID + ": " + params.Reason,
		Metadata: map[string]any{
			"original_journal_entry_id": original.ID,
			"original_source_type":      original.SourceType,
			"original_source_id":        original.SourceID,
			"movement_type":             policy.MovementType,
			"reversal_type":             params.ReversalType,
			"reason":                    params.Reason,
			"product_state_required":    policy.ProductStateRequired,
		},
		Lines: lines,
	})
	if err != nil {
		return JournalReversal{}, err
	}

	reversal, err = completeReversal(ctx, tx, reversal.ID, entry.ID)
	if err != nil {
		return JournalReversal{}, err
	}
	originalCopy := original
	reversalCopy := entry
	reversal.OriginalEntry = &originalCopy
	reversal.ReversalEntry = &reversalCopy

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

func findJournalEntryForUpdate(ctx context.Context, tx pgx.Tx, entryID string) (domain.LedgerJournalEntry, error) {
	row := tx.QueryRow(ctx, `
		SELECT id::text, event_type, source_type, source_id, COALESCE(idempotency_key, ''),
			COALESCE(description, ''), metadata, created_at
		FROM ledger_journal_entries
		WHERE id = $1
		FOR UPDATE
	`, entryID)
	entry, err := scanEntry(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.LedgerJournalEntry{}, domain.ErrNotFound
	}
	if err != nil {
		return domain.LedgerJournalEntry{}, err
	}
	entry.Lines, err = (&Repository{}).listLines(ctx, tx, entry.ID)
	return entry, err
}

func findReversalPolicy(ctx context.Context, tx pgx.Tx, sourceType string) (reversalPolicy, error) {
	var policy reversalPolicy
	err := tx.QueryRow(ctx, `
		SELECT source_type, movement_type, product_state_required, allowed_reversal_types, notes
		FROM ledger_reversal_policies
		WHERE source_type = $1
	`, sourceType).Scan(
		&policy.SourceType,
		&policy.MovementType,
		&policy.ProductStateRequired,
		&policy.AllowedTypes,
		&policy.Notes,
	)
	if errors.Is(err, pgx.ErrNoRows) {
		return reversalPolicy{}, fmt.Errorf("%w: no reversal policy configured for source_type", domain.ErrValidation)
	}
	return policy, err
}

func insertPendingReversal(ctx context.Context, tx pgx.Tx, originalID string, params ReversalParams, policy reversalPolicy, metadata []byte) (JournalReversal, bool, error) {
	row := tx.QueryRow(ctx, `
		INSERT INTO ledger_journal_reversals (
			original_journal_entry_id, requested_by_admin_user_id, source_type, movement_type,
			reversal_type, reason, metadata
		)
		VALUES ($1, NULLIF($2, '')::uuid, $3, $4, $5, $6, $7)
		ON CONFLICT (original_journal_entry_id) DO NOTHING
		RETURNING id::text, original_journal_entry_id::text, COALESCE(reversal_journal_entry_id::text, ''),
			COALESCE(requested_by_admin_user_id::text, ''), source_type, movement_type, reversal_type,
			status, reason, metadata, created_at, completed_at
	`, originalID, params.AdminUserID, policy.SourceType, policy.MovementType, params.ReversalType, params.Reason, metadata)
	reversal, err := scanReversal(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return JournalReversal{}, false, nil
	}
	return reversal, true, err
}

func completeReversal(ctx context.Context, tx pgx.Tx, reversalID, entryID string) (JournalReversal, error) {
	row := tx.QueryRow(ctx, reversalSelect()+`
		WHERE r.id = $1
	`, reversalID)
	before, err := scanReversal(row)
	if err != nil {
		return JournalReversal{}, err
	}
	row = tx.QueryRow(ctx, `
		UPDATE ledger_journal_reversals
		SET reversal_journal_entry_id = $2,
			status = 'completed',
			completed_at = now()
		WHERE id = $1
		RETURNING id::text, original_journal_entry_id::text, COALESCE(reversal_journal_entry_id::text, ''),
			COALESCE(requested_by_admin_user_id::text, ''), source_type, movement_type, reversal_type,
			status, reason, metadata, created_at, completed_at
	`, before.ID, entryID)
	return scanReversal(row)
}

func findReversalByOriginal(ctx context.Context, tx pgx.Tx, originalID string) (JournalReversal, error) {
	row := tx.QueryRow(ctx, reversalSelect()+`
		WHERE r.original_journal_entry_id = $1
	`, originalID)
	item, err := scanReversal(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return JournalReversal{}, domain.ErrNotFound
	}
	return item, err
}

func hydrateReversalEntries(ctx context.Context, tx pgx.Tx, item *JournalReversal) error {
	original, err := findJournalEntryByID(ctx, tx, item.OriginalJournalEntryID)
	if err != nil {
		return err
	}
	item.OriginalEntry = &original
	if item.ReversalJournalEntryID == "" {
		return nil
	}
	reversal, err := findJournalEntryByID(ctx, tx, item.ReversalJournalEntryID)
	if err != nil {
		return err
	}
	item.ReversalEntry = &reversal
	return nil
}

func findJournalEntryByID(ctx context.Context, q txer, entryID string) (domain.LedgerJournalEntry, error) {
	row := q.QueryRow(ctx, `
		SELECT id::text, event_type, source_type, source_id, COALESCE(idempotency_key, ''),
			COALESCE(description, ''), metadata, created_at
		FROM ledger_journal_entries
		WHERE id = $1
	`, entryID)
	entry, err := scanEntry(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.LedgerJournalEntry{}, domain.ErrNotFound
	}
	if err != nil {
		return domain.LedgerJournalEntry{}, err
	}
	entry.Lines, err = (&Repository{}).listLines(ctx, q, entry.ID)
	return entry, err
}

func reversalSelect() string {
	return `
		SELECT r.id::text, r.original_journal_entry_id::text, COALESCE(r.reversal_journal_entry_id::text, ''),
			COALESCE(r.requested_by_admin_user_id::text, ''), r.source_type, r.movement_type,
			r.reversal_type, r.status, r.reason, r.metadata, r.created_at, r.completed_at
		FROM ledger_journal_reversals r
	`
}

func scanReversal(row scanner) (JournalReversal, error) {
	var item JournalReversal
	var completedAt sql.NullTime
	err := row.Scan(
		&item.ID,
		&item.OriginalJournalEntryID,
		&item.ReversalJournalEntryID,
		&item.RequestedByAdminUserID,
		&item.SourceType,
		&item.MovementType,
		&item.ReversalType,
		&item.Status,
		&item.Reason,
		&item.Metadata,
		&item.CreatedAt,
		&completedAt,
	)
	if completedAt.Valid {
		item.CompletedAt = &completedAt.Time
	}
	return item, err
}

func normalizeReversalParams(params *ReversalParams) error {
	if err := domain.ValidateUUID("journal_entry_id", params.OriginalJournalEntryID); err != nil {
		return err
	}
	if params.AdminUserID != "" {
		if err := domain.ValidateUUID("admin_user_id", params.AdminUserID); err != nil {
			return err
		}
	}
	params.Reason = strings.TrimSpace(params.Reason)
	if len(params.Reason) < 8 || len(params.Reason) > 500 {
		return fmt.Errorf("%w: reason must be between 8 and 500 characters", domain.ErrValidation)
	}
	params.ReversalType = strings.ToLower(strings.TrimSpace(params.ReversalType))
	if params.ReversalType == "" {
		params.ReversalType = "correction"
	}
	switch params.ReversalType {
	case "correction", "provider_error", "customer_refund", "fraud", "operational":
		return nil
	default:
		return fmt.Errorf("%w: invalid reversal_type", domain.ErrValidation)
	}
}

func reverseDirection(direction string) string {
	if direction == "debit" {
		return "credit"
	}
	return "debit"
}

func containsString(values []string, needle string) bool {
	for _, value := range values {
		if value == needle {
			return true
		}
	}
	return false
}
