from typing import List, Optional from sqlalchemy.orm import Session from sqlalchemy import func, case import ipaddress from app.models.network import Network as NetworkModel from app.models.network import IPAddress as IPAddressModel from app.models.network import IPStatus from app.schemas.network import NetworkCreate, NetworkUpdate, NetworkStats class NetworkService: """网段管理服务""" @staticmethod def calculate_total_ips(cidr: str) -> int: """计算网段的总IP数量""" network = ipaddress.ip_network(cidr, strict=False) return network.num_addresses @staticmethod def get_network_addresses(cidr: str) -> List[str]: """获取网段内所有IP地址列表""" network = ipaddress.ip_network(cidr, strict=False) return [str(ip) for ip in network.hosts()] @staticmethod def get_by_id(db: Session, network_id: int) -> Optional[NetworkModel]: """根据ID获取网段""" return db.query(NetworkModel).filter(NetworkModel.id == network_id).first() @staticmethod def get_by_cidr(db: Session, cidr: str) -> Optional[NetworkModel]: """根据CIDR获取网段""" return db.query(NetworkModel).filter(NetworkModel.cidr == cidr).first() @staticmethod def get_list( db: Session, skip: int = 0, limit: int = 100, group_name: Optional[str] = None ) -> tuple[int, List[NetworkModel]]: """获取网段列表""" query = db.query(NetworkModel) if group_name: query = query.filter(NetworkModel.group_name == group_name) total = query.count() items = query.order_by(NetworkModel.id.desc()).offset(skip).limit(limit).all() return total, items @staticmethod def create(db: Session, network_in: NetworkCreate) -> NetworkModel: """创建新网段""" total_ips = NetworkService.calculate_total_ips(network_in.cidr) db_network = NetworkModel( cidr=network_in.cidr, name=network_in.name, description=network_in.description, group_name=network_in.group_name, gateway=network_in.gateway, vlan_id=network_in.vlan_id, total_ips=total_ips, used_ips=0, reserved_ips=0 ) db.add(db_network) db.flush() # 自动创建该网段下的所有IP记录 NetworkService._create_ip_addresses(db, db_network) db.commit() db.refresh(db_network) return db_network @staticmethod def _create_ip_addresses(db: Session, network: NetworkModel): """为网段创建所有IP记录""" ip_addresses = NetworkService.get_network_addresses(network.cidr) ip_records = [] for ip in ip_addresses: # 标记网关为保留状态 status = IPStatus.RESERVED if ip == network.gateway else IPStatus.AVAILABLE ip_records.append( IPAddressModel( network_id=network.id, ip_address=ip, status=status ) ) db.bulk_save_objects(ip_records) # 更新统计 reserved_count = sum(1 for ip in ip_addresses if ip == network.gateway) network.reserved_ips = reserved_count @staticmethod def update(db: Session, network_id: int, network_in: NetworkUpdate) -> Optional[NetworkModel]: """更新网段信息""" db_network = NetworkService.get_by_id(db, network_id) if not db_network: return None update_data = network_in.model_dump(exclude_unset=True) for field, value in update_data.items(): setattr(db_network, field, value) db.commit() db.refresh(db_network) return db_network @staticmethod def delete(db: Session, network_id: int) -> bool: """删除网段""" db_network = NetworkService.get_by_id(db, network_id) if not db_network: return False db.delete(db_network) db.commit() return True @staticmethod def get_stats(db: Session, network_id: int) -> Optional[NetworkStats]: """获取网段统计信息""" db_network = NetworkService.get_by_id(db, network_id) if not db_network: return None # 实时计算IP状态 - 使用SUM(CASE)兼容MySQL stats = db.query( func.sum(case((IPAddressModel.status == IPStatus.ONLINE, 1), else_=0)).label('online'), func.sum(case((IPAddressModel.status == IPStatus.OFFLINE, 1), else_=0)).label('offline'), func.sum(case((IPAddressModel.status == IPStatus.RESERVED, 1), else_=0)).label('reserved'), ).filter(IPAddressModel.network_id == network_id).first() used_ips = int(stats.online or 0) + int(stats.offline or 0) reserved_ips = int(stats.reserved or 0) available_ips = db_network.total_ips - used_ips - reserved_ips usage_percent = (used_ips / db_network.total_ips * 100) if db_network.total_ips > 0 else 0 # 更新数据库统计 db_network.used_ips = used_ips db_network.reserved_ips = reserved_ips db.commit() return NetworkStats( id=db_network.id, cidr=db_network.cidr, name=db_network.name, total_ips=db_network.total_ips, used_ips=used_ips, reserved_ips=reserved_ips, available_ips=available_ips, usage_percent=round(usage_percent, 2), group_name=db_network.group_name, vlan_id=db_network.vlan_id ) @staticmethod def get_all_stats(db: Session) -> List[NetworkStats]: """获取所有网段的统计信息""" networks = db.query(NetworkModel).all() stats_list = [] for network in networks: stats = NetworkService.get_stats(db, network.id) if stats: stats_list.append(stats) return stats_list