feat: add typed auth record variants
This commit is contained in:
+101
-1
@@ -3,7 +3,9 @@ from __future__ import annotations
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
from typing import Annotated, Any, Literal, Protocol
|
||||
|
||||
from pydantic import AnyUrl, BaseModel, Field
|
||||
|
||||
AUTH_ID_PATTERN = r"^[A-Za-z0-9_][A-Za-z0-9_.-]*$"
|
||||
|
||||
@@ -49,9 +51,107 @@ class AuthStore(Protocol):
|
||||
def load_auth(self, auth_ref: str) -> AuthRecord | None: ...
|
||||
|
||||
|
||||
class BearerAuth(BaseModel):
|
||||
kind: Literal["bearer"] = "bearer"
|
||||
access_token: str
|
||||
|
||||
|
||||
class HeaderAuth(BaseModel):
|
||||
kind: Literal["headers"] = "headers"
|
||||
headers: dict[str, str]
|
||||
|
||||
|
||||
class EnvAuth(BaseModel):
|
||||
kind: Literal["env"] = "env"
|
||||
env: dict[str, str]
|
||||
|
||||
|
||||
class OAuthRefreshTokenAuth(BaseModel):
|
||||
kind: Literal["oauth_refresh_token"] = "oauth_refresh_token"
|
||||
client_id: str
|
||||
client_secret: str
|
||||
refresh_token: str
|
||||
token_url: AnyUrl
|
||||
scopes: tuple[str, ...] = ()
|
||||
|
||||
|
||||
class OpaqueAuth(BaseModel):
|
||||
kind: Literal["opaque"] = "opaque"
|
||||
scheme: str
|
||||
payload: dict[str, object] = Field(default_factory=dict)
|
||||
|
||||
|
||||
AuthVariant = Annotated[
|
||||
BearerAuth | HeaderAuth | EnvAuth | OAuthRefreshTokenAuth | OpaqueAuth,
|
||||
Field(discriminator="kind"),
|
||||
]
|
||||
|
||||
|
||||
class StoredAuthRecord(BaseModel):
|
||||
id: str
|
||||
auth: AuthVariant
|
||||
metadata: dict[str, object] = Field(default_factory=dict)
|
||||
|
||||
def model_post_init(self, __context: Any) -> None:
|
||||
validate_auth_id(self.id)
|
||||
|
||||
|
||||
def auth_record_from_compat(
|
||||
*,
|
||||
id: str,
|
||||
scheme: str,
|
||||
payload: Mapping[str, object],
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
) -> StoredAuthRecord:
|
||||
"""Create a StoredAuthRecord from legacy scheme + payload shape."""
|
||||
payload_dict = dict(payload)
|
||||
metadata_dict = dict(metadata or {})
|
||||
match scheme:
|
||||
case "bearer":
|
||||
token = payload_dict.get("token") or payload_dict.get("access_token")
|
||||
if not isinstance(token, str) or not token:
|
||||
raise ValueError("bearer token is required")
|
||||
auth: AuthVariant = BearerAuth(access_token=token)
|
||||
case "headers":
|
||||
raw_headers = payload_dict.get("headers", {})
|
||||
headers = (
|
||||
{
|
||||
str(key): str(value)
|
||||
for key, value in raw_headers.items()
|
||||
if isinstance(key, str) and isinstance(value, str)
|
||||
}
|
||||
if isinstance(raw_headers, dict)
|
||||
else {}
|
||||
)
|
||||
auth = HeaderAuth(headers=headers)
|
||||
case "env":
|
||||
raw_env = payload_dict.get("env", {})
|
||||
env = (
|
||||
{
|
||||
str(key): str(value)
|
||||
for key, value in raw_env.items()
|
||||
if isinstance(key, str) and isinstance(value, str)
|
||||
}
|
||||
if isinstance(raw_env, dict)
|
||||
else {}
|
||||
)
|
||||
auth = EnvAuth(env=env)
|
||||
case _:
|
||||
auth = OpaqueAuth(scheme=scheme, payload=payload_dict)
|
||||
return StoredAuthRecord(id=id, auth=auth, metadata=metadata_dict)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AUTH_ID_PATTERN",
|
||||
"AuthRecord",
|
||||
"AuthStore",
|
||||
"AuthVariant",
|
||||
"BearerAuth",
|
||||
"EnvAuth",
|
||||
"HeaderAuth",
|
||||
"OAuthRefreshTokenAuth",
|
||||
"OpaqueAuth",
|
||||
"StoredAuthRecord",
|
||||
"auth_record_from_compat",
|
||||
"validate_auth_id",
|
||||
]
|
||||
|
||||
Reference in New Issue
Block a user