from __future__ import annotations

from decimal import Decimal

import pytest

from jarvis_finance.fx.rates import get_fx_rate_to_chf, recheck_transaction_fx_status, update_fx_rates, upsert_fx_rate
from jarvis_finance.fx.overrides import set_manual_fx_override
from jarvis_finance.fx.providers import MockFxProvider
from jarvis_finance.market_data.prices import MockEquityPriceProvider, refresh_market_prices, store_market_price
from jarvis_finance.market_data.mappings import confirm_instrument_price_mapping, ensure_instrument_price_mapping_quality
from jarvis_finance.ledger.positions import calculate_positions
from jarvis_finance.storage.database import connect_memory
from jarvis_finance.storage.migrations import apply_migrations


def setup_conn():
    conn = connect_memory()
    apply_migrations(conn)
    conn.execute("INSERT INTO platforms(platform_id,name,platform_type,default_currency,created_at) VALUES('p1','Broker','broker','CHF','now')")
    conn.execute("INSERT INTO accounts(account_id,platform_id,account_name,account_type,currency,created_at) VALUES('a1','p1','Account','brokerage','CHF','now')")
    conn.execute("INSERT INTO instruments(instrument_id,asset_class,name,ticker,isin,exchange,currency,created_at) VALUES('chf1','ETF','CHF ETF','CHETF','CH0000000001','SWX','CHF','now')")
    conn.execute("INSERT INTO instruments(instrument_id,asset_class,name,ticker,isin,exchange,currency,created_at) VALUES('usd1','ETF','USD ETF','USETF','US0000000001','NYSE','USD','now')")
    return conn


def test_chf_fx_is_not_needed_and_missing_alert_can_be_resolved() -> None:
    conn = setup_conn()
    conn.execute("""
        INSERT INTO transactions(transaction_id, transaction_type, account_id, instrument_id, trade_date, quantity, currency_original, fx_status, source_type, is_confirmed, quality_status, created_at)
        VALUES('tx-chf','initial_position_snapshot','a1','chf1','2025-12-31','1','CHF','missing','test',1,'warning','now')
    """)
    conn.execute("INSERT INTO alerts(alert_id,priority,category,entity_type,entity_id,rule_id,message,status,created_at) VALUES('al-chf','kritisch','ledger','instrument','chf1','missing_fx','missing','active','now')")

    assert get_fx_rate_to_chf(conn, base_currency='CHF') == Decimal('1')
    row = conn.execute("SELECT fx_rate_to_chf, fx_source, fx_status FROM transactions WHERE transaction_id='tx-chf'").fetchone()
    assert row['fx_rate_to_chf'] == '1'
    assert row['fx_source'] == 'not_needed'
    assert row['fx_status'] == 'not_needed'


def test_recheck_transaction_fx_status_corrects_chf_false_positive_with_audit() -> None:
    conn = setup_conn()
    conn.execute("""
        INSERT INTO transactions(transaction_id, transaction_type, account_id, instrument_id, trade_date, quantity, currency_original, fx_status, source_type, is_confirmed, quality_status, created_at)
        VALUES('tx-chf','initial_position_snapshot','a1','chf1','2025-12-31','1','CHF','missing','true_wealth',1,'warning','now')
    """)
    conn.execute("INSERT INTO alerts(alert_id,priority,category,entity_type,entity_id,rule_id,message,status,created_at) VALUES('al-chf','kritisch','ledger','instrument','chf1','missing_fx','missing','active','now')")

    result = recheck_transaction_fx_status(conn, source_type='true_wealth')

    assert result.transaction_count == 1
    assert result.chf_positions == 1
    assert result.false_positive_missing_fx_alerts == 1
    assert result.corrected_transactions == 1
    assert result.resolved_alerts == 1
    assert result.audit_events == 1
    assert conn.execute("SELECT status FROM alerts WHERE alert_id='al-chf'").fetchone()['status'] == 'resolved'
    assert conn.execute("SELECT COUNT(*) AS c FROM audit_log WHERE action='set_chf_fx_not_needed'").fetchone()['c'] == 1


def test_update_fx_rates_stores_usd_eur_and_resolves_missing_fx() -> None:
    conn = setup_conn()
    conn.execute("""
        INSERT INTO transactions(transaction_id, transaction_type, account_id, instrument_id, trade_date, quantity, currency_original, fx_status, source_type, is_confirmed, quality_status, created_at)
        VALUES('tx-usd','initial_position_snapshot','a1','usd1','2025-12-31','1','USD','missing','test',1,'warning','now')
    """)
    conn.execute("INSERT INTO alerts(alert_id,priority,category,entity_type,entity_id,rule_id,message,status,created_at) VALUES('al-usd','kritisch','ledger','instrument','usd1','missing_fx','missing','active','now')")
    provider = MockFxProvider({('USD', 'CHF', '2025-12-31'): Decimal('0.9000'), ('EUR', 'CHF', '2025-12-31'): Decimal('0.9500')})

    result = update_fx_rates(conn, provider=provider, currencies=['USD', 'EUR'], rate_date='2025-12-31', resolve_fixed=True)

    assert result.updated_count == 2
    assert get_fx_rate_to_chf(conn, base_currency='USD', rate_date='2025-12-31') == Decimal('0.9000')
    assert conn.execute("SELECT status FROM alerts WHERE alert_id='al-usd'").fetchone()["status"] == "resolved"


