from typing import Optional, List, Any from sqlalchemy.orm import Session from datetime import datetime, timedelta import json from app.models.audit import AuditLog, AuditAction, AuditResource from app.models.auth import User # ===== 兼容层:老 API(字符串 action/resource)映射到新枚举 ===== # 字符串 action -> 枚举 + 资源类型;用于兼容各 endpoint 里硬编码的 "network.create" 等 _ACTION_MAP = { # network "network.create": (AuditAction.CREATE, AuditResource.NETWORK), "network.update": (AuditAction.UPDATE, AuditResource.NETWORK), "network.delete": (AuditAction.DELETE, AuditResource.NETWORK), # ip "ip.update": (AuditAction.UPDATE, AuditResource.IP_ADDRESS), "ip.reserve": (AuditAction.UPDATE, AuditResource.IP_ADDRESS), "ip.unreserve": (AuditAction.UPDATE, AuditResource.IP_ADDRESS), # scan "scan.trigger": (AuditAction.SCAN, AuditResource.SCAN_TASK), # snmp "snmp.credential.create": (AuditAction.CREATE, AuditResource.SNMP_CREDENTIAL), "snmp.credential.update": (AuditAction.UPDATE, AuditResource.SNMP_CREDENTIAL), "snmp.credential.delete": (AuditAction.DELETE, AuditResource.SNMP_CREDENTIAL), "snmp.device.create": (AuditAction.CREATE, AuditResource.SNMP_DEVICE), "snmp.device.update": (AuditAction.UPDATE, AuditResource.SNMP_DEVICE), "snmp.device.delete": (AuditAction.DELETE, AuditResource.SNMP_DEVICE), # alert "alert.acknowledge": (AuditAction.UPDATE, AuditResource.ALERT), "alert.resolve": (AuditAction.UPDATE, AuditResource.ALERT), # user "user.login": (AuditAction.LOGIN, AuditResource.USER), "user.login.failed": (AuditAction.LOGIN, AuditResource.USER), "user.logout": (AuditAction.LOGOUT, AuditResource.USER), "user.create": (AuditAction.CREATE, AuditResource.USER), "user.update_status": (AuditAction.UPDATE, AuditResource.USER), "user.unlock": (AuditAction.UPDATE, AuditResource.USER), "user.reset_password": (AuditAction.UPDATE, AuditResource.USER), "user.delete": (AuditAction.DELETE, AuditResource.USER), "user.change_password": (AuditAction.UPDATE, AuditResource.USER), } _RESOURCE_MAP = { "network": AuditResource.NETWORK, "ip": AuditResource.IP_ADDRESS, "snmp_credential": AuditResource.SNMP_CREDENTIAL, "snmp_device": AuditResource.SNMP_DEVICE, "user": AuditResource.USER, "alert": AuditResource.ALERT, "scan": AuditResource.SCAN_TASK, } class AuditService: """审计日志服务""" @staticmethod def record( db: Session, *, action: str, user=None, resource_type: Optional[str] = None, resource_id: Optional[Any] = None, resource_name: Optional[str] = None, method: Optional[str] = None, path: Optional[str] = None, ip_address: Optional[str] = None, user_agent: Optional[str] = None, status: str = "success", detail: Optional[Any] = None, ) -> Optional["AuditLog"]: """ 兼容层:将老版 API(字符串 action/resource)映射到新版 log()。 所有 endpoint 都通过本方法写审计,避免直接调用新签名造成 500。 任意环节失败都安全降级(不抛异常),审计写入失败不影响主业务。 """ try: # action 字符串 -> (枚举, 资源枚举) ;未知 action 用 OTHER + SYSTEM enum_action, default_resource = _ACTION_MAP.get( action, (AuditAction.OTHER, AuditResource.SYSTEM) ) resource_enum = _RESOURCE_MAP.get(resource_type, default_resource) # 详情拆分:dict 拆出 old/new;其余当 description old_value = None new_value = None description = None error_message = None success = status != "failed" if isinstance(detail, dict): old_value = detail.get("before") new_value = detail.get("after") # 把 action 原始字符串 + 资源名放到 description 里,保留细节 desc_parts = [action] if resource_name: desc_parts.append(resource_name) if status == "failed": error_message = detail.get("error") or detail.get("message") or json.dumps(detail, ensure_ascii=False, default=str) description = " | ".join(desc_parts) else: description = f"{action} | {resource_name or ''}" if status == "failed" and detail: error_message = str(detail) # user 不能 None 传给 log();fallback 到 User(id=None, username='anonymous') # 但 log() 会用 user.id / user.username / user.real_name,如果都是 None 就 OK if user is None: anon = User(id=None, username="anonymous", real_name=None) else: anon = user return AuditService.log( db=db, user=anon, action=enum_action, resource=resource_enum, resource_id=str(resource_id) if resource_id is not None else None, description=description[:500] if description else None, old_value=old_value, new_value=new_value, user_ip=ip_address, user_agent=(user_agent or "")[:500], request_method=method, request_path=path, success=success, error_message=error_message, ) except Exception as e: db.rollback() print(f"[audit] record() failed action={action}: {e}") return None @staticmethod def log(db: Session, user: User, action: AuditAction, resource: AuditResource, resource_id: Optional[str] = None, description: Optional[str] = None, old_value: Optional[dict] = None, new_value: Optional[dict] = None, changed_fields: Optional[List[str]] = None, user_ip: Optional[str] = None, user_agent: Optional[str] = None, request_method: Optional[str] = None, request_path: Optional[str] = None, success: bool = True, error_message: Optional[str] = None) -> AuditLog: """ 记录审计日志 Args: db: 数据库会话 user: 操作用户 action: 操作类型 resource: 资源类型 resource_id: 资源ID description: 操作描述 old_value: 旧值 new_value: 新值 changed_fields: 变更的字段列表 user_ip: 用户IP user_agent: 用户代理 request_method: 请求方法 request_path: 请求路径 success: 操作是否成功 error_message: 错误信息 Returns: AuditLog 实例 """ audit_log = AuditLog( user_id=user.id if user else None, username=user.username if user else None, real_name=user.real_name if user else None, action=action, resource=resource, resource_id=str(resource_id) if resource_id else None, description=description, old_value=json.dumps(old_value, ensure_ascii=False, default=str) if old_value else None, new_value=json.dumps(new_value, ensure_ascii=False, default=str) if new_value else None, changed_fields=','.join(changed_fields) if changed_fields else None, user_ip=user_ip, user_agent=user_agent, request_method=request_method, request_path=request_path, success=1 if success else 0, error_message=error_message ) db.add(audit_log) db.commit() db.refresh(audit_log) return audit_log @staticmethod def get_logs(db: Session, user_id: Optional[int] = None, action: Optional[AuditAction] = None, resource: Optional[AuditResource] = None, success: Optional[bool] = None, start_time: Optional[datetime] = None, end_time: Optional[datetime] = None, skip: int = 0, limit: int = 100) -> tuple[int, List[AuditLog]]: """ 查询审计日志 Args: db: 数据库会话 user_id: 按用户ID过滤 action: 按操作类型过滤 resource: 按资源类型过滤 success: 按操作结果过滤 start_time: 开始时间 end_time: 结束时间 skip: 跳过记录数 limit: 返回记录数 Returns: (总数, 记录列表) """ query = db.query(AuditLog) if user_id: query = query.filter(AuditLog.user_id == user_id) if action: query = query.filter(AuditLog.action == action) if resource: query = query.filter(AuditLog.resource == resource) if success is not None: query = query.filter(AuditLog.success == (1 if success else 0)) if start_time: query = query.filter(AuditLog.created_at >= start_time) if end_time: query = query.filter(AuditLog.created_at <= end_time) total = query.count() items = query.order_by(AuditLog.created_at.desc()).offset(skip).limit(limit).all() return total, items @staticmethod def get_user_actions(db: Session, user_id: int, limit: int = 50) -> List[AuditLog]: """获取指定用户的操作日志""" return db.query(AuditLog)\ .filter(AuditLog.user_id == user_id)\ .order_by(AuditLog.created_at.desc())\ .limit(limit)\ .all() @staticmethod def get_resource_logs(db: Session, resource: AuditResource, resource_id: str, limit: int = 50) -> List[AuditLog]: """获取指定资源的操作历史""" return db.query(AuditLog)\ .filter(AuditLog.resource == resource, AuditLog.resource_id == str(resource_id))\ .order_by(AuditLog.created_at.desc())\ .limit(limit)\ .all() @staticmethod def get_statistics(db: Session, days: int = 30) -> dict: """ 获取审计统计信息 Args: db: 数据库会话 days: 统计天数 Returns: 统计信息字典 """ start_date = datetime.utcnow() - timedelta(days=days) # 总操作数 total_ops = db.query(AuditLog).filter(AuditLog.created_at >= start_date).count() # 成功/失败统计 success_ops = db.query(AuditLog).filter( AuditLog.created_at >= start_date, AuditLog.success == 1 ).count() failed_ops = total_ops - success_ops # 按操作类型统计 action_stats = {} for action in AuditAction: count = db.query(AuditLog).filter( AuditLog.created_at >= start_date, AuditLog.action == action ).count() if count > 0: action_stats[action.value] = count # 按资源类型统计 resource_stats = {} for resource in AuditResource: count = db.query(AuditLog).filter( AuditLog.created_at >= start_date, AuditLog.resource == resource ).count() if count > 0: resource_stats[resource.value] = count # 活跃用户数 active_users = db.query(AuditLog.user_id).filter( AuditLog.created_at >= start_date ).distinct().count() return { "period_days": days, "total_operations": total_ops, "successful_operations": success_ops, "failed_operations": failed_ops, "success_rate": round(success_ops / total_ops * 100, 2) if total_ops > 0 else 0, "by_action": action_stats, "by_resource": resource_stats, "active_users_count": active_users } @staticmethod def clean_old_logs(db: Session, days: int = 90) -> int: """ 清理指定天数之前的日志 Args: db: 数据库会话 days: 保留天数 Returns: 删除的记录数 """ cutoff_date = datetime.utcnow() - timedelta(days=days) deleted = db.query(AuditLog).filter(AuditLog.created_at < cutoff_date).delete() db.commit() return deleted class AuditLogger: """审计日志装饰器/上下文管理器""" def __init__(self, db: Session, user: User, resource: AuditResource, resource_id: Optional[str] = None): self.db = db self.user = user self.resource = resource self.resource_id = resource_id self.old_value = None def set_old_value(self, value: dict): """设置旧值(用于更新操作)""" self.old_value = value def log(self, action: AuditAction, description: str, new_value: Optional[dict] = None, changed_fields: Optional[List[str]] = None): """记录日志""" AuditService.log( db=self.db, user=self.user, action=action, resource=self.resource, resource_id=self.resource_id, description=description, old_value=self.old_value, new_value=new_value, changed_fields=changed_fields )