refactor(oauth): LinuxDo 备用端点回退、Basic Auth 认证,修复 session 外访问 ORM 对象

- LinuxDo provider: token/userinfo 请求增加 backup 端点自动回退
- LinuxDo provider: token 请求改用 HTTP Basic Auth 认证
- 授权 URL 构建: scope 为空时不再发送该参数
- OAuthService: 引入 OAuthAuthenticatedUser 快照,避免 DB session 关闭后访问 ORM 对象
- OAuthService: _handle_login_sync 设置 expire_on_commit=False 防止属性过期
- 新增 LinuxDo provider 单元测试(Basic Auth、端点回退)
- 新增 _handle_login_sync 返回快照的集成测试
This commit is contained in:
fawney19
2026-03-19 20:32:33 +08:00
parent f573110725
commit e4ebd5cca1
5 changed files with 456 additions and 141 deletions

View File

@@ -74,10 +74,12 @@ class OAuthProviderBase(ABC):
"response_type": "code", "response_type": "code",
"client_id": client_id, "client_id": client_id,
"redirect_uri": redirect_uri, "redirect_uri": redirect_uri,
"scope": self.get_effective_scopes(config),
"state": state, "state": state,
} }
) )
scopes = self.get_effective_scopes(config)
if scopes:
query["scope"] = scopes
return urlunparse(parsed._replace(query=urlencode(query))) return urlunparse(parsed._replace(query=urlencode(query)))

View File

@@ -1,6 +1,11 @@
from __future__ import annotations from __future__ import annotations
import base64
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from urllib.parse import urlparse
import httpx
from src.core.logger import logger from src.core.logger import logger
from src.services.auth.oauth.base import OAuthProviderBase from src.services.auth.oauth.base import OAuthProviderBase
@@ -41,12 +46,50 @@ class LinuxDoOAuthProvider(OAuthProviderBase):
authorization_url = "https://connect.linux.do/oauth2/authorize" authorization_url = "https://connect.linux.do/oauth2/authorize"
token_url = "https://connect.linux.do/oauth2/token" token_url = "https://connect.linux.do/oauth2/token"
userinfo_url = "https://connect.linux.do/api/user" userinfo_url = "https://connect.linux.do/api/user"
backup_token_url = "https://connect.linuxdo.org/oauth2/token"
backup_userinfo_url = "https://connect.linuxdo.org/api/user"
# LinuxDo 不需要 scope # LinuxDo 不需要 scope
default_scopes = () default_scopes = ()
@staticmethod
def _build_basic_auth_header(client_id: str, client_secret: str) -> str:
credentials = f"{client_id}:{client_secret}".encode("utf-8")
return f"Basic {base64.b64encode(credentials).decode('ascii')}"
@staticmethod
def _build_candidate_urls(primary_url: str, backup_url: str) -> list[str]:
parsed = urlparse(primary_url)
host = (parsed.hostname or "").lower().rstrip(".")
backup_path = urlparse(backup_url).path
urls = [primary_url]
if parsed.scheme == "https" and host == "connect.linux.do" and parsed.path == backup_path:
urls.append(backup_url)
return urls
@staticmethod
async def _request_with_fallback(
candidate_urls: list[str],
request_fn: Callable[[str], Awaitable[httpx.Response]],
error_code: str,
label: str,
) -> httpx.Response:
resp: httpx.Response | None = None
for idx, url in enumerate(candidate_urls):
try:
resp = await request_fn(url)
break
except httpx.HTTPError as exc:
if idx < len(candidate_urls) - 1:
logger.warning("LinuxDo {} 端点不可达,尝试备用端点: {} ({})", label, url, exc)
continue
logger.warning("LinuxDo {} 请求失败: {} ({})", label, url, exc)
raise OAuthFlowError(error_code, "transport_error") from exc
if resp is None:
raise OAuthFlowError(error_code, "no_response")
return resp
async def exchange_code(self, config: OAuthProvider, code: str) -> OAuthToken: async def exchange_code(self, config: OAuthProvider, code: str) -> OAuthToken:
url = self.get_effective_token_url(config)
client_secret = config.get_client_secret() client_secret = config.get_client_secret()
if not client_secret: if not client_secret:
raise OAuthFlowError("provider_unavailable", "client_secret 未配置") raise OAuthFlowError("provider_unavailable", "client_secret 未配置")
@@ -56,15 +99,26 @@ class LinuxDoOAuthProvider(OAuthProviderBase):
if not redirect_uri or not client_id: if not redirect_uri or not client_id:
raise OAuthFlowError("provider_unavailable", "redirect_uri/client_id 未配置") raise OAuthFlowError("provider_unavailable", "redirect_uri/client_id 未配置")
resp = await self._http_post_form( candidate_urls = self._build_candidate_urls(
self.get_effective_token_url(config), self.backup_token_url
)
headers = {
"Authorization": self._build_basic_auth_header(client_id, client_secret),
"Accept": "application/json",
}
resp = await self._request_with_fallback(
candidate_urls,
lambda url: self._http_post_form(
url, url,
data={ data={
"grant_type": "authorization_code", "grant_type": "authorization_code",
"code": code, "code": code,
"redirect_uri": redirect_uri, "redirect_uri": redirect_uri,
"client_id": client_id,
"client_secret": client_secret,
}, },
headers=headers,
),
error_code="token_exchange_failed",
label="token",
) )
if resp.status_code >= 400: if resp.status_code >= 400:
@@ -87,8 +141,15 @@ class LinuxDoOAuthProvider(OAuthProviderBase):
) )
async def get_user_info(self, config: OAuthProvider, access_token: str) -> OAuthUserInfo: async def get_user_info(self, config: OAuthProvider, access_token: str) -> OAuthUserInfo:
url = self.get_effective_userinfo_url(config) candidate_urls = self._build_candidate_urls(
resp = await self._http_get(url, headers={"Authorization": f"Bearer {access_token}"}) self.get_effective_userinfo_url(config), self.backup_userinfo_url
)
resp = await self._request_with_fallback(
candidate_urls,
lambda url: self._http_get(url, headers={"Authorization": f"Bearer {access_token}"}),
error_code="userinfo_fetch_failed",
label="userinfo",
)
if resp.status_code >= 400: if resp.status_code >= 400:
logger.warning("LinuxDo userinfo 获取失败: status={}", resp.status_code) logger.warning("LinuxDo userinfo 获取失败: status={}", resp.status_code)

