Files
Aether/_deprecated_py_src/services/auth/session_service.py
fawney19 1d9c77522a refactor: 移除 Python 后端源码,全面迁移至 Rust gateway 架构
- 删除全部 Python 源码 (src/) 及 Alembic 迁移脚本,归档至 _deprecated_py_src/
- 重构 Rust gateway ai_pipeline: 拆分 planner/finalize 模块,新增 contracts/adaptation 层
- 重组 handlers 模块为 admin/public/proxy/internal/shared 子模块结构
- 新增 executor 模块,引入 Rust 原生数据库迁移 (aether-data/migrations)
- 简化 CI/Docker 构建流程,移除 base image 二级构建,统一为单一 app image
- 移除 Python 相关基础设施文件 (entrypoint.sh, gunicorn_conf.py, Dockerfile.base)
2026-04-03 16:26:16 +08:00

468 lines
16 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import re
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Mapping
from fastapi import HTTPException, Request, status
from sqlalchemy import and_, or_
from sqlalchemy.orm import Session
from src.core.logger import logger
from src.models.database import AuditEventType, User, UserSession
from src.services.system.audit import AuditService
SESSION_TOUCH_INTERVAL_SECONDS = 300
CLIENT_DEVICE_ID_HEADER = "X-Client-Device-Id"
MAX_SESSIONS_PER_USER = 20
TERMINAL_SESSION_RETENTION_DAYS = 30
_DEVICE_ID_PATTERN = re.compile(r"^[a-zA-Z0-9\-_]{1,128}$")
def _strip_hint(value: str | None) -> str | None:
if value is None:
return None
cleaned = value.strip().strip('"').strip()
return cleaned or None
def _parse_browser(user_agent: str) -> tuple[str | None, str | None]:
patterns = [
("Edge", r"Edg/([\d.]+)"),
("Opera", r"OPR/([\d.]+)"),
("Chrome", r"Chrome/([\d.]+)"),
("Firefox", r"Firefox/([\d.]+)"),
("Safari", r"Version/([\d.]+).*Safari"),
]
for name, pattern in patterns:
match = re.search(pattern, user_agent)
if match:
return name, match.group(1)
return None, None
def _parse_os(
user_agent: str, client_hints: Mapping[str, str | None]
) -> tuple[str | None, str | None]:
platform = _strip_hint(client_hints.get("sec-ch-ua-platform"))
platform_version = _strip_hint(client_hints.get("sec-ch-ua-platform-version"))
if platform:
return platform, platform_version
patterns = [
("Windows", r"Windows NT ([\d.]+)"),
("macOS", r"Mac OS X ([\d_]+)"),
("iOS", r"(?:iPhone OS|CPU OS) ([\d_]+)"),
("Android", r"Android ([\d.]+)"),
("Linux", r"Linux"),
]
for name, pattern in patterns:
match = re.search(pattern, user_agent)
if match:
version = match.group(1).replace("_", ".") if match.lastindex else None
return name, version
return None, None
def _parse_device_type(user_agent: str, client_hints: Mapping[str, str | None]) -> str:
ch_mobile = _strip_hint(client_hints.get("sec-ch-ua-mobile"))
ua_lower = user_agent.lower()
if ch_mobile == "?1":
return "mobile"
if "ipad" in ua_lower or "tablet" in ua_lower:
return "tablet"
if any(marker in ua_lower for marker in ("iphone", "android", "mobile")):
return "mobile"
if any(marker in ua_lower for marker in ("macintosh", "windows", "linux", "x11")):
return "desktop"
return "unknown"
def _build_device_label(
*,
browser_name: str | None,
os_name: str | None,
device_model: str | None,
device_type: str,
) -> str:
if device_model:
return device_model
if browser_name and os_name:
return f"{browser_name} / {os_name}"
if browser_name:
return browser_name
if os_name:
return os_name
if device_type == "mobile":
return "移动设备"
if device_type == "tablet":
return "平板设备"
if device_type == "desktop":
return "桌面设备"
return "未知设备"
@dataclass(frozen=True)
class SessionClientContext:
client_device_id: str
device_label: str
device_type: str
browser_name: str | None
browser_version: str | None
os_name: str | None
os_version: str | None
device_model: str | None
client_hints: dict[str, str | None]
ip_address: str | None
user_agent: str
class SessionService:
"""用户设备会话服务。"""
@staticmethod
def _utcnow() -> datetime:
return datetime.now(timezone.utc)
@staticmethod
def _active_sessions_query(db: Session, *, user_id: str) -> Any:
return db.query(UserSession).filter(
UserSession.user_id == user_id,
UserSession.revoked_at.is_(None),
UserSession.expires_at > SessionService._utcnow(),
)
@staticmethod
def cleanup_user_sessions(db: Session, *, user_id: str) -> int:
"""清理该用户已进入终态且超过保留期的会话记录。"""
cutoff = SessionService._utcnow() - timedelta(days=TERMINAL_SESSION_RETENTION_DAYS)
deleted = (
db.query(UserSession)
.filter(UserSession.user_id == user_id)
.filter(
or_(
UserSession.expires_at < cutoff,
and_(UserSession.revoked_at.is_not(None), UserSession.revoked_at < cutoff),
)
)
.delete(synchronize_session=False)
)
if deleted:
db.flush()
return int(deleted or 0)
@staticmethod
def _normalize_device_id(raw: str) -> str | None:
"""校验并规范化 device id非法值返回 None。"""
cleaned = raw.strip()[:128]
if not cleaned:
return None
if _DEVICE_ID_PATTERN.match(cleaned):
return cleaned
# 不符合格式的视为无效
return None
@staticmethod
def extract_client_device_id(request: Request) -> str:
header_value = request.headers.get(CLIENT_DEVICE_ID_HEADER)
if header_value:
normalized = SessionService._normalize_device_id(header_value)
if normalized:
return normalized
query_value = request.query_params.get("client_device_id")
if query_value:
normalized = SessionService._normalize_device_id(query_value)
if normalized:
return normalized
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="缺少或无效的设备标识",
)
@staticmethod
def build_client_context(
*,
client_device_id: str,
client_ip: str | None,
user_agent: str,
headers: Mapping[str, str],
) -> SessionClientContext:
normalized_headers = {str(key).lower(): value for key, value in headers.items()}
client_hints = {
"sec-ch-ua": normalized_headers.get("sec-ch-ua"),
"sec-ch-ua-platform": normalized_headers.get("sec-ch-ua-platform"),
"sec-ch-ua-platform-version": normalized_headers.get("sec-ch-ua-platform-version"),
"sec-ch-ua-model": normalized_headers.get("sec-ch-ua-model"),
"sec-ch-ua-mobile": normalized_headers.get("sec-ch-ua-mobile"),
}
browser_name, browser_version = _parse_browser(user_agent)
os_name, os_version = _parse_os(user_agent, client_hints)
device_model = _strip_hint(client_hints.get("sec-ch-ua-model"))
device_type = _parse_device_type(user_agent, client_hints)
device_label = _build_device_label(
browser_name=browser_name,
os_name=os_name,
device_model=device_model,
device_type=device_type,
)
return SessionClientContext(
client_device_id=client_device_id,
device_label=device_label,
device_type=device_type,
browser_name=browser_name,
browser_version=browser_version,
os_name=os_name,
os_version=os_version,
device_model=device_model,
client_hints=client_hints,
ip_address=client_ip,
user_agent=user_agent,
)
@staticmethod
def get_active_session(db: Session, session_id: str, user_id: str) -> UserSession | None:
session = (
db.query(UserSession)
.filter(UserSession.id == session_id, UserSession.user_id == user_id)
.first()
)
if not session:
return None
if session.is_revoked or session.is_expired:
return None
return session
@staticmethod
def assert_session_device_matches(session: UserSession, client_device_id: str) -> None:
if session.client_device_id != client_device_id:
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="设备标识与登录会话不匹配",
)
@staticmethod
def create_session(
db: Session,
*,
user: User,
session_id: str,
refresh_token: str,
expires_at: datetime,
client: SessionClientContext,
revoke_existing_same_device: bool = True,
) -> UserSession:
now = SessionService._utcnow()
SessionService.cleanup_user_sessions(db, user_id=user.id)
if revoke_existing_same_device:
existing_sessions = (
SessionService._active_sessions_query(db, user_id=user.id)
.filter(UserSession.client_device_id == client.client_device_id)
.all()
)
for existing in existing_sessions:
existing.revoked_at = now
existing.revoke_reason = "replaced_by_new_login"
existing.updated_at = now
# 如果活跃会话数超出限制,淘汰最旧的会话
active_count = SessionService._active_sessions_query(db, user_id=user.id).count()
if active_count >= MAX_SESSIONS_PER_USER:
oldest_sessions = (
SessionService._active_sessions_query(db, user_id=user.id)
.order_by(UserSession.last_seen_at.asc())
.limit(active_count - MAX_SESSIONS_PER_USER + 1)
.all()
)
for old_session in oldest_sessions:
old_session.revoked_at = now
old_session.revoke_reason = "session_limit_exceeded"
old_session.updated_at = now
session = UserSession(
id=session_id,
user_id=user.id,
client_device_id=client.client_device_id,
device_label=client.device_label,
device_type=client.device_type,
browser_name=client.browser_name,
browser_version=client.browser_version,
os_name=client.os_name,
os_version=client.os_version,
device_model=client.device_model,
ip_address=client.ip_address,
user_agent=client.user_agent[:1000],
client_hints=client.client_hints,
last_seen_at=now,
expires_at=expires_at,
revoked_at=None,
revoke_reason=None,
)
session.set_refresh_token(refresh_token)
db.add(session)
db.flush()
return session
@staticmethod
def rotate_refresh_token(
session: UserSession,
*,
refresh_token: str,
expires_at: datetime,
client_ip: str | None,
user_agent: str,
) -> None:
session.set_refresh_token(refresh_token)
session.expires_at = expires_at
session.last_seen_at = datetime.now(timezone.utc)
if client_ip:
session.ip_address = client_ip
if user_agent:
session.user_agent = user_agent[:1000]
@staticmethod
def touch_session(
session: UserSession,
*,
client_ip: str | None,
user_agent: str,
) -> bool:
now = datetime.now(timezone.utc)
last_seen_at = session.last_seen_at
if last_seen_at.tzinfo is None:
last_seen_at = last_seen_at.replace(tzinfo=timezone.utc)
if (now - last_seen_at).total_seconds() < SESSION_TOUCH_INTERVAL_SECONDS:
return False
session.last_seen_at = now
if client_ip:
session.ip_address = client_ip
if user_agent:
session.user_agent = user_agent[:1000]
return True
@staticmethod
def revoke_session(
db: Session,
*,
session: UserSession,
reason: str,
audit_user_id: str | None = None,
ip_address: str | None = None,
user_agent: str | None = None,
) -> None:
if session.revoked_at is not None:
return
now = datetime.now(timezone.utc)
session.revoked_at = now
session.revoke_reason = reason[:100]
session.updated_at = now
if reason == "refresh_token_reused":
AuditService.log_event(
db=db,
event_type=AuditEventType.SUSPICIOUS_ACTIVITY,
description="Detected refresh token reuse; session revoked",
user_id=audit_user_id or session.user_id,
ip_address=ip_address,
user_agent=user_agent,
metadata={
"session_id": session.id,
"client_device_id": session.client_device_id,
"reason": reason,
},
)
@staticmethod
def revoke_all_user_sessions(
db: Session,
*,
user_id: str,
reason: str,
exclude_session_id: str | None = None,
) -> int:
now = SessionService._utcnow()
SessionService.cleanup_user_sessions(db, user_id=user_id)
sessions = SessionService._active_sessions_query(db, user_id=user_id).all()
count = 0
for session in sessions:
if exclude_session_id and session.id == exclude_session_id:
continue
session.revoked_at = now
session.revoke_reason = reason[:100]
session.updated_at = now
count += 1
return count
@staticmethod
def list_user_sessions(db: Session, *, user_id: str) -> list[UserSession]:
SessionService.cleanup_user_sessions(db, user_id=user_id)
return (
SessionService._active_sessions_query(db, user_id=user_id)
.order_by(UserSession.last_seen_at.desc(), UserSession.created_at.desc())
.all()
)
@staticmethod
def update_session_label(session: UserSession, device_label: str) -> None:
normalized = device_label.strip()
if not normalized:
raise ValueError("设备名称不能为空")
session.device_label = normalized[:120]
session.updated_at = datetime.now(timezone.utc)
@staticmethod
def get_session_for_user(
db: Session,
*,
user_id: str,
session_id: str,
lock_for_update: bool = False,
) -> UserSession | None:
query = db.query(UserSession).filter(
UserSession.user_id == user_id,
UserSession.id == session_id,
)
if lock_for_update:
# 串行化 refresh token 轮换,避免并发刷新把合法会话误判为重放攻击。
query = query.with_for_update()
return query.first()
@staticmethod
def validate_refresh_session(
db: Session,
*,
user_id: str,
session_id: str,
refresh_token: str,
ip_address: str | None = None,
user_agent: str | None = None,
) -> tuple[UserSession | None, bool]:
session = SessionService.get_session_for_user(
db,
user_id=user_id,
session_id=session_id,
lock_for_update=True,
)
if not session or session.is_revoked or session.is_expired:
return None, False
is_valid, is_prev = session.verify_refresh_token(refresh_token)
if not is_valid:
logger.warning("Refresh token mismatch for session {}", session_id)
SessionService.revoke_session(
db,
session=session,
reason="refresh_token_reused",
audit_user_id=user_id,
ip_address=ip_address,
user_agent=user_agent,
)
db.flush()
return None, False
if is_prev:
logger.info("Grace window hit for session {} (prev token used)", session_id)
return session, is_prev