"""
Утилиты для работы с PyTorch: CUDA OOM fallback, device management.

Использование:
    from utils.torch_utils import with_cpu_fallback, safe_device_call
"""

import functools
import logging
import torch
import config

logger = logging.getLogger('AI_Strategy')


def with_cpu_fallback(func):
    """Декоратор: при CUDA OOM переключает на CPU и повторяет вызов.

    Использует глобальный DEVICE из config для переключения.
    После успешного вызова на CPU, оставляет устройство CPU для последующих вызовов.

    Пример:
        @with_cpu_fallback
        def load_model(ticker):
            ...
    """
    @functools.wraps(func)
    def wrapper(*args, **kwargs):
        try:
            return func(*args, **kwargs)
        except (RuntimeError, torch.OutOfMemoryError) as e:
            if 'out of memory' in str(e).lower() or 'CUDA' in str(e):
                logger.warning(f'⚠ CUDA OOM в {func.__name__}, переключаем на CPU')
                _switch_to_cpu()
                # Повторная попытка
                return func(*args, **kwargs)
            else:
                raise
    return wrapper


def _switch_to_cpu():
    """Переключает глобальное устройство на CPU и очищает GPU память."""
    config.set_device('inference')
    import gc
    gc.collect()
    try:
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
    except Exception:
        pass
    logger.info('Устройство переключено на CPU, GPU память очищена')


def safe_device_call(func, *args, fallback_cpu: bool = True, **kwargs):
    """Безопасный вызов функции с fallback на CPU при OOM.
    
    Args:
        func: вызываемая функция
        fallback_cpu: если True — при OOM повторяет на CPU
        *args, **kwargs: аргументы для func
    
    Returns:
        Результат func(*args, **kwargs)
    
    Raises:
        RuntimeError: если OOM даже на CPU
    """
    try:
        return func(*args, **kwargs)
    except (RuntimeError, torch.OutOfMemoryError) as e:
        if fallback_cpu and ('out of memory' in str(e).lower() or 'CUDA' in str(e)):
            logger.warning(f'⚠ CUDA OOM в {getattr(func, "__name__", "unknown")}, переключаем на CPU')
            _switch_to_cpu()
            return func(*args, **kwargs)
        raise
