import asyncio
from typing import Dict, List, Optional
from datetime import datetime, timedelta
import logging
from motor.motor_asyncio import AsyncIOMotorClient
import os

from snmp_manager import snmp_manager

logger = logging.getLogger(__name__)

class DeviceMonitor:
    """Real-time device monitoring system"""
    
    def __init__(self, db):
        self.db = db
        self.monitoring_tasks = {}
        self.running = False
    
    async def start_monitoring(self):
        """Start monitoring all devices"""
        self.running = True
        logger.info("Device monitoring started")
        
        while self.running:
            try:
                # Get all devices that should be monitored
                devices = await self.db.devices.find({
                    'snmp_enabled': True,
                    'status': {'$ne': 'offline'}
                }).to_list(1000)
                
                # Monitor devices concurrently
                tasks = []
                for device in devices[:50]:  # Limit to 50 concurrent monitors
                    tasks.append(self.monitor_device(device))
                
                if tasks:
                    await asyncio.gather(*tasks, return_exceptions=True)
                
                # Wait before next monitoring cycle
                await asyncio.sleep(60)  # Monitor every 60 seconds
                
            except Exception as e:
                logger.error(f"Monitoring loop error: {e}")
                await asyncio.sleep(60)
    
    async def monitor_device(self, device: Dict):
        """Monitor single device"""
        try:
            device_id = device['id']
            ip = device['ip']
            community = device.get('snmp_community', 'public')
            version = device.get('snmp_version', 'v2c')
            
            # Get device metrics
            metrics = await snmp_manager.get_full_device_info(
                ip, version, community
            )
            
            # Save metrics to database
            metric_doc = {
                'device_id': device_id,
                'ip': ip,
                'timestamp': datetime.utcnow(),
                'cpu': metrics.get('cpu'),
                'memory': metrics.get('memory'),
                'temperature': metrics.get('temperature'),
                'interfaces': metrics.get('interfaces', []),
                'system': metrics.get('system', {})
            }
            
            await self.db.device_metrics.insert_one(metric_doc)
            
            # Update device status
            update_data = {
                'last_seen': datetime.utcnow(),
                'status': 'online',
                'cpu_usage': metrics.get('cpu'),
                'memory_usage': metrics.get('memory', {}).get('percentage'),
                'temperature': metrics.get('temperature'),
            }
            
            await self.db.devices.update_one(
                {'id': device_id},
                {'$set': update_data}
            )
            
            # Check for alerts
            await self.check_alerts(device, metrics)
            
            logger.debug(f"Monitored device {ip}: CPU={metrics.get('cpu')}%")
            
        except Exception as e:
            logger.error(f"Error monitoring device {device.get('ip')}: {e}")
            
            # Mark device as offline if monitoring fails
            await self.db.devices.update_one(
                {'id': device.get('id')},
                {'$set': {'status': 'offline', 'last_seen': datetime.utcnow()}}
            )
    
    async def check_alerts(self, device: Dict, metrics: Dict):
        """Check if device metrics trigger any alerts"""
        alerts = []
        device_id = device['id']
        
        # CPU Alert
        cpu = metrics.get('cpu')
        if cpu and cpu > 80:
            alerts.append({
                'device_id': device_id,
                'type': 'cpu_high',
                'severity': 'warning' if cpu < 90 else 'critical',
                'message': f"High CPU usage: {cpu}%",
                'value': cpu,
                'threshold': 80
            })
        
        # Memory Alert
        memory = metrics.get('memory', {})
        mem_usage = memory.get('percentage', 0)
        if mem_usage > 85:
            alerts.append({
                'device_id': device_id,
                'type': 'memory_high',
                'severity': 'warning' if mem_usage < 95 else 'critical',
                'message': f"High memory usage: {mem_usage:.1f}%",
                'value': mem_usage,
                'threshold': 85
            })
        
        # Temperature Alert
        temp = metrics.get('temperature')
        if temp and temp > 70:
            alerts.append({
                'device_id': device_id,
                'type': 'temperature_high',
                'severity': 'warning' if temp < 80 else 'critical',
                'message': f"High temperature: {temp}°C",
                'value': temp,
                'threshold': 70
            })
        
        # Interface Down Alert
        interfaces = metrics.get('interfaces', [])
        for iface in interfaces:
            if iface.get('admin_status') == 'up' and iface.get('status') == 'down':
                alerts.append({
                    'device_id': device_id,
                    'type': 'interface_down',
                    'severity': 'warning',
                    'message': f"Interface {iface['name']} is down",
                    'interface': iface['name']
                })
        
        # Save alerts
        if alerts:
            timestamp = datetime.utcnow()
            for alert in alerts:
                alert['timestamp'] = timestamp
                alert['ip'] = device['ip']
                alert['hostname'] = device.get('hostname', 'Unknown')
                alert['acknowledged'] = False
                
                # Check if alert already exists (avoid duplicates)
                existing = await self.db.alerts.find_one({
                    'device_id': device_id,
                    'type': alert['type'],
                    'acknowledged': False
                })
                
                if not existing:
                    await self.db.alerts.insert_one(alert)
                    logger.info(f"Alert created: {alert['type']} for {device['ip']}")
    
    async def get_device_metrics(self, device_id: str, hours: int = 24) -> List[Dict]:
        """Get historical metrics for a device"""
        since = datetime.utcnow() - timedelta(hours=hours)
        
        metrics = await self.db.device_metrics.find({
            'device_id': device_id,
            'timestamp': {'$gte': since}
        }).sort('timestamp', 1).to_list(1000)
        
        return metrics
    
    async def get_interface_bandwidth(self, device_id: str, interface: str, 
                                     hours: int = 1) -> Dict:
        """Calculate interface bandwidth usage"""
        metrics = await self.get_device_metrics(device_id, hours)
        
        bandwidth_data = []
        
        for i in range(1, len(metrics)):
            prev = metrics[i-1]
            curr = metrics[i]
            
            # Find interface in both samples
            prev_iface = next((iface for iface in prev.get('interfaces', []) 
                             if iface['name'] == interface), None)
            curr_iface = next((iface for iface in curr.get('interfaces', []) 
                             if iface['name'] == interface), None)
            
            if prev_iface and curr_iface:
                time_diff = (curr['timestamp'] - prev['timestamp']).total_seconds()
                if time_diff > 0:
                    in_bps = (curr_iface['in_octets'] - prev_iface['in_octets']) * 8 / time_diff
                    out_bps = (curr_iface['out_octets'] - prev_iface['out_octets']) * 8 / time_diff
                    
                    bandwidth_data.append({
                        'timestamp': curr['timestamp'],
                        'in_bps': max(0, in_bps),
                        'out_bps': max(0, out_bps)
                    })
        
        return {
            'interface': interface,
            'data': bandwidth_data
        }
    
    def stop_monitoring(self):
        """Stop monitoring"""
        self.running = False
        logger.info("Device monitoring stopped")

# Global monitor instance (will be initialized in server.py)
device_monitor = None
