"""主案、子方案只读查询、DTO和查看审计。""" from __future__ import annotations import logging from datetime import datetime, timezone from typing import Any, Mapping from sqlalchemy import String, cast, or_, select from dms.common.enums import ( AuditAction, AuditTarget, DocumentStatus, DocumentType, RoleCode, SecurityLevel, VisibilityType, ) from dms.common.errors import ( DocumentViewForbiddenError, InvalidArgumentError, MainPlanNotFoundError, ResourceNotFoundError, SecurityLevelForbiddenError, ) from dms.common.pagination import PageRequest, page_result from dms.common.response import serialize_id from dms.common.time_range import parse_datetime, parse_updated_range from dms.database.transaction import transaction from dms.extensions import db from dms.models import Document, Permission from dms.security.auth_context import AuthContext, get_auth_context from dms.services.audit_service import business_audit from dms.services.authorization_service import ( PlanAccess, evaluate_plan_access, plan_allowed_actions, ) logger = logging.getLogger(__name__) PLAN_TYPES = (DocumentType.MAIN.value, DocumentType.SUB_PLAN.value) DOCUMENT_SORTS = { "documentName": Document.document_name, "createdAt": Document.created_at, "updatedAt": Document.updated_at, "viewCount": Document.view_count, "downloadCount": Document.download_count, } def parse_string_id(value: str, *, field: str = "id") -> int: if not value.isdecimal() or int(value) <= 0: raise InvalidArgumentError(f"{field}必须是正整数形式的字符串ID") return int(value) def _integer( params: Mapping[str, str], name: str, default: int, ) -> int: raw = params.get(name) if raw is None: return default if not raw.isdecimal(): raise InvalidArgumentError(f"{name}必须是正整数") return int(raw) def _page(params: Mapping[str, str]) -> PageRequest: return PageRequest( page=_integer(params, "page", 1), page_size=_integer(params, "pageSize", 20), ) def _enum(value: str | None, enum_type, name: str) -> str | None: if value is None or value == "": return None try: return enum_type(value).value except ValueError as exc: raise InvalidArgumentError(f"{name}不是有效枚举值") from exc def _date(value: str | None, name: str) -> datetime | None: return parse_datetime(value, name=name) def _sort( params: Mapping[str, str], *, allowed: dict[str, Any] = DOCUMENT_SORTS, default: str = "updatedAt", ) -> tuple[Any, str]: sort_by = params.get("sortBy", default) direction = params.get("sortDirection", "desc").lower() if sort_by not in allowed: raise InvalidArgumentError("sortBy不是允许的排序字段") if direction not in {"asc", "desc"}: raise InvalidArgumentError("sortDirection必须是asc或desc") column = allowed[sort_by] return (column.asc() if direction == "asc" else column.desc()), direction def _iso(value: datetime) -> str: return value.replace(tzinfo=timezone.utc).isoformat(timespec="milliseconds").replace( "+00:00", "Z" ) def _permission_summary( access: PlanAccess, document: Document, ) -> tuple[dict[str, object], str]: source = access.source if source is None: return { "organizationCount": 0, "userCount": 0, "inheritedFromMainPlan": False, }, "未配置" permissions = db.session.scalars( select(Permission).where( Permission.document_id == source.id, Permission.is_deleted.is_(False), ) ).all() organization_count = sum(item.subject_type == "ORG" for item in permissions) user_count = sum(item.subject_type == "USER" for item in permissions) visibility = VisibilityType(source.visibility_type) if visibility == VisibilityType.ALL_AUTHENTICATED: visibility_summary = "全部已登录用户" else: names = [item.subject_name for item in permissions if item.can_view] visibility_summary = "、".join(names) if names else "未配置" return { "organizationCount": organization_count, "userCount": user_count, "inheritedFromMainPlan": document.document_type == DocumentType.SUB_PLAN.value, }, visibility_summary def document_summary( document: Document, context: AuthContext, access: PlanAccess, ) -> dict[str, object]: _, visibility_summary = _permission_summary(access, document) source = access.source or document return { "id": serialize_id(document.id), "documentName": document.document_name, "summary": document.summary, "documentType": document.document_type, "status": document.document_status, "securityLevel": document.security_level, "visibilityType": source.visibility_type, "visibilitySummary": visibility_summary, "attachmentType": document.attachment_type, "categoryId": serialize_id(document.category_id), "categoryName": document.category_name, "categoryPath": document.category_path, "parentDocumentId": serialize_id(document.parent_document_id), "rootDocumentId": serialize_id( document.id if document.document_type == DocumentType.MAIN.value else document.root_document_id ), "tags": document.tags or [], "fileExtension": document.file_extension, "childCount": document.child_count, "attachmentCount": document.attachment_count, "viewCount": document.view_count, "downloadCount": document.download_count, "createdByName": document.created_by_name, "createdAt": _iso(document.created_at), "updatedAt": _iso(document.updated_at), "rowVersion": document.row_version, "allowedActions": plan_allowed_actions(document, context, access), } def document_detail( document: Document, context: AuthContext, access: PlanAccess, ) -> dict[str, object]: result = document_summary(document, context, access) permission_summary, _ = _permission_summary(access, document) result.update( { "originalFileName": document.original_file_name, "mimeType": document.mime_type, "fileSize": document.file_size, "fileHash": document.file_hash, "permissionSummary": permission_summary, } ) return result def _require_plan_access(document: Document, context: AuthContext) -> PlanAccess: access = evaluate_plan_access(document, context) if access.allowed: return access if access.reason == "SECURITY": raise SecurityLevelForbiddenError() raise DocumentViewForbiddenError() def _keyword(statement, keyword: str | None): if not keyword or not keyword.strip(): return statement pattern = f"%{keyword.strip()}%" return statement.where( or_( Document.document_name.like(pattern), Document.summary.like(pattern), Document.search_text.like(pattern), cast(Document.tags, String).like(pattern), ) ) def list_documents(params: Mapping[str, str]) -> dict[str, object]: context = get_auth_context() if context.role_code == RoleCode.AUDITOR: raise DocumentViewForbiddenError("审计员默认无方案浏览权限") page = _page(params) raw_types = params.get("documentType", "MAIN,SUB_PLAN") document_types = [item.strip() for item in raw_types.split(",") if item.strip()] if ( not document_types or any(item not in PLAN_TYPES for item in document_types) or len(set(document_types)) != len(document_types) ): raise InvalidArgumentError("documentType只允许MAIN、SUB_PLAN或二者组合") statement = select(Document).where( Document.is_deleted.is_(False), Document.document_type.in_(document_types), ) category_id = params.get("categoryId") if category_id: statement = statement.where( Document.category_id == parse_string_id(category_id, field="categoryId") ) statement = _keyword(statement, params.get("keyword")) visibility_filter = _enum( params.get("visibilityType"), VisibilityType, "visibilityType" ) for name, enum_type, column in ( ("securityLevel", SecurityLevel, Document.security_level), ("status", DocumentStatus, Document.document_status), ): value = _enum(params.get(name), enum_type, name) if value: statement = statement.where(column == value) updated_from, updated_to = parse_updated_range(params) if updated_from: statement = statement.where(Document.updated_at >= updated_from) if updated_to: statement = statement.where(Document.updated_at <= updated_to) order, _ = _sort(params) documents = db.session.scalars(statement.order_by(order, Document.id.asc())).all() visible: list[tuple[Document, PlanAccess]] = [] for document in documents: access = evaluate_plan_access(document, context) if ( access.allowed and ( visibility_filter is None or ( access.source is not None and access.source.visibility_type == visibility_filter ) ) ): visible.append((document, access)) total = len(visible) selected = visible[page.offset : page.offset + page.page_size] return page_result( [ document_summary(document, context, access) for document, access in selected ], page=page.page, page_size=page.page_size, total=total, ) def _active_plan(document_id: int) -> Document: document = db.session.scalar( select(Document).where( Document.id == document_id, Document.is_deleted.is_(False), Document.document_type.in_(PLAN_TYPES), ) ) if document is None: raise ResourceNotFoundError("方案文档不存在") return document def _record_view(document: Document) -> int: old_count = document.view_count document_id = document.id document_name = document.document_name target = ( AuditTarget.ATTACHMENT if document.document_type == DocumentType.ATTACHMENT.value else AuditTarget.DOCUMENT ) db.session.rollback() try: with transaction() as session: current = session.scalar( select(Document) .where( Document.id == document_id, Document.is_deleted.is_(False), ) .with_for_update() ) if current is None: raise ResourceNotFoundError("文档不存在") current.view_count += 1 session.add( business_audit( action=AuditAction.VIEW_DOCUMENT, target=target, target_id=document_id, target_name=document_name, detail={"documentType": current.document_type}, ) ) new_count = current.view_count return new_count except Exception: db.session.rollback() logger.exception("文档查看计数或审计写入失败:document_id=%s", document_id) return old_count def get_document(document_id: int) -> dict[str, object]: context = get_auth_context() document = _active_plan(document_id) access = _require_plan_access(document, context) result = document_detail(document, context, access) result["viewCount"] = _record_view(document) return result def list_sub_plans( main_document_id: int, params: Mapping[str, str], ) -> dict[str, object]: context = get_auth_context() main = db.session.scalar( select(Document).where( Document.id == main_document_id, Document.document_type == DocumentType.MAIN.value, Document.is_deleted.is_(False), ) ) if main is None: raise MainPlanNotFoundError() _require_plan_access(main, context) page = _page(params) statement = select(Document).where( Document.parent_document_id == main.id, Document.document_type == DocumentType.SUB_PLAN.value, Document.is_deleted.is_(False), ) statement = _keyword(statement, params.get("keyword")) status = _enum(params.get("status"), DocumentStatus, "status") if status: statement = statement.where(Document.document_status == status) updated_from, updated_to = parse_updated_range(params) if updated_from: statement = statement.where(Document.updated_at >= updated_from) if updated_to: statement = statement.where(Document.updated_at <= updated_to) order, _ = _sort(params) children = db.session.scalars( statement.order_by(order, Document.id.asc()) ).all() visible: list[tuple[Document, PlanAccess]] = [] for child in children: access = evaluate_plan_access(child, context) if access.allowed: visible.append((child, access)) total = len(visible) selected = visible[page.offset : page.offset + page.page_size] return page_result( [document_summary(child, context, access) for child, access in selected], page=page.page, page_size=page.page_size, total=total, ) __all__ = [ "DOCUMENT_SORTS", "_date", "_enum", "_iso", "_keyword", "_page", "_record_view", "_sort", "document_detail", "document_summary", "get_document", "list_documents", "list_sub_plans", "parse_string_id", ]