import logging
from typing import Dict, List, Optional
from netmiko import ConnectHandler
from netmiko.exceptions import NetMikoTimeoutException, NetMikoAuthenticationException

logger = logging.getLogger(__name__)

class CiscoSSHManager:
    """Cisco Device SSH Manager"""
    
    DEVICE_TYPES = {
        'ios': 'cisco_ios',
        'xe': 'cisco_xe',
        'xr': 'cisco_xr',
        'nxos': 'cisco_nxos',
        'asa': 'cisco_asa',
    }
    
    def __init__(self):
        self.connections = {}
    
    def connect(self, ip: str, username: str, password: str, 
                device_type: str = 'ios', port: int = 22,
                secret: Optional[str] = None) -> ConnectHandler:
        """Connect to Cisco device"""
        key = f"{ip}:{port}"
        
        if key not in self.connections:
            try:
                device = {
                    'device_type': self.DEVICE_TYPES.get(device_type, 'cisco_ios'),
                    'host': ip,
                    'username': username,
                    'password': password,
                    'port': port,
                    'timeout': 30,
                    'session_timeout': 60,
                }
                
                if secret:
                    device['secret'] = secret
                
                connection = ConnectHandler(**device)
                self.connections[key] = connection
                logger.info(f"Connected to Cisco device {ip}")
                return connection
                
            except NetMikoAuthenticationException as e:
                logger.error(f"Authentication failed for {ip}: {e}")
                raise Exception("Authentication failed")
            except NetMikoTimeoutException as e:
                logger.error(f"Connection timeout for {ip}: {e}")
                raise Exception("Connection timeout")
            except Exception as e:
                logger.error(f"Failed to connect to {ip}: {e}")
                raise
        
        return self.connections[key]
    
    def disconnect(self, ip: str, port: int = 22):
        """Disconnect from device"""
        key = f"{ip}:{port}"
        if key in self.connections:
            try:
                self.connections[key].disconnect()
                del self.connections[key]
            except:
                pass
    
    def execute_command(self, ip: str, username: str, password: str,
                       command: str, device_type: str = 'ios',
                       enable: bool = False, secret: Optional[str] = None) -> str:
        """Execute command on device"""
        try:
            connection = self.connect(ip, username, password, device_type, secret=secret)
            
            if enable and secret:
                connection.enable()
            
            output = connection.send_command(command)
            return output
            
        except Exception as e:
            logger.error(f"Failed to execute command on {ip}: {e}")
            raise
    
    def execute_config_commands(self, ip: str, username: str, password: str,
                                commands: List[str], device_type: str = 'ios',
                                secret: Optional[str] = None) -> str:
        """Execute configuration commands"""
        try:
            connection = self.connect(ip, username, password, device_type, secret=secret)
            
            if secret:
                connection.enable()
            
            output = connection.send_config_set(commands)
            connection.save_config()
            
            return output
            
        except Exception as e:
            logger.error(f"Failed to execute config commands on {ip}: {e}")
            raise
    
    def get_running_config(self, ip: str, username: str, password: str,
                          device_type: str = 'ios', secret: Optional[str] = None) -> str:
        """Get running configuration"""
        return self.execute_command(
            ip, username, password, 
            'show running-config',
            device_type, enable=True, secret=secret
        )
    
    def get_version(self, ip: str, username: str, password: str,
                   device_type: str = 'ios', secret: Optional[str] = None) -> Dict:
        """Get version information"""
        try:
            output = self.execute_command(
                ip, username, password,
                'show version',
                device_type, secret=secret
            )
            
            # Parse version info
            version_info = {
                'raw_output': output,
                'hostname': self._extract_hostname(output),
                'version': self._extract_version(output),
                'uptime': self._extract_uptime(output),
                'model': self._extract_model(output),
            }
            
            return version_info
            
        except Exception as e:
            logger.error(f"Failed to get version from {ip}: {e}")
            raise
    
    def get_interfaces(self, ip: str, username: str, password: str,
                      device_type: str = 'ios', secret: Optional[str] = None) -> str:
        """Get interface status"""
        return self.execute_command(
            ip, username, password,
            'show ip interface brief',
            device_type, secret=secret
        )
    
    def get_arp_table(self, ip: str, username: str, password: str,
                     device_type: str = 'ios', secret: Optional[str] = None) -> str:
        """Get ARP table"""
        return self.execute_command(
            ip, username, password,
            'show ip arp',
            device_type, secret=secret
        )
    
    def get_mac_address_table(self, ip: str, username: str, password: str,
                             device_type: str = 'ios', secret: Optional[str] = None) -> str:
        """Get MAC address table"""
        return self.execute_command(
            ip, username, password,
            'show mac address-table',
            device_type, secret=secret
        )
    
    def get_vlan_info(self, ip: str, username: str, password: str,
                     device_type: str = 'ios', secret: Optional[str] = None) -> str:
        """Get VLAN information"""
        return self.execute_command(
            ip, username, password,
            'show vlan brief',
            device_type, secret=secret
        )
    
    def backup_config(self, ip: str, username: str, password: str,
                     device_type: str = 'ios', secret: Optional[str] = None) -> str:
        """Backup running configuration"""
        return self.get_running_config(ip, username, password, device_type, secret)
    
    # Helper methods for parsing
    def _extract_hostname(self, output: str) -> str:
        for line in output.split('\n'):
            if 'hostname' in line.lower():
                parts = line.split()
                if len(parts) >= 2:
                    return parts[1]
        return "Unknown"
    
    def _extract_version(self, output: str) -> str:
        for line in output.split('\n'):
            if 'Version' in line or 'version' in line:
                return line.strip()
        return "Unknown"
    
    def _extract_uptime(self, output: str) -> str:
        for line in output.split('\n'):
            if 'uptime' in line.lower():
                return line.strip()
        return "Unknown"
    
    def _extract_model(self, output: str) -> str:
        for line in output.split('\n'):
            if 'cisco' in line.lower() and ('processor' in line.lower() or 'model' in line.lower()):
                return line.strip()
        return "Unknown"

# Global instance
cisco_ssh_manager = CiscoSSHManager()
