현재 작업 전체 반영
This commit is contained in:
parent
5560638e54
commit
c0dddab594
85 changed files with 11322 additions and 539 deletions
|
|
@ -11,6 +11,9 @@ from __future__ import annotations
|
|||
|
||||
import base64
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -34,11 +37,14 @@ from ..saml import (
|
|||
)
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
GOOGLE_AUTHORIZE_URL = "https://accounts.google.com/o/oauth2/v2/auth"
|
||||
GOOGLE_TOKEN_URL = "https://oauth2.googleapis.com/token"
|
||||
GOOGLE_TOKENINFO_URL = "https://oauth2.googleapis.com/tokeninfo"
|
||||
OAUTH_STATE_TTL_SECONDS = 10 * 60
|
||||
OAUTH_STATE_COOKIE_NAME = "__Host-vignette_oauth_state"
|
||||
DEV_OAUTH_STATE_COOKIE_NAME = "vignette_oauth_state"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
|
|
@ -201,6 +207,54 @@ def _role_for_saml_identity(identity: SamlIdentity) -> Role:
|
|||
return _role_for_email(identity.email)
|
||||
|
||||
|
||||
def _split_cohort_values(value: str | None) -> list[str]:
|
||||
if not value:
|
||||
return []
|
||||
return [item.strip() for item in value.split(",") if item.strip()]
|
||||
|
||||
|
||||
def _append_unique(items: list[str], values: list[str]) -> None:
|
||||
seen = {item.lower() for item in items}
|
||||
for value in values:
|
||||
key = value.lower()
|
||||
if key and key not in seen:
|
||||
items.append(value)
|
||||
seen.add(key)
|
||||
|
||||
|
||||
def _configured_cohort_ids(
|
||||
*,
|
||||
email: str,
|
||||
hosted_domain: str | None = None,
|
||||
claim_hint: str | None = None,
|
||||
) -> list[str]:
|
||||
normalized_email = _normalize_email(email)
|
||||
domain = _email_domain(normalized_email)
|
||||
hd = _normalize_domain(hosted_domain)
|
||||
email_map = {
|
||||
_normalize_email(key): value
|
||||
for key, value in settings.auth_email_cohort_map.items()
|
||||
if _normalize_email(key)
|
||||
}
|
||||
domain_map = {
|
||||
_normalize_domain(key): value
|
||||
for key, value in settings.auth_domain_cohort_map.items()
|
||||
if _normalize_domain(key)
|
||||
}
|
||||
cohorts: list[str] = []
|
||||
_append_unique(cohorts, _split_cohort_values(email_map.get(normalized_email)))
|
||||
_append_unique(cohorts, _split_cohort_values(domain_map.get(domain)))
|
||||
if hd and hd != domain:
|
||||
_append_unique(cohorts, _split_cohort_values(domain_map.get(hd)))
|
||||
_append_unique(cohorts, _split_cohort_values(claim_hint))
|
||||
return cohorts
|
||||
|
||||
|
||||
def _provider_external_id(provider: str, subject: str | None, email: str) -> str:
|
||||
value = (subject or "").strip() or _normalize_email(email)
|
||||
return f"{provider}:{value.lower()}"
|
||||
|
||||
|
||||
def _safe_next_path(next_path: str | None) -> str:
|
||||
if not next_path or not next_path.startswith("/") or next_path.startswith("//"):
|
||||
return "/"
|
||||
|
|
@ -254,7 +308,12 @@ def _is_dev_login_allowed_origin(origin: str) -> bool:
|
|||
|
||||
def _configured_frontend_origins() -> list[str]:
|
||||
origins: list[str] = []
|
||||
for value in [settings.frontend_base_url, *settings.frontend_origin_map.values(), *settings.cors_origins]:
|
||||
for value in [
|
||||
settings.frontend_base_url,
|
||||
*settings.frontend_origin_map.values(),
|
||||
*settings.cors_origins,
|
||||
*settings.auth_dev_login_extra_origins,
|
||||
]:
|
||||
origin = _url_origin(value)
|
||||
if origin and origin not in origins:
|
||||
origins.append(origin)
|
||||
|
|
@ -276,6 +335,17 @@ def _frontend_origin_for_request(request: Request | None = None) -> str:
|
|||
hostname = host.rsplit(":", 1)[0].lower() if host else ""
|
||||
if mapped_origin := _frontend_origin_map().get(hostname):
|
||||
return mapped_origin
|
||||
forwarded_proto = (
|
||||
request.headers.get("x-forwarded-proto", "").split(",", 1)[0].strip().lower()
|
||||
)
|
||||
if forwarded_host and forwarded_proto in {"http", "https"}:
|
||||
forwarded_origin = _url_origin(f"{forwarded_proto}://{host}")
|
||||
if (
|
||||
settings.environment == "dev"
|
||||
and forwarded_origin
|
||||
and _is_dev_login_allowed_origin(forwarded_origin)
|
||||
):
|
||||
return forwarded_origin
|
||||
if hostname in {"localhost", "127.0.0.1", "::1"}:
|
||||
return fallback
|
||||
|
||||
|
|
@ -299,6 +369,93 @@ def _pkce_challenge(verifier: str) -> str:
|
|||
return base64.urlsafe_b64encode(digest).decode("ascii").rstrip("=")
|
||||
|
||||
|
||||
def _b64url_encode(value: bytes) -> str:
|
||||
return base64.urlsafe_b64encode(value).decode("ascii").rstrip("=")
|
||||
|
||||
|
||||
def _b64url_decode(value: str) -> bytes:
|
||||
padded = value + ("=" * (-len(value) % 4))
|
||||
return base64.urlsafe_b64decode(padded.encode("ascii"))
|
||||
|
||||
|
||||
def _oauth_state_signature(payload: str) -> str:
|
||||
digest = hmac.new(
|
||||
settings.session_secret.encode("utf-8"),
|
||||
f"oauth-state:{payload}".encode("utf-8"),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
return _b64url_encode(digest)
|
||||
|
||||
|
||||
def _oauth_code_verifier_for_state(state: str) -> str:
|
||||
digest = hmac.new(
|
||||
settings.session_secret.encode("utf-8"),
|
||||
f"oauth-pkce:{state}".encode("utf-8"),
|
||||
hashlib.sha256,
|
||||
).digest()
|
||||
return _b64url_encode(digest)
|
||||
|
||||
|
||||
def _build_oauth_state(next_path: str | None) -> tuple[str, OAuthState]:
|
||||
created_at = time.time()
|
||||
safe_next = _safe_next_path(next_path)
|
||||
payload = {
|
||||
"iat": created_at,
|
||||
"next": safe_next,
|
||||
"nonce": secrets.token_urlsafe(24),
|
||||
}
|
||||
payload_blob = _b64url_encode(
|
||||
json.dumps(payload, separators=(",", ":"), sort_keys=True).encode("utf-8")
|
||||
)
|
||||
state = f"{payload_blob}.{_oauth_state_signature(payload_blob)}"
|
||||
return state, OAuthState(
|
||||
code_verifier=_oauth_code_verifier_for_state(state),
|
||||
next_path=safe_next,
|
||||
created_at=created_at,
|
||||
)
|
||||
|
||||
|
||||
def _oauth_state_from_signed_token(state: str | None) -> OAuthState | None:
|
||||
if not state or "." not in state:
|
||||
return None
|
||||
payload_blob, signature = state.rsplit(".", 1)
|
||||
if not payload_blob or not signature:
|
||||
return None
|
||||
if not hmac.compare_digest(signature, _oauth_state_signature(payload_blob)):
|
||||
return None
|
||||
|
||||
try:
|
||||
payload = json.loads(_b64url_decode(payload_blob).decode("utf-8"))
|
||||
created_at = float(payload.get("iat", 0))
|
||||
except (ValueError, TypeError, json.JSONDecodeError):
|
||||
return None
|
||||
|
||||
if created_at <= 0 or created_at < time.time() - OAUTH_STATE_TTL_SECONDS:
|
||||
return None
|
||||
if not isinstance(payload.get("nonce"), str):
|
||||
return None
|
||||
|
||||
next_path = payload.get("next")
|
||||
if not isinstance(next_path, str):
|
||||
return None
|
||||
return OAuthState(
|
||||
code_verifier=_oauth_code_verifier_for_state(state),
|
||||
next_path=_safe_next_path(next_path),
|
||||
created_at=created_at,
|
||||
)
|
||||
|
||||
|
||||
def _oauth_state_for_callback(state: str | None, cookie_state: str | None) -> OAuthState | None:
|
||||
if not state:
|
||||
return None
|
||||
if cookie_state != state:
|
||||
return None
|
||||
stored = _oauth_states.pop(state, None)
|
||||
if stored is not None:
|
||||
return stored
|
||||
return _oauth_state_from_signed_token(state)
|
||||
|
||||
|
||||
def _prune_oauth_states() -> None:
|
||||
cutoff = time.time() - OAUTH_STATE_TTL_SECONDS
|
||||
stale = [key for key, value in _oauth_states.items() if value.created_at < cutoff]
|
||||
|
|
@ -342,6 +499,46 @@ def _set_session_cookie(response: Response, sid: str) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _set_oauth_state_cookie(response: Response, state: str) -> None:
|
||||
response.set_cookie(
|
||||
key=OAUTH_STATE_COOKIE_NAME,
|
||||
value=state,
|
||||
max_age=OAUTH_STATE_TTL_SECONDS,
|
||||
httponly=True,
|
||||
secure=True,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
if settings.environment == "dev":
|
||||
response.set_cookie(
|
||||
key=DEV_OAUTH_STATE_COOKIE_NAME,
|
||||
value=state,
|
||||
max_age=OAUTH_STATE_TTL_SECONDS,
|
||||
httponly=True,
|
||||
secure=False,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
|
||||
|
||||
def _delete_oauth_state_cookie(response: Response) -> None:
|
||||
response.delete_cookie(
|
||||
OAUTH_STATE_COOKIE_NAME,
|
||||
httponly=True,
|
||||
secure=True,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
if settings.environment == "dev":
|
||||
response.delete_cookie(
|
||||
DEV_OAUTH_STATE_COOKIE_NAME,
|
||||
httponly=True,
|
||||
secure=False,
|
||||
samesite="lax",
|
||||
path="/",
|
||||
)
|
||||
|
||||
|
||||
def _delete_session_cookie(response: Response) -> None:
|
||||
response.delete_cookie(
|
||||
settings.cookie_name,
|
||||
|
|
@ -375,13 +572,28 @@ def _frontend_login_redirect(reason: str, request: Request) -> RedirectResponse:
|
|||
return RedirectResponse(f"{base_url}/login?{urlencode({'oauth': reason})}", status_code=302)
|
||||
|
||||
|
||||
def _oauth_callback_error(reason: str, request: Request) -> RedirectResponse:
|
||||
response = _frontend_login_redirect(reason, request)
|
||||
_delete_oauth_state_cookie(response)
|
||||
return response
|
||||
|
||||
|
||||
def _log_oauth_callback_failure(request: Request, reason: str, **fields: object) -> None:
|
||||
"""Log OAuth callback failures without authorization codes, tokens, or raw user IDs."""
|
||||
host = request.headers.get("x-forwarded-host") or request.headers.get("host")
|
||||
logger.warning(
|
||||
"google_oauth_callback_failed reason=%s host=%s forwarded_proto=%s details=%s",
|
||||
reason,
|
||||
host,
|
||||
request.headers.get("x-forwarded-proto"),
|
||||
fields,
|
||||
)
|
||||
|
||||
|
||||
def _dev_login_available(request: Request) -> bool:
|
||||
if settings.environment != "dev" or not settings.auth_dev_login_enabled:
|
||||
return False
|
||||
|
||||
if _dev_login_extra_origins():
|
||||
return True
|
||||
|
||||
saw_browser_origin = False
|
||||
for header_name in ("origin", "referer"):
|
||||
origin = _url_origin(request.headers.get(header_name))
|
||||
|
|
@ -389,16 +601,28 @@ def _dev_login_available(request: Request) -> bool:
|
|||
saw_browser_origin = True
|
||||
if not _is_dev_login_allowed_origin(origin):
|
||||
return False
|
||||
|
||||
if not saw_browser_origin and _dev_login_extra_origins():
|
||||
if saw_browser_origin:
|
||||
return True
|
||||
|
||||
forwarded_host = request.headers.get("x-forwarded-host")
|
||||
host = (forwarded_host or request.headers.get("host") or "").split(",", 1)[0].strip()
|
||||
origin = _url_origin(f"http://{host}") if host else None
|
||||
forwarded_proto = (
|
||||
request.headers.get("x-forwarded-proto", "").split(",", 1)[0].strip().lower()
|
||||
)
|
||||
scheme = forwarded_proto if forwarded_host and forwarded_proto in {"http", "https"} else "http"
|
||||
origin = _url_origin(f"{scheme}://{host}") if host else None
|
||||
return bool(origin and _is_dev_login_allowed_origin(origin))
|
||||
|
||||
|
||||
def _dev_oauth_redirect_unavailable(request: Request) -> bool:
|
||||
redirect_origin = _url_origin(settings.oauth_redirect_uri)
|
||||
return bool(
|
||||
_dev_login_available(request)
|
||||
and redirect_origin
|
||||
and not _is_local_origin(redirect_origin)
|
||||
)
|
||||
|
||||
|
||||
@router.get("/config", response_model=AuthConfigResponse)
|
||||
async def auth_config(request: Request) -> AuthConfigResponse:
|
||||
"""Return non-secret login configuration for the browser login screen."""
|
||||
|
|
@ -449,15 +673,12 @@ async def login(
|
|||
return _frontend_login_redirect("unsupported_provider", request)
|
||||
if not _google_configured():
|
||||
return _frontend_login_redirect("not_configured", request)
|
||||
if _dev_oauth_redirect_unavailable(request):
|
||||
return _frontend_login_redirect("local_oauth_unavailable", request)
|
||||
|
||||
_prune_oauth_states()
|
||||
state = secrets.token_urlsafe(32)
|
||||
verifier = secrets.token_urlsafe(64)
|
||||
_oauth_states[state] = OAuthState(
|
||||
code_verifier=verifier,
|
||||
next_path=_safe_next_path(next),
|
||||
created_at=time.time(),
|
||||
)
|
||||
state, stored_state = _build_oauth_state(next)
|
||||
_oauth_states[state] = stored_state
|
||||
|
||||
params = {
|
||||
"client_id": settings.oauth_google_client_id,
|
||||
|
|
@ -465,11 +686,13 @@ async def login(
|
|||
"response_type": "code",
|
||||
"scope": "openid email profile",
|
||||
"state": state,
|
||||
"code_challenge": _pkce_challenge(verifier),
|
||||
"code_challenge": _pkce_challenge(stored_state.code_verifier),
|
||||
"code_challenge_method": "S256",
|
||||
"prompt": "select_account",
|
||||
}
|
||||
return RedirectResponse(f"{GOOGLE_AUTHORIZE_URL}?{urlencode(params)}", status_code=302)
|
||||
response = RedirectResponse(f"{GOOGLE_AUTHORIZE_URL}?{urlencode(params)}", status_code=302)
|
||||
_set_oauth_state_cookie(response, state)
|
||||
return response
|
||||
|
||||
|
||||
@router.get("/callback")
|
||||
|
|
@ -477,15 +700,45 @@ async def callback(
|
|||
request: Request,
|
||||
code: Annotated[Optional[str], Query()] = None,
|
||||
state: Annotated[Optional[str], Query()] = None,
|
||||
error: Annotated[Optional[str], Query()] = None,
|
||||
error_description: Annotated[Optional[str], Query()] = None,
|
||||
oauth_state_cookie: Annotated[Optional[str], Cookie(alias=OAUTH_STATE_COOKIE_NAME)] = None,
|
||||
dev_oauth_state_cookie: Annotated[Optional[str], Cookie(alias=DEV_OAUTH_STATE_COOKIE_NAME)] = None,
|
||||
) -> RedirectResponse:
|
||||
"""Exchange Google auth code, validate identity, and issue a BFF cookie."""
|
||||
if error:
|
||||
reason = "access_denied" if error == "access_denied" else "provider_error"
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
reason,
|
||||
provider_error=error,
|
||||
has_state=bool(state),
|
||||
has_error_description=bool(error_description),
|
||||
)
|
||||
return _oauth_callback_error(reason, request)
|
||||
|
||||
if not code or not state:
|
||||
return _frontend_login_redirect("missing_callback", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"missing_callback",
|
||||
has_code=bool(code),
|
||||
has_state=bool(state),
|
||||
)
|
||||
return _oauth_callback_error("missing_callback", request)
|
||||
|
||||
_prune_oauth_states()
|
||||
stored = _oauth_states.pop(state, None)
|
||||
cookie_state = oauth_state_cookie or (
|
||||
dev_oauth_state_cookie if settings.environment == "dev" else None
|
||||
)
|
||||
stored = _oauth_state_for_callback(state, cookie_state)
|
||||
if stored is None:
|
||||
return _frontend_login_redirect("invalid_state", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"invalid_state",
|
||||
has_cookie=bool(cookie_state),
|
||||
state_in_memory=state in _oauth_states if state else False,
|
||||
)
|
||||
return _oauth_callback_error("invalid_state", request)
|
||||
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
token_res = await client.post(
|
||||
|
|
@ -501,22 +754,45 @@ async def callback(
|
|||
headers={"Accept": "application/json"},
|
||||
)
|
||||
if token_res.status_code >= 400:
|
||||
return _frontend_login_redirect("token_exchange_failed", request)
|
||||
token_error: object
|
||||
try:
|
||||
token_body = token_res.json()
|
||||
token_error = {
|
||||
"error": token_body.get("error"),
|
||||
"error_description": token_body.get("error_description"),
|
||||
}
|
||||
except Exception:
|
||||
token_error = "non_json_error"
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"token_exchange_failed",
|
||||
status_code=token_res.status_code,
|
||||
token_error=token_error,
|
||||
)
|
||||
return _oauth_callback_error("token_exchange_failed", request)
|
||||
token_payload = token_res.json()
|
||||
id_token = token_payload.get("id_token")
|
||||
if not isinstance(id_token, str) or not id_token:
|
||||
return _frontend_login_redirect("id_token_missing", request)
|
||||
_log_oauth_callback_failure(request, "id_token_missing")
|
||||
return _oauth_callback_error("id_token_missing", request)
|
||||
|
||||
info_res = await client.get(GOOGLE_TOKENINFO_URL, params={"id_token": id_token})
|
||||
if info_res.status_code >= 400:
|
||||
return _frontend_login_redirect("id_token_invalid", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"id_token_invalid",
|
||||
status_code=info_res.status_code,
|
||||
)
|
||||
return _oauth_callback_error("id_token_invalid", request)
|
||||
claims = info_res.json()
|
||||
|
||||
if claims.get("aud") != settings.oauth_google_client_id:
|
||||
return _frontend_login_redirect("audience_mismatch", request)
|
||||
_log_oauth_callback_failure(request, "audience_mismatch")
|
||||
return _oauth_callback_error("audience_mismatch", request)
|
||||
issuer = claims.get("iss")
|
||||
if issuer not in {"accounts.google.com", "https://accounts.google.com"}:
|
||||
return _frontend_login_redirect("issuer_mismatch", request)
|
||||
_log_oauth_callback_failure(request, "issuer_mismatch", issuer=issuer)
|
||||
return _oauth_callback_error("issuer_mismatch", request)
|
||||
|
||||
try:
|
||||
email = validate_google_identity_domain(
|
||||
|
|
@ -525,21 +801,39 @@ async def callback(
|
|||
hosted_domain=claims.get("hd"),
|
||||
)
|
||||
except HTTPException:
|
||||
return _frontend_login_redirect("domain_not_allowed", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"domain_not_allowed",
|
||||
email_domain=_email_domain(str(claims.get("email") or "")),
|
||||
hosted_domain=_normalize_domain(str(claims.get("hd") or "")),
|
||||
)
|
||||
return _oauth_callback_error("domain_not_allowed", request)
|
||||
role = _role_for_email(email)
|
||||
display_name = str(claims.get("name") or email)
|
||||
cohort_ids = _configured_cohort_ids(
|
||||
email=email,
|
||||
hosted_domain=str(claims.get("hd") or ""),
|
||||
)
|
||||
external_id = _provider_external_id("google", str(claims.get("sub") or ""), email)
|
||||
try:
|
||||
sid, _ = await create_session(
|
||||
email=email,
|
||||
display_name=display_name,
|
||||
role=role.value,
|
||||
cohort_ids=[],
|
||||
cohort_ids=cohort_ids,
|
||||
external_id=external_id,
|
||||
)
|
||||
except InactiveUserError as exc:
|
||||
return _frontend_login_redirect("inactive_user", request)
|
||||
_log_oauth_callback_failure(
|
||||
request,
|
||||
"inactive_user",
|
||||
email_domain=_email_domain(email),
|
||||
)
|
||||
return _oauth_callback_error("inactive_user", request)
|
||||
|
||||
response = RedirectResponse(_frontend_url(stored.next_path, request), status_code=302)
|
||||
_set_session_cookie(response, sid)
|
||||
_delete_oauth_state_cookie(response)
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -580,12 +874,15 @@ async def saml_acs(request: Request) -> RedirectResponse:
|
|||
return _frontend_login_redirect("saml_assertion_invalid", request)
|
||||
|
||||
role = _role_for_saml_identity(identity)
|
||||
cohort_ids = _configured_cohort_ids(email=email, claim_hint=identity.cohort_hint)
|
||||
external_id = _provider_external_id("saml", identity.subject, email)
|
||||
try:
|
||||
sid, _ = await create_session(
|
||||
email=email,
|
||||
display_name=identity.display_name or email,
|
||||
role=role.value,
|
||||
cohort_ids=[],
|
||||
cohort_ids=cohort_ids,
|
||||
external_id=external_id,
|
||||
)
|
||||
except InactiveUserError:
|
||||
return _frontend_login_redirect("inactive_user", request)
|
||||
|
|
@ -615,7 +912,8 @@ async def dev_login(request: Request, body: DevLoginRequest, response: Response)
|
|||
email=email,
|
||||
display_name=body.display_name or email,
|
||||
role=body.role,
|
||||
cohort_ids=[],
|
||||
cohort_ids=_configured_cohort_ids(email=email),
|
||||
external_id=_provider_external_id("dev", email, email),
|
||||
)
|
||||
except InactiveUserError as exc:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="user is inactive") from exc
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue