"""TunnelSession 的 DAO。""" from __future__ import annotations from datetime import datetime from sqlalchemy import select, update from sqlalchemy.orm import Session from ..models.tunnel_session import TunnelSession class TunnelSessionDAO: def __init__(self, db: Session) -> None: self.db = db def create(self, session: TunnelSession) -> TunnelSession: self.db.add(session) self.db.commit() self.db.refresh(session) return session def get_active_by_user(self, user_name: str) -> TunnelSession | None: """返回该 user 当前活跃的隧道会话(至多一条)。""" stmt = ( select(TunnelSession) .where(TunnelSession.user_name == user_name) .where(TunnelSession.status == "active") .order_by(TunnelSession.started_at.desc()) .limit(1) ) return self.db.scalars(stmt).first() def list_active(self) -> list[TunnelSession]: stmt = select(TunnelSession).where(TunnelSession.status == "active") return list(self.db.scalars(stmt).all()) def close(self, session: TunnelSession) -> None: """标记会话结束。""" session.ended_at = datetime.now() session.status = "closed" self.db.commit() def close_active_by_user(self, user_name: str) -> int: """关闭该 user 所有 active 会话(断开清理用),返回关闭条数。""" stmt = ( update(TunnelSession) .where(TunnelSession.user_name == user_name) .where(TunnelSession.status == "active") .values(status="closed", ended_at=datetime.now()) ) result = self.db.execute(stmt) self.db.commit() return result.rowcount or 0