Files
taiji-AI-PAD/services/mcp-server/app/device_auth.py
T
chenchenandClaude Opus 4.8 b754af8ba1 fix(mcp-server): device-code 一次性作废失效修复(consume 集群 Redis 删除未生效)
初版 consume_device_code 用裸 redis_client.delete() 且吞异常,在集群 Azure
Redis 上删除未生效,导致一个 device_code 换发 token 后仍能在每次 >interval
的轮询继续换发新 token —— 违反 RFC 8628 一次性语义与验收「换一次后再用→拒绝」。

初测二次轮询都在 slow_down 窗口内(<5s)被限流响应遮住,未暴露;>5s 公网
真实轮询复测才暴露。

修复:consume 改用已验证可靠的 _set_keepttl 置 status=consumed(token 端点
签发前硬检查 consumed → expired_token),并 best-effort 删除 device_code +
device_user_code 两个 key。即使集群删除失败,状态位硬拦截。

复测(公网 APIM 真实路径,间隔 >5s):首 poll 签发 → 二/三次 poll 均
expired_token,不再重复签发。

镜像 device-code-fix2-20260722-arm64 @sha256:716c2e2d 已部署生产 3/3 Running。

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-22 15:24:27 +08:00

187 lines
7.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Heicode headless 设备登录(device-code / RFC 8628)—— Redis 状态机。
对接契约见 `Docs/Heicode-headless设备登录-device-code-给mcp-server的对接需求.md`。
设计(与 magic_link 一致的一次性 Redis 模式):
- authorize 生成 device_code(高熵,设备侧保密)+ user_code(短、去混淆字符,展示给用户)。
- 两个 key(同 TTL,默认 600s):
device_code:{device_code} → {status, user_code, client, user_id?, channel_id?, last_poll}
device_user_code:{user_code} → device_code (approve 反查)
- 状态:pending → approved / denied;token 换发一次后置 consumed(再用即拒)。
- 更新状态用 SET ... KEEPTTL,保留原到期时间(不因批准而续命)。
"""
from __future__ import annotations
import json
import os
import secrets
import time
from typing import Optional, Tuple
import structlog
from app.state import get_state
logger = structlog.get_logger(__name__)
# ===== 常量(契约 §2/§4)=====
DEVICE_CODE_TTL_SECONDS = 600 # device_code / user_code 有效期
DEVICE_POLL_INTERVAL_SECONDS = 5 # 轮询最小间隔
DEVICE_VERIFICATION_URI = os.getenv("DEVICE_VERIFICATION_URI", "https://code.heicode.cc/device")
_DC_PREFIX = "device_code:"
_UC_PREFIX = "device_user_code:"
_MAX_REDIS_RETRIES = 3
# user_code 字符集:去掉易混字符 0/O/1/I(契约 §4)
_UC_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
def generate_device_code() -> str:
"""高熵、不可猜;仅设备侧持有。"""
return secrets.token_urlsafe(32)
def generate_user_code() -> str:
"""展示给用户的短码,形如 WDJB-MJHT。"""
s = "".join(secrets.choice(_UC_ALPHABET) for _ in range(8))
return f"{s[:4]}-{s[4:]}"
async def _get(key: str) -> Optional[str]:
state = get_state()
if not state.redis_client:
logger.warning("device_auth_redis_unavailable", op="get")
return None
for attempt in range(_MAX_REDIS_RETRIES):
try:
return await state.redis_client.get(key)
except Exception as e: # noqa: BLE001
if "MOVED" in str(e) and attempt < _MAX_REDIS_RETRIES - 1:
import asyncio
await asyncio.sleep(0.1)
continue
raise
return None
async def _setex(key: str, ttl: int, value: str) -> bool:
state = get_state()
if not state.redis_client:
logger.warning("device_auth_redis_unavailable", op="setex")
return False
try:
await state.redis_client.setex(key, ttl, value)
return True
except Exception as e: # noqa: BLE001
logger.error("device_auth_redis_setex_failed", key_prefix=key[:24], error=str(e))
return False
async def _set_keepttl(key: str, value: str) -> bool:
"""更新值、保留原 TTL(状态流转不续命)。"""
state = get_state()
if not state.redis_client:
return False
try:
await state.redis_client.set(key, value, keepttl=True)
return True
except TypeError:
# 兼容不支持 keepttl 的客户端:读剩余 TTL 后 setex
ttl = await state.redis_client.ttl(key)
await state.redis_client.setex(key, ttl if ttl and ttl > 0 else DEVICE_CODE_TTL_SECONDS, value)
return True
except Exception as e: # noqa: BLE001
logger.error("device_auth_redis_set_failed", key_prefix=key[:24], error=str(e))
return False
async def create_device_authorization(client: Optional[str]) -> Tuple[Optional[str], Optional[str]]:
"""生成并存储 device_code + user_code。返回 (device_code, user_code);失败返回 (None, None)。"""
device_code = generate_device_code()
user_code = generate_user_code()
payload = json.dumps({
"status": "pending",
"user_code": user_code,
"client": (client or "")[:64],
"created": time.time(),
"last_poll": 0.0,
})
ok1 = await _setex(_DC_PREFIX + device_code, DEVICE_CODE_TTL_SECONDS, payload)
ok2 = await _setex(_UC_PREFIX + user_code, DEVICE_CODE_TTL_SECONDS, device_code)
if not (ok1 and ok2):
return None, None
return device_code, user_code
async def get_device_state(device_code: str) -> Optional[dict]:
"""按 device_code 取状态(不存在=过期/无效 → None)。"""
if not device_code:
return None
raw = await _get(_DC_PREFIX + device_code)
if not raw:
return None
try:
return json.loads(raw)
except (ValueError, TypeError):
return None
async def lookup_by_user_code(user_code: str) -> Tuple[Optional[str], Optional[dict]]:
"""approve 反查:user_code → (device_code, state)。"""
if not user_code:
return None, None
device_code = await _get(_UC_PREFIX + user_code)
if not device_code:
return None, None
if isinstance(device_code, bytes):
device_code = device_code.decode()
state = await get_device_state(device_code)
return device_code, state
async def set_device_decision(device_code: str, state: dict, approved: bool,
user_id: Optional[str] = None, channel_id: Optional[str] = None) -> bool:
"""批准/拒绝:写状态 + 绑定批准人身份(仅来自批准人 token,不取设备侧输入)。保留原 TTL。"""
state = dict(state)
state["status"] = "approved" if approved else "denied"
if approved:
state["user_id"] = user_id
state["channel_id"] = channel_id
return await _set_keepttl(_DC_PREFIX + device_code, json.dumps(state))
async def mark_poll(device_code: str, state: dict) -> bool:
"""记录本次轮询时刻(用于 slow_down 判定)。保留原 TTL。"""
state = dict(state)
state["last_poll"] = time.time()
return await _set_keepttl(_DC_PREFIX + device_code, json.dumps(state))
async def consume_device_code(device_code: str, state: Optional[dict] = None) -> None:
"""token 换发后作废该 device_code(一次性)。
先用**已验证可靠**的 _set_keepttl 把状态置 consumed(token 端点在签发前会
检查 status==consumed → expired_token,这是硬拦截),再 best-effort 删除 key。
裸 delete 在集群 Redis 上可能抛 MOVED 被吞掉,故不能只依赖删除。
"""
# 1) 置 consumed(权威拦截;保留原 TTL,不续命)
if state is not None:
st = dict(state)
st["status"] = "consumed"
ok = await _set_keepttl(_DC_PREFIX + device_code, json.dumps(st))
if not ok:
logger.error("device_auth_consume_mark_failed", key_prefix=(_DC_PREFIX + device_code)[:24])
# 2) best-effort 物理删除(清 Redis,同时清 user_code 反查)
rs = get_state()
if not rs.redis_client:
return
try:
uc = (state or {}).get("user_code") if state else None
await rs.redis_client.delete(_DC_PREFIX + device_code)
if uc:
await rs.redis_client.delete(_UC_PREFIX + uc)
except Exception as e: # noqa: BLE001
logger.warning("device_auth_consume_delete_failed", error=str(e))