import numpy as np

from wormhole_proof.core.bridge import CCABridge


def test_cca_bridge_recovers_linear_mapping():
    rng = np.random.default_rng(7)
    donor = rng.normal(size=(256, 6))
    true_map = rng.normal(size=(6, 5))
    receiver = donor @ true_map + 0.05 * rng.normal(size=(256, 5))

    bridge = CCABridge(rank=5, regularization=1e-4)
    fit = bridge.fit(donor, receiver)

    delta = rng.normal(size=6)
    pred = bridge.map_delta(delta)
    target = true_map.T @ delta

    assert np.allclose(pred, target, atol=0.12)
    assert fit.alignment_error < 0.1
