release(main): 同步 dev 最新候选改动
Storage Persistence Guard / storage-persistence-audit (push) Has been cancelled
Client Quality Gates / Shared client and Web (push) Has been cancelled
Client Quality Gates / macOS Desktop (push) Has been cancelled
Client Quality Gates / Shared client and Web (pull_request) Has been cancelled
Client Quality Gates / macOS Desktop (pull_request) Has been cancelled
Storage Persistence Guard / storage-persistence-audit (pull_request) Has been cancelled

This commit is contained in:
Cheng Zhou
2026-07-16 17:15:50 +08:00
parent 32167fba02
commit d5279b124f
393 changed files with 51630 additions and 9711 deletions
+34 -30
View File
@@ -77,6 +77,7 @@ FAQ_REPLY_ATTACHMENT_PERMISSION_BY_ACTION = {
"read": "faq:read",
"delete": "faq_attachments:delete",
}
SHARED_LIBRARY_ENTITY_TYPES = {"precaution", "faq_replies"}
STARTUP_AUTH_ATTACHMENT_ENTITY_TYPES = {
"startup_kickoff",
"startup_kickoff_minutes",
@@ -239,16 +240,17 @@ async def upload_attachment(
content_type=file.content_type,
uploaded_by=current_user.id,
)
await audit_crud.log_action(
db,
study_id=study_id,
entity_type=entity_type,
entity_id=entity_id,
action="UPLOAD_FILE",
detail=f"文件已上传:{file.filename}",
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
if entity_type not in SHARED_LIBRARY_ENTITY_TYPES:
await audit_crud.log_action(
db,
study_id=study_id,
entity_type=entity_type,
entity_id=entity_id,
action="UPLOAD_FILE",
detail=f"文件已上传:{file.filename}",
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
return AttachmentRead(
id=attachment.id,
filename=attachment.filename,
@@ -442,16 +444,17 @@ async def global_delete_attachment(
if not can_delete:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="无权限删除附件")
await attachment_crud.soft_delete_attachment(db, attachment)
await audit_crud.log_action(
db,
study_id=attachment.study_id,
entity_type=attachment.entity_type,
entity_id=attachment.entity_id,
action="DELETE_ATTACHMENT",
detail=f"文件已删除:{attachment.filename}",
operator_id=user.id,
operator_role=await get_operator_role_label(db, attachment.study_id, user),
)
if attachment.entity_type not in SHARED_LIBRARY_ENTITY_TYPES:
await audit_crud.log_action(
db,
study_id=attachment.study_id,
entity_type=attachment.entity_type,
entity_id=attachment.entity_id,
action="DELETE_ATTACHMENT",
detail=f"文件已删除:{attachment.filename}",
operator_id=user.id,
operator_role=await get_operator_role_label(db, attachment.study_id, user),
)
@router.delete(
@@ -494,13 +497,14 @@ async def delete_attachment(
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="无权限删除附件")
await attachment_crud.soft_delete_attachment(db, attachment)
await audit_crud.log_action(
db,
study_id=study_id,
entity_type=entity_type,
entity_id=entity_id,
action="DELETE_ATTACHMENT",
detail=f"文件已删除:{attachment.filename}",
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
if entity_type not in SHARED_LIBRARY_ENTITY_TYPES:
await audit_crud.log_action(
db,
study_id=study_id,
entity_type=entity_type,
entity_id=entity_id,
action="DELETE_ATTACHMENT",
detail=f"文件已删除:{attachment.filename}",
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
+51 -1
View File
@@ -1,13 +1,18 @@
import json
import uuid
from datetime import datetime
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_operator_role_label, get_current_user, get_db_session, require_api_permission
from app.crud import audit as audit_crud
from app.crud import study as study_crud
from app.models.user import User
from app.schemas.audit import AuditEventCreate, AuditLogRead
from app.services.ip_location import resolve_ip_location
router = APIRouter()
@@ -40,6 +45,10 @@ async def list_audit_logs(
entity_id: uuid.UUID | None = None,
action: str | None = None,
operator_id: uuid.UUID | None = None,
client_ip: str | None = None,
client_type: str | None = None,
start_time: datetime | None = None,
end_time: datetime | None = None,
skip: int = 0,
limit: int = 100,
db: AsyncSession = Depends(get_db_session),
@@ -51,10 +60,51 @@ async def list_audit_logs(
entity_id=entity_id,
action=action,
operator_id=operator_id,
client_ip=client_ip,
client_type=client_type,
start_time=start_time,
end_time=end_time,
skip=skip,
limit=limit,
)
return list(logs)
operator_ids = {log.operator_id for log in logs}
operator_map: dict[uuid.UUID, User] = {}
if operator_ids:
result = await db.execute(select(User).where(User.id.in_(operator_ids)))
operator_map = {user.id: user for user in result.scalars().all()}
items: list[AuditLogRead] = []
for log in logs:
ip_location = resolve_ip_location(log.client_ip)
operator = operator_map.get(log.operator_id)
items.append(
AuditLogRead(
id=log.id,
study_id=log.study_id,
entity_type=log.entity_type,
entity_id=log.entity_id,
action=log.action,
detail=log.detail,
operator_id=log.operator_id,
operator_name=operator.full_name if operator else None,
operator_email=operator.email if operator else None,
operator_role=log.operator_role,
client_ip=log.client_ip,
ip_location=ip_location.location,
ip_country=ip_location.country,
ip_province=ip_location.province,
ip_city=ip_location.city,
ip_isp=ip_location.isp,
user_agent=log.user_agent,
client_type=log.client_type,
client_version=log.client_version,
client_platform=log.client_platform,
build_channel=log.build_channel,
build_commit=log.build_commit,
created_at=log.created_at,
)
)
return items
@router.post(
+108 -12
View File
@@ -1,5 +1,6 @@
from datetime import datetime, timedelta, timezone
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from dataclasses import dataclass
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response, status
from fastapi import File, UploadFile
from pydantic import BaseModel, EmailStr, Field
from sqlalchemy.ext.asyncio import AsyncSession
@@ -8,7 +9,8 @@ import uuid
from app.core.config import settings
from app.core.login_crypto import create_login_challenge, decrypt_login_payload, get_public_key_pem
from app.core.security import create_access_token, decode_token_allow_expired, oauth2_scheme, verify_password
from app.core.request_context import resolve_client_ip, resolve_ctms_client_type
from app.core.security import create_access_token, decode_token, decode_token_allow_expired, oauth2_scheme, verify_password
from app.core.deps import get_current_user, get_db_session
from app.crud import user as user_crud
from app.models.user import UserStatus
@@ -25,6 +27,12 @@ from app.schemas.email_settings import (
)
from app.schemas.user import Token, UserRead, UserRegisterRequest, UserSelfUpdate, UserUpdate
from app.services import email_service
from app.services.user_login_sessions import (
create_login_session,
end_login_session,
session_id_from_payload,
touch_login_session,
)
from fastapi.responses import FileResponse
@@ -70,12 +78,59 @@ AVATAR_ALLOWED_CONTENT_TYPES = {
}
def issue_user_token(db_user) -> Token:
@dataclass(frozen=True)
class SessionPolicy:
access_minutes: int
absolute_max_seconds: int
def normalize_session_client_type(value: str | None) -> str:
return "desktop" if (value or "").strip().lower() == "desktop" else "web"
def get_session_policy_for_client_type(client_type: str) -> SessionPolicy:
if client_type == "desktop":
max_seconds = settings.DESKTOP_SESSION_MAX_DAYS * 24 * 3600
return SessionPolicy(
access_minutes=settings.DESKTOP_SESSION_MAX_DAYS * 24 * 60,
absolute_max_seconds=max_seconds,
)
return SessionPolicy(
access_minutes=settings.JWT_EXPIRE_MINUTES,
absolute_max_seconds=settings.ABSOLUTE_SESSION_MAX_HOURS * 3600,
)
def get_request_session_client_type(request: Request) -> str:
return normalize_session_client_type(resolve_ctms_client_type(request.headers))
def policy_expires_at(issued_at: datetime, session_start: datetime, policy: SessionPolicy) -> datetime:
access_expires_at = issued_at + timedelta(minutes=policy.access_minutes)
session_expires_at = session_start + timedelta(seconds=policy.absolute_max_seconds)
return min(access_expires_at, session_expires_at)
async def issue_user_token(db_user, request: Request, db: AsyncSession) -> Token:
session_start = datetime.now(timezone.utc)
session_id = uuid.uuid4()
client_type = get_request_session_client_type(request)
policy = get_session_policy_for_client_type(client_type)
access_token = create_access_token(
user_id=str(db_user.id),
expires_minutes=None,
expires_minutes=policy.access_minutes,
session_start=session_start,
max_age_seconds=policy.absolute_max_seconds,
issued_at=session_start,
client_type=client_type,
session_id=str(session_id),
)
await create_login_session(
db,
session_id=session_id,
user_id=db_user.id,
request=request,
login_at=session_start,
)
return Token(access_token=access_token, token_type="bearer")
@@ -240,23 +295,23 @@ async def get_login_key() -> LoginKeyResponse:
@router.post("/login", response_model=Token)
async def login_for_access_token(
payload: LoginRequest, db: AsyncSession = Depends(get_db_session)
payload: LoginRequest, request: Request, db: AsyncSession = Depends(get_db_session)
) -> Token:
db_user = await authenticate_encrypted_password(payload, db)
ensure_user_active(db_user)
return issue_user_token(db_user)
return await issue_user_token(db_user, request, db)
@router.post("/dev-login", response_model=Token)
async def dev_login_for_access_token(
payload: DevLoginRequest, db: AsyncSession = Depends(get_db_session)
payload: DevLoginRequest, request: Request, db: AsyncSession = Depends(get_db_session)
) -> Token:
if settings.ENV != "development":
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not found")
db_user = await authenticate_plain_password(payload, db)
ensure_user_active(db_user)
return issue_user_token(db_user)
return await issue_user_token(db_user, request, db)
@router.get("/me", response_model=UserRead)
@@ -285,22 +340,63 @@ async def extend_access_token(
if db_user.status != UserStatus.ACTIVE:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="账号已停用")
session_start_ts = payload.get("orig_iat") or payload.get("iat")
policy = get_session_policy_for_client_type(normalize_session_client_type(payload.get("client_type")))
if session_start_ts:
max_seconds = settings.ABSOLUTE_SESSION_MAX_HOURS * 3600
if now_ts - int(session_start_ts) > max_seconds:
if now_ts - int(session_start_ts) > policy.absolute_max_seconds:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="会话已到期,请重新登录")
session_start = datetime.fromtimestamp(int(session_start_ts), tz=timezone.utc)
else:
session_start = datetime.now(timezone.utc)
issued_at = datetime.now(timezone.utc)
new_token = create_access_token(
user_id=str(db_user.id),
expires_minutes=None,
expires_minutes=policy.access_minutes,
session_start=session_start,
max_age_seconds=policy.absolute_max_seconds,
issued_at=issued_at,
client_type=normalize_session_client_type(payload.get("client_type")),
session_id=str(session_id_from_payload(payload)),
)
expires_at = datetime.now(timezone.utc) + timedelta(minutes=settings.JWT_EXPIRE_MINUTES)
expires_at = policy_expires_at(issued_at, session_start, policy)
return ExtendResponse(accessToken=new_token, expiresAt=expires_at)
@router.post("/session/heartbeat")
async def heartbeat_login_session(
request: Request,
token: str = Depends(oauth2_scheme),
current_user=Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> dict:
session = await touch_login_session(
db,
user_id=current_user.id,
payload=decode_token(token),
request=request,
)
if session is None:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="登录会话已结束")
return {
"status": "online",
"last_seen_at": session.last_seen_at.isoformat(),
"client_ip": resolve_client_ip(request),
}
@router.post("/session/logout", status_code=status.HTTP_204_NO_CONTENT)
async def logout_login_session(
token: str = Depends(oauth2_scheme),
current_user=Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> Response:
await end_login_session(
db,
user_id=current_user.id,
payload=decode_token(token),
)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@router.patch("/me", response_model=UserRead)
async def update_me(
payload: UserSelfUpdate,
+629
View File
@@ -0,0 +1,629 @@
from __future__ import annotations
import uuid
from fastapi import APIRouter, Depends, File, Form, Header, Request, Response, UploadFile, status
from fastapi.responses import FileResponse
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_user, get_db_session, require_api_permission
from app.schemas.collaboration import (
CollaborationCallbackPayload,
CollaborationCandidateRead,
CollaborationEditorConfigRead,
CollaborationEditRequestRead,
CollaborationEditRequestResolve,
CollaborationExportRecord,
CollaborationFileCreate,
CollaborationFileRead,
CollaborationFileUpdate,
CollaborationFolderCreate,
CollaborationFolderRead,
CollaborationFolderUpdate,
CollaborationMemberRead,
CollaborationMemberUpsert,
CollaborationOwnershipTransferRequest,
CollaborationPublicEditorConfigRequest,
CollaborationPublicShareMetadata,
CollaborationRestoreRequest,
CollaborationRevisionCopyRequest,
CollaborationRevisionRead,
CollaborationRevisionUpdate,
CollaborationShareAccessGrant,
CollaborationShareLinkRead,
CollaborationShareLinkUpdate,
CollaborationSharePasswordRequest,
)
from app.schemas.onlyoffice import OnlyOfficePreviewConfigRead
from app.services import collaboration_service, collaboration_share_service, onlyoffice_collaboration_service, onlyoffice_service
router = APIRouter()
public_router = APIRouter()
internal_router = APIRouter(include_in_schema=False)
@router.get(
"/folders",
response_model=list[CollaborationFolderRead],
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def list_folders(study_id: uuid.UUID, db: AsyncSession = Depends(get_db_session)):
return await collaboration_service.list_folders(db, study_id)
@router.post(
"/folders",
response_model=CollaborationFolderRead,
status_code=status.HTTP_201_CREATED,
dependencies=[Depends(require_api_permission("collaboration:manage"))],
)
async def create_folder(
study_id: uuid.UUID,
payload: CollaborationFolderCreate,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
return await collaboration_service.create_folder(db, study_id, payload, current_user)
@router.patch(
"/folders/{folder_id}",
response_model=CollaborationFolderRead,
dependencies=[Depends(require_api_permission("collaboration:manage"))],
)
async def update_folder(
study_id: uuid.UUID,
folder_id: uuid.UUID,
payload: CollaborationFolderUpdate,
db: AsyncSession = Depends(get_db_session),
):
return await collaboration_service.update_folder(db, study_id, folder_id, payload)
@router.delete(
"/folders/{folder_id}",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_api_permission("collaboration:manage"))],
)
async def delete_folder(study_id: uuid.UUID, folder_id: uuid.UUID, db: AsyncSession = Depends(get_db_session)):
await collaboration_service.delete_folder(db, study_id, folder_id)
@router.get(
"/files",
response_model=list[CollaborationFileRead],
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def list_files(
study_id: uuid.UUID,
folder_id: uuid.UUID | None = None,
keyword: str | None = None,
deleted: bool = False,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
return await collaboration_service.list_files(
db, study_id, current_user, folder_id=folder_id, keyword=keyword, deleted=deleted
)
@router.post(
"/files",
response_model=CollaborationFileRead,
status_code=status.HTTP_201_CREATED,
dependencies=[Depends(require_api_permission("collaboration:create"))],
)
async def create_file(
study_id: uuid.UUID,
payload: CollaborationFileCreate,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.create_blank_file(db, study_id, payload, current_user)
return await collaboration_service.file_read(db, item, current_user)
@router.post(
"/files/import",
response_model=CollaborationFileRead,
status_code=status.HTTP_201_CREATED,
dependencies=[Depends(require_api_permission("collaboration:create"))],
)
async def import_file(
study_id: uuid.UUID,
file: UploadFile = File(...),
folder_id: uuid.UUID | None = Form(None),
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.import_file(db, study_id, folder_id, file, current_user)
return await collaboration_service.file_read(db, item, current_user)
@router.post(
"/files/{file_id}/copy",
response_model=CollaborationFileRead,
status_code=status.HTTP_201_CREATED,
dependencies=[Depends(require_api_permission("collaboration:create"))],
)
async def copy_file(
study_id: uuid.UUID,
file_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
source = await collaboration_service.get_file_or_404(db, study_id, file_id)
item = await collaboration_service.copy_file(db, source, current_user)
return await collaboration_service.file_read(db, item, current_user)
@router.get(
"/files/{file_id}",
response_model=CollaborationFileRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def get_file(
study_id: uuid.UUID,
file_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.file_read(db, item, current_user)
@router.get(
"/files/{file_id}/download",
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def download_file(
study_id: uuid.UUID,
file_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
revision = await collaboration_service.prepare_download(db, item, current_user)
return FileResponse(
path=revision.file_uri,
media_type=revision.mime_type,
filename=item.title,
content_disposition_type="attachment",
headers={"Cache-Control": "no-store"},
)
@router.patch(
"/files/{file_id}",
response_model=CollaborationFileRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def update_file(
study_id: uuid.UUID,
file_id: uuid.UUID,
payload: CollaborationFileUpdate,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
item = await collaboration_service.update_file(db, item, payload, current_user)
return await collaboration_service.file_read(db, item, current_user)
@router.delete(
"/files/{file_id}",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def trash_file(
study_id: uuid.UUID,
file_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
await collaboration_service.move_to_trash(db, item, current_user)
@router.post(
"/files/{file_id}/restore",
response_model=CollaborationFileRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def restore_file(
study_id: uuid.UUID,
file_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id, include_deleted=True)
item = await collaboration_service.restore_file(db, item, current_user)
return await collaboration_service.file_read(db, item, current_user)
@router.get(
"/files/{file_id}/members",
response_model=list[CollaborationMemberRead],
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def list_members(study_id: uuid.UUID, file_id: uuid.UUID, db: AsyncSession = Depends(get_db_session)):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.list_members(db, item)
@router.put(
"/files/{file_id}/members",
response_model=CollaborationMemberRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def upsert_member(
study_id: uuid.UUID,
file_id: uuid.UUID,
payload: CollaborationMemberUpsert,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.upsert_member(db, item, payload, current_user)
@router.delete(
"/files/{file_id}/members/{user_id}",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def remove_member(
study_id: uuid.UUID,
file_id: uuid.UUID,
user_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
await collaboration_service.remove_member(db, item, user_id, current_user)
@router.post(
"/files/{file_id}/edit-requests",
response_model=CollaborationEditRequestRead,
status_code=status.HTTP_201_CREATED,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def create_edit_request(
study_id: uuid.UUID,
file_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.create_edit_request(db, item, current_user)
@router.get(
"/files/{file_id}/edit-requests",
response_model=list[CollaborationEditRequestRead],
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def list_edit_requests(
study_id: uuid.UUID,
file_id: uuid.UUID,
pending_only: bool = True,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.list_edit_requests(db, item, current_user, pending_only=pending_only)
@router.post(
"/files/{file_id}/edit-requests/{request_id}/resolve",
response_model=CollaborationEditRequestRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def resolve_edit_request(
study_id: uuid.UUID,
file_id: uuid.UUID,
request_id: uuid.UUID,
payload: CollaborationEditRequestResolve,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.resolve_edit_request(db, item, request_id, payload, current_user)
@router.post(
"/files/{file_id}/transfer-ownership",
response_model=CollaborationFileRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def transfer_ownership(
study_id: uuid.UUID,
file_id: uuid.UUID,
payload: CollaborationOwnershipTransferRequest,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
item = await collaboration_service.transfer_ownership(db, item, payload, current_user)
return await collaboration_service.file_read(db, item, current_user)
@router.get(
"/files/{file_id}/share-link",
response_model=CollaborationShareLinkRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def get_share_link(
study_id: uuid.UUID,
file_id: uuid.UUID,
response: Response,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
result = await collaboration_share_service.get_share_link(db, item, current_user)
response.headers["Cache-Control"] = "no-store"
return result
@router.put(
"/files/{file_id}/share-link",
response_model=CollaborationShareLinkRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def update_share_link(
study_id: uuid.UUID,
file_id: uuid.UUID,
payload: CollaborationShareLinkUpdate,
response: Response,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
result = await collaboration_share_service.update_share_link(db, item, payload, current_user)
response.headers["Cache-Control"] = "no-store"
return result
@router.get(
"/member-candidates",
response_model=list[CollaborationCandidateRead],
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def list_member_candidates(study_id: uuid.UUID, db: AsyncSession = Depends(get_db_session)):
return await collaboration_service.list_candidates(db, study_id)
@router.get(
"/files/{file_id}/revisions",
response_model=list[CollaborationRevisionRead],
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def list_revisions(study_id: uuid.UUID, file_id: uuid.UUID, db: AsyncSession = Depends(get_db_session)):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.list_revisions(db, item)
@router.patch(
"/files/{file_id}/revisions/{revision_id}",
response_model=CollaborationRevisionRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def update_revision(
study_id: uuid.UUID,
file_id: uuid.UUID,
revision_id: uuid.UUID,
payload: CollaborationRevisionUpdate,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.update_revision(db, item, revision_id, payload, current_user)
@router.delete(
"/files/{file_id}/revisions/{revision_id}",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def delete_revision(
study_id: uuid.UUID,
file_id: uuid.UUID,
revision_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
await collaboration_service.delete_revision(db, item, revision_id, current_user)
@router.post(
"/files/{file_id}/revisions/{revision_id}/copy",
response_model=CollaborationFileRead,
status_code=status.HTTP_201_CREATED,
dependencies=[
Depends(require_api_permission("collaboration:create")),
Depends(require_api_permission("collaboration:read")),
],
)
async def copy_revision(
study_id: uuid.UUID,
file_id: uuid.UUID,
revision_id: uuid.UUID,
payload: CollaborationRevisionCopyRequest,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
copied = await collaboration_service.copy_revision(db, item, revision_id, payload, current_user)
return await collaboration_service.file_read(db, copied, current_user)
@router.get(
"/files/{file_id}/revisions/{revision_id}/preview-config",
response_model=OnlyOfficePreviewConfigRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def get_revision_preview_config(
study_id: uuid.UUID,
file_id: uuid.UUID,
revision_id: uuid.UUID,
response: Response,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
revision = await collaboration_service.prepare_revision_preview(db, item, revision_id, current_user)
await onlyoffice_service.ensure_onlyoffice_available()
result = onlyoffice_service.build_preview_config(
resource_type="collaboration_revision",
resource_id=revision.id,
file_name=item.title,
file_hash=revision.file_hash,
user_id=current_user.id,
user_name=current_user.full_name,
)
response.headers["Cache-Control"] = "no-store"
return result
@router.post(
"/files/{file_id}/revisions/{revision_id}/restore",
response_model=CollaborationRevisionRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def restore_revision(
study_id: uuid.UUID,
file_id: uuid.UUID,
revision_id: uuid.UUID,
payload: CollaborationRestoreRequest,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
return await collaboration_service.restore_revision(
db, item, revision_id, current_user, payload.change_summary
)
@router.get(
"/files/{file_id}/editor-config",
response_model=CollaborationEditorConfigRead,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def get_editor_config(
study_id: uuid.UUID,
file_id: uuid.UUID,
response: Response,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
result = await onlyoffice_collaboration_service.build_editor_config(db, item, current_user)
response.headers["Cache-Control"] = "no-store"
return result
@router.post(
"/files/{file_id}/exports",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def record_export(
study_id: uuid.UUID,
file_id: uuid.UUID,
payload: CollaborationExportRecord,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
await collaboration_service.record_export(db, item, current_user, payload.file_type)
@router.post(
"/files/{file_id}/downloads",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_api_permission("collaboration:read"))],
)
async def record_download(
study_id: uuid.UUID,
file_id: uuid.UUID,
payload: CollaborationExportRecord,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
):
item = await collaboration_service.get_file_or_404(db, study_id, file_id)
await collaboration_service.record_download(db, item, current_user, payload.file_type)
return Response(status_code=status.HTTP_204_NO_CONTENT)
@public_router.get("/metadata", response_model=CollaborationPublicShareMetadata)
async def get_public_share_metadata(
response: Response,
x_ctms_share_token: str | None = Header(default=None, alias="X-CTMS-Share-Token"),
db: AsyncSession = Depends(get_db_session),
):
result = await collaboration_share_service.public_metadata(db, x_ctms_share_token)
response.headers["Cache-Control"] = "no-store"
return result
@public_router.post("/access", response_model=CollaborationShareAccessGrant)
async def verify_public_share_password(
payload: CollaborationSharePasswordRequest,
response: Response,
x_ctms_share_token: str | None = Header(default=None, alias="X-CTMS-Share-Token"),
db: AsyncSession = Depends(get_db_session),
):
result = await collaboration_share_service.verify_share_password(
db, x_ctms_share_token, payload.password
)
response.headers["Cache-Control"] = "no-store"
return result
@public_router.post("/editor-config", response_model=CollaborationEditorConfigRead)
async def get_public_share_editor_config(
payload: CollaborationPublicEditorConfigRequest,
response: Response,
x_ctms_share_token: str | None = Header(default=None, alias="X-CTMS-Share-Token"),
db: AsyncSession = Depends(get_db_session),
):
link, item = await collaboration_share_service.resolve_active_share(db, x_ctms_share_token)
collaboration_share_service.validate_access_grant(link, payload.access_token)
result = await onlyoffice_collaboration_service.build_shared_editor_config(
db,
item,
link,
client_id=payload.client_id,
display_name=payload.display_name,
)
response.headers["Cache-Control"] = "no-store"
return result
@internal_router.get("/internal/onlyoffice/collaboration/sessions/{session_id}/content")
async def get_session_content(session_id: uuid.UUID, request: Request, db: AsyncSession = Depends(get_db_session)):
revision, item = await onlyoffice_collaboration_service.get_session_content(
db, session_id, request.headers.get("AuthorizationJwt")
)
return FileResponse(
path=revision.file_uri,
media_type=revision.mime_type,
filename=item.title,
content_disposition_type="inline",
)
@internal_router.post("/internal/onlyoffice/collaboration/sessions/{session_id}/callback")
async def collaboration_callback(
session_id: uuid.UUID,
payload: CollaborationCallbackPayload,
request: Request,
db: AsyncSession = Depends(get_db_session),
):
onlyoffice_collaboration_service.validate_callback_token(
request.headers.get("AuthorizationJwt"), payload
)
return await onlyoffice_collaboration_service.process_callback(db, session_id, payload)
+3 -3
View File
@@ -75,12 +75,12 @@ async def acknowledge_notifications(
)
@router.post("/{distribution_id}/read", status_code=status.HTTP_204_NO_CONTENT)
@router.post("/{notification_id}/read", status_code=status.HTTP_204_NO_CONTENT)
async def mark_notification_read(
distribution_id: uuid.UUID,
notification_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
) -> None:
await desktop_notification_service.mark_notification_read(
db, current_user.id, distribution_id
db, current_user.id, notification_id
)
+1 -36
View File
@@ -1,15 +1,11 @@
import uuid
import json
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_user, get_db_session, get_operator_role_label, is_system_admin, require_study_not_locked, require_api_permission
from app.core.project_permissions import role_has_api_permission
from app.crud import audit as audit_crud
from app.core.deps import get_current_user, get_db_session, is_system_admin, require_study_not_locked, require_api_permission
from app.crud import faq_category as category_crud
from app.crud import faq_item as item_crud
from app.crud import member as member_crud
from app.schemas.common import PaginatedResponse
from app.schemas.faq import CategoryCreate, CategoryRead, CategoryUpdate
from app.utils.pagination import paginate
@@ -47,16 +43,6 @@ async def create_category(
if dup:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="该图标已被其他分类使用")
category = await category_crud.create_category(db, payload)
await audit_crud.log_action(
db,
study_id=payload.study_id,
entity_type="faq_category",
entity_id=category.id,
action="CREATE_FAQ_CATEGORY",
detail=json.dumps({"targetName": category.name, "description": f"创建“{category.name}”分类"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, payload.study_id, current_user),
)
return CategoryRead.model_validate(category)
@@ -112,16 +98,6 @@ async def update_category(
if not target_study_id:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="必须提供项目 ID")
updated = await category_crud.update_category(db, category, payload)
await audit_crud.log_action(
db,
study_id=updated.study_id,
entity_type="faq_category",
entity_id=category_id,
action="UPDATE_FAQ_CATEGORY",
detail=json.dumps({"targetName": updated.name, "description": f"更新“{updated.name}”分类"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, updated.study_id, current_user),
)
return CategoryRead.model_validate(updated)
@@ -148,16 +124,5 @@ async def delete_category(
item_count = await item_crud.count_items_by_category(db, category_id)
if item_count > 0:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="分类下存在 FAQ,无法删除")
category_name = category.name
await db.delete(category)
await db.commit()
await audit_crud.log_action(
db,
study_id=category.study_id,
entity_type="faq_category",
entity_id=category_id,
action="DELETE_FAQ_CATEGORY",
detail=json.dumps({"targetName": category_name, "description": f"删除“{category_name}”分类"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, category.study_id, current_user),
)
+1 -68
View File
@@ -1,12 +1,9 @@
import json
import uuid
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_user, get_db_session, get_operator_role_label, is_system_admin, require_study_not_locked, require_api_permission
from app.core.project_permissions import role_has_api_permission
from app.crud import audit as audit_crud
from app.core.deps import get_current_user, get_db_session, is_system_admin, require_study_not_locked, require_api_permission
from app.crud import faq_category as category_crud
from app.crud import faq_item as faq_crud
from app.crud import faq_reply as reply_crud
@@ -31,13 +28,6 @@ def _is_system_admin(current_user) -> bool:
return is_system_admin(current_user)
def _compact_text(value: str | None, max_length: int = 40) -> str:
text = " ".join(str(value or "").split())
if len(text) <= max_length:
return text
return f"{text[:max_length]}..."
@router.post(
"/",
response_model=FaqRead,
@@ -74,17 +64,6 @@ async def create_faq(
reply_in=FaqReplyCreate(content=payload.answer),
)
await faq_crud.set_status(db, item.id, "PROCESSING")
question_name = _compact_text(item.question)
await audit_crud.log_action(
db,
study_id=payload.study_id,
entity_type="faq_item",
entity_id=item.id,
action="CREATE_FAQ_ITEM",
detail=json.dumps({"targetName": question_name, "description": f"创建医学咨询问题“{question_name}"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, payload.study_id, current_user),
)
return FaqRead.model_validate(item)
@@ -178,19 +157,6 @@ async def update_faq(
if not item:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="FAQ 不存在")
updated = await faq_crud.update_item(db, item, payload)
action = "UPDATE_FAQ_ITEM"
question_name = _compact_text(updated.question)
detail = json.dumps({"targetName": question_name, "description": f"更新医学咨询问题“{question_name}"}, ensure_ascii=False)
await audit_crud.log_action(
db,
study_id=item.study_id,
entity_type="faq_item",
entity_id=item_id,
action=action,
detail=detail,
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, item.study_id, current_user),
)
return FaqRead.model_validate(updated)
@@ -337,17 +303,6 @@ async def create_reply(
if item.status != "RESOLVED":
await faq_crud.set_status(db, item.id, "PROCESSING")
await faq_crud.touch_item(db, item.id)
question_name = _compact_text(item.question)
await audit_crud.log_action(
db,
study_id=item.study_id,
entity_type="faq_reply",
entity_id=reply.id,
action="CREATE_FAQ_REPLY",
detail=json.dumps({"targetName": question_name, "description": f"回复医学咨询问题“{question_name}"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, item.study_id, current_user),
)
data = FaqReplyRead.model_validate(reply)
if quote:
if quote.is_deleted:
@@ -380,20 +335,9 @@ async def delete_faq(
item = await faq_crud.get_item(db, item_id)
if not item:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="FAQ 不存在")
question_name = _compact_text(item.question)
await reply_crud.delete_replies_by_faq_id(db, item.id)
await db.delete(item)
await db.commit()
await audit_crud.log_action(
db,
study_id=item.study_id,
entity_type="faq_item",
entity_id=item_id,
action="DELETE_FAQ_ITEM",
detail=json.dumps({"targetName": question_name, "description": f"删除医学咨询问题“{question_name}"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, item.study_id, current_user),
)
@router.delete(
@@ -415,7 +359,6 @@ async def delete_reply(
item = await faq_crud.get_item(db, item_id)
if not item:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="FAQ 不存在")
question_name = _compact_text(item.question)
reply = await reply_crud.get_reply(db, reply_id)
if not reply or reply.faq_id != item.id:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="回复不存在")
@@ -438,13 +381,3 @@ async def delete_reply(
resolved_by_confirm=False,
)
await faq_crud.touch_item(db, item.id)
await audit_crud.log_action(
db,
study_id=item.study_id,
entity_type="faq_reply",
entity_id=reply_id,
action="DELETE_FAQ_REPLY",
detail=json.dumps({"targetName": question_name, "description": f"删除医学咨询问题“{question_name}”的回复"}, ensure_ascii=False),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, item.study_id, current_user),
)
+65 -3
View File
@@ -1,11 +1,11 @@
import uuid
from fastapi import APIRouter, Depends
from fastapi import APIRouter, Depends, Response, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_db_session, get_current_user, require_study_member
from app.schemas.notification import NotificationItem
from app.services import document_service
from app.schemas.notification import GeneralNotificationFeed, GeneralNotificationRead, NotificationItem
from app.services import document_service, notification_service, project_reminder_service
router = APIRouter()
@@ -29,3 +29,65 @@ async def list_notifications(
skip=skip,
limit=limit,
)
@router.get(
"/notifications/feed",
response_model=GeneralNotificationFeed,
dependencies=[Depends(require_study_member())],
)
async def list_general_notifications(
study_id: uuid.UUID,
skip: int = 0,
limit: int = 10,
category: str | None = None,
unread_only: bool = False,
requires_action: bool | None = None,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
) -> GeneralNotificationFeed:
# Fail closed: stale reminders may contain details for a permission that was just revoked.
await project_reminder_service.sync_project_reminders(db, study_id, current_user)
return await notification_service.list_feed(
db,
study_id=study_id,
recipient_id=current_user.id,
skip=skip,
limit=limit,
category=category,
unread_only=unread_only,
requires_action=requires_action,
)
@router.post(
"/notifications/{notification_id}/read",
response_model=GeneralNotificationRead,
dependencies=[Depends(require_study_member())],
)
async def mark_general_notification_read(
study_id: uuid.UUID,
notification_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
) -> GeneralNotificationRead:
return await notification_service.mark_read(
db,
study_id=study_id,
recipient_id=current_user.id,
notification_id=notification_id,
)
@router.post(
"/notifications/read-all",
status_code=status.HTTP_204_NO_CONTENT,
dependencies=[Depends(require_study_member())],
)
async def mark_all_general_notifications_read(
study_id: uuid.UUID,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
) -> Response:
await notification_service.mark_all_read(db, study_id=study_id, recipient_id=current_user.id)
return Response(status_code=status.HTTP_204_NO_CONTENT)
+195
View File
@@ -0,0 +1,195 @@
from __future__ import annotations
import os
import uuid
from pathlib import Path
from urllib.parse import quote
from fastapi import APIRouter, Depends, Request, Response, status
from fastapi.responses import FileResponse
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.v1.attachments import _ensure_attachment_permission, _ensure_study_exists
from app.core.deps import get_current_user, get_db_session
from app.crud import attachment as attachment_crud
from app.crud import document as document_crud
from app.crud import document_version as version_crud
from app.models.collaboration import CollaborationRevision
from app.schemas.onlyoffice import OnlyOfficePreviewConfigRead
from app.services import document_service, onlyoffice_service
router = APIRouter()
internal_router = APIRouter(include_in_schema=False)
def _content_disposition(filename: str) -> str:
fallback = "".join(
character if 32 <= ord(character) < 127 and character not in {'"', "\\"} else "_"
for character in filename
) or "document"
encoded = quote(filename, safe="")
return f'inline; filename="{fallback}"; filename*=UTF-8\'\'{encoded}'
@router.get(
"/attachments/{attachment_id}/config",
response_model=OnlyOfficePreviewConfigRead,
)
async def get_attachment_preview_config(
attachment_id: uuid.UUID,
response: Response,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
) -> OnlyOfficePreviewConfigRead:
attachment = await attachment_crud.get_attachment(db, attachment_id)
if not attachment:
raise onlyoffice_service.onlyoffice_error(
"ATTACHMENT_NOT_FOUND", "附件不存在", status.HTTP_404_NOT_FOUND
)
await _ensure_study_exists(db, attachment.study_id)
await _ensure_attachment_permission(
db,
attachment.study_id,
attachment.entity_type,
attachment.entity_id,
"read",
current_user,
)
if not os.path.exists(attachment.file_path):
raise onlyoffice_service.onlyoffice_error(
"ATTACHMENT_FILE_NOT_FOUND", "服务器未找到文件", status.HTTP_404_NOT_FOUND
)
if not onlyoffice_service.office_format_for_filename(attachment.filename):
raise onlyoffice_service.onlyoffice_error(
"ONLYOFFICE_FORMAT_UNSUPPORTED",
"该文件格式不支持 Office 在线预览",
status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
)
await onlyoffice_service.ensure_onlyoffice_available()
result = onlyoffice_service.build_preview_config(
resource_type="attachment",
resource_id=attachment.id,
file_name=attachment.filename,
user_id=current_user.id,
user_name=current_user.full_name,
)
response.headers["Cache-Control"] = "no-store"
return result
@router.get(
"/versions/{version_id}/config",
response_model=OnlyOfficePreviewConfigRead,
)
async def get_version_preview_config(
version_id: uuid.UUID,
response: Response,
db: AsyncSession = Depends(get_db_session),
current_user=Depends(get_current_user),
) -> OnlyOfficePreviewConfigRead:
version = await version_crud.get(db, version_id)
if not version:
raise onlyoffice_service.onlyoffice_error(
"DOCUMENT_VERSION_NOT_FOUND", "版本不存在", status.HTTP_404_NOT_FOUND
)
document = await document_crud.get(db, version.document_id)
if not document:
raise onlyoffice_service.onlyoffice_error(
"DOCUMENT_NOT_FOUND", "文档不存在", status.HTTP_404_NOT_FOUND
)
await document_service._ensure_study_access(db, document.trial_id, current_user, action="view")
file_path = Path(version.file_uri)
if not file_path.exists():
raise onlyoffice_service.onlyoffice_error(
"DOCUMENT_FILE_NOT_FOUND", "文件不存在", status.HTTP_404_NOT_FOUND
)
file_name = version.original_filename or document_service._legacy_download_filename(version, document)
if not onlyoffice_service.office_format_for_filename(file_name):
raise onlyoffice_service.onlyoffice_error(
"ONLYOFFICE_FORMAT_UNSUPPORTED",
"该文件格式不支持 Office 在线预览",
status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
)
await onlyoffice_service.ensure_onlyoffice_available()
result = onlyoffice_service.build_preview_config(
resource_type="version",
resource_id=version.id,
file_name=file_name,
file_hash=version.file_hash,
user_id=current_user.id,
user_name=current_user.full_name,
)
response.headers["Cache-Control"] = "no-store"
return result
def _authorize_internal_file_request(request: Request, expected_url: str) -> None:
onlyoffice_service.validate_outbox_token(request.headers.get("AuthorizationJwt"), expected_url)
@internal_router.get("/internal/onlyoffice/attachments/{attachment_id}/content")
async def get_internal_attachment_content(
attachment_id: uuid.UUID,
request: Request,
db: AsyncSession = Depends(get_db_session),
) -> FileResponse:
expected_url = onlyoffice_service.onlyoffice_content_url("attachment", attachment_id)
_authorize_internal_file_request(request, expected_url)
attachment = await attachment_crud.get_attachment(db, attachment_id)
if not attachment or not os.path.exists(attachment.file_path):
raise onlyoffice_service.onlyoffice_error(
"ATTACHMENT_FILE_NOT_FOUND", "文件不存在", status.HTTP_404_NOT_FOUND
)
return FileResponse(
path=attachment.file_path,
media_type=attachment.content_type or "application/octet-stream",
headers={"Content-Disposition": _content_disposition(attachment.filename)},
)
@internal_router.get("/internal/onlyoffice/versions/{version_id}/content")
async def get_internal_version_content(
version_id: uuid.UUID,
request: Request,
db: AsyncSession = Depends(get_db_session),
) -> FileResponse:
expected_url = onlyoffice_service.onlyoffice_content_url("version", version_id)
_authorize_internal_file_request(request, expected_url)
version = await version_crud.get(db, version_id)
if not version:
raise onlyoffice_service.onlyoffice_error(
"DOCUMENT_VERSION_NOT_FOUND", "版本不存在", status.HTTP_404_NOT_FOUND
)
document = await document_crud.get(db, version.document_id)
file_path = Path(version.file_uri)
if not document or not file_path.exists():
raise onlyoffice_service.onlyoffice_error(
"DOCUMENT_FILE_NOT_FOUND", "文件不存在", status.HTTP_404_NOT_FOUND
)
file_name = version.original_filename or document_service._legacy_download_filename(version, document)
return FileResponse(
path=str(file_path),
media_type=version.mime_type or "application/octet-stream",
headers={"Content-Disposition": _content_disposition(file_name)},
)
@internal_router.get("/internal/onlyoffice/collaboration-revisions/{revision_id}/content")
async def get_internal_collaboration_revision_content(
revision_id: uuid.UUID,
request: Request,
db: AsyncSession = Depends(get_db_session),
) -> FileResponse:
expected_url = onlyoffice_service.onlyoffice_content_url("collaboration_revision", revision_id)
_authorize_internal_file_request(request, expected_url)
revision = await db.get(CollaborationRevision, revision_id)
file_path = Path(revision.file_uri) if revision else None
if not revision or getattr(revision, "deleted_at", None) is not None or not file_path or not file_path.exists():
raise onlyoffice_service.onlyoffice_error(
"COLLABORATION_REVISION_NOT_FOUND", "协作修订不存在", status.HTTP_404_NOT_FOUND
)
return FileResponse(
path=str(file_path),
media_type=revision.mime_type or "application/octet-stream",
headers={"Content-Disposition": _content_disposition(revision.original_filename)},
)
File diff suppressed because it is too large Load Diff
+1 -39
View File
@@ -1,11 +1,9 @@
import uuid
import json
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_operator_role_label, get_current_user, get_db_session, require_study_not_locked, require_api_permission
from app.crud import audit as audit_crud
from app.core.deps import get_current_user, get_db_session, require_study_not_locked, require_api_permission
from app.crud import precaution as precaution_crud
from app.crud import site as site_crud
from app.crud import study as study_crud
@@ -29,11 +27,6 @@ async def _ensure_site_name_active(db: AsyncSession, study_id: uuid.UUID, site_n
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="中心已停用")
def _precaution_audit_detail(action: str, precaution) -> str:
title = str(precaution.title or "").strip() or "注意事项"
return json.dumps({"targetName": title, "description": f"{action}注意事项“{title}"}, ensure_ascii=False)
@router.post(
"/precautions",
response_model=PrecautionRead,
@@ -49,16 +42,6 @@ async def create_precaution(
await _ensure_study_exists(db, study_id)
await _ensure_site_name_active(db, study_id, precaution_in.site_name)
precaution = await precaution_crud.create_precaution(db, study_id, precaution_in, created_by=current_user.id)
await audit_crud.log_action(
db,
study_id=study_id,
entity_type="precaution",
entity_id=precaution.id,
action="CREATE_PRECAUTION",
detail=_precaution_audit_detail("创建", precaution),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
return PrecautionRead.model_validate(precaution)
@@ -115,16 +98,6 @@ async def update_precaution(
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="注意事项不存在")
await _ensure_site_name_active(db, study_id, precaution.site_name)
precaution = await precaution_crud.update_precaution(db, precaution, precaution_in)
await audit_crud.log_action(
db,
study_id=study_id,
entity_type="precaution",
entity_id=precaution_id,
action="UPDATE_PRECAUTION",
detail=_precaution_audit_detail("更新", precaution),
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
return PrecautionRead.model_validate(precaution)
@@ -144,15 +117,4 @@ async def delete_precaution(
if not precaution or precaution.study_id != study_id:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="注意事项不存在")
await _ensure_site_name_active(db, study_id, precaution.site_name)
precaution_detail = _precaution_audit_detail("删除", precaution)
await precaution_crud.delete_precaution(db, precaution)
await audit_crud.log_action(
db,
study_id=study_id,
entity_type="precaution",
entity_id=precaution_id,
action="DELETE_PRECAUTION",
detail=precaution_detail,
operator_id=current_user.id,
operator_role=await get_operator_role_label(db, study_id, current_user),
)
+4 -1
View File
@@ -1,6 +1,6 @@
from fastapi import APIRouter
from app.api.v1 import auth, users, admin_email_settings, studies, sites, members, attachments, audit_logs, dashboard, subjects, visits, aes, finance_dashboard, fees_contracts, drug_shipments, material_equipments, project_milestones, startup, precautions, subject_histories, subject_pds, study_subject_pds, faq_categories, faqs, documents, etmf, overview, notifications, desktop_notifications, monitoring_visit_issues, api_permissions, permission_monitoring, permission_templates, system_permissions, study_active_roles
from app.api.v1 import auth, users, admin_email_settings, studies, sites, members, attachments, audit_logs, dashboard, subjects, visits, aes, finance_dashboard, fees_contracts, drug_shipments, material_equipments, project_milestones, startup, precautions, subject_histories, subject_pds, study_subject_pds, faq_categories, faqs, documents, etmf, overview, notifications, desktop_notifications, monitoring_visit_issues, api_permissions, permission_monitoring, permission_templates, system_permissions, study_active_roles, onlyoffice, collaboration
api_router = APIRouter()
@@ -17,6 +17,7 @@ api_router.include_router(api_permissions.router, tags=["api-permissions"])
api_router.include_router(api_permissions.study_router, prefix="/studies/{study_id}", tags=["api-permissions"])
api_router.include_router(attachments.router, prefix="/studies/{study_id}/{entity_type}/{entity_id}/attachments", tags=["attachments"])
api_router.include_router(attachments.global_router, prefix="/attachments", tags=["attachments"])
api_router.include_router(onlyoffice.router, prefix="/onlyoffice", tags=["onlyoffice"])
api_router.include_router(audit_logs.router, prefix="/studies/{study_id}/audit-logs", tags=["audit-logs"])
api_router.include_router(dashboard.router, prefix="/studies/{study_id}/dashboard", tags=["dashboard"])
api_router.include_router(subjects.router, prefix="/studies/{study_id}/subjects", tags=["subjects"])
@@ -29,6 +30,8 @@ api_router.include_router(material_equipments.router, prefix="/studies/{study_id
api_router.include_router(project_milestones.router, prefix="/studies/{study_id}/project", tags=["project-milestones"])
api_router.include_router(startup.router, prefix="/studies/{study_id}/startup", tags=["startup"])
api_router.include_router(precautions.router, prefix="/studies/{study_id}/shared-library", tags=["precautions"])
api_router.include_router(collaboration.router, prefix="/studies/{study_id}/collaboration", tags=["collaboration"])
api_router.include_router(collaboration.public_router, prefix="/collaboration/shares", tags=["collaboration-shares"])
api_router.include_router(monitoring_visit_issues.router, prefix="/studies/{study_id}/monitoring", tags=["monitoring-visit-issues"])
api_router.include_router(subject_histories.router, prefix="/studies/{study_id}/subjects/{subject_id}", tags=["subject-histories"])
api_router.include_router(subject_pds.router, prefix="/studies/{study_id}/subjects/{subject_id}", tags=["subject-pds"])
+42 -5
View File
@@ -1,6 +1,6 @@
import uuid
from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
@@ -9,7 +9,8 @@ from app.schemas.common import PaginatedResponse
from app.crud import user as user_crud
from app.crud import member as member_crud
from app.utils.pagination import paginate
from app.schemas.user import UserCreate, UserRead, UserStatus, UserUpdate
from app.schemas.user import LoginStatus, UserCreate, UserLoginActivityRead, UserRead, UserStatus, UserUpdate
from app.services.user_login_sessions import get_login_summaries, list_login_activities, login_activity_payload
router = APIRouter()
@@ -20,12 +21,48 @@ async def list_users(
limit: int = 100,
keyword: str | None = Query(default=None),
user_status: UserStatus | None = Query(default=None, alias="status"),
login_status: LoginStatus | None = Query(default=None),
db: AsyncSession = Depends(get_db_session),
current_user=Depends(require_roles(["ADMIN"])),
) -> PaginatedResponse[UserRead]:
users = await user_crud.list_users(db, skip=skip, limit=limit, keyword=keyword, status=user_status)
total_users = await user_crud.count_users(db, keyword=keyword, status=user_status)
return paginate(list(users), total=total_users)
users = await user_crud.list_users(
db, skip=skip, limit=limit, keyword=keyword, status=user_status, login_status=login_status
)
total_users = await user_crud.count_users(
db, keyword=keyword, status=user_status, login_status=login_status
)
summaries = await get_login_summaries(db, [user.id for user in users])
items = []
for user in users:
summary = summaries.get(user.id)
item = UserRead.model_validate(user).model_dump()
if summary:
item.update(
{
"login_status": summary.status,
"last_login_at": summary.last_login_at,
"last_seen_at": summary.last_seen_at,
"last_client_type": summary.client_type,
"active_session_count": summary.active_session_count,
}
)
items.append(item)
return paginate(items, total=total_users)
@router.get("/{user_id}/login-activities", response_model=list[UserLoginActivityRead])
async def read_user_login_activities(
user_id: uuid.UUID,
response: Response,
limit: int = Query(default=30, ge=1, le=100),
db: AsyncSession = Depends(get_db_session),
current_user=Depends(require_roles(["ADMIN"])),
) -> list[UserLoginActivityRead]:
response.headers["Cache-Control"] = "no-store"
if not await user_crud.get_by_id(db, user_id):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在")
activities = await list_login_activities(db, user_id=user_id, limit=limit)
return [login_activity_payload(activity) for activity in activities]
@router.post("/", response_model=UserRead, status_code=status.HTTP_201_CREATED)
+18
View File
@@ -12,6 +12,7 @@ from app.crud import study as study_crud
from app.crud import visit as visit_crud
from app.schemas.subject import SubjectUpdate
from app.schemas.visit import EarlyTerminationCreate, VisitCreate, VisitRead, VisitUpdate
from app.services import notification_service
router = APIRouter()
@@ -42,6 +43,17 @@ async def _ensure_subject_active(db: AsyncSession, subject) -> None:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="中心已停用")
async def _resolve_visit_reminders(db: AsyncSession, visit_ids: list[uuid.UUID]) -> None:
for visit_id in set(visit_ids):
await notification_service.resolve_source_notifications(
db,
source_type="SUBJECT_VISIT_WINDOW",
source_id=str(visit_id),
)
if visit_ids:
await db.commit()
@router.get(
"/",
response_model=list[VisitRead],
@@ -145,6 +157,8 @@ async def create_early_termination(
if subject.baseline_date and termination_in.termination_date < subject.baseline_date:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="提前终止日期不能早于基线/治疗日期")
existing_visit_ids = [visit.id for visit in await visit_crud.list_visits(db, subject_id)]
try:
visit = await visit_crud.create_early_termination_visit(
db,
@@ -164,6 +178,7 @@ async def create_early_termination(
drop_reason=reason,
),
)
await _resolve_visit_reminders(db, existing_visit_ids)
await audit_crud.log_action(
db,
study_id=study_id,
@@ -209,6 +224,8 @@ async def update_visit(
old_status = visit.status
updated = await visit_crud.update_visit(db, visit, visit_in)
await subject_crud.sync_subject_status(db, subject)
if updated.actual_date is not None or updated.status in {"DONE", "CANCELLED"}:
await _resolve_visit_reminders(db, [updated.id])
detail = None
if visit_in.status:
detail = json.dumps(
@@ -253,6 +270,7 @@ async def delete_visit(
visit_detail = _visit_audit_detail("删除", subject, visit)
await visit_crud.delete_visit(db, visit)
await subject_crud.sync_subject_status(db, subject)
await _resolve_visit_reminders(db, [visit_id])
await audit_crud.log_action(
db,
study_id=study_id,
+43
View File
@@ -643,6 +643,43 @@ API_ENDPOINT_PERMISSIONS = {
"description": "删除文档",
"default_roles": ["PM"],
},
# 共享库在线协作(文档级邀请权限在接口权限之后继续校验)
"collaboration:create": {
"module": "shared_library",
"action": "write",
"description": "创建在线协作文件",
"default_roles": ["PM", "CRA", "PV", "QA"],
},
"collaboration:read": {
"module": "shared_library",
"action": "read",
"description": "查看在线协作文件",
"default_roles": ["PM", "CRA", "PV", "QA", "CTA"],
},
"collaboration:edit": {
"module": "shared_library",
"action": "write",
"description": "编辑受邀在线协作文件",
"default_roles": ["PM", "CRA", "PV", "QA", "CTA"],
},
"collaboration:manage": {
"module": "shared_library",
"action": "write",
"description": "管理在线协作文件成员与目录",
"default_roles": ["PM", "CRA", "PV", "QA"],
},
"collaboration:export": {
"module": "shared_library",
"action": "export",
"description": "另存为在线协作文件副本",
"default_roles": ["PM", "CRA", "PV", "QA"],
},
"collaboration:delete": {
"module": "shared_library",
"action": "write",
"description": "移入或恢复在线协作文件",
"default_roles": ["PM", "CRA", "PV", "QA"],
},
}
def _operation_read_candidates(operation_key: str) -> list[str]:
@@ -843,6 +880,7 @@ OPERATION_TO_ENDPOINTS: dict[str, dict[str, list[str]]] = {
"precautions:read",
"faq:read",
"faq_category:read",
"collaboration:read",
],
"write": [
"precautions:create",
@@ -858,6 +896,11 @@ OPERATION_TO_ENDPOINTS: dict[str, dict[str, list[str]]] = {
"faq_reply:create",
"faq_reply:delete",
"faq_attachments:delete",
"collaboration:create",
"collaboration:edit",
"collaboration:manage",
"collaboration:export",
"collaboration:delete",
],
},
}
+51
View File
@@ -1,5 +1,6 @@
from functools import lru_cache
from typing import Literal, Optional
from urllib.parse import urlsplit
from pydantic import Field
from pydantic_settings import BaseSettings, SettingsConfigDict
@@ -19,6 +20,7 @@ class Settings(BaseSettings):
JWT_EXPIRE_MINUTES: int = 60
JWT_EXTEND_GRACE_SECONDS: int = 120
ABSOLUTE_SESSION_MAX_HOURS: int = 8
DESKTOP_SESSION_MAX_DAYS: int = 30
LOGIN_RSA_PRIVATE_KEY: Optional[str] = None
LOGIN_RSA_PUBLIC_KEY: Optional[str] = None
LOGIN_RSA_KEY_ID: str = "default"
@@ -32,6 +34,31 @@ class Settings(BaseSettings):
)
IP2REGION_XDB_PATH: Optional[str] = None
IP2REGION_IPV6_XDB_PATH: Optional[str] = None
TRUSTED_PROXY_CIDRS: str = "127.0.0.1/32,::1/128,172.16.0.0/12"
MONITORING_SERVER_PUBLIC_IP: Optional[str] = None
MONITORING_PUBLIC_IP_DISCOVERY_URLS: str = (
"https://api64.ipify.org,https://icanhazip.com"
)
MONITORING_PUBLIC_IP_DISCOVERY_TIMEOUT_SECONDS: float = Field(default=2.5, ge=0.5, le=10)
MONITORING_SERVER_LOCATION_CACHE_SECONDS: int = Field(default=86400, ge=300, le=604800)
MONITORING_IP_GEO_FALLBACK_ENABLED: bool = True
MONITORING_IP_GEO_FALLBACK_API_KEY: Optional[str] = None
MONITORING_IP_GEO_FALLBACK_TIMEOUT_SECONDS: float = Field(default=2.5, ge=0.5, le=10)
MONITORING_IP_GEO_FALLBACK_CACHE_SECONDS: int = Field(default=604800, ge=3600, le=2592000)
MONITORING_IP_GEO_FALLBACK_MAX_LOOKUPS: int = Field(default=10, ge=1, le=100)
MONITORING_ACCESS_LOG_RETENTION_DAYS: int = Field(default=90, ge=7, le=3650)
MONITORING_METRIC_RETENTION_DAYS: int = Field(default=400, ge=30, le=3650)
MONITORING_RETENTION_INTERVAL_SECONDS: int = Field(default=86400, ge=60, le=604800)
USER_LOGIN_ACTIVITY_RETENTION_DAYS: int = Field(default=180, ge=30, le=3650)
USER_SESSION_ONLINE_SECONDS: int = Field(default=300, ge=60, le=3600)
NOTIFICATION_SYNC_INTERVAL_SECONDS: int = Field(default=300, ge=60, le=3600)
ONLYOFFICE_ENABLED: bool = False
ONLYOFFICE_JWT_SECRET: Optional[str] = None
ONLYOFFICE_INTERNAL_URL: str = "http://onlyoffice"
ONLYOFFICE_STORAGE_BASE_URL: str = "http://backend:8000"
ONLYOFFICE_INSTANCE_ID: Optional[str] = None
ONLYOFFICE_CONFIG_TTL_SECONDS: int = Field(default=300, ge=60, le=900)
COLLABORATION_MAX_FILE_BYTES: int = Field(default=50 * 1024 * 1024, ge=1024, le=500 * 1024 * 1024)
@lru_cache
@@ -42,6 +69,30 @@ def get_settings() -> Settings:
settings = get_settings()
def validate_onlyoffice_configuration() -> None:
if not settings.ONLYOFFICE_ENABLED:
return
secret = (settings.ONLYOFFICE_JWT_SECRET or "").strip()
instance_id = (settings.ONLYOFFICE_INSTANCE_ID or "").strip()
if secret == settings.JWT_SECRET_KEY:
raise RuntimeError("ONLYOFFICE_JWT_SECRET must not reuse JWT_SECRET_KEY")
if len(secret) < 32:
raise RuntimeError("ONLYOFFICE_JWT_SECRET must contain at least 32 characters when ONLYOFFICE is enabled")
if not instance_id:
raise RuntimeError("ONLYOFFICE_INSTANCE_ID is required when ONLYOFFICE is enabled")
for name, value in (
("ONLYOFFICE_INTERNAL_URL", settings.ONLYOFFICE_INTERNAL_URL),
("ONLYOFFICE_STORAGE_BASE_URL", settings.ONLYOFFICE_STORAGE_BASE_URL),
):
parsed = urlsplit(value.strip())
if parsed.scheme.lower() not in {"http", "https"} or not parsed.hostname:
raise RuntimeError(f"{name} must be an HTTP(S) URL")
if parsed.username or parsed.password or parsed.query or parsed.fragment:
raise RuntimeError(f"{name} must not contain credentials, a query, or a fragment")
if parsed.path not in {"", "/"}:
raise RuntimeError(f"{name} must not contain a path")
def get_cors_allowed_origins() -> list[str]:
return [
origin.strip()
+13 -3
View File
@@ -286,8 +286,9 @@ def _enqueue_permission_log(
writer = get_log_writer()
if writer:
forwarded = request.headers.get("x-forwarded-for")
ip = forwarded.split(",")[0].strip() if forwarded else (request.client.host if request.client else None)
from app.core.request_context import build_request_audit_context, get_request_audit_context
context = get_request_audit_context() or build_request_audit_context(request)
writer.enqueue({
"study_id": study_id,
"user_id": user_id,
@@ -295,7 +296,16 @@ def _enqueue_permission_log(
"role": role,
"allowed": allowed,
"elapsed_ms": elapsed_ms,
"ip_address": ip,
"ip_address": context.client_ip,
"user_agent": context.user_agent,
"client_type": context.client_type,
"client_version": context.client_version,
"client_platform": context.client_platform,
"build_channel": context.build_channel,
"build_commit": context.build_commit,
"request_headers": context.request_headers,
"request_snapshot": context.request_snapshot,
"request_id": context.request_id,
})
+254
View File
@@ -0,0 +1,254 @@
"""Request-scoped metadata used by server-side audit writers."""
from __future__ import annotations
import ipaddress
import re
import uuid
from contextvars import ContextVar, Token
from dataclasses import dataclass
from typing import Any
from app.core.config import settings
@dataclass(frozen=True)
class RequestAuditContext:
request_id: str | None = None
client_ip: str | None = None
user_agent: str | None = None
client_type: str | None = None
client_version: str | None = None
client_platform: str | None = None
build_channel: str | None = None
build_commit: str | None = None
request_headers: dict[str, str] | None = None
request_snapshot: dict[str, Any] | None = None
_request_audit_context: ContextVar[RequestAuditContext | None] = ContextVar(
"request_audit_context",
default=None,
)
_SENSITIVE_TEXT_PATTERN = re.compile(
r"(?i)(bearer\s+)[A-Za-z0-9._~+/=-]+|((?:access_)?token|authorization|password|passwd|secret|credential|api[-_]?key|session)=([^&\s]+)"
)
_SENSITIVE_HEADER_NAME_PATTERN = re.compile(
r"(?i)(authorization|cookie|set-cookie|token|password|passwd|secret|credential|api[-_]?key|session)"
)
_SAFE_REQUEST_HEADER_NAMES = frozenset(
{
"accept",
"accept-language",
"content-type",
"host",
"origin",
"user-agent",
"x-request-id",
"x-correlation-id",
}
)
_SAFE_REQUEST_HEADER_PREFIXES = ("x-ctms-",)
_SENSITIVE_PARAMETER_NAME_PATTERN = re.compile(
r"(?i)(token|authorization|password|passwd|secret|credential|api[-_]?key|session|"
r"subject|participant|patient|user_?name|full_?name|email|phone|mobile|id_?card|"
r"identity|certificate|contact|address)"
)
_KNOWN_CLIENT_TYPES = frozenset({"web", "desktop"})
_CTMS_CLIENT_SOURCE_TO_TYPE = {
"ctms-web": "web",
"ctms-desktop": "desktop",
}
def _clean_audit_value(value: str | None, max_length: int) -> str | None:
if value is None:
return None
cleaned = _SENSITIVE_TEXT_PATTERN.sub(lambda m: f"{m.group(1) or m.group(2) + '='}[redacted]", value.strip())
if not cleaned:
return None
return cleaned[:max_length]
def _clean_header_value(value: str | None, max_length: int) -> str | None:
return _clean_audit_value(value, max_length)
def resolve_ctms_client_type(headers: Any) -> str | None:
source = _clean_header_value(headers.get("x-ctms-client-source"), 32)
if source:
mapped_type = _CTMS_CLIENT_SOURCE_TO_TYPE.get(source.strip().lower())
if mapped_type:
return mapped_type
client_type = _clean_header_value(headers.get("x-ctms-client-type"), 16)
if not client_type:
return None
normalized = client_type.strip().lower()
return normalized if normalized in _KNOWN_CLIENT_TYPES else normalized[:16]
def build_sanitized_request_headers(request: Any) -> dict[str, str] | None:
captured: dict[str, str] = {}
for raw_name, raw_value in request.headers.items():
name = str(raw_name).strip().lower()
if not name:
continue
if _SENSITIVE_HEADER_NAME_PATTERN.search(name):
captured[name] = "[redacted]"
continue
if name not in _SAFE_REQUEST_HEADER_NAMES and not name.startswith(_SAFE_REQUEST_HEADER_PREFIXES):
continue
value = _clean_header_value(str(raw_value), 500)
if value:
captured[name] = value
return captured or None
def _sanitize_named_value(raw_name: Any, raw_value: Any, max_length: int = 500) -> dict[str, str] | None:
name = _clean_audit_value(str(raw_name), 120)
if not name:
return None
if _SENSITIVE_PARAMETER_NAME_PATTERN.search(name):
return {"name": name, "value": "[redacted]"}
value = _clean_audit_value(str(raw_value), max_length)
if value is None:
return {"name": name, "value": ""}
return {"name": name, "value": value}
def _query_param_items(request: Any) -> list[dict[str, str]] | None:
query_params = getattr(request, "query_params", None)
if not query_params:
return None
if hasattr(query_params, "multi_items"):
raw_items = query_params.multi_items()
else:
raw_items = query_params.items()
items = [
item
for item in (_sanitize_named_value(name, value) for name, value in raw_items)
if item is not None
][:50]
return items or None
def _sanitized_query_string(items: list[dict[str, str]] | None) -> str | None:
if not items:
return None
value = "&".join(f"{item['name']}={item['value']}" for item in items)
return value[:2000] or None
def _request_client_snapshot(request: Any) -> dict[str, str | int | None] | None:
client = getattr(request, "client", None)
if not client:
return None
host = _clean_audit_value(getattr(client, "host", None), 120)
port = getattr(client, "port", None)
if host is None and port is None:
return None
return {"host": host, "port": port}
def build_request_snapshot(request: Any, *, request_id: str | None = None) -> dict[str, Any]:
headers = getattr(request, "headers", {})
scope = getattr(request, "scope", {}) or {}
url = getattr(request, "url", None)
query_params = _query_param_items(request)
path = getattr(url, "path", None) or scope.get("path")
method = getattr(request, "method", None) or scope.get("method")
snapshot: dict[str, Any] = {
"request_id": request_id,
"method": _clean_audit_value(str(method).upper() if method else None, 12),
"path": _clean_audit_value(path, 500),
"query_string": _sanitized_query_string(query_params),
"query_params": query_params,
"headers": build_sanitized_request_headers(request),
"http_version": _clean_audit_value(scope.get("http_version"), 16),
"scheme": _clean_audit_value(getattr(url, "scheme", None) or scope.get("scheme"), 16),
"client": _request_client_snapshot(request),
"content": {
"type": _clean_audit_value(headers.get("content-type"), 200),
"length": _clean_audit_value(headers.get("content-length"), 32),
},
"body": {
"captured": False,
"reason": "body_not_captured_by_audit_policy",
},
}
return {key: value for key, value in snapshot.items() if value not in (None, {}, [])}
def _parse_ip(value: str | None):
try:
return ipaddress.ip_address((value or "").strip())
except ValueError:
return None
def _trusted_proxy_networks():
networks = []
for value in settings.TRUSTED_PROXY_CIDRS.split(","):
candidate = value.strip()
if not candidate:
continue
try:
networks.append(ipaddress.ip_network(candidate, strict=False))
except ValueError:
continue
return tuple(networks)
def _is_trusted_proxy(value: str | None) -> bool:
address = _parse_ip(value)
return bool(address and any(address in network for network in _trusted_proxy_networks()))
def resolve_client_ip(request: Any) -> str | None:
peer_ip = request.client.host if request.client else None
if not _is_trusted_proxy(peer_ip):
return _clean_header_value(peer_ip, 45)
forwarded = request.headers.get("x-forwarded-for")
if forwarded:
chain = [part.strip() for part in forwarded.split(",") if _parse_ip(part)]
chain.append(str(peer_ip))
for candidate in reversed(chain):
if not _is_trusted_proxy(candidate):
return _clean_header_value(candidate, 45)
real_ip = request.headers.get("x-real-ip")
if _parse_ip(real_ip):
return _clean_header_value(real_ip.strip(), 45)
return _clean_header_value(peer_ip, 45)
def build_request_audit_context(request: Any) -> RequestAuditContext:
headers = request.headers
request_id = str(uuid.uuid4())
return RequestAuditContext(
request_id=request_id,
client_ip=resolve_client_ip(request),
user_agent=_clean_header_value(headers.get("user-agent"), 500),
client_type=resolve_ctms_client_type(headers),
client_version=_clean_header_value(headers.get("x-ctms-client-version"), 32),
client_platform=_clean_header_value(headers.get("x-ctms-client-platform"), 16),
build_channel=_clean_header_value(headers.get("x-ctms-build-channel"), 16),
build_commit=_clean_header_value(headers.get("x-ctms-build-commit"), 64),
request_headers=build_sanitized_request_headers(request),
request_snapshot=build_request_snapshot(request, request_id=request_id),
)
def set_request_audit_context(context: RequestAuditContext) -> Token[RequestAuditContext | None]:
return _request_audit_context.set(context)
def reset_request_audit_context(token: Token[RequestAuditContext | None]) -> None:
_request_audit_context.reset(token)
def get_request_audit_context() -> RequestAuditContext | None:
return _request_audit_context.get()
+17 -1
View File
@@ -19,16 +19,32 @@ def create_access_token(
user_id: str,
expires_minutes: Optional[int] = None,
session_start: Optional[datetime] = None,
max_age_seconds: Optional[int] = None,
issued_at: Optional[datetime] = None,
client_type: Optional[str] = None,
session_id: Optional[str] = None,
) -> str:
now = datetime.now(timezone.utc)
now = issued_at or datetime.now(timezone.utc)
if now.tzinfo is None:
now = now.replace(tzinfo=timezone.utc)
expire = now + timedelta(minutes=expires_minutes or settings.JWT_EXPIRE_MINUTES)
session_start_time = session_start or now
if session_start_time.tzinfo is None:
session_start_time = session_start_time.replace(tzinfo=timezone.utc)
if max_age_seconds is not None:
session_expire = session_start_time + timedelta(seconds=max_age_seconds)
if expire > session_expire:
expire = session_expire
to_encode: Dict[str, Any] = {
"sub": user_id,
"exp": expire,
"iat": int(now.timestamp()),
"orig_iat": int(session_start_time.timestamp()),
}
if client_type:
to_encode["client_type"] = client_type
if session_id:
to_encode["sid"] = session_id
return jwt.encode(to_encode, settings.JWT_SECRET_KEY, algorithm=ALGORITHM)
+39
View File
@@ -4,8 +4,10 @@ import uuid
from typing import Sequence
from sqlalchemy import select
from sqlalchemy import or_
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.request_context import get_request_audit_context
from app.models.audit_log import AuditLog
@@ -19,8 +21,16 @@ async def log_action(
detail: str | None,
operator_id: uuid.UUID,
operator_role: str,
client_ip: str | None = None,
user_agent: str | None = None,
client_type: str | None = None,
client_version: str | None = None,
client_platform: str | None = None,
build_channel: str | None = None,
build_commit: str | None = None,
auto_commit: bool = True,
) -> AuditLog:
request_context = get_request_audit_context()
log = AuditLog(
study_id=study_id,
entity_type=entity_type,
@@ -29,6 +39,13 @@ async def log_action(
detail=detail,
operator_id=operator_id,
operator_role=operator_role,
client_ip=client_ip if client_ip is not None else (request_context.client_ip if request_context else None),
user_agent=user_agent if user_agent is not None else (request_context.user_agent if request_context else None),
client_type=client_type if client_type is not None else (request_context.client_type if request_context else None),
client_version=client_version if client_version is not None else (request_context.client_version if request_context else None),
client_platform=client_platform if client_platform is not None else (request_context.client_platform if request_context else None),
build_channel=build_channel if build_channel is not None else (request_context.build_channel if request_context else None),
build_commit=build_commit if build_commit is not None else (request_context.build_commit if request_context else None),
)
db.add(log)
if auto_commit:
@@ -47,6 +64,10 @@ async def list_logs(
entity_id: uuid.UUID | None = None,
action: str | None = None,
operator_id: uuid.UUID | None = None,
client_ip: str | None = None,
client_type: str | None = None,
start_time=None,
end_time=None,
skip: int = 0,
limit: int = 100,
) -> Sequence[AuditLog]:
@@ -59,6 +80,24 @@ async def list_logs(
stmt = stmt.where(AuditLog.action == action)
if operator_id:
stmt = stmt.where(AuditLog.operator_id == operator_id)
if client_ip:
stmt = stmt.where(AuditLog.client_ip.ilike(f"%{client_ip.strip()}%"))
if client_type:
normalized_client_type = client_type.strip().lower()
if normalized_client_type == "unknown":
stmt = stmt.where(
or_(
AuditLog.client_type.is_(None),
AuditLog.client_type == "",
~AuditLog.client_type.in_(("web", "desktop")),
)
)
else:
stmt = stmt.where(AuditLog.client_type == normalized_client_type)
if start_time:
stmt = stmt.where(AuditLog.created_at >= start_time)
if end_time:
stmt = stmt.where(AuditLog.created_at <= end_time)
stmt = stmt.order_by(AuditLog.created_at.desc()).offset(skip).limit(limit)
result = await db.execute(stmt)
return result.scalars().all()
+32 -4
View File
@@ -1,7 +1,8 @@
from __future__ import annotations
import uuid
from typing import Sequence
from datetime import datetime, timedelta, timezone
from typing import Literal, Sequence
from sqlalchemy import delete, func, or_, select, update
from sqlalchemy.ext.asyncio import AsyncSession
@@ -11,12 +12,14 @@ from app.core.config import (
PROTECTED_ADMIN_DEFAULT_PASSWORD,
PROTECTED_ADMIN_EMAIL,
PROTECTED_ADMIN_FULL_NAME,
settings,
)
from app.core.security import hash_password
from app.models.audit_log import AuditLog
from app.models.permission_access_log import PermissionAccessLog
from app.models.study_member import StudyMember
from app.models.user import User, UserStatus
from app.models.user_login_session import UserLoginSession
from app.schemas.user import UserCreate, UserRegisterRequest, UserUpdate
@@ -81,7 +84,13 @@ async def update_user(db: AsyncSession, user: User, user_in: UserUpdate) -> User
return user
def _apply_user_filters(query, *, keyword: str | None = None, status: UserStatus | None = None):
def _apply_user_filters(
query,
*,
keyword: str | None = None,
status: UserStatus | None = None,
login_status: Literal["ONLINE", "OFFLINE"] | None = None,
):
if keyword:
pattern = f"%{keyword.strip()}%"
query = query.where(
@@ -93,6 +102,18 @@ def _apply_user_filters(query, *, keyword: str | None = None, status: UserStatus
)
if status is not None:
query = query.where(User.status == status)
if login_status is not None:
cutoff = datetime.now(timezone.utc) - timedelta(seconds=settings.USER_SESSION_ONLINE_SECONDS)
has_online_session = (
select(UserLoginSession.id)
.where(
UserLoginSession.user_id == User.id,
UserLoginSession.ended_at.is_(None),
UserLoginSession.last_seen_at >= cutoff,
)
.exists()
)
query = query.where(has_online_session if login_status == "ONLINE" else ~has_online_session)
return query
@@ -103,8 +124,9 @@ async def list_users(
*,
keyword: str | None = None,
status: UserStatus | None = None,
login_status: Literal["ONLINE", "OFFLINE"] | None = None,
) -> Sequence[User]:
query = _apply_user_filters(select(User), keyword=keyword, status=status)
query = _apply_user_filters(select(User), keyword=keyword, status=status, login_status=login_status)
result = await db.execute(query.order_by(User.created_at.desc()).offset(skip).limit(limit))
return result.scalars().all()
@@ -114,8 +136,14 @@ async def count_users(
*,
keyword: str | None = None,
status: UserStatus | None = None,
login_status: Literal["ONLINE", "OFFLINE"] | None = None,
) -> int:
query = _apply_user_filters(select(func.count()).select_from(User), keyword=keyword, status=status)
query = _apply_user_filters(
select(func.count()).select_from(User),
keyword=keyword,
status=status,
login_status=login_status,
)
result = await db.execute(query)
return int(result.scalar_one() or 0)
+13
View File
@@ -10,6 +10,16 @@ from app.models.audit_log import AuditLog # noqa: F401
from app.models.etmf import EtmfNode # noqa: F401
from app.models.document import Document # noqa: F401
from app.models.document_version import DocumentVersion # noqa: F401
from app.models.collaboration import ( # noqa: F401
CollaborationCallbackReceipt,
CollaborationEditRequest,
CollaborationFile,
CollaborationFolder,
CollaborationMember,
CollaborationRevision,
CollaborationSession,
CollaborationShareLink,
)
from app.models.distribution import Distribution # noqa: F401
from app.models.acknowledgement import Acknowledgement # noqa: F401
from app.models.milestone import Milestone # noqa: F401
@@ -42,8 +52,11 @@ from app.models.permission_access_log import PermissionAccessLog # noqa: F401
from app.models.permission_metric_snapshot import PermissionMetricSnapshot # noqa: F401
from app.models.permission_template import PermissionTemplate, PermissionTemplateVersion # noqa: F401
from app.models.security_access_log import SecurityAccessLog # noqa: F401
from app.models.source_location_snapshot import SourceLocationSnapshot # noqa: F401
from app.models.user_login_session import UserLoginSession # noqa: F401
from app.models.desktop_notification import ( # noqa: F401
DesktopNotificationDelivery,
DesktopNotificationSubscription,
)
from app.models.notification import Notification # noqa: F401
from app.models.email_settings import EmailVerificationCode, SystemEmailSettings # noqa: F401
+126 -61
View File
@@ -5,26 +5,44 @@ import time
from contextlib import asynccontextmanager
from collections import defaultdict
from fastapi import FastAPI
from fastapi import Depends, FastAPI
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from sqlalchemy import text
from app.api.v1.router import api_router
from app.core.config import get_cors_allowed_origins, settings
from app.api.v1.onlyoffice import internal_router as onlyoffice_internal_router
from app.api.v1.collaboration import internal_router as collaboration_internal_router
from app.core.config import get_cors_allowed_origins, settings, validate_onlyoffice_configuration
from app.core.exceptions import register_exception_handlers
from app.core.login_crypto import validate_login_crypto_configuration
from app.crud.user import ensure_admin_exists
from app.db.base import Base
from app.db.session import SessionLocal, engine
from app.services.visit_scheduler import run_daily_lost_visit_job
from app.services.notification_scheduler import run_notification_sync_job
from app.services.permission_log_writer import start_log_writer, stop_log_writer
from app.services.permission_metric_aggregator import run_hourly_metric_aggregation
from app.services.source_location_aggregator import run_hourly_source_location_aggregation
from app.services.monitoring_server_location import resolve_monitoring_server_location
from app.services.monitoring_retention import run_monitoring_retention
from app.services.security_access_log_writer import (
get_security_log_writer,
start_security_log_writer,
stop_security_log_writer,
)
from app.core.security import decode_token
from app.core.deps import get_db_session
from app.core.request_context import (
build_request_audit_context,
build_request_snapshot,
build_sanitized_request_headers,
get_request_audit_context,
reset_request_audit_context,
resolve_ctms_client_type,
resolve_client_ip,
set_request_audit_context,
)
logger = logging.getLogger("ctms.setup_config")
UUID_RE = re.compile(
@@ -36,8 +54,12 @@ setup_config_stats: dict[str, int] = defaultdict(int)
@asynccontextmanager
async def lifespan(_: FastAPI):
stop_event = asyncio.Event()
await resolve_monitoring_server_location()
scheduler_task = asyncio.create_task(run_daily_lost_visit_job(stop_event))
aggregator_task = asyncio.create_task(run_hourly_metric_aggregation(stop_event))
source_location_aggregator_task = asyncio.create_task(
run_hourly_source_location_aggregation(stop_event)
)
await start_log_writer()
await start_security_log_writer()
if settings.ENV == "development":
@@ -46,12 +68,17 @@ async def lifespan(_: FastAPI):
await conn.run_sync(Base.metadata.create_all)
async with SessionLocal() as session:
await ensure_admin_exists(session)
notification_scheduler_task = asyncio.create_task(run_notification_sync_job(stop_event))
retention_task = asyncio.create_task(run_monitoring_retention(stop_event))
yield
stop_event.set()
await stop_log_writer()
await stop_security_log_writer()
await scheduler_task
await aggregator_task
await source_location_aggregator_task
await notification_scheduler_task
await retention_task
async def _ensure_legacy_primary_keys(conn) -> None:
@@ -89,6 +116,7 @@ async def _ensure_legacy_primary_keys(conn) -> None:
def create_app() -> FastAPI:
validate_login_crypto_configuration()
validate_onlyoffice_configuration()
app = FastAPI(
title="CTMS 后端 API",
description="临床试验项目管理系统后端接口文档",
@@ -123,74 +151,85 @@ def create_app() -> FastAPI:
"Accept",
"Authorization",
"Content-Type",
"X-CTMS-Client-Source",
"X-CTMS-Client-Type",
"X-CTMS-Client-Version",
"X-CTMS-Client-Platform",
"X-CTMS-Build-Channel",
"X-CTMS-Build-Commit",
"X-Request-ID",
"X-Correlation-ID",
],
expose_headers=["X-Request-ID"],
)
@app.middleware("http")
async def setup_config_monitoring_middleware(request, call_next):
path = request.url.path
is_setup_config_path = "/api/v1/studies/" in path and "/setup-config" in path
should_security_log = path.startswith("/api/")
should_security_log = path.startswith("/api/") and path != "/api/v1/auth/session/heartbeat"
started_at = time.perf_counter()
audit_context = build_request_audit_context(request)
audit_context_token = set_request_audit_context(audit_context)
status_code = 500
try:
response = await call_next(request)
except Exception:
try:
response = await call_next(request)
except Exception:
if is_setup_config_path:
normalized_path = UUID_RE.sub("{study_id}", path)
duration_ms = int((time.perf_counter() - started_at) * 1000)
key = f"{request.method} {normalized_path} 5xx"
setup_config_stats[key] += 1
logger.exception(
"setup_config_request_error method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
500,
duration_ms,
)
if should_security_log:
_enqueue_security_access_log(request, path, status_code, started_at)
raise
status_code = int(response.status_code)
if audit_context.request_id:
response.headers["X-Request-ID"] = audit_context.request_id
if is_setup_config_path:
normalized_path = UUID_RE.sub("{study_id}", path)
duration_ms = int((time.perf_counter() - started_at) * 1000)
key = f"{request.method} {normalized_path} 5xx"
status_bucket = f"{status_code // 100}xx"
key = f"{request.method} {normalized_path} {status_bucket}"
setup_config_stats[key] += 1
logger.exception(
"setup_config_request_error method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
500,
duration_ms,
)
if status_code >= 500:
logger.error(
"setup_config_request method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
status_code,
duration_ms,
)
elif status_code >= 400:
logger.warning(
"setup_config_request method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
status_code,
duration_ms,
)
else:
logger.info(
"setup_config_request method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
status_code,
duration_ms,
)
if should_security_log:
_enqueue_security_access_log(request, path, status_code, started_at)
raise
status_code = int(response.status_code)
if is_setup_config_path:
normalized_path = UUID_RE.sub("{study_id}", path)
duration_ms = int((time.perf_counter() - started_at) * 1000)
status_bucket = f"{status_code // 100}xx"
key = f"{request.method} {normalized_path} {status_bucket}"
setup_config_stats[key] += 1
if status_code >= 500:
logger.error(
"setup_config_request method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
status_code,
duration_ms,
)
elif status_code >= 400:
logger.warning(
"setup_config_request method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
status_code,
duration_ms,
)
else:
logger.info(
"setup_config_request method=%s path=%s status=%s duration_ms=%s",
request.method,
normalized_path,
status_code,
duration_ms,
)
if should_security_log:
_enqueue_security_access_log(request, path, status_code, started_at)
return response
return response
finally:
reset_request_audit_context(audit_context_token)
register_exception_handlers(app)
@@ -218,6 +257,36 @@ def create_app() -> FastAPI:
async def health() -> dict[str, str]:
return {"status": "ok"}
@app.get(
"/readyz",
tags=["health"],
summary="服务就绪检查",
description="验证应用可访问数据库;失败时返回 503。",
)
async def readiness(db=Depends(get_db_session)):
started_at = time.perf_counter()
try:
await db.execute(text("SELECT 1"))
except Exception as exc:
logger.exception("Readiness database check failed")
return JSONResponse(
status_code=503,
content={
"status": "not_ready",
"database": {
"status": "unhealthy",
"error_type": type(exc).__name__,
},
},
)
return {
"status": "ready",
"database": {
"status": "healthy",
"latency_ms": round((time.perf_counter() - started_at) * 1000, 2),
},
}
@app.get(
"/health/setup-config-stats",
tags=["health"],
@@ -234,22 +303,14 @@ def create_app() -> FastAPI:
return {"totals": totals, "by_endpoint": summary}
app.include_router(api_router, prefix="/api/v1")
app.include_router(onlyoffice_internal_router)
app.include_router(collaboration_internal_router)
return app
app = create_app()
def _resolve_client_ip(request) -> str | None:
forwarded = request.headers.get("x-forwarded-for")
if forwarded:
return forwarded.split(",")[0].strip()
real_ip = request.headers.get("x-real-ip")
if real_ip:
return real_ip.strip()
return request.client.host if request.client else None
def _resolve_auth_context(request) -> tuple[str, str | None]:
authorization = request.headers.get("authorization") or ""
if not authorization.lower().startswith("bearer "):
@@ -270,19 +331,23 @@ def _enqueue_security_access_log(request, path: str, status_code: int, started_a
if not writer:
return
auth_status, user_identifier = _resolve_auth_context(request)
context = get_request_audit_context()
writer.enqueue(
{
"method": request.method,
"path": path,
"status_code": status_code,
"elapsed_ms": round((time.perf_counter() - started_at) * 1000, 2),
"client_ip": _resolve_client_ip(request),
"client_ip": resolve_client_ip(request),
"user_agent": request.headers.get("user-agent"),
"client_type": request.headers.get("x-ctms-client-type"),
"client_type": resolve_ctms_client_type(request.headers),
"client_version": request.headers.get("x-ctms-client-version"),
"client_platform": request.headers.get("x-ctms-client-platform"),
"build_channel": request.headers.get("x-ctms-build-channel"),
"build_commit": request.headers.get("x-ctms-build-commit"),
"request_headers": context.request_headers if context else build_sanitized_request_headers(request),
"request_snapshot": context.request_snapshot if context else build_request_snapshot(request),
"request_id": context.request_id if context else None,
"auth_status": auth_status,
"user_identifier": user_identifier,
}
+13 -1
View File
@@ -13,7 +13,12 @@ from app.db.base_class import Base
class AuditLog(Base):
__tablename__ = "audit_logs"
__table_args__ = (Index("ix_audit_logs_entity_at", "entity_type", "entity_id", "created_at"),)
__table_args__ = (
Index("ix_audit_logs_entity_at", "entity_type", "entity_id", "created_at"),
Index("ix_audit_logs_operator_created", "operator_id", "created_at"),
Index("ix_audit_logs_client_ip_created", "client_ip", "created_at"),
Index("ix_audit_logs_client_source_created", "client_type", "created_at"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
study_id: Mapped[Optional[uuid.UUID]] = mapped_column(UUID(as_uuid=True), ForeignKey("studies.id"), nullable=True)
@@ -23,4 +28,11 @@ class AuditLog(Base):
detail: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
operator_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
operator_role: Mapped[str] = mapped_column(String(50), nullable=False)
client_ip: Mapped[Optional[str]] = mapped_column(String(45), nullable=True)
user_agent: Mapped[Optional[str]] = mapped_column(String(500), nullable=True)
client_type: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
client_version: Mapped[Optional[str]] = mapped_column(String(32), nullable=True)
client_platform: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
build_channel: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
build_commit: Mapped[Optional[str]] = mapped_column(String(64), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
+212
View File
@@ -0,0 +1,212 @@
from __future__ import annotations
import uuid
from datetime import datetime
from typing import Optional
from sqlalchemy import BigInteger, Boolean, DateTime, ForeignKey, Index, Integer, String, Text, UniqueConstraint, func, text
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import Mapped, mapped_column
from app.db.base_class import Base
class CollaborationFolder(Base):
__tablename__ = "collaboration_folders"
__table_args__ = (
Index("ix_collaboration_folders_study_parent", "study_id", "parent_id"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
study_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("studies.id"), nullable=False)
parent_id: Mapped[Optional[uuid.UUID]] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_folders.id", ondelete="SET NULL"), nullable=True
)
name: Mapped[str] = mapped_column(String(120), nullable=False)
sort_order: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
created_by: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
deleted_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
)
class CollaborationFile(Base):
__tablename__ = "collaboration_files"
__table_args__ = (
Index("ix_collaboration_files_study_folder", "study_id", "folder_id"),
Index("ix_collaboration_files_study_status", "study_id", "status"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
study_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("studies.id"), nullable=False)
folder_id: Mapped[Optional[uuid.UUID]] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_folders.id", ondelete="SET NULL"), nullable=True
)
title: Mapped[str] = mapped_column(String(255), nullable=False)
file_type: Mapped[str] = mapped_column(String(16), nullable=False)
extension: Mapped[str] = mapped_column(String(16), nullable=False)
status: Mapped[str] = mapped_column(String(20), nullable=False, default="ACTIVE", server_default="ACTIVE")
owner_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
current_revision_id: Mapped[Optional[uuid.UUID]] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_revisions.id", ondelete="SET NULL"), nullable=True
)
generation: Mapped[int] = mapped_column(Integer, nullable=False, default=1, server_default="1")
allow_export: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False, server_default="false")
allow_edit_request: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False, server_default="false")
allow_sheet_structure_edit: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True, server_default="true")
sheet_structure_protection_backup: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
deleted_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
)
class CollaborationMember(Base):
__tablename__ = "collaboration_members"
__table_args__ = (
UniqueConstraint("file_id", "user_id", name="uq_collaboration_member_file_user"),
Index("ix_collaboration_members_user", "user_id"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
file_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_files.id", ondelete="CASCADE"), nullable=False
)
user_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
role: Mapped[str] = mapped_column(String(16), nullable=False)
invited_by: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
)
class CollaborationEditRequest(Base):
__tablename__ = "collaboration_edit_requests"
__table_args__ = (
Index("ix_collaboration_edit_requests_file_status", "file_id", "status"),
Index(
"uq_collaboration_edit_requests_pending_user",
"file_id",
"requester_id",
unique=True,
postgresql_where=text("status = 'PENDING'"),
),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
file_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_files.id", ondelete="CASCADE"), nullable=False
)
requester_id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
status: Mapped[str] = mapped_column(String(16), nullable=False, default="PENDING", server_default="PENDING")
resolved_by: Mapped[Optional[uuid.UUID]] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=True)
resolved_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
)
class CollaborationShareLink(Base):
__tablename__ = "collaboration_share_links"
__table_args__ = (
UniqueConstraint("file_id", name="uq_collaboration_share_link_file"),
Index("ix_collaboration_share_links_enabled_expiry", "enabled", "expires_at"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
file_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_files.id", ondelete="CASCADE"), nullable=False
)
enabled: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False, server_default="false")
access_mode: Mapped[str] = mapped_column(String(12), nullable=False, default="VIEW", server_default="VIEW")
expiry_policy: Mapped[str] = mapped_column(
String(16), nullable=False, default="SEVEN_DAYS", server_default="SEVEN_DAYS"
)
expires_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
password_hash: Mapped[Optional[str]] = mapped_column(String(255), nullable=True)
token_version: Mapped[int] = mapped_column(Integer, nullable=False, default=1, server_default="1")
failed_attempts: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default="0")
last_failed_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
locked_until: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
created_by: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
updated_by: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
)
class CollaborationRevision(Base):
__tablename__ = "collaboration_revisions"
__table_args__ = (
UniqueConstraint("file_id", "revision_no", name="uq_collaboration_revision_file_no"),
Index("ix_collaboration_revisions_file_created", "file_id", "created_at"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
file_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_files.id", ondelete="CASCADE"), nullable=False
)
revision_no: Mapped[int] = mapped_column(Integer, nullable=False)
parent_revision_id: Mapped[Optional[uuid.UUID]] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_revisions.id", ondelete="SET NULL"), nullable=True
)
file_uri: Mapped[str] = mapped_column(String(500), nullable=False)
original_filename: Mapped[str] = mapped_column(String(255), nullable=False)
file_hash: Mapped[str] = mapped_column(String(128), nullable=False)
file_size: Mapped[int] = mapped_column(BigInteger, nullable=False)
mime_type: Mapped[str] = mapped_column(String(100), nullable=False)
source: Mapped[str] = mapped_column(String(24), nullable=False)
change_summary: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
created_by: Mapped[Optional[uuid.UUID]] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
deleted_by: Mapped[Optional[uuid.UUID]] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=True)
deleted_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
class CollaborationSession(Base):
__tablename__ = "collaboration_sessions"
__table_args__ = (
UniqueConstraint("document_key", name="uq_collaboration_session_document_key"),
Index("ix_collaboration_sessions_file_status", "file_id", "status"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
file_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_files.id", ondelete="CASCADE"), nullable=False
)
base_revision_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_revisions.id"), nullable=False
)
document_key: Mapped[str] = mapped_column(String(128), nullable=False)
generation: Mapped[int] = mapped_column(Integer, nullable=False)
status: Mapped[str] = mapped_column(String(20), nullable=False, default="ACTIVE", server_default="ACTIVE")
started_by: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), ForeignKey("users.id"), nullable=False)
active_users: Mapped[Optional[str]] = mapped_column(Text, nullable=True)
last_callback_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
closed_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
class CollaborationCallbackReceipt(Base):
__tablename__ = "collaboration_callback_receipts"
__table_args__ = (
UniqueConstraint("session_id", "fingerprint", name="uq_collaboration_callback_session_fingerprint"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
session_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_sessions.id", ondelete="CASCADE"), nullable=False
)
fingerprint: Mapped[str] = mapped_column(String(128), nullable=False)
callback_status: Mapped[int] = mapped_column(Integer, nullable=False)
result: Mapped[str] = mapped_column(String(24), nullable=False)
revision_id: Mapped[Optional[uuid.UUID]] = mapped_column(
UUID(as_uuid=True), ForeignKey("collaboration_revisions.id", ondelete="SET NULL"), nullable=True
)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
+6 -2
View File
@@ -28,6 +28,7 @@ class DesktopNotificationDelivery(Base):
__tablename__ = "desktop_notification_deliveries"
__table_args__ = (
UniqueConstraint("user_id", "distribution_id", name="uq_desktop_notification_user_distribution"),
UniqueConstraint("user_id", "notification_id", name="uq_desktop_notification_user_notification"),
Index("ix_desktop_notification_claim", "user_id", "delivered_at", "claimed_at"),
)
@@ -35,8 +36,11 @@ class DesktopNotificationDelivery(Base):
user_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("users.id", ondelete="CASCADE"), nullable=False
)
distribution_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("distributions.id", ondelete="CASCADE"), nullable=False
distribution_id: Mapped[uuid.UUID | None] = mapped_column(
UUID(as_uuid=True), ForeignKey("distributions.id", ondelete="CASCADE"), nullable=True
)
notification_id: Mapped[uuid.UUID | None] = mapped_column(
UUID(as_uuid=True), ForeignKey("notifications.id", ondelete="CASCADE"), nullable=True
)
claim_token: Mapped[uuid.UUID | None] = mapped_column(UUID(as_uuid=True), nullable=True)
claimed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
+1
View File
@@ -47,6 +47,7 @@ class DocumentVersion(Base):
effective_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
superseded_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True)
file_uri: Mapped[str] = mapped_column(String(500), nullable=False)
original_filename: Mapped[Optional[str]] = mapped_column(String(255), nullable=True)
file_hash: Mapped[str] = mapped_column(String(128), nullable=False)
file_size: Mapped[int] = mapped_column(BigInteger, nullable=False)
mime_type: Mapped[Optional[str]] = mapped_column(String(100), nullable=True)
+53
View File
@@ -0,0 +1,53 @@
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import Boolean, DateTime, ForeignKey, Index, String, Text, UniqueConstraint, func
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import Mapped, mapped_column
from app.db.base_class import Base
class Notification(Base):
__tablename__ = "notifications"
__table_args__ = (
UniqueConstraint("recipient_id", "dedupe_key", name="uq_notifications_recipient_dedupe"),
Index(
"ix_notifications_recipient_study_state",
"recipient_id",
"study_id",
"resolved_at",
"read_at",
"created_at",
),
Index("ix_notifications_source", "source_type", "source_id"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
study_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("studies.id", ondelete="CASCADE"), nullable=False
)
recipient_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True), ForeignKey("users.id", ondelete="CASCADE"), nullable=False
)
category: Mapped[str] = mapped_column(String(64), nullable=False)
priority: Mapped[str] = mapped_column(String(16), nullable=False, default="NORMAL", server_default="NORMAL")
title: Mapped[str] = mapped_column(String(180), nullable=False)
message: Mapped[str] = mapped_column(String(500), nullable=False)
action_path: Mapped[str | None] = mapped_column(Text, nullable=True)
source_type: Mapped[str] = mapped_column(String(64), nullable=False)
source_id: Mapped[str] = mapped_column(String(100), nullable=False)
source_version: Mapped[str | None] = mapped_column(String(100), nullable=True)
dedupe_key: Mapped[str] = mapped_column(String(255), nullable=False)
requires_action: Mapped[bool] = mapped_column(
Boolean, nullable=False, default=True, server_default="true"
)
due_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
read_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
resolved_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now(), onupdate=func.now()
)
+16 -2
View File
@@ -6,13 +6,15 @@ import uuid
from datetime import datetime
from typing import Optional
from sqlalchemy import Boolean, DateTime, Float, ForeignKey, Index, String
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy import JSON, Boolean, DateTime, Float, ForeignKey, Index, String
from sqlalchemy.dialects.postgresql import JSONB, UUID
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.sql import func
from app.db.base_class import Base
JSONB_TYPE = JSON().with_variant(JSONB, "postgresql")
class PermissionAccessLog(Base):
__tablename__ = "permission_access_logs"
@@ -22,6 +24,9 @@ class PermissionAccessLog(Base):
Index("ix_perm_log_endpoint_created", "endpoint_key", "created_at"),
Index("ix_perm_log_created_at", "created_at"),
Index("ix_perm_log_allowed", "allowed", "created_at"),
Index("ix_perm_log_ip_created", "ip_address", "created_at"),
Index("ix_perm_log_client_source_created", "client_type", "created_at"),
Index("ix_perm_log_request_id", "request_id"),
)
id: Mapped[uuid.UUID] = mapped_column(
@@ -38,6 +43,15 @@ class PermissionAccessLog(Base):
allowed: Mapped[bool] = mapped_column(Boolean, nullable=False)
elapsed_ms: Mapped[float] = mapped_column(Float, nullable=False)
ip_address: Mapped[Optional[str]] = mapped_column(String(45), nullable=True)
user_agent: Mapped[Optional[str]] = mapped_column(String(500), nullable=True)
client_type: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
client_version: Mapped[Optional[str]] = mapped_column(String(32), nullable=True)
client_platform: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
build_channel: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
build_commit: Mapped[Optional[str]] = mapped_column(String(64), nullable=True)
request_headers: Mapped[Optional[dict]] = mapped_column(JSONB_TYPE, nullable=True)
request_snapshot: Mapped[Optional[dict]] = mapped_column(JSONB_TYPE, nullable=True)
request_id: Mapped[Optional[str]] = mapped_column(String(36), nullable=True)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now()
)
+12 -2
View File
@@ -6,13 +6,15 @@ import uuid
from datetime import datetime
from typing import Optional
from sqlalchemy import DateTime, Float, Index, Integer, String
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy import JSON, DateTime, Float, Index, Integer, String
from sqlalchemy.dialects.postgresql import JSONB, UUID
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.sql import func
from app.db.base_class import Base
JSONB_TYPE = JSON().with_variant(JSONB, "postgresql")
class SecurityAccessLog(Base):
__tablename__ = "security_access_logs"
@@ -21,6 +23,9 @@ class SecurityAccessLog(Base):
Index("ix_security_log_ip_created", "client_ip", "created_at"),
Index("ix_security_log_status_created", "status_code", "created_at"),
Index("ix_security_log_auth_created", "auth_status", "created_at"),
Index("ix_security_log_request_id", "request_id"),
Index("ix_security_log_category_created", "category", "created_at"),
Index("ix_security_log_severity_created", "severity", "created_at"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
@@ -35,8 +40,13 @@ class SecurityAccessLog(Base):
client_platform: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
build_channel: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
build_commit: Mapped[Optional[str]] = mapped_column(String(64), nullable=True)
request_headers: Mapped[Optional[dict]] = mapped_column(JSONB_TYPE, nullable=True)
request_snapshot: Mapped[Optional[dict]] = mapped_column(JSONB_TYPE, nullable=True)
auth_status: Mapped[str] = mapped_column(String(30), nullable=False)
user_identifier: Mapped[Optional[str]] = mapped_column(String(80), nullable=True)
request_id: Mapped[Optional[str]] = mapped_column(String(36), nullable=True)
category: Mapped[Optional[str]] = mapped_column(String(30), nullable=True)
severity: Mapped[Optional[str]] = mapped_column(String(16), nullable=True)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now()
)
@@ -0,0 +1,48 @@
"""Hourly source-location rollup for monitoring analytics."""
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import DateTime, Float, Index, Integer, String, UniqueConstraint
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.sql import func
from app.db.base_class import Base
class SourceLocationSnapshot(Base):
__tablename__ = "source_location_snapshots"
__table_args__ = (
UniqueConstraint("bucket_time", "ip_hash", "user_hash", name="uq_source_location_bucket_identity"),
Index("ix_source_location_bucket", "bucket_time"),
Index("ix_source_location_country_bucket", "country_code", "bucket_time"),
Index("ix_source_location_risk_bucket", "high_risk_count", "bucket_time"),
)
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True, default=uuid.uuid4)
bucket_time: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
ip_hash: Mapped[str] = mapped_column(String(64), nullable=False)
user_hash: Mapped[str] = mapped_column(String(64), nullable=False, default="")
country: Mapped[str] = mapped_column(String(100), nullable=False, default="")
country_code: Mapped[str] = mapped_column(String(16), nullable=False, default="")
province: Mapped[str] = mapped_column(String(100), nullable=False, default="")
region_code: Mapped[str] = mapped_column(String(24), nullable=False, default="")
city: Mapped[str] = mapped_column(String(100), nullable=False, default="")
isp: Mapped[str] = mapped_column(String(160), nullable=False, default="")
location: Mapped[str] = mapped_column(String(320), nullable=False, default="")
longitude: Mapped[float | None] = mapped_column(Float, nullable=True)
latitude: Mapped[float | None] = mapped_column(Float, nullable=True)
accuracy_level: Mapped[str] = mapped_column(String(16), nullable=False, default="unknown")
allowed_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
denied_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
security_event_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
high_risk_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
auth_failure_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
first_seen_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
last_seen_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, server_default=func.now()
)
+33
View File
@@ -0,0 +1,33 @@
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import DateTime, ForeignKey, String, func
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import Mapped, mapped_column
from app.db.base_class import Base
class UserLoginSession(Base):
"""Server-side login activity record without storing credentials."""
__tablename__ = "user_login_sessions"
id: Mapped[uuid.UUID] = mapped_column(UUID(as_uuid=True), primary_key=True)
user_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
index=True,
)
client_type: Mapped[str] = mapped_column(String(16), nullable=False, default="web")
client_platform: Mapped[str | None] = mapped_column(String(32), nullable=True)
client_version: Mapped[str | None] = mapped_column(String(64), nullable=True)
client_source: Mapped[str | None] = mapped_column(String(32), nullable=True)
login_ip: Mapped[str | None] = mapped_column(String(45), nullable=True)
login_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
last_seen_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False, server_default=func.now())
ended_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
end_reason: Mapped[str | None] = mapped_column(String(32), nullable=True)
+17
View File
@@ -7,10 +7,27 @@ from pydantic import BaseModel, ConfigDict
class AuditLogRead(BaseModel):
id: uuid.UUID
study_id: Optional[uuid.UUID] = None
entity_type: str
entity_id: Optional[uuid.UUID] = None
action: str
detail: Optional[str]
operator_id: uuid.UUID
operator_name: Optional[str] = None
operator_email: Optional[str] = None
operator_role: str
client_ip: Optional[str] = None
ip_location: str = ""
ip_country: str = ""
ip_province: str = ""
ip_city: str = ""
ip_isp: str = ""
user_agent: Optional[str] = None
client_type: Optional[str] = None
client_version: Optional[str] = None
client_platform: Optional[str] = None
build_channel: Optional[str] = None
build_commit: Optional[str] = None
created_at: datetime
model_config = ConfigDict(from_attributes=True)
+285
View File
@@ -0,0 +1,285 @@
import uuid
from datetime import datetime
from typing import Literal, Optional
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
CollaborationFileType = Literal["word", "cell", "slide"]
CollaborationMemberRole = Literal["EDITOR", "MANAGER"]
CollaborationFileStatus = Literal["ACTIVE", "ARCHIVED", "DELETED"]
CollaborationShareAccessMode = Literal["VIEW", "EDIT"]
CollaborationShareExpiryPolicy = Literal["ONE_DAY", "SEVEN_DAYS", "THIRTY_DAYS", "PERMANENT"]
CollaborationEditRequestStatus = Literal["PENDING", "APPROVED", "REJECTED"]
class CollaborationFolderCreate(BaseModel):
name: str = Field(min_length=1, max_length=120)
parent_id: Optional[uuid.UUID] = None
sort_order: int = 0
@field_validator("name")
@classmethod
def normalize_name(cls, value: str) -> str:
value = value.strip()
if not value:
raise ValueError("文件夹名称不能为空")
return value
class CollaborationFolderUpdate(BaseModel):
name: Optional[str] = Field(default=None, min_length=1, max_length=120)
parent_id: Optional[uuid.UUID] = None
sort_order: Optional[int] = None
class CollaborationFolderRead(BaseModel):
id: uuid.UUID
study_id: uuid.UUID
parent_id: Optional[uuid.UUID]
name: str
sort_order: int
created_by: uuid.UUID
deleted_at: Optional[datetime]
created_at: datetime
updated_at: datetime
model_config = ConfigDict(from_attributes=True)
class CollaborationFileCreate(BaseModel):
title: str = Field(min_length=1, max_length=240)
file_type: CollaborationFileType
folder_id: Optional[uuid.UUID] = None
class CollaborationFileUpdate(BaseModel):
title: Optional[str] = Field(default=None, min_length=1, max_length=240)
folder_id: Optional[uuid.UUID] = None
status: Optional[Literal["ACTIVE", "ARCHIVED"]] = None
allow_export: Optional[bool] = None
allow_edit_request: Optional[bool] = None
allow_sheet_structure_edit: Optional[bool] = None
class CollaborationRevisionRead(BaseModel):
id: uuid.UUID
file_id: uuid.UUID
revision_no: int
parent_revision_id: Optional[uuid.UUID]
original_filename: str
file_hash: str
file_size: int
mime_type: str
source: str
change_summary: Optional[str]
created_by: Optional[uuid.UUID]
created_by_name: Optional[str] = None
created_by_avatar_url: Optional[str] = None
created_at: datetime
model_config = ConfigDict(from_attributes=True)
class CollaborationRevisionUpdate(BaseModel):
change_summary: str = Field(min_length=1, max_length=240)
@field_validator("change_summary")
@classmethod
def normalize_change_summary(cls, value: str) -> str:
value = value.strip()
if not value:
raise ValueError("版本名称不能为空")
return value
class CollaborationRevisionCopyRequest(BaseModel):
title: str = Field(min_length=1, max_length=240)
folder_id: Optional[uuid.UUID] = None
@field_validator("title")
@classmethod
def normalize_title(cls, value: str) -> str:
value = value.strip()
if not value:
raise ValueError("文件名不能为空")
return value
class CollaborationFileCollaboratorRead(BaseModel):
user_id: uuid.UUID
full_name: str
role: CollaborationMemberRole
avatar_url: Optional[str] = None
class CollaborationFileRead(BaseModel):
id: uuid.UUID
study_id: uuid.UUID
folder_id: Optional[uuid.UUID]
title: str
file_type: CollaborationFileType
extension: str
status: CollaborationFileStatus
owner_id: uuid.UUID
current_revision_id: Optional[uuid.UUID]
generation: int
allow_export: bool = False
allow_edit_request: bool = False
allow_sheet_structure_edit: bool = True
deleted_at: Optional[datetime]
created_at: datetime
updated_at: datetime
owner_name: Optional[str] = None
folder_name: Optional[str] = None
current_revision_no: Optional[int] = None
current_revision_file_size: Optional[int] = None
current_revision_mime_type: Optional[str] = None
current_revision_created_at: Optional[datetime] = None
collaboration_role: Optional[CollaborationMemberRole] = None
collaborators: list[CollaborationFileCollaboratorRead] = Field(default_factory=list)
can_edit: bool = False
can_manage: bool = False
can_export: bool = False
can_request_edit: bool = False
edit_request_status: Optional[CollaborationEditRequestStatus] = None
can_transfer_ownership: bool = False
model_config = ConfigDict(from_attributes=True)
class CollaborationMemberUpsert(BaseModel):
user_id: uuid.UUID
role: CollaborationMemberRole
class CollaborationMemberRead(BaseModel):
id: uuid.UUID
file_id: uuid.UUID
user_id: uuid.UUID
role: CollaborationMemberRole
invited_by: uuid.UUID
full_name: str
email: str
created_at: datetime
class CollaborationEditRequestRead(BaseModel):
id: uuid.UUID
file_id: uuid.UUID
requester_id: uuid.UUID
requester_name: str
requester_email: str
status: CollaborationEditRequestStatus
resolved_by: Optional[uuid.UUID]
resolved_at: Optional[datetime]
created_at: datetime
class CollaborationEditRequestResolve(BaseModel):
status: Literal["APPROVED", "REJECTED"]
class CollaborationOwnershipTransferRequest(BaseModel):
new_owner_id: uuid.UUID
class CollaborationShareLinkUpdate(BaseModel):
enabled: bool
access_mode: CollaborationShareAccessMode = "VIEW"
expiry_policy: CollaborationShareExpiryPolicy = "SEVEN_DAYS"
password_mode: Literal["KEEP", "SET", "CLEAR"] = "KEEP"
password: Optional[str] = Field(default=None, min_length=4, max_length=64)
@model_validator(mode="after")
def validate_password_change(self):
if self.password_mode == "SET" and not self.password:
raise ValueError("设置链接密码时必须提供密码")
if self.password_mode != "SET" and self.password is not None:
raise ValueError("仅在设置链接密码时允许提交密码")
return self
class CollaborationShareLinkRead(BaseModel):
id: uuid.UUID
file_id: uuid.UUID
enabled: bool
access_mode: CollaborationShareAccessMode
expiry_policy: CollaborationShareExpiryPolicy
expires_at: Optional[datetime]
has_password: bool
share_path: str = "/collaboration/share"
share_token: Optional[str] = None
created_at: datetime
updated_at: datetime
class CollaborationPublicShareMetadata(BaseModel):
file_name: str
file_type: CollaborationFileType
access_mode: Literal["view", "edit"]
allow_export: bool
requires_password: bool
expires_at: Optional[datetime]
class CollaborationSharePasswordRequest(BaseModel):
password: str = Field(min_length=1, max_length=64)
class CollaborationShareAccessGrant(BaseModel):
access_token: str
expires_at: datetime
class CollaborationPublicEditorConfigRequest(BaseModel):
access_token: Optional[str] = Field(default=None, max_length=2048)
client_id: str = Field(min_length=8, max_length=64, pattern=r"^[A-Za-z0-9_-]+$")
display_name: str = Field(default="链接访客", min_length=1, max_length=40)
@field_validator("display_name")
@classmethod
def normalize_display_name(cls, value: str) -> str:
return value.strip() or "链接访客"
class CollaborationCandidateRead(BaseModel):
user_id: uuid.UUID
full_name: str
email: str
role_in_study: str
can_be_editor: bool
can_be_manager: bool
class CollaborationEditorConfigRead(BaseModel):
file_id: uuid.UUID
file_name: str
access_mode: Literal["view", "edit"]
can_save_as: bool = False
can_download: bool = False
can_request_edit: bool = False
host_path: str = "/onlyoffice-host.html"
expires_at: datetime
config: dict
class CollaborationExportRecord(BaseModel):
file_type: str = Field(min_length=1, max_length=16, pattern=r"^[a-z0-9]+$")
class CollaborationCallbackPayload(BaseModel):
key: str
status: int
url: Optional[str] = None
changesurl: Optional[str] = None
filetype: Optional[str] = None
forcesavetype: Optional[int] = None
userdata: Optional[str] = None
users: list[str] = Field(default_factory=list)
actions: list[dict] = Field(default_factory=list)
history: Optional[dict] = None
class CollaborationRestoreRequest(BaseModel):
change_summary: Optional[str] = Field(default=None, max_length=500)
+1
View File
@@ -36,6 +36,7 @@ class DocumentVersionRead(BaseModel):
effective_at: Optional[datetime] = None
superseded_at: Optional[datetime] = None
file_uri: str
original_filename: Optional[str] = None
file_hash: str
file_size: int
mime_type: Optional[str] = None
+27 -1
View File
@@ -37,12 +37,38 @@ class DesktopNotificationClaimRequest(BaseModel):
limit: int = 20
class GeneralNotificationRead(BaseModel):
id: uuid.UUID
study_id: uuid.UUID
recipient_id: uuid.UUID
category: str
priority: str
title: str
message: str
action_path: str | None = None
source_type: str
source_id: str
requires_action: bool = True
due_at: datetime | None = None
read_at: datetime | None = None
resolved_at: datetime | None = None
created_at: datetime
model_config = ConfigDict(from_attributes=True)
class DesktopNotificationClaimResponse(BaseModel):
claim_token: uuid.UUID | None = None
lease_expires_at: datetime | None = None
items: list[NotificationItem]
items: list[GeneralNotificationRead]
class DesktopNotificationAckRequest(BaseModel):
claim_token: uuid.UUID
delivered_ids: list[uuid.UUID]
class GeneralNotificationFeed(BaseModel):
unread_count: int
total_count: int
items: list[GeneralNotificationRead]
+14
View File
@@ -0,0 +1,14 @@
from datetime import datetime
from typing import Any, Literal
import uuid
from pydantic import BaseModel
class OnlyOfficePreviewConfigRead(BaseModel):
resource_type: Literal["attachment", "version", "collaboration_revision"]
resource_id: uuid.UUID
file_name: str
host_path: str = "/onlyoffice-host.html"
expires_at: datetime
config: dict[str, Any]
+23
View File
@@ -6,6 +6,7 @@ from typing import Literal, Optional
from pydantic import BaseModel, ConfigDict, EmailStr, Field, field_validator
UserStatus = Literal["PENDING", "ACTIVE", "REJECTED", "DISABLED"]
LoginStatus = Literal["ONLINE", "OFFLINE"]
PASSWORD_REGEX = re.compile(r"^(?=.*[A-Za-z])(?=.*\d).{8,}$")
@@ -60,6 +61,11 @@ class UserRead(BaseModel):
approved_at: Optional[datetime] = None
approved_by: Optional[uuid.UUID] = None
avatar_url: Optional[str] = None
login_status: LoginStatus = "OFFLINE"
last_login_at: Optional[datetime] = None
last_seen_at: Optional[datetime] = None
last_client_type: Optional[Literal["web", "desktop"]] = None
active_session_count: int = 0
model_config = ConfigDict(from_attributes=True)
@@ -78,6 +84,23 @@ class UserResponse(UserRead):
pass
class UserLoginActivityRead(BaseModel):
id: uuid.UUID
client_type: Literal["web", "desktop"]
client_platform: Optional[str] = None
client_version: Optional[str] = None
client_source: Optional[str] = None
login_ip: Optional[str] = None
ip_location: Optional[str] = None
login_at: datetime
last_seen_at: datetime
ended_at: Optional[datetime] = None
end_reason: Optional[str] = None
activity_status: Literal["ONLINE", "OFFLINE", "ENDED"]
model_config = ConfigDict(from_attributes=True)
class UserSelfUpdate(_PasswordValidator):
full_name: Optional[str] = None
clinical_department: Optional[str] = None
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,289 @@
from __future__ import annotations
import base64
import binascii
import hashlib
import hmac
import uuid
from datetime import datetime, timedelta, timezone
from anyio import to_thread
from fastapi import HTTPException, status
from jose import JWTError, jwt
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.core.security import hash_password, verify_password
from app.models.collaboration import CollaborationFile, CollaborationShareLink
from app.schemas.collaboration import (
CollaborationPublicShareMetadata,
CollaborationShareAccessGrant,
CollaborationShareLinkRead,
CollaborationShareLinkUpdate,
)
from app.services import collaboration_service
SHARE_PATH = "/collaboration/share"
SHARE_ACCESS_TTL_SECONDS = 30 * 60
PASSWORD_FAILURE_WINDOW = timedelta(minutes=15)
PASSWORD_LOCK_DURATION = timedelta(minutes=15)
PASSWORD_FAILURE_LIMIT = 5
_ACCESS_PURPOSE = "ctms-collaboration-share-access"
_EXPIRY_DURATIONS = {
"ONE_DAY": timedelta(days=1),
"SEVEN_DAYS": timedelta(days=7),
"THIRTY_DAYS": timedelta(days=30),
"PERMANENT": None,
}
def _now() -> datetime:
return datetime.now(timezone.utc)
def _signing_key() -> bytes:
return hmac.new(
settings.JWT_SECRET_KEY.encode("utf-8"),
b"ctms-collaboration-share-v1",
hashlib.sha256,
).digest()
def _b64encode(value: bytes) -> str:
return base64.urlsafe_b64encode(value).rstrip(b"=").decode("ascii")
def _b64decode(value: str) -> bytes:
padding = "=" * (-len(value) % 4)
return base64.urlsafe_b64decode(f"{value}{padding}".encode("ascii"))
def share_token(link: CollaborationShareLink) -> str:
payload = f"{link.id}.{link.token_version}".encode("ascii")
signature = hmac.new(_signing_key(), payload, hashlib.sha256).digest()
return f"{_b64encode(payload)}.{_b64encode(signature)}"
def _decode_share_token(value: str | None) -> tuple[uuid.UUID, int]:
token = (value or "").strip()
if not token or len(token) > 256 or token.count(".") != 1:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="共享链接不存在或已失效")
encoded_payload, encoded_signature = token.split(".", 1)
try:
payload = _b64decode(encoded_payload)
actual_signature = _b64decode(encoded_signature)
if (
not hmac.compare_digest(_b64encode(payload), encoded_payload)
or not hmac.compare_digest(_b64encode(actual_signature), encoded_signature)
):
raise ValueError("non-canonical token encoding")
expected_signature = hmac.new(_signing_key(), payload, hashlib.sha256).digest()
if not hmac.compare_digest(actual_signature, expected_signature):
raise ValueError("signature mismatch")
raw_id, raw_version = payload.decode("ascii").split(".", 1)
return uuid.UUID(raw_id), int(raw_version)
except (ValueError, UnicodeError, TypeError, binascii.Error) as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="共享链接不存在或已失效") from exc
def _expiry_for_policy(policy: str, now: datetime) -> datetime | None:
duration = _EXPIRY_DURATIONS[policy]
return now + duration if duration else None
def _is_expired(link: CollaborationShareLink, now: datetime | None = None) -> bool:
return bool(link.expires_at and link.expires_at <= (now or _now()))
async def _link_for_file(
db: AsyncSession,
item: CollaborationFile,
user,
*,
create: bool,
lock: bool = False,
) -> CollaborationShareLink | None:
await collaboration_service.require_file_manager(db, item, user)
statement = select(CollaborationShareLink).where(CollaborationShareLink.file_id == item.id)
if lock:
statement = statement.with_for_update()
link = await db.scalar(statement)
if link or not create:
return link
now = _now()
link = CollaborationShareLink(
file_id=item.id,
enabled=False,
access_mode="VIEW",
expiry_policy="SEVEN_DAYS",
expires_at=now + timedelta(days=7),
created_by=user.id,
updated_by=user.id,
)
db.add(link)
await db.flush()
return link
def share_link_read(link: CollaborationShareLink) -> CollaborationShareLinkRead:
return CollaborationShareLinkRead(
id=link.id,
file_id=link.file_id,
enabled=link.enabled,
access_mode=link.access_mode,
expiry_policy=link.expiry_policy,
expires_at=link.expires_at,
has_password=bool(link.password_hash),
share_path=SHARE_PATH,
share_token=share_token(link) if link.enabled else None,
created_at=link.created_at,
updated_at=link.updated_at,
)
async def get_share_link(
db: AsyncSession, item: CollaborationFile, user
) -> CollaborationShareLinkRead:
link = await _link_for_file(db, item, user, create=True)
assert link is not None
await db.commit()
await db.refresh(link)
return share_link_read(link)
async def update_share_link(
db: AsyncSession,
item: CollaborationFile,
payload: CollaborationShareLinkUpdate,
user,
) -> CollaborationShareLinkRead:
link = await _link_for_file(db, item, user, create=True, lock=True)
assert link is not None
now = _now()
link.enabled = payload.enabled
link.access_mode = payload.access_mode
link.expiry_policy = payload.expiry_policy
link.expires_at = _expiry_for_policy(payload.expiry_policy, now)
link.updated_by = user.id
if payload.password_mode == "SET":
link.password_hash = await to_thread.run_sync(hash_password, payload.password or "")
link.failed_attempts = 0
link.last_failed_at = None
link.locked_until = None
elif payload.password_mode == "CLEAR":
link.password_hash = None
link.failed_attempts = 0
link.last_failed_at = None
link.locked_until = None
await db.commit()
await db.refresh(link)
return share_link_read(link)
async def resolve_active_share(
db: AsyncSession,
token: str | None,
*,
lock: bool = False,
) -> tuple[CollaborationShareLink, CollaborationFile]:
link_id, version = _decode_share_token(token)
statement = select(CollaborationShareLink).where(CollaborationShareLink.id == link_id)
if lock:
statement = statement.with_for_update()
link = await db.scalar(statement)
if not link or link.token_version != version or not link.enabled:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="共享链接不存在或已失效")
if _is_expired(link):
raise HTTPException(status_code=status.HTTP_410_GONE, detail="共享链接已过期")
item = await db.get(CollaborationFile, link.file_id)
if not item or item.deleted_at or item.status != "ACTIVE":
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="共享文件不存在或已停止共享")
return link, item
async def public_metadata(
db: AsyncSession, token: str | None
) -> CollaborationPublicShareMetadata:
link, item = await resolve_active_share(db, token)
return CollaborationPublicShareMetadata(
file_name=item.title,
file_type=item.file_type,
access_mode="edit" if link.access_mode == "EDIT" else "view",
allow_export=item.allow_export,
requires_password=bool(link.password_hash),
expires_at=link.expires_at,
)
def _grant_token(link: CollaborationShareLink) -> CollaborationShareAccessGrant:
now = _now()
expires_at = now + timedelta(seconds=SHARE_ACCESS_TTL_SECONDS)
if link.expires_at and link.expires_at < expires_at:
expires_at = link.expires_at
value = jwt.encode(
{
"purpose": _ACCESS_PURPOSE,
"sub": str(link.id),
"ver": link.token_version,
"iat": int(now.timestamp()),
"exp": int(expires_at.timestamp()),
},
_signing_key().hex(),
algorithm="HS256",
)
return CollaborationShareAccessGrant(access_token=value, expires_at=expires_at)
async def verify_share_password(
db: AsyncSession,
token: str | None,
password: str,
) -> CollaborationShareAccessGrant:
link, _ = await resolve_active_share(db, token, lock=True)
if not link.password_hash:
return _grant_token(link)
now = _now()
if link.locked_until and link.locked_until > now:
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="密码尝试次数过多,请稍后再试",
)
if link.last_failed_at and now - link.last_failed_at > PASSWORD_FAILURE_WINDOW:
link.failed_attempts = 0
valid = await to_thread.run_sync(verify_password, password, link.password_hash)
if not valid:
link.failed_attempts += 1
link.last_failed_at = now
if link.failed_attempts >= PASSWORD_FAILURE_LIMIT:
link.locked_until = now + PASSWORD_LOCK_DURATION
await db.commit()
if link.locked_until:
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="密码尝试次数过多,请稍后再试",
)
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="链接密码不正确")
link.failed_attempts = 0
link.last_failed_at = None
link.locked_until = None
await db.commit()
return _grant_token(link)
def validate_access_grant(link: CollaborationShareLink, value: str | None) -> None:
if not link.password_hash:
return
if not value:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="请输入链接密码")
try:
payload = jwt.decode(value, _signing_key().hex(), algorithms=["HS256"])
if (
payload.get("purpose") != _ACCESS_PURPOSE
or not hmac.compare_digest(str(payload.get("sub") or ""), str(link.id))
or payload.get("ver") != link.token_version
):
raise JWTError("share grant mismatch")
except JWTError as exc:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="链接访问凭证已失效") from exc
@@ -10,12 +10,9 @@ from app.models.desktop_notification import (
DesktopNotificationDelivery,
DesktopNotificationSubscription,
)
from app.models.distribution import Distribution, DistributionStatus, DistributionTargetType
from app.models.document import Document
from app.models.document_version import DocumentVersion
from app.models.study import Study
from app.models.notification import Notification
from app.models.study_member import StudyMember
from app.schemas.notification import NotificationItem
from app.models.user import User
CLAIM_LEASE = timedelta(minutes=5)
@@ -47,46 +44,28 @@ async def set_subscription(
def _eligible_query(user_id: uuid.UUID, enabled_at: datetime, lease_cutoff: datetime):
role_target_exists = exists(
active_membership = exists(
select(StudyMember.id).where(
StudyMember.study_id == Document.trial_id,
StudyMember.study_id == Notification.study_id,
StudyMember.user_id == user_id,
StudyMember.is_active.is_(True),
StudyMember.role_in_study == Distribution.target_id,
)
)
target_matches = or_(
and_(
Distribution.target_type == DistributionTargetType.USER,
Distribution.target_id == str(user_id),
),
and_(
Distribution.target_type == DistributionTargetType.ROLE,
role_target_exists,
),
)
return (
select(
Distribution,
Document,
DocumentVersion,
Study,
DesktopNotificationDelivery,
)
.join(Document, Distribution.document_id == Document.id)
.join(DocumentVersion, Distribution.version_id == DocumentVersion.id)
.join(Study, Document.trial_id == Study.id)
select(Notification, DesktopNotificationDelivery)
.outerjoin(
DesktopNotificationDelivery,
and_(
DesktopNotificationDelivery.distribution_id == Distribution.id,
DesktopNotificationDelivery.notification_id == Notification.id,
DesktopNotificationDelivery.user_id == user_id,
),
)
.where(
Distribution.status == DistributionStatus.ACTIVE,
Distribution.created_at >= enabled_at,
target_matches,
Notification.recipient_id == user_id,
Notification.created_at >= enabled_at,
Notification.resolved_at.is_(None),
Notification.read_at.is_(None),
active_membership,
or_(
DesktopNotificationDelivery.id.is_(None),
and_(
@@ -98,8 +77,8 @@ def _eligible_query(user_id: uuid.UUID, enabled_at: datetime, lease_cutoff: date
),
),
)
.order_by(Distribution.created_at.asc())
.with_for_update(of=Distribution, skip_locked=True)
.order_by(Notification.created_at.asc())
.with_for_update(of=Notification, skip_locked=True)
)
@@ -107,45 +86,46 @@ async def claim_notifications(
db: AsyncSession,
user_id: uuid.UUID,
limit: int,
) -> tuple[uuid.UUID | None, datetime | None, list[NotificationItem]]:
) -> tuple[uuid.UUID | None, datetime | None, list[Notification]]:
subscription = await get_subscription(db, user_id)
if not subscription or not subscription.enabled or not subscription.enabled_at:
return None, None, []
# Reconcile every active project before selecting desktop deliveries so permission
# changes fail closed and time-based reminders do not depend on opening the web Feed.
user = await db.get(User, user_id)
if user is None or not user.is_active:
return None, None, []
study_ids = (await db.scalars(
select(StudyMember.study_id).where(
StudyMember.user_id == user_id,
StudyMember.is_active.is_(True),
)
)).all()
from app.services.project_reminder_service import sync_project_reminders
for study_id in set(study_ids):
await sync_project_reminders(db, study_id, user)
now = datetime.now(timezone.utc)
token = uuid.uuid4()
rows = (
await db.execute(
_eligible_query(user_id, subscription.enabled_at, now - CLAIM_LEASE).limit(max(1, min(limit, 50)))
_eligible_query(user_id, subscription.enabled_at, now - CLAIM_LEASE)
.limit(max(1, min(limit, 50)))
)
).all()
items: list[NotificationItem] = []
for distribution, document, version, study, delivery in rows:
items: list[Notification] = []
for notification, delivery in rows:
if delivery is None:
delivery = DesktopNotificationDelivery(
user_id=user_id,
distribution_id=distribution.id,
notification_id=notification.id,
)
db.add(delivery)
delivery.claim_token = token
delivery.claimed_at = now
items.append(
NotificationItem(
id=distribution.id,
document_id=document.id,
version_id=version.id,
document_title=document.title,
document_no=document.doc_no,
version_no=version.version_no,
change_summary=version.change_summary,
effective_at=version.effective_at,
created_at=distribution.created_at,
study_id=study.id,
study_name=study.name,
delivered_at=delivery.delivered_at,
read_at=delivery.read_at,
)
)
items.append(notification)
await db.commit()
if not items:
return None, None, []
@@ -164,7 +144,7 @@ async def acknowledge_notifications(
.where(
DesktopNotificationDelivery.user_id == user_id,
DesktopNotificationDelivery.claim_token == claim_token,
DesktopNotificationDelivery.distribution_id.in_(delivered_ids),
DesktopNotificationDelivery.notification_id.in_(delivered_ids),
)
.values(delivered_at=datetime.now(timezone.utc))
)
@@ -174,21 +154,17 @@ async def acknowledge_notifications(
async def mark_notification_read(
db: AsyncSession,
user_id: uuid.UUID,
distribution_id: uuid.UUID,
notification_id: uuid.UUID,
) -> None:
delivery = (
await db.execute(
select(DesktopNotificationDelivery).where(
DesktopNotificationDelivery.user_id == user_id,
DesktopNotificationDelivery.distribution_id == distribution_id,
)
)
).scalar_one_or_none()
if delivery is None:
delivery = DesktopNotificationDelivery(
user_id=user_id,
distribution_id=distribution_id,
)
db.add(delivery)
delivery.read_at = datetime.now(timezone.utc)
item = await db.scalar(select(Notification).where(
Notification.id == notification_id,
Notification.recipient_id == user_id,
Notification.resolved_at.is_(None),
))
if item is None:
return
now = datetime.now(timezone.utc)
item.read_at = item.read_at or now
if not item.requires_action:
item.resolved_at = item.resolved_at or now
await db.commit()
+51 -6
View File
@@ -6,11 +6,12 @@ import uuid
from datetime import datetime, timezone
from pathlib import Path
from typing import Iterable
from urllib.parse import quote
import aiofiles
from fastapi import HTTPException, UploadFile, status
from fastapi.responses import FileResponse
from sqlalchemy import and_, delete as sa_delete, or_, select, update as sa_update
from sqlalchemy import String, and_, cast, delete as sa_delete, or_, select, update as sa_update
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_cra_site_scope
@@ -31,6 +32,7 @@ from app.models.distribution import Distribution, DistributionStatus, Distributi
from app.models.document import Document, DocumentScopeType, DocumentStatus
from app.models.document_version import DocumentVersion, DocumentVersionStatus
from app.models.desktop_notification import DesktopNotificationDelivery
from app.models.notification import Notification
from app.models.study import Study
from app.schemas.acknowledgement import AcknowledgementCreate
from app.schemas.distribution import DistributionCreate, DistributionRead, DistributionStats
@@ -38,6 +40,7 @@ from app.schemas.document import DocumentCreate, DocumentDetail, DocumentSummary
from app.schemas.document_version import DocumentVersionRead, DocumentVersionSummary
from app.schemas.notification import NotificationItem
from app.schemas.user import UserDisplay
from app.services import notification_service
UPLOAD_ROOT = Path(__file__).resolve().parent.parent / "uploads" / "documents"
@@ -52,6 +55,29 @@ DOCUMENT_ACTION_PERMISSIONS = {
}
def _safe_original_filename(value: str | None) -> str:
filename = (value or "").replace("\\", "/").rsplit("/", 1)[-1].strip()
filename = "".join(character for character in filename if ord(character) >= 32 and ord(character) != 127)
return filename[:255] or "document"
def _content_disposition(filename: str, disposition: str = "attachment") -> str:
fallback = "".join(character if 32 <= ord(character) < 127 and character not in {'"', "\\"} else "_" for character in filename)
fallback = fallback or "download"
encoded = quote(filename, safe="")
return f'{disposition}; filename="{fallback}"; filename*=UTF-8\'\'{encoded}'
def _legacy_download_filename(version: DocumentVersion, document: Document) -> str:
"""Return a readable fallback for rows created before original_filename existed."""
stored_name = Path(version.file_uri).name
suffix = Path(stored_name).suffix
title = _safe_original_filename(document.title)
if suffix and title.lower().endswith(suffix.lower()):
return title
return f"{title}{suffix}"
def _audit_detail(before: dict | None, after: dict | None) -> str:
payload = {"before": before, "after": after}
return json.dumps(payload, ensure_ascii=True)
@@ -356,7 +382,8 @@ async def create_version(
file_hash = hashlib.sha256(content).hexdigest()
dest_dir = UPLOAD_ROOT / str(document_id)
dest_dir.mkdir(parents=True, exist_ok=True)
unique_name = f"{uuid.uuid4()}{Path(file.filename).suffix}"
original_filename = _safe_original_filename(file.filename)
unique_name = f"{uuid.uuid4()}{Path(original_filename).suffix}"
dest_path = dest_dir / unique_name
async with aiofiles.open(dest_path, "wb") as out_file:
await out_file.write(content)
@@ -377,6 +404,7 @@ async def create_version(
status=DocumentVersionStatus.EFFECTIVE,
effective_at=effective_at,
file_uri=str(dest_path),
original_filename=original_filename,
file_hash=file_hash,
file_size=len(content),
mime_type=file.content_type,
@@ -552,12 +580,11 @@ async def get_version_download_response(
file_path = Path(version.file_uri)
if not file_path.exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="文件不存在")
filename = file_path.name
filename = version.original_filename or _legacy_download_filename(version, doc)
return FileResponse(
path=str(file_path),
filename=filename,
media_type=version.mime_type or "application/octet-stream",
headers={"Content-Disposition": f'inline; filename="{filename}"'},
headers={"Content-Disposition": _content_disposition(filename)},
)
@@ -663,6 +690,10 @@ async def create_acknowledgement(
if distribution.target_type == DistributionTargetType.ROLE:
if distribution.target_id not in ("ADMIN" if is_system_admin(current_user) else "", getattr(membership, "role_in_study", "")):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="不在分发范围内")
if distribution.target_type == DistributionTargetType.SITE and not is_system_admin(current_user):
site_ids = await site_crud.list_ids_by_contact_user(db, doc.trial_id, current_user.id)
if distribution.target_id not in {str(site_id) for site_id in site_ids}:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="不在分发中心范围内")
if payload.ack_type != AcknowledgementType.RECEIVED:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail="仅支持已接收回执")
@@ -697,6 +728,12 @@ async def create_acknowledgement(
operator_role=await get_operator_role_label(db, doc.trial_id, current_user),
)
)
await notification_service.resolve_source_notifications(
db,
source_type="DOCUMENT_DISTRIBUTION",
source_id=str(distribution.id),
recipient_id=current_user.id,
)
await db.commit()
await db.refresh(ack)
return ack
@@ -787,10 +824,18 @@ async def list_distribution_notifications(
.join(DocumentVersion, Distribution.version_id == DocumentVersion.id)
.join(Document, Distribution.document_id == Document.id)
.join(Study, Document.trial_id == Study.id)
.outerjoin(
Notification,
and_(
Notification.source_type == "DOCUMENT_DISTRIBUTION",
Notification.source_id == cast(Distribution.id, String),
Notification.recipient_id == current_user.id,
),
)
.outerjoin(
DesktopNotificationDelivery,
and_(
DesktopNotificationDelivery.distribution_id == Distribution.id,
DesktopNotificationDelivery.notification_id == Notification.id,
DesktopNotificationDelivery.user_id == current_user.id,
),
)
@@ -0,0 +1,175 @@
"""Canonical map metadata for monitoring source locations.
ip2region provides administrative names but no coordinates. This module keeps
the normalization and centroid contract on the server so web and desktop
clients render the same location semantics.
"""
from __future__ import annotations
from dataclasses import dataclass
from app.services.ip_location import IpLocation
@dataclass(frozen=True)
class GeoLocationMetadata:
country: str
country_code: str
region_code: str
longitude: float | None
latitude: float | None
accuracy_level: str
COUNTRY_ALIASES = {
"中国": "China",
"China": "China",
"Mainland China": "China",
"中国香港": "China",
"中国澳门": "China",
"中国台湾": "China",
"美国": "United States",
"United States": "United States",
"United States of America": "United States",
"USA": "United States",
"土耳其": "Türkiye",
"Turkey": "Türkiye",
"Türkiye": "Türkiye",
}
COUNTRY_CODES = {
"China": "CN",
"United States": "US",
"Netherlands": "NL",
"Türkiye": "TR",
"Australia": "AU",
"Japan": "JP",
"Singapore": "SG",
"Germany": "DE",
"France": "FR",
"United Kingdom": "GB",
"Canada": "CA",
"India": "IN",
"Russia": "RU",
}
COUNTRY_CENTROIDS = {
"China": (104.1954, 35.8617),
"United States": (-95.7129, 37.0902),
"Netherlands": (5.2913, 52.1326),
"Türkiye": (35.2433, 38.9637),
"Australia": (133.7751, -25.2744),
"Japan": (138.2529, 36.2048),
"Singapore": (103.8198, 1.3521),
"Germany": (10.4515, 51.1657),
"France": (2.2137, 46.2276),
"United Kingdom": (-3.436, 55.3781),
"Canada": (-106.3468, 56.1304),
"India": (78.9629, 20.5937),
"Russia": (105.3188, 61.524),
}
CHINA_REGION_METADATA = {
"北京市": ("CN-BJ", 116.4074, 39.9042),
"天津市": ("CN-TJ", 117.2008, 39.0842),
"河北省": ("CN-HE", 114.5025, 38.0455),
"山西省": ("CN-SX", 112.5492, 37.857),
"内蒙古自治区": ("CN-NM", 111.6708, 40.8183),
"辽宁省": ("CN-LN", 123.4315, 41.8057),
"吉林省": ("CN-JL", 125.3245, 43.8868),
"黑龙江省": ("CN-HL", 126.6424, 45.7567),
"上海市": ("CN-SH", 121.4737, 31.2304),
"江苏省": ("CN-JS", 118.7633, 32.0617),
"浙江省": ("CN-ZJ", 120.1551, 30.2741),
"安徽省": ("CN-AH", 117.2272, 31.8206),
"福建省": ("CN-FJ", 119.2965, 26.0745),
"江西省": ("CN-JX", 115.8582, 28.682),
"山东省": ("CN-SD", 117.1201, 36.6512),
"河南省": ("CN-HA", 113.6254, 34.7466),
"湖北省": ("CN-HB", 114.3055, 30.5928),
"湖南省": ("CN-HN", 112.9388, 28.2282),
"广东省": ("CN-GD", 113.2644, 23.1291),
"广西壮族自治区": ("CN-GX", 108.3669, 22.817),
"海南省": ("CN-HI", 110.3312, 20.0311),
"重庆": ("CN-CQ", 106.5516, 29.563),
"重庆市": ("CN-CQ", 106.5516, 29.563),
"四川省": ("CN-SC", 104.0665, 30.5723),
"贵州省": ("CN-GZ", 106.6302, 26.647),
"云南省": ("CN-YN", 102.8329, 24.8801),
"西藏自治区": ("CN-XZ", 91.1322, 29.6604),
"陕西省": ("CN-SN", 108.9398, 34.3416),
"甘肃省": ("CN-GS", 103.8343, 36.0611),
"青海省": ("CN-QH", 101.7782, 36.6171),
"宁夏回族自治区": ("CN-NX", 106.2309, 38.4872),
"新疆维吾尔自治区": ("CN-XJ", 87.6168, 43.8256),
"台湾省": ("CN-TW", 121.5654, 25.033),
"香港特别行政区": ("CN-HK", 114.1694, 22.3193),
"澳门特别行政区": ("CN-MO", 113.5439, 22.1987),
}
CITY_CENTROIDS = {
"南京": (118.7969, 32.0603),
"南京市": (118.7969, 32.0603),
"San Jose": (-121.8863, 37.3382),
"South Holland": (4.493, 52.0208),
"Istanbul": (28.9784, 41.0082),
}
def resolve_geo_location_metadata(location: IpLocation) -> GeoLocationMetadata:
if location.location in {"局域网", "本机"}:
return GeoLocationMetadata(
country="",
country_code="PRIVATE",
region_code="PRIVATE",
longitude=None,
latitude=None,
accuracy_level="private",
)
country = COUNTRY_ALIASES.get(location.country, location.country)
country_code = COUNTRY_CODES.get(country, "")
city_coordinate = CITY_CENTROIDS.get(location.city)
if city_coordinate:
return GeoLocationMetadata(
country=country,
country_code=country_code,
region_code="",
longitude=city_coordinate[0],
latitude=city_coordinate[1],
accuracy_level="city",
)
region_metadata = CHINA_REGION_METADATA.get(location.province)
if country in {"China", "中国"} and region_metadata:
region_code, longitude, latitude = region_metadata
return GeoLocationMetadata(
country="China",
country_code="CN",
region_code=region_code,
longitude=longitude,
latitude=latitude,
accuracy_level="region",
)
country_coordinate = COUNTRY_CENTROIDS.get(country)
if country_coordinate:
return GeoLocationMetadata(
country=country,
country_code=country_code,
region_code="",
longitude=country_coordinate[0],
latitude=country_coordinate[1],
accuracy_level="country",
)
return GeoLocationMetadata(
country=country,
country_code=country_code,
region_code="",
longitude=None,
latitude=None,
accuracy_level="unknown",
)
@@ -0,0 +1,174 @@
"""Bounded third-party coordinate fallback for public source IPs.
IPAddress.my identifies IP2Location.io as its data provider. We use the
provider's documented JSON API instead of scraping the public HTML page.
Only globally routable addresses are eligible, and results are cached so the
monitoring UI does not turn into a per-refresh third-party lookup fan-out.
"""
from __future__ import annotations
import asyncio
import ipaddress
import logging
import math
import time
from dataclasses import dataclass
from typing import Any, Iterable
import httpx
from app.core.config import settings
from app.services.geo_location_metadata import GeoLocationMetadata
from app.services.ip_location import IpLocation
logger = logging.getLogger("ctms.ip_geolocation_fallback")
_API_URL = "https://api.ip2location.io/"
_MAX_RESPONSE_BYTES = 64 * 1024
_MAX_CACHE_ENTRIES = 4096
_NEGATIVE_CACHE_SECONDS = 3600
_cache: dict[str, tuple[float, "ExternalIpLocation | None"]] = {}
_cache_lock = asyncio.Lock()
def _text(value: Any, limit: int = 160) -> str:
candidate = str(value or "").strip()
return "" if candidate in {"", "-", "0"} else candidate[:limit]
def _coordinate(value: Any, *, minimum: float, maximum: float) -> float | None:
try:
number = float(value)
except (TypeError, ValueError):
return None
return number if math.isfinite(number) and minimum <= number <= maximum else None
def _global_ip(value: str | None) -> str | None:
try:
address = ipaddress.ip_address((value or "").strip())
except ValueError:
return None
return str(address) if address.is_global else None
@dataclass(frozen=True)
class ExternalIpLocation:
ip_address: str
country: str
country_code: str
region: str
city: str
isp: str
longitude: float
latitude: float
def merge_ip_location(self, local: IpLocation) -> IpLocation:
country = self.country or local.country
province = self.region or local.province
city = self.city or local.city
isp = self.isp or local.isp
location = " / ".join(part for part in [country, province, city, isp] if part) or local.location
return IpLocation(location=location, country=country, province=province, city=city, isp=isp)
def to_metadata(self) -> GeoLocationMetadata:
return GeoLocationMetadata(
country=self.country,
country_code=self.country_code,
region_code="",
longitude=self.longitude,
latitude=self.latitude,
accuracy_level="city" if self.city else "country",
)
def _parse_response(expected_ip: str, payload: Any) -> ExternalIpLocation | None:
if not isinstance(payload, dict):
return None
response_ip = _global_ip(_text(payload.get("ip"), 45))
if response_ip != expected_ip:
return None
latitude = _coordinate(payload.get("latitude"), minimum=-90, maximum=90)
longitude = _coordinate(payload.get("longitude"), minimum=-180, maximum=180)
if latitude is None or longitude is None:
return None
country_code = _text(payload.get("country_code"), 2).upper()
if len(country_code) != 2:
country_code = ""
return ExternalIpLocation(
ip_address=expected_ip,
country=_text(payload.get("country_name"), 100),
country_code=country_code,
region=_text(payload.get("region_name"), 100),
city=_text(payload.get("city_name"), 100),
isp=_text(payload.get("isp"), 160),
longitude=longitude,
latitude=latitude,
)
async def _fetch_one(
client: httpx.AsyncClient,
semaphore: asyncio.Semaphore,
ip_address: str,
) -> ExternalIpLocation | None:
async with semaphore:
try:
response = await client.get(_API_URL, params={"ip": ip_address, "format": "json"})
response.raise_for_status()
if len(response.content) > _MAX_RESPONSE_BYTES:
return None
return _parse_response(ip_address, response.json())
except (httpx.HTTPError, ValueError):
logger.warning("External IP coordinate fallback unavailable")
return None
async def resolve_external_ip_locations(ip_addresses: Iterable[str]) -> dict[str, ExternalIpLocation]:
if not settings.MONITORING_IP_GEO_FALLBACK_ENABLED or settings.ENV == "test":
return {}
candidates = list(dict.fromkeys(filter(None, (_global_ip(value) for value in ip_addresses))))
if not candidates:
return {}
now = time.monotonic()
resolved: dict[str, ExternalIpLocation] = {}
pending: list[str] = []
async with _cache_lock:
for ip_address in candidates:
cached = _cache.get(ip_address)
if cached and cached[0] > now:
if cached[1] is not None:
resolved[ip_address] = cached[1]
continue
pending.append(ip_address)
pending = pending[: settings.MONITORING_IP_GEO_FALLBACK_MAX_LOOKUPS]
if not pending:
return resolved
headers = {"Accept": "application/json", "User-Agent": "CTMS-IP-Coordinate-Fallback/1.0"}
api_key = (settings.MONITORING_IP_GEO_FALLBACK_API_KEY or "").strip()
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
timeout = httpx.Timeout(settings.MONITORING_IP_GEO_FALLBACK_TIMEOUT_SECONDS)
semaphore = asyncio.Semaphore(min(5, len(pending)))
async with httpx.AsyncClient(timeout=timeout, follow_redirects=False, headers=headers) as client:
fetched = await asyncio.gather(*(_fetch_one(client, semaphore, item) for item in pending))
now = time.monotonic()
async with _cache_lock:
for ip_address, item in zip(pending, fetched):
ttl = settings.MONITORING_IP_GEO_FALLBACK_CACHE_SECONDS if item else _NEGATIVE_CACHE_SECONDS
_cache[ip_address] = (now + ttl, item)
if item is not None:
resolved[ip_address] = item
while len(_cache) > _MAX_CACHE_ENTRIES:
_cache.pop(next(iter(_cache)))
return resolved
def reset_ip_geolocation_fallback_cache() -> None:
_cache.clear()
+2
View File
@@ -9,6 +9,7 @@ import ipaddress
import logging
import sys
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path
from typing import Optional
@@ -125,5 +126,6 @@ class Ip2RegionResolver:
_resolver = Ip2RegionResolver()
@lru_cache(maxsize=8192)
def resolve_ip_location(ip: str | None) -> IpLocation:
return _resolver.lookup(ip)
@@ -0,0 +1,151 @@
"""监测数据留存清理任务。"""
from __future__ import annotations
import asyncio
import logging
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy import delete
from app.core.config import settings
from app.db.session import SessionLocal
from app.models.permission_access_log import PermissionAccessLog
from app.models.permission_metric_snapshot import PermissionMetricSnapshot
from app.models.security_access_log import SecurityAccessLog
from app.models.source_location_snapshot import SourceLocationSnapshot
from app.models.user_login_session import UserLoginSession
logger = logging.getLogger("ctms.monitoring_retention")
@dataclass
class MonitoringRetentionState:
running: bool = False
last_run_started_at: datetime | None = None
last_success_at: datetime | None = None
last_error_at: datetime | None = None
last_error_type: str | None = None
error_count: int = 0
last_deleted: dict[str, int] = field(default_factory=dict)
total_deleted: dict[str, int] = field(default_factory=dict)
def snapshot(self) -> dict[str, Any]:
return {
"running": self.running,
"access_log_retention_days": settings.MONITORING_ACCESS_LOG_RETENTION_DAYS,
"metric_retention_days": settings.MONITORING_METRIC_RETENTION_DAYS,
"login_activity_retention_days": settings.USER_LOGIN_ACTIVITY_RETENTION_DAYS,
"interval_seconds": settings.MONITORING_RETENTION_INTERVAL_SECONDS,
"last_run_started_at": (
self.last_run_started_at.isoformat() if self.last_run_started_at else None
),
"last_success_at": self.last_success_at.isoformat() if self.last_success_at else None,
"last_error_at": self.last_error_at.isoformat() if self.last_error_at else None,
"last_error_type": self.last_error_type,
"error_count": self.error_count,
"last_deleted": dict(self.last_deleted),
"total_deleted": dict(self.total_deleted),
}
_state = MonitoringRetentionState()
def get_monitoring_retention_status() -> dict[str, Any]:
return _state.snapshot()
def _deleted_count(result: Any) -> int:
rowcount = getattr(result, "rowcount", 0)
return max(0, int(rowcount or 0))
async def purge_expired_monitoring_data(
*, now: datetime | None = None
) -> dict[str, int]:
reference_time = now or datetime.now(timezone.utc)
access_cutoff = reference_time - timedelta(
days=settings.MONITORING_ACCESS_LOG_RETENTION_DAYS
)
metric_cutoff = reference_time - timedelta(
days=settings.MONITORING_METRIC_RETENTION_DAYS
)
login_activity_cutoff = reference_time - timedelta(
days=settings.USER_LOGIN_ACTIVITY_RETENTION_DAYS
)
async with SessionLocal() as session:
permission_result = await session.execute(
delete(PermissionAccessLog).where(PermissionAccessLog.created_at < access_cutoff)
)
security_result = await session.execute(
delete(SecurityAccessLog).where(SecurityAccessLog.created_at < access_cutoff)
)
metric_result = await session.execute(
delete(PermissionMetricSnapshot).where(
PermissionMetricSnapshot.bucket_time < metric_cutoff
)
)
source_location_result = await session.execute(
delete(SourceLocationSnapshot).where(
SourceLocationSnapshot.bucket_time < access_cutoff
)
)
login_session_result = await session.execute(
delete(UserLoginSession).where(UserLoginSession.last_seen_at < login_activity_cutoff)
)
await session.commit()
return {
"permission_access_logs": _deleted_count(permission_result),
"security_access_logs": _deleted_count(security_result),
"permission_metric_snapshots": _deleted_count(metric_result),
"source_location_snapshots": _deleted_count(source_location_result),
"user_login_sessions": _deleted_count(login_session_result),
}
async def run_monitoring_retention_once(
*, now: datetime | None = None
) -> dict[str, int]:
_state.last_run_started_at = now or datetime.now(timezone.utc)
try:
deleted = await purge_expired_monitoring_data(now=now)
except Exception as exc:
_state.last_error_at = datetime.now(timezone.utc)
_state.last_error_type = type(exc).__name__
_state.error_count += 1
raise
_state.last_success_at = datetime.now(timezone.utc)
_state.last_deleted = deleted
for key, count in deleted.items():
_state.total_deleted[key] = _state.total_deleted.get(key, 0) + count
return deleted
async def run_monitoring_retention(stop_event: asyncio.Event) -> None:
logger.info("Monitoring retention task started")
_state.running = True
try:
while not stop_event.is_set():
try:
deleted = await run_monitoring_retention_once()
if any(deleted.values()):
logger.info("Purged expired monitoring data: %s", deleted)
except Exception:
logger.exception("Failed to purge expired monitoring data")
try:
await asyncio.wait_for(
stop_event.wait(),
timeout=settings.MONITORING_RETENTION_INTERVAL_SECONDS,
)
except asyncio.TimeoutError:
continue
finally:
_state.running = False
logger.info("Monitoring retention task stopped")
@@ -0,0 +1,170 @@
"""Resolve the monitoring deployment's public network location.
The map server marker is derived from the deployment's public IP instead of a
frontend or configuration coordinate. An explicit public IP is accepted for
restricted networks; otherwise the public frontend hostname and, finally, a
small set of public-IP discovery endpoints are used. Results are cached so the
monitoring API never performs per-request network discovery.
"""
from __future__ import annotations
import asyncio
import ipaddress
import logging
import socket
import time
from dataclasses import dataclass
from urllib.parse import urlparse
import httpx
from app.core.config import settings
from app.services.geo_location_metadata import resolve_geo_location_metadata
from app.services.ip_geolocation_fallback import resolve_external_ip_locations
from app.services.ip_location import resolve_ip_location
logger = logging.getLogger("ctms.monitoring_server_location")
@dataclass(frozen=True)
class MonitoringServerLocation:
name: str
longitude: float
latitude: float
accuracy_level: str
resolution_source: str
def to_public_dict(self) -> dict:
return {
"name": self.name,
"longitude": self.longitude,
"latitude": self.latitude,
"accuracy_level": self.accuracy_level,
"resolution_source": self.resolution_source,
}
_cached_location: MonitoringServerLocation | None = None
_cache_initialized = False
_cache_expires_at = 0.0
_cache_lock = asyncio.Lock()
def _normalize_public_ip(value: str | None) -> str | None:
candidate = (value or "").strip().splitlines()[0] if (value or "").strip() else ""
try:
address = ipaddress.ip_address(candidate)
except ValueError:
return None
return str(address) if address.is_global else None
async def _resolve_frontend_public_ip() -> str | None:
hostname = urlparse(settings.FRONTEND_PUBLIC_URL).hostname
if not hostname:
return None
literal_ip = _normalize_public_ip(hostname)
if literal_ip:
return literal_ip
try:
loop = asyncio.get_running_loop()
records = await loop.getaddrinfo(hostname, None, type=socket.SOCK_STREAM)
except (OSError, socket.gaierror):
return None
candidates = []
for _, _, _, _, socket_address in records:
candidate = _normalize_public_ip(str(socket_address[0]))
if candidate and candidate not in candidates:
candidates.append(candidate)
return next((item for item in candidates if ":" not in item), candidates[0] if candidates else None)
async def _discover_public_ip() -> tuple[str | None, str]:
configured_ip = _normalize_public_ip(settings.MONITORING_SERVER_PUBLIC_IP)
if configured_ip:
return configured_ip, "configured_public_ip"
frontend_ip = await _resolve_frontend_public_ip()
if frontend_ip:
return frontend_ip, "frontend_dns"
if settings.ENV == "test":
return None, "unavailable"
urls = [
item.strip()
for item in settings.MONITORING_PUBLIC_IP_DISCOVERY_URLS.split(",")
if item.strip()
]
timeout = httpx.Timeout(settings.MONITORING_PUBLIC_IP_DISCOVERY_TIMEOUT_SECONDS)
async with httpx.AsyncClient(timeout=timeout, follow_redirects=False) as client:
for url in urls:
try:
response = await client.get(url, headers={"Accept": "text/plain"})
response.raise_for_status()
except httpx.HTTPError:
logger.warning("Public IP discovery endpoint unavailable", extra={"endpoint": url})
continue
discovered_ip = _normalize_public_ip(response.text[:128])
if discovered_ip:
return discovered_ip, "public_ip_discovery"
return None, "unavailable"
async def resolve_monitoring_server_location(*, force: bool = False) -> MonitoringServerLocation | None:
global _cached_location, _cache_initialized, _cache_expires_at
now = time.monotonic()
if not force and _cache_initialized and now < _cache_expires_at:
return _cached_location
async with _cache_lock:
now = time.monotonic()
if not force and _cache_initialized and now < _cache_expires_at:
return _cached_location
public_ip, resolution_source = await _discover_public_ip()
location = None
if public_ip:
ip_location = resolve_ip_location(public_ip)
metadata = resolve_geo_location_metadata(ip_location)
if metadata.longitude is None or metadata.latitude is None:
external_location = (await resolve_external_ip_locations([public_ip])).get(public_ip)
if external_location is not None:
ip_location = external_location.merge_ip_location(ip_location)
metadata = external_location.to_metadata()
if metadata.longitude is not None and metadata.latitude is not None:
country = metadata.country or ip_location.country
name = " / ".join(
part for part in [country, ip_location.province, ip_location.city] if part
) or "公网服务器"
location = MonitoringServerLocation(
name=name,
longitude=metadata.longitude,
latitude=metadata.latitude,
accuracy_level=metadata.accuracy_level,
resolution_source=resolution_source,
)
logger.info(
"Monitoring server location resolved",
extra={"resolution_source": resolution_source, "accuracy_level": metadata.accuracy_level},
)
else:
logger.warning("Deployment public IP has no usable geo coordinates")
else:
logger.warning("Unable to discover the deployment public IP")
_cached_location = location
_cache_initialized = True
cache_seconds = settings.MONITORING_SERVER_LOCATION_CACHE_SECONDS if location else 300
_cache_expires_at = now + cache_seconds
return location
def reset_monitoring_server_location_cache() -> None:
"""Reset process-local state for configuration reloads and tests."""
global _cached_location, _cache_initialized, _cache_expires_at
_cached_location = None
_cache_initialized = False
_cache_expires_at = 0.0
@@ -0,0 +1,78 @@
"""Server-side materialization for time and state driven business reminders."""
from __future__ import annotations
import asyncio
import logging
from sqlalchemy import select, text
from app.core.config import settings
from app.db.session import SessionLocal
from app.models.study_member import StudyMember
from app.models.user import User, UserStatus
from app.services.project_reminder_service import sync_project_reminders
logger = logging.getLogger("ctms.notifications")
NOTIFICATION_SYNC_LOCK_KEY = 0x43544D53 # "CTMS"
async def sync_all_project_reminders_once() -> int:
"""Synchronize every active project member once, with a cross-worker advisory lock."""
async with SessionLocal() as lock_session:
acquired = bool(await lock_session.scalar(
text("SELECT pg_try_advisory_lock(:lock_key)"),
{"lock_key": NOTIFICATION_SYNC_LOCK_KEY},
))
if not acquired:
return 0
try:
memberships = (await lock_session.execute(
select(StudyMember.study_id, StudyMember.user_id)
.join(User, User.id == StudyMember.user_id)
.where(
StudyMember.is_active.is_(True),
User.status == UserStatus.ACTIVE,
)
)).all()
await lock_session.commit()
synced = 0
for study_id, user_id in memberships:
async with SessionLocal() as session:
try:
user = await session.get(User, user_id)
if user is None or not user.is_active:
continue
await sync_project_reminders(session, study_id, user)
synced += 1
except Exception:
await session.rollback()
logger.exception(
"Failed to synchronize reminders for study=%s user=%s",
study_id,
user_id,
)
return synced
finally:
await lock_session.execute(
text("SELECT pg_advisory_unlock(:lock_key)"),
{"lock_key": NOTIFICATION_SYNC_LOCK_KEY},
)
async def run_notification_sync_job(stop_event: asyncio.Event) -> None:
logger.info("Notification synchronization task started")
while not stop_event.is_set():
try:
synced = await sync_all_project_reminders_once()
logger.debug("Synchronized reminders for %s project memberships", synced)
except Exception:
logger.exception("Notification synchronization task failed")
try:
await asyncio.wait_for(
stop_event.wait(),
timeout=settings.NOTIFICATION_SYNC_INTERVAL_SECONDS,
)
except asyncio.TimeoutError:
continue
logger.info("Notification synchronization task stopped")
@@ -0,0 +1,366 @@
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from typing import Iterable
from fastapi import HTTPException, status
from sqlalchemy import func, select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.desktop_notification import DesktopNotificationDelivery
from app.models.notification import Notification
from app.schemas.notification import GeneralNotificationFeed
async def _requeue_desktop_delivery(db: AsyncSession, notification_id: uuid.UUID) -> None:
"""Allow an escalated/reopened reminder to be delivered again by the desktop channel."""
await db.execute(
update(DesktopNotificationDelivery)
.where(DesktopNotificationDelivery.notification_id == notification_id)
.values(
claim_token=None,
claimed_at=None,
delivered_at=None,
)
)
async def create_recipient_notifications(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_ids: Iterable[uuid.UUID],
category: str,
priority: str,
title: str,
message: str,
action_path: str | None,
source_type: str,
source_id: str,
dedupe_key: str,
source_version: str | None = None,
requires_action: bool = True,
due_at: datetime | None = None,
) -> None:
recipients = set(recipient_ids)
if not recipients:
return
existing = set((await db.scalars(
select(Notification.recipient_id).where(
Notification.recipient_id.in_(recipients),
Notification.dedupe_key == dedupe_key,
)
)).all())
for recipient_id in recipients - existing:
db.add(Notification(
study_id=study_id,
recipient_id=recipient_id,
category=category,
priority=priority,
title=title,
message=message,
action_path=action_path,
source_type=source_type,
source_id=source_id,
source_version=source_version,
dedupe_key=dedupe_key,
requires_action=requires_action,
due_at=due_at,
))
async def sync_state_notification(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
category: str,
priority: str,
title: str,
message: str,
action_path: str,
source_type: str,
source_id: str,
source_version: str,
active: bool,
due_at: datetime | None = None,
requires_action: bool = True,
dedupe_key: str | None = None,
) -> None:
"""Synchronize one actionable business state into a recipient reminder."""
dedupe_key = dedupe_key or f"state:{source_type}:{source_id}"
item = await db.scalar(select(Notification).where(
Notification.recipient_id == recipient_id,
Notification.dedupe_key == dedupe_key,
))
now = datetime.now(timezone.utc)
if not active:
if item and item.resolved_at is None:
item.resolved_at = now
item.read_at = item.read_at or now
return
if item is None:
db.add(Notification(
study_id=study_id,
recipient_id=recipient_id,
category=category,
priority=priority,
title=title,
message=message,
action_path=action_path,
source_type=source_type,
source_id=source_id,
source_version=source_version,
dedupe_key=dedupe_key,
requires_action=requires_action,
due_at=due_at,
))
return
should_notify_again = item.resolved_at is not None or item.source_version != source_version
item.study_id = study_id
item.category = category
item.priority = priority
item.title = title
item.message = message
item.action_path = action_path
item.source_version = source_version
item.requires_action = requires_action
item.due_at = due_at
if should_notify_again:
item.read_at = None
item.created_at = now
await _requeue_desktop_delivery(db, item.id)
item.resolved_at = None
async def sync_aggregate_notification(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
category: str,
priority: str,
title: str,
message: str,
action_path: str,
source_type: str,
source_id: str,
count: int,
due_at: datetime | None = None,
) -> None:
dedupe_key = f"aggregate:{source_type}:{source_id}"
item = await db.scalar(select(Notification).where(
Notification.recipient_id == recipient_id,
Notification.dedupe_key == dedupe_key,
))
now = datetime.now(timezone.utc)
if count <= 0:
if item and item.resolved_at is None:
item.resolved_at = now
item.read_at = item.read_at or now
return
version = str(count)
if item is None:
db.add(Notification(
study_id=study_id,
recipient_id=recipient_id,
category=category,
priority=priority,
title=title,
message=message,
action_path=action_path,
source_type=source_type,
source_id=source_id,
source_version=version,
dedupe_key=dedupe_key,
requires_action=True,
due_at=due_at,
))
return
previous_count = int(item.source_version or 0)
item.category = category
item.priority = priority
item.title = title
item.message = message
item.action_path = action_path
item.source_version = version
item.due_at = due_at
if item.resolved_at is not None or count > previous_count:
item.read_at = None
item.created_at = now
await _requeue_desktop_delivery(db, item.id)
item.resolved_at = None
async def resolve_missing_recipient_sources(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
source_type: str,
active_source_ids: set[str],
) -> None:
filters = [
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.source_type == source_type,
Notification.resolved_at.is_(None),
]
if active_source_ids:
filters.append(Notification.source_id.not_in(active_source_ids))
now = datetime.now(timezone.utc)
await db.execute(
update(Notification)
.where(*filters)
.values(resolved_at=now, read_at=func.coalesce(Notification.read_at, now))
)
async def resolve_source_notifications(
db: AsyncSession,
*,
source_type: str,
source_id: str,
recipient_id: uuid.UUID | None = None,
) -> None:
filters = [
Notification.source_type == source_type,
Notification.source_id == source_id,
Notification.resolved_at.is_(None),
]
if recipient_id is not None:
filters.append(Notification.recipient_id == recipient_id)
now = datetime.now(timezone.utc)
await db.execute(
update(Notification)
.where(*filters)
.values(resolved_at=now, read_at=func.coalesce(Notification.read_at, now))
)
async def resolve_recipient_study_notifications(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
) -> None:
now = datetime.now(timezone.utc)
await db.execute(
update(Notification)
.where(
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.resolved_at.is_(None),
)
.values(resolved_at=now, read_at=func.coalesce(Notification.read_at, now))
)
async def list_feed(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
skip: int = 0,
limit: int = 10,
category: str | None = None,
unread_only: bool = False,
requires_action: bool | None = None,
) -> GeneralNotificationFeed:
active_filter = (
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.resolved_at.is_(None),
)
unread_count = int(await db.scalar(
select(func.count(Notification.id)).where(*active_filter, Notification.read_at.is_(None))
) or 0)
filters = [
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.resolved_at.is_(None),
]
if category:
filters.append(Notification.category == category)
if unread_only:
filters.append(Notification.read_at.is_(None))
if requires_action is not None:
filters.append(Notification.requires_action.is_(requires_action))
total_count = int(await db.scalar(
select(func.count(Notification.id)).where(*filters)
) or 0)
items = (await db.scalars(
select(Notification)
.where(*filters)
.order_by(
Notification.read_at.is_not(None),
Notification.created_at.desc(),
)
.offset(max(0, skip))
.limit(max(1, min(limit, 100)))
)).all()
return GeneralNotificationFeed(
unread_count=unread_count,
total_count=total_count,
items=list(items),
)
async def mark_read(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
notification_id: uuid.UUID,
) -> Notification:
item = await db.scalar(select(Notification).where(
Notification.id == notification_id,
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.resolved_at.is_(None),
))
if item is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="通知不存在")
if item.read_at is None:
now = datetime.now(timezone.utc)
item.read_at = now
if not item.requires_action:
item.resolved_at = now
await db.commit()
await db.refresh(item)
return item
async def mark_all_read(
db: AsyncSession,
*,
study_id: uuid.UUID,
recipient_id: uuid.UUID,
) -> None:
now = datetime.now(timezone.utc)
await db.execute(
update(Notification)
.where(
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.resolved_at.is_(None),
Notification.read_at.is_(None),
Notification.requires_action.is_(False),
)
.values(read_at=now, resolved_at=now)
)
await db.execute(
update(Notification)
.where(
Notification.study_id == study_id,
Notification.recipient_id == recipient_id,
Notification.resolved_at.is_(None),
Notification.read_at.is_(None),
Notification.requires_action.is_(True),
)
.values(read_at=now)
)
await db.commit()
@@ -0,0 +1,417 @@
from __future__ import annotations
import hashlib
import hmac
import json
import uuid
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any
from urllib.parse import urlsplit, urlunsplit
import httpx
from fastapi import HTTPException, Request, status
from jose import JWTError, jwt
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.models.collaboration import (
CollaborationCallbackReceipt,
CollaborationFile,
CollaborationRevision,
CollaborationSession,
CollaborationShareLink,
)
from app.models.user import User
from app.schemas.collaboration import CollaborationCallbackPayload, CollaborationEditorConfigRead
from app.services import collaboration_service, onlyoffice_service
def collaboration_document_key(file_id: uuid.UUID, generation: int) -> str:
fingerprint = f"{settings.ONLYOFFICE_INSTANCE_ID or ''}:collaboration:{file_id}:{generation}"
return f"ctms-collab-{hashlib.sha256(fingerprint.encode('utf-8')).hexdigest()}"
def _content_url(session_id: uuid.UUID) -> str:
return (
f"{settings.ONLYOFFICE_STORAGE_BASE_URL.rstrip('/')}"
f"/internal/onlyoffice/collaboration/sessions/{session_id}/content"
)
def _callback_url(session_id: uuid.UUID) -> str:
return (
f"{settings.ONLYOFFICE_STORAGE_BASE_URL.rstrip('/')}"
f"/internal/onlyoffice/collaboration/sessions/{session_id}/callback"
)
async def _active_session(
db: AsyncSession, item: CollaborationFile, user_id: uuid.UUID
) -> CollaborationSession:
session = await db.scalar(
select(CollaborationSession).where(
CollaborationSession.file_id == item.id,
CollaborationSession.generation == item.generation,
).order_by(CollaborationSession.created_at.desc())
)
if session:
if session.status != "ACTIVE":
session.status = "ACTIVE"
session.closed_at = None
await db.commit()
await db.refresh(session)
return session
if not item.current_revision_id:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="协作文件尚无可编辑内容")
session = CollaborationSession(
file_id=item.id,
base_revision_id=item.current_revision_id,
document_key=collaboration_document_key(item.id, item.generation),
generation=item.generation,
started_by=user_id,
)
db.add(session)
await db.commit()
await db.refresh(session)
return session
async def build_editor_config(
db: AsyncSession, item: CollaborationFile, user
) -> CollaborationEditorConfigRead:
await onlyoffice_service.ensure_onlyoffice_available()
revision = await db.get(CollaborationRevision, item.current_revision_id)
if not revision or not Path(revision.file_uri).exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="协作文件内容不存在")
can_edit = await collaboration_service.can_edit_file(db, item, user)
can_request_edit = await collaboration_service.can_request_edit_file(db, item, user)
can_download = await collaboration_service.can_export_file(db, item, user)
can_save_as = can_download and await collaboration_service.can_create_file(db, item, user)
session = await _active_session(db, item, user.id)
now = datetime.now(timezone.utc)
expires_at = now + timedelta(seconds=settings.ONLYOFFICE_CONFIG_TTL_SECONDS)
config: dict[str, Any] = {
"type": "desktop",
"documentType": item.file_type,
"document": {
"fileType": item.extension,
"key": session.document_key,
"title": item.title,
"url": _content_url(session.id),
"permissions": {
"chat": False,
"copy": can_download,
"comment": can_edit,
"download": can_download,
# In view mode ONLYOFFICE displays "Edit current file" only
# when edit=true and onRequestEditRights is registered. CTMS
# handles that event as an approval request, not an escalation.
"edit": can_edit or can_request_edit,
"fillForms": False,
"modifyContentControl": can_edit,
"modifyFilter": can_edit,
"print": can_download,
"protect": False,
"review": False,
},
},
"editorConfig": {
"callbackUrl": _callback_url(session.id),
"coEditing": {"mode": "fast", "change": False},
"customization": {
"autosave": True,
"chat": False,
"comments": can_edit,
"forcesave": can_edit,
"help": False,
"plugins": False,
},
"lang": "zh-CN",
"mode": "edit" if can_edit else "view",
"user": {"id": str(user.id), "name": user.full_name},
},
}
config["token"] = jwt.encode(
{**config, "iat": int(now.timestamp()), "exp": int(expires_at.timestamp())},
settings.ONLYOFFICE_JWT_SECRET or "",
algorithm="HS256",
)
return CollaborationEditorConfigRead(
file_id=item.id,
file_name=item.title,
access_mode="edit" if can_edit else "view",
can_save_as=can_save_as,
can_download=can_download,
can_request_edit=can_request_edit,
expires_at=expires_at,
config=config,
)
async def build_shared_editor_config(
db: AsyncSession,
item: CollaborationFile,
link: CollaborationShareLink,
*,
client_id: str,
display_name: str,
) -> CollaborationEditorConfigRead:
await onlyoffice_service.ensure_onlyoffice_available()
revision = await db.get(CollaborationRevision, item.current_revision_id)
if not revision or not Path(revision.file_uri).exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="共享文件内容不存在")
can_edit = link.access_mode == "EDIT"
session = await _active_session(db, item, item.owner_id)
now = datetime.now(timezone.utc)
expires_at = now + timedelta(seconds=settings.ONLYOFFICE_CONFIG_TTL_SECONDS)
if link.expires_at and link.expires_at < expires_at:
expires_at = link.expires_at
external_user_id = f"share-{link.id.hex[:12]}-{client_id[:32]}"
config: dict[str, Any] = {
"type": "desktop",
"documentType": item.file_type,
"document": {
"fileType": item.extension,
"key": session.document_key,
"title": item.title,
"url": _content_url(session.id),
"permissions": {
"chat": False,
"copy": item.allow_export,
"comment": can_edit,
"download": item.allow_export,
"edit": can_edit,
"fillForms": False,
"modifyContentControl": can_edit,
"modifyFilter": can_edit,
"print": item.allow_export,
"protect": False,
"review": False,
},
},
"editorConfig": {
"callbackUrl": _callback_url(session.id),
"coEditing": {"mode": "fast", "change": False},
"customization": {
"autosave": can_edit,
"chat": False,
"comments": can_edit,
"forcesave": can_edit,
"help": False,
"plugins": False,
},
"lang": "zh-CN",
"mode": "edit" if can_edit else "view",
"user": {"id": external_user_id, "name": display_name},
},
}
config["token"] = jwt.encode(
{**config, "iat": int(now.timestamp()), "exp": int(expires_at.timestamp())},
settings.ONLYOFFICE_JWT_SECRET or "",
algorithm="HS256",
)
return CollaborationEditorConfigRead(
file_id=item.id,
file_name=item.title,
access_mode="edit" if can_edit else "view",
can_save_as=False,
can_download=item.allow_export,
expires_at=expires_at,
config=config,
)
async def get_session_content(
db: AsyncSession, session_id: uuid.UUID, authorization: str | None
) -> tuple[CollaborationRevision, CollaborationFile]:
session = await db.get(CollaborationSession, session_id)
if not session:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="协作会话不存在")
onlyoffice_service.validate_outbox_token(authorization, _content_url(session_id))
revision = await db.get(CollaborationRevision, session.base_revision_id)
item = await db.get(CollaborationFile, session.file_id)
if not revision or not item or not Path(revision.file_uri).exists():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="协作文件内容不存在")
return revision, item
def validate_callback_token(token: str | None, payload: CollaborationCallbackPayload) -> dict[str, Any]:
if not token:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="缺少 ONLYOFFICE 回调签名")
value = token.strip()
if " " in value:
scheme, credential = value.split(" ", 1)
if scheme.lower() != "bearer":
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="ONLYOFFICE 回调签名格式无效")
value = credential.strip()
try:
decoded = onlyoffice_service.decode_onlyoffice_token(value)
except JWTError as exc:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="ONLYOFFICE 回调签名无效") from exc
signed = decoded.get("payload") if isinstance(decoded.get("payload"), dict) else decoded
if not isinstance(signed, dict):
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="ONLYOFFICE 回调载荷无效")
signed_key = signed.get("key")
signed_status = signed.get("status")
if not isinstance(signed_key, str) or not hmac.compare_digest(signed_key, payload.key):
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="ONLYOFFICE 回调 key 不匹配")
if not isinstance(signed_status, int) or signed_status != payload.status:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="ONLYOFFICE 回调状态不匹配")
if payload.url:
signed_url = signed.get("url")
if not isinstance(signed_url, str) or not hmac.compare_digest(signed_url, payload.url):
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="ONLYOFFICE 回调文件地址不匹配")
return decoded
def _callback_fingerprint(payload: CollaborationCallbackPayload) -> str:
normalized = payload.model_dump(mode="json", exclude_none=True)
return hashlib.sha256(json.dumps(normalized, sort_keys=True, separators=(",", ":")).encode()).hexdigest()
def _url_origin_matches(actual, expected) -> bool:
actual_port = actual.port or (443 if actual.scheme == "https" else 80)
expected_port = expected.port or (443 if expected.scheme == "https" else 80)
return (
actual.scheme == expected.scheme
and actual.hostname
and actual.hostname.lower() == (expected.hostname or "").lower()
and actual_port == expected_port
)
def _validate_result_url(url: str) -> str:
actual = urlsplit(url)
if (
actual.scheme not in {"http", "https"}
or actual.username
or actual.password
or actual.fragment
or not actual.hostname
):
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail="ONLYOFFICE 保存地址不受信任")
internal = urlsplit(settings.ONLYOFFICE_INTERNAL_URL.rstrip("/"))
if _url_origin_matches(actual, internal):
return urlunsplit((internal.scheme, internal.netloc, actual.path, actual.query, ""))
public = urlsplit(settings.FRONTEND_PUBLIC_URL.rstrip("/"))
proxy_prefix = "/onlyoffice/"
if not _url_origin_matches(actual, public) or not actual.path.startswith(proxy_prefix):
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail="ONLYOFFICE 保存地址不受信任")
internal_path = f"{internal.path.rstrip('/')}/{actual.path[len(proxy_prefix):]}"
return urlunsplit((internal.scheme, internal.netloc, internal_path, actual.query, ""))
async def _download_result(url: str) -> bytes:
download_url = _validate_result_url(url)
try:
async with httpx.AsyncClient(timeout=30.0, follow_redirects=False) as client:
async with client.stream("GET", download_url) as response:
if response.status_code != status.HTTP_200_OK:
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="ONLYOFFICE 保存文件下载失败")
content = bytearray()
async for chunk in response.aiter_bytes():
content.extend(chunk)
if len(content) > settings.COLLABORATION_MAX_FILE_BYTES:
raise HTTPException(status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, detail="ONLYOFFICE 保存文件超出限制")
except httpx.HTTPError as exc:
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="ONLYOFFICE 保存文件下载失败") from exc
return bytes(content)
async def _callback_user(
db: AsyncSession, payload: CollaborationCallbackPayload, session: CollaborationSession
) -> User | None:
has_public_share_user = False
for value in payload.users:
if value.startswith("share-"):
has_public_share_user = True
continue
try:
user = await db.get(User, uuid.UUID(value))
except (ValueError, TypeError):
user = None
if user:
return user
if has_public_share_user:
return None
user = await db.get(User, session.started_by)
if not user:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="协作会话用户不存在")
return user
async def process_callback(
db: AsyncSession,
session_id: uuid.UUID,
payload: CollaborationCallbackPayload,
) -> dict[str, int]:
session = await db.scalar(
select(CollaborationSession).where(CollaborationSession.id == session_id).with_for_update()
)
if not session:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="协作会话不存在")
if not hmac.compare_digest(session.document_key, payload.key):
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="协作会话 key 不匹配")
fingerprint = _callback_fingerprint(payload)
duplicate = await db.scalar(select(CollaborationCallbackReceipt.id).where(
CollaborationCallbackReceipt.session_id == session.id,
CollaborationCallbackReceipt.fingerprint == fingerprint,
))
if duplicate:
return {"error": 0}
item = await db.get(CollaborationFile, session.file_id)
if not item:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="协作文件不存在")
session.last_callback_at = datetime.now(timezone.utc)
session.active_users = json.dumps(payload.users, ensure_ascii=True)
result = "ACKNOWLEDGED"
saved_revision_id = None
if payload.status in {2, 6}:
if not payload.url:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail="ONLYOFFICE 保存回调缺少文件地址")
if session.generation != item.generation:
result = "STALE"
else:
content = await _download_result(payload.url)
actor = await _callback_user(db, payload, session)
source = "SESSION_CLOSE" if payload.status == 2 else "FORCE_SAVE"
if actor is None:
source = "SHARE_SESSION_CLOSE" if payload.status == 2 else "SHARE_FORCE_SAVE"
revision, created = await collaboration_service.append_revision(
db, item, content, source=source, created_by=actor.id if actor else None
)
saved_revision_id = revision.id
result = "SAVED" if created else "UNCHANGED"
if payload.status == 2:
item.generation += 1
session.status = "CLOSED"
session.closed_at = datetime.now(timezone.utc)
else:
# 强制保存不结束当前共同编辑会话;同步基线可保证 Document
# Server 缓存重建时仍从最近一次持久化内容恢复。
session.base_revision_id = revision.id
elif payload.status == 4:
session.status = "CLOSED"
session.closed_at = datetime.now(timezone.utc)
result = "UNCHANGED"
elif payload.status in {3, 7}:
session.status = "ERROR"
result = "ERROR"
db.add(CollaborationCallbackReceipt(
session_id=session.id,
fingerprint=fingerprint,
callback_status=payload.status,
result=result,
revision_id=saved_revision_id,
))
await db.commit()
# ONLYOFFICE 要求回调处理器在接收并记录状态后固定确认成功。
# status 3/7 表示文档服务自身保存失败,不应通过 error=1 制造重试环。
return {"error": 0}
+235
View File
@@ -0,0 +1,235 @@
from __future__ import annotations
import asyncio
import hashlib
import hmac
import time
import uuid
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, Literal
import httpx
from fastapi import status
from jose import JWTError, jwt
from app.core.config import settings
from app.core.exceptions import AppException
from app.schemas.onlyoffice import OnlyOfficePreviewConfigRead
OnlyOfficeDocumentType = Literal["word", "cell", "slide"]
OnlyOfficeResourceType = Literal["attachment", "version", "collaboration_revision"]
WORD_FORMATS = frozenset({
"doc", "docx", "docm", "dot", "dotx", "dotm", "odt", "ott", "rtf", "txt", "wps", "wpt",
})
CELL_FORMATS = frozenset({
"xls", "xlsx", "xlsm", "xlsb", "xlt", "xltx", "xltm", "ods", "ots", "csv", "et", "ett",
})
SLIDE_FORMATS = frozenset({
"ppt", "pptx", "pptm", "pps", "ppsx", "ppsm", "pot", "potx", "potm", "odp", "otp", "dps", "dpt",
})
_health_lock = asyncio.Lock()
_health_checked_at = 0.0
_health_available = False
_HEALTH_CACHE_SECONDS = 15.0
def onlyoffice_error(code: str, message: str, status_code: int) -> AppException:
return AppException(code=code, message=message, status_code=status_code)
def office_format_for_filename(filename: str) -> tuple[str, OnlyOfficeDocumentType] | None:
suffix = Path(filename).suffix.lower().lstrip(".")
if suffix in WORD_FORMATS:
return suffix, "word"
if suffix in CELL_FORMATS:
return suffix, "cell"
if suffix in SLIDE_FORMATS:
return suffix, "slide"
return None
def onlyoffice_content_url(resource_type: OnlyOfficeResourceType, resource_id: uuid.UUID) -> str:
plural = {
"attachment": "attachments",
"version": "versions",
"collaboration_revision": "collaboration-revisions",
}[resource_type]
return (
f"{settings.ONLYOFFICE_STORAGE_BASE_URL.rstrip('/')}"
f"/internal/onlyoffice/{plural}/{resource_id}/content"
)
def onlyoffice_document_key(
resource_type: OnlyOfficeResourceType,
resource_id: uuid.UUID,
*,
file_hash: str | None = None,
) -> str:
instance_id = (settings.ONLYOFFICE_INSTANCE_ID or "").strip()
fingerprint = f"{instance_id}{resource_type}{resource_id}"
if resource_type == "version":
fingerprint = f"{fingerprint}{file_hash or ''}"
return f"ctms-{hashlib.sha256(fingerprint.encode('utf-8')).hexdigest()}"
async def ensure_onlyoffice_available() -> None:
global _health_available, _health_checked_at
if not settings.ONLYOFFICE_ENABLED:
raise onlyoffice_error(
"ONLYOFFICE_DISABLED",
"Office 预览服务尚未启用",
status.HTTP_503_SERVICE_UNAVAILABLE,
)
now = time.monotonic()
if now - _health_checked_at < _HEALTH_CACHE_SECONDS:
if _health_available:
return
raise onlyoffice_error(
"ONLYOFFICE_UNAVAILABLE",
"Office 预览服务暂不可用,请稍后重试",
status.HTTP_503_SERVICE_UNAVAILABLE,
)
async with _health_lock:
now = time.monotonic()
if now - _health_checked_at >= _HEALTH_CACHE_SECONDS:
available = False
try:
async with httpx.AsyncClient(timeout=2.0, follow_redirects=False) as client:
response = await client.get(f"{settings.ONLYOFFICE_INTERNAL_URL.rstrip('/')}/healthcheck")
available = response.status_code == status.HTTP_200_OK
except httpx.HTTPError:
available = False
_health_available = available
_health_checked_at = now
if not _health_available:
raise onlyoffice_error(
"ONLYOFFICE_UNAVAILABLE",
"Office 预览服务暂不可用,请稍后重试",
status.HTTP_503_SERVICE_UNAVAILABLE,
)
def build_preview_config(
*,
resource_type: OnlyOfficeResourceType,
resource_id: uuid.UUID,
file_name: str,
user_id: uuid.UUID,
user_name: str,
file_hash: str | None = None,
) -> OnlyOfficePreviewConfigRead:
format_info = office_format_for_filename(file_name)
if not format_info:
raise onlyoffice_error(
"ONLYOFFICE_FORMAT_UNSUPPORTED",
"该文件格式不支持 Office 在线预览",
status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
)
file_type, document_type = format_info
now = datetime.now(timezone.utc)
expires_at = now + timedelta(seconds=settings.ONLYOFFICE_CONFIG_TTL_SECONDS)
config: dict[str, Any] = {
"type": "desktop",
"documentType": document_type,
"document": {
"fileType": file_type,
"key": onlyoffice_document_key(resource_type, resource_id, file_hash=file_hash),
"title": file_name,
"url": onlyoffice_content_url(resource_type, resource_id),
"permissions": {
"copy": False,
"comment": False,
"download": False,
"edit": False,
"fillForms": False,
"modifyContentControl": False,
"modifyFilter": False,
"print": False,
"review": False,
},
},
"editorConfig": {
"coEditing": {"mode": "strict", "change": False},
"customization": {
"chat": False,
"comments": False,
"forcesave": False,
},
"lang": "zh-CN",
"mode": "view",
"user": {"id": str(user_id), "name": user_name},
},
}
token_payload = {
**config,
"iat": int(now.timestamp()),
"exp": int(expires_at.timestamp()),
}
config["token"] = jwt.encode(
token_payload,
settings.ONLYOFFICE_JWT_SECRET or "",
algorithm="HS256",
)
return OnlyOfficePreviewConfigRead(
resource_type=resource_type,
resource_id=resource_id,
file_name=file_name,
expires_at=expires_at,
config=config,
)
def validate_outbox_token(token: str | None, expected_url: str) -> dict[str, Any]:
if not token:
raise onlyoffice_error(
"ONLYOFFICE_SOURCE_UNAUTHORIZED",
"无法验证 Office 文件请求",
status.HTTP_401_UNAUTHORIZED,
)
token = token.strip()
if " " in token:
scheme, credential = token.split(" ", 1)
if scheme.lower() != "bearer" or not credential.strip():
raise onlyoffice_error(
"ONLYOFFICE_SOURCE_UNAUTHORIZED",
"无法验证 Office 文件请求",
status.HTTP_401_UNAUTHORIZED,
)
token = credential.strip()
try:
payload = decode_onlyoffice_token(token)
except JWTError as exc:
raise onlyoffice_error(
"ONLYOFFICE_SOURCE_UNAUTHORIZED",
"无法验证 Office 文件请求",
status.HTTP_401_UNAUTHORIZED,
) from exc
request_payload = payload.get("payload")
token_url = request_payload.get("url") if isinstance(request_payload, dict) else None
if not isinstance(token_url, str) or not hmac.compare_digest(token_url, expected_url):
raise onlyoffice_error(
"ONLYOFFICE_SOURCE_URL_MISMATCH",
"Office 文件请求地址不匹配",
status.HTTP_401_UNAUTHORIZED,
)
return payload
def decode_onlyoffice_token(token: str) -> dict[str, Any]:
header = jwt.get_unverified_header(token)
if header.get("alg") != "HS256":
raise JWTError("unexpected algorithm")
return jwt.decode(
token,
settings.ONLYOFFICE_JWT_SECRET or "",
algorithms=["HS256"],
options={"verify_aud": False},
)
+106 -40
View File
@@ -13,6 +13,7 @@ import asyncio
import logging
import uuid
from datetime import datetime, timezone
from typing import Any
from app.db.session import SessionLocal
from app.models.permission_access_log import PermissionAccessLog
@@ -22,85 +23,150 @@ logger = logging.getLogger("ctms.permission_log_writer")
BATCH_SIZE = 50
FLUSH_INTERVAL = 5.0
QUEUE_MAX_SIZE = 10000
WRITE_ATTEMPTS = 2
_QUEUE_STOP = object()
class PermissionLogWriter:
def __init__(self):
self._queue: asyncio.Queue[dict] = asyncio.Queue(maxsize=QUEUE_MAX_SIZE)
self._queue: asyncio.Queue[dict | object] = asyncio.Queue(maxsize=QUEUE_MAX_SIZE)
self._task: asyncio.Task | None = None
self._stopping = False
self._started_at: datetime | None = None
self._last_success_at: datetime | None = None
self._last_error_at: datetime | None = None
self._last_error_type: str | None = None
self._accepted_entries = 0
self._written_entries = 0
self._dropped_entries = 0
self._failed_batches = 0
self._failed_entries = 0
self._retry_count = 0
def enqueue(self, entry: dict) -> None:
if self._stopping:
self._dropped_entries += 1
logger.warning("Permission log writer is stopping, dropping entry")
return
try:
self._queue.put_nowait(entry)
self._accepted_entries += 1
except asyncio.QueueFull:
self._dropped_entries += 1
logger.warning("Permission log queue full, dropping entry")
async def start(self) -> None:
if self._task and not self._task.done():
return
self._stopping = False
self._started_at = datetime.now(timezone.utc)
self._task = asyncio.create_task(self._flush_loop())
logger.info("PermissionLogWriter started")
async def stop(self) -> None:
if self._task:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
await self._drain()
if self._task and not self._task.done():
self._stopping = True
await self._queue.put(_QUEUE_STOP)
await self._task
self._task = None
logger.info("PermissionLogWriter stopped")
async def _flush_loop(self) -> None:
while True:
batch = await self._collect_batch()
batch, should_stop = await self._collect_batch()
if batch:
await self._write_batch(batch)
if should_stop:
break
async def _collect_batch(self) -> list[dict]:
async def _collect_batch(self) -> tuple[list[dict], bool]:
batch: list[dict] = []
try:
first = await asyncio.wait_for(self._queue.get(), timeout=FLUSH_INTERVAL)
if first is _QUEUE_STOP:
return batch, True
batch.append(first)
except asyncio.TimeoutError:
return batch
return batch, False
while len(batch) < BATCH_SIZE:
try:
item = self._queue.get_nowait()
if item is _QUEUE_STOP:
return batch, True
batch.append(item)
except asyncio.QueueEmpty:
break
return batch
return batch, False
async def _write_batch(self, batch: list[dict]) -> None:
try:
async with SessionLocal() as session:
for entry in batch:
log = PermissionAccessLog(
id=uuid.uuid4(),
study_id=entry["study_id"],
user_id=entry["user_id"],
endpoint_key=entry["endpoint_key"],
role=entry["role"],
allowed=entry["allowed"],
elapsed_ms=entry["elapsed_ms"],
ip_address=entry.get("ip_address"),
created_at=entry.get("created_at", datetime.now(timezone.utc)),
)
session.add(log)
await session.commit()
except Exception:
logger.exception("Failed to write permission access log batch (%d entries)", len(batch))
async def _drain(self) -> None:
batch: list[dict] = []
while not self._queue.empty():
async def _write_batch(self, batch: list[dict]) -> bool:
for attempt in range(WRITE_ATTEMPTS):
try:
batch.append(self._queue.get_nowait())
except asyncio.QueueEmpty:
break
if batch:
await self._write_batch(batch)
async with SessionLocal() as session:
for entry in batch:
log = PermissionAccessLog(
id=uuid.uuid4(),
study_id=entry["study_id"],
user_id=entry["user_id"],
endpoint_key=entry["endpoint_key"],
role=entry["role"],
allowed=entry["allowed"],
elapsed_ms=entry["elapsed_ms"],
ip_address=entry.get("ip_address"),
user_agent=entry.get("user_agent"),
client_type=entry.get("client_type"),
client_version=entry.get("client_version"),
client_platform=entry.get("client_platform"),
build_channel=entry.get("build_channel"),
build_commit=entry.get("build_commit"),
request_headers=entry.get("request_headers"),
request_snapshot=entry.get("request_snapshot"),
request_id=entry.get("request_id"),
created_at=entry.get("created_at", datetime.now(timezone.utc)),
)
session.add(log)
await session.commit()
except Exception as exc:
self._last_error_at = datetime.now(timezone.utc)
self._last_error_type = type(exc).__name__
if attempt + 1 < WRITE_ATTEMPTS:
self._retry_count += 1
logger.warning(
"Permission access log batch write failed; retrying (%d entries)",
len(batch),
)
await asyncio.sleep(0)
continue
self._failed_batches += 1
self._failed_entries += len(batch)
logger.exception(
"Failed to write permission access log batch after retries (%d entries)",
len(batch),
)
return False
self._written_entries += len(batch)
self._last_success_at = datetime.now(timezone.utc)
return True
return False
def stats(self) -> dict[str, Any]:
return {
"running": bool(self._task and not self._task.done()),
"stopping": self._stopping,
"queue_size": self._queue.qsize(),
"queue_capacity": self._queue.maxsize,
"accepted_entries": self._accepted_entries,
"written_entries": self._written_entries,
"dropped_entries": self._dropped_entries,
"failed_batches": self._failed_batches,
"failed_entries": self._failed_entries,
"retry_count": self._retry_count,
"started_at": self._started_at.isoformat() if self._started_at else None,
"last_success_at": self._last_success_at.isoformat() if self._last_success_at else None,
"last_error_at": self._last_error_at.isoformat() if self._last_error_at else None,
"last_error_type": self._last_error_type,
}
_writer: PermissionLogWriter | None = None
@@ -0,0 +1,541 @@
from __future__ import annotations
import uuid
from datetime import date, datetime, time, timedelta, timezone
from zoneinfo import ZoneInfo
from sqlalchemy import and_, func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_cra_site_scope, is_system_admin
from app.core.project_permissions import role_has_api_permission
from app.crud import member as member_crud
from app.crud import site as site_crud
from app.models.acknowledgement import Acknowledgement, AcknowledgementType
from app.models.ae import AdverseEvent
from app.models.collaboration import CollaborationEditRequest, CollaborationFile, CollaborationMember
from app.models.distribution import Distribution, DistributionStatus, DistributionTargetType
from app.models.document import Document, DocumentStatus
from app.models.document_version import DocumentVersion
from app.models.milestone import Milestone
from app.models.monitoring_visit_issue import MonitoringVisitIssue
from app.models.site import Site
from app.models.subject import Subject
from app.models.user import User
from app.models.visit import Visit
from app.services import notification_service
BUSINESS_TIMEZONE = ZoneInfo("Asia/Shanghai")
RISK_DUE_SOON_DAYS = 3
MILESTONE_DUE_SOON_DAYS = 7
VISIT_WINDOW_DUE_SOON_DAYS = 3
def _date_due_at(value: date | None) -> datetime | None:
if value is None:
return None
return datetime.combine(value, time.max, tzinfo=BUSINESS_TIMEZONE).astimezone(timezone.utc)
def _as_utc(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
async def _count_and_earliest(db: AsyncSession, model, *filters):
return (await db.execute(
select(func.count(model.id), func.min(model.report_due_date if model is AdverseEvent else model.due_at))
.where(*filters)
)).one()
async def _sync_risk_reminders(
db: AsyncSession,
*,
study_id: uuid.UUID,
user,
role: str,
current_time: datetime,
) -> None:
can_read_aes = is_system_admin(user) or await role_has_api_permission(
db, study_id, role, "subject_aes:read"
)
can_read_monitoring = is_system_admin(user) or await role_has_api_permission(
db, study_id, role, "monitoring_issues:read"
)
cra_scope = await get_cra_site_scope(db, study_id, user)
site_ids = cra_scope[0] if cra_scope else None
today = current_time.astimezone(BUSINESS_TIMEZONE).date()
due_soon_date = today + timedelta(days=RISK_DUE_SOON_DAYS)
ae_scope = [AdverseEvent.study_id == study_id, AdverseEvent.status != "CLOSED"]
if site_ids is not None:
ae_scope.append(AdverseEvent.site_id.in_(site_ids))
overdue_aes = 0
overdue_ae_due = None
due_soon_aes = 0
due_soon_ae_due = None
if can_read_aes:
overdue_aes, overdue_ae_due = await _count_and_earliest(
db,
AdverseEvent,
*ae_scope,
AdverseEvent.report_due_date.is_not(None),
AdverseEvent.report_due_date < today,
)
due_soon_aes, due_soon_ae_due = await _count_and_earliest(
db,
AdverseEvent,
*ae_scope,
AdverseEvent.report_due_date >= today,
AdverseEvent.report_due_date <= due_soon_date,
)
await notification_service.sync_aggregate_notification(
db,
study_id=study_id,
recipient_id=user.id,
category="RISK_OVERDUE_AE",
priority="URGENT",
title="逾期 AE 待处理",
message=f"当前有 {overdue_aes} 条逾期 AE 需要跟进",
action_path="/risk-issues/sae",
source_type="RISK_OVERDUE_AE",
source_id=str(study_id),
count=int(overdue_aes) if can_read_aes else 0,
due_at=_date_due_at(overdue_ae_due),
)
await notification_service.sync_aggregate_notification(
db,
study_id=study_id,
recipient_id=user.id,
category="RISK_AE_DUE_SOON",
priority="HIGH",
title="AE 上报时限临近",
message=f"未来 {RISK_DUE_SOON_DAYS} 天内有 {due_soon_aes} 条 AE 到达上报时限",
action_path="/risk-issues/sae",
source_type="RISK_AE_DUE_SOON",
source_id=str(study_id),
count=int(due_soon_aes) if can_read_aes else 0,
due_at=_date_due_at(due_soon_ae_due),
)
monitoring_scope = [
MonitoringVisitIssue.study_id == study_id,
MonitoringVisitIssue.status == "OPEN",
MonitoringVisitIssue.due_at.is_not(None),
]
if site_ids is not None:
monitoring_scope.append(MonitoringVisitIssue.site_id.in_(site_ids))
overdue_monitoring = 0
overdue_monitoring_due = None
due_soon_monitoring = 0
due_soon_monitoring_due = None
if can_read_monitoring:
overdue_monitoring, overdue_monitoring_due = await _count_and_earliest(
db,
MonitoringVisitIssue,
*monitoring_scope,
MonitoringVisitIssue.due_at < current_time,
)
due_soon_monitoring, due_soon_monitoring_due = await _count_and_earliest(
db,
MonitoringVisitIssue,
*monitoring_scope,
MonitoringVisitIssue.due_at >= current_time,
MonitoringVisitIssue.due_at <= current_time + timedelta(days=RISK_DUE_SOON_DAYS),
)
await notification_service.sync_aggregate_notification(
db,
study_id=study_id,
recipient_id=user.id,
category="RISK_OVERDUE_MONITORING",
priority="HIGH",
title="监查问题已逾期",
message=f"当前有 {overdue_monitoring} 条监查问题已超过计划解决日期",
action_path="/risk-issues/monitoring-visits",
source_type="RISK_OVERDUE_MONITORING",
source_id=str(study_id),
count=int(overdue_monitoring) if can_read_monitoring else 0,
due_at=_as_utc(overdue_monitoring_due),
)
await notification_service.sync_aggregate_notification(
db,
study_id=study_id,
recipient_id=user.id,
category="RISK_MONITORING_DUE_SOON",
priority="HIGH",
title="监查问题整改时限临近",
message=f"未来 {RISK_DUE_SOON_DAYS} 天内有 {due_soon_monitoring} 条监查问题到期",
action_path="/risk-issues/monitoring-visits",
source_type="RISK_MONITORING_DUE_SOON",
source_id=str(study_id),
count=int(due_soon_monitoring) if can_read_monitoring else 0,
due_at=_as_utc(due_soon_monitoring_due),
)
async def _sync_document_distribution_reminders(
db: AsyncSession,
*,
study_id: uuid.UUID,
user,
role: str,
current_time: datetime,
) -> None:
can_read_documents = is_system_admin(user) or await role_has_api_permission(
db, study_id, role, "documents:read"
)
if not can_read_documents:
await notification_service.resolve_missing_recipient_sources(
db,
study_id=study_id,
recipient_id=user.id,
source_type="DOCUMENT_DISTRIBUTION",
active_source_ids=set(),
)
return
site_ids = await site_crud.list_ids_by_contact_user(db, study_id, user.id)
target_filters = [
and_(
Distribution.target_type == DistributionTargetType.USER,
Distribution.target_id == str(user.id),
),
and_(
Distribution.target_type == DistributionTargetType.ROLE,
Distribution.target_id == role,
),
]
if site_ids:
target_filters.append(and_(
Distribution.target_type == DistributionTargetType.SITE,
Distribution.target_id.in_({str(site_id) for site_id in site_ids}),
))
rows = (await db.execute(
select(Distribution, Document, DocumentVersion)
.join(Document, Distribution.document_id == Document.id)
.join(DocumentVersion, Distribution.version_id == DocumentVersion.id)
.outerjoin(
Acknowledgement,
and_(
Acknowledgement.distribution_id == Distribution.id,
Acknowledgement.user_id == user.id,
Acknowledgement.ack_type == AcknowledgementType.RECEIVED,
),
)
.where(
Document.trial_id == study_id,
Document.status == DocumentStatus.ACTIVE,
Distribution.status == DistributionStatus.ACTIVE,
Acknowledgement.id.is_(None),
or_(*target_filters),
)
)).all()
active_source_ids: set[str] = set()
for distribution, document, version in rows:
source_id = str(distribution.id)
active_source_ids.add(source_id)
distribution_due_at = _as_utc(distribution.due_at)
if distribution_due_at and distribution_due_at < current_time:
stage = "OVERDUE"
category = "DOCUMENT_ACK_OVERDUE"
priority = "URGENT"
title = "文件回执已逾期"
elif distribution_due_at and distribution_due_at <= current_time + timedelta(days=RISK_DUE_SOON_DAYS):
stage = "DUE_SOON"
category = "DOCUMENT_ACK_DUE_SOON"
priority = "HIGH"
title = "文件回执即将到期"
else:
stage = "PENDING"
category = "DOCUMENT_DISTRIBUTION"
priority = "NORMAL"
title = "有新的文件版本待接收"
due_version = distribution_due_at.isoformat() if distribution_due_at else "none"
await notification_service.sync_state_notification(
db,
study_id=study_id,
recipient_id=user.id,
category=category,
priority=priority,
title=title,
message=f"{document.title}{version.version_no} 已分发,请完成接收回执",
action_path=f"/documents/{document.id}",
source_type="DOCUMENT_DISTRIBUTION",
source_id=source_id,
source_version=f"{stage}:{due_version}",
active=True,
due_at=distribution_due_at,
)
await notification_service.resolve_missing_recipient_sources(
db,
study_id=study_id,
recipient_id=user.id,
source_type="DOCUMENT_DISTRIBUTION",
active_source_ids=active_source_ids,
)
async def _sync_milestone_reminders(
db: AsyncSession,
*,
study_id: uuid.UUID,
user,
role: str,
current_time: datetime,
) -> None:
can_read = is_system_admin(user) or await role_has_api_permission(
db, study_id, role, "project_milestones:read"
)
active_source_ids: set[str] = set()
if can_read:
today = current_time.astimezone(BUSINESS_TIMEZONE).date()
due_soon_date = today + timedelta(days=MILESTONE_DUE_SOON_DAYS)
milestones = (await db.scalars(
select(Milestone).where(
Milestone.study_id == study_id,
Milestone.owner_id == user.id,
Milestone.status != "DONE",
)
)).all()
for milestone in milestones:
due_date = milestone.adjusted_end_date or milestone.planned_date
if due_date is None or due_date > due_soon_date:
continue
source_id = str(milestone.id)
active_source_ids.add(source_id)
overdue = due_date < today
await notification_service.sync_state_notification(
db,
study_id=study_id,
recipient_id=user.id,
category="MILESTONE_OVERDUE" if overdue else "MILESTONE_DUE_SOON",
priority="HIGH" if overdue else "NORMAL",
title="项目里程碑已逾期" if overdue else "项目里程碑即将到期",
message=f"{milestone.name}”计划日期为 {due_date.isoformat()}",
action_path="/project/milestones",
source_type="PROJECT_MILESTONE",
source_id=source_id,
source_version=f"{'OVERDUE' if overdue else 'DUE_SOON'}:{due_date.isoformat()}:{milestone.status}",
active=True,
due_at=_date_due_at(due_date),
)
await notification_service.resolve_missing_recipient_sources(
db,
study_id=study_id,
recipient_id=user.id,
source_type="PROJECT_MILESTONE",
active_source_ids=active_source_ids,
)
async def _sync_visit_window_reminders(
db: AsyncSession,
*,
study_id: uuid.UUID,
user,
role: str,
current_time: datetime,
) -> None:
"""Project active visit windows to the members who can actually maintain them."""
can_manage_visits = is_system_admin(user) or await role_has_api_permission(
db, study_id, role, "visits:update"
)
active_source_ids: set[str] = set()
if can_manage_visits:
cra_scope = await get_cra_site_scope(db, study_id, user)
site_ids = cra_scope[0] if cra_scope else None
filters = [
Visit.study_id == study_id,
Visit.actual_date.is_(None),
Visit.status.not_in(["DONE", "CANCELLED"]),
Subject.status.not_in(["COMPLETED", "DROPPED"]),
Site.is_active.is_(True),
]
if site_ids is not None:
filters.append(Subject.site_id.in_(site_ids))
rows = (await db.execute(
select(Visit, Subject)
.join(Subject, Subject.id == Visit.subject_id)
.join(Site, Site.id == Subject.site_id)
.where(*filters)
)).all()
today = current_time.astimezone(BUSINESS_TIMEZONE).date()
due_soon_date = today + timedelta(days=VISIT_WINDOW_DUE_SOON_DAYS)
for visit, subject in rows:
window_start = visit.window_start or visit.planned_date
window_end = visit.window_end or visit.planned_date
if window_start is None or window_end is None or window_start > window_end:
continue
if window_end < today:
category = "VISIT_WINDOW_MISSED"
priority = "HIGH"
title = "受试者访视已错过窗口"
message = "一个受试者访视已错过计划窗口,请进入详情核对并处理"
stage = "MISSED"
elif window_start <= due_soon_date:
category = "VISIT_WINDOW_DUE_SOON"
priority = "NORMAL"
title = "受试者访视窗口临近"
message = "一个受试者访视已进入近期窗口,请及时核对访视安排"
stage = "DUE_SOON"
else:
continue
source_id = str(visit.id)
active_source_ids.add(source_id)
await notification_service.sync_state_notification(
db,
study_id=study_id,
recipient_id=user.id,
category=category,
priority=priority,
title=title,
message=message,
action_path=f"/subjects/{subject.id}",
source_type="SUBJECT_VISIT_WINDOW",
source_id=source_id,
source_version=f"{stage}:{window_start.isoformat()}:{window_end.isoformat()}",
active=True,
due_at=_date_due_at(window_end),
)
await notification_service.resolve_missing_recipient_sources(
db,
study_id=study_id,
recipient_id=user.id,
source_type="SUBJECT_VISIT_WINDOW",
active_source_ids=active_source_ids,
)
async def _sync_collaboration_edit_request_reminders(
db: AsyncSession,
*,
study_id: uuid.UUID,
user,
role: str,
) -> None:
can_read = is_system_admin(user) or await role_has_api_permission(
db, study_id, role, "collaboration:read"
)
if not can_read:
await notification_service.resolve_missing_recipient_sources(
db,
study_id=study_id,
recipient_id=user.id,
source_type="COLLABORATION_EDIT_REQUEST",
active_source_ids=set(),
)
return
manages_file = or_(
CollaborationFile.owner_id == user.id,
select(CollaborationMember.id).where(
CollaborationMember.file_id == CollaborationFile.id,
CollaborationMember.user_id == user.id,
CollaborationMember.role == "MANAGER",
).exists(),
)
rows = (await db.execute(
select(CollaborationEditRequest, CollaborationFile, User)
.join(CollaborationFile, CollaborationFile.id == CollaborationEditRequest.file_id)
.join(User, User.id == CollaborationEditRequest.requester_id)
.where(
CollaborationFile.study_id == study_id,
CollaborationFile.status == "ACTIVE",
CollaborationFile.deleted_at.is_(None),
CollaborationEditRequest.status == "PENDING",
manages_file,
)
)).all()
active_source_ids: set[str] = set()
for request, item, requester in rows:
source_id = str(request.id)
active_source_ids.add(source_id)
requester_name = str(requester.full_name or requester.email or "项目成员")
await notification_service.sync_state_notification(
db,
study_id=study_id,
recipient_id=user.id,
category="COLLABORATION_EDIT_REQUEST",
priority="NORMAL",
title="新的编辑权限申请",
message=f"{requester_name} 申请编辑“{item.title}",
action_path=f"/knowledge/collaboration?editRequestFile={item.id}",
source_type="COLLABORATION_EDIT_REQUEST",
source_id=source_id,
source_version="PENDING",
active=True,
dedupe_key=f"collaboration-edit-request:{request.id}",
)
await notification_service.resolve_missing_recipient_sources(
db,
study_id=study_id,
recipient_id=user.id,
source_type="COLLABORATION_EDIT_REQUEST",
active_source_ids=active_source_ids,
)
async def sync_project_reminders(
db: AsyncSession,
study_id: uuid.UUID,
user,
*,
now: datetime | None = None,
) -> None:
membership = await member_crud.get_member(db, study_id, user.id)
if not membership or not membership.is_active:
await notification_service.resolve_recipient_study_notifications(
db,
study_id=study_id,
recipient_id=user.id,
)
await db.commit()
return
role = membership.role_in_study
current_time = now or datetime.now(timezone.utc)
if current_time.tzinfo is None:
current_time = current_time.replace(tzinfo=timezone.utc)
else:
current_time = current_time.astimezone(timezone.utc)
await _sync_risk_reminders(
db,
study_id=study_id,
user=user,
role=role,
current_time=current_time,
)
await _sync_document_distribution_reminders(
db,
study_id=study_id,
user=user,
role=role,
current_time=current_time,
)
await _sync_milestone_reminders(
db,
study_id=study_id,
user=user,
role=role,
current_time=current_time,
)
await _sync_visit_window_reminders(
db,
study_id=study_id,
user=user,
role=role,
current_time=current_time,
)
await _sync_collaboration_edit_request_reminders(
db,
study_id=study_id,
user=user,
role=role,
)
await db.commit()
@@ -6,99 +6,170 @@ import asyncio
import logging
import uuid
from datetime import datetime, timezone
from typing import Any
from app.db.session import SessionLocal
from app.models.security_access_log import SecurityAccessLog
from app.services.ip_location import resolve_ip_location
from app.services.security_events import classify_security_event
logger = logging.getLogger("ctms.security_access_log_writer")
BATCH_SIZE = 100
FLUSH_INTERVAL = 3.0
QUEUE_MAX_SIZE = 20000
WRITE_ATTEMPTS = 2
_QUEUE_STOP = object()
class SecurityAccessLogWriter:
def __init__(self) -> None:
self._queue: asyncio.Queue[dict] = asyncio.Queue(maxsize=QUEUE_MAX_SIZE)
self._queue: asyncio.Queue[dict | object] = asyncio.Queue(maxsize=QUEUE_MAX_SIZE)
self._task: asyncio.Task | None = None
self._stopping = False
self._started_at: datetime | None = None
self._last_success_at: datetime | None = None
self._last_error_at: datetime | None = None
self._last_error_type: str | None = None
self._accepted_entries = 0
self._written_entries = 0
self._dropped_entries = 0
self._failed_batches = 0
self._failed_entries = 0
self._retry_count = 0
def enqueue(self, entry: dict) -> None:
if self._stopping:
self._dropped_entries += 1
logger.warning("Security access log writer is stopping, dropping entry")
return
try:
self._queue.put_nowait(entry)
self._accepted_entries += 1
except asyncio.QueueFull:
self._dropped_entries += 1
logger.warning("Security access log queue full, dropping entry")
async def start(self) -> None:
if self._task and not self._task.done():
return
self._stopping = False
self._started_at = datetime.now(timezone.utc)
self._task = asyncio.create_task(self._flush_loop())
logger.info("SecurityAccessLogWriter started")
async def stop(self) -> None:
if self._task:
self._task.cancel()
try:
await self._task
except asyncio.CancelledError:
pass
await self._drain()
if self._task and not self._task.done():
self._stopping = True
await self._queue.put(_QUEUE_STOP)
await self._task
self._task = None
logger.info("SecurityAccessLogWriter stopped")
async def _flush_loop(self) -> None:
while True:
batch = await self._collect_batch()
batch, should_stop = await self._collect_batch()
if batch:
await self._write_batch(batch)
if should_stop:
break
async def _collect_batch(self) -> list[dict]:
async def _collect_batch(self) -> tuple[list[dict], bool]:
batch: list[dict] = []
try:
first = await asyncio.wait_for(self._queue.get(), timeout=FLUSH_INTERVAL)
if first is _QUEUE_STOP:
return batch, True
batch.append(first)
except asyncio.TimeoutError:
return batch
return batch, False
while len(batch) < BATCH_SIZE:
try:
batch.append(self._queue.get_nowait())
item = self._queue.get_nowait()
if item is _QUEUE_STOP:
return batch, True
batch.append(item)
except asyncio.QueueEmpty:
break
return batch
return batch, False
async def _write_batch(self, batch: list[dict]) -> None:
try:
async with SessionLocal() as session:
for entry in batch:
session.add(
SecurityAccessLog(
id=uuid.uuid4(),
method=entry["method"],
async def _write_batch(self, batch: list[dict]) -> bool:
for attempt in range(WRITE_ATTEMPTS):
try:
async with SessionLocal() as session:
for entry in batch:
classification = classify_security_event(
path=entry["path"],
status_code=entry["status_code"],
elapsed_ms=entry["elapsed_ms"],
client_ip=entry.get("client_ip"),
user_agent=entry.get("user_agent"),
client_type=entry.get("client_type"),
client_version=entry.get("client_version"),
client_platform=entry.get("client_platform"),
build_channel=entry.get("build_channel"),
build_commit=entry.get("build_commit"),
auth_status=entry["auth_status"],
user_identifier=entry.get("user_identifier"),
created_at=entry.get("created_at", datetime.now(timezone.utc)),
ip_location=resolve_ip_location(entry.get("client_ip")),
)
session.add(
SecurityAccessLog(
id=uuid.uuid4(),
method=entry["method"],
path=entry["path"],
status_code=entry["status_code"],
elapsed_ms=entry["elapsed_ms"],
client_ip=entry.get("client_ip"),
user_agent=entry.get("user_agent"),
client_type=entry.get("client_type"),
client_version=entry.get("client_version"),
client_platform=entry.get("client_platform"),
build_channel=entry.get("build_channel"),
build_commit=entry.get("build_commit"),
request_headers=entry.get("request_headers"),
request_snapshot=entry.get("request_snapshot"),
request_id=entry.get("request_id"),
category=classification["category"],
severity=classification["severity"],
auth_status=entry["auth_status"],
user_identifier=entry.get("user_identifier"),
created_at=entry.get("created_at", datetime.now(timezone.utc)),
)
)
await session.commit()
except Exception as exc:
self._last_error_at = datetime.now(timezone.utc)
self._last_error_type = type(exc).__name__
if attempt + 1 < WRITE_ATTEMPTS:
self._retry_count += 1
logger.warning(
"Security access log batch write failed; retrying (%d entries)",
len(batch),
)
await session.commit()
except Exception:
logger.exception("Failed to write security access log batch (%d entries)", len(batch))
await asyncio.sleep(0)
continue
self._failed_batches += 1
self._failed_entries += len(batch)
logger.exception(
"Failed to write security access log batch after retries (%d entries)",
len(batch),
)
return False
self._written_entries += len(batch)
self._last_success_at = datetime.now(timezone.utc)
return True
return False
async def _drain(self) -> None:
batch: list[dict] = []
while not self._queue.empty():
try:
batch.append(self._queue.get_nowait())
except asyncio.QueueEmpty:
break
if batch:
await self._write_batch(batch)
def stats(self) -> dict[str, Any]:
return {
"running": bool(self._task and not self._task.done()),
"stopping": self._stopping,
"queue_size": self._queue.qsize(),
"queue_capacity": self._queue.maxsize,
"accepted_entries": self._accepted_entries,
"written_entries": self._written_entries,
"dropped_entries": self._dropped_entries,
"failed_batches": self._failed_batches,
"failed_entries": self._failed_entries,
"retry_count": self._retry_count,
"started_at": self._started_at.isoformat() if self._started_at else None,
"last_success_at": self._last_success_at.isoformat() if self._last_success_at else None,
"last_error_at": self._last_error_at.isoformat() if self._last_error_at else None,
"last_error_type": self._last_error_type,
}
_writer: SecurityAccessLogWriter | None = None
+37
View File
@@ -0,0 +1,37 @@
"""Security event classification shared by writers and monitoring APIs."""
from __future__ import annotations
from app.services.ip_location import IpLocation
SENSITIVE_PROBE_MARKERS = (
"/.env",
".env",
"/.git",
".git/config",
"backup",
"config.php",
"wp-config",
"database.yml",
)
def classify_security_event(
*,
path: str,
status_code: int,
auth_status: str,
ip_location: IpLocation | None = None,
) -> dict[str, str]:
normalized_path = (path or "").lower()
if any(marker in normalized_path for marker in SENSITIVE_PROBE_MARKERS):
return {"category": "PROBE", "severity": "CRITICAL"}
if status_code >= 500:
return {"category": "SERVER_ERROR", "severity": "HIGH"}
if auth_status == "INVALID_TOKEN":
return {"category": "INVALID_TOKEN", "severity": "MEDIUM"}
if auth_status == "ANONYMOUS" and status_code in {401, 403}:
return {"category": "ANONYMOUS_API", "severity": "MEDIUM"}
if status_code == 404:
return {"category": "NOT_FOUND_NOISE", "severity": "LOW"}
return {"category": "OTHER", "severity": "LOW"}
@@ -0,0 +1,292 @@
"""Hourly source-location aggregation and timeline reads."""
from __future__ import annotations
import asyncio
import hashlib
import hmac
import logging
import uuid
from collections import defaultdict
from datetime import datetime, timedelta, timezone
from typing import Literal
from sqlalchemy import delete, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.db.session import SessionLocal
from app.models.permission_access_log import PermissionAccessLog
from app.models.security_access_log import SecurityAccessLog
from app.models.source_location_snapshot import SourceLocationSnapshot
from app.services.geo_location_metadata import resolve_geo_location_metadata
from app.services.ip_geolocation_fallback import resolve_external_ip_locations
from app.services.ip_location import resolve_ip_location
logger = logging.getLogger("ctms.source_location_aggregator")
def _normalize_datetime(value: datetime) -> datetime:
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
def _identity_hash(kind: str, value: str | None) -> str:
if not value:
return ""
payload = f"source-location:{kind}:{value}".encode("utf-8")
return hmac.new(settings.JWT_SECRET_KEY.encode("utf-8"), payload, hashlib.sha256).hexdigest()
async def aggregate_source_location_hour(bucket_start: datetime, bucket_end: datetime) -> int:
bucket_start = _normalize_datetime(bucket_start)
bucket_end = _normalize_datetime(bucket_end)
async with SessionLocal() as session:
matching_security_request = select(SecurityAccessLog.id).where(
PermissionAccessLog.request_id.is_not(None),
SecurityAccessLog.request_id == PermissionAccessLog.request_id,
).exists()
permission_rows = (
await session.execute(
select(
PermissionAccessLog.ip_address,
PermissionAccessLog.user_id,
func.count().filter(PermissionAccessLog.allowed.is_(True)).label("allowed_count"),
func.count().filter(PermissionAccessLog.allowed.is_(False)).label("denied_count"),
func.min(PermissionAccessLog.created_at).label("first_seen_at"),
func.max(PermissionAccessLog.created_at).label("last_seen_at"),
)
.where(
PermissionAccessLog.created_at >= bucket_start,
PermissionAccessLog.created_at < bucket_end,
PermissionAccessLog.ip_address.is_not(None),
~matching_security_request,
)
.group_by(PermissionAccessLog.ip_address, PermissionAccessLog.user_id)
)
).all()
security_rows = (
await session.execute(
select(
SecurityAccessLog.client_ip,
SecurityAccessLog.user_identifier,
func.count().filter(SecurityAccessLog.status_code < 400).label("allowed_count"),
func.count().filter(SecurityAccessLog.status_code >= 400).label("denied_count"),
func.count()
.filter(
SecurityAccessLog.category.is_not(None),
~SecurityAccessLog.category.in_(("OTHER", "NOT_FOUND_NOISE")),
)
.label("security_event_count"),
func.count()
.filter(SecurityAccessLog.severity.in_(("HIGH", "CRITICAL")))
.label("high_risk_count"),
func.count()
.filter(
SecurityAccessLog.status_code >= 400,
SecurityAccessLog.auth_status.in_(("INVALID_TOKEN", "ANONYMOUS")),
)
.label("auth_failure_count"),
func.min(SecurityAccessLog.created_at).label("first_seen_at"),
func.max(SecurityAccessLog.created_at).label("last_seen_at"),
)
.where(
SecurityAccessLog.created_at >= bucket_start,
SecurityAccessLog.created_at < bucket_end,
SecurityAccessLog.client_ip.is_not(None),
)
.group_by(SecurityAccessLog.client_ip, SecurityAccessLog.user_identifier)
)
).all()
aggregates: dict[tuple[str, str], dict] = {}
def merge_row(
ip_address: str,
user_identity: str,
allowed_count: int,
denied_count: int,
first_seen_at: datetime,
last_seen_at: datetime,
*,
security_event_count: int = 0,
high_risk_count: int = 0,
auth_failure_count: int = 0,
) -> None:
key = (ip_address, user_identity)
row = aggregates.setdefault(
key,
{
"ip_address": ip_address,
"user_identity": user_identity,
"allowed_count": 0,
"denied_count": 0,
"security_event_count": 0,
"high_risk_count": 0,
"auth_failure_count": 0,
"first_seen_at": _normalize_datetime(first_seen_at),
"last_seen_at": _normalize_datetime(last_seen_at),
},
)
row["allowed_count"] += int(allowed_count or 0)
row["denied_count"] += int(denied_count or 0)
row["security_event_count"] += int(security_event_count or 0)
row["high_risk_count"] += int(high_risk_count or 0)
row["auth_failure_count"] += int(auth_failure_count or 0)
row["first_seen_at"] = min(row["first_seen_at"], _normalize_datetime(first_seen_at))
row["last_seen_at"] = max(row["last_seen_at"], _normalize_datetime(last_seen_at))
for ip_address, user_id, allowed, denied, first_seen, last_seen in permission_rows:
merge_row(str(ip_address), str(user_id or ""), allowed, denied, first_seen, last_seen)
for ip_address, user_identifier, allowed, denied, security_events, high_risk, auth_failures, first_seen, last_seen in security_rows:
merge_row(
str(ip_address),
str(user_identifier or ""),
allowed,
denied,
first_seen,
last_seen,
security_event_count=security_events,
high_risk_count=high_risk,
auth_failure_count=auth_failures,
)
resolved_locations = {
ip_address: resolve_ip_location(ip_address)
for ip_address, _user_identity in aggregates
}
resolved_metadata = {
ip_address: resolve_geo_location_metadata(ip_info)
for ip_address, ip_info in resolved_locations.items()
}
missing_coordinate_ips = [
ip_address
for ip_address, metadata in resolved_metadata.items()
if metadata.accuracy_level != "private"
and (metadata.longitude is None or metadata.latitude is None)
]
external_locations = await resolve_external_ip_locations(missing_coordinate_ips)
for ip_address, external_location in external_locations.items():
resolved_locations[ip_address] = external_location.merge_ip_location(resolved_locations[ip_address])
resolved_metadata[ip_address] = external_location.to_metadata()
await session.execute(
delete(SourceLocationSnapshot).where(SourceLocationSnapshot.bucket_time == bucket_start)
)
for row in aggregates.values():
ip_info = resolved_locations[row["ip_address"]]
metadata = resolved_metadata[row["ip_address"]]
longitude = metadata.longitude
latitude = metadata.latitude
country = metadata.country or ip_info.country
location = " / ".join(
part for part in [country, ip_info.province, ip_info.city] if part
) or ip_info.location or "未知"
session.add(
SourceLocationSnapshot(
id=uuid.uuid4(),
bucket_time=bucket_start,
ip_hash=_identity_hash("ip", row["ip_address"]),
user_hash=_identity_hash("user", row["user_identity"]),
country=country,
country_code=metadata.country_code,
province=ip_info.province,
region_code=metadata.region_code,
city=ip_info.city,
isp=ip_info.isp,
location=location,
longitude=longitude,
latitude=latitude,
accuracy_level=metadata.accuracy_level,
allowed_count=row["allowed_count"],
denied_count=row["denied_count"],
security_event_count=row["security_event_count"],
high_risk_count=row["high_risk_count"],
auth_failure_count=row["auth_failure_count"],
first_seen_at=row["first_seen_at"],
last_seen_at=row["last_seen_at"],
)
)
await session.commit()
return len(aggregates)
async def get_source_location_timeline(
db: AsyncSession,
*,
start_at: datetime,
end_at: datetime,
granularity: Literal["hour", "day"],
) -> list[dict]:
rows = (
await db.execute(
select(SourceLocationSnapshot)
.where(
SourceLocationSnapshot.bucket_time >= start_at,
SourceLocationSnapshot.bucket_time < end_at,
)
.order_by(SourceLocationSnapshot.bucket_time)
)
).scalars().all()
buckets: dict[datetime, dict] = defaultdict(
lambda: {
"allowed_count": 0,
"denied_count": 0,
"security_event_count": 0,
"high_risk_count": 0,
"ip_hashes": set(),
"user_hashes": set(),
}
)
for row in rows:
bucket = _normalize_datetime(row.bucket_time)
if granularity == "day":
bucket = bucket.replace(hour=0, minute=0, second=0, microsecond=0)
else:
bucket = bucket.replace(minute=0, second=0, microsecond=0)
item = buckets[bucket]
item["allowed_count"] += row.allowed_count
item["denied_count"] += row.denied_count
item["security_event_count"] += row.security_event_count
item["high_risk_count"] += row.high_risk_count
item["ip_hashes"].add(row.ip_hash)
if row.user_hash:
item["user_hashes"].add(row.user_hash)
return [
{
"bucket_time": bucket.isoformat(),
"total_count": item["allowed_count"] + item["denied_count"],
"allowed_count": item["allowed_count"],
"denied_count": item["denied_count"],
"security_event_count": item["security_event_count"],
"high_risk_count": item["high_risk_count"],
"unique_ip_count": len(item["ip_hashes"]),
"unique_user_count": len(item["user_hashes"]),
}
for bucket, item in sorted(buckets.items())
]
async def run_hourly_source_location_aggregation(stop_event: asyncio.Event) -> None:
logger.info("Source location aggregator started")
now = datetime.now(timezone.utc)
completed_hour = now.replace(minute=0, second=0, microsecond=0)
try:
await aggregate_source_location_hour(completed_hour - timedelta(hours=1), completed_hour)
except Exception:
logger.exception("Failed to backfill source location snapshot for %s", completed_hour)
while not stop_event.is_set():
now = datetime.now(timezone.utc)
next_hour = now.replace(minute=0, second=0, microsecond=0) + timedelta(hours=1)
try:
await asyncio.wait_for(stop_event.wait(), timeout=(next_hour - now).total_seconds())
break
except asyncio.TimeoutError:
pass
try:
await aggregate_source_location_hour(next_hour - timedelta(hours=1), next_hour)
except Exception:
logger.exception("Failed to aggregate source locations for %s", next_hour)
logger.info("Source location aggregator stopped")
+233
View File
@@ -0,0 +1,233 @@
"""Server-authoritative login session activity for the admin account list."""
from __future__ import annotations
import uuid
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Iterable
from sqlalchemy import select, update
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import settings
from app.core.request_context import resolve_client_ip
from app.models.user_login_session import UserLoginSession
from app.services.ip_location import resolve_ip_location
@dataclass(frozen=True)
class UserLoginSummary:
status: str = "OFFLINE"
last_login_at: datetime | None = None
last_seen_at: datetime | None = None
client_type: str | None = None
active_session_count: int = 0
def _now() -> datetime:
return datetime.now(timezone.utc)
def _as_utc(value: datetime) -> datetime:
return value.replace(tzinfo=timezone.utc) if value.tzinfo is None else value.astimezone(timezone.utc)
def _text_header(headers: Any, name: str, limit: int) -> str | None:
value = (headers.get(name) or "").strip()
return value[:limit] or None
def session_id_from_payload(payload: dict[str, Any]) -> uuid.UUID:
raw_session_id = payload.get("sid")
try:
return uuid.UUID(str(raw_session_id))
except (TypeError, ValueError):
# Tokens issued before session tracking are mapped to a deterministic
# legacy session without persisting the token or a token-derived value.
legacy_key = ":".join(
[
str(payload.get("sub") or ""),
str(payload.get("orig_iat") or payload.get("iat") or ""),
str(payload.get("client_type") or "web"),
]
)
return uuid.uuid5(uuid.NAMESPACE_URL, f"ctms:legacy-session:{legacy_key}")
def session_client_type(payload: dict[str, Any], headers: Any) -> str:
from_payload = (payload.get("client_type") or "").strip().lower()
from_header = (headers.get("x-ctms-client-type") or "").strip().lower()
return "desktop" if "desktop" in {from_payload, from_header} else "web"
def login_activity_status(session: UserLoginSession, *, reference_time: datetime | None = None) -> str:
if session.ended_at is not None:
return "ENDED"
cutoff = (reference_time or _now()) - timedelta(seconds=settings.USER_SESSION_ONLINE_SECONDS)
return "ONLINE" if _as_utc(session.last_seen_at) >= cutoff else "OFFLINE"
def login_activity_payload(
session: UserLoginSession,
*,
reference_time: datetime | None = None,
) -> dict[str, Any]:
ip_location = resolve_ip_location(session.login_ip) if session.login_ip else None
return {
"id": session.id,
"client_type": session.client_type,
"client_platform": session.client_platform,
"client_version": session.client_version,
"client_source": session.client_source,
"login_ip": session.login_ip,
"ip_location": ip_location.location if ip_location else None,
"login_at": session.login_at,
"last_seen_at": session.last_seen_at,
"ended_at": session.ended_at,
"end_reason": session.end_reason,
"activity_status": login_activity_status(session, reference_time=reference_time),
}
async def create_login_session(
db: AsyncSession,
*,
session_id: uuid.UUID,
user_id: uuid.UUID,
request: Any,
login_at: datetime | None = None,
) -> UserLoginSession:
occurred_at = login_at or _now()
session = UserLoginSession(
id=session_id,
user_id=user_id,
client_type=session_client_type({}, request.headers),
client_platform=_text_header(request.headers, "x-ctms-client-platform", 32),
client_version=_text_header(request.headers, "x-ctms-client-version", 64),
client_source=_text_header(request.headers, "x-ctms-client-source", 32),
login_ip=resolve_client_ip(request),
login_at=occurred_at,
last_seen_at=occurred_at,
)
db.add(session)
await db.commit()
await db.refresh(session)
return session
async def touch_login_session(
db: AsyncSession,
*,
user_id: uuid.UUID,
payload: dict[str, Any],
request: Any,
) -> UserLoginSession | None:
session_id = session_id_from_payload(payload)
session = await db.get(UserLoginSession, session_id)
now = _now()
if session is None:
session = UserLoginSession(
id=session_id,
user_id=user_id,
client_type=session_client_type(payload, request.headers),
client_platform=_text_header(request.headers, "x-ctms-client-platform", 32),
client_version=_text_header(request.headers, "x-ctms-client-version", 64),
client_source=_text_header(request.headers, "x-ctms-client-source", 32),
login_ip=resolve_client_ip(request),
login_at=now,
last_seen_at=now,
)
db.add(session)
elif session.user_id != user_id or session.ended_at is not None:
return None
else:
session.last_seen_at = now
session.client_platform = _text_header(request.headers, "x-ctms-client-platform", 32) or session.client_platform
session.client_version = _text_header(request.headers, "x-ctms-client-version", 64) or session.client_version
await db.commit()
await db.refresh(session)
return session
async def end_login_session(
db: AsyncSession,
*,
user_id: uuid.UUID,
payload: dict[str, Any],
reason: str = "logout",
) -> bool:
session_id = session_id_from_payload(payload)
result = await db.execute(
update(UserLoginSession)
.where(
UserLoginSession.id == session_id,
UserLoginSession.user_id == user_id,
UserLoginSession.ended_at.is_(None),
)
.values(ended_at=_now(), end_reason=reason, last_seen_at=_now())
)
await db.commit()
return bool(result.rowcount)
async def get_login_summaries(
db: AsyncSession,
user_ids: Iterable[uuid.UUID],
) -> dict[uuid.UUID, UserLoginSummary]:
ids = list(user_ids)
if not ids:
return {}
rows = (
await db.execute(
select(UserLoginSession)
.where(UserLoginSession.user_id.in_(ids))
.order_by(UserLoginSession.user_id, UserLoginSession.login_at.desc())
)
).scalars().all()
cutoff = _now() - timedelta(seconds=settings.USER_SESSION_ONLINE_SECONDS)
summaries: dict[uuid.UUID, UserLoginSummary] = {}
mutable: dict[uuid.UUID, dict[str, Any]] = {}
for row in rows:
current = mutable.setdefault(
row.user_id,
{
"last_login_at": row.login_at,
"last_seen_at": row.last_seen_at,
"client_type": row.client_type,
"active_sources": set(),
},
)
row_last_seen_at = _as_utc(row.last_seen_at)
if row_last_seen_at and (
current["last_seen_at"] is None or row_last_seen_at > _as_utc(current["last_seen_at"])
):
current["last_seen_at"] = row_last_seen_at
if row.ended_at is None and row_last_seen_at >= cutoff:
source_key = (row.client_type, row.login_ip) if row.login_ip else ("session", row.id)
current["active_sources"].add(source_key)
for user_id, item in mutable.items():
active_session_count = len(item["active_sources"])
summaries[user_id] = UserLoginSummary(
status="ONLINE" if active_session_count else "OFFLINE",
last_login_at=item["last_login_at"],
last_seen_at=item["last_seen_at"],
client_type=item["client_type"],
active_session_count=active_session_count,
)
return summaries
async def list_login_activities(
db: AsyncSession,
*,
user_id: uuid.UUID,
limit: int,
) -> list[UserLoginSession]:
result = await db.execute(
select(UserLoginSession)
.where(UserLoginSession.user_id == user_id)
.order_by(UserLoginSession.login_at.desc())
.limit(limit)
)
return list(result.scalars().all())