"""Shim selecting the appropriate proto language model implementation.

MONIKA requires PyTorch to function. This module imports the PyTorch-based
implementation. Fallback mode exists only for unit tests.
"""

from __future__ import annotations

try:  # pragma: no cover - import guard
    import torch  # noqa: F401
    # Test if torch actually works by trying to access a basic attribute
    _ = torch.__version__
    _TORCH_AVAILABLE = True
except ModuleNotFoundError as e:  # pragma: no cover - PyTorch not installed
    _TORCH_AVAILABLE = False
    raise RuntimeError(
        "PyTorch is not installed. MONIKA requires PyTorch to function.\n"
        "Install with: pip install torch"
    ) from e
except (OSError, AttributeError) as e:  # pragma: no cover - PyTorch installed but broken
    _TORCH_AVAILABLE = False
    raise RuntimeError(
        f"PyTorch import failed: {e}\n"
        "MONIKA requires a working PyTorch installation.\n"
        "This may be due to:\n"
        "  - Python version mismatch (PyTorch may not support Python 3.13 yet)\n"
        "  - Missing Visual C++ Redistributables\n"
        "  - CUDA driver/runtime mismatch\n"
        "  - Corrupted PyTorch installation\n"
        "\nTo fix, try:\n"
        "  pip uninstall torch torchvision torchaudio\n"
        "  pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu131"
    ) from e

if getattr(torch, "__SALIENT_STUB__", False):  # pragma: no cover - fallback to stub
    _TORCH_AVAILABLE = False
    raise RuntimeError("PyTorch stub detected - this should not happen in production")

if not _TORCH_AVAILABLE:  # pragma: no cover - should not happen due to exceptions above
    # Only import fallback for unit tests that explicitly need it
    from ._fallback import CheckpointRecord, ProtoLanguageModel, TrainingConfig
else:  # pragma: no cover - exercised when torch is installed
    from ._torch_impl import ProtoLanguageModel
    from ._fallback import TrainingConfig
    from .checkpoints import CheckpointRecord

__all__ = ["TrainingConfig", "ProtoLanguageModel", "CheckpointRecord"]
