feat(desktop): 支持桌面端三十天免登录

This commit is contained in:
Cheng Zhou
2026-07-02 09:06:24 +08:00
parent 8c8327df92
commit b8c5c4123a
9 changed files with 334 additions and 19 deletions
+54 -11
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
@@ -70,12 +71,50 @@ 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(request.headers.get("x-ctms-client-type"))
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)
def issue_user_token(db_user, request: Request) -> Token:
session_start = datetime.now(timezone.utc)
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,
)
return Token(access_token=access_token, token_type="bearer")
@@ -240,23 +279,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 issue_user_token(db_user, request)
@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 issue_user_token(db_user, request)
@router.get("/me", response_model=UserRead)
@@ -285,19 +324,23 @@ 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")),
)
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)
+1
View File
@@ -19,6 +19,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"
+14 -1
View File
@@ -19,16 +19,29 @@ 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,
) -> 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
return jwt.encode(to_encode, settings.JWT_SECRET_KEY, algorithm=ALGORITHM)
+69 -1
View File
@@ -16,7 +16,7 @@ import os
from app.main import create_app
from app.core.config import settings
from app.core.deps import get_db_session
from app.core.security import hash_password, verify_password
from app.core.security import create_access_token, decode_token_allow_expired, hash_password, verify_password
from app.crud import user as user_crud
from app.db.base_class import Base
from app.models.audit_log import AuditLog
@@ -422,6 +422,74 @@ async def test_registered_user_can_login_after_email_verification(client_and_db)
assert resp.json()["access_token"]
@pytest.mark.asyncio
async def test_desktop_login_uses_30_day_token_without_changing_web_login(client_and_db):
client, _ = client_and_db
web_resp = await encrypted_login(client, "admin@test.com", "admin123")
desktop_resp = await client.post(
"/api/v1/auth/login",
json=await encrypted_auth_payload(client, "admin@test.com", "admin123"),
headers={"X-CTMS-Client-Type": "desktop"},
)
assert web_resp.status_code == 200
assert desktop_resp.status_code == 200
web_payload = decode_token_allow_expired(web_resp.json()["access_token"])
desktop_payload = decode_token_allow_expired(desktop_resp.json()["access_token"])
assert web_payload["client_type"] == "web"
assert desktop_payload["client_type"] == "desktop"
assert web_payload["exp"] - web_payload["iat"] == settings.JWT_EXPIRE_MINUTES * 60
assert desktop_payload["exp"] - desktop_payload["iat"] == settings.DESKTOP_SESSION_MAX_DAYS * 24 * 3600
@pytest.mark.asyncio
async def test_web_token_extension_cannot_be_upgraded_with_desktop_header(client_and_db):
client, _ = client_and_db
web_resp = await encrypted_login(client, "admin@test.com", "admin123")
token = web_resp.json()["access_token"]
resp = await client.post(
"/api/v1/auth/extend",
headers={
"Authorization": f"Bearer {token}",
"X-CTMS-Client-Type": "desktop",
},
)
assert resp.status_code == 200
payload = decode_token_allow_expired(resp.json()["accessToken"])
assert payload["client_type"] == "web"
assert payload["exp"] - payload["iat"] == settings.JWT_EXPIRE_MINUTES * 60
@pytest.mark.asyncio
async def test_desktop_token_extension_rejects_sessions_after_30_days(client_and_db):
client, SessionLocal = client_and_db
async with SessionLocal() as session:
admin = await user_crud.get_by_email(session, "admin@test.com")
session_start = datetime.now(timezone.utc) - timedelta(days=settings.DESKTOP_SESSION_MAX_DAYS, seconds=1)
token = create_access_token(
user_id=str(admin.id),
expires_minutes=settings.DESKTOP_SESSION_MAX_DAYS * 24 * 60,
session_start=session_start,
max_age_seconds=settings.DESKTOP_SESSION_MAX_DAYS * 24 * 3600,
client_type="desktop",
)
resp = await client.post(
"/api/v1/auth/extend",
headers={
"Authorization": f"Bearer {token}",
"X-CTMS-Client-Type": "desktop",
},
)
assert resp.status_code == 401
assert "会话已到期" in resp.json().get("detail", "")
@pytest.mark.asyncio
async def test_admin_created_user_is_active_by_default(client_and_db):
client, SessionLocal = client_and_db