Files
ipam/backend/app/services/network_service.py
T
Your Name 13c457acb3 fix(network): 修复网段 total_ips 包含网络地址和广播地址的问题
问题原因:
- calculate_total_ips 使用 network.num_addresses 计算总IP数量
- 该值包含网络地址(第一个IP)和广播地址(最后一个IP)
- 而 _create_ip_addresses 使用 hosts() 生成IP列表,排除这两个地址
- 导致 total_ips 比实际存储的IP记录多2个
- 例: /24 显示 256 个IP,实际只有 254 条记录

解决方案:
1. calculate_total_ips 改为 num_addresses - 2(排除网络地址和广播地址)
2. 对于 /31 /32 等没有网络/广播地址的网段,特殊处理
3. 提供 fix_network_total_ips.py 迁移脚本修复已有数据
2026-07-23 15:33:13 +08:00

178 lines
6.0 KiB
Python

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 - 2 if network.num_addresses > 2 else 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