import hashlib
import sys

import numpy as np
import pandas as pd
import pytest

from strategies.mlp_feature_groups import (
    GROUPED_UNIT_ALLOCATION,
    HYBRID_GROUPED_DENSE_ROW_INDICES,
    HYBRID_GROUPED_SPECIALIST_ALLOCATION,
    grouped_first_layer_mask,
    hybrid_grouped_first_layer_mask,
    hybrid_grouped_structure_metadata,
    random_sparse_first_layer_mask,
    random_sparse_structure_metadata,
)
from strategies.strategy_mlp_scores import FEATURE_COLS
from tools import run_mlp_train, train_mlp
from tools.train_mlp import flatten_layers, pretrain, unflatten_layers
from tools.run_mlp_train import effective_artifact_tag


def test_grouped_mask_routes_each_first_layer_unit_to_one_signal_family():
    mask = grouped_first_layer_mask(FEATURE_COLS, hidden_size=16)

    assert mask.shape == (16, len(FEATURE_COLS))
    assert tuple(mask.sum(axis=0) > 0) == (True,) * len(FEATURE_COLS)
    for row, family in zip(mask, GROUPED_UNIT_ALLOCATION):
        expected = {i for i, feature in enumerate(FEATURE_COLS) if feature in family.features}
        assert set(np.flatnonzero(row)) == expected


def test_grouped_mask_rejects_unsupported_width_and_unknown_features():
    with pytest.raises(ValueError, match="16"):
        grouped_first_layer_mask(FEATURE_COLS, hidden_size=8)
    with pytest.raises(ValueError, match="unknown"):
        grouped_first_layer_mask([*FEATURE_COLS, "future_feature"], hidden_size=16)


def test_random_sparse_mask_matches_grouped_row_sparsity_and_is_reproducible():
    grouped = grouped_first_layer_mask(FEATURE_COLS, hidden_size=16)
    first = random_sparse_first_layer_mask(FEATURE_COLS, hidden_size=16, mask_seed=13001)
    again = random_sparse_first_layer_mask(FEATURE_COLS, hidden_size=16, mask_seed=13001)
    different = random_sparse_first_layer_mask(FEATURE_COLS, hidden_size=16, mask_seed=13102)

    assert first.shape == (16, len(FEATURE_COLS))
    assert first.sum() == grouped.sum() == 149
    np.testing.assert_array_equal(first.sum(axis=1), grouped.sum(axis=1))
    np.testing.assert_array_equal(first, again)
    assert not np.array_equal(first, different)


def test_random_sparse_metadata_records_reproducible_mask_provenance():
    metadata = random_sparse_structure_metadata(FEATURE_COLS, hidden_size=16, mask_seed=13001)

    assert metadata["name"] == "random_sparse_first_layer_v1"
    assert metadata["mask_seed"] == 13001
    assert metadata["active_edges"] == 149
    assert metadata["row_active_edges"] == [13, 13, 13, 10, 10, 10, 11, 11, 11, 4, 10, 10, 10, 6, 6, 1]
    assert len(metadata["row_feature_indices"]) == 16


def test_hybrid_grouped_mask_preserves_specialists_and_dense_cross_family_rows():
    mask = hybrid_grouped_first_layer_mask(FEATURE_COLS, hidden_size=16)

    assert mask.shape == (16, len(FEATURE_COLS))
    assert int(mask.sum()) == 508
    assert int(mask.size - mask.sum()) == 372
    for row, family in enumerate(HYBRID_GROUPED_SPECIALIST_ALLOCATION):
        expected = {i for i, feature in enumerate(FEATURE_COLS) if feature in family.features}
        assert set(np.flatnonzero(mask[row])) == expected
    for row in HYBRID_GROUPED_DENSE_ROW_INDICES:
        assert mask[row].all()
    with pytest.raises(ValueError, match="16"):
        hybrid_grouped_first_layer_mask(FEATURE_COLS, hidden_size=8)


def test_hybrid_grouped_metadata_audits_exact_v1_mask():
    metadata = hybrid_grouped_structure_metadata(FEATURE_COLS, hidden_size=16)

    assert metadata["name"] == "hybrid_grouped_first_layer_v1"
    assert metadata["version"] == 1
    assert metadata["specialist_row_families"] == {
        "0": "local_technical", "1": "local_technical", "2": "on_chain",
        "3": "macro", "4": "candlesticks", "5": "rsid_structure",
        "6": "market_flow", "7": "cross_asset",
    }
    assert metadata["dense_row_indices"] == list(range(8, 16))
    assert metadata["active_edges"] == 508
    assert metadata["blocked_edges"] == 372
    assert metadata["row_active_edges"] == [13, 13, 10, 11, 4, 10, 6, 1] + [55] * 8
    assert metadata["feature_cols_sha256"] == hashlib.sha256("\x1f".join(FEATURE_COLS).encode()).hexdigest()
    mask = hybrid_grouped_first_layer_mask(FEATURE_COLS, hidden_size=16)
    assert metadata["row_feature_indices"] == [np.flatnonzero(row).tolist() for row in mask]


