From 9526c0cb1b796fd52300fa254f5a97967c8c55fd Mon Sep 17 00:00:00 2001 From: "fengchuanhn@gmail.com" Date: Fri, 22 May 2026 18:13:07 +0800 Subject: [PATCH] 11 --- app/api/v1/auth.py | 6 ++++++ app/services/auth.py | 43 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 49 insertions(+) diff --git a/app/api/v1/auth.py b/app/api/v1/auth.py index eb54fb1..e2cd3a7 100644 --- a/app/api/v1/auth.py +++ b/app/api/v1/auth.py @@ -22,13 +22,19 @@ async def register( db: DbSession, redis: RedisClient, device_serial: str | None = Header(default=None, convert_underscores=False), + oem_id: str | None = Header(default=None, convert_underscores=False), + agent_id: str | None = Header(default=None, convert_underscores=False), ) -> ApiResponse[TokenResponse]: try: + parsed_oem_id = auth_service.parse_id_header(oem_id, field_name="oem_id") + parsed_agent_id = auth_service.parse_id_header(agent_id, field_name="agent_id") result = await auth_service.register_user( db, username=body.username, password=body.password, device_serial=device_serial, + oem_id=parsed_oem_id, + agent_id=parsed_agent_id, ) except auth_service.AuthError as exc: return ApiResponse(ok=False, message=exc.message) diff --git a/app/services/auth.py b/app/services/auth.py index 101e66b..7c03397 100644 --- a/app/services/auth.py +++ b/app/services/auth.py @@ -4,6 +4,7 @@ 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, @@ -29,6 +30,40 @@ def _normalize_device_serial(value: str | None) -> str: 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 @@ -61,6 +96,8 @@ async def register_user( 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: @@ -72,12 +109,18 @@ async def register_user( 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()