package crypto

import (
	"context"
	"errors"

	"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"
)

type Repository struct {
	db *pgxpool.Pool
}

type CreateWalletParams struct {
	UserID           string
	Name             string
	CustodyProvider  string
	ExternalWalletID string
}

type AddAssetParams struct {
	UserID         string
	CryptoWalletID string
	CryptoAssetID  string
}

type CreateAddressParams struct {
	UserID         string
	CryptoWalletID string
	CryptoAssetID  string
	Address        string
	AddressTag     string
}

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

func (r *Repository) ListAssets(ctx context.Context) ([]domain.CryptoAsset, error) {
	rows, err := r.db.Query(ctx, `
		SELECT id::text, symbol, name, network, asset_type, COALESCE(contract_address, ''), decimals::int, enabled, created_at
		FROM crypto_assets
		WHERE enabled = true
		ORDER BY symbol, network
	`)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	assets := []domain.CryptoAsset{}
	for rows.Next() {
		asset, err := scanAsset(rows)
		if err != nil {
			return nil, err
		}
		assets = append(assets, asset)
	}
	return assets, rows.Err()
}

func (r *Repository) FindAsset(ctx context.Context, assetID string) (domain.CryptoAsset, error) {
	row := r.db.QueryRow(ctx, `
		SELECT id::text, symbol, name, network, asset_type, COALESCE(contract_address, ''), decimals::int, enabled, created_at
		FROM crypto_assets
		WHERE id = $1 AND enabled = true
	`, assetID)

	asset, err := scanAsset(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.CryptoAsset{}, domain.ErrNotFound
	}
	return asset, err
}

func (r *Repository) CreateWallet(ctx context.Context, params CreateWalletParams) (domain.CryptoWallet, error) {
	row := r.db.QueryRow(ctx, `
		INSERT INTO crypto_wallets (user_id, name, custody_provider, external_wallet_id)
		VALUES ($1, $2, $3, $4)
		RETURNING id::text, user_id::text, name, custody_provider, external_wallet_id, status, created_at, updated_at
	`, params.UserID, params.Name, params.CustodyProvider, params.ExternalWalletID)

	wallet, err := scanWallet(row)
	if err != nil {
		if isUniqueViolation(err) {
			return domain.CryptoWallet{}, domain.ErrConflict
		}
		return domain.CryptoWallet{}, err
	}
	return wallet, nil
}

func (r *Repository) ListWallets(ctx context.Context, userID string) ([]domain.CryptoWallet, error) {
	rows, err := r.db.Query(ctx, `
		SELECT id::text, user_id::text, name, custody_provider, external_wallet_id, status, created_at, updated_at
		FROM crypto_wallets
		WHERE user_id = $1
		ORDER BY created_at DESC
	`, userID)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	wallets := []domain.CryptoWallet{}
	for rows.Next() {
		wallet, err := scanWallet(rows)
		if err != nil {
			return nil, err
		}
		wallets = append(wallets, wallet)
	}
	if err := rows.Err(); err != nil {
		return nil, err
	}

	for i := range wallets {
		balances, err := r.ListBalances(ctx, userID, wallets[i].ID)
		if err != nil {
			return nil, err
		}
		addresses, err := r.ListAddresses(ctx, userID, wallets[i].ID)
		if err != nil {
			return nil, err
		}
		wallets[i].Balances = balances
		wallets[i].Addresses = addresses
	}

	return wallets, nil
}

func (r *Repository) FindOwned(ctx context.Context, userID, walletID string) (domain.CryptoWallet, error) {
	row := r.db.QueryRow(ctx, `
		SELECT id::text, user_id::text, name, custody_provider, external_wallet_id, status, created_at, updated_at
		FROM crypto_wallets
		WHERE id = $1 AND user_id = $2
	`, walletID, userID)

	wallet, err := scanWallet(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.CryptoWallet{}, domain.ErrNotFound
	}
	if err != nil {
		return domain.CryptoWallet{}, err
	}

	wallet.Balances, err = r.ListBalances(ctx, userID, walletID)
	if err != nil {
		return domain.CryptoWallet{}, err
	}
	wallet.Addresses, err = r.ListAddresses(ctx, userID, walletID)
	if err != nil {
		return domain.CryptoWallet{}, err
	}

	return wallet, nil
}

func (r *Repository) AddAsset(ctx context.Context, params AddAssetParams) (domain.CryptoWalletBalance, error) {
	row := r.db.QueryRow(ctx, `
		INSERT INTO crypto_wallet_balances (crypto_wallet_id, crypto_asset_id)
		SELECT cw.id, ca.id
		FROM crypto_wallets cw
		JOIN crypto_assets ca ON ca.id = $3 AND ca.enabled = true
		WHERE cw.id = $2 AND cw.user_id = $1 AND cw.status = 'active'
		ON CONFLICT (crypto_wallet_id, crypto_asset_id) DO UPDATE
		SET updated_at = crypto_wallet_balances.updated_at
		RETURNING id::text, crypto_wallet_id::text, crypto_asset_id::text,
			(SELECT symbol FROM crypto_assets WHERE id = crypto_asset_id),
			(SELECT network FROM crypto_assets WHERE id = crypto_asset_id),
			available_amount_base_units::text, reserved_amount_base_units::text, created_at, updated_at
	`, params.UserID, params.CryptoWalletID, params.CryptoAssetID)

	balance, err := scanBalance(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.CryptoWalletBalance{}, domain.ErrNotFound
	}
	return balance, err
}