def test_hybrid_phase1_and_cma_masking_preserve_specialist_blocks():
    rng = np.random.default_rng(11)
    mask = hybrid_grouped_first_layer_mask(FEATURE_COLS, hidden_size=16)
    X = rng.normal(size=(72, len(FEATURE_COLS)))
    y = rng.normal(size=72)
    valid = np.ones(72, dtype=bool)
    times = pd.Series(pd.date_range("2020-01-01", periods=72, freq="h"))

    layers, _ = pretrain(
        X, y, valid, times, [16, 8], epochs=2, seed=13,
        val_start=times.iloc[60], val_end=times.iloc[-1], log=lambda _: None,
        first_layer_mask=mask,
    )
    assert np.count_nonzero(layers[0][0][~mask]) == 0
    assert mask[8:].all()

    masks = [mask, np.ones((8, 16), dtype=bool), np.ones((1, 8), dtype=bool)]
    theta = flatten_layers(layers, masks)
    restored = unflatten_layers(theta, [55, 16, 8, 1], masks)
    assert len(theta) == 669
    assert len(flatten_layers(layers)) == 1041
    assert np.count_nonzero(restored[0][0][~mask]) == 0


def test_masked_parameter_roundtrip_never_restores_blocked_connections():
    mask = grouped_first_layer_mask(FEATURE_COLS, hidden_size=16)
    rng = np.random.default_rng(4)
    layers = [
        (rng.normal(size=(16, len(FEATURE_COLS))), rng.normal(size=16)),
        (rng.normal(size=(8, 16)), rng.normal(size=8)),
        (rng.normal(size=(1, 8)), rng.normal(size=1)),
    ]
    masks = [mask, np.ones((8, 16), dtype=bool), np.ones((1, 8), dtype=bool)]

    theta = flatten_layers(layers, masks)
    restored = unflatten_layers(theta, [55, 16, 8, 1], masks)

    assert np.all(restored[0][0][~mask] == 0.0)
    np.testing.assert_array_equal(restored[0][0][mask], layers[0][0][mask])


def test_grouped_training_uses_an_isolated_default_artifact_tag():
    assert effective_artifact_tag("bb50", "all", "grouped") == "bb50_grouped"
    assert effective_artifact_tag("bb50", "bull", "grouped") == "bb50_grouped_bull"
    assert effective_artifact_tag("grouped_v1", "all", "grouped") == "grouped_v1"
    assert effective_artifact_tag("bb50", "all", "dense") == "bb50"
    assert effective_artifact_tag("bb50", "all", "random_sparse") == "bb50_random_sparse"
    assert effective_artifact_tag("bb50", "bull", "random_sparse") == "bb50_random_sparse_bull"
    assert effective_artifact_tag("grouped_v1", "all", "random_sparse") == "grouped_v1"
    assert effective_artifact_tag("bb50", "all", "hybrid_grouped") == "bb50_hybrid_grouped"
    assert effective_artifact_tag("bb50", "bull", "hybrid_grouped") == "bb50_hybrid_grouped_bull"
    assert effective_artifact_tag("hybrid_grouped_v1", "all", "hybrid_grouped") == "hybrid_grouped_v1"
    with pytest.raises(ValueError, match="reserved"):
        effective_artifact_tag("grouped_v1", "all", "hybrid_grouped")
    with pytest.raises(ValueError, match="reserved"):
        effective_artifact_tag("random_sparse_v1", "all", "hybrid_grouped")
    with pytest.raises(ValueError, match="reserved"):
        effective_artifact_tag("grouped_v1_bull", "all", "hybrid_grouped")


def test_hybrid_direct_trainer_requires_an_isolated_output(monkeypatch):
    monkeypatch.setattr(sys, "argv", [
        "train_mlp.py", "--data", "unused.csv", "--input-structure", "hybrid_grouped",
    ])
    with pytest.raises(SystemExit, match="2"):
        train_mlp.main()


def test_hybrid_orchestrator_rejects_shared_sweep_and_preset_paths(monkeypatch):
    monkeypatch.setattr(sys, "argv", [
        "run_mlp_train.py", "--input-structure", "hybrid_grouped", "--sweep",
    ])
    with pytest.raises(SystemExit, match="2"):
        run_mlp_train.main()
