"""B7审计日志和统计查询;本模块不写入审计表。""" from __future__ import annotations from datetime import UTC, datetime, time, timedelta from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from flask import current_app from sqlalchemy import case, func, or_, select from werkzeug.datastructures import MultiDict from dms.common.enums import AuditAction, AuditResult, AuditTarget from dms.common.errors import InternalError, InvalidArgumentError from dms.common.pagination import PageRequest, page_result from dms.common.response import serialize_id from dms.extensions import db from dms.models import AuditLog, Document AUDIT_LIST_PARAMETERS = { "keyword", "createdFrom", "createdTo", "userId", "organizationId", "actionType", "operationResult", "targetType", "targetId", "page", "pageSize", "sortDirection", } TIME_PARAMETERS = {"createdFrom", "createdTo"} def _validate_parameters(params: MultiDict, allowed: set[str]) -> None: unknown = set(params) - allowed if unknown: raise InvalidArgumentError(f"不支持的查询参数:{sorted(unknown)[0]}") duplicate = next( (key for key in params if len(params.getlist(key)) != 1), None, ) if duplicate is not None: raise InvalidArgumentError(f"查询参数{duplicate}不能重复") def _positive_integer( params: MultiDict, name: str, *, default: int | None = None, maximum: int | None = None, ) -> int | None: raw = params.get(name) if raw is None: return default if not raw.isdecimal() or int(raw) < 1: raise InvalidArgumentError(f"{name}必须是正整数字符串") value = int(raw) if maximum is not None and value > maximum: raise InvalidArgumentError(f"{name}不得大于{maximum}") return value def _enum_value(params: MultiDict, name: str, enum_type) -> str | None: raw = params.get(name) if raw is None: return None try: return enum_type(raw).value except ValueError as exc: raise InvalidArgumentError(f"{name}不是有效枚举值") from exc def _utc_parameter(params: MultiDict, name: str) -> datetime | None: raw = params.get(name) if raw is None: return None if not raw or not raw.endswith("Z"): raise InvalidArgumentError(f"{name}必须是带Z的ISO 8601 UTC时间") try: parsed = datetime.fromisoformat(f"{raw[:-1]}+00:00") except ValueError as exc: raise InvalidArgumentError(f"{name}必须是带Z的ISO 8601 UTC时间") from exc if parsed.tzinfo is None or parsed.utcoffset() != timedelta(0): raise InvalidArgumentError(f"{name}必须是带Z的ISO 8601 UTC时间") return parsed.astimezone(UTC).replace(tzinfo=None) def _time_filters(params: MultiDict) -> list: created_from = _utc_parameter(params, "createdFrom") created_to = _utc_parameter(params, "createdTo") if ( created_from is not None and created_to is not None and created_from >= created_to ): raise InvalidArgumentError("createdFrom必须早于createdTo") filters = [] if created_from is not None: filters.append(AuditLog.created_at >= created_from) if created_to is not None: filters.append(AuditLog.created_at < created_to) return filters def _format_utc(value: datetime) -> str: return ( value.replace(tzinfo=UTC) .isoformat(timespec="milliseconds") .replace("+00:00", "Z") ) def _escape_like(value: str) -> str: return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") def _audit_summary(item: AuditLog) -> dict: return { "id": serialize_id(item.id), "username": item.username, "realName": item.real_name, "organizationName": item.organization_name, "actionType": item.action_type, "targetType": item.target_type, "targetId": serialize_id(item.target_id), "targetName": item.target_name, "operationResult": item.operation_result, "failureReason": item.failure_reason, "clientIp": item.client_ip, "requestId": item.request_id, "createdAt": _format_utc(item.created_at), } def audit_log_page(params: MultiDict) -> dict: _validate_parameters(params, AUDIT_LIST_PARAMETERS) page = _positive_integer(params, "page", default=1) page_size = _positive_integer( params, "pageSize", default=20, maximum=100 ) assert page is not None and page_size is not None page_request = PageRequest(page=page, page_size=page_size) sort_direction = params.get("sortDirection", "desc") if sort_direction not in {"asc", "desc"}: raise InvalidArgumentError("sortDirection只允许asc或desc") filters = _time_filters(params) for parameter, column in ( ("userId", AuditLog.user_id), ("organizationId", AuditLog.organization_id), ("targetId", AuditLog.target_id), ): value = _positive_integer(params, parameter) if value is not None: filters.append(column == value) for parameter, enum_type, column in ( ("actionType", AuditAction, AuditLog.action_type), ("operationResult", AuditResult, AuditLog.operation_result), ("targetType", AuditTarget, AuditLog.target_type), ): value = _enum_value(params, parameter, enum_type) if value is not None: filters.append(column == value) keyword = (params.get("keyword") or "").strip() if keyword: pattern = f"%{_escape_like(keyword)}%" filters.append( or_( *[ column.like(pattern, escape="\\") for column in ( AuditLog.username, AuditLog.real_name, AuditLog.organization_name, AuditLog.target_name, AuditLog.failure_reason, AuditLog.request_id, ) ] ) ) total = db.session.scalar( select(func.count(AuditLog.id)).where(*filters) ) or 0 direction = ( (AuditLog.created_at.asc(), AuditLog.id.asc()) if sort_direction == "asc" else (AuditLog.created_at.desc(), AuditLog.id.desc()) ) items = db.session.scalars( select(AuditLog) .where(*filters) .order_by(*direction) .offset(page_request.offset) .limit(page_request.page_size) ).all() return page_result( [_audit_summary(item) for item in items], page=page_request.page, page_size=page_request.page_size, total=int(total), ) def _utc_now() -> datetime: return datetime.now(UTC) def _business_day_utc_range() -> tuple[datetime, datetime]: timezone_name = current_app.config["DMS_BUSINESS_TIMEZONE"] try: business_timezone = ZoneInfo(timezone_name) except ZoneInfoNotFoundError as exc: raise InternalError("业务时区配置无效") from exc local_now = _utc_now().astimezone(business_timezone) start_local = datetime.combine( local_now.date(), time.min, tzinfo=business_timezone ) end_local = start_local + timedelta(days=1) return ( start_local.astimezone(UTC).replace(tzinfo=None), end_local.astimezone(UTC).replace(tzinfo=None), ) def document_statistics(params: MultiDict) -> dict: _validate_parameters(params, set()) counts = db.session.execute( select( func.coalesce( func.sum(case((Document.document_type == "MAIN", 1), else_=0)), 0, ), func.coalesce( func.sum( case((Document.document_type == "SUB_PLAN", 1), else_=0) ), 0, ), func.coalesce( func.sum( case((Document.document_type == "ATTACHMENT", 1), else_=0) ), 0, ), ).where(Document.is_deleted.is_(False)) ).one() start_utc, end_utc = _business_day_utc_range() today_viewers = db.session.scalar( select(func.count(func.distinct(AuditLog.user_id))).where( AuditLog.action_type == AuditAction.VIEW_DOCUMENT.value, AuditLog.operation_result == AuditResult.SUCCESS.value, AuditLog.user_id.is_not(None), AuditLog.created_at >= start_utc, AuditLog.created_at < end_utc, ) ) or 0 main_count, sub_count, attachment_count = map(int, counts) return { "planDocumentCount": main_count + sub_count, "mainPlanCount": main_count, "subPlanCount": sub_count, "attachmentCount": attachment_count, "todayViewerCount": int(today_viewers), } def _period_expression(granularity: str): if granularity == "DAY": return func.date_format(AuditLog.created_at, "%Y-%m-%d") if granularity == "MONTH": return func.date_format(AuditLog.created_at, "%Y-%m") monday = func.from_days( func.to_days(AuditLog.created_at) - func.weekday(AuditLog.created_at) ) return func.date_format(monday, "%Y-%m-%d") def trend_statistics(params: MultiDict) -> dict: _validate_parameters(params, TIME_PARAMETERS | {"granularity"}) granularity = params.get("granularity") if granularity not in {"DAY", "WEEK", "MONTH"}: raise InvalidArgumentError("granularity只允许DAY、WEEK或MONTH且为必填") filters = _time_filters(params) period = _period_expression(granularity).label("period") operation_count = func.count(AuditLog.id).label("operation_count") rows = db.session.execute( select(period, operation_count) .where(*filters) .group_by(period) .order_by(period.asc()) ).all() return { "items": [ {"period": row.period, "operationCount": int(row.operation_count)} for row in rows ] } def action_statistics(params: MultiDict) -> dict: _validate_parameters(params, TIME_PARAMETERS) filters = _time_filters(params) count = func.count(AuditLog.id).label("action_count") rows = db.session.execute( select(AuditLog.action_type, count) .where(*filters) .group_by(AuditLog.action_type) .order_by(count.desc(), AuditLog.action_type.asc()) ).all() return { "items": [ {"actionType": row.action_type, "count": int(row.action_count)} for row in rows ] } def active_user_statistics(params: MultiDict) -> dict: _validate_parameters(params, TIME_PARAMETERS | {"limit"}) limit = _positive_integer(params, "limit", default=10, maximum=100) assert limit is not None filters = _time_filters(params) ranked = ( select( AuditLog.user_id.label("user_id"), AuditLog.real_name.label("real_name"), AuditLog.organization_name.label("organization_name"), func.count() .over(partition_by=AuditLog.user_id) .label("operation_count"), func.row_number() .over( partition_by=AuditLog.user_id, order_by=(AuditLog.created_at.desc(), AuditLog.id.desc()), ) .label("snapshot_rank"), ) .where(AuditLog.user_id.is_not(None), *filters) .subquery() ) rows = db.session.execute( select( ranked.c.user_id, ranked.c.real_name, ranked.c.organization_name, ranked.c.operation_count, ) .where(ranked.c.snapshot_rank == 1) .order_by(ranked.c.operation_count.desc(), ranked.c.user_id.asc()) .limit(limit) ).all() return { "items": [ { "userId": serialize_id(row.user_id), "realName": row.real_name, "organizationName": row.organization_name, "operationCount": int(row.operation_count), } for row in rows ] }