mfa.py 103 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007100810091010101110121013101410151016101710181019102010211022102310241025102610271028102910301031103210331034103510361037103810391040104110421043104410451046104710481049105010511052105310541055105610571058105910601061106210631064106510661067106810691070107110721073107410751076107710781079108010811082108310841085108610871088108910901091109210931094109510961097109810991100110111021103110411051106110711081109111011111112111311141115111611171118111911201121112211231124112511261127112811291130113111321133113411351136113711381139114011411142114311441145114611471148114911501151115211531154115511561157115811591160116111621163116411651166116711681169117011711172117311741175117611771178117911801181118211831184118511861187118811891190119111921193119411951196119711981199120012011202120312041205120612071208120912101211121212131214121512161217121812191220122112221223122412251226122712281229123012311232123312341235123612371238123912401241124212431244124512461247124812491250125112521253125412551256125712581259126012611262126312641265126612671268126912701271127212731274127512761277127812791280128112821283128412851286128712881289129012911292129312941295129612971298129913001301130213031304130513061307130813091310131113121313131413151316131713181319132013211322132313241325132613271328132913301331133213331334133513361337133813391340134113421343134413451346134713481349135013511352135313541355135613571358135913601361136213631364136513661367136813691370137113721373137413751376137713781379138013811382138313841385138613871388138913901391139213931394139513961397139813991400140114021403140414051406140714081409141014111412141314141415141614171418141914201421142214231424142514261427142814291430143114321433143414351436143714381439144014411442144314441445144614471448144914501451145214531454145514561457145814591460146114621463146414651466146714681469147014711472147314741475147614771478147914801481148214831484148514861487148814891490149114921493149414951496149714981499150015011502150315041505150615071508150915101511151215131514151515161517151815191520152115221523152415251526152715281529153015311532153315341535153615371538153915401541154215431544154515461547154815491550155115521553155415551556155715581559156015611562156315641565156615671568156915701571157215731574157515761577157815791580158115821583158415851586158715881589159015911592159315941595159615971598159916001601160216031604160516061607160816091610161116121613161416151616161716181619162016211622162316241625162616271628162916301631163216331634163516361637163816391640164116421643164416451646164716481649165016511652165316541655165616571658165916601661166216631664166516661667166816691670167116721673167416751676167716781679168016811682168316841685168616871688168916901691169216931694169516961697169816991700170117021703170417051706170717081709171017111712171317141715171617171718171917201721172217231724172517261727172817291730173117321733173417351736173717381739174017411742174317441745174617471748174917501751175217531754175517561757175817591760176117621763176417651766176717681769177017711772177317741775177617771778177917801781178217831784178517861787178817891790179117921793179417951796179717981799180018011802180318041805180618071808180918101811181218131814181518161817181818191820182118221823182418251826182718281829183018311832183318341835183618371838183918401841184218431844184518461847184818491850185118521853185418551856185718581859186018611862186318641865186618671868186918701871187218731874187518761877187818791880188118821883188418851886188718881889189018911892189318941895189618971898189919001901190219031904190519061907190819091910191119121913191419151916191719181919192019211922192319241925192619271928192919301931193219331934193519361937193819391940194119421943194419451946194719481949195019511952195319541955195619571958195919601961196219631964196519661967196819691970197119721973197419751976197719781979198019811982198319841985198619871988198919901991199219931994199519961997199819992000200120022003200420052006200720082009201020112012201320142015201620172018201920202021202220232024202520262027202820292030203120322033203420352036203720382039204020412042204320442045204620472048204920502051205220532054205520562057205820592060206120622063206420652066206720682069207020712072207320742075207620772078207920802081208220832084208520862087208820892090209120922093209420952096209720982099210021012102210321042105210621072108210921102111211221132114211521162117211821192120212121222123212421252126212721282129213021312132213321342135213621372138213921402141214221432144214521462147214821492150215121522153215421552156215721582159216021612162216321642165216621672168216921702171217221732174217521762177217821792180218121822183218421852186218721882189219021912192219321942195219621972198219922002201220222032204220522062207220822092210221122122213221422152216221722182219222022212222222322242225222622272228222922302231223222332234223522362237223822392240224122422243224422452246224722482249225022512252225322542255225622572258225922602261226222632264226522662267226822692270227122722273227422752276227722782279228022812282228322842285228622872288228922902291229222932294229522962297229822992300230123022303230423052306230723082309231023112312231323142315231623172318231923202321
  1. """2FA (TOTP + Email OTP) and OIDC authentication routes.
  2. Security model
  3. --------------
  4. * Pre-auth tokens : secrets.token_urlsafe(32) stored in-memory with a 5-minute TTL.
  5. They are single-use and do NOT grant access to any protected resource.
  6. * TOTP codes : verified with pyotp (30-second window, ±1 step tolerance).
  7. * Email OTP codes : 6-digit numeric, hashed with pbkdf2_sha256, 10-minute TTL,
  8. max 5 failed attempts per code before invalidation.
  9. * Backup codes : 10 × 8-char alphanumeric codes, each stored as pbkdf2_sha256 hash,
  10. single-use.
  11. * OIDC state : secrets.token_urlsafe(32) bound to provider_id + nonce, 10-minute TTL.
  12. * OIDC exchange : secrets.token_urlsafe(32), 2-minute TTL, single-use.
  13. * Rate limiting : max 5 failed 2FA verification attempts per user within 15 minutes.
  14. """
  15. from __future__ import annotations
  16. import base64
  17. import hashlib
  18. import io
  19. import logging
  20. import os
  21. import re
  22. import secrets
  23. import string
  24. import urllib.parse
  25. from datetime import datetime, timedelta, timezone
  26. import httpx
  27. import jwt
  28. import pyotp
  29. from fastapi import APIRouter, Body, Depends, Header, HTTPException, Query, Request, Response, status
  30. from fastapi.responses import RedirectResponse
  31. from jwt import PyJWKClient
  32. from passlib.context import CryptContext
  33. from sqlalchemy import delete, select, update
  34. from sqlalchemy.ext.asyncio import AsyncSession
  35. from sqlalchemy.orm import selectinload, undefer
  36. from backend.app.api.routes._oidc_helpers import assert_safe_public_https_url
  37. from backend.app.api.routes.settings import get_setting, set_setting
  38. from backend.app.core.auth import (
  39. RequirePermissionIfAuthEnabled,
  40. create_access_token,
  41. get_current_active_user,
  42. get_user_by_email,
  43. get_user_by_username,
  44. is_auth_enabled,
  45. resolve_session_max_minutes,
  46. verify_password,
  47. )
  48. from backend.app.core.database import get_db
  49. from backend.app.core.permissions import Permission
  50. from backend.app.models.auth_ephemeral import AuthEphemeralToken, AuthRateLimitEvent, EventType, TokenType
  51. from backend.app.models.group import Group
  52. from backend.app.models.oidc_provider import OIDCProvider, UserOIDCLink
  53. from backend.app.models.user import User
  54. from backend.app.models.user_otp_code import UserOTPCode
  55. from backend.app.models.user_totp import UserTOTP
  56. from backend.app.schemas.auth import (
  57. AUTO_LINK_REQUIREMENTS_ERROR,
  58. AdminDisable2FARequest,
  59. BackupCodesResponse,
  60. EmailOTPDisableRequest,
  61. EmailOTPEnableConfirmRequest,
  62. EmailOTPSendRequest,
  63. GroupBrief,
  64. LoginResponse,
  65. OIDCAuthorizeResponse,
  66. OIDCExchangeRequest,
  67. OIDCLinkResponse,
  68. OIDCProviderCreate,
  69. OIDCProviderPublicResponse,
  70. OIDCProviderResponse,
  71. OIDCProviderUpdate,
  72. TOTPDisableRequest,
  73. TOTPEnableRequest,
  74. TOTPEnableResponse,
  75. TOTPSetupRequest,
  76. TOTPSetupResponse,
  77. TwoFAStatusResponse,
  78. TwoFAVerifyRequest,
  79. TwoFAVerifyResponse,
  80. UserResponse,
  81. )
  82. from backend.app.services.email_service import get_smtp_settings, send_email
  83. from backend.app.services.oidc_icon import OIDCIconError, fetch_icon
  84. logger = logging.getLogger(__name__)
  85. def _redact_url_for_log(url: str) -> str:
  86. """Return ``scheme://host/path`` with query string and fragment stripped.
  87. Admin-supplied icon URLs are usually CDN paths, but nothing stops an
  88. admin from pasting a presigned URL whose query string carries an
  89. ``X-Amz-Signature`` / OAuth token / etc. Operators need a forensic
  90. trail without those secrets ending up in log files.
  91. """
  92. try:
  93. parsed = urllib.parse.urlparse(url)
  94. except ValueError:
  95. return "<unparseable>"
  96. netloc = parsed.netloc or "<no-host>"
  97. return f"{parsed.scheme}://{netloc}{parsed.path}"
  98. async def _fetch_icon_or_400(icon_url: str) -> tuple[bytes, str, str]:
  99. """Validate URL + fetch icon, mapping any failure to HTTPException(400).
  100. Centralises the SSRF guard + fetcher invocation so create/update/refresh
  101. all behave identically — admin always gets a 400 with a precise reason,
  102. never a 500 / opaque server error.
  103. Both failure paths log at WARNING so operators have a forensic trail
  104. later — without these log lines the admin's UI toast was the only
  105. record of the failure (#1333 review).
  106. """
  107. try:
  108. assert_safe_public_https_url(icon_url)
  109. except ValueError as exc:
  110. logger.warning("OIDC icon URL rejected by SSRF guard: url=%s reason=%s", _redact_url_for_log(icon_url), exc)
  111. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
  112. try:
  113. return await fetch_icon(icon_url)
  114. except OIDCIconError as exc:
  115. logger.warning("OIDC icon fetch failed: url=%s reason=%s", _redact_url_for_log(icon_url), exc)
  116. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
  117. def _build_provider_response(provider: OIDCProvider) -> OIDCProviderResponse:
  118. """Build OIDCProviderResponse via ``from_attributes``. The required
  119. ``has_icon`` field is supplied by ``OIDCProvider.has_icon`` (a property
  120. reading the non-deferred ``icon_content_type`` column)."""
  121. return OIDCProviderResponse.model_validate(provider)
  122. def _etag_matches(if_none_match: str | None, etag_raw: str | None) -> bool:
  123. """RFC 7232 §3.2 If-None-Match comparison.
  124. Supports:
  125. * ``*`` wildcard — matches any current representation when the resource
  126. exists (and it does here; we wouldn't have an etag otherwise).
  127. * Multiple comma-separated tokens.
  128. * Weak-validator prefix ``W/`` (RFC 7232 §2.3) — accepted on GET since
  129. cached representations of a static byte-blob are byte-identical.
  130. Returns False on missing header or missing stored etag.
  131. """
  132. if not if_none_match or not etag_raw:
  133. return False
  134. quoted = f'"{etag_raw}"'
  135. tokens = [t.strip() for t in if_none_match.split(",")]
  136. if "*" in tokens:
  137. return True
  138. return any(tok.removeprefix("W/") == quoted for tok in tokens)
  139. def _as_utc(dt: datetime) -> datetime:
  140. """Return *dt* with UTC timezone attached.
  141. SQLite/aiosqlite strips timezone info when reading DateTime(timezone=True)
  142. columns back – the stored value is always UTC, so we just re-attach the
  143. info when doing Python-level comparisons.
  144. """
  145. return dt if dt.tzinfo is not None else dt.replace(tzinfo=timezone.utc)
  146. # ---------------------------------------------------------------------------
  147. # Passlib context (same scheme as auth.py)
  148. # ---------------------------------------------------------------------------
  149. pwd_context = CryptContext(schemes=["pbkdf2_sha256"], deprecated="auto")
  150. # ---------------------------------------------------------------------------
  151. # TTL / rate-limit constants
  152. # ---------------------------------------------------------------------------
  153. MAX_2FA_ATTEMPTS = 5
  154. MAX_LOGIN_ATTEMPTS = 10
  155. LOCKOUT_WINDOW = timedelta(minutes=15)
  156. MAX_EMAIL_OTP_SENDS = 3
  157. EMAIL_OTP_SEND_WINDOW = timedelta(minutes=10)
  158. PRE_AUTH_TOKEN_TTL = timedelta(minutes=5)
  159. OIDC_STATE_TTL = timedelta(minutes=10)
  160. OIDC_EXCHANGE_TTL = timedelta(minutes=2)
  161. # ---------------------------------------------------------------------------
  162. # Router
  163. # ---------------------------------------------------------------------------
  164. router = APIRouter(prefix="/auth", tags=["2fa", "oidc"])
  165. # ---------------------------------------------------------------------------
  166. # Helper: user response
  167. # ---------------------------------------------------------------------------
  168. def _user_to_response(user: User) -> UserResponse:
  169. return UserResponse(
  170. id=user.id,
  171. username=user.username,
  172. email=user.email,
  173. role=user.role,
  174. is_active=user.is_active,
  175. is_admin=user.is_admin,
  176. groups=[GroupBrief(id=g.id, name=g.name) for g in user.groups],
  177. permissions=sorted(user.get_permissions()),
  178. created_at=user.created_at.isoformat(),
  179. )
  180. # ---------------------------------------------------------------------------
  181. # Helper: QR code generation
  182. # ---------------------------------------------------------------------------
  183. def _generate_totp_qr_b64(provisioning_uri: str) -> str:
  184. """Generate a base64-encoded PNG QR code for the given TOTP provisioning URI."""
  185. import qrcode # type: ignore
  186. qr = qrcode.QRCode(box_size=6, border=2)
  187. qr.add_data(provisioning_uri)
  188. qr.make(fit=True)
  189. img = qr.make_image(fill_color="black", back_color="white")
  190. buf = io.BytesIO()
  191. img.save(buf, format="PNG")
  192. return base64.b64encode(buf.getvalue()).decode()
  193. # ---------------------------------------------------------------------------
  194. # Helper: backup code generation
  195. # ---------------------------------------------------------------------------
  196. def _generate_backup_codes() -> tuple[list[str], list[str]]:
  197. """Return (plain_codes, hashed_codes) — 10 codes of 8 alphanumeric chars each."""
  198. alphabet = string.ascii_uppercase + string.digits
  199. plain = ["".join(secrets.choice(alphabet) for _ in range(8)) for _ in range(10)]
  200. hashed = [pwd_context.hash(c) for c in plain]
  201. return plain, hashed
  202. # ---------------------------------------------------------------------------
  203. # DB-backed pre-auth token helpers
  204. # ---------------------------------------------------------------------------
  205. async def create_pre_auth_token(db: AsyncSession, username: str, challenge_id: str | None = None) -> str:
  206. """Create a single-use pre-auth token stored in the DB.
  207. Pass ``challenge_id`` (from the HttpOnly 2fa_challenge cookie) to bind the
  208. token to the originating browser session. The same value must be present as
  209. a cookie on every subsequent call that consumes this token.
  210. """
  211. now = datetime.now(timezone.utc)
  212. # Prune expired tokens opportunistically (keep table small)
  213. await db.execute(
  214. delete(AuthEphemeralToken).where(
  215. AuthEphemeralToken.token_type == TokenType.PRE_AUTH,
  216. AuthEphemeralToken.expires_at < now,
  217. )
  218. )
  219. token = secrets.token_urlsafe(32)
  220. db.add(
  221. AuthEphemeralToken(
  222. token=token,
  223. token_type=TokenType.PRE_AUTH,
  224. username=username,
  225. challenge_id=challenge_id,
  226. expires_at=now + PRE_AUTH_TOKEN_TTL,
  227. )
  228. )
  229. await db.commit()
  230. return token
  231. async def consume_pre_auth_token(db: AsyncSession, token: str, challenge_id: str | None = None) -> str | None:
  232. """Atomically validate and consume a pre-auth token. Returns username or None.
  233. Uses DELETE...RETURNING so two concurrent requests with the same token cannot
  234. both succeed — only the first DELETE finds the row.
  235. M5: When challenge_id is provided, also enforces the cookie-binding constraint
  236. so a stolen token cannot be replayed from a different browser session.
  237. """
  238. now = datetime.now(timezone.utc)
  239. result = await db.execute(
  240. delete(AuthEphemeralToken)
  241. .where(
  242. AuthEphemeralToken.token == token,
  243. AuthEphemeralToken.token_type == TokenType.PRE_AUTH,
  244. AuthEphemeralToken.expires_at > now,
  245. )
  246. .returning(AuthEphemeralToken.username, AuthEphemeralToken.challenge_id)
  247. )
  248. row = result.one_or_none()
  249. if row is None:
  250. return None
  251. username, stored_challenge_id = row
  252. # Enforce client binding: if the token was issued with a challenge_id,
  253. # the caller must supply the matching value.
  254. if stored_challenge_id is not None and stored_challenge_id != challenge_id:
  255. await db.rollback()
  256. return None
  257. await db.commit()
  258. return username
  259. async def peek_pre_auth_token(db: AsyncSession, token: str, challenge_id: str | None = None) -> str | None:
  260. """Validate a pre-auth token and return the username WITHOUT consuming it.
  261. When the stored token has a ``challenge_id`` (client-binding cookie), the
  262. caller must supply the matching value. A mismatch is treated as an invalid
  263. token — no information leakage about whether the token itself exists.
  264. """
  265. now = datetime.now(timezone.utc)
  266. result = await db.execute(
  267. select(AuthEphemeralToken).where(
  268. AuthEphemeralToken.token == token,
  269. AuthEphemeralToken.token_type == TokenType.PRE_AUTH,
  270. AuthEphemeralToken.expires_at > now,
  271. )
  272. )
  273. eph = result.scalar_one_or_none()
  274. if eph is None:
  275. return None
  276. # Enforce client binding: if the token was issued with a challenge_id the
  277. # cookie must match. Treat a mismatch as if the token doesn't exist.
  278. if eph.challenge_id is not None and eph.challenge_id != challenge_id:
  279. return None
  280. return eph.username
  281. # ---------------------------------------------------------------------------
  282. # DB-backed rate-limiting helpers
  283. # ---------------------------------------------------------------------------
  284. async def check_rate_limit(
  285. db: AsyncSession,
  286. username: str,
  287. event_type: str = EventType.TWO_FA_ATTEMPT,
  288. max_attempts: int = MAX_2FA_ATTEMPTS,
  289. ) -> None:
  290. """Raise HTTP 429 if the user has exceeded the failed attempt limit.
  291. The username is normalised to lower-case so case-variant attempts
  292. (which all resolve to the same user) share the same rate-limit bucket.
  293. L-2: Known TOCTOU — the SELECT (count) and the subsequent INSERT
  294. (record_failed_attempt) are not atomic. Two concurrent requests can both
  295. read a count below the threshold and both proceed. This is an inherent
  296. trade-off of the event-log rate-limit pattern: fixing it would require
  297. a serialising lock (SELECT FOR UPDATE on a dedicated counter row), which
  298. adds contention and is not worth it for a soft rate-limit whose window is
  299. already measured in minutes. In practice the race window is microseconds
  300. and the limit can be slightly exceeded only under precise concurrent timing.
  301. """
  302. username_key = username.lower()
  303. now = datetime.now(timezone.utc)
  304. cutoff = now - LOCKOUT_WINDOW
  305. result = await db.execute(
  306. select(AuthRateLimitEvent).where(
  307. AuthRateLimitEvent.username == username_key,
  308. AuthRateLimitEvent.event_type == event_type,
  309. AuthRateLimitEvent.occurred_at > cutoff,
  310. )
  311. )
  312. recent_count = len(result.scalars().all())
  313. if recent_count >= max_attempts:
  314. raise HTTPException(
  315. status_code=status.HTTP_429_TOO_MANY_REQUESTS,
  316. detail="Too many failed attempts. Please try again later.",
  317. )
  318. async def record_failed_attempt(db: AsyncSession, username: str, event_type: str = EventType.TWO_FA_ATTEMPT) -> None:
  319. """Record a failed attempt for rate-limiting purposes."""
  320. db.add(AuthRateLimitEvent(username=username.lower(), event_type=event_type))
  321. await db.commit()
  322. async def clear_failed_attempts(db: AsyncSession, username: str, event_type: str = EventType.TWO_FA_ATTEMPT) -> None:
  323. """Delete all recorded failed attempts for a user on successful verification."""
  324. await db.execute(
  325. delete(AuthRateLimitEvent).where(
  326. AuthRateLimitEvent.username == username.lower(),
  327. AuthRateLimitEvent.event_type == event_type,
  328. )
  329. )
  330. await db.commit()
  331. async def check_email_otp_send_rate(db: AsyncSession, username: str) -> None:
  332. """Raise HTTP 429 if the user has requested too many OTP emails recently.
  333. I1: This function only *checks* the limit. The caller is responsible for
  334. recording the slot via ``record_email_otp_send`` **after** the email has
  335. been sent successfully. This prevents failed sends from consuming a slot
  336. (wasting the user's quota) and makes it impossible to farm rate-limit events
  337. without actually triggering a send.
  338. """
  339. username_key = username.lower()
  340. now = datetime.now(timezone.utc)
  341. cutoff = now - EMAIL_OTP_SEND_WINDOW
  342. result = await db.execute(
  343. select(AuthRateLimitEvent).where(
  344. AuthRateLimitEvent.username == username_key,
  345. AuthRateLimitEvent.event_type == EventType.EMAIL_SEND,
  346. AuthRateLimitEvent.occurred_at > cutoff,
  347. )
  348. )
  349. recent_count = len(result.scalars().all())
  350. if recent_count >= MAX_EMAIL_OTP_SENDS:
  351. raise HTTPException(
  352. status_code=status.HTTP_429_TOO_MANY_REQUESTS,
  353. detail=f"Too many OTP email requests. Please wait {EMAIL_OTP_SEND_WINDOW.seconds // 60} minutes.",
  354. )
  355. async def record_email_otp_send(db: AsyncSession, username: str) -> None:
  356. """Record a successful OTP email send for rate-limiting purposes (I1).
  357. Must be called *after* the email has been sent successfully so that failed
  358. sends do not consume a slot from the user's quota.
  359. """
  360. db.add(AuthRateLimitEvent(username=username.lower(), event_type=EventType.EMAIL_SEND))
  361. await db.commit()
  362. # ---------------------------------------------------------------------------
  363. # TOTP replay-protection helper
  364. # ---------------------------------------------------------------------------
  365. def _assert_totp_not_replayed(totp_obj: pyotp.TOTP, totp_record: UserTOTP, code: str) -> None:
  366. """Raise HTTP 400 if this TOTP code was already accepted in its time window.
  367. M3 fix: store the counter of the *accepted* code rather than the current
  368. wall-clock counter. With valid_window=1, pyotp accepts codes from the
  369. previous 30-second step. Using timecode(now) would store the wrong counter
  370. when the previous-window code is accepted, allowing immediate replay.
  371. """
  372. # Determine which time-step the accepted code belongs to.
  373. now = datetime.now(timezone.utc)
  374. accepted_counter: int | None = None
  375. for offset in (0, -1): # current window first, then previous
  376. candidate_time = now.timestamp() + offset * totp_obj.interval
  377. candidate_counter = totp_obj.timecode(datetime.fromtimestamp(candidate_time, tz=timezone.utc))
  378. if totp_obj.at(candidate_counter) == code:
  379. accepted_counter = candidate_counter
  380. break
  381. if accepted_counter is None:
  382. accepted_counter = totp_obj.timecode(now) # fallback (should not happen after verify())
  383. totp_record.accept_counter(accepted_counter)
  384. # ---------------------------------------------------------------------------
  385. # OIDC helpers
  386. # ---------------------------------------------------------------------------
  387. _EMAIL_SHAPE_RE = re.compile(r"[^\s@]+@[^\s@]+\.[^\s@]+")
  388. def _is_valid_email_shaped(value: str | None) -> bool:
  389. # SEC-2: shape check for non-standard claims (upn, preferred_username).
  390. # Requires local@domain.tld — rejects "@", "x@", "@domain", "x@nodot".
  391. if not value or len(value) > 255:
  392. return False
  393. return _EMAIL_SHAPE_RE.fullmatch(value) is not None
  394. def _enforce_auto_link_safety(provider: OIDCProvider) -> None:
  395. """Raise HTTP 422 if auto_link_existing_accounts is on with an unsafe combined state.
  396. SEC-1: only Fall B (email_claim='email' + require_email_verified=False) is unsafe —
  397. an attacker-controlled IdP could present an unverified email that matches a local account.
  398. Fall C (custom claim) never performs an email_verified check, so auto_link is safe there.
  399. Called after ORM construction (create) and after the setattr loop (update).
  400. """
  401. if provider.auto_link_existing_accounts and provider.email_claim == "email" and not provider.require_email_verified:
  402. raise HTTPException(
  403. status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
  404. detail=AUTO_LINK_REQUIREMENTS_ERROR,
  405. )
  406. def _resolve_provider_email(provider: OIDCProvider, claims: dict, provider_sub: str) -> str | None:
  407. """Extract and normalise the email address from OIDC ID-token claims.
  408. Implements three resolution paths (Fall A/B/C):
  409. Fall C — custom email_claim (!= "email"): shape-check only, no email_verified gate.
  410. Recommended for Azure Entra ID (preferred_username or upn).
  411. Fall A — email_claim="email" + require_email_verified=True: strict, email_verified must be True.
  412. Fall B — email_claim="email" + require_email_verified=False: permissive, explicit False drops email.
  413. Returns a lowercase-stripped email string, or None when the claim is absent/invalid.
  414. """
  415. provider_id = provider.id
  416. raw_claim_value = claims.get(provider.email_claim)
  417. if raw_claim_value is not None and not isinstance(raw_claim_value, str):
  418. # TYPE-GUARD: non-string claim (e.g. list, int) would raise AttributeError on .lower().
  419. logger.warning(
  420. "OIDC provider %d: email_claim %r has unexpected type %s for sub=%r, ignoring",
  421. provider_id,
  422. provider.email_claim,
  423. type(raw_claim_value).__name__,
  424. provider_sub,
  425. )
  426. raw_claim_value = None
  427. raw_email: str | None = raw_claim_value.lower().strip() if raw_claim_value else None
  428. if provider.email_claim != "email":
  429. # Fall C: custom claim (preferred_username, upn, …) — no email_verified check.
  430. # SEC-2: _is_valid_email_shaped instead of bare '"@" in value'.
  431. # Recommended for Azure Entra ID: set email_claim="preferred_username" or "upn".
  432. if raw_email and _is_valid_email_shaped(raw_email):
  433. return raw_email
  434. if raw_email:
  435. logger.warning(
  436. "OIDC provider %d: email_claim %r value failed shape check for sub=%r, ignoring",
  437. provider_id,
  438. provider.email_claim,
  439. provider_sub,
  440. )
  441. return None
  442. email_verified = claims.get("email_verified")
  443. if provider.require_email_verified:
  444. # Fall A: standard C1-Guard — fail closed unless email_verified is True.
  445. # SEC-2: apply shape check to standard email claim — providers may set
  446. # email_verified=True on non-email values (e.g. numeric user IDs).
  447. # SEC-3 normalisation applies; existing mixed-case provider_email records
  448. # were normalised to lowercase by run_migrations at startup.
  449. if raw_email and not _is_valid_email_shaped(raw_email):
  450. logger.warning(
  451. "OIDC provider %d: email claim failed shape check for sub=%r, ignoring",
  452. provider_id,
  453. provider_sub,
  454. )
  455. return None
  456. if email_verified is True:
  457. return raw_email
  458. if raw_email:
  459. logger.info(
  460. "OIDC provider %d: ignoring email for sub=%r because email_verified=%r",
  461. provider_id,
  462. provider_sub,
  463. email_verified,
  464. )
  465. return None
  466. # Fall B: permissive — explicit False drops email, absent/None keeps it.
  467. # Required for Azure Entra ID which never sends email_verified.
  468. # SEC-2: apply shape check before the email_verified=False drop so malformed
  469. # values are rejected regardless of the email_verified claim.
  470. if raw_email and not _is_valid_email_shaped(raw_email):
  471. logger.warning(
  472. "OIDC provider %d: email claim failed shape check for sub=%r, ignoring",
  473. provider_id,
  474. provider_sub,
  475. )
  476. return None
  477. if email_verified is False:
  478. return None
  479. if email_verified is not True:
  480. # SEC-5: log only when the permissive path actually fires (ev absent/None),
  481. # not on every successful login.
  482. logger.info(
  483. "OIDC provider %r (%d): accepting email for sub=%r without email_verified claim (permissive mode)",
  484. provider.name,
  485. provider.id,
  486. provider_sub,
  487. )
  488. return raw_email
  489. def _resolve_standard_email_for_user_record(provider: OIDCProvider, claims: dict, provider_sub: str) -> str | None:
  490. """Resolve the standard 'email' claim for populating a newly-created User.email.
  491. Issue #1569: when an operator sets email_claim to a non-email identity claim
  492. (e.g. preferred_username on Authentik), the primary _resolve_provider_email
  493. returns None because the identity value isn't email-shaped. This helper lets
  494. the auto-create-users path still capture the user's real email from the
  495. standard 'email' claim that the IdP usually sends alongside.
  496. This is NOT a substitute for _resolve_provider_email and does NOT feed
  497. auto_link_existing_accounts — that gate stays on the primary resolver, so
  498. the GHSA Fall-B/C security guards remain intact.
  499. Applies the same Fall A/B shape + email_verified logic as the primary
  500. resolver does for the standard 'email' claim.
  501. """
  502. raw_claim_value = claims.get("email")
  503. if raw_claim_value is not None and not isinstance(raw_claim_value, str):
  504. logger.warning(
  505. "OIDC provider %d: standard 'email' claim has unexpected type %s for sub=%r, ignoring",
  506. provider.id,
  507. type(raw_claim_value).__name__,
  508. provider_sub,
  509. )
  510. return None
  511. raw_email = raw_claim_value.lower().strip() if raw_claim_value else None
  512. if not raw_email:
  513. return None
  514. if not _is_valid_email_shaped(raw_email):
  515. logger.warning(
  516. "OIDC provider %d: standard 'email' claim failed shape check for sub=%r, ignoring",
  517. provider.id,
  518. provider_sub,
  519. )
  520. return None
  521. email_verified = claims.get("email_verified")
  522. if provider.require_email_verified:
  523. if email_verified is True:
  524. return raw_email
  525. logger.info(
  526. "OIDC provider %d: ignoring fallback email for sub=%r because email_verified=%r",
  527. provider.id,
  528. provider_sub,
  529. email_verified,
  530. )
  531. return None
  532. if email_verified is False:
  533. return None
  534. return raw_email
  535. # ---------------------------------------------------------------------------
  536. # Settings helpers (email 2FA flag)
  537. # ---------------------------------------------------------------------------
  538. async def _get_email_2fa_enabled(db: AsyncSession, user_id: int) -> bool:
  539. val = await get_setting(db, f"user_{user_id}_email_2fa_enabled")
  540. return val == "true"
  541. async def _set_email_2fa_enabled(db: AsyncSession, user_id: int, enabled: bool) -> None:
  542. await set_setting(db, f"user_{user_id}_email_2fa_enabled", "true" if enabled else "false")
  543. # ===========================================================================
  544. # 2FA Endpoints
  545. # ===========================================================================
  546. @router.get("/2fa/status", response_model=TwoFAStatusResponse)
  547. async def get_2fa_status(
  548. current_user: User = Depends(get_current_active_user),
  549. db: AsyncSession = Depends(get_db),
  550. ) -> TwoFAStatusResponse:
  551. """Return the current 2FA configuration for the authenticated user."""
  552. result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == current_user.id))
  553. totp_record = result.scalar_one_or_none()
  554. totp_enabled = totp_record is not None and totp_record.is_enabled
  555. backup_codes_remaining = len(totp_record.backup_code_hashes) if totp_record else 0
  556. email_otp_enabled = await _get_email_2fa_enabled(db, current_user.id)
  557. return TwoFAStatusResponse(
  558. totp_enabled=totp_enabled,
  559. email_otp_enabled=email_otp_enabled,
  560. backup_codes_remaining=backup_codes_remaining,
  561. )
  562. @router.post("/2fa/totp/setup", response_model=TOTPSetupResponse)
  563. async def setup_totp(
  564. body: TOTPSetupRequest | None = Body(default=None),
  565. current_user: User = Depends(get_current_active_user),
  566. db: AsyncSession = Depends(get_db),
  567. ) -> TOTPSetupResponse:
  568. """Initiate TOTP setup: generates a new secret and QR code.
  569. Creates (or replaces) a pending UserTOTP record with is_enabled=False.
  570. The caller must confirm with POST /auth/2fa/totp/enable.
  571. M-R7-A: If an *active* TOTP is already configured, the caller must supply
  572. the current TOTP code in the request body to confirm intent before the
  573. secret is overwritten (prevents silently locking out the real user).
  574. """
  575. if not await is_auth_enabled(db):
  576. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Authentication is not enabled")
  577. # Upsert a pending TOTP record (is_enabled=False)
  578. existing = (await db.execute(select(UserTOTP).where(UserTOTP.user_id == current_user.id))).scalar_one_or_none()
  579. # M-R7-A: Guard against silent TOTP replacement when one is already active.
  580. if existing and existing.is_enabled:
  581. await check_rate_limit(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  582. supplied_code = (body.code if body else None) or ""
  583. # S4: narrow the RuntimeError catch to ONLY the property access — that
  584. # is the single line that raises on key-loss. The previous wide try
  585. # block also covered record_failed_attempt, clear_failed_attempts,
  586. # and _assert_totp_not_replayed, so a future RuntimeError from any
  587. # of those would have been misreported as "TOTP secret unavailable".
  588. try:
  589. secret_plain = existing.secret
  590. except RuntimeError:
  591. logger.exception("TOTP decryption failed for user_id=%s", current_user.id)
  592. raise HTTPException(
  593. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
  594. detail="TOTP secret unavailable",
  595. )
  596. if not pyotp.TOTP(secret_plain).verify(supplied_code, valid_window=1):
  597. await record_failed_attempt(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  598. raise HTTPException(
  599. status_code=status.HTTP_400_BAD_REQUEST,
  600. detail="Current TOTP code required to replace an active authenticator",
  601. )
  602. await clear_failed_attempts(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  603. _assert_totp_not_replayed(pyotp.TOTP(secret_plain), existing, supplied_code)
  604. await db.flush() # L-3: persist last_totp_counter immediately to block replay
  605. secret = pyotp.random_base32()
  606. totp = pyotp.TOTP(secret)
  607. provisioning_uri = totp.provisioning_uri(name=current_user.username, issuer_name="Bambuddy")
  608. qr_b64 = _generate_totp_qr_b64(provisioning_uri)
  609. if existing:
  610. existing.secret = secret
  611. existing.is_enabled = False
  612. existing.backup_code_hashes = []
  613. else:
  614. db.add(UserTOTP(user_id=current_user.id, secret=secret, is_enabled=False))
  615. await db.commit()
  616. return TOTPSetupResponse(secret=secret, qr_code_b64=qr_b64, issuer="Bambuddy")
  617. @router.post("/2fa/totp/enable", response_model=TOTPEnableResponse)
  618. async def enable_totp(
  619. body: TOTPEnableRequest,
  620. current_user: User = Depends(get_current_active_user),
  621. db: AsyncSession = Depends(get_db),
  622. ) -> TOTPEnableResponse:
  623. """Confirm TOTP setup by verifying a code from the authenticator app.
  624. On success, enables TOTP and returns 10 single-use backup codes (shown once).
  625. L-R7-A: Rate-limited to prevent brute-forcing the 6-digit confirmation code.
  626. """
  627. # L-R7-A: Rate-limit the enable step to prevent brute-forcing the 6-digit code.
  628. await check_rate_limit(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  629. result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == current_user.id))
  630. totp_record = result.scalar_one_or_none()
  631. if not totp_record:
  632. raise HTTPException(
  633. status_code=status.HTTP_400_BAD_REQUEST, detail="TOTP setup not initiated. Call /auth/2fa/totp/setup first."
  634. )
  635. try:
  636. totp_verify = pyotp.TOTP(totp_record.secret).verify(body.code, valid_window=1)
  637. except RuntimeError:
  638. logger.exception("TOTP decryption failed for user_id=%s", totp_record.user_id)
  639. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="TOTP secret unavailable")
  640. if not totp_verify:
  641. await record_failed_attempt(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  642. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid TOTP code")
  643. await clear_failed_attempts(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  644. plain_codes, hashed_codes = _generate_backup_codes()
  645. totp_record.is_enabled = True
  646. totp_record.backup_code_hashes = hashed_codes
  647. await db.commit()
  648. return TOTPEnableResponse(
  649. message="TOTP enabled successfully. Store your backup codes in a safe place.",
  650. backup_codes=plain_codes,
  651. )
  652. @router.post("/2fa/totp/disable")
  653. async def disable_totp(
  654. body: TOTPDisableRequest,
  655. current_user: User = Depends(get_current_active_user),
  656. db: AsyncSession = Depends(get_db),
  657. ) -> dict:
  658. """Disable TOTP by verifying a valid TOTP code or a backup code.
  659. I10: Rate-limited to prevent backup-code brute-forcing from a hijacked session.
  660. """
  661. await check_rate_limit(db, current_user.username)
  662. result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == current_user.id))
  663. totp_record = result.scalar_one_or_none()
  664. if not totp_record or not totp_record.is_enabled:
  665. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="TOTP is not enabled")
  666. # Accept either a valid TOTP code or a valid backup code. When the secret
  667. # cannot be decrypted (encryption key lost), fall through to the backup-
  668. # code path so the user can still disable 2FA with their printed codes.
  669. totp_obj: pyotp.TOTP | None = None
  670. code_valid = False
  671. decryption_failed = False
  672. try:
  673. totp_obj = pyotp.TOTP(totp_record.secret)
  674. code_valid = totp_obj.verify(body.code, valid_window=1)
  675. except RuntimeError:
  676. # S3: track that the failure was server-side so we don't penalise
  677. # the user with a fail-counter increment for a problem they can't fix.
  678. decryption_failed = True
  679. logger.exception(
  680. "TOTP decryption failed for user_id=%s — falling through to backup-code check",
  681. totp_record.user_id,
  682. )
  683. if code_valid and totp_obj is not None:
  684. _assert_totp_not_replayed(totp_obj, totp_record, body.code)
  685. await db.flush() # L-3: persist last_totp_counter immediately to block replay
  686. else:
  687. # Check backup codes — always iterate all entries (L-R9-A: no early break
  688. # to avoid timing oracle based on code position in the list).
  689. for hashed in totp_record.backup_code_hashes:
  690. if pwd_context.verify(body.code, hashed):
  691. code_valid = True
  692. if not code_valid:
  693. # S3: skip the fail-counter debit when the cause was a server-side
  694. # decryption failure (key loss / rotation). The user submitted a
  695. # wrong backup code on top of a broken TOTP, but locking them out
  696. # of the recovery path for an admin's mistake is not the right move.
  697. if not decryption_failed:
  698. await record_failed_attempt(db, current_user.username)
  699. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid code")
  700. await db.execute(delete(UserTOTP).where(UserTOTP.user_id == current_user.id))
  701. await db.commit()
  702. return {"message": "TOTP disabled"}
  703. @router.post("/2fa/totp/regenerate-backup-codes", response_model=BackupCodesResponse)
  704. async def regenerate_backup_codes(
  705. body: TOTPDisableRequest,
  706. current_user: User = Depends(get_current_active_user),
  707. db: AsyncSession = Depends(get_db),
  708. ) -> BackupCodesResponse:
  709. """Generate 10 new backup codes. Requires a valid TOTP code OR a backup code.
  710. M10: Accepts backup codes for consistency with disable_totp — users who have
  711. lost their authenticator app but still have backup codes can regenerate.
  712. Rate-limited to prevent brute-forcing from a hijacked session.
  713. """
  714. await check_rate_limit(db, current_user.username)
  715. result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == current_user.id))
  716. totp_record = result.scalar_one_or_none()
  717. if not totp_record or not totp_record.is_enabled:
  718. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="TOTP is not enabled")
  719. # Same recovery contract as disable_totp: when the TOTP secret cannot be
  720. # decrypted, fall through to the backup-code branch so the user can
  721. # rotate their codes with a printed backup code.
  722. totp_obj: pyotp.TOTP | None = None
  723. code_valid = False
  724. decryption_failed = False
  725. try:
  726. totp_obj = pyotp.TOTP(totp_record.secret)
  727. code_valid = totp_obj.verify(body.code, valid_window=1)
  728. except RuntimeError:
  729. # S3: track server-side failure so we skip the fail-counter debit.
  730. decryption_failed = True
  731. logger.exception(
  732. "TOTP decryption failed for user_id=%s — falling through to backup-code check",
  733. totp_record.user_id,
  734. )
  735. if code_valid and totp_obj is not None:
  736. _assert_totp_not_replayed(totp_obj, totp_record, body.code)
  737. await db.flush() # L-3: persist last_totp_counter immediately to block replay
  738. else:
  739. # Accept a backup code as an alternative (M10)
  740. matched_index: int | None = None
  741. for idx, hashed in enumerate(totp_record.backup_code_hashes):
  742. if pwd_context.verify(body.code, hashed) and matched_index is None:
  743. matched_index = idx
  744. if matched_index is None:
  745. # S3: skip fail-counter debit when the cause was a server-side
  746. # decryption failure (key loss / rotation).
  747. if not decryption_failed:
  748. await record_failed_attempt(db, current_user.username)
  749. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid TOTP or backup code")
  750. # Remove the used backup code
  751. totp_record.backup_code_hashes = [c for i, c in enumerate(totp_record.backup_code_hashes) if i != matched_index]
  752. plain_codes, hashed_codes = _generate_backup_codes()
  753. totp_record.backup_code_hashes = hashed_codes
  754. await db.commit()
  755. return BackupCodesResponse(
  756. backup_codes=plain_codes,
  757. message="Backup codes regenerated. Store them safely — they will not be shown again.",
  758. )
  759. @router.post("/2fa/email/enable")
  760. async def enable_email_otp(
  761. current_user: User = Depends(get_current_active_user),
  762. db: AsyncSession = Depends(get_db),
  763. ) -> dict:
  764. """Step 1 of email OTP enable: send a verification code to the user's email.
  765. C5: Proof of possession — the user must prove they control the registered email
  766. address before email 2FA is activated. Returns a ``setup_token`` that must be
  767. passed to POST /auth/2fa/email/enable/confirm together with the received code.
  768. H-3: Rate-limited to prevent email flooding via repeated calls to this endpoint.
  769. """
  770. await check_email_otp_send_rate(db, current_user.username)
  771. if not current_user.email:
  772. raise HTTPException(
  773. status_code=status.HTTP_400_BAD_REQUEST,
  774. detail="You must have an email address configured to enable email OTP 2FA",
  775. )
  776. smtp_settings = await get_smtp_settings(db)
  777. if not smtp_settings:
  778. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Email service is not configured")
  779. # Generate and store the setup token (reuse AuthEphemeralToken with type "email_otp_setup")
  780. now = datetime.now(timezone.utc)
  781. # Prune any existing pending setup tokens for this user
  782. await db.execute(
  783. delete(AuthEphemeralToken).where(
  784. AuthEphemeralToken.token_type == TokenType.EMAIL_OTP_SETUP,
  785. AuthEphemeralToken.username == current_user.username,
  786. )
  787. )
  788. code = str(secrets.randbelow(1_000_000)).zfill(6)
  789. code_hash = pwd_context.hash(code)
  790. setup_token = secrets.token_urlsafe(32)
  791. db.add(
  792. AuthEphemeralToken(
  793. token=setup_token,
  794. token_type=TokenType.EMAIL_OTP_SETUP,
  795. username=current_user.username,
  796. # Reuse the nonce field to store the code hash
  797. nonce=code_hash,
  798. expires_at=now + timedelta(minutes=10),
  799. )
  800. )
  801. await db.commit()
  802. try:
  803. send_email(
  804. smtp_settings=smtp_settings,
  805. to_email=current_user.email,
  806. subject="Verify your Bambuddy email address for 2FA",
  807. body_text=(
  808. f"Your Bambuddy email 2FA setup code is: {code}\n\n"
  809. "Enter this code to confirm email-based two-factor authentication.\n"
  810. "The code expires in 10 minutes."
  811. ),
  812. body_html=(
  813. "<p>To enable <strong>email-based two-factor authentication</strong> on your Bambuddy account, "
  814. "enter the code below:</p>"
  815. f"<h2 style='letter-spacing:4px'>{code}</h2>"
  816. "<p>The code expires in <strong>10 minutes</strong>. "
  817. "If you did not request this, you can safely ignore this email.</p>"
  818. ),
  819. )
  820. await record_email_otp_send(db, current_user.username)
  821. except Exception as exc:
  822. logger.error("Failed to send email OTP setup code to user_id=%d: %s", current_user.id, exc)
  823. raise HTTPException(
  824. status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to send verification email"
  825. )
  826. return {"message": "Verification code sent to your email address", "setup_token": setup_token}
  827. @router.post("/2fa/email/enable/confirm")
  828. async def confirm_enable_email_otp(
  829. body: EmailOTPEnableConfirmRequest,
  830. current_user: User = Depends(get_current_active_user),
  831. db: AsyncSession = Depends(get_db),
  832. ) -> dict:
  833. """Step 2 of email OTP enable: verify the code and activate email 2FA.
  834. H-2 fix: Uses peek-then-consume so a wrong code does NOT burn the setup token.
  835. The token is only deleted after successful code verification, allowing retries
  836. up to the rate limit (5 attempts / 15 min).
  837. M4: Rate-limited to prevent brute-forcing the 6-digit setup code.
  838. """
  839. await check_rate_limit(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  840. now = datetime.now(timezone.utc)
  841. # --- Peek: validate token without consuming ---
  842. peek_result = await db.execute(
  843. select(AuthEphemeralToken).where(
  844. AuthEphemeralToken.token == body.setup_token,
  845. AuthEphemeralToken.token_type == TokenType.EMAIL_OTP_SETUP,
  846. AuthEphemeralToken.username == current_user.username,
  847. AuthEphemeralToken.expires_at > now,
  848. )
  849. )
  850. eph = peek_result.scalar_one_or_none()
  851. if eph is None:
  852. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid or expired setup token")
  853. code_hash = eph.nonce # code hash stored in the nonce field
  854. # --- Verify code before consuming the token ---
  855. if not pwd_context.verify(body.code, code_hash):
  856. await record_failed_attempt(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  857. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid verification code")
  858. # --- Atomically consume the token now that the code is correct ---
  859. # DELETE...RETURNING prevents a concurrent request from using the same token.
  860. del_result = await db.execute(
  861. delete(AuthEphemeralToken)
  862. .where(
  863. AuthEphemeralToken.token == body.setup_token,
  864. AuthEphemeralToken.token_type == TokenType.EMAIL_OTP_SETUP,
  865. AuthEphemeralToken.username == current_user.username,
  866. )
  867. .returning(AuthEphemeralToken.id)
  868. )
  869. if del_result.one_or_none() is None:
  870. # Concurrent request consumed it between peek and delete — treat as invalid.
  871. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid or expired setup token")
  872. await clear_failed_attempts(db, current_user.username, event_type=EventType.TWO_FA_ATTEMPT)
  873. await _set_email_2fa_enabled(db, current_user.id, True)
  874. await db.commit()
  875. return {"message": "Email OTP 2FA enabled"}
  876. @router.post("/2fa/email/disable")
  877. async def disable_email_otp(
  878. body: EmailOTPDisableRequest,
  879. current_user: User = Depends(get_current_active_user),
  880. db: AsyncSession = Depends(get_db),
  881. ) -> dict:
  882. """Disable email-based OTP 2FA for the current user.
  883. C6: Re-authentication required — the caller must supply their account password
  884. to prevent a hijacked session from silently removing a second factor.
  885. LDAP/OIDC-only users (no local password) are exempt from this check.
  886. H-2: Rate-limited to prevent brute-forcing the password via this endpoint.
  887. """
  888. await check_rate_limit(db, current_user.username)
  889. if current_user.password_hash:
  890. if not verify_password(body.password, current_user.password_hash):
  891. await record_failed_attempt(db, current_user.username)
  892. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid password")
  893. await _set_email_2fa_enabled(db, current_user.id, False)
  894. await db.commit()
  895. return {"message": "Email OTP 2FA disabled"}
  896. @router.post("/2fa/email/send")
  897. async def send_email_otp(
  898. request: Request,
  899. body: EmailOTPSendRequest,
  900. db: AsyncSession = Depends(get_db),
  901. ) -> dict:
  902. """Send a 6-digit OTP code to the user's email address.
  903. Requires a valid pre_auth_token obtained during the login flow.
  904. """
  905. # Peek (validate without consuming) first so a rate-limit rejection does not
  906. # permanently burn the caller's pre-auth token.
  907. challenge_id = request.cookies.get("2fa_challenge")
  908. username = await peek_pre_auth_token(db, body.pre_auth_token, challenge_id=challenge_id)
  909. if not username:
  910. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired pre-auth token")
  911. # Enforce rate limit BEFORE consuming the token to prevent OTP email flooding.
  912. await check_email_otp_send_rate(db, username)
  913. user = await get_user_by_username(db, username)
  914. if not user or not user.is_active:
  915. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found or inactive")
  916. if not user.email:
  917. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="User has no email address configured")
  918. smtp_settings = await get_smtp_settings(db)
  919. if not smtp_settings:
  920. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Email service is not configured")
  921. # Invalidate all existing unused OTP codes for this user (staged, not yet committed)
  922. await db.execute(
  923. UserOTPCode.__table__.update() # type: ignore[attr-defined]
  924. .where(UserOTPCode.user_id == user.id)
  925. .where(UserOTPCode.used.is_(False))
  926. .values(used=True)
  927. )
  928. # Generate a 6-digit code and stage the record (not committed yet)
  929. code = str(secrets.randbelow(1_000_000)).zfill(6)
  930. code_hash = pwd_context.hash(code)
  931. expires_at = datetime.now(timezone.utc) + timedelta(minutes=UserOTPCode.OTP_TTL_MINUTES)
  932. otp_record = UserOTPCode(
  933. user_id=user.id,
  934. code_hash=code_hash,
  935. attempts=0,
  936. used=False,
  937. expires_at=expires_at,
  938. )
  939. db.add(otp_record)
  940. # M2: Send the email BEFORE consuming the pre-auth token.
  941. # If the send fails we raise an exception here; the session is uncommitted so
  942. # the OTP record is discarded and the original token remains valid for retry.
  943. try:
  944. send_email(
  945. smtp_settings=smtp_settings,
  946. to_email=user.email,
  947. subject="Your Bambuddy verification code",
  948. body_text=f"Your Bambuddy login code is: {code}\n\nThis code expires in {UserOTPCode.OTP_TTL_MINUTES} minutes and can only be used once.",
  949. body_html=(
  950. f"<p>Your <strong>Bambuddy</strong> login verification code is:</p>"
  951. f"<h2 style='letter-spacing:4px'>{code}</h2>"
  952. f"<p>This code expires in <strong>{UserOTPCode.OTP_TTL_MINUTES} minutes</strong> and can only be used once.</p>"
  953. f"<p>If you did not request this code, you can safely ignore this email.</p>"
  954. ),
  955. )
  956. await record_email_otp_send(db, username)
  957. except Exception as exc:
  958. logger.error("Failed to send OTP email to user_id=%d: %s", user.id, exc)
  959. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to send OTP email")
  960. # Email sent — now atomically consume the old token (this also commits the
  961. # staged OTP record) and issue a fresh token for the verify step.
  962. consumed = await consume_pre_auth_token(db, body.pre_auth_token, challenge_id=challenge_id)
  963. if not consumed:
  964. # Raced with another request or token just expired — treat as invalid.
  965. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired pre-auth token")
  966. # Re-issue a fresh pre-auth token bound to the same cookie so the binding
  967. # carries forward through the email → verify step.
  968. fresh_token = await create_pre_auth_token(db, username, challenge_id=challenge_id)
  969. # Return the fresh pre-auth token so the frontend can proceed to verify
  970. return {"message": "Code sent to your email address", "pre_auth_token": fresh_token}
  971. @router.post("/2fa/verify", response_model=TwoFAVerifyResponse)
  972. async def verify_2fa(
  973. request: Request,
  974. body: TwoFAVerifyRequest,
  975. db: AsyncSession = Depends(get_db),
  976. ) -> TwoFAVerifyResponse:
  977. """Verify a 2FA code and exchange the pre_auth_token for a full JWT.
  978. Accepted methods: ``totp``, ``email``, ``backup``.
  979. The pre_auth_token is NOT consumed on failed verification attempts so the
  980. user can retry without restarting the login flow. It is only consumed once
  981. verification succeeds, preventing token replay after success.
  982. """
  983. # Peek without consuming — bad codes must not burn the session token.
  984. # Pass the HttpOnly challenge cookie so the binding check is enforced.
  985. challenge_id = request.cookies.get("2fa_challenge")
  986. username = await peek_pre_auth_token(db, body.pre_auth_token, challenge_id=challenge_id)
  987. if not username:
  988. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired pre-auth token")
  989. await check_rate_limit(db, username)
  990. user = await get_user_by_username(db, username)
  991. if not user or not user.is_active:
  992. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found or inactive")
  993. method = body.method
  994. if method == "totp":
  995. result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == user.id))
  996. totp_record = result.scalar_one_or_none()
  997. if not totp_record or not totp_record.is_enabled:
  998. await record_failed_attempt(db, username)
  999. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="TOTP is not enabled for this user")
  1000. try:
  1001. totp_obj = pyotp.TOTP(totp_record.secret)
  1002. except RuntimeError:
  1003. logger.exception("TOTP decryption failed for user_id=%s", totp_record.user_id)
  1004. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="TOTP secret unavailable")
  1005. if not totp_obj.verify(body.code, valid_window=1):
  1006. await record_failed_attempt(db, username)
  1007. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid TOTP code")
  1008. _assert_totp_not_replayed(totp_obj, totp_record, body.code)
  1009. await db.flush() # L-3: persist last_totp_counter immediately to block replay
  1010. elif method == "email":
  1011. now = datetime.now(timezone.utc)
  1012. result = await db.execute(
  1013. select(UserOTPCode)
  1014. .where(UserOTPCode.user_id == user.id)
  1015. .where(UserOTPCode.used.is_(False))
  1016. .where(UserOTPCode.expires_at > now)
  1017. .order_by(UserOTPCode.created_at.desc())
  1018. )
  1019. otp_record = result.scalar_one_or_none()
  1020. if not otp_record:
  1021. await record_failed_attempt(db, username)
  1022. raise HTTPException(
  1023. status_code=status.HTTP_401_UNAUTHORIZED, detail="No valid OTP code found. Request a new one."
  1024. )
  1025. if otp_record.attempts >= UserOTPCode.MAX_ATTEMPTS:
  1026. otp_record.consume()
  1027. await db.commit()
  1028. await record_failed_attempt(db, username)
  1029. raise HTTPException(
  1030. status_code=status.HTTP_401_UNAUTHORIZED, detail="OTP code has been invalidated after too many attempts"
  1031. )
  1032. if not pwd_context.verify(body.code, otp_record.code_hash):
  1033. otp_record.attempts += 1
  1034. await db.commit()
  1035. await record_failed_attempt(db, username)
  1036. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid OTP code")
  1037. otp_record.consume()
  1038. await db.commit()
  1039. else: # method == "backup"
  1040. result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == user.id))
  1041. totp_record = result.scalar_one_or_none()
  1042. if not totp_record or not totp_record.is_enabled:
  1043. await record_failed_attempt(db, username)
  1044. raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="TOTP is not enabled for this user")
  1045. # Always iterate all codes — no early break (L-R9-A: constant iteration
  1046. # count prevents timing oracle based on used-code position in the list).
  1047. matched_index: int | None = None
  1048. for idx, hashed in enumerate(totp_record.backup_code_hashes):
  1049. if pwd_context.verify(body.code, hashed) and matched_index is None:
  1050. matched_index = idx
  1051. if matched_index is None:
  1052. await record_failed_attempt(db, username)
  1053. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid backup code")
  1054. # M1: Consume the pre-auth token FIRST (atomic single-use enforcement).
  1055. # Only if that succeeds do we remove the backup code — this prevents a race
  1056. # where two concurrent requests both pass code verification but only one
  1057. # should be granted a session.
  1058. consumed_username = await consume_pre_auth_token(db, body.pre_auth_token, challenge_id=challenge_id)
  1059. if not consumed_username:
  1060. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired pre-auth token")
  1061. # Remove the used backup code now that the token is atomically consumed.
  1062. updated_codes = [c for i, c in enumerate(totp_record.backup_code_hashes) if i != matched_index]
  1063. totp_record.backup_code_hashes = updated_codes
  1064. await db.commit()
  1065. await clear_failed_attempts(db, username)
  1066. access_token = create_access_token(
  1067. data={"sub": user.username},
  1068. expires_delta=timedelta(minutes=await resolve_session_max_minutes(db)),
  1069. )
  1070. result = await db.execute(select(User).where(User.id == user.id).options(selectinload(User.groups)))
  1071. user = result.scalar_one()
  1072. return TwoFAVerifyResponse(access_token=access_token, token_type="bearer", user=_user_to_response(user))
  1073. # Verification succeeded (TOTP or email) — consume the pre-auth token.
  1074. # C-1: Check the return value; if None the token was already consumed by a
  1075. # concurrent request (race condition) — reject to prevent double-use.
  1076. consumed_username = await consume_pre_auth_token(db, body.pre_auth_token, challenge_id=challenge_id)
  1077. if not consumed_username:
  1078. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired pre-auth token")
  1079. await clear_failed_attempts(db, username)
  1080. access_token = create_access_token(
  1081. data={"sub": user.username},
  1082. expires_delta=timedelta(minutes=await resolve_session_max_minutes(db)),
  1083. )
  1084. # Reload with groups for permission calculation
  1085. result = await db.execute(select(User).where(User.id == user.id).options(selectinload(User.groups)))
  1086. user = result.scalar_one()
  1087. return TwoFAVerifyResponse(
  1088. access_token=access_token,
  1089. token_type="bearer",
  1090. user=_user_to_response(user),
  1091. )
  1092. @router.delete("/2fa/admin/{user_id}")
  1093. async def admin_disable_2fa(
  1094. user_id: int,
  1095. body: AdminDisable2FARequest = Body(default_factory=AdminDisable2FARequest),
  1096. current_user: User | None = RequirePermissionIfAuthEnabled(Permission.USERS_UPDATE),
  1097. db: AsyncSession = Depends(get_db),
  1098. ) -> dict:
  1099. """Admin endpoint: disable all 2FA for a given user.
  1100. Nit 3: Requires the admin's own password as a re-auth step (matching how
  1101. disable_email_otp protects a user's own 2FA removal). OIDC/LDAP-only admins
  1102. (no local password_hash) are exempt.
  1103. """
  1104. # Nit 3: Re-auth — admin must supply their own password.
  1105. if current_user and current_user.password_hash:
  1106. if not body.admin_password or not verify_password(body.admin_password, current_user.password_hash):
  1107. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Admin password required")
  1108. # Delete TOTP record
  1109. await db.execute(delete(UserTOTP).where(UserTOTP.user_id == user_id))
  1110. # Disable email 2FA setting
  1111. await _set_email_2fa_enabled(db, user_id, False)
  1112. # Invalidate all OTP codes
  1113. await db.execute(
  1114. UserOTPCode.__table__.update() # type: ignore[attr-defined]
  1115. .where(UserOTPCode.user_id == user_id)
  1116. .values(used=True)
  1117. )
  1118. # I2: Invalidate existing JWTs for the target user by bumping password_changed_at.
  1119. # Without this, a stolen token remains valid after 2FA removal.
  1120. target_user = (await db.execute(select(User).where(User.id == user_id))).scalar_one_or_none()
  1121. if target_user:
  1122. target_user.password_changed_at = datetime.now(timezone.utc)
  1123. await db.commit()
  1124. actor = current_user.username if current_user else "anonymous"
  1125. logger.info("Admin %s disabled all 2FA for user_id=%d", actor, user_id)
  1126. return {"message": "2FA disabled for user"}
  1127. # ===========================================================================
  1128. # OIDC Endpoints
  1129. # ===========================================================================
  1130. @router.get("/oidc/providers", response_model=list[OIDCProviderPublicResponse])
  1131. async def list_oidc_providers(
  1132. db: AsyncSession = Depends(get_db),
  1133. ) -> list[OIDCProviderPublicResponse]:
  1134. """List all enabled OIDC providers (public).
  1135. The login page renders icons via /oidc/providers/{id}/icon — `icon_data`
  1136. stays deferred so this list query never pulls the BLOB.
  1137. #3107: returns the slim public shape only. The login page needs id, name,
  1138. has_icon and is_autologin; the full response (scopes, claims, and now
  1139. group_claim / group_mapping) is served by the permission-gated
  1140. /oidc/providers/all below, so an unauthenticated caller cannot learn
  1141. which IdP group name maps to which Bambuddy group.
  1142. """
  1143. result = await db.execute(select(OIDCProvider).where(OIDCProvider.is_enabled.is_(True)))
  1144. providers = result.scalars().all()
  1145. return [OIDCProviderPublicResponse.model_validate(p) for p in providers]
  1146. @router.get("/oidc/providers/all", response_model=list[OIDCProviderResponse])
  1147. async def list_all_oidc_providers(
  1148. _: User | None = RequirePermissionIfAuthEnabled(Permission.SETTINGS_READ),
  1149. db: AsyncSession = Depends(get_db),
  1150. ) -> list[OIDCProviderResponse]:
  1151. """List ALL OIDC providers including disabled ones (admin only)."""
  1152. result2 = await db.execute(select(OIDCProvider))
  1153. providers = result2.scalars().all()
  1154. return [_build_provider_response(p) for p in providers]
  1155. @router.post("/oidc/providers", response_model=OIDCProviderResponse, status_code=status.HTTP_201_CREATED)
  1156. async def create_oidc_provider(
  1157. body: OIDCProviderCreate,
  1158. _: User | None = RequirePermissionIfAuthEnabled(Permission.SETTINGS_UPDATE),
  1159. db: AsyncSession = Depends(get_db),
  1160. ) -> OIDCProviderResponse:
  1161. """Create a new OIDC provider (admin only).
  1162. If `icon_url` is supplied, the icon is fetched server-side and cached in
  1163. the BLOB columns (#1333). A fetch failure aborts the create with 400 —
  1164. no half-configured provider is left in the DB.
  1165. """
  1166. if body.default_group_id is not None:
  1167. grp_chk = await db.execute(select(Group).where(Group.id == body.default_group_id))
  1168. if not grp_chk.scalar_one_or_none():
  1169. raise HTTPException(
  1170. status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
  1171. detail="default_group_id references a non-existent group",
  1172. )
  1173. # #3107 — every mapping value must reference an existing Bambuddy group,
  1174. # same answer default_group_id gets. Checked as a set: a mapping with three
  1175. # entries naming the same group is one lookup, not three.
  1176. if body.group_mapping:
  1177. missing_groups = await _missing_group_names(db, set(body.group_mapping.values()))
  1178. if missing_groups:
  1179. raise HTTPException(
  1180. status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
  1181. detail=f"group_mapping references non-existent groups: {', '.join(sorted(missing_groups))}",
  1182. )
  1183. # Fetch the icon BEFORE creating the row so a failure leaves the DB clean.
  1184. icon_data: bytes | None = None
  1185. icon_content_type: str | None = None
  1186. icon_etag: str | None = None
  1187. if body.icon_url:
  1188. icon_data, icon_content_type, icon_etag = await _fetch_icon_or_400(body.icon_url)
  1189. provider = OIDCProvider(
  1190. name=body.name,
  1191. issuer_url=body.issuer_url.rstrip("/"),
  1192. client_id=body.client_id,
  1193. client_secret=body.client_secret,
  1194. scopes=body.scopes,
  1195. is_enabled=body.is_enabled,
  1196. auto_create_users=body.auto_create_users,
  1197. auto_link_existing_accounts=body.auto_link_existing_accounts,
  1198. email_claim=body.email_claim,
  1199. require_email_verified=body.require_email_verified,
  1200. group_claim=body.group_claim,
  1201. group_mapping=body.group_mapping,
  1202. icon_url=body.icon_url,
  1203. icon_data=icon_data,
  1204. icon_content_type=icon_content_type,
  1205. icon_etag=icon_etag,
  1206. default_group_id=body.default_group_id,
  1207. is_autologin=body.is_autologin,
  1208. )
  1209. # SEC-1 + SEC-6: runtime guard mirrors the OIDCProviderCreate model_validator in schemas/auth.py.
  1210. # Catches any future path that bypasses Pydantic validation (direct ORM, scripts).
  1211. _enforce_auto_link_safety(provider)
  1212. db.add(provider)
  1213. # #1589: at most one provider may be the autologin target. When a new one
  1214. # is created with the flag set, clear it on all others first so the
  1215. # session still satisfies the invariant after add.
  1216. if body.is_autologin:
  1217. await db.execute(update(OIDCProvider).where(OIDCProvider.is_autologin.is_(True)).values(is_autologin=False))
  1218. await db.commit()
  1219. await db.refresh(provider)
  1220. return _build_provider_response(provider)
  1221. def _refuse_if_env_managed(provider: OIDCProvider) -> None:
  1222. """Startup rewrites this provider from BAMBUDDY_OIDC_* on every boot, so an
  1223. edit here would be accepted and then silently reverted at the next restart.
  1224. BAMBUDDY_LOCAL_LOGIN (#1589) remains the recovery path if it becomes
  1225. unusable, so refusing outright cannot lock anyone out."""
  1226. if provider.is_env_managed:
  1227. raise HTTPException(
  1228. status_code=status.HTTP_409_CONFLICT,
  1229. detail="This OIDC provider is managed by environment variables and cannot be modified.",
  1230. )
  1231. async def _missing_group_names(db: AsyncSession, names: set[str]) -> set[str]:
  1232. """#3107 — names from a group_mapping with no matching Bambuddy group.
  1233. Exact match only, same as the sync's own ``Group.name.in_()`` lookup and
  1234. the env path's ``Group.name == target``: the sync resolves names exactly,
  1235. so admitting a case-variant here would pass validation only to have the
  1236. sync silently never grant it — the exact failure this check exists to
  1237. catch at save time, in front of the admin, instead of at login time.
  1238. """
  1239. if not names:
  1240. return set()
  1241. exact = set((await db.execute(select(Group.name).where(Group.name.in_(names)))).scalars().all())
  1242. return names - exact
  1243. @router.put("/oidc/providers/{provider_id}", response_model=OIDCProviderResponse)
  1244. async def update_oidc_provider(
  1245. provider_id: int,
  1246. body: OIDCProviderUpdate,
  1247. _: User | None = RequirePermissionIfAuthEnabled(Permission.SETTINGS_UPDATE),
  1248. db: AsyncSession = Depends(get_db),
  1249. ) -> OIDCProviderResponse:
  1250. """Update an existing OIDC provider (admin only).
  1251. Icon refetch fires when:
  1252. 1. The submitted `icon_url` differs from the stored one (URL changed), OR
  1253. 2. The submitted `icon_url` equals the stored one AND `icon_content_type`
  1254. is NULL — this is the upgrade-path edge case: old providers carry
  1255. `icon_url` but no cached bytes until the admin first saves them.
  1256. On fetch failure the request aborts with 400 *before* commit, so the
  1257. existing cached bytes (if any) remain untouched.
  1258. """
  1259. result2 = await db.execute(select(OIDCProvider).where(OIDCProvider.id == provider_id))
  1260. provider = result2.scalar_one_or_none()
  1261. if not provider:
  1262. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Provider not found")
  1263. _refuse_if_env_managed(provider)
  1264. if body.default_group_id is not None:
  1265. grp_chk = await db.execute(select(Group).where(Group.id == body.default_group_id))
  1266. if not grp_chk.scalar_one_or_none():
  1267. raise HTTPException(
  1268. status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
  1269. detail="default_group_id references a non-existent group",
  1270. )
  1271. # #3107 — same existence check as the create route, on the submitted
  1272. # mapping only. A null group_mapping (field absent) leaves the stored one
  1273. # alone; an explicit {} empties it, and emptiness needs no group lookup.
  1274. if body.group_mapping:
  1275. missing_groups = await _missing_group_names(db, set(body.group_mapping.values()))
  1276. if missing_groups:
  1277. raise HTTPException(
  1278. status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
  1279. detail=f"group_mapping references non-existent groups: {', '.join(sorted(missing_groups))}",
  1280. )
  1281. dumped = body.model_dump(exclude_none=True)
  1282. # Decide whether an icon refetch is needed BEFORE mutating the ORM object,
  1283. # so the comparison sees provider.icon_url / icon_content_type as they are
  1284. # in the database.
  1285. new_icon_url = dumped.get("icon_url")
  1286. needs_icon_refetch = new_icon_url is not None and (
  1287. new_icon_url != provider.icon_url or provider.icon_content_type is None
  1288. )
  1289. # Fetch FIRST. If the upstream is unreachable or SSRF-blocked, _fetch_icon_or_400
  1290. # raises HTTPException(400) here — provider attributes are still untouched, so
  1291. # the in-memory ORM object stays consistent on the way out (and the DB row is
  1292. # safe regardless via get_db()'s rollback).
  1293. fetched_icon: tuple[bytes, str, str] | None = None
  1294. if needs_icon_refetch:
  1295. fetched_icon = await _fetch_icon_or_400(new_icon_url)
  1296. # Explicit `icon_url: null` in the PUT body means "clear the icon".
  1297. # The exclude_none=True dump above drops None values, which would
  1298. # otherwise silently ignore this request. Check model_fields_set on
  1299. # the unfiltered body to distinguish "client cleared it" from "client
  1300. # didn't include this field at all".
  1301. if "icon_url" in body.model_fields_set and body.icon_url is None:
  1302. provider.icon_url = None
  1303. provider.icon_data = None
  1304. provider.icon_content_type = None
  1305. provider.icon_etag = None
  1306. for field, value in dumped.items():
  1307. if field == "issuer_url" and value:
  1308. value = value.rstrip("/")
  1309. setattr(provider, field, value)
  1310. if fetched_icon is not None:
  1311. provider.icon_data, provider.icon_content_type, provider.icon_etag = fetched_icon
  1312. # SEC-1 + SEC-6: Combined-State-Guard after setattr loop.
  1313. # Checks the final in-memory state (DB values + newly set values combined) to catch
  1314. # partial updates that each pass schema validation individually but are unsafe together.
  1315. _enforce_auto_link_safety(provider)
  1316. # #1589: at most one provider may be the autologin target. Clear the flag
  1317. # on every other provider when this one becomes the autologin. Excludes
  1318. # the current row so SQLAlchemy doesn't fight our in-memory set above.
  1319. if body.is_autologin is True:
  1320. await db.execute(
  1321. update(OIDCProvider)
  1322. .where(OIDCProvider.id != provider.id, OIDCProvider.is_autologin.is_(True))
  1323. .values(is_autologin=False)
  1324. )
  1325. await db.commit()
  1326. await db.refresh(provider)
  1327. return _build_provider_response(provider)
  1328. @router.delete("/oidc/providers/{provider_id}")
  1329. async def delete_oidc_provider(
  1330. provider_id: int,
  1331. _: User | None = RequirePermissionIfAuthEnabled(Permission.SETTINGS_UPDATE),
  1332. db: AsyncSession = Depends(get_db),
  1333. ) -> dict:
  1334. """Delete an OIDC provider and all its user links (admin only)."""
  1335. result2 = await db.execute(select(OIDCProvider).where(OIDCProvider.id == provider_id))
  1336. provider = result2.scalar_one_or_none()
  1337. if not provider:
  1338. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Provider not found")
  1339. _refuse_if_env_managed(provider)
  1340. await db.delete(provider)
  1341. await db.commit()
  1342. return {"message": "Provider deleted"}
  1343. # ---------------------------------------------------------------------------
  1344. # OIDC provider icon proxy (#1333)
  1345. # ---------------------------------------------------------------------------
  1346. @router.get("/oidc/providers/{provider_id}/icon")
  1347. async def get_oidc_provider_icon(
  1348. provider_id: int,
  1349. if_none_match: str | None = Header(default=None, alias="If-None-Match"),
  1350. db: AsyncSession = Depends(get_db),
  1351. ) -> Response:
  1352. """Serve the cached icon for an enabled OIDC provider (public, no auth).
  1353. Unauthenticated because ``<img>`` tags cannot send Authorization headers
  1354. and the login page renders these icons before the user is signed in — the
  1355. same justification as ``/api/v1/makerworld/thumbnail``. The SSRF guard
  1356. runs at admin-config time (create/update/refresh), not here.
  1357. Disabled providers respond 404 to avoid leaking their existence to
  1358. anonymous callers (mirrors ``GET /oidc/providers`` which filters on
  1359. ``is_enabled``).
  1360. """
  1361. result = await db.execute(
  1362. select(OIDCProvider)
  1363. .options(undefer(OIDCProvider.icon_data))
  1364. .where(OIDCProvider.id == provider_id, OIDCProvider.is_enabled.is_(True))
  1365. )
  1366. provider = result.scalar_one_or_none()
  1367. if provider is None or provider.icon_content_type is None or provider.icon_data is None:
  1368. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Icon not found")
  1369. etag_value = f'"{provider.icon_etag}"'
  1370. cache_headers = {"ETag": etag_value, "Cache-Control": "public, max-age=3600"}
  1371. if _etag_matches(if_none_match, provider.icon_etag):
  1372. return Response(status_code=status.HTTP_304_NOT_MODIFIED, headers=cache_headers)
  1373. return Response(
  1374. content=provider.icon_data,
  1375. media_type=provider.icon_content_type,
  1376. headers=cache_headers,
  1377. )
  1378. @router.delete("/oidc/providers/{provider_id}/icon", status_code=status.HTTP_204_NO_CONTENT)
  1379. async def delete_oidc_provider_icon(
  1380. provider_id: int,
  1381. _: User | None = RequirePermissionIfAuthEnabled(Permission.SETTINGS_UPDATE),
  1382. db: AsyncSession = Depends(get_db),
  1383. ) -> Response:
  1384. """Remove the icon entirely for a provider (admin only).
  1385. Clears all four icon columns — ``icon_url`` plus the three cached-bytes
  1386. columns. "Remove icon" means the whole record is gone, not just the
  1387. cache; without this the admin form would still show the URL while
  1388. the login page rendered a blank fallback (confusing half-state).
  1389. To re-add an icon the admin re-types the URL in the edit form.
  1390. """
  1391. result = await db.execute(select(OIDCProvider).where(OIDCProvider.id == provider_id))
  1392. provider = result.scalar_one_or_none()
  1393. if provider is None:
  1394. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Provider not found")
  1395. _refuse_if_env_managed(provider)
  1396. # Setting deferred columns is safe — no read happens, just a write.
  1397. provider.icon_url = None
  1398. provider.icon_data = None
  1399. provider.icon_content_type = None
  1400. provider.icon_etag = None
  1401. await db.commit()
  1402. return Response(status_code=status.HTTP_204_NO_CONTENT)
  1403. @router.post("/oidc/providers/{provider_id}/icon/refresh", response_model=OIDCProviderResponse)
  1404. async def refresh_oidc_provider_icon(
  1405. provider_id: int,
  1406. _: User | None = RequirePermissionIfAuthEnabled(Permission.SETTINGS_UPDATE),
  1407. db: AsyncSession = Depends(get_db),
  1408. ) -> OIDCProviderResponse:
  1409. """Refetch the icon from the stored `icon_url` (admin only).
  1410. Used when:
  1411. - The IdP changed its icon and the admin wants Bambuddy to pick up the
  1412. new bytes.
  1413. - An upgrade left the provider with an `icon_url` but no cached bytes
  1414. (covered automatically by `update_oidc_provider` too, but this gives
  1415. the UI an explicit "Refresh" button).
  1416. Failure to refetch returns 400 *before* commit, so the previously cached
  1417. bytes survive intact.
  1418. """
  1419. result = await db.execute(select(OIDCProvider).where(OIDCProvider.id == provider_id))
  1420. provider = result.scalar_one_or_none()
  1421. if provider is None:
  1422. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Provider not found")
  1423. _refuse_if_env_managed(provider)
  1424. if not provider.icon_url:
  1425. raise HTTPException(
  1426. status_code=status.HTTP_400_BAD_REQUEST,
  1427. detail="Provider has no icon_url to refresh",
  1428. )
  1429. icon_data, icon_content_type, icon_etag = await _fetch_icon_or_400(provider.icon_url)
  1430. provider.icon_data = icon_data
  1431. provider.icon_content_type = icon_content_type
  1432. provider.icon_etag = icon_etag
  1433. await db.commit()
  1434. await db.refresh(provider)
  1435. return _build_provider_response(provider)
  1436. @router.get("/oidc/authorize/{provider_id}", response_model=OIDCAuthorizeResponse)
  1437. async def oidc_authorize(
  1438. provider_id: int,
  1439. db: AsyncSession = Depends(get_db),
  1440. ) -> OIDCAuthorizeResponse:
  1441. """Return the OIDC authorization URL for the given provider."""
  1442. result = await db.execute(
  1443. select(OIDCProvider).where(OIDCProvider.id == provider_id).where(OIDCProvider.is_enabled.is_(True))
  1444. )
  1445. provider = result.scalar_one_or_none()
  1446. if not provider:
  1447. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Provider not found or not enabled")
  1448. # Fetch discovery document
  1449. discovery_url = f"{provider.issuer_url.rstrip('/')}/.well-known/openid-configuration"
  1450. try:
  1451. async with httpx.AsyncClient(timeout=10) as client:
  1452. resp = await client.get(discovery_url)
  1453. resp.raise_for_status()
  1454. discovery = resp.json()
  1455. except Exception as exc:
  1456. logger.error("Failed to fetch OIDC discovery for provider %d: %s", provider_id, exc)
  1457. raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="Failed to fetch OIDC discovery document")
  1458. authorization_endpoint = discovery.get("authorization_endpoint")
  1459. if not authorization_endpoint:
  1460. raise HTTPException(
  1461. status_code=status.HTTP_502_BAD_GATEWAY, detail="OIDC discovery document missing authorization_endpoint"
  1462. )
  1463. # B2: SSRF guard — reject non-HTTP(S) schemes in the authorization endpoint
  1464. if not authorization_endpoint.startswith(("https://", "http://")):
  1465. logger.warning("OIDC discovery authorization_endpoint has invalid scheme: %s", authorization_endpoint)
  1466. raise HTTPException(
  1467. status_code=status.HTTP_502_BAD_GATEWAY,
  1468. detail="OIDC discovery document contains invalid authorization_endpoint",
  1469. )
  1470. external_url = await _get_base_external_url(db)
  1471. redirect_uri = f"{external_url}/api/v1/auth/oidc/callback"
  1472. now = datetime.now(timezone.utc)
  1473. # Prune expired OIDC states from the DB
  1474. await db.execute(
  1475. delete(AuthEphemeralToken).where(
  1476. AuthEphemeralToken.token_type == TokenType.OIDC_STATE,
  1477. AuthEphemeralToken.expires_at < now,
  1478. )
  1479. )
  1480. state = secrets.token_urlsafe(32)
  1481. nonce = secrets.token_urlsafe(32)
  1482. # PKCE (S256) – required by PocketID and recommended for all OIDC flows
  1483. code_verifier = secrets.token_urlsafe(48) # 64-char URL-safe string
  1484. code_challenge = base64.urlsafe_b64encode(hashlib.sha256(code_verifier.encode()).digest()).rstrip(b"=").decode()
  1485. db.add(
  1486. AuthEphemeralToken(
  1487. token=state,
  1488. token_type=TokenType.OIDC_STATE,
  1489. provider_id=provider_id,
  1490. nonce=nonce,
  1491. code_verifier=code_verifier,
  1492. expires_at=now + OIDC_STATE_TTL,
  1493. )
  1494. )
  1495. await db.commit()
  1496. params = urllib.parse.urlencode(
  1497. {
  1498. "response_type": "code",
  1499. "client_id": provider.client_id,
  1500. "redirect_uri": redirect_uri,
  1501. "scope": provider.scopes,
  1502. "state": state,
  1503. "nonce": nonce,
  1504. "code_challenge": code_challenge,
  1505. "code_challenge_method": "S256",
  1506. }
  1507. )
  1508. auth_url = f"{authorization_endpoint}?{params}"
  1509. return OIDCAuthorizeResponse(auth_url=auth_url)
  1510. @router.get("/oidc/callback")
  1511. async def oidc_callback(
  1512. code: str | None = Query(default=None, max_length=2048),
  1513. state: str | None = Query(default=None, max_length=2048),
  1514. error: str | None = Query(default=None, max_length=256),
  1515. db: AsyncSession = Depends(get_db),
  1516. ) -> RedirectResponse:
  1517. """Handle the OIDC authorization code callback from the identity provider."""
  1518. external_url = await _get_base_external_url(db)
  1519. frontend_error_url = f"{external_url}/?oidc_error="
  1520. try:
  1521. if error:
  1522. logger.warning("OIDC callback received error: %s", error)
  1523. return RedirectResponse(url=f"{frontend_error_url}oidc_provider_error", status_code=302)
  1524. if not code or not state:
  1525. return RedirectResponse(url=f"{frontend_error_url}missing_parameters", status_code=302)
  1526. # Atomically validate and consume OIDC state from DB (I6: single-use enforcement).
  1527. # DELETE...RETURNING ensures concurrent callbacks with the same state token
  1528. # cannot both succeed — only the first DELETE finds the row.
  1529. now = datetime.now(timezone.utc)
  1530. state_del = await db.execute(
  1531. delete(AuthEphemeralToken)
  1532. .where(
  1533. AuthEphemeralToken.token == state,
  1534. AuthEphemeralToken.token_type == TokenType.OIDC_STATE,
  1535. AuthEphemeralToken.expires_at > now, # reject expired tokens atomically
  1536. )
  1537. .returning(
  1538. AuthEphemeralToken.provider_id,
  1539. AuthEphemeralToken.nonce,
  1540. AuthEphemeralToken.code_verifier,
  1541. )
  1542. )
  1543. state_row = state_del.one_or_none()
  1544. if state_row is None:
  1545. await db.rollback()
  1546. return RedirectResponse(url=f"{frontend_error_url}invalid_state", status_code=302)
  1547. provider_id, nonce, code_verifier = state_row
  1548. await db.commit()
  1549. # Load provider
  1550. result = await db.execute(select(OIDCProvider).where(OIDCProvider.id == provider_id))
  1551. provider = result.scalar_one_or_none()
  1552. if not provider:
  1553. return RedirectResponse(url=f"{frontend_error_url}provider_not_found", status_code=302)
  1554. redirect_uri = f"{external_url}/api/v1/auth/oidc/callback"
  1555. # ── Step 1: Fetch discovery document ────────────────────────────────
  1556. discovery_url = f"{provider.issuer_url.rstrip('/')}/.well-known/openid-configuration"
  1557. try:
  1558. async with httpx.AsyncClient(timeout=10) as client:
  1559. disc_resp = await client.get(discovery_url)
  1560. disc_resp.raise_for_status()
  1561. discovery = disc_resp.json()
  1562. except Exception as exc:
  1563. logger.error("OIDC discovery fetch failed for provider %d: %s", provider_id, exc)
  1564. return RedirectResponse(url=f"{frontend_error_url}discovery_failed", status_code=302)
  1565. token_endpoint = discovery.get("token_endpoint")
  1566. jwks_uri = discovery.get("jwks_uri")
  1567. if not token_endpoint or not jwks_uri:
  1568. return RedirectResponse(url=f"{frontend_error_url}invalid_discovery_document", status_code=302)
  1569. # L-R7-C: Reject non-HTTP(S) URLs in the discovery document to prevent
  1570. # SSRF via crafted responses (e.g. file://, gopher://, internal schemes).
  1571. if not token_endpoint.startswith(("https://", "http://")) or not jwks_uri.startswith(("https://", "http://")):
  1572. logger.warning(
  1573. "OIDC discovery document contains non-HTTP URL(s): token=%s jwks=%s", token_endpoint, jwks_uri
  1574. )
  1575. return RedirectResponse(url=f"{frontend_error_url}invalid_discovery_document", status_code=302)
  1576. # ── Step 2: Exchange authorization code for tokens ───────────────────
  1577. token_form: dict[str, str] = {
  1578. "grant_type": "authorization_code",
  1579. "code": code,
  1580. "redirect_uri": redirect_uri,
  1581. "client_id": provider.client_id,
  1582. }
  1583. if provider.client_secret:
  1584. token_form["client_secret"] = provider.client_secret
  1585. if code_verifier:
  1586. token_form["code_verifier"] = code_verifier
  1587. try:
  1588. async with httpx.AsyncClient(timeout=15) as client:
  1589. token_resp = await client.post(
  1590. token_endpoint,
  1591. data=token_form,
  1592. headers={"Accept": "application/json"},
  1593. )
  1594. except Exception as exc:
  1595. logger.error("OIDC token exchange request failed for provider %d: %s", provider_id, exc)
  1596. return RedirectResponse(url=f"{frontend_error_url}token_exchange_network_error", status_code=302)
  1597. if not token_resp.is_success:
  1598. try:
  1599. err_body = token_resp.json()
  1600. oidc_err = err_body.get("error", "")
  1601. oidc_desc = err_body.get("error_description", "")
  1602. except Exception:
  1603. oidc_err = ""
  1604. oidc_desc = token_resp.text[:200]
  1605. logger.error(
  1606. "OIDC token exchange HTTP %d for provider %d. redirect_uri=%r error=%r desc=%r",
  1607. token_resp.status_code,
  1608. provider_id,
  1609. redirect_uri,
  1610. oidc_err,
  1611. oidc_desc,
  1612. )
  1613. # Encode the OIDC error code into the redirect so the user sees it in the toast.
  1614. # URL-encode the value to prevent query-parameter injection from provider responses.
  1615. raw_err = oidc_err[:40] if oidc_err else str(token_resp.status_code)
  1616. safe_err = urllib.parse.quote(raw_err, safe="")
  1617. return RedirectResponse(
  1618. url=f"{frontend_error_url}token_exchange_{safe_err}",
  1619. status_code=302,
  1620. )
  1621. try:
  1622. token_data = token_resp.json()
  1623. except Exception as exc:
  1624. logger.error("OIDC token exchange non-JSON response for provider %d: %s", provider_id, exc)
  1625. return RedirectResponse(url=f"{frontend_error_url}token_exchange_bad_response", status_code=302)
  1626. id_token = token_data.get("id_token")
  1627. if not id_token:
  1628. # Only log the keys present — values may contain secrets (access_token, etc.)
  1629. logger.error(
  1630. "OIDC token response missing id_token for provider %d; keys present: %s",
  1631. provider_id,
  1632. list(token_data.keys()),
  1633. )
  1634. return RedirectResponse(url=f"{frontend_error_url}no_id_token", status_code=302)
  1635. # ── Step 3: Fetch JWKS and validate ID token ─────────────────────────
  1636. # Use the issuer from the discovery document as the canonical value (OIDC Core
  1637. # §3.1.3.7 requires iss == discovery issuer exactly). We strip trailing slashes
  1638. # from both sides because some providers (e.g. Authentik, older PocketID versions)
  1639. # are inconsistent between the discovery issuer and the JWT iss claim.
  1640. discovery_issuer: str = discovery.get("issuer", provider.issuer_url).rstrip("/")
  1641. try:
  1642. async with httpx.AsyncClient(timeout=10) as jwks_http:
  1643. jwks_resp = await jwks_http.get(jwks_uri)
  1644. jwks_resp.raise_for_status()
  1645. jwks_data = jwks_resp.json()
  1646. jwks_client = PyJWKClient(jwks_uri)
  1647. jwks_client.fetch_data = lambda: jwks_data # type: ignore[method-assign]
  1648. signing_key = jwks_client.get_signing_key_from_jwt(id_token)
  1649. # M-3: Decode without built-in issuer check, then compare normalised
  1650. # (both sides rstrip("/")) to handle providers like Authentik that include
  1651. # a trailing slash in iss but not in the discovery issuer, or vice-versa.
  1652. claims = jwt.decode(
  1653. id_token,
  1654. signing_key.key,
  1655. algorithms=["RS256", "ES256", "RS384", "ES384", "RS512"],
  1656. audience=provider.client_id,
  1657. options={"verify_iss": False},
  1658. )
  1659. token_iss = claims.get("iss", "").rstrip("/")
  1660. if token_iss != discovery_issuer:
  1661. raise jwt.exceptions.InvalidIssuerError("Invalid issuer")
  1662. except Exception as exc:
  1663. logger.error("OIDC JWT validation failed for provider %d: %s", provider_id, exc, exc_info=True)
  1664. return RedirectResponse(url=f"{frontend_error_url}token_validation_failed", status_code=302)
  1665. # Verify nonce — fail closed: we always send a nonce, so the provider must echo it.
  1666. # Skipping the check when nonce is absent would allow CSRF on non-nonce providers.
  1667. token_nonce = claims.get("nonce")
  1668. if token_nonce is None or token_nonce != nonce:
  1669. logger.warning("OIDC nonce mismatch for provider %d (present=%r)", provider_id, token_nonce is not None)
  1670. return RedirectResponse(url=f"{frontend_error_url}nonce_mismatch", status_code=302)
  1671. provider_sub: str = claims.get("sub", "")
  1672. if not provider_sub:
  1673. return RedirectResponse(url=f"{frontend_error_url}missing_sub_claim", status_code=302)
  1674. # SEC-3: resolve email via Fall A/B/C logic (see _resolve_provider_email).
  1675. provider_email = _resolve_provider_email(provider, claims, provider_sub)
  1676. # ── Step 4: Resolve / create user ────────────────────────────────────
  1677. try:
  1678. # 1. Look up existing OIDC link
  1679. link_result = await db.execute(
  1680. select(UserOIDCLink)
  1681. .where(UserOIDCLink.provider_id == provider_id)
  1682. .where(UserOIDCLink.provider_user_id == provider_sub)
  1683. )
  1684. link = link_result.scalar_one_or_none()
  1685. user: User | None = None
  1686. if link:
  1687. # Existing link → load the linked user
  1688. user_result = await db.execute(
  1689. select(User).where(User.id == link.user_id).options(selectinload(User.groups))
  1690. )
  1691. user = user_result.scalar_one_or_none()
  1692. else:
  1693. # 2. No OIDC link yet — check for an existing user with the same email.
  1694. # Use case-insensitive matching (func.lower) so that "User@Example.com"
  1695. # and "user@example.com" are treated as the same identity, preventing
  1696. # an attacker-controlled IdP from bypassing the auto-link guard by
  1697. # registering the target email with different casing.
  1698. email_user: User | None = None
  1699. if provider_email:
  1700. email_user = await get_user_by_email(db, provider_email)
  1701. if email_user and provider.auto_link_existing_accounts:
  1702. # M-4: Only auto-link when the provider has auto_link_existing_accounts
  1703. # enabled. Operators can disable this to require explicit account linking,
  1704. # preventing an attacker-controlled IdP from hijacking local accounts.
  1705. #
  1706. # M-NEW-6: Refuse auto-link if the target user already has any OIDC
  1707. # link (to any provider). Without this guard an attacker who controls
  1708. # a second OIDC provider with auto_link enabled could add themselves as
  1709. # a second IdP for a user that already authenticates via a legitimate
  1710. # provider, effectively taking over the account.
  1711. existing_links_result = await db.execute(
  1712. select(UserOIDCLink).where(UserOIDCLink.user_id == email_user.id)
  1713. )
  1714. has_existing_oidc_link = existing_links_result.scalar_one_or_none() is not None
  1715. if has_existing_oidc_link:
  1716. logger.warning(
  1717. "Auto-link rejected for user '%s': already linked to another OIDC provider",
  1718. email_user.username,
  1719. )
  1720. return RedirectResponse(url=f"{frontend_error_url}no_linked_account", status_code=302)
  1721. db.add(
  1722. UserOIDCLink(
  1723. user_id=email_user.id,
  1724. provider_id=provider_id,
  1725. provider_user_id=provider_sub,
  1726. provider_email=provider_email,
  1727. )
  1728. )
  1729. await db.commit()
  1730. user = email_user
  1731. logger.info(
  1732. "Auto-linked existing user '%s' to OIDC provider %d via email match",
  1733. email_user.username,
  1734. provider_id,
  1735. )
  1736. elif provider.auto_create_users:
  1737. # 3. No existing user — create one
  1738. if provider_email:
  1739. raw = provider_email.split("@")[0]
  1740. else:
  1741. # Prefer a human-readable IdP claim over the opaque sub.
  1742. # isinstance guards are required: claims may carry non-string
  1743. # values (e.g. a list) that would break .strip().
  1744. # Sanitization is applied per-candidate so that a value that
  1745. # strips to empty (e.g. "!!!") correctly falls through to the
  1746. # next candidate rather than silently becoming "oidcuser".
  1747. _pref = claims.get("preferred_username")
  1748. _name = claims.get("name")
  1749. raw = ""
  1750. if isinstance(_pref, str):
  1751. raw = re.sub(r"[^a-zA-Z0-9._-]", "", _pref.strip())[:30]
  1752. if not raw and isinstance(_name, str):
  1753. raw = re.sub(r"[^a-zA-Z0-9._-]", "", _name.strip())[:30]
  1754. if not raw:
  1755. raw = provider_sub[:30]
  1756. candidate = re.sub(r"[^a-zA-Z0-9._-]", "", raw)[:30] or "oidcuser"
  1757. # Issue #1569: when email_claim is configured to a non-email
  1758. # identity claim (e.g. preferred_username on Authentik), the
  1759. # primary resolver returns None for the email field because the
  1760. # identity value isn't email-shaped. Fall back to the standard
  1761. # 'email' claim for User.email so the operator can split
  1762. # username-from-preferred_username and email-from-email.
  1763. # The auto-link gate above stays on provider_email, so the
  1764. # GHSA Fall-B/C guards remain intact.
  1765. user_email_for_storage = provider_email
  1766. if user_email_for_storage is None and provider.email_claim != "email":
  1767. user_email_for_storage = _resolve_standard_email_for_user_record(provider, claims, provider_sub)
  1768. username = candidate
  1769. counter = 1
  1770. while True:
  1771. existing = await get_user_by_username(db, username)
  1772. if not existing:
  1773. break
  1774. username = f"{candidate}{counter}"
  1775. counter += 1
  1776. # I9: Assign new OIDC users to a group before flush — accessing
  1777. # new_user.groups after a flush triggers a lazy-load which fails
  1778. # in async context. Resolution order:
  1779. # 1. provider.default_group_id (operator-configured)
  1780. # 2. "Viewers" (system fallback for read-only access)
  1781. # 3. no group (last resort if Viewers was deleted)
  1782. # SQLite does not enforce ON DELETE SET NULL, so a dangling
  1783. # default_group_id returns None here and falls through to Viewers.
  1784. default_group: Group | None = None
  1785. if provider.default_group_id is not None:
  1786. dg_result = await db.execute(select(Group).where(Group.id == provider.default_group_id))
  1787. default_group = dg_result.scalar_one_or_none()
  1788. if default_group is None:
  1789. viewers_result = await db.execute(select(Group).where(Group.name == "Viewers"))
  1790. default_group = viewers_result.scalar_one_or_none()
  1791. new_user = User(
  1792. username=username,
  1793. email=user_email_for_storage,
  1794. # M-1: auth_source="oidc" prevents local password-reset flow
  1795. # for users who should only authenticate via OIDC.
  1796. auth_source="oidc",
  1797. password_hash=None, # OIDC users never use password auth
  1798. role="user",
  1799. is_active=True,
  1800. groups=[default_group] if default_group else [],
  1801. )
  1802. db.add(new_user)
  1803. await db.flush()
  1804. db.add(
  1805. UserOIDCLink(
  1806. user_id=new_user.id,
  1807. provider_id=provider_id,
  1808. provider_user_id=provider_sub,
  1809. provider_email=user_email_for_storage,
  1810. )
  1811. )
  1812. await db.commit()
  1813. user_result = await db.execute(
  1814. select(User).where(User.id == new_user.id).options(selectinload(User.groups))
  1815. )
  1816. user = user_result.scalar_one()
  1817. logger.info("Auto-created user '%s' via OIDC provider %d", username, provider_id)
  1818. else:
  1819. return RedirectResponse(url=f"{frontend_error_url}no_linked_account", status_code=302)
  1820. if not user or not user.is_active:
  1821. return RedirectResponse(url=f"{frontend_error_url}account_inactive", status_code=302)
  1822. # #3107 — apply the provider's group mapping on every login, not
  1823. # just at account creation. Same managed-slice contract as the
  1824. # LDAP sync (#1292): only groups named in group_mapping values are
  1825. # touched, manual assignments elsewhere survive. A sync failure is
  1826. # logged inside and never blocks the login.
  1827. if provider.group_mapping:
  1828. from backend.app.services.oidc_group_sync import sync_oidc_user_groups
  1829. await sync_oidc_user_groups(
  1830. db,
  1831. user,
  1832. group_claim=provider.group_claim,
  1833. group_mapping=provider.group_mapping,
  1834. claims=claims,
  1835. )
  1836. # Post-rollback guard, not token freshness: nothing below reads
  1837. # user.groups (the callback only puts the username on the exchange
  1838. # token; /oidc/exchange re-selects the user with selectinload).
  1839. # ALL attributes, not just groups: a sync that raised called
  1840. # db.rollback(), which expires every loaded object in the
  1841. # session — the next username=user.username below would then
  1842. # lazy-load and raise MissingGreenlet, landing the user on
  1843. # ?oidc_error=user_resolution_failed, i.e. a failed sync would
  1844. # block the login after all. Inside the mapping branch so
  1845. # sync-off installs do not pay the query.
  1846. await db.refresh(user)
  1847. # Issue an OIDC exchange token (short-lived, single-use) stored in DB.
  1848. # I7: Opportunistically prune expired exchange tokens to keep the table small.
  1849. now2 = datetime.now(timezone.utc)
  1850. await db.execute(
  1851. delete(AuthEphemeralToken).where(
  1852. AuthEphemeralToken.token_type == TokenType.OIDC_EXCHANGE,
  1853. AuthEphemeralToken.expires_at < now2,
  1854. )
  1855. )
  1856. exchange_token = secrets.token_urlsafe(32)
  1857. db.add(
  1858. AuthEphemeralToken(
  1859. token=exchange_token,
  1860. token_type=TokenType.OIDC_EXCHANGE,
  1861. username=user.username,
  1862. expires_at=now2 + OIDC_EXCHANGE_TTL,
  1863. )
  1864. )
  1865. await db.commit()
  1866. # H-4: Use a URL fragment (#) instead of a query parameter so the exchange
  1867. # token is never sent to the server in the Referer header or server logs.
  1868. return RedirectResponse(url=f"{external_url}/login#oidc_token={exchange_token}", status_code=302)
  1869. except Exception as exc:
  1870. logger.error("OIDC user resolution failed for provider %d: %s", provider_id, exc, exc_info=True)
  1871. try:
  1872. await db.rollback()
  1873. except Exception as rb_exc:
  1874. logger.error("DB rollback failed after OIDC user-resolution error: %s", rb_exc, exc_info=True)
  1875. return RedirectResponse(url=f"{frontend_error_url}user_resolution_failed", status_code=302)
  1876. except Exception as exc:
  1877. # L-1: Log the exception class name internally but never expose it in the
  1878. # redirect URL — leaking exception names aids attacker reconnaissance.
  1879. logger.error("Unexpected error in OIDC callback (%s): %s", type(exc).__name__, exc, exc_info=True)
  1880. try:
  1881. return RedirectResponse(url=f"{frontend_error_url}internal_error", status_code=302)
  1882. except Exception as redirect_exc:
  1883. logger.error("Failed to construct error redirect in OIDC callback: %s", redirect_exc, exc_info=True)
  1884. raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="OIDC callback failed")
  1885. @router.post("/oidc/exchange", response_model=LoginResponse)
  1886. async def oidc_exchange(
  1887. body: OIDCExchangeRequest,
  1888. raw_request: Request,
  1889. response: Response,
  1890. db: AsyncSession = Depends(get_db),
  1891. ) -> LoginResponse:
  1892. """Exchange an OIDC exchange token (from the callback redirect) for a full JWT.
  1893. C4: If the resolved user has 2FA enabled the exchange returns a pre_auth_token
  1894. (requires_2fa=True) instead of a full JWT. The frontend must then complete the
  1895. 2FA step exactly as it would after a password-based login.
  1896. """
  1897. now = datetime.now(timezone.utc)
  1898. # Atomically consume the exchange token (DELETE...RETURNING prevents replay).
  1899. consume_result = await db.execute(
  1900. delete(AuthEphemeralToken)
  1901. .where(
  1902. AuthEphemeralToken.token == body.oidc_token,
  1903. AuthEphemeralToken.token_type == TokenType.OIDC_EXCHANGE,
  1904. AuthEphemeralToken.expires_at > now, # reject expired tokens atomically
  1905. )
  1906. .returning(AuthEphemeralToken.username)
  1907. )
  1908. row = consume_result.one_or_none()
  1909. if row is None:
  1910. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid or expired OIDC exchange token")
  1911. (username,) = row
  1912. await db.commit()
  1913. user = await get_user_by_username(db, username)
  1914. if not user or not user.is_active:
  1915. raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="User not found or inactive")
  1916. # Reload with groups
  1917. result = await db.execute(select(User).where(User.id == user.id).options(selectinload(User.groups)))
  1918. user = result.scalar_one()
  1919. # C4: Check whether the user has any 2FA method enabled.
  1920. totp_result = await db.execute(select(UserTOTP).where(UserTOTP.user_id == user.id))
  1921. totp_record = totp_result.scalar_one_or_none()
  1922. totp_enabled = totp_record is not None and totp_record.is_enabled
  1923. email_2fa_enabled = await _get_email_2fa_enabled(db, user.id)
  1924. if totp_enabled or email_2fa_enabled:
  1925. # User has 2FA — issue a pre_auth_token bound to this browser session via
  1926. # an HttpOnly cookie (H-A: mirrors the cookie-binding done in auth.py:login).
  1927. two_fa_methods: list[str] = []
  1928. if totp_enabled:
  1929. two_fa_methods.append("totp")
  1930. if email_2fa_enabled:
  1931. two_fa_methods.append("email")
  1932. if totp_enabled:
  1933. two_fa_methods.append("backup")
  1934. challenge_id = secrets.token_urlsafe(32)
  1935. pre_auth_token = await create_pre_auth_token(db, user.username, challenge_id=challenge_id)
  1936. response.set_cookie(
  1937. key="2fa_challenge",
  1938. value=challenge_id,
  1939. httponly=True,
  1940. secure=raw_request.url.scheme == "https",
  1941. samesite="lax",
  1942. max_age=300,
  1943. path="/api/v1/auth/2fa",
  1944. )
  1945. return LoginResponse(
  1946. requires_2fa=True,
  1947. pre_auth_token=pre_auth_token,
  1948. two_fa_methods=two_fa_methods,
  1949. user=_user_to_response(user),
  1950. )
  1951. access_token = create_access_token(
  1952. data={"sub": user.username},
  1953. expires_delta=timedelta(minutes=await resolve_session_max_minutes(db)),
  1954. )
  1955. return LoginResponse(
  1956. access_token=access_token,
  1957. token_type="bearer",
  1958. user=_user_to_response(user),
  1959. requires_2fa=False,
  1960. )
  1961. @router.get("/oidc/links", response_model=list[OIDCLinkResponse])
  1962. async def list_oidc_links(
  1963. current_user: User = Depends(get_current_active_user),
  1964. db: AsyncSession = Depends(get_db),
  1965. ) -> list[OIDCLinkResponse]:
  1966. """List all OIDC provider links for the current user."""
  1967. result = await db.execute(
  1968. select(UserOIDCLink).where(UserOIDCLink.user_id == current_user.id).options(selectinload(UserOIDCLink.provider))
  1969. )
  1970. links = result.scalars().all()
  1971. # Defensive null-check on link.provider: on PostgreSQL the FK cascade
  1972. # ensures provider exists, but SQLite ships with FK enforcement off, so
  1973. # a deleted provider could in theory leave the link briefly orphan until
  1974. # the next init_db() cleanup runs. Returning "<deleted>" instead of
  1975. # crashing keeps the endpoint usable in that edge case (#1285 follow-up).
  1976. return [
  1977. OIDCLinkResponse(
  1978. id=link.id,
  1979. provider_id=link.provider_id,
  1980. provider_name=link.provider.name if link.provider else "<deleted>",
  1981. provider_email=link.provider_email,
  1982. created_at=link.created_at.isoformat(),
  1983. )
  1984. for link in links
  1985. ]
  1986. @router.delete("/oidc/links/{provider_id}")
  1987. async def remove_oidc_link(
  1988. provider_id: int,
  1989. current_user: User = Depends(get_current_active_user),
  1990. db: AsyncSession = Depends(get_db),
  1991. ) -> dict:
  1992. """Remove the OIDC link between the current user and a provider."""
  1993. result = await db.execute(
  1994. select(UserOIDCLink)
  1995. .where(UserOIDCLink.user_id == current_user.id)
  1996. .where(UserOIDCLink.provider_id == provider_id)
  1997. )
  1998. link = result.scalar_one_or_none()
  1999. if not link:
  2000. raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="OIDC link not found")
  2001. await db.delete(link)
  2002. await db.commit()
  2003. return {"message": "OIDC link removed"}
  2004. # ---------------------------------------------------------------------------
  2005. # Internal helpers
  2006. # ---------------------------------------------------------------------------
  2007. async def _get_base_external_url(db: AsyncSession) -> str:
  2008. """Return the base external URL (no trailing slash, no /login suffix)."""
  2009. external_url = await get_setting(db, "external_url")
  2010. if external_url:
  2011. return external_url.rstrip("/")
  2012. return os.environ.get("APP_URL", "http://localhost:5173").rstrip("/")