def test_manual_fx_override_requires_note_and_audit() -> None:
    conn = setup_conn()
    with pytest.raises(ValueError, match="note"):
        set_manual_fx_override(conn, base_currency='USD', quote_currency='CHF', rate_date='2025-12-31', rate=Decimal('0.9'), note='')

    rate_id = set_manual_fx_override(conn, base_currency='USD', quote_currency='CHF', rate_date='2025-12-31', rate=Decimal('0.9'), note='bank statement')

    assert rate_id
    assert conn.execute("SELECT COUNT(*) AS n FROM audit_log WHERE action='manual_fx_override'").fetchone()["n"] == 1


def test_instrument_price_mapping_requires_provider_symbol_and_exchange_review() -> None:
    conn = setup_conn()
    conn.execute("UPDATE instruments SET exchange=NULL WHERE instrument_id='usd1'")

    result = ensure_instrument_price_mapping_quality(conn, instrument_id='usd1')

    assert result.mapping_status == 'missing_provider_symbol'
    assert 'ticker_without_exchange' in result.quality_flags
    assert conn.execute("SELECT COUNT(*) AS n FROM alerts WHERE rule_id='missing_provider_symbol' AND status='active'").fetchone()["n"] == 1
    assert conn.execute("SELECT COUNT(*) AS n FROM alerts WHERE rule_id='ambiguous_instrument_mapping' AND status='active'").fetchone()["n"] == 1


def test_confirm_instrument_price_mapping_writes_audit() -> None:
    conn = setup_conn()

    mapping_id = confirm_instrument_price_mapping(conn, instrument_id='usd1', provider='mock', provider_symbol='USETF.N', provider_market='NYSE', confidence='manual', note='confirmed from broker')

    mapping = conn.execute("SELECT * FROM instrument_price_mappings WHERE mapping_id=?", (mapping_id,)).fetchone()
    assert mapping['mapping_status'] == 'mapped'
    assert mapping['provider_symbol'] == 'USETF.N'
    assert conn.execute("SELECT COUNT(*) AS n FROM audit_log WHERE action='confirm_instrument_price_mapping'").fetchone()["n"] == 1


def test_refresh_market_prices_with_mock_provider_stores_and_resolves_alert() -> None:
    conn = setup_conn()
    confirm_instrument_price_mapping(conn, instrument_id='usd1', provider='mock', provider_symbol='USETF.N', provider_market='NYSE', confidence='manual', note='test mapping')
    conn.execute("INSERT INTO alerts(alert_id,priority,category,entity_type,entity_id,rule_id,message,status,created_at) VALUES('mp1','warnung','market_data','instrument','usd1','missing_market_price','missing','active','now')")
    provider = MockEquityPriceProvider({'USETF.N': Decimal('10.25')})

    result = refresh_market_prices(conn, provider=provider, asset_class='etf', price_date='2026-01-02')

    assert result.updated_count == 1
    price = conn.execute("SELECT close, quality_status FROM market_prices WHERE instrument_id='usd1'").fetchone()
    assert price['close'] == '10.25'
    assert price['quality_status'] == 'fresh'
    assert conn.execute("SELECT status FROM alerts WHERE alert_id='mp1'").fetchone()["status"] == "resolved"


def test_missing_and_stale_market_price_alerts_are_deduped() -> None:
    conn = setup_conn()
    confirm_instrument_price_mapping(conn, instrument_id='usd1', provider='mock', provider_symbol='USETF.N', provider_market='NYSE', confidence='manual', note='test mapping')
    provider = MockEquityPriceProvider({'USETF.N': None})

    result = refresh_market_prices(conn, provider=provider, asset_class='etf', price_date='2026-01-02')
    result2 = refresh_market_prices(conn, provider=provider, asset_class='etf', price_date='2026-01-02')

    assert result.warning_count == 1
    assert result2.warning_count == 1
    assert conn.execute("SELECT COUNT(*) AS n FROM alerts WHERE rule_id='missing_market_price' AND status='active'").fetchone()["n"] == 1


def test_position_valuation_requires_market_price_and_fx_without_fake_total_return() -> None:
    conn = setup_conn()
    conn.execute("""
        INSERT INTO transactions(transaction_id, transaction_type, account_id, instrument_id, trade_date, quantity, currency_original, fx_status, source_type, is_confirmed, quality_status, created_at)
        VALUES('tx-usd','initial_position_snapshot','a1','usd1','2025-12-31','2','USD','missing','test',1,'warning','now')
    """)
    store_market_price(conn, instrument_id='usd1', price_date='2026-01-02', close=Decimal('10'), currency='USD', provider='mock', provider_symbol='USETF.N', quality_status='fresh')

    pos = calculate_positions(conn).positions[('a1', 'usd1')]

    assert pos.market_value_chf is None
    assert pos.total_return_chf is None
    assert 'missing_fx' in pos.quality_warnings

    upsert_fx_rate(conn, base_currency='USD', quote_currency='CHF', rate_date='2026-01-02', rate=Decimal('0.9'), provider='mock', rate_type='close', quality_status='fresh')
    pos2 = calculate_positions(conn).positions[('a1', 'usd1')]
    assert pos2.market_value_chf is not None
    assert pos2.total_return_chf is None  # cost basis still incomplete for snapshot-only import
    assert 'cost_basis_uncertain' in pos2.quality_warnings
