11
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user