"""UploadSession 的 DAO。""" from __future__ import annotations from datetime import datetime, timedelta from sqlalchemy import select from sqlalchemy.orm import Session from ..models.upload_session import UploadSession class UploadSessionDAO: def __init__(self, db: Session) -> None: self.db = db def create(self, session: UploadSession) -> UploadSession: self.db.add(session) self.db.commit() self.db.refresh(session) return session def get_by_upload_id(self, upload_id: str) -> UploadSession | None: stmt = ( select(UploadSession) .where(UploadSession.upload_id == upload_id) .limit(1) ) return self.db.scalars(stmt).first() def mark_uploaded(self, session: UploadSession, chunks: list[int]) -> UploadSession: """更新已上传分片集合(整体替换,避免并发追加丢失)。""" session.uploaded_chunks = sorted(set(chunks)) self.db.commit() self.db.refresh(session) return session def mark_completed(self, session: UploadSession, file_id: int) -> UploadSession: session.file_id = file_id session.status = "completed" self.db.commit() self.db.refresh(session) return session def list_stale(self, ttl_seconds: int) -> list[UploadSession]: """返回 pending 且 updated_at 早于 cutoff 的会话(被放弃的上传)。""" cutoff = datetime.now() - timedelta(seconds=ttl_seconds) stmt = ( select(UploadSession) .where(UploadSession.status == "pending") .where(UploadSession.updated_at < cutoff) ) return list(self.db.scalars(stmt).all()) def delete(self, session: UploadSession) -> None: self.db.delete(session) self.db.commit()