from __future__ import annotations

from decimal import Decimal

import pytest

from jarvis_finance.fx.overrides import set_manual_fx_override
from jarvis_finance.fx.providers import HistoricalFxUnavailable, MockFxProvider
from jarvis_finance.fx.rates import update_fx_rates, upsert_fx_rate
from jarvis_finance.ledger.positions import calculate_positions
from jarvis_finance.market_data.instruments import ensure_instrument_metadata_quality, update_instrument_metadata
from jarvis_finance.market_data.mappings import confirm_instrument_price_mapping, ensure_instrument_price_mapping_quality
from jarvis_finance.market_data.prices import MockEquityPriceProvider, refresh_market_prices, store_market_price
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('usd1','ETF','USD ETF','USETF','US0000000001','NYSE','USD','now')")
    return conn


def add_snapshot(conn, *, fx_status='missing', fx_rate=None):
    conn.execute(
        """
        INSERT INTO transactions(transaction_id, transaction_type, account_id, instrument_id, trade_date, quantity, currency_original, fx_status, fx_rate_to_chf, source_type, is_confirmed, quality_status, created_at)
        VALUES('tx-usd','initial_position_snapshot','a1','usd1','2025-12-31','2','USD',?,?, 'true_wealth',1,'warning','now')
        """,
        (fx_status, fx_rate),
    )


def test_provider_symbol_missing_alert_and_confirm_resolves_with_audit() -> None:
    conn = setup_conn()
    result = ensure_instrument_price_mapping_quality(conn, instrument_id='usd1')
    assert result.mapping_status == 'missing_provider_symbol'
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='missing_provider_symbol' AND status='active'").fetchone()['c'] == 1

    confirm_instrument_price_mapping(conn, instrument_id='usd1', provider='mock', provider_symbol='USETF.N', provider_market='NYSE', confidence='manual', note='confirmed in pilot', hedge_status='unhedged', instrument_status='active')

    assert conn.execute("SELECT COUNT(*) AS c FROM audit_log WHERE action='confirm_instrument_price_mapping'").fetchone()['c'] == 1
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='missing_provider_symbol' AND status='active'").fetchone()['c'] == 0


def test_hedge_status_unknown_warns_and_change_audits() -> None:
    conn = setup_conn()
    ensure_instrument_metadata_quality(conn, instrument_id='usd1')
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='hedge_status_unknown' AND status='active'").fetchone()['c'] == 1

    update_instrument_metadata(conn, instrument_id='usd1', hedge_status='hedged', hedged_to_currency='CHF', note='manual hedge confirmation')

    inst = conn.execute("SELECT hedge_status,is_currency_hedged,hedged_to_currency FROM instruments WHERE instrument_id='usd1'").fetchone()
    assert inst['hedge_status'] == 'hedged'
    assert inst['is_currency_hedged'] == 1
    assert conn.execute("SELECT COUNT(*) AS c FROM audit_log WHERE action='update_instrument_metadata'").fetchone()['c'] == 1
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='hedge_status_unknown' AND status='active'").fetchone()['c'] == 0


def test_chf_hedged_instrument_marks_fx_attribution_not_free_fx_pnl() -> None:
    conn = setup_conn()
    add_snapshot(conn, fx_status='ok', fx_rate='0.9')
    update_instrument_metadata(conn, instrument_id='usd1', hedge_status='hedged', hedged_to_currency='CHF', instrument_status='active', note='hedged review')
    store_market_price(conn, instrument_id='usd1', price_date='2026-01-02', close=Decimal('10'), currency='USD', provider='mock', provider_symbol='USETF.N', price_timestamp='2026-01-02T20:00:00+00:00')
    upsert_fx_rate(conn, base_currency='USD', quote_currency='CHF', rate_date='2026-01-02', rate=Decimal('0.91'), provider='mock', rate_type='close')

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

    assert pos.market_value_chf is not None
    assert 'currency_hedged_fx_attribution' in pos.quality_warnings
    assert pos.total_return_chf is None


def test_historical_fx_unavailable_sets_manual_override_required_without_crash() -> None:
    conn = setup_conn()
    add_snapshot(conn)
    provider = MockFxProvider({}, unavailable={('USD', 'CHF', '2025-12-31')})

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

    assert result.warning_count == 1
    row = conn.execute("SELECT quality_status,error_status FROM fx_rates WHERE base_currency='USD' AND quote_currency='CHF'").fetchone()
    assert row['quality_status'] == 'manual_override_required'
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='historical_fx_unavailable' AND status='active'").fetchone()['c'] == 1
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='manual_override_required' AND status='active'").fetchone()['c'] == 1


def test_manual_fx_override_requires_note_audits_and_resolves_manual_required() -> None:
    conn = setup_conn()
    add_snapshot(conn)
    conn.execute("INSERT INTO alerts(alert_id,priority,category,entity_type,entity_id,rule_id,message,status,created_at) VALUES('m1','kritisch','fx','instrument','usd1','manual_override_required','manual','active','now')")
    with pytest.raises(ValueError):
        set_manual_fx_override(conn, base_currency='USD', quote_currency='CHF', rate_date='2025-12-31', rate=Decimal('0.9'), note='')

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

    assert conn.execute("SELECT COUNT(*) AS c FROM audit_log WHERE action='manual_fx_override'").fetchone()['c'] == 1
    assert conn.execute("SELECT status FROM alerts WHERE alert_id='m1'").fetchone()['status'] == 'resolved'


