From 96e68e1e89f4e299700794fd800dac829e605284 Mon Sep 17 00:00:00 2001 From: fxbin Date: Sun, 27 Sep 2026 23:27:21 +0800 Subject: [PATCH] =?UTF-8?q?feat(auth):=20=E9=9C=80=E8=A6=81=20step-up=20?= =?UTF-8?q?=E7=9A=84=E7=AC=AC=E4=B8=89=E6=96=B9=E8=B4=A6=E5=8F=B7=E6=98=BE?= =?UTF-8?q?=E5=BC=8F=E7=BB=91=E5=AE=9A=E6=B5=81=E7=A8=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/app/api/v1/oauth.py | 119 ++++++++- backend/app/services/auth_service.py | 54 ++++- backend/tests_oauth_patch/test_oauth_bind.py | 240 +++++++++++++++++++ 3 files changed, 409 insertions(+), 4 deletions(-) create mode 100644 backend/tests_oauth_patch/test_oauth_bind.py diff --git a/backend/app/api/v1/oauth.py b/backend/app/api/v1/oauth.py index 5546d9cc..be6be8cc 100644 --- a/backend/app/api/v1/oauth.py +++ b/backend/app/api/v1/oauth.py @@ -11,22 +11,27 @@ from __future__ import annotations import logging +import time from typing import Any from urllib.parse import urlencode, urlsplit from fastapi import APIRouter, Depends, HTTPException, Request, status from fastapi.responses import RedirectResponse +from pydantic import BaseModel from sqlalchemy.ext.asyncio import AsyncSession -from app.api.v1.auth import _set_auth_cookies +from app.api.v1.auth import _set_auth_cookies, get_current_user, get_optional_current_user from app.core.config import settings from app.core.database import get_db from app.core.oauth import ENABLED_PROVIDERS, oauth from app.core.request_utils import client_ip from app.services.auth_service import ( OAuthAccountConflictError, + OAuthBindConflictError, create_session, get_or_create_oauth_user, + link_oauth_identity_to_user, + verify_password, ) router = APIRouter(prefix="/auth/oauth", tags=["auth"]) @@ -97,6 +102,67 @@ async def oauth_login(request: Request, provider: str): return await client.authorize_redirect(request, redirect_uri) +# ── 显式绑定流程(#64:需要 step-up 的管理员/已登录用户第三方绑定)──── + + +BIND_INTENT_TTL_SECONDS = 600 + + +class OAuthBindStartRequest(BaseModel): + """发起绑定的 step-up 请求体:重新输入当前账号密码。""" + + password: str + + +def _pop_bind_intent(request: Request, provider: str) -> dict | None: + """取出并校验绑定意图:provider 匹配且未过期才有效。""" + try: + raw = request.session.pop("oauth_bind", None) + except Exception: + return None + if not isinstance(raw, dict): + return None + if raw.get("provider") != provider or float(raw.get("exp", 0)) < time.time(): + return None + return raw + + +@router.post("/{provider}/bind/start") +async def oauth_bind_start( + provider: str, + payload: OAuthBindStartRequest, + request: Request, + db: AsyncSession = Depends(get_db), + current_user=Depends(get_current_user), +): + """发起第三方账号绑定(#64):step-up 密码验证通过后进入 provider 授权。 + + 绑定意图写入服务端 session(与 OAuth state 同生命周期,TTL 10 分钟), + 回调时据此走绑定分支而非登录分支。 + """ + if not _is_provider_enabled(provider): + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"OAuth provider '{provider}' not enabled") + client = oauth.create_client(provider) + if client is None: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"OAuth provider '{provider}' not registered") + + # step-up:重新验证当前账号密码(未设密码的账号无法走此流程) + if not current_user.password_hash or not verify_password(payload.password, current_user.password_hash): + logger.warning( + "OAuth bind start rejected: step-up password mismatch user_id=%s", getattr(current_user, "id", None) + ) + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="密码验证失败,无法发起绑定") + + request.session["oauth_bind"] = { + "user_id": current_user.id, + "provider": provider, + "exp": time.time() + BIND_INTENT_TTL_SECONDS, + } + redirect_uri = _backend_callback_url(request, provider) + resp = await client.authorize_redirect(request, redirect_uri) + return {"authorize_url": resp.headers["location"]} + + # ── provider 回调 ───────────────────────────────────────────────── @@ -128,6 +194,19 @@ async def oauth_callback(request: Request, provider: str, db: AsyncSession = Dep if not email or not provider_user_id: return _frontend_redirect(error="OAuth 身份信息不完整(缺少邮箱或用户 ID)") + bind_intent = _pop_bind_intent(request, provider) + if bind_intent is not None: + return await _handle_bind_callback( + request, + db, + provider=provider, + intent=bind_intent, + provider_user_id=str(provider_user_id), + email=email, + email_verified=email_verified, + display_name=display_name, + ) + try: user = await get_or_create_oauth_user( db, @@ -158,6 +237,44 @@ async def oauth_callback(request: Request, provider: str, db: AsyncSession = Dep return response +async def _handle_bind_callback( + request: Request, + db: AsyncSession, + *, + provider: str, + intent: dict, + provider_user_id: str, + email: str, + email_verified: bool, + display_name: str | None, +) -> RedirectResponse: + """绑定分支回调:会话本人 + 已验证邮箱 + 无冲突 → 落 user_oauth_account。""" + current = await get_optional_current_user(request, authorization=None, db=db) + if current is None or current.id != intent.get("user_id"): + logger.warning( + "OAuth bind rejected: session mismatch provider=%s intent_user=%s", provider, intent.get("user_id") + ) + return _frontend_redirect(error="绑定会话已失效或与当前登录不一致,请重新发起绑定") + if not email_verified: + return _frontend_redirect(error="该第三方身份邮箱未验证,无法绑定") + try: + account = await link_oauth_identity_to_user( + db, + user=current, + provider=provider, + provider_user_id=provider_user_id, + email=email, + email_verified=email_verified, + display_name=display_name, + ) + except (OAuthBindConflictError, ValueError) as exc: + logger.warning("OAuth bind failed: provider=%s user_id=%s exc=%s", provider, current.id, exc) + return _frontend_redirect(error=str(exc)) + logger.info("OAuth identity bound: provider=%s user_id=%d account_id=%d", provider, current.id, account.id) + fragment = urlencode({"provider": provider, "bind": "linked"}) + return _frontend_redirect(fragment=fragment) + + # ── 已启用 provider 列表(前端据此渲染按钮)────────────────────────── diff --git a/backend/app/services/auth_service.py b/backend/app/services/auth_service.py index c9e3e87f..9c083a37 100644 --- a/backend/app/services/auth_service.py +++ b/backend/app/services/auth_service.py @@ -106,6 +106,10 @@ async def create_user( return user +class OAuthBindConflictError(Exception): + """显式绑定流程冲突:同 provider 已绑定 / 身份已被其它账号绑定。""" + + class OAuthAccountConflictError(Exception): """OAuth 登录时邮箱与现有账号冲突但无法自动合并。 @@ -124,6 +128,52 @@ async def get_oauth_account(db: AsyncSession, *, provider: str, provider_user_id return result.scalar_one_or_none() +async def link_oauth_identity_to_user( + db: AsyncSession, + *, + user: User, + provider: str, + provider_user_id: str, + email: str, + email_verified: bool, + display_name: str | None = None, +) -> UserOAuthAccount: + """把已验证的 OAuth 身份显式挂到指定账号(#64,管理员同样适用)。 + + 与 get_or_create_oauth_user 的自动合并互补:调用方(API 层)必须已完成 + step-up 验证并确认会话本人;本函数只负责冲突判定与落库,两种冲突一律 + 抛 OAuthBindConflictError(目标账号已有同 provider 绑定需先解绑、身份 + 已被其它账号绑定拒绝),防止置换攻击。 + """ + if not email_verified: + raise ValueError("该第三方身份邮箱未验证,无法绑定") + + own_stmt = select(UserOAuthAccount).where( + UserOAuthAccount.user_id == user.id, + UserOAuthAccount.provider == provider, + ) + own = (await db.execute(own_stmt)).scalar_one_or_none() + if own is not None: + raise OAuthBindConflictError("该账号已绑定此登录方式,请先解绑后再重新绑定") + + taken = await get_oauth_account(db, provider=provider, provider_user_id=provider_user_id) + if taken is not None: + raise OAuthBindConflictError("该第三方身份已绑定其它账号,无法重复绑定") + + account = UserOAuthAccount( + user_id=user.id, + provider=provider, + provider_user_id=provider_user_id, + provider_email=email, + email_verified=True, + display_name=display_name, + ) + db.add(account) + await db.commit() + await db.refresh(account) + return account + + async def get_or_create_oauth_user( db: AsyncSession, *, @@ -263,9 +313,7 @@ async def ensure_admin_user( # promoted to admin merely by matching ADMIN_EMAIL. if not user.password_hash: raise ValueError("Admin seed refuses to promote a passwordless existing account") - oauth_link = await db.scalar( - select(UserOAuthAccount.id).where(UserOAuthAccount.user_id == user.id).limit(1) - ) + oauth_link = await db.scalar(select(UserOAuthAccount.id).where(UserOAuthAccount.user_id == user.id).limit(1)) if oauth_link is not None: raise ValueError("Admin seed refuses to promote an OAuth-linked existing account") if not user: diff --git a/backend/tests_oauth_patch/test_oauth_bind.py b/backend/tests_oauth_patch/test_oauth_bind.py new file mode 100644 index 00000000..f98ada97 --- /dev/null +++ b/backend/tests_oauth_patch/test_oauth_bind.py @@ -0,0 +1,240 @@ +"""#64 显式绑定流程回归(隔离 mock,不连任何数据库)。 + +覆盖 issue #64 验收四场景 + 会话不匹配: +- 正常绑定:step-up 密码 → bind/start 下发授权 URL → 回调落绑定、不建登录会话; +- 未 step-up 拒绝:密码错误 401,不产生绑定意图; +- 已有绑定冲突:link 抛 OAuthBindConflictError → 错误回跳; +- 未验证邮箱拒绝:与 #62 语义一致,在 userinfo 阶段即拒绝; +- 发起人与当前会话不一致:拒绝绑定。 +""" + +from __future__ import annotations + +import os +from datetime import UTC, datetime, timedelta +from types import SimpleNamespace +from unittest.mock import AsyncMock +from urllib.parse import parse_qs, unquote, urlsplit + +os.environ.setdefault("DATABASE_URL", "postgresql+asyncpg://dummy:dummy@localhost:5432/unused") + +import httpx # noqa: E402 +import pytest # noqa: E402 +from fastapi import FastAPI # noqa: E402 +from starlette.middleware.sessions import SessionMiddleware # noqa: E402 + +from app.api.v1 import oauth as oauth_routes # noqa: E402 +from app.services.auth_service import OAuthBindConflictError # noqa: E402 + +PUBLIC_ORIGIN = "https://topic.example.com" + + +class FakeResponse: + def __init__(self, data): + self.data = data + + def raise_for_status(self): + return None + + def json(self): + return self.data + + +class FakeGithub: + def __init__(self, emails=None): + self.callback_uri = None + self.emails = ( + emails if emails is not None else [{"email": "private@example.com", "primary": True, "verified": True}] + ) + + async def get(self, endpoint, *, token): + if endpoint == "user": + return FakeResponse({"id": 42, "name": "Example", "login": "example", "email": None}) + if endpoint == "user/emails": + return FakeResponse(self.emails) + raise AssertionError(endpoint) + + async def authorize_redirect(self, request, redirect_uri): + self.callback_uri = redirect_uri + request.session["fake_state"] = "test-state" + from fastapi.responses import RedirectResponse + + return RedirectResponse("https://github.com/login/oauth/authorize?state=test-state") + + async def authorize_access_token(self, request): + assert request.session.get("fake_state") == "test-state" + return {"access_token": "dummy", "token_type": "bearer"} + + +def make_bind_app(monkeypatch, provider_client, *, user, link=None, verify_ok=True, callback_user=None): + """组装带依赖覆盖的 OAuth app(绑定流程专用)。""" + monkeypatch.setattr(oauth_routes.settings, "SITE_BASE_URL", PUBLIC_ORIGIN) + monkeypatch.setattr(oauth_routes.settings, "OAUTH_FRONTEND_REDIRECT_URL", PUBLIC_ORIGIN + "/oauth/callback") + monkeypatch.setattr(oauth_routes.settings, "APP_ENV", "production") + monkeypatch.setattr(oauth_routes.settings, "AUTH_COOKIE_SECURE", True) + monkeypatch.setattr(oauth_routes, "ENABLED_PROVIDERS", ["github"]) + monkeypatch.setattr(oauth_routes.oauth, "create_client", lambda _: provider_client) + + get_user = AsyncMock(return_value=SimpleNamespace(id=101, email="private@example.com")) + create_session = AsyncMock( + return_value=( + "dummy-auth-cookie", + SimpleNamespace(expires_at=datetime.now(UTC) + timedelta(hours=1)), + ) + ) + monkeypatch.setattr(oauth_routes, "get_or_create_oauth_user", get_user) + monkeypatch.setattr(oauth_routes, "create_session", create_session) + + monkeypatch.setattr(oauth_routes, "verify_password", lambda pwd, hash: verify_ok) + link_mock = link if link is not None else AsyncMock(return_value=SimpleNamespace(id=7)) + monkeypatch.setattr(oauth_routes, "link_oauth_identity_to_user", link_mock) + monkeypatch.setattr(oauth_routes, "get_optional_current_user", AsyncMock(return_value=callback_user or user)) + + app = FastAPI() + app.add_middleware( + SessionMiddleware, + secret_key="dummy-test-only-signing-secret", + session_cookie="topiceye_oauth_state", + https_only=True, + same_site="lax", + ) + app.include_router(oauth_routes.router, prefix="/api/v1") + app.dependency_overrides[oauth_routes.get_current_user] = lambda: user + + async def fake_db(): + yield object() + + app.dependency_overrides[oauth_routes.get_db] = fake_db + return app, get_user, create_session, link_mock + + +def _query(location: str) -> dict: + return {k: v[0] for k, v in parse_qs(urlsplit(location).query).items()} + + +ADMIN = SimpleNamespace(id=101, email="admin@example.com", password_hash="x", role="admin") + + +@pytest.mark.asyncio +async def test_bind_start_rejects_wrong_password(monkeypatch): + """未 step-up 拒绝:密码验证失败 401。""" + provider = FakeGithub() + app, *_ = make_bind_app(monkeypatch, provider, user=ADMIN, verify_ok=False) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url=PUBLIC_ORIGIN, + ) as client: + resp = await client.post("/api/v1/auth/oauth/github/bind/start", json={"password": "wrong"}) + assert resp.status_code == 401 + assert "step-up" in resp.json()["detail"] or "密码" in resp.json()["detail"] + + +@pytest.mark.asyncio +async def test_bind_full_flow_links_without_login_session(monkeypatch): + """正常绑定:start → 回调绑定成功,fragment 带 bind=linked 且不建登录会话。""" + provider = FakeGithub() + app, get_user, create_session, link = make_bind_app(monkeypatch, provider, user=ADMIN) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url=PUBLIC_ORIGIN, + follow_redirects=False, + ) as client: + start = await client.post("/api/v1/auth/oauth/github/bind/start", json={"password": "correct"}) + assert start.status_code == 200 + assert start.json()["authorize_url"].startswith("https://github.com/login/oauth/authorize") + + callback = await client.get("/api/v1/auth/oauth/github/callback?code=dummy&state=test-state") + assert callback.status_code == 302 + location = callback.headers["location"] + assert location.startswith(PUBLIC_ORIGIN + "/oauth/callback#") + fragment = location.split("#", 1)[1] + assert "bind=linked" in fragment + assert "token=" not in fragment + + link.assert_awaited_once() + kwargs = link.await_args.kwargs + assert kwargs["provider"] == "github" + assert kwargs["provider_user_id"] == "42" + assert kwargs["email_verified"] is True + # 绑定分支不触发登录:不建用户、不建会话 + get_user.assert_not_awaited() + create_session.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bind_callback_conflict_returns_error(monkeypatch): + """已有绑定冲突:link 抛冲突 → 错误回跳,且不建登录会话。""" + provider = FakeGithub() + link = AsyncMock(side_effect=OAuthBindConflictError("该账号已绑定此登录方式,请先解绑后再重新绑定")) + app, get_user, create_session, _ = make_bind_app(monkeypatch, provider, user=ADMIN, link=link) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url=PUBLIC_ORIGIN, + follow_redirects=False, + ) as client: + await client.post("/api/v1/auth/oauth/github/bind/start", json={"password": "correct"}) + callback = await client.get("/api/v1/auth/oauth/github/callback?code=dummy&state=test-state") + assert callback.status_code == 302 + q = _query(callback.headers["location"]) + assert "已绑定" in unquote(q.get("error", "")) + get_user.assert_not_awaited() + create_session.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bind_callback_unverified_email_rejected(monkeypatch): + """未验证邮箱拒绝(#62 语义一致):userinfo 阶段即拒绝,link 不被调用。""" + provider = FakeGithub(emails=[{"email": "nope@example.com", "primary": True, "verified": False}]) + app, _, _, link = make_bind_app(monkeypatch, provider, user=ADMIN) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url=PUBLIC_ORIGIN, + follow_redirects=False, + ) as client: + await client.post("/api/v1/auth/oauth/github/bind/start", json={"password": "correct"}) + callback = await client.get("/api/v1/auth/oauth/github/callback?code=dummy&state=test-state") + assert callback.status_code == 302 + q = _query(callback.headers["location"]) + assert "邮箱" in unquote(q.get("error", "")) + link.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bind_callback_session_mismatch_rejected(monkeypatch): + """发起人与当前会话不一致(intent user 101 ≠ 当前 202):拒绝绑定。""" + provider = FakeGithub() + other = SimpleNamespace(id=202, email="other@example.com", password_hash="x") + # 发起人 = ADMIN(101),回调时会话本人 = other(202) + app, _, _, link = make_bind_app(monkeypatch, provider, user=ADMIN, callback_user=other) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url=PUBLIC_ORIGIN, + follow_redirects=False, + ) as client: + # ADMIN(101)发起,但回调时 get_optional_current_user 返回 202 + start = await client.post("/api/v1/auth/oauth/github/bind/start", json={"password": "correct"}) + assert start.status_code == 200 + callback = await client.get("/api/v1/auth/oauth/github/callback?code=dummy&state=test-state") + assert callback.status_code == 302 + q = _query(callback.headers["location"]) + assert "重新发起绑定" in unquote(q.get("error", "")), callback.headers["location"] + link.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_callback_without_bind_intent_falls_back_to_login(monkeypatch): + """无绑定意图(未经过 bind/start)时回调走原登录分支。""" + provider = FakeGithub() + app, get_user, create_session, link = make_bind_app(monkeypatch, provider, user=ADMIN) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url=PUBLIC_ORIGIN, + follow_redirects=False, + ) as client: + await client.get("/api/v1/auth/oauth/github/login") + callback = await client.get("/api/v1/auth/oauth/github/callback?code=dummy&state=test-state") + assert callback.status_code == 302 + assert "token=" not in callback.headers["location"] + get_user.assert_awaited_once() + create_session.assert_awaited_once() + link.assert_not_awaited()