View File

@@ -46,14 +46,25 @@ class OAuthCallbackResult:
refresh_token: str | None = field(default=None, repr=False) refresh_token: str | None = field(default=None, repr=False)
@dataclass(frozen=True)
class OAuthAuthenticatedUser:
user_id: str
email: str | None
role: UserRole
created_at: datetime | None
class OAuthService: class OAuthService:
"""OAuth 核心业务服务v1""" """OAuth 核心业务服务v1"""
@staticmethod @staticmethod
def _handle_login_sync(provider_type: str, oauth_user: OAuthUserInfo) -> User: def _handle_login_sync(provider_type: str, oauth_user: OAuthUserInfo) -> OAuthAuthenticatedUser:
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
with get_db_context() as db: with get_db_context() as db:
original_expire_on_commit = getattr(db, "expire_on_commit", True)
db.expire_on_commit = False
try:
existing_link = ( existing_link = (
db.query(UserOAuthLink) db.query(UserOAuthLink)
.filter( .filter(
@@ -70,8 +81,14 @@ class OAuthService:
linked_user.last_login_at = now linked_user.last_login_at = now
existing_link.last_login_at = now existing_link.last_login_at = now
db.commit() db.commit()
db.expunge(linked_user) assert linked_user.id is not None
return linked_user assert linked_user.role is not None
return OAuthAuthenticatedUser(
user_id=linked_user.id,
email=linked_user.email,
role=linked_user.role,
created_at=linked_user.created_at,
)
enable_registration = SystemConfigService.get_config( enable_registration = SystemConfigService.get_config(
db, "enable_registration", default=False db, "enable_registration", default=False
@@ -181,13 +198,26 @@ class OAuthService:
existing_user.last_login_at = now existing_user.last_login_at = now
existing_link.last_login_at = now existing_link.last_login_at = now
db.commit() db.commit()
db.expunge(existing_user) assert existing_user.id is not None
return existing_user assert existing_user.role is not None
return OAuthAuthenticatedUser(
user_id=existing_user.id,
email=existing_user.email,
role=existing_user.role,
created_at=existing_user.created_at,
)
raise OAuthFlowError("oauth_already_bound") raise OAuthFlowError("oauth_already_bound")
raise OAuthFlowError("provider_error", "link_create_failed") raise OAuthFlowError("provider_error", "link_create_failed")
db.expunge(user) assert user.role is not None
return user return OAuthAuthenticatedUser(
user_id=user.id,
email=user.email,
role=user.role,
created_at=user.created_at,
)
finally:
db.expire_on_commit = original_expire_on_commit
@staticmethod @staticmethod
def _handle_bind_sync( def _handle_bind_sync(
@@ -747,13 +777,18 @@ class OAuthService:
) )
) )
assert user.id is not None db_user = db.query(User).filter(User.id == user.user_id).first()
assert user.role is not None if not db_user or not db_user.is_active or db_user.is_deleted:
return OAuthCallbackResult(
redirect_url=OAuthService._build_frontend_error_redirect(
frontend_callback_url, error_code="account_disabled"
)
)
session_id = str(uuid.uuid4()) session_id = str(uuid.uuid4())
access_token = AuthService.create_access_token( access_token = AuthService.create_access_token(
data={ data={
"user_id": user.id, "user_id": user.user_id,
"role": user.role.value, "role": user.role.value,
"created_at": user.created_at.isoformat() if user.created_at else None, "created_at": user.created_at.isoformat() if user.created_at else None,
"session_id": session_id, "session_id": session_id,
@@ -761,7 +796,7 @@ class OAuthService:
) )
refresh_token = AuthService.create_refresh_token( refresh_token = AuthService.create_refresh_token(
data={ data={
"user_id": user.id, "user_id": user.user_id,
"created_at": user.created_at.isoformat() if user.created_at else None, "created_at": user.created_at.isoformat() if user.created_at else None,
"session_id": session_id, "session_id": session_id,
"jti": str(uuid.uuid4()), "jti": str(uuid.uuid4()),
@@ -775,7 +810,7 @@ class OAuthService:
) )
SessionService.create_session( SessionService.create_session(
db, db,
user=user, user=db_user,
session_id=session_id, session_id=session_id,
refresh_token=refresh_token, refresh_token=refresh_token,
expires_at=AuthService.get_refresh_token_expiry(), expires_at=AuthService.get_refresh_token_expiry(),
@@ -793,14 +828,13 @@ class OAuthService:
@staticmethod @staticmethod
async def _handle_login( async def _handle_login(
db: Session, *, config: OAuthProvider, oauth_user: OAuthUserInfo db: Session, *, config: OAuthProvider, oauth_user: OAuthUserInfo
) -> User: ) -> OAuthAuthenticatedUser:
user = await run_in_threadpool( user = await run_in_threadpool(
OAuthService._handle_login_sync, OAuthService._handle_login_sync,
config.provider_type, config.provider_type,
oauth_user, oauth_user,
) )
assert user.id is not None await UserCacheService.invalidate_user_cache(user.user_id, user.email)
await UserCacheService.invalidate_user_cache(user.id, user.email)
return user return user
@staticmethod @staticmethod

View File

@@ -0,0 +1,138 @@
from __future__ import annotations
from types import SimpleNamespace
from urllib.parse import parse_qs, urlparse
import httpx
import pytest
from src.services.auth.oauth.providers.linuxdo import LinuxDoOAuthProvider
def _make_config() -> SimpleNamespace:
return SimpleNamespace(
client_id="client-id",
redirect_uri="https://api.example.com/api/oauth/linuxdo/callback",
authorization_url_override=None,
token_url_override=None,
userinfo_url_override=None,
scopes=None,
get_client_secret=lambda: "client-secret",
)
def test_linuxdo_authorization_url_omits_empty_scope() -> None:
provider = LinuxDoOAuthProvider()
url = provider.get_authorization_url(_make_config(), "state-1")
parsed = urlparse(url)
params = parse_qs(parsed.query, keep_blank_values=True)
assert parsed.netloc == "connect.linux.do"
assert "scope" not in params
assert params["state"] == ["state-1"]
@pytest.mark.asyncio
async def test_linuxdo_exchange_code_uses_basic_auth(monkeypatch: pytest.MonkeyPatch) -> None:
provider = LinuxDoOAuthProvider()
config = _make_config()
captured: dict[str, object] = {}
async def _fake_post_form(
url: str,
data: dict[str, str],
*,
timeout_seconds: float = 5.0,
headers: dict[str, str] | None = None,
) -> httpx.Response:
captured["url"] = url
captured["data"] = data
captured["headers"] = headers or {}
captured["timeout_seconds"] = timeout_seconds
return httpx.Response(
200,
json={"access_token": "access-1", "token_type": "bearer"},
request=httpx.Request("POST", url),
)
monkeypatch.setattr(provider, "_http_post_form", _fake_post_form)
token = await provider.exchange_code(config, "code-1")
assert token.access_token == "access-1"
assert captured["url"] == provider.token_url
assert captured["data"] == {
"grant_type": "authorization_code",
"code": "code-1",
"redirect_uri": config.redirect_uri,
}
assert captured["headers"] == {
"Authorization": provider._build_basic_auth_header("client-id", "client-secret"),
"Accept": "application/json",
}
@pytest.mark.asyncio
async def test_linuxdo_exchange_code_falls_back_to_backup_endpoint(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = LinuxDoOAuthProvider()
config = _make_config()
called_urls: list[str] = []
async def _fake_post_form(
url: str,
data: dict[str, str],
*,
timeout_seconds: float = 5.0,
headers: dict[str, str] | None = None,
) -> httpx.Response:
called_urls.append(url)
if len(called_urls) == 1:
raise httpx.ConnectError("network down", request=httpx.Request("POST", url))
return httpx.Response(
200,
json={"access_token": "access-2", "token_type": "bearer"},
request=httpx.Request("POST", url),
)
monkeypatch.setattr(provider, "_http_post_form", _fake_post_form)
token = await provider.exchange_code(config, "code-2")
assert token.access_token == "access-2"
assert called_urls == [provider.token_url, provider.backup_token_url]
@pytest.mark.asyncio
async def test_linuxdo_userinfo_falls_back_to_backup_endpoint(
monkeypatch: pytest.MonkeyPatch,
) -> None:
provider = LinuxDoOAuthProvider()
config = _make_config()
called_urls: list[str] = []
async def _fake_get(
url: str,
*,
timeout_seconds: float = 5.0,
headers: dict[str, str] | None = None,
) -> httpx.Response:
called_urls.append(url)
if len(called_urls) == 1:
raise httpx.ConnectError("network down", request=httpx.Request("GET", url))
return httpx.Response(
200,
json={"id": 42, "username": "neo", "email": "Neo@Linux.Do"},
request=httpx.Request("GET", url),
)
monkeypatch.setattr(provider, "_http_get", _fake_get)
user = await provider.get_user_info(config, "access-token")
assert user.id == "42"
assert user.username == "neo"
assert user.email == "neo@linux.do"
assert called_urls == [provider.userinfo_url, provider.backup_userinfo_url]

View File

@@ -1,12 +1,16 @@
from __future__ import annotations from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
from urllib.parse import parse_qs, urlparse from urllib.parse import parse_qs, urlparse
import pytest import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from src.core.enums import AuthSource from src.core.enums import AuthSource, UserRole
from src.models.database import Base, OAuthProvider, User, UserOAuthLink
from src.services.auth.oauth.service import OAuthService from src.services.auth.oauth.service import OAuthService
from src.services.auth.oauth.state import OAuthStateData from src.services.auth.oauth.state import OAuthStateData
@@ -127,3 +131,79 @@ async def test_handle_callback_allows_bind_state_without_device_id(
parsed = urlparse(result.redirect_url) parsed = urlparse(result.redirect_url)
assert result.refresh_token is None assert result.refresh_token is None
assert parse_qs(parsed.query)["oauth_bound"] == ["GitHub"] assert parse_qs(parsed.query)["oauth_bound"] == ["GitHub"]
def test_handle_login_sync_returns_user_snapshot_outside_db_session(
monkeypatch: pytest.MonkeyPatch,
) -> None:
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(
engine,
tables=[User.__table__, OAuthProvider.__table__, UserOAuthLink.__table__],
)
SessionLocal = sessionmaker(bind=engine)
with SessionLocal() as db:
provider = OAuthProvider(
provider_type="linuxdo",
display_name="Linux Do",
client_id="client-id",
redirect_uri="https://api.example.com/api/oauth/linuxdo/callback",
frontend_callback_url="https://app.example.com/auth/callback",
is_enabled=True,
)
user = User(
email="user@example.com",
email_verified=True,
username="tester",
auth_source=AuthSource.OAUTH,
role=UserRole.USER,
is_active=True,
is_deleted=False,
)
db.add(provider)
db.add(user)
db.flush()
user_id = str(user.id)
db.add(
UserOAuthLink(
user_id=user_id,
provider_type="linuxdo",
provider_user_id="linuxdo-user-1",
provider_username="tester",
provider_email="user@example.com",
extra_data={},
)
)
db.commit()
@contextmanager
def _fake_get_db_context():
db = SessionLocal()
try:
yield db
db.commit()
except Exception:
db.rollback()
raise
finally:
db.close()
monkeypatch.setattr("src.services.auth.oauth.service.get_db_context", _fake_get_db_context)
snapshot = OAuthService._handle_login_sync(
"linuxdo",
SimpleNamespace(
id="linuxdo-user-1",
username="tester",
email="user@example.com",
email_verified=True,
raw={},
),
)
assert snapshot.user_id == user_id
assert snapshot.email == "user@example.com"
assert snapshot.role == UserRole.USER
engine.dispose()