from __future__ import annotations

from pathlib import Path

from jarvis_finance.imports.instruments_importer import import_instruments_csv
from jarvis_finance.storage.database import connect_memory
from jarvis_finance.storage.migrations import apply_migrations


def test_instruments_import_commit_and_idempotent(tmp_path: Path) -> None:
    path = tmp_path / "instruments.csv"
    path.write_text(
        "asset_class,name,ticker,isin,exchange,currency,country,sector,data_provider_primary,notes\n"
        "equity,Demo Global AG,DGA,CH0000000001,SIX,CHF,CH,Technology,dummy,Synthetic only\n",
        encoding="utf-8",
    )
    conn = connect_memory()
    apply_migrations(conn)

    first = import_instruments_csv(conn, path, commit=True, source_filename="instruments.csv")
    second = import_instruments_csv(conn, path, commit=True, source_filename="instruments.csv")

    assert first.rows_new == 1
    assert second.rows_existing == 1
    assert conn.execute("SELECT COUNT(*) AS n FROM instruments").fetchone()["n"] == 1


def test_instruments_dry_run_does_not_write(tmp_path: Path) -> None:
    path = tmp_path / "instruments.csv"
    path.write_text(
        "asset_class,name,ticker,isin,exchange,currency,country,sector,data_provider_primary,notes\n"
        "etf,Demo World ETF,DWLD,CH0000000002,SIX,CHF,CH,Diversified,dummy,Synthetic only\n",
        encoding="utf-8",
    )
    conn = connect_memory()
    apply_migrations(conn)

    result = import_instruments_csv(conn, path, commit=False, source_filename="instruments.csv")

    assert result.rows_new == 1
    assert result.status == "dry_run_ok"
    assert conn.execute("SELECT COUNT(*) AS n FROM instruments").fetchone()["n"] == 0


def test_instruments_import_collects_validation_errors(tmp_path: Path) -> None:
    path = tmp_path / "bad_instruments.csv"
    path.write_text(
        "asset_class,name,ticker,isin,exchange,currency\n"
        "equity,,DGA,CH0000000001,SIX,CHF\n",
        encoding="utf-8",
    )
    conn = connect_memory()
    apply_migrations(conn)

    result = import_instruments_csv(conn, path, commit=True, source_filename="bad_instruments.csv")

    assert result.rows_failed == 1
    assert result.status == "failed"
    assert "name" in result.errors[0]
