feat: add typed auth record variants

This commit is contained in:
lda
2026-06-13 02:00:19 +07:00 Verified
parent c3e931e30c
commit 93fdda732d
2 changed files with 180 additions and 1 deletions
+101 -1
View File
@@ -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",
]