Files
ipam/backend/app/tasks/scan_tasks.py
T
Your Name 8775c9836d feat(snmp): 实现网络设备自动轮询定时任务
之前只有 arp/mac/interface_poll_interval 字段但没有实际调度,
SNMP 设备添加后永远不会被自动采集。

新增 poll_snmp_devices Celery 定时任务:
- 遍历所有活跃且配置了 SNMP 凭据的设备
- 按各设备的 poll_interval 判断是否到期,到期才轮询
- 用 asyncio.run 执行异步的 SNMPService.poll_device
- Beat 每 60 秒触发检查,新设备 last_polled_at 为空会立即轮询
2026-08-12 10:56:40 +08:00

287 lines
10 KiB
Python

from typing import List, Dict, Any, Optional
from celery import Task
import logging
from app.core.database import SessionLocal
from app.services.network_service import NetworkService
from app.services.enhanced_scan_service import EnhancedScanService
from app.services.scan_service import ScanService
from app.models.network import TaskStatus, ScanTask, IPAddress
logger = logging.getLogger(__name__)
def register_tasks(celery_app):
"""
注册任务到 Celery 应用
"""
@celery_app.task(bind=True)
def scan_network_task(self, network_id: int,
enable_ping: bool = True,
enable_arp: bool = True,
enable_dns: bool = True,
enable_netbios: bool = True,
timeout: int = 2):
"""
Celery 异步任务:扫描指定网段
"""
task_id = self.request.id
db = SessionLocal()
try:
# 获取网段信息
network = NetworkService.get_by_id(db, network_id)
if not network:
logger.error(f"网段 {network_id} 不存在")
return {'status': 'failed', 'error': '网段不存在'}
logger.info(f"开始扫描网段: {network.cidr}")
# 更新任务状态为运行中
# 查找已存在的任务记录
db_task = db.query(ScanTask).filter(ScanTask.celery_task_id == task_id).first()
if not db_task:
db_task = ScanService.create_scan_task(db, 'ping', network_id)
db_task.celery_task_id = task_id
db.commit()
db.refresh(db_task)
ScanService.update_task_status(db, db_task.id, TaskStatus.RUNNING, progress=0)
# 执行扫描
results = EnhancedScanService.scan_network(
network.cidr,
enable_ping=enable_ping,
enable_arp=enable_arp,
enable_dns=enable_dns,
enable_netbios=enable_netbios,
timeout=timeout
)
# 批量更新IP信息
stats = EnhancedScanService.bulk_update_ips_from_scan(db, results)
# 更新任务状态
ScanService.update_task_status(
db,
db_task.id,
TaskStatus.COMPLETED,
progress=100,
total_count=stats['total_scanned'],
success_count=stats['online_count']
)
logger.info(f"网段 {network.cidr} 扫描完成: 在线 {stats['online_count']} 台, "
f"发现MAC {stats['mac_found_count']} 个, 主机名 {stats['hostname_found_count']}")
return {
'status': 'completed',
'network_id': network_id,
'cidr': network.cidr,
'statistics': stats
}
except Exception as e:
logger.error(f"扫描失败: {str(e)}", exc_info=True)
if 'db_task' in locals() and db_task:
ScanService.update_task_status(
db,
db_task.id,
TaskStatus.FAILED,
error_message=str(e)
)
raise
finally:
db.close()
@celery_app.task(bind=True)
def scan_single_ip_task(self, ip_address: str,
enable_ping: bool = True,
enable_arp: bool = True,
enable_dns: bool = True):
"""
Celery 异步任务:扫描单个IP
"""
db = SessionLocal()
try:
result = EnhancedScanService.comprehensive_scan(
ip_address,
enable_ping=enable_ping,
enable_arp=enable_arp,
enable_dns=enable_dns
)
# 更新IP信息
EnhancedScanService.update_ip_from_scan_result(db, result)
db.commit()
return result
finally:
db.close()
@celery_app.task(bind=True)
def full_network_scan(self, enable_ping: bool = True,
enable_arp: bool = True,
enable_dns: bool = True,
enable_netbios: bool = True):
"""
Celery 定时任务:扫描所有网段
"""
db = SessionLocal()
try:
_, all_networks = NetworkService.get_list(db, limit=1000)
total_stats = {
'total_networks': len(all_networks),
'total_ips': 0,
'online_count': 0,
'mac_found_count': 0,
'hostname_found_count': 0
}
logger.info(f"开始全量扫描: {len(all_networks)} 个网段")
for network in all_networks:
try:
results = EnhancedScanService.scan_network(
network.cidr,
enable_ping=enable_ping,
enable_arp=enable_arp,
enable_dns=enable_dns,
enable_netbios=enable_netbios
)
stats = EnhancedScanService.bulk_update_ips_from_scan(db, results)
total_stats['total_ips'] += stats['total_scanned']
total_stats['online_count'] += stats['online_count']
total_stats['mac_found_count'] += stats['mac_found_count']
total_stats['hostname_found_count'] += stats['hostname_found_count']
logger.info(f"网段 {network.cidr} 扫描完成: 在线 {stats['online_count']}")
except Exception as e:
logger.error(f"网段 {network.cidr} 扫描失败: {str(e)}")
logger.info(f"全量扫描完成: {total_stats}")
return {
'status': 'completed',
'statistics': total_stats
}
finally:
db.close()
@celery_app.task(bind=True)
def update_all_statistics(self):
"""
定时任务:更新所有网段统计信息
"""
db = SessionLocal()
try:
NetworkService.get_all_stats(db)
logger.info("所有网段统计信息已更新")
return {'status': 'completed'}
finally:
db.close()
@celery_app.task(bind=True)
def quick_status_update(self, ip_addresses: List[str]):
"""
快速更新多个IP的状态(用于实时扫描)
"""
db = SessionLocal()
try:
results = []
for ip in ip_addresses:
result = EnhancedScanService.comprehensive_scan(
ip, enable_ping=True, enable_arp=True, enable_dns=False,
enable_netbios=True
)
EnhancedScanService.update_ip_from_scan_result(db, result)
results.append(result)
db.commit()
return {
'scanned_count': len(ip_addresses),
'online_count': sum(1 for r in results if r.get('success')),
'results': results
}
finally:
db.close()
@celery_app.task(bind=True)
def poll_snmp_devices(self):
"""
定时任务:自动轮询所有活跃且配置了 SNMP 凭据的网络设备。
每个设备按自己的 arp/mac/interface_poll_interval(秒)判断是否到达
轮询间隔,到期才执行 SNMPService.poll_device(采集 ARP 表、MAC 地址
表、接口信息)。由 Celery Beat 每 60 秒触发一次检查。
"""
import asyncio
from datetime import datetime
from app.models.snmp import NetworkDevice
from app.services.snmp_service import SNMPService
db = SessionLocal()
try:
devices = db.query(NetworkDevice).filter(
NetworkDevice.is_active == True,
NetworkDevice.snmp_credential_id.isnot(None),
).all()
now = datetime.utcnow()
def _as_naive(dt):
"""把 aware/naive datetime 统一转 naive(本地比较用)"""
if dt is None:
return None
if dt.tzinfo is not None:
return dt.replace(tzinfo=None)
return dt
polled, skipped = 0, 0
for dev in devices:
# 取三种采集间隔的最小值作为该设备的轮询周期
intervals = [
dev.arp_poll_interval,
dev.mac_poll_interval,
dev.interface_poll_interval,
]
interval = min([i for i in intervals if i], default=300)
last = _as_naive(dev.last_polled_at)
if last is not None:
elapsed = (now - last).total_seconds()
if elapsed < interval:
skipped += 1
continue
try:
# poll_device 是 async 函数,用 asyncio.run 在同步任务中执行
asyncio.run(SNMPService.poll_device(db, dev.id))
polled += 1
logger.info(f"SNMP 自动轮询完成: {dev.name} ({dev.ip_address})")
except Exception as e:
logger.error(f"SNMP 自动轮询失败 {dev.name} ({dev.ip_address}): {e}", exc_info=True)
logger.info(f"SNMP 自动轮询检查完成: 设备总数 {len(devices)}, 本轮轮询 {polled}, 未到期跳过 {skipped}")
return {'status': 'completed', 'total': len(devices), 'polled': polled, 'skipped': skipped}
finally:
db.close()
# 返回任务函数供外部使用
return {
'scan_network_task': scan_network_task,
'scan_single_ip_task': scan_single_ip_task,
'full_network_scan': full_network_scan,
'update_all_statistics': update_all_statistics,
'quick_status_update': quick_status_update,
'poll_snmp_devices': poll_snmp_devices,
}