2026-01-19 03:19:17 +08:00
|
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
|
|
import json
|
|
|
|
|
|
import secrets
|
|
|
|
|
|
import time
|
2026-02-01 17:28:00 +08:00
|
|
|
|
from collections.abc import Awaitable
|
2026-01-19 03:19:17 +08:00
|
|
|
|
from dataclasses import dataclass
|
2026-01-30 03:10:21 +08:00
|
|
|
|
from typing import Any, cast
|
2026-01-19 03:19:17 +08:00
|
|
|
|
|
|
|
|
|
|
from redis.asyncio import Redis
|
|
|
|
|
|
|
|
|
|
|
|
OAUTH_STATE_TTL_SECONDS = 600
|
|
|
|
|
|
OAUTH_STATE_KEY_PREFIX = "oauth_state:"
|
|
|
|
|
|
|
2026-01-19 19:34:31 +08:00
|
|
|
|
# OAuth bind token: 用于安全地在浏览器跳转时传递用户身份
|
|
|
|
|
|
# 短期有效(5分钟),一次性使用
|
|
|
|
|
|
OAUTH_BIND_TOKEN_TTL_SECONDS = 300
|
|
|
|
|
|
OAUTH_BIND_TOKEN_KEY_PREFIX = "oauth_bind_token:"
|
|
|
|
|
|
|
2026-01-19 03:19:17 +08:00
|
|
|
|
|
|
|
|
|
|
CONSUME_STATE_SCRIPT = r"""
|
|
|
|
|
|
local value = redis.call("GET", KEYS[1])
|
|
|
|
|
|
if value then
|
|
|
|
|
|
redis.call("DEL", KEYS[1])
|
|
|
|
|
|
end
|
|
|
|
|
|
return value
|
|
|
|
|
|
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
|
|
|
|
class OAuthStateData:
|
|
|
|
|
|
nonce: str
|
|
|
|
|
|
provider_type: str
|
|
|
|
|
|
action: str # "login" | "bind"
|
2026-01-30 03:10:21 +08:00
|
|
|
|
user_id: str | None
|
2026-03-17 16:34:09 +08:00
|
|
|
|
client_device_id: str | None
|
2026-01-19 03:19:17 +08:00
|
|
|
|
created_at: int
|
|
|
|
|
|
|
|
|
|
|
|
@classmethod
|
2026-01-30 03:10:21 +08:00
|
|
|
|
def from_dict(cls, data: dict[str, Any]) -> OAuthStateData:
|
2026-01-19 03:19:17 +08:00
|
|
|
|
return cls(
|
|
|
|
|
|
nonce=str(data.get("nonce") or ""),
|
|
|
|
|
|
provider_type=str(data.get("provider_type") or ""),
|
|
|
|
|
|
action=str(data.get("action") or ""),
|
|
|
|
|
|
user_id=data.get("user_id"),
|
2026-03-17 16:34:09 +08:00
|
|
|
|
client_device_id=data.get("client_device_id"),
|
2026-01-19 03:19:17 +08:00
|
|
|
|
created_at=int(data.get("created_at") or 0),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _state_key(nonce: str) -> str:
|
|
|
|
|
|
return f"{OAUTH_STATE_KEY_PREFIX}{nonce}"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def create_oauth_state(
|
2026-03-17 16:34:09 +08:00
|
|
|
|
redis: Redis,
|
|
|
|
|
|
*,
|
|
|
|
|
|
provider_type: str,
|
|
|
|
|
|
action: str,
|
|
|
|
|
|
user_id: str | None = None,
|
|
|
|
|
|
client_device_id: str | None = None,
|
2026-01-19 03:19:17 +08:00
|
|
|
|
) -> str:
|
|
|
|
|
|
nonce = secrets.token_urlsafe(24)
|
|
|
|
|
|
data = {
|
|
|
|
|
|
"nonce": nonce,
|
|
|
|
|
|
"provider_type": provider_type,
|
|
|
|
|
|
"action": action,
|
|
|
|
|
|
"user_id": user_id,
|
2026-03-17 16:34:09 +08:00
|
|
|
|
"client_device_id": client_device_id,
|
2026-01-19 03:19:17 +08:00
|
|
|
|
"created_at": int(time.time()),
|
|
|
|
|
|
}
|
|
|
|
|
|
await redis.setex(_state_key(nonce), OAUTH_STATE_TTL_SECONDS, json.dumps(data))
|
|
|
|
|
|
return nonce
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-01-30 03:10:21 +08:00
|
|
|
|
async def consume_oauth_state(redis: Redis, nonce: str) -> OAuthStateData | None:
|
2026-01-19 03:19:17 +08:00
|
|
|
|
if not nonce:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
key = _state_key(nonce)
|
|
|
|
|
|
# redis-py 的类型标注在 sync/async 之间会出现 Union;这里明确按 async 处理。
|
2026-01-30 03:10:21 +08:00
|
|
|
|
raw = await cast(Awaitable[str | None], redis.eval(CONSUME_STATE_SCRIPT, 1, key))
|
2026-01-19 03:19:17 +08:00
|
|
|
|
if not raw:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
|
parsed = json.loads(raw)
|
|
|
|
|
|
except json.JSONDecodeError:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
return OAuthStateData.from_dict(parsed)
|
2026-01-19 19:34:31 +08:00
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
|
|
|
|
class OAuthBindTokenData:
|
|
|
|
|
|
"""OAuth 绑定临时令牌数据,用于浏览器跳转场景的安全认证"""
|
|
|
|
|
|
|
|
|
|
|
|
token: str
|
|
|
|
|
|
user_id: str
|
|
|
|
|
|
provider_type: str
|
|
|
|
|
|
created_at: int
|
|
|
|
|
|
|
|
|
|
|
|
@classmethod
|
2026-01-30 03:10:21 +08:00
|
|
|
|
def from_dict(cls, data: dict[str, Any]) -> OAuthBindTokenData:
|
2026-01-19 19:34:31 +08:00
|
|
|
|
return cls(
|
|
|
|
|
|
token=str(data.get("token") or ""),
|
|
|
|
|
|
user_id=str(data.get("user_id") or ""),
|
|
|
|
|
|
provider_type=str(data.get("provider_type") or ""),
|
|
|
|
|
|
created_at=int(data.get("created_at") or 0),
|
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def _bind_token_key(token: str) -> str:
|
|
|
|
|
|
return f"{OAUTH_BIND_TOKEN_KEY_PREFIX}{token}"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
async def create_oauth_bind_token(redis: Redis, *, user_id: str, provider_type: str) -> str:
|
|
|
|
|
|
"""创建一次性 OAuth 绑定令牌,用于浏览器跳转场景"""
|
|
|
|
|
|
token = secrets.token_urlsafe(32)
|
|
|
|
|
|
data = {
|
|
|
|
|
|
"token": token,
|
|
|
|
|
|
"user_id": user_id,
|
|
|
|
|
|
"provider_type": provider_type,
|
|
|
|
|
|
"created_at": int(time.time()),
|
|
|
|
|
|
}
|
|
|
|
|
|
await redis.setex(_bind_token_key(token), OAUTH_BIND_TOKEN_TTL_SECONDS, json.dumps(data))
|
|
|
|
|
|
return token
|
|
|
|
|
|
|
|
|
|
|
|
|
2026-01-30 03:10:21 +08:00
|
|
|
|
async def consume_oauth_bind_token(redis: Redis, token: str) -> OAuthBindTokenData | None:
|
2026-01-19 19:34:31 +08:00
|
|
|
|
"""消费(验证并删除)OAuth 绑定令牌,返回令牌数据或 None"""
|
|
|
|
|
|
if not token:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
key = _bind_token_key(token)
|
2026-01-30 03:10:21 +08:00
|
|
|
|
raw = await cast(Awaitable[str | None], redis.eval(CONSUME_STATE_SCRIPT, 1, key))
|
2026-01-19 19:34:31 +08:00
|
|
|
|
if not raw:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
|
parsed = json.loads(raw)
|
|
|
|
|
|
except json.JSONDecodeError:
|
|
|
|
|
|
return None
|
|
|
|
|
|
|
|
|
|
|
|
return OAuthBindTokenData.from_dict(parsed)
|