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, 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, 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): """ 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 ) 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) 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() # 返回任务函数供外部使用 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 }