현재 작업 전체 반영

This commit is contained in:
Yun Chan 2026-06-27 16:08:41 +09:00
parent 5560638e54
commit c0dddab594
85 changed files with 11322 additions and 539 deletions

View file

@ -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