Files
yaoyaoai/app/services/auth.py
fengchuanhn@gmail.com 9526c0cb1b 11
2026-05-22 18:13:07 +08:00

200 lines
5.5 KiB
Python

from datetime import datetime, timezone
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import get_settings
from app.core.roles import ROLE_AGENT, ROLE_OEM
from app.core.heartbeat_verify import compute_server_verify_code
from app.core.security import (
create_access_token,
hash_password,
session_redis_key,
verify_password,
)
from app.models.user import User
from app.schemas.auth import HeartbeatResponse, TokenResponse, UserPublic
settings = get_settings()
class AuthError(Exception):
def __init__(self, message: str) -> None:
self.message = message
super().__init__(message)
def _normalize_device_serial(value: str | None) -> str:
if not value:
return ""
return value.strip()[:255]
def parse_id_header(value: str | None, *, field_name: str) -> int | None:
if value is None or not str(value).strip():
return None
try:
parsed = int(str(value).strip())
except ValueError as exc:
raise AuthError(f"{field_name} 无效") from exc
if parsed < 1:
raise AuthError(f"{field_name} 无效")
return parsed
async def _validate_register_owner_ids(
db: AsyncSession,
*,
oem_id: int | None,
agent_id: int | None,
) -> tuple[int | None, int | None]:
resolved_oem_id = oem_id
if agent_id is not None:
agent = await db.get(User, agent_id)
if agent is None or agent.role_id != ROLE_AGENT:
raise AuthError("代理无效")
if resolved_oem_id is None:
resolved_oem_id = agent.oem_id
elif agent.oem_id != resolved_oem_id:
raise AuthError("代理与 OEM 不匹配")
if resolved_oem_id is not None:
oem = await db.get(User, resolved_oem_id)
if oem is None or oem.role_id != ROLE_OEM:
raise AuthError("OEM 无效")
return resolved_oem_id, agent_id
def _vip_is_active(vip_end_time: datetime | None) -> bool:
if vip_end_time is None:
return False
now = datetime.now(timezone.utc)
end = vip_end_time
if end.tzinfo is None:
end = end.replace(tzinfo=timezone.utc)
return end > now
async def _sync_device_serial(
db: AsyncSession,
user: User,
incoming: str,
*,
mismatch_message: str,
) -> None:
stored = user.device_serial or ""
if not stored:
if incoming:
user.device_serial = incoming
await db.flush()
elif incoming != stored:
raise AuthError(mismatch_message)
async def register_user(
db: AsyncSession,
*,
username: str,
password: str,
device_serial: str | None = None,
oem_id: int | None = None,
agent_id: int | None = None,
) -> TokenResponse:
existing = await db.scalar(select(User).where(User.username == username))
if existing is not None:
raise AuthError("用户名已存在")
serial = _normalize_device_serial(device_serial)
if serial:
bound = await db.scalar(select(User).where(User.device_serial == serial))
if bound is not None:
raise AuthError("该电脑已绑定其它用户")
resolved_oem_id, resolved_agent_id = await _validate_register_owner_ids(
db, oem_id=oem_id, agent_id=agent_id
)
user = User(
username=username,
password_hash=hash_password(password),
role_id=0,
phone=None,
device_serial=serial,
oem_id=resolved_oem_id,
agent_id=resolved_agent_id,
)
db.add(user)
await db.flush()
await db.refresh(user)
token = create_access_token(str(user.id))
return TokenResponse(
access_token=token,
user=UserPublic.model_validate(user),
)
async def login_user(
db: AsyncSession,
*,
username: str,
password: str,
device_serial: str | None = None,
) -> TokenResponse:
user = await db.scalar(select(User).where(User.username == username))
if user is None or not verify_password(password, user.password_hash):
raise AuthError("用户名或密码错误")
incoming = _normalize_device_serial(device_serial)
await _sync_device_serial(
db, user, incoming, mismatch_message="设备不匹配,无法登录"
)
token = create_access_token(str(user.id))
return TokenResponse(
access_token=token,
user=UserPublic.model_validate(user),
)
async def process_heartbeat(
db: AsyncSession,
user: User,
device_serial: str | None,
verify_code: str,
) -> HeartbeatResponse:
incoming = _normalize_device_serial(device_serial)
await _sync_device_serial(db, user, incoming, mismatch_message="设备不匹配")
if not _vip_is_active(user.vip_end_time):
raise AuthError("会员已过期或未开通")
code = verify_code.strip()
if not code:
raise AuthError("verify_code 无效")
return HeartbeatResponse(
vip_end_time=user.vip_end_time,
vip_active=True,
server_verify_code=compute_server_verify_code(code),
)
async def store_session(redis_client, token: str, user_id: int) -> None:
ttl = settings.access_token_expire_minutes * 60
await redis_client.setex(session_redis_key(token), ttl, str(user_id))
async def revoke_session(redis_client, token: str) -> None:
await redis_client.delete(session_redis_key(token))
async def session_user_id(redis_client, token: str) -> int | None:
raw = await redis_client.get(session_redis_key(token))
if raw is None:
return None
return int(raw)
async def get_user_by_id(db: AsyncSession, user_id: int) -> User | None:
return await db.get(User, user_id)