def test_market_price_resolves_missing_and_delisted_is_excluded_without_alert_flood() -> None:
    conn = setup_conn()
    confirm_instrument_price_mapping(conn, instrument_id='usd1', provider='mock', provider_symbol='USETF.N', provider_market='NYSE', confidence='manual', note='map', hedge_status='unhedged', instrument_status='active')
    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')")
    result = refresh_market_prices(conn, provider=MockEquityPriceProvider({'USETF.N': Decimal('10')}), asset_class='etf', price_date='2026-01-02')
    assert result.updated_count == 1
    assert conn.execute("SELECT status FROM alerts WHERE alert_id='mp1'").fetchone()['status'] == 'resolved'

    update_instrument_metadata(conn, instrument_id='usd1', instrument_status='delisted', valuation_policy='exclude_from_auto_price_update', note='status review')
    result2 = refresh_market_prices(conn, provider=MockEquityPriceProvider({'USETF.N': None}), asset_class='etf', price_date='2026-01-03')
    result3 = refresh_market_prices(conn, provider=MockEquityPriceProvider({'USETF.N': None}), asset_class='etf', price_date='2026-01-03')
    assert result2.excluded_count == 1
    assert result3.excluded_count == 1
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='delisted_or_suspended' AND status='active'").fetchone()['c'] == 1


def test_timestamp_alignment_and_stale_valuation_alert() -> None:
    conn = setup_conn()
    add_snapshot(conn, fx_status='ok', fx_rate='0.9')
    update_instrument_metadata(conn, instrument_id='usd1', hedge_status='unhedged', instrument_status='active', note='review')
    store_market_price(conn, instrument_id='usd1', price_date='2026-01-02', close=Decimal('10'), currency='USD', provider='mock', provider_symbol='USETF.N', price_timestamp='2026-01-02T20:00:00+00:00')
    upsert_fx_rate(conn, base_currency='USD', quote_currency='CHF', rate_date='2026-01-02', rate=Decimal('0.9'), provider='mock', rate_type='close')
    pos = calculate_positions(conn).positions[('a1', 'usd1')]
    assert pos.valuation_timestamp_status in {'aligned', 'acceptable'}

    conn.execute("DELETE FROM fx_rates")
    conn.execute("UPDATE alerts SET status='resolved'")
    upsert_fx_rate(conn, base_currency='USD', quote_currency='CHF', rate_date='2026-01-01', rate=Decimal('0.9'), provider='mock', rate_type='close')
    pos2 = calculate_positions(conn, max_price_fx_time_delta_hours=1).positions[('a1', 'usd1')]
    assert pos2.valuation_timestamp_status == 'stale_mismatch'
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='stale_valuation' AND status='active'").fetchone()['c'] == 1


def test_corporate_action_large_price_move_flags_no_auto_split_correction() -> None:
    conn = setup_conn()
    store_market_price(conn, instrument_id='usd1', price_date='2026-01-01', close=Decimal('100'), currency='USD', provider='mock', provider_symbol='USETF.N')
    store_market_price(conn, instrument_id='usd1', price_date='2026-01-02', close=Decimal('50'), currency='USD', provider='mock', provider_symbol='USETF.N')

    inst = conn.execute("SELECT corporate_action_status,split_or_corporate_action_review_required FROM instruments WHERE instrument_id='usd1'").fetchone()
    assert inst['corporate_action_status'] == 'suspected'
    assert inst['split_or_corporate_action_review_required'] == 1
    assert conn.execute("SELECT COUNT(*) AS c FROM alerts WHERE rule_id='corporate_action_suspected' AND status='active'").fetchone()['c'] == 1


def test_valuation_statuses_and_total_return_remains_incomplete_for_snapshot() -> None:
    conn = setup_conn()
    add_snapshot(conn, fx_status='ok', fx_rate='0.9')
    update_instrument_metadata(conn, instrument_id='usd1', hedge_status='unhedged', instrument_status='active', corporate_action_status='none_known', note='review')
    pos_missing = calculate_positions(conn).positions[('a1', 'usd1')]
    assert pos_missing.valuation_status == 'partially_valuable'

    store_market_price(conn, instrument_id='usd1', price_date='2026-01-02', close=Decimal('10'), currency='USD', provider='mock', provider_symbol='USETF.N', price_timestamp='2026-01-02T20:00:00+00:00')
    upsert_fx_rate(conn, base_currency='USD', quote_currency='CHF', rate_date='2026-01-02', rate=Decimal('0.9'), provider='mock', rate_type='close')
    pos = calculate_positions(conn).positions[('a1', 'usd1')]
    assert pos.valuation_status == 'valuable'
    assert pos.total_return_chf is None
    assert 'snapshot_only' in pos.quality_warnings

    update_instrument_metadata(conn, instrument_id='usd1', instrument_status='delisted', valuation_policy='exclude_from_auto_price_update', note='exclude')
    pos_blocked = calculate_positions(conn).positions[('a1', 'usd1')]
    assert pos_blocked.valuation_status == 'not_valuable'