func (r *Repository) ListBalances(ctx context.Context, userID, walletID string) ([]domain.CryptoWalletBalance, error) {
	rows, err := r.db.Query(ctx, `
		SELECT cb.id::text, cb.crypto_wallet_id::text, cb.crypto_asset_id::text,
			ca.symbol, ca.network, cb.available_amount_base_units::text, cb.reserved_amount_base_units::text,
			cb.created_at, cb.updated_at
		FROM crypto_wallet_balances cb
		JOIN crypto_wallets cw ON cw.id = cb.crypto_wallet_id
		JOIN crypto_assets ca ON ca.id = cb.crypto_asset_id
		WHERE cb.crypto_wallet_id = $1 AND cw.user_id = $2
		ORDER BY ca.symbol, ca.network
	`, walletID, userID)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	balances := []domain.CryptoWalletBalance{}
	for rows.Next() {
		balance, err := scanBalance(rows)
		if err != nil {
			return nil, err
		}
		balances = append(balances, balance)
	}
	return balances, rows.Err()
}

func (r *Repository) CreateAddress(ctx context.Context, params CreateAddressParams) (domain.CryptoAddress, error) {
	row := r.db.QueryRow(ctx, `
		INSERT INTO crypto_addresses (crypto_wallet_id, crypto_asset_id, network, address, address_tag)
		SELECT cw.id, ca.id, ca.network, $4, NULLIF($5, '')
		FROM crypto_wallets cw
		JOIN crypto_assets ca ON ca.id = $3 AND ca.enabled = true
		WHERE cw.id = $2 AND cw.user_id = $1 AND cw.status = 'active'
		ON CONFLICT (crypto_wallet_id, crypto_asset_id, network) DO UPDATE
		SET address = crypto_addresses.address
		RETURNING id::text, crypto_wallet_id::text, crypto_asset_id::text,
			(SELECT symbol FROM crypto_assets WHERE id = crypto_asset_id),
			network, address, COALESCE(address_tag, ''), status, created_at
	`, params.UserID, params.CryptoWalletID, params.CryptoAssetID, params.Address, params.AddressTag)

	address, err := scanAddress(row)
	if errors.Is(err, pgx.ErrNoRows) {
		return domain.CryptoAddress{}, domain.ErrNotFound
	}
	if err != nil {
		if isUniqueViolation(err) {
			return domain.CryptoAddress{}, domain.ErrConflict
		}
		return domain.CryptoAddress{}, err
	}
	return address, nil
}

func (r *Repository) ListAddresses(ctx context.Context, userID, walletID string) ([]domain.CryptoAddress, error) {
	rows, err := r.db.Query(ctx, `
		SELECT caa.id::text, caa.crypto_wallet_id::text, caa.crypto_asset_id::text,
			ca.symbol, caa.network, caa.address, COALESCE(caa.address_tag, ''), caa.status, caa.created_at
		FROM crypto_addresses caa
		JOIN crypto_wallets cw ON cw.id = caa.crypto_wallet_id
		JOIN crypto_assets ca ON ca.id = caa.crypto_asset_id
		WHERE caa.crypto_wallet_id = $1 AND cw.user_id = $2
		ORDER BY ca.symbol, caa.network
	`, walletID, userID)
	if err != nil {
		return nil, err
	}
	defer rows.Close()

	addresses := []domain.CryptoAddress{}
	for rows.Next() {
		address, err := scanAddress(rows)
		if err != nil {
			return nil, err
		}
		addresses = append(addresses, address)
	}
	return addresses, rows.Err()
}

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

func scanAsset(row scanner) (domain.CryptoAsset, error) {
	var asset domain.CryptoAsset
	err := row.Scan(&asset.ID, &asset.Symbol, &asset.Name, &asset.Network, &asset.AssetType, &asset.ContractAddress, &asset.Decimals, &asset.Enabled, &asset.CreatedAt)
	return asset, err
}

func scanWallet(row scanner) (domain.CryptoWallet, error) {
	var wallet domain.CryptoWallet
	err := row.Scan(&wallet.ID, &wallet.UserID, &wallet.Name, &wallet.CustodyProvider, &wallet.ExternalWalletID, &wallet.Status, &wallet.CreatedAt, &wallet.UpdatedAt)
	return wallet, err
}

func scanBalance(row scanner) (domain.CryptoWalletBalance, error) {
	var balance domain.CryptoWalletBalance
	err := row.Scan(
		&balance.ID,
		&balance.CryptoWalletID,
		&balance.CryptoAssetID,
		&balance.Symbol,
		&balance.Network,
		&balance.AvailableAmountBaseUnits,
		&balance.ReservedAmountBaseUnits,
		&balance.CreatedAt,
		&balance.UpdatedAt,
	)
	return balance, err
}

func scanAddress(row scanner) (domain.CryptoAddress, error) {
	var address domain.CryptoAddress
	err := row.Scan(
		&address.ID,
		&address.CryptoWalletID,
		&address.CryptoAssetID,
		&address.Symbol,
		&address.Network,
		&address.Address,
		&address.AddressTag,
		&address.Status,
		&address.CreatedAt,
	)
	return address, err
}

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