feat: add oauth login flow helpers
This commit is contained in:
@@ -0,0 +1,88 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
from wf_api.auth import OAuthRefreshTokenAuth, StoredAuthRecord
|
||||
from wf_config import OAuthProviderConfig
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OAuthLoginResult:
|
||||
refresh_token: str | None
|
||||
subject: str | None = None
|
||||
scopes: tuple[str, ...] = ()
|
||||
|
||||
|
||||
def build_oauth_record(
|
||||
*,
|
||||
auth_ref: str,
|
||||
provider_name: str,
|
||||
provider: OAuthProviderConfig,
|
||||
client_id: str,
|
||||
client_secret: str | None,
|
||||
result: OAuthLoginResult,
|
||||
) -> StoredAuthRecord:
|
||||
if not result.refresh_token:
|
||||
raise ValueError("OAuth login did not return a refresh token")
|
||||
metadata: dict[str, object] = {"provider": provider_name}
|
||||
if result.subject:
|
||||
metadata["subject"] = result.subject
|
||||
if result.scopes:
|
||||
metadata["scopes"] = list(result.scopes)
|
||||
return StoredAuthRecord(
|
||||
id=auth_ref,
|
||||
auth=OAuthRefreshTokenAuth(
|
||||
client_id=client_id,
|
||||
client_secret=client_secret or "",
|
||||
refresh_token=result.refresh_token,
|
||||
token_url=provider.token_url,
|
||||
scopes=tuple(result.scopes or provider.scopes),
|
||||
),
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
||||
class OAuthClientLike(Protocol):
|
||||
def create_authorization_url(
|
||||
self, auth_url: str, **kwargs: object
|
||||
) -> tuple[str, str]: ...
|
||||
|
||||
async def fetch_token(
|
||||
self, token_url: str, authorization_response: str
|
||||
) -> dict[str, object]: ...
|
||||
|
||||
|
||||
OAuthClientFactory = Callable[..., OAuthClientLike]
|
||||
|
||||
|
||||
class OAuthCodeLoginFlow:
|
||||
def __init__(self, client_factory: OAuthClientFactory) -> None:
|
||||
self._client_factory = client_factory
|
||||
|
||||
async def login_with_authorization_response(
|
||||
self,
|
||||
*,
|
||||
provider: OAuthProviderConfig,
|
||||
client_id: str,
|
||||
client_secret: str | None,
|
||||
authorization_response: str,
|
||||
) -> OAuthLoginResult:
|
||||
client = self._client_factory(
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
scope=" ".join(provider.scopes),
|
||||
code_challenge_method="S256",
|
||||
)
|
||||
client.create_authorization_url(str(provider.auth_url))
|
||||
token = await client.fetch_token(
|
||||
str(provider.token_url),
|
||||
authorization_response=authorization_response,
|
||||
)
|
||||
refresh_token = token.get("refresh_token")
|
||||
if refresh_token is not None and not isinstance(refresh_token, str):
|
||||
raise ValueError("OAuth refresh_token must be a string")
|
||||
raw_scope = token.get("scope")
|
||||
scopes = tuple(str(raw_scope).split()) if raw_scope else provider.scopes
|
||||
return OAuthLoginResult(refresh_token=refresh_token, scopes=scopes)
|
||||
@@ -0,0 +1,99 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from wf_api.auth import OAuthRefreshTokenAuth
|
||||
from wf_cli.oauth import OAuthCodeLoginFlow, OAuthLoginResult, build_oauth_record
|
||||
from wf_config import OAuthProviderConfig
|
||||
|
||||
|
||||
def test_build_oauth_record_creates_refresh_token_auth() -> None:
|
||||
provider = OAuthProviderConfig(
|
||||
kind="oauth_authorization_code_pkce",
|
||||
auth_url="https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url="https://oauth2.googleapis.com/token",
|
||||
client_id_env="GOOGLE_OAUTH_CLIENT_ID",
|
||||
client_secret_env="GOOGLE_OAUTH_CLIENT_SECRET",
|
||||
scopes=("https://www.googleapis.com/auth/drive.readonly",),
|
||||
)
|
||||
result = OAuthLoginResult(
|
||||
refresh_token="refresh",
|
||||
subject="[email protected]",
|
||||
scopes=("https://www.googleapis.com/auth/drive.readonly",),
|
||||
)
|
||||
|
||||
record = build_oauth_record(
|
||||
auth_ref="google.drive.personal",
|
||||
provider_name="google",
|
||||
provider=provider,
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
result=result,
|
||||
)
|
||||
|
||||
assert record.id == "google.drive.personal"
|
||||
assert isinstance(record.auth, OAuthRefreshTokenAuth)
|
||||
assert record.auth.client_id == "client"
|
||||
assert record.auth.client_secret == "secret"
|
||||
assert record.auth.refresh_token == "refresh"
|
||||
assert record.metadata["provider"] == "google"
|
||||
assert record.metadata["subject"] == "[email protected]"
|
||||
|
||||
|
||||
def test_build_oauth_record_rejects_missing_refresh_token() -> None:
|
||||
provider = OAuthProviderConfig(
|
||||
kind="oauth_authorization_code_pkce",
|
||||
auth_url="https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url="https://oauth2.googleapis.com/token",
|
||||
client_id_env="GOOGLE_OAUTH_CLIENT_ID",
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="refresh token"):
|
||||
build_oauth_record(
|
||||
auth_ref="google.drive.personal",
|
||||
provider_name="google",
|
||||
provider=provider,
|
||||
client_id="client",
|
||||
client_secret=None,
|
||||
result=OAuthLoginResult(refresh_token=None),
|
||||
)
|
||||
|
||||
|
||||
class _FakeOAuthClient:
|
||||
def __init__(self) -> None:
|
||||
self.authorization_url = "https://auth.example/authorize?state=abc"
|
||||
self.fetch_calls: list[str] = []
|
||||
|
||||
def create_authorization_url(self, auth_url: str, **kwargs: object) -> tuple[str, str]:
|
||||
assert auth_url == "https://accounts.google.com/o/oauth2/v2/auth"
|
||||
return self.authorization_url, "state-123"
|
||||
|
||||
async def fetch_token(self, token_url: str, authorization_response: str) -> dict[str, object]:
|
||||
self.fetch_calls.append(authorization_response)
|
||||
return {
|
||||
"refresh_token": "refresh",
|
||||
"scope": "https://www.googleapis.com/auth/drive.readonly",
|
||||
}
|
||||
|
||||
|
||||
async def test_oauth_code_login_flow_uses_injected_client() -> None:
|
||||
provider = OAuthProviderConfig(
|
||||
kind="oauth_authorization_code_pkce",
|
||||
auth_url="https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url="https://oauth2.googleapis.com/token",
|
||||
client_id_env="GOOGLE_OAUTH_CLIENT_ID",
|
||||
scopes=("https://www.googleapis.com/auth/drive.readonly",),
|
||||
)
|
||||
client = _FakeOAuthClient()
|
||||
flow = OAuthCodeLoginFlow(client_factory=lambda **kwargs: client)
|
||||
|
||||
result = await flow.login_with_authorization_response(
|
||||
provider=provider,
|
||||
client_id="client",
|
||||
client_secret=None,
|
||||
authorization_response="http://127.0.0.1/callback?code=abc&state=state-123",
|
||||
)
|
||||
|
||||
assert result.refresh_token == "refresh"
|
||||
assert result.scopes == ("https://www.googleapis.com/auth/drive.readonly",)
|
||||
assert client.fetch_calls == ["http://127.0.0.1/callback?code=abc&state=state-123"]
|
||||
Reference in New Issue
Block a user