/
/
1"""Authentication manager for Music Assistant webserver."""
2
3from __future__ import annotations
4
5import asyncio
6import contextlib
7import hashlib
8import logging
9import secrets
10from collections.abc import Callable, Collection, Mapping
11from datetime import datetime, timedelta
12from sqlite3 import IntegrityError, OperationalError
13from typing import TYPE_CHECKING, Any, cast
14
15import jwt as pyjwt
16from music_assistant_models.auth import (
17 AuthProviderType,
18 AuthToken,
19 Scope,
20 User,
21 UserAuthProvider,
22 UserRole,
23)
24from music_assistant_models.errors import (
25 AuthenticationRequired,
26 InsufficientPermissions,
27 InvalidDataError,
28)
29
30from music_assistant.constants import (
31 CONF_PLAYERS,
32 CONF_PROVIDERS,
33 DB_TABLE_PLAYLOG,
34 HOMEASSISTANT_SYSTEM_USER,
35 MASS_LOGGER_NAME,
36)
37from music_assistant.controllers.webserver.helpers.auth_middleware import (
38 ROLE_SCOPES,
39 get_current_client_id,
40 get_current_peer_address,
41 get_current_token,
42 get_current_user,
43 has_scope,
44)
45from music_assistant.controllers.webserver.helpers.auth_providers import (
46 AuthResult,
47 BuiltinLoginProvider,
48 HomeAssistantOAuthProvider,
49 HomeAssistantProviderConfig,
50 LoginProvider,
51 LoginRateLimiter,
52 normalize_username,
53)
54from music_assistant.helpers.api import api_command
55from music_assistant.helpers.database import DatabaseConnection
56from music_assistant.helpers.datetime import utc
57from music_assistant.helpers.json import json_dumps, json_loads
58from music_assistant.helpers.jwt_auth import JWTHelper
59
60if TYPE_CHECKING:
61 from music_assistant.controllers.webserver import WebserverController
62 from music_assistant.providers.hass import HomeAssistantProvider
63
64LOGGER = logging.getLogger(f"{MASS_LOGGER_NAME}.auth")
65
66# Database schema version
67DB_SCHEMA_VERSION = 5
68
69# Token expiration constants (in days)
70TOKEN_SHORT_LIVED_EXPIRATION = 30 # Short-lived tokens (auto-renewing on use)
71TOKEN_LONG_LIVED_EXPIRATION = 365 # Long-lived tokens (1 year, no auto-renewal)
72# Max days a sliding short-lived session may live from creation before re-auth.
73TOKEN_ABSOLUTE_MAX_EXPIRATION = 90
74TOKEN_GUEST_EXPIRATION = 1 # Guest sessions: short fixed lifetime, no renewal
75# Days before the absolute cap at which the HA integration token is rotated
76HA_TOKEN_ROTATION_MARGIN = 7
77# Minimum age of a token's stored last_used_at before token activity is persisted again
78TOKEN_ACTIVITY_PERSIST_INTERVAL = timedelta(hours=1)
79
80HA_TOKEN_SETTING_KEY = "ha_integration_token"
81HA_TOKEN_NAME = "Home Assistant Integration"
82
83# Join code constants (short codes for QR/link-based login)
84JOIN_CODE_LENGTH = 12
85JOIN_CODE_CHARSET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" # No I/O/0/1 for readability
86JOIN_CODE_DEFAULT_EXPIRY_HOURS = 8
87# Failed exchanges are throttled per calling websocket connection, so one guest fumbling a
88# stale QR code cannot lock out every other guest at a party. Callers that reach the API
89# without a connection identity (the JSON RPC endpoint, in-process callers) share one bucket.
90JOIN_CODE_ANONYMOUS_RATE_LIMIT_KEY = "no-connection"
91# Second, server-wide bucket that backstops the per-connection buckets, since a client can
92# start a new connection (and thus a new bucket) at will. The join code itself is what makes
93# guessing infeasible (12 chars over a 32 symbol alphabet is ~2^60, valid for hours), so this
94# ceiling is deliberately far above any plausible party-scale burst of legitimate failures.
95JOIN_CODE_GLOBAL_RATE_LIMIT_KEY = "all-connections"
96JOIN_CODE_GLOBAL_FAILURE_CEILING = 1000
97JOIN_CODE_GLOBAL_COOLDOWN_SECONDS = 60
98
99
100class AuthenticationManager:
101 """Manager for authentication and user management (part of webserver controller)."""
102
103 def __init__(self, webserver: WebserverController) -> None:
104 """
105 Initialize the authentication manager.
106
107 :param webserver: WebserverController instance.
108 """
109 self.webserver = webserver
110 self.mass = webserver.mass
111 self.database: DatabaseConnection = None # type: ignore[assignment]
112 self.login_providers: dict[str, LoginProvider] = {}
113 self.logger = LOGGER
114 self._has_users: bool = False
115 self.jwt_helper: JWTHelper = None # type: ignore[assignment]
116 self._join_code_rate_limiter = LoginRateLimiter(subject="client")
117 self._join_code_global_rate_limiter = LoginRateLimiter(
118 delay_tiers=((JOIN_CODE_GLOBAL_FAILURE_CEILING, JOIN_CODE_GLOBAL_COOLDOWN_SECONDS),),
119 warn_threshold=JOIN_CODE_GLOBAL_FAILURE_CEILING,
120 alert_threshold=JOIN_CODE_GLOBAL_FAILURE_CEILING * 2,
121 subject="join_codes",
122 )
123 # Stops concurrent exchanges from passing the rate limit check before failures land
124 self._join_code_exchange_lock = asyncio.Lock()
125 # Serialises the read-modify-write of the user access filters
126 self._user_filter_lock = asyncio.Lock()
127 self._access_revoked_callbacks: list[Callable[[User], None]] = []
128
129 async def setup(self) -> None:
130 """Initialize the authentication manager."""
131 # Setup database
132 db_path = self.mass.storage_path + "/auth.db"
133 self.database = DatabaseConnection(db_path)
134 await self.database.setup()
135
136 # Create database schema and handle migrations
137 await self._setup_database()
138
139 # Initialize JWT helper with secret key
140 jwt_secret = await self._get_or_create_jwt_secret()
141 self.jwt_helper = JWTHelper(jwt_secret)
142
143 # Setup login providers
144 await self._setup_login_providers()
145
146 self._has_users = await self._has_non_system_users()
147
148 # migrate the Home Assistant system user of pre-existing installs to the service role
149 await self._migrate_system_user_role()
150
151 # repair filters that were left pointing at removed providers/players
152 await self._prune_stale_user_filters()
153
154 self._schedule_join_code_cleanup()
155
156 self.logger.info(
157 "Authentication manager initialized (providers=%d)", len(self.login_providers)
158 )
159
160 async def close(self) -> None:
161 """Cleanup on exit."""
162 if self.database:
163 await self.database.close()
164
165 @property
166 def has_users(self) -> bool:
167 """Check if any users exist in the system."""
168 return self._has_users
169
170 async def authenticate_with_credentials(
171 self, provider_id: str, credentials: dict[str, Any]
172 ) -> AuthResult:
173 """
174 Authenticate a user with credentials.
175
176 :param provider_id: The login provider ID.
177 :param credentials: Provider-specific credentials.
178 """
179 provider = self.login_providers.get(provider_id)
180 if not provider:
181 return AuthResult(success=False, error="Invalid provider")
182
183 return await provider.authenticate(credentials)
184
185 async def authenticate_with_token(self, token: str) -> User | None:
186 """
187 Authenticate a user with an access token (JWT or legacy).
188
189 Supports both JWT tokens and legacy hash-based tokens for backward compatibility.
190
191 :param token: The access token (JWT or legacy hash token).
192 """
193 # Try to decode as JWT first
194 try:
195 payload = self.jwt_helper.decode_token(token, verify_exp=True)
196 token_id = payload.get("jti")
197 token_user_id = payload.get("sub")
198
199 if not token_id or not token_user_id:
200 return None
201
202 token_row = await self.database.get_row("auth_tokens", {"token_id": token_id})
203 if not token_row:
204 return None
205
206 # Database is source of truth for token metadata, not the (immutable) JWT payload.
207 # A payload/row mismatch means a tampered or stale token: reject rather than trust it.
208 if token_user_id != token_row["user_id"]:
209 return None
210 is_long_lived = bool(token_row["is_long_lived"])
211
212 # Database expiration is source of truth
213 if token_row["expires_at"]:
214 db_expires_at = datetime.fromisoformat(token_row["expires_at"])
215 if utc() > db_expires_at:
216 await self.database.delete("auth_tokens", {"token_id": token_id})
217 return None
218
219 user = await self.get_user(token_row["user_id"])
220 if not user:
221 return None
222
223 updates = await self._refresh_token_expiration(token_row, user, is_long_lived)
224 if updates is None:
225 return None
226 if updates:
227 await self.database.update("auth_tokens", {"token_id": token_id}, updates)
228
229 return user
230
231 except pyjwt.ExpiredSignatureError:
232 if token_id := self.jwt_helper.get_token_id(token):
233 await self.database.delete("auth_tokens", {"token_id": token_id})
234 return None
235 except pyjwt.InvalidTokenError:
236 self.logger.debug("Token is not a valid JWT, trying legacy hash lookup")
237 except Exception as err:
238 self.logger.debug("Error decoding JWT token: %s, trying legacy hash lookup", err)
239
240 # Fallback to legacy hash-based token lookup
241 token_hash = hashlib.sha256(token.encode()).hexdigest()
242 token_row = await self.database.get_row("auth_tokens", {"token_hash": token_hash})
243 if not token_row:
244 return None
245
246 # Check if token is expired
247 if token_row["expires_at"]:
248 expires_at = datetime.fromisoformat(token_row["expires_at"])
249 if utc() > expires_at:
250 # Token expired, delete it
251 await self.database.delete("auth_tokens", {"token_id": token_row["token_id"]})
252 return None
253
254 user = await self.get_user(token_row["user_id"])
255 if not user:
256 return None
257
258 is_long_lived = bool(token_row["is_long_lived"])
259 legacy_updates = await self._refresh_token_expiration(token_row, user, is_long_lived)
260 if legacy_updates is None:
261 return None
262 if legacy_updates:
263 await self.database.update(
264 "auth_tokens", {"token_id": token_row["token_id"]}, legacy_updates
265 )
266
267 return user
268
269 async def get_token_id_from_token(self, token: str) -> str | None:
270 """
271 Get token_id from a token string (for tracking revocation).
272
273 :param token: The access token (JWT or legacy hash token).
274 :return: The token_id or None if token not found.
275 """
276 # Try to extract from JWT first
277 if token_id := self.jwt_helper.get_token_id(token):
278 return token_id
279
280 # Fallback: Hash-based lookup for legacy tokens
281 token_hash = hashlib.sha256(token.encode()).hexdigest()
282 token_row = await self.database.get_row("auth_tokens", {"token_hash": token_hash})
283 if not token_row:
284 return None
285 return str(token_row["token_id"])
286
287 @api_command("auth/user", required_scope=Scope.USERS_READ)
288 async def get_user(self, user_id: str) -> User | None:
289 """
290 Get user by ID (requires the users.read scope).
291
292 :param user_id: The user ID.
293 :return: User object or None if not found.
294 """
295 user_row = await self.database.get_row("users", {"user_id": user_id})
296 if not user_row or not user_row["enabled"]:
297 return None
298
299 return User(
300 user_id=user_row["user_id"],
301 username=user_row["username"],
302 role=user_row["role"],
303 enabled=bool(user_row["enabled"]),
304 created_at=datetime.fromisoformat(user_row["created_at"]),
305 display_name=user_row["display_name"],
306 avatar_url=user_row["avatar_url"],
307 preferences=json_loads(user_row["preferences"]),
308 player_filter=json_loads(user_row["player_filter"]),
309 provider_filter=json_loads(user_row["provider_filter"]),
310 )
311
312 async def get_user_by_username(self, username: str) -> User | None:
313 """
314 Get user by username.
315
316 :param username: The username.
317 :return: User object or None if not found.
318 """
319 username = normalize_username(username)
320
321 user_row = await self.database.get_row("users", {"username": username})
322 if not user_row:
323 return None
324
325 return await self.get_user(user_row["user_id"])
326
327 async def get_user_by_provider_link(
328 self, provider_type: AuthProviderType, provider_user_id: str
329 ) -> User | None:
330 """
331 Get user by their provider link.
332
333 :param provider_type: The auth provider type.
334 :param provider_user_id: The user ID from the provider.
335 """
336 link_row = await self.database.get_row(
337 "user_auth_providers",
338 {
339 "provider_type": provider_type.value,
340 "provider_user_id": provider_user_id,
341 },
342 )
343 if not link_row:
344 return None
345
346 return await self.get_user(link_row["user_id"])
347
348 async def create_user(
349 self,
350 username: str,
351 role: UserRole = UserRole.USER,
352 display_name: str | None = None,
353 avatar_url: str | None = None,
354 preferences: dict[str, Any] | None = None,
355 player_filter: list[str] | None = None,
356 provider_filter: list[str] | None = None,
357 ) -> User:
358 """
359 Create a new user.
360
361 :param username: The username.
362 :param role: The user role (default: USER).
363 :param display_name: Optional display name.
364 :param avatar_url: Optional avatar URL.
365 :param preferences: Optional user preferences dict.
366 :param player_filter: Optional list of player IDs user has access to.
367 :param provider_filter: Optional list of provider instance IDs user has access to.
368 """
369 normalized_username = normalize_username(username)
370
371 # Check if this is the first non-system user
372 is_first_user = not await self._has_non_system_users()
373
374 user_id = secrets.token_urlsafe(32)
375 created_at = utc()
376 if preferences is None:
377 preferences = {}
378 if player_filter is None:
379 player_filter = []
380 if provider_filter is None:
381 provider_filter = []
382
383 user_data = {
384 "user_id": user_id,
385 "username": normalized_username,
386 "role": role.value,
387 "enabled": True,
388 "created_at": created_at.isoformat(),
389 "display_name": display_name,
390 "avatar_url": avatar_url,
391 "preferences": json_dumps(preferences),
392 "player_filter": json_dumps(player_filter),
393 "provider_filter": json_dumps(provider_filter),
394 }
395
396 await self.database.insert("users", user_data)
397
398 user = User(
399 user_id=user_id,
400 username=normalized_username,
401 role=role,
402 enabled=True,
403 created_at=created_at,
404 display_name=display_name,
405 avatar_url=avatar_url,
406 preferences=preferences,
407 player_filter=player_filter,
408 provider_filter=provider_filter,
409 )
410
411 # If this is the first non-system user, migrate playlog entries to them
412 if is_first_user and normalized_username != HOMEASSISTANT_SYSTEM_USER:
413 self._has_users = True
414 await self._migrate_playlog_to_first_user(user_id)
415
416 return user
417
418 async def get_homeassistant_system_user(self) -> User:
419 """
420 Get or create the Home Assistant system user.
421
422 This is a special system user created automatically for Home Assistant integration.
423 It bypasses normal authentication but is restricted to the ingress webserver.
424
425 :return: The Home Assistant system user.
426 """
427 username = HOMEASSISTANT_SYSTEM_USER
428 display_name = "Home Assistant Integration"
429 role = UserRole.SERVICE
430
431 normalized_username = normalize_username(username)
432
433 # Try to find existing user by username
434 user_row = await self.database.get_row("users", {"username": normalized_username})
435 if user_row:
436 # Use get_user to ensure preferences are parsed correctly
437 user = await self.get_user(user_row["user_id"])
438 assert user is not None # User exists in DB, so get_user must return it
439 return user
440
441 # Create new system user
442 user = await self.create_user(
443 username=username,
444 role=role,
445 display_name=display_name,
446 )
447 self.logger.debug("Created Home Assistant system user: %s (role: %s)", username, role.value)
448 return user
449
450 async def get_homeassistant_system_user_token(self) -> str:
451 """
452 Get the auth token to announce to the Home Assistant integration.
453
454 Returns the same (still valid) token on repeated calls so re-announcing it via
455 Supervisor discovery is idempotent for the HA integration. A replacement is only
456 minted when the current token is missing, expired or revoked, or shortly before
457 it reaches its absolute lifetime cap - allowing seamless rotation as HA reloads
458 with the newly announced token while the old one is still accepted.
459
460 :return: Authentication token for the Home Assistant system user.
461 """
462 system_user = await self.get_homeassistant_system_user()
463
464 # Keep the plain token in settings for re-announcing; the jwt_secret next to it can mint any token anyway
465 if token_row := await self.database.get_row("settings", {"key": HA_TOKEN_SETTING_KEY}):
466 token = str(token_row["value"])
467 if await self._can_reuse_ha_integration_token(token, system_user):
468 return token
469
470 # A superseded token stays valid until expiry, so HA keeps working until it reloads
471 token = await self.create_token(
472 user=system_user,
473 name=HA_TOKEN_NAME,
474 is_long_lived=False,
475 )
476 await self.database.insert_or_replace(
477 "settings",
478 {"key": HA_TOKEN_SETTING_KEY, "value": token, "type": "string"},
479 )
480 now = utc()
481 for old_row in await self.database.get_rows(
482 "auth_tokens", {"user_id": system_user.user_id, "name": HA_TOKEN_NAME}
483 ):
484 if old_row["expires_at"] and datetime.fromisoformat(old_row["expires_at"]) <= now:
485 await self.database.delete("auth_tokens", {"token_id": old_row["token_id"]})
486 await self.database.commit()
487 return token
488
489 async def link_user_to_provider(
490 self,
491 user: User,
492 provider_type: AuthProviderType,
493 provider_user_id: str,
494 ) -> UserAuthProvider:
495 """
496 Link a user to an authentication provider.
497
498 If a link already exists for this provider/provider_user_id, returns the existing link.
499
500 :param user: The user to link.
501 :param provider_type: The provider type.
502 :param provider_user_id: The user ID from the provider (e.g., password hash, OAuth ID).
503 """
504 # Check if a link already exists for this provider/provider_user_id
505 existing_link = await self.database.get_row(
506 "user_auth_providers",
507 {
508 "provider_type": provider_type.value,
509 "provider_user_id": provider_user_id,
510 },
511 )
512
513 if existing_link:
514 # Link already exists - return it
515 return UserAuthProvider(
516 link_id=existing_link["link_id"],
517 user_id=existing_link["user_id"],
518 provider_type=AuthProviderType(existing_link["provider_type"]),
519 provider_user_id=existing_link["provider_user_id"],
520 created_at=datetime.fromisoformat(existing_link["created_at"]),
521 )
522
523 # Create new link
524 link_id = secrets.token_urlsafe(32)
525 created_at = utc()
526 link_data = {
527 "link_id": link_id,
528 "user_id": user.user_id,
529 "provider_type": provider_type.value,
530 "provider_user_id": provider_user_id,
531 "created_at": created_at.isoformat(),
532 }
533
534 await self.database.insert("user_auth_providers", link_data)
535
536 return UserAuthProvider(
537 link_id=link_id,
538 user_id=user.user_id,
539 provider_type=provider_type,
540 provider_user_id=provider_user_id,
541 created_at=created_at,
542 )
543
544 async def update_user(
545 self,
546 user: User,
547 username: str | None = None,
548 display_name: str | None = None,
549 avatar_url: str | None = None,
550 ) -> User:
551 """
552 Update a user's profile information.
553
554 :param user: The user to update.
555 :param username: New username (optional).
556 :param display_name: New display name (optional).
557 :param avatar_url: New avatar URL (optional).
558 """
559 updates = {}
560 if username is not None:
561 # Normalize username for case-insensitive authentication
562 updates["username"] = normalize_username(username)
563 if display_name is not None:
564 updates["display_name"] = display_name
565 if avatar_url is not None:
566 updates["avatar_url"] = avatar_url
567
568 if updates:
569 await self.database.update("users", {"user_id": user.user_id}, updates)
570
571 # Return updated user
572 updated_user = await self.get_user(user.user_id)
573 assert updated_user is not None # User exists, so get_user must return it
574 return updated_user
575
576 async def update_user_preferences(
577 self,
578 user: User,
579 preferences: dict[str, Any],
580 ) -> User:
581 """
582 Update a user's preferences.
583
584 :param user: The user to update.
585 :param preferences: New preferences dict (completely replaces existing preferences).
586 """
587 # Verify user exists
588 current_user = await self.get_user(user.user_id)
589 if not current_user:
590 raise ValueError(f"User {user.user_id} not found")
591
592 # Update database with new preferences (complete replacement)
593 await self.database.update(
594 "users",
595 {"user_id": user.user_id},
596 {"preferences": json_dumps(preferences)},
597 )
598
599 # Return updated user
600 updated_user = await self.get_user(user.user_id)
601 assert updated_user is not None # User exists, so get_user must return it
602 return updated_user
603
604 async def update_provider_link(
605 self,
606 user: User,
607 provider_type: AuthProviderType,
608 provider_user_id: str,
609 ) -> None:
610 """
611 Update a user's provider link (e.g., change password).
612
613 :param user: The user.
614 :param provider_type: The provider type.
615 :param provider_user_id: The new provider user ID (e.g., new password hash).
616 """
617 # Find existing link
618 link_row = await self.database.get_row(
619 "user_auth_providers",
620 {
621 "user_id": user.user_id,
622 "provider_type": provider_type.value,
623 },
624 )
625
626 if link_row:
627 # Update existing link
628 await self.database.update(
629 "user_auth_providers",
630 {"link_id": link_row["link_id"]},
631 {"provider_user_id": provider_user_id},
632 )
633 else:
634 # Create new link
635 await self.link_user_to_provider(user, provider_type, provider_user_id)
636
637 async def create_token(self, user: User, name: str, is_long_lived: bool = False) -> str:
638 """
639 Create a new JWT access token for a user.
640
641 :param user: The user to create the token for.
642 :param name: A name/description for the token (e.g., device name).
643 :param is_long_lived: Whether this is a long-lived token (default: False).
644 Short-lived tokens (False): Auto-renewing on use, expire after 30 days of inactivity,
645 capped at an absolute maximum lifetime from creation (see TOKEN_ABSOLUTE_MAX_EXPIRATION).
646 Tokens for guest users get a short fixed lifetime instead and never renew.
647 Long-lived tokens (True): No auto-renewal, expire after 1 year.
648 :return: JWT token string.
649 """
650 # Generate unique token ID
651 token_id = secrets.token_urlsafe(32)
652
653 # Calculate expiration based on token type
654 created_at = utc()
655 if is_long_lived:
656 # Long-lived tokens expire after 1 year (no auto-renewal)
657 expires_at = created_at + timedelta(days=TOKEN_LONG_LIVED_EXPIRATION)
658 jwt_expires_at = expires_at
659 elif user.role == UserRole.GUEST:
660 expires_at = created_at + timedelta(days=TOKEN_GUEST_EXPIRATION)
661 jwt_expires_at = expires_at
662 else:
663 # Short-lived tokens expire after 30 days (with auto-renewal on use)
664 expires_at = created_at + timedelta(days=TOKEN_SHORT_LIVED_EXPIRATION)
665 # The exp claim must carry the absolute cap, or it would cut off sliding renewals
666 jwt_expires_at = created_at + timedelta(days=TOKEN_ABSOLUTE_MAX_EXPIRATION)
667
668 # Generate JWT token
669 token = self.jwt_helper.encode_token(
670 user=user,
671 token_id=token_id,
672 token_name=name,
673 expires_at=jwt_expires_at,
674 is_long_lived=is_long_lived,
675 )
676
677 # Store token hash in database for revocation checking
678 token_hash = hashlib.sha256(token.encode()).hexdigest()
679 token_data = {
680 "token_id": token_id,
681 "user_id": user.user_id,
682 "token_hash": token_hash,
683 "name": name,
684 "created_at": created_at.isoformat(),
685 "expires_at": expires_at.isoformat(),
686 "is_long_lived": 1 if is_long_lived else 0,
687 }
688 await self.database.insert("auth_tokens", token_data)
689
690 return token
691
692 @api_command("auth/token/revoke")
693 async def revoke_token(self, token_id: str) -> None:
694 """
695 Revoke an auth token.
696
697 :param token_id: The token ID to revoke.
698 """
699 user = get_current_user()
700 if not user:
701 raise AuthenticationRequired("Not authenticated")
702
703 token_row = await self.database.get_row("auth_tokens", {"token_id": token_id})
704 if not token_row:
705 raise InvalidDataError("Token not found")
706
707 # Check permissions - users can only revoke their own tokens
708 # unless they hold the users.manage scope
709 if token_row["user_id"] != user.user_id and not has_scope(user, Scope.USERS_MANAGE):
710 raise InsufficientPermissions("You can only revoke your own tokens")
711
712 await self.database.delete("auth_tokens", {"token_id": token_id})
713
714 # Disconnect any WebSocket connections using this token
715 self.webserver.disconnect_websockets_for_token(token_id)
716
717 self.logger.info(
718 "Token revoked by user '%s' (token_id=%s)",
719 user.username,
720 token_id,
721 )
722
723 def subscribe_user_access_revoked(self, callback: Callable[[User], None]) -> Callable[[], None]:
724 """
725 Subscribe to a user's access being withdrawn.
726
727 Fires on deliberate access withdrawal: bulk token revocation
728 (revoke_tokens_for_user), account disable, and account deletion. Revoking a
729 single token (e.g. a logout) does not fire it, so credentials bound to the
730 account survive a plain logout.
731
732 :param callback: Called with the affected user.
733 :return: Callable that removes the subscription.
734 """
735 self._access_revoked_callbacks.append(callback)
736
737 def _unsubscribe() -> None:
738 with contextlib.suppress(ValueError):
739 self._access_revoked_callbacks.remove(callback)
740
741 return _unsubscribe
742
743 async def revoke_tokens_for_user(self, user: User) -> int:
744 """
745 Revoke all auth tokens for a user.
746
747 This is an internal method for programmatic use (e.g., when disabling guest access).
748 Unlike revoke_token(), this does not require an authenticated user context.
749
750 :param user: The user whose tokens should be revoked.
751 :return: Number of tokens revoked.
752 """
753 token_rows = await self.database.get_rows("auth_tokens", {"user_id": user.user_id})
754
755 # Disconnect any WebSocket connections using these tokens
756 for token_row in token_rows:
757 self.webserver.disconnect_websockets_for_token(token_row["token_id"])
758
759 if token_rows:
760 # Delete all tokens in one go
761 await self.database.execute(
762 "DELETE FROM auth_tokens WHERE user_id = :user_id",
763 {"user_id": user.user_id},
764 )
765 await self.database.commit()
766 self.logger.info("Revoked %d token(s) for user '%s'", len(token_rows), user.username)
767
768 # Notify even with no tokens left: subscribers may hold credentials tied to
769 # this user's access that must be withdrawn regardless.
770 self._notify_user_access_revoked(user)
771
772 return len(token_rows)
773
774 @api_command("auth/tokens")
775 async def get_user_tokens(self, user_id: str | None = None) -> list[AuthToken]:
776 """
777 Get current user's auth tokens or another user's tokens (admin only).
778
779 The last_used_at timestamp is persisted at most once per hour, so it may lag
780 actual token usage by up to an hour.
781
782 :param user_id: Optional user ID to get tokens for (admin only).
783 :return: List of auth tokens.
784 """
785 current_user = get_current_user()
786 if not current_user:
787 return []
788
789 # If user_id is provided and different from current user,
790 # require the users.manage scope
791 if user_id and user_id != current_user.user_id:
792 if not has_scope(current_user, Scope.USERS_MANAGE):
793 return []
794 target_user = await self.get_user(user_id)
795 if not target_user:
796 return []
797 else:
798 target_user = current_user
799
800 token_rows = await self.database.get_rows(
801 "auth_tokens", {"user_id": target_user.user_id}, limit=100
802 )
803 return [AuthToken.from_dict(dict(row)) for row in token_rows]
804
805 @api_command("auth/users", required_scope=Scope.USERS_READ)
806 async def list_users(self) -> list[User]:
807 """
808 Get all users (requires the users.read scope).
809
810 System users are excluded from the list.
811
812 :return: List of user objects.
813 """
814 user_rows = await self.database.get_rows("users", limit=1000)
815 users = []
816 for row in user_rows:
817 # Skip system users
818 if row["username"] == HOMEASSISTANT_SYSTEM_USER:
819 continue
820 users.append(
821 User(
822 user_id=row["user_id"],
823 username=row["username"],
824 role=row["role"],
825 enabled=bool(row["enabled"]),
826 created_at=datetime.fromisoformat(row["created_at"]),
827 display_name=row["display_name"],
828 avatar_url=row["avatar_url"],
829 preferences=json_loads(row["preferences"]),
830 player_filter=json_loads(row["player_filter"]),
831 provider_filter=json_loads(row["provider_filter"]),
832 )
833 )
834 return users
835
836 async def update_user_role(self, user_id: str, new_role: UserRole, admin_user: User) -> bool:
837 """
838 Update a user's role (requires the users.manage scope).
839
840 :param user_id: The user ID to update.
841 :param new_role: The new role to assign.
842 :param admin_user: The user performing the action.
843 """
844 if not has_scope(admin_user, Scope.USERS_MANAGE):
845 return False
846
847 user_row = await self.database.get_row("users", {"user_id": user_id})
848 if not user_row:
849 return False
850
851 old_role = user_row["role"]
852 await self.database.update(
853 "users",
854 {"user_id": user_id},
855 {"role": new_role.value},
856 )
857 self.logger.info(
858 "User role changed: '%s' from '%s' to '%s' by admin '%s'",
859 user_row["username"],
860 old_role,
861 new_role.value,
862 admin_user.username,
863 )
864 return True
865
866 @api_command("auth/user/enable", required_scope=Scope.USERS_MANAGE)
867 async def enable_user(self, user_id: str) -> None:
868 """
869 Enable user account (admin only).
870
871 :param user_id: The user ID.
872 """
873 await self.database.update(
874 "users",
875 {"user_id": user_id},
876 {"enabled": 1},
877 )
878 self.logger.info("User account enabled (user_id=%s)", user_id)
879
880 @api_command("auth/user/disable", required_scope=Scope.USERS_MANAGE)
881 async def disable_user(self, user_id: str) -> None:
882 """
883 Disable user account (admin only).
884
885 :param user_id: The user ID.
886 """
887 admin_user = get_current_user()
888 if not admin_user:
889 raise AuthenticationRequired("Not authenticated")
890
891 # Cannot disable yourself
892 if user_id == admin_user.user_id:
893 raise InvalidDataError("Cannot disable your own account")
894
895 # Look up the user before disabling (get_user hides disabled accounts)
896 user_row = await self.database.get_row("users", {"user_id": user_id})
897 if not user_row:
898 raise InvalidDataError("User not found")
899
900 await self.database.update(
901 "users",
902 {"user_id": user_id},
903 {"enabled": 0},
904 )
905
906 # Disconnect all WebSocket connections for this user
907 self.webserver.disconnect_websockets_for_user(user_id)
908
909 # A disabled account's tokens stop authenticating, so credentials bound to its
910 # access must be withdrawn with them (they return on the next login after enable).
911 self._notify_user_access_revoked(
912 User(user_id=user_row["user_id"], username=user_row["username"], role=user_row["role"])
913 )
914
915 self.logger.info("User account disabled (user_id=%s)", user_id)
916
917 async def get_login_providers(self) -> list[dict[str, Any]]:
918 """Get list of available login providers (dynamically checks for HA provider)."""
919 # Sync HA OAuth provider with HA provider availability
920 await self._sync_ha_oauth_provider()
921
922 providers = []
923 for provider_id, provider in self.login_providers.items():
924 providers.append(
925 {
926 "provider_id": provider_id,
927 "provider_type": provider.provider_type.value,
928 "requires_redirect": provider.requires_redirect,
929 }
930 )
931 return providers
932
933 @api_command("auth/login", authenticated=False)
934 async def login(
935 self,
936 username: str | None = None,
937 password: str | None = None,
938 provider_id: str = "builtin",
939 device_name: str | None = None,
940 **extra_credentials: Any,
941 ) -> dict[str, Any]:
942 """
943 Authenticate user with credentials via WebSocket.
944
945 This command allows clients to authenticate over the WebSocket connection
946 using username/password or other provider-specific credentials.
947
948 :param username: Username for authentication (for builtin provider).
949 :param password: Password for authentication (for builtin provider).
950 :param provider_id: The login provider ID (defaults to "builtin").
951 :param device_name: Optional device name for the token (e.g., "iPhone 15", "Desktop PC").
952 :param extra_credentials: Additional provider-specific credentials.
953 :return: Authentication result with access token if successful.
954 """
955 # Build credentials dict from parameters
956 credentials: dict[str, Any] = {}
957 if username is not None:
958 credentials["username"] = username
959 if password is not None:
960 credentials["password"] = password
961 credentials.update(extra_credentials)
962
963 auth_result = await self.authenticate_with_credentials(provider_id, credentials)
964
965 if not auth_result.success:
966 self.logger.warning(
967 "Login failed for username '%s' via provider '%s'",
968 username or "<not provided>",
969 provider_id,
970 )
971 return {
972 "success": False,
973 "error": auth_result.error or "Authentication failed",
974 }
975
976 if not auth_result.user:
977 return {
978 "success": False,
979 "error": "Authentication failed: no user returned",
980 }
981
982 # Create short-lived access token with device name if provided
983 token_name = device_name or f"WebSocket Session - {auth_result.user.username}"
984 token = await self.create_token(
985 auth_result.user,
986 is_long_lived=False,
987 name=token_name,
988 )
989
990 self.logger.info(
991 "User '%s' logged in via provider '%s'",
992 auth_result.user.username,
993 provider_id,
994 )
995
996 return {
997 "success": True,
998 "access_token": token,
999 "user": {
1000 "user_id": auth_result.user.user_id,
1001 "username": auth_result.user.username,
1002 "display_name": auth_result.user.display_name,
1003 "role": auth_result.user.role,
1004 },
1005 }
1006
1007 @api_command("auth/providers", authenticated=False)
1008 async def get_providers(self) -> list[dict[str, Any]]:
1009 """
1010 Get list of available authentication providers.
1011
1012 Returns information about all available login providers including
1013 whether they require OAuth redirect flow.
1014 """
1015 return await self.get_login_providers()
1016
1017 @api_command("auth/authorization_url", authenticated=False)
1018 async def get_auth_url(
1019 self,
1020 provider_id: str,
1021 return_url: str | None = None,
1022 ) -> dict[str, str | None]:
1023 """
1024 Get OAuth authorization URL for authentication.
1025
1026 For OAuth providers (like Home Assistant), this returns the URL that
1027 the user should visit in their browser to authorize the application.
1028
1029 :param provider_id: The provider ID (e.g., "hass").
1030 :param return_url: URL to redirect to after OAuth completes.
1031 :return: Dictionary with authorization_url.
1032 """
1033 auth_url = await self.get_authorization_url(provider_id, return_url)
1034 if not auth_url:
1035 return {
1036 "authorization_url": None,
1037 "error": "Provider does not support OAuth or does not exist",
1038 }
1039
1040 return {
1041 "authorization_url": auth_url,
1042 }
1043
1044 async def get_authorization_url(
1045 self, provider_id: str, return_url: str | None = None
1046 ) -> str | None:
1047 """
1048 Get OAuth authorization URL for a provider.
1049
1050 :param provider_id: The provider ID.
1051 :param return_url: Optional URL to redirect to after successful login.
1052 """
1053 provider = self.login_providers.get(provider_id)
1054 if not provider or not provider.requires_redirect:
1055 return None
1056
1057 # Build callback redirect_uri
1058 redirect_uri = f"{self.webserver.base_url}/auth/callback?provider_id={provider_id}"
1059 return await provider.get_authorization_url(redirect_uri, return_url)
1060
1061 async def handle_oauth_callback(
1062 self, provider_id: str, code: str, state: str, redirect_uri: str
1063 ) -> AuthResult:
1064 """
1065 Handle OAuth callback.
1066
1067 :param provider_id: The provider ID.
1068 :param code: OAuth authorization code.
1069 :param state: OAuth state parameter.
1070 :param redirect_uri: The callback URL.
1071 """
1072 provider = self.login_providers.get(provider_id)
1073 if not provider:
1074 return AuthResult(success=False, error="Invalid provider")
1075
1076 return await provider.handle_oauth_callback(code, state, redirect_uri)
1077
1078 @api_command("auth/token/create")
1079 async def create_long_lived_token(self, name: str, user_id: str | None = None) -> str:
1080 """
1081 Create a new long-lived access token for current user or another user (admin only).
1082
1083 Long-lived tokens are intended for external integrations and API access.
1084 They expire after 1 year and do NOT auto-renew on use.
1085
1086 Short-lived tokens (for regular user sessions) are only created during login
1087 and auto-renew on each use (sliding 30-day expiration window).
1088
1089 Long-lived tokens cannot be created for guest accounts.
1090
1091 :param name: The name/description for the token (e.g., "Home Assistant", "Mobile App").
1092 :param user_id: Optional user ID to create token for (admin only).
1093 :return: The created token string.
1094 """
1095 current_user = get_current_user()
1096 if not current_user:
1097 raise AuthenticationRequired("Not authenticated")
1098
1099 # If user_id is provided and different from current user,
1100 # require the users.manage scope
1101 if user_id and user_id != current_user.user_id:
1102 if not has_scope(current_user, Scope.USERS_MANAGE):
1103 raise InsufficientPermissions(
1104 "The users.manage scope is required to create tokens for other users"
1105 )
1106 target_user = await self.get_user(user_id)
1107 if not target_user:
1108 raise InvalidDataError("User not found")
1109 else:
1110 target_user = current_user
1111
1112 # Guest access is temporary by design, deny tokens that would outlive it
1113 if target_user.role == UserRole.GUEST:
1114 raise InsufficientPermissions("Long-lived tokens cannot be created for guest accounts")
1115
1116 # Create a long-lived token (only long-lived tokens can be created via this command)
1117 token = await self.create_token(target_user, name, is_long_lived=True)
1118 self.logger.info("Created long-lived token '%s' for user '%s'", name, target_user.username)
1119 return token
1120
1121 @api_command("auth/user/create", required_scope=Scope.USERS_MANAGE)
1122 async def create_user_with_api(
1123 self,
1124 username: str,
1125 password: str,
1126 role: str = "user",
1127 display_name: str | None = None,
1128 avatar_url: str | None = None,
1129 player_filter: list[str] | None = None,
1130 provider_filter: list[str] | None = None,
1131 ) -> User:
1132 """
1133 Create a new user with built-in authentication (admin only).
1134
1135 :param username: The username (minimum 2 characters).
1136 :param password: The password (minimum 8 characters).
1137 :param role: User role - "admin" or "user" (default: "user").
1138 :param display_name: Optional display name.
1139 :param avatar_url: Optional avatar URL.
1140 :param player_filter: Optional list of player IDs user has access to.
1141 :param provider_filter: Optional list of provider instance IDs user has access to.
1142 :return: Created user object.
1143 """
1144 # Validation
1145 if not username or len(username) < 2:
1146 raise InvalidDataError("Username must be at least 2 characters")
1147
1148 if not password or len(password) < 8:
1149 raise InvalidDataError("Password must be at least 8 characters")
1150
1151 # Validate role
1152 try:
1153 user_role = UserRole(role)
1154 except ValueError as err:
1155 raise InvalidDataError("Invalid role. Must be 'admin' or 'user'") from err
1156
1157 # Get built-in provider
1158 builtin_provider = self.login_providers.get("builtin")
1159 if not builtin_provider or not isinstance(builtin_provider, BuiltinLoginProvider):
1160 raise InvalidDataError("Built-in auth provider not available")
1161
1162 # Create user with password
1163 user = await builtin_provider.create_user_with_password(
1164 username,
1165 password,
1166 role=user_role,
1167 player_filter=player_filter,
1168 provider_filter=provider_filter,
1169 )
1170
1171 # Update optional fields if provided
1172 if display_name or avatar_url:
1173 updated_user = await self.update_user(
1174 user, display_name=display_name, avatar_url=avatar_url
1175 )
1176 if updated_user:
1177 user = updated_user
1178
1179 self.logger.info("User created by admin: %s (role: %s)", username, role)
1180 return user
1181
1182 @api_command("auth/user/delete", required_scope=Scope.USERS_MANAGE)
1183 async def delete_user(self, user_id: str) -> None:
1184 """
1185 Delete user account (admin only).
1186
1187 :param user_id: The user ID.
1188 """
1189 admin_user = get_current_user()
1190 if not admin_user:
1191 raise AuthenticationRequired("Not authenticated")
1192
1193 # Don't allow deleting yourself
1194 if user_id == admin_user.user_id:
1195 raise InvalidDataError("Cannot delete your own account")
1196
1197 # Look up the username before deleting
1198 user_row = await self.database.get_row("users", {"user_id": user_id})
1199 if not user_row:
1200 raise InvalidDataError("User not found")
1201
1202 # Delete user from database
1203 await self.database.delete("users", {"user_id": user_id})
1204 await self.database.commit()
1205
1206 # Disconnect all WebSocket connections for this user
1207 self.webserver.disconnect_websockets_for_user(user_id)
1208
1209 # Deletion cascades the user's tokens away, so it must announce the access
1210 # withdrawal itself for credentials bound to this user.
1211 self._notify_user_access_revoked(
1212 User(user_id=user_row["user_id"], username=user_row["username"], role=user_row["role"])
1213 )
1214
1215 self.logger.info(
1216 "User '%s' deleted by admin '%s'",
1217 user_row["username"],
1218 admin_user.username,
1219 )
1220
1221 @api_command("auth/me")
1222 async def get_current_user_info(self) -> User:
1223 """Get current authenticated user information."""
1224 current_user_obj = get_current_user()
1225 if not current_user_obj:
1226 raise AuthenticationRequired("Not authenticated")
1227 return current_user_obj
1228
1229 @api_command("auth/scopes")
1230 async def get_role_scopes(self) -> dict[str, list[str]]:
1231 """Get the scopes granted to each of the builtin user roles."""
1232 return {
1233 str(role): sorted(str(scope) for scope in scopes)
1234 for role, scopes in ROLE_SCOPES.items()
1235 }
1236
1237 async def update_user_filters(
1238 self,
1239 target_user: User,
1240 player_filter: list[str] | None,
1241 provider_filter: list[str] | None,
1242 ) -> User:
1243 """Update user player and provider filters (helper method)."""
1244 updates = {}
1245 if player_filter is not None:
1246 updates["player_filter"] = json_dumps(player_filter)
1247 if provider_filter is not None:
1248 updates["provider_filter"] = json_dumps(provider_filter)
1249
1250 if updates:
1251 await self.database.update("users", {"user_id": target_user.user_id}, updates)
1252 # Refresh target user to get updated filters
1253 refreshed_user = await self.get_user(target_user.user_id)
1254 if not refreshed_user:
1255 raise InvalidDataError("Failed to refresh user after filter update")
1256 return refreshed_user
1257 return target_user
1258
1259 async def remove_from_user_filters(
1260 self,
1261 provider_instance_ids: Collection[str] = (),
1262 player_ids: Collection[str] = (),
1263 ) -> None:
1264 """
1265 Remove the given providers and/or players from the access filters of all users.
1266
1267 Call this when a provider or player is permanently removed, so no user is left with
1268 an access filter that points at something that no longer exists.
1269
1270 :param provider_instance_ids: Instance IDs of the removed providers.
1271 :param player_ids: IDs of the removed players.
1272 """
1273 await self._rewrite_user_filters(
1274 keep_provider=(lambda x: x not in provider_instance_ids)
1275 if provider_instance_ids
1276 else None,
1277 keep_player=(lambda x: x not in player_ids) if player_ids else None,
1278 )
1279
1280 @api_command("auth/user/update")
1281 async def update_user_profile(
1282 self,
1283 user_id: str | None = None,
1284 username: str | None = None,
1285 display_name: str | None = None,
1286 avatar_url: str | None = None,
1287 password: str | None = None,
1288 role: str | None = None,
1289 preferences: dict[str, Any] | None = None,
1290 player_filter: list[str] | None = None,
1291 provider_filter: list[str] | None = None,
1292 ) -> User:
1293 """
1294 Update user profile information.
1295
1296 Users can update their own profile. Admins can update any user including role and password.
1297
1298 :param user_id: User ID to update (optional, defaults to current user).
1299 :param username: New username (optional).
1300 :param display_name: New display name (optional).
1301 :param avatar_url: New avatar URL (optional).
1302 :param password: New password (optional, minimum 8 characters).
1303 :param role: New role - "admin" or "user" (optional, set by admin only).
1304 :param preferences: User preferences dict (completely replaces existing, optional).
1305 :param player_filter: List of player IDs user has access to (set by admin only, optional).
1306 :param provider_filter: List of provider instance IDs user has access to (set by admin only, optional).
1307 :return: Updated user object.
1308 """
1309 current_user_obj = get_current_user()
1310 if not current_user_obj:
1311 raise AuthenticationRequired("Not authenticated")
1312
1313 # Determine target user
1314 may_manage_users = has_scope(current_user_obj, Scope.USERS_MANAGE)
1315 if user_id and user_id != current_user_obj.user_id:
1316 # Updating another user - requires the users.manage scope
1317 if not may_manage_users:
1318 raise InsufficientPermissions(
1319 "The users.manage scope is required to update other users"
1320 )
1321 target_user = await self.get_user(user_id)
1322 if not target_user:
1323 raise InvalidDataError("User not found")
1324 else:
1325 # Updating own profile
1326 target_user = current_user_obj
1327
1328 # Update role (requires the users.manage scope)
1329 if role:
1330 if not may_manage_users:
1331 raise InsufficientPermissions(
1332 "The users.manage scope is required to update user roles"
1333 )
1334
1335 try:
1336 new_role = UserRole(role)
1337 except ValueError as err:
1338 raise InvalidDataError("Invalid role. Must be 'admin' or 'user'") from err
1339
1340 success = await self.update_user_role(target_user.user_id, new_role, current_user_obj)
1341 if not success:
1342 raise InvalidDataError("Failed to update role")
1343
1344 # Refresh target user to get updated role
1345 refreshed_user = await self.get_user(target_user.user_id)
1346 if not refreshed_user:
1347 raise InvalidDataError("Failed to refresh user after role update")
1348 target_user = refreshed_user
1349
1350 # Update basic profile fields
1351 if username or display_name or avatar_url:
1352 updated_user = await self.update_user(
1353 target_user,
1354 username=username,
1355 display_name=display_name,
1356 avatar_url=avatar_url,
1357 )
1358 if not updated_user:
1359 raise InvalidDataError("Failed to update user profile")
1360 target_user = updated_user
1361
1362 # Update preferences if provided
1363 if preferences is not None:
1364 target_user = await self.update_user_preferences(target_user, preferences)
1365
1366 # Update player_filter and provider_filter (requires the users.manage scope)
1367 if player_filter is not None or provider_filter is not None:
1368 if not may_manage_users:
1369 raise InsufficientPermissions(
1370 "The users.manage scope is required to update player/provider filters"
1371 )
1372 target_user = await self.update_user_filters(
1373 target_user, player_filter, provider_filter
1374 )
1375
1376 # Update password if provided
1377 if password:
1378 await self._update_profile_password(
1379 target_user, password, may_manage_users, current_user_obj
1380 )
1381
1382 return target_user
1383
1384 @api_command("auth/logout")
1385 async def logout(self) -> None:
1386 """Logout current user by revoking the current token."""
1387 user = get_current_user()
1388 if not user:
1389 raise AuthenticationRequired("Not authenticated")
1390
1391 # Get current token from context
1392 token = get_current_token()
1393 if not token:
1394 raise InvalidDataError("No token in context")
1395
1396 # Find and revoke the token
1397 token_hash = hashlib.sha256(token.encode()).hexdigest()
1398 token_row = await self.database.get_row("auth_tokens", {"token_hash": token_hash})
1399 if token_row:
1400 await self.database.delete("auth_tokens", {"token_id": token_row["token_id"]})
1401
1402 # Disconnect any WebSocket connections using this token
1403 self.webserver.disconnect_websockets_for_token(token_row["token_id"])
1404
1405 self.logger.info("User '%s' logged out", user.username)
1406
1407 @api_command("auth/user/providers")
1408 async def get_my_providers(self) -> list[dict[str, Any]]:
1409 """
1410 Get current user's linked authentication providers.
1411
1412 :return: List of provider links.
1413 """
1414 user = get_current_user()
1415 if not user:
1416 return []
1417
1418 # Get provider links from database
1419 rows = await self.database.get_rows("user_auth_providers", {"user_id": user.user_id})
1420 providers = [UserAuthProvider.from_dict(dict(row)) for row in rows]
1421 return [p.to_dict() for p in providers]
1422
1423 @api_command("auth/user/unlink_provider", required_scope=Scope.USERS_MANAGE)
1424 async def unlink_provider(self, user_id: str, provider_type: str) -> bool:
1425 """
1426 Unlink authentication provider from user (admin only).
1427
1428 :param user_id: The user ID.
1429 :param provider_type: Provider type to unlink.
1430 :return: True if successful.
1431 """
1432 await self.database.delete(
1433 "user_auth_providers", {"user_id": user_id, "provider_type": provider_type}
1434 )
1435 await self.database.commit()
1436
1437 self.logger.info(
1438 "Auth provider '%s' unlinked from user (user_id=%s)",
1439 provider_type,
1440 user_id,
1441 )
1442 return True
1443
1444 # ==================== Join Code Methods ====================
1445
1446 async def generate_join_code(
1447 self,
1448 user: User,
1449 expires_in_hours: int = JOIN_CODE_DEFAULT_EXPIRY_HOURS,
1450 max_uses: int = 1,
1451 device_name: str = "Short Code Login",
1452 ) -> tuple[str, datetime]:
1453 """
1454 Generate a short join code for link/QR-based login.
1455
1456 This creates a short alphanumeric code that can be exchanged for a JWT token.
1457 Used for features like the party provider guest access, device pairing,
1458 or other short-code authentication flows.
1459
1460 :param user: The guest user that tokens created from this code will belong to.
1461 :param expires_in_hours: Hours until code expires (default: 8).
1462 :param max_uses: Maximum number of uses (0 = unlimited).
1463 :param device_name: Device name for tokens created with this code.
1464 :return: Tuple of (code, expires_at datetime).
1465 """
1466 if expires_in_hours <= 0:
1467 raise ValueError("expires_in_hours must be positive")
1468 if max_uses < 0:
1469 raise ValueError("max_uses must be non-negative (0 = unlimited)")
1470 if user.role != UserRole.GUEST:
1471 raise ValueError("Join codes can only be generated for guest accounts")
1472
1473 now = utc()
1474 expires_at = now + timedelta(hours=expires_in_hours)
1475
1476 for _ in range(3): # Try up to 3 times to avoid code collisions
1477 code = "".join(secrets.choice(JOIN_CODE_CHARSET) for _ in range(JOIN_CODE_LENGTH))
1478 code_data = {
1479 "code_id": secrets.token_urlsafe(32),
1480 "code": code,
1481 "user_id": user.user_id,
1482 "created_at": now.isoformat(),
1483 "expires_at": expires_at.isoformat(),
1484 "max_uses": max_uses,
1485 "use_count": 0,
1486 "device_name": device_name,
1487 }
1488 try:
1489 await self.database.insert("join_codes", code_data)
1490 await self.database.commit()
1491 self.logger.info(
1492 "Join code generated for user %s (expires: %s, max_uses: %s)",
1493 user.username,
1494 expires_at,
1495 max_uses,
1496 )
1497 return code, expires_at
1498 except IntegrityError:
1499 self.logger.warning("Join code collision, retrying...")
1500 continue
1501
1502 raise RuntimeError("Failed to generate a unique join code after 3 attempts")
1503
1504 async def revoke_join_codes(self, user: User) -> int:
1505 """
1506 Revoke all join codes for a user.
1507
1508 :param user: The user whose join codes should be revoked.
1509 :return: Number of codes revoked.
1510 """
1511 cursor = await self.database.execute(
1512 "DELETE FROM join_codes WHERE user_id = :user_id",
1513 {"user_id": user.user_id},
1514 )
1515 await self.database.commit()
1516
1517 count = int(cursor.rowcount)
1518 if count > 0:
1519 self.logger.info("Revoked %d join code(s) for user %s", count, user.username)
1520 return count
1521
1522 async def get_active_join_code(self, user: User) -> str | None:
1523 """
1524 Get the most recently created, non-expired join code for a user.
1525
1526 :param user: The user to look up codes for.
1527 :return: The join code string if found, None otherwise.
1528 """
1529 now = utc()
1530 cursor = await self.database.execute(
1531 """
1532 SELECT code FROM join_codes
1533 WHERE user_id = :user_id
1534 AND expires_at > :now
1535 AND (max_uses = 0 OR use_count < max_uses)
1536 ORDER BY created_at DESC
1537 LIMIT 1
1538 """,
1539 {"user_id": user.user_id, "now": now.isoformat()},
1540 )
1541 row = await cursor.fetchone()
1542 return str(row["code"]) if row else None
1543
1544 async def get_join_code_expiry(self, code: str, user: User | None = None) -> datetime | None:
1545 """
1546 Get the expiry datetime for an active join code.
1547
1548 :param code: The join code to look up.
1549 :param user: Optional user that must own the join code.
1550 :return: The expiry datetime if the code is active, None otherwise.
1551 """
1552 query = """
1553 SELECT expires_at FROM join_codes
1554 WHERE code = :code
1555 AND expires_at > :now
1556 AND (max_uses = 0 OR use_count < max_uses)
1557 """
1558 params: dict[str, Any] = {"code": code.upper(), "now": utc().isoformat()}
1559 if user is not None:
1560 query += "AND user_id = :user_id "
1561 params["user_id"] = user.user_id
1562 cursor = await self.database.execute(query + "LIMIT 1", params)
1563 row = await cursor.fetchone()
1564 return datetime.fromisoformat(str(row["expires_at"])) if row else None
1565
1566 @api_command("auth/join_code/exchange", authenticated=False)
1567 async def exchange_join_code(self, code: str) -> dict[str, Any]:
1568 """
1569 Exchange a join code for an access token (public API).
1570
1571 This is the public API endpoint for short-code authentication.
1572 Clients call this with a code (e.g., from QR scan or link) to receive a JWT token.
1573
1574 :param code: The short join code.
1575 :return: Authentication result with access token if successful.
1576 """
1577 rate_limit_key, key_is_exclusive = _join_code_rate_limit_key()
1578 async with self._join_code_exchange_lock:
1579 if throttled := await self._check_join_code_rate_limit(rate_limit_key):
1580 return throttled
1581
1582 token = await self._exchange_join_code(code)
1583
1584 if not token:
1585 await self._join_code_rate_limiter.record_failed_attempt(rate_limit_key)
1586 await self._join_code_global_rate_limiter.record_failed_attempt(
1587 JOIN_CODE_GLOBAL_RATE_LIMIT_KEY
1588 )
1589 return {
1590 "success": False,
1591 "error": "Invalid or expired join code",
1592 }
1593
1594 # A bucket is only cleared when it belongs to one caller alone, so presenting a
1595 # valid code never lifts the throttle for anyone else.
1596 if key_is_exclusive:
1597 await self._join_code_rate_limiter.clear_attempts(rate_limit_key)
1598
1599 # Decode token to get user info
1600 try:
1601 payload = self.jwt_helper.decode_token(token)
1602 return {
1603 "success": True,
1604 "access_token": token,
1605 "user": {
1606 "user_id": payload.get("sub"),
1607 "username": payload.get("username"),
1608 "role": payload.get("role"),
1609 },
1610 }
1611 except pyjwt.InvalidTokenError:
1612 return {
1613 "success": False,
1614 "error": "Failed to create access token",
1615 }
1616
1617 @api_command("auth/join_codes", required_scope=Scope.USERS_MANAGE)
1618 async def list_join_codes(self, user_id: str | None = None) -> list[dict[str, Any]]:
1619 """
1620 List join codes, optionally filtered by user (admin only).
1621
1622 :param user_id: Optional user ID to filter codes for.
1623 :return: List of join code records.
1624 """
1625 filter_args = {"user_id": user_id} if user_id else None
1626 rows = await self.database.get_rows("join_codes", filter_args, limit=100)
1627 return [dict(row) for row in rows]
1628
1629 @api_command("auth/join_code/revoke", required_scope=Scope.USERS_MANAGE)
1630 async def revoke_join_code(self, code_id: str) -> None:
1631 """
1632 Revoke a specific join code (admin only).
1633
1634 :param code_id: The code ID to revoke.
1635 """
1636 code_row = await self.database.get_row("join_codes", {"code_id": code_id})
1637 if not code_row:
1638 raise InvalidDataError("Join code not found")
1639
1640 await self.database.delete("join_codes", {"code_id": code_id})
1641 await self.database.commit()
1642 self.logger.info("Join code revoked (code_id=%s)", code_id)
1643
1644 async def _setup_database(self) -> None:
1645 """Set up database schema and handle migrations."""
1646 # Always create tables if they don't exist
1647 await self._create_database_tables()
1648
1649 # Check current schema version
1650 try:
1651 if db_row := await self.database.get_row("settings", {"key": "schema_version"}):
1652 prev_version = int(db_row["value"])
1653 else:
1654 prev_version = DB_SCHEMA_VERSION
1655 except KeyError, ValueError, Exception:
1656 # settings table doesn't exist yet or other error
1657 prev_version = 0
1658
1659 # Perform migration if needed
1660 if prev_version < DB_SCHEMA_VERSION:
1661 self.logger.warning(
1662 "Performing database migration from schema version %s to %s",
1663 prev_version,
1664 DB_SCHEMA_VERSION,
1665 )
1666 await self._migrate_database(prev_version)
1667
1668 # Store current schema version
1669 await self.database.insert_or_replace(
1670 "settings",
1671 {"key": "schema_version", "value": str(DB_SCHEMA_VERSION), "type": "int"},
1672 )
1673
1674 # Create indexes
1675 await self._create_database_indexes()
1676 await self.database.commit()
1677
1678 async def _create_database_tables(self) -> None:
1679 """Create database tables."""
1680 # Settings table (for schema version and other settings)
1681 await self.database.execute(
1682 """
1683 CREATE TABLE IF NOT EXISTS settings (
1684 key TEXT PRIMARY KEY,
1685 value TEXT,
1686 type TEXT
1687 )
1688 """
1689 )
1690 # Users table
1691 await self.database.execute(
1692 """
1693 CREATE TABLE IF NOT EXISTS users (
1694 user_id TEXT PRIMARY KEY,
1695 username TEXT NOT NULL UNIQUE,
1696 role TEXT NOT NULL,
1697 enabled INTEGER NOT NULL DEFAULT 1,
1698 created_at TEXT NOT NULL,
1699 display_name TEXT,
1700 avatar_url TEXT,
1701 preferences json NOT NULL DEFAULT '{}',
1702 player_filter json NOT NULL DEFAULT '[]',
1703 provider_filter json NOT NULL DEFAULT '[]'
1704 )
1705 """
1706 )
1707 # User auth provider links (many-to-many)
1708 await self.database.execute(
1709 """
1710 CREATE TABLE IF NOT EXISTS user_auth_providers (
1711 link_id TEXT PRIMARY KEY,
1712 user_id TEXT NOT NULL,
1713 provider_type TEXT NOT NULL,
1714 provider_user_id TEXT NOT NULL,
1715 created_at TEXT NOT NULL,
1716 UNIQUE(provider_type, provider_user_id),
1717 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1718 )
1719 """
1720 )
1721 # Auth tokens table
1722 await self.database.execute(
1723 """
1724 CREATE TABLE IF NOT EXISTS auth_tokens (
1725 token_id TEXT PRIMARY KEY,
1726 user_id TEXT NOT NULL,
1727 token_hash TEXT NOT NULL UNIQUE,
1728 name TEXT NOT NULL,
1729 created_at TEXT NOT NULL,
1730 expires_at TEXT,
1731 last_used_at TEXT,
1732 is_long_lived INTEGER NOT NULL DEFAULT 0,
1733 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1734 )
1735 """
1736 )
1737 # Join codes table (for short code to JWT exchange, used by providers like party)
1738 await self.database.execute(
1739 """
1740 CREATE TABLE IF NOT EXISTS join_codes (
1741 code_id TEXT PRIMARY KEY,
1742 code TEXT NOT NULL UNIQUE,
1743 user_id TEXT NOT NULL,
1744 created_at TEXT NOT NULL,
1745 expires_at TEXT NOT NULL,
1746 max_uses INTEGER DEFAULT 0,
1747 use_count INTEGER DEFAULT 0,
1748 last_used_at TEXT,
1749 device_name TEXT,
1750 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1751 )
1752 """
1753 )
1754 await self.database.commit()
1755
1756 async def _create_database_indexes(self) -> None:
1757 """Create database indexes."""
1758 await self.database.execute(
1759 "CREATE INDEX IF NOT EXISTS idx_user_auth_providers_user "
1760 "ON user_auth_providers(user_id)"
1761 )
1762 await self.database.execute(
1763 "CREATE INDEX IF NOT EXISTS idx_user_auth_providers_provider "
1764 "ON user_auth_providers(provider_type, provider_user_id)"
1765 )
1766 await self.database.execute(
1767 "CREATE INDEX IF NOT EXISTS idx_tokens_user ON auth_tokens(user_id)"
1768 )
1769 await self.database.execute(
1770 "CREATE INDEX IF NOT EXISTS idx_tokens_hash ON auth_tokens(token_hash)"
1771 )
1772 await self.database.execute(
1773 "CREATE INDEX IF NOT EXISTS idx_join_codes_user ON join_codes(user_id)"
1774 )
1775
1776 async def _migrate_database(self, from_version: int) -> None:
1777 """
1778 Perform database migration.
1779
1780 :param from_version: The schema version to migrate from.
1781 """
1782 self.logger.info(
1783 "Migrating auth database from version %s to %s", from_version, DB_SCHEMA_VERSION
1784 )
1785 # Migration to version 2: Recreate tables due to password salt breaking change
1786 if from_version < 2:
1787 # Drop all auth-related tables
1788 await self.database.execute("DROP TABLE IF EXISTS auth_tokens")
1789 await self.database.execute("DROP TABLE IF EXISTS user_auth_providers")
1790 await self.database.execute("DROP TABLE IF EXISTS users")
1791 await self.database.commit()
1792
1793 # Recreate tables with current schema
1794 await self._create_database_tables()
1795
1796 # Migration to version 3: Add player_filter and provider_filter columns
1797 if from_version < 3:
1798 with contextlib.suppress(OperationalError):
1799 # Column(s) may already exist
1800 await self.database.execute(
1801 "ALTER TABLE users ADD COLUMN player_filter json NOT NULL DEFAULT '[]'"
1802 )
1803 await self.database.execute(
1804 "ALTER TABLE users ADD COLUMN provider_filter json NOT NULL DEFAULT '[]'"
1805 )
1806 await self.database.commit()
1807
1808 # Migration to version 4: Make usernames case-insensitive by converting to lowercase
1809 if from_version < 4:
1810 await self.database.execute("UPDATE users SET username = LOWER(username)")
1811 await self.database.commit()
1812
1813 # Migration to version 5: Add join codes table
1814 if from_version < 5:
1815 await self.database.execute(
1816 """
1817 CREATE TABLE IF NOT EXISTS join_codes (
1818 code_id TEXT PRIMARY KEY,
1819 code TEXT NOT NULL UNIQUE,
1820 user_id TEXT NOT NULL,
1821 created_at TEXT NOT NULL,
1822 expires_at TEXT NOT NULL,
1823 max_uses INTEGER DEFAULT 0,
1824 use_count INTEGER DEFAULT 0,
1825 last_used_at TEXT,
1826 device_name TEXT,
1827 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1828 )
1829 """
1830 )
1831 await self.database.commit()
1832
1833 async def _get_or_create_jwt_secret(self) -> str:
1834 """
1835 Get or create JWT secret key from database.
1836
1837 :return: JWT secret key for signing tokens.
1838 """
1839 # Try to get existing secret
1840 if secret_row := await self.database.get_row("settings", {"key": "jwt_secret"}):
1841 return str(secret_row["value"])
1842
1843 # Generate new secret
1844 jwt_secret = JWTHelper.generate_secret_key()
1845
1846 # Store in database
1847 await self.database.insert_or_replace(
1848 "settings",
1849 {"key": "jwt_secret", "value": jwt_secret, "type": "string"},
1850 )
1851 await self.database.commit()
1852
1853 self.logger.info("Generated new JWT secret key")
1854 return jwt_secret
1855
1856 async def _setup_login_providers(self) -> None:
1857 """Set up available login providers based on configuration."""
1858 # Always enable built-in provider
1859 self.login_providers["builtin"] = BuiltinLoginProvider(self.mass, "builtin", {})
1860
1861 # Home Assistant OAuth provider
1862 # Automatically enabled if HA provider (plugin) is configured
1863 ha_provider = None
1864 for provider in self.mass.providers:
1865 if provider.domain == "hass" and provider.available:
1866 ha_provider = provider
1867 break
1868
1869 if ha_provider:
1870 ha_provider = cast("HomeAssistantProvider", ha_provider)
1871 ha_url = ha_provider.url
1872 if not ha_url:
1873 self.logger.warning(
1874 "Home Assistant provider has no URL configured, "
1875 "Home Assistant OAuth login is not available"
1876 )
1877 return
1878 ha_config: HomeAssistantProviderConfig = {"ha_url": ha_url}
1879 self.login_providers["homeassistant"] = HomeAssistantOAuthProvider(
1880 self.mass, "homeassistant", ha_config
1881 )
1882 self.logger.info(
1883 "Home Assistant OAuth provider enabled (using URL from HA provider: %s)",
1884 ha_url,
1885 )
1886
1887 async def _sync_ha_oauth_provider(self) -> None:
1888 """
1889 Sync HA OAuth provider with HA provider availability (dynamic check).
1890
1891 Adds the provider if HA is available, removes it if HA is not available.
1892 """
1893 # Find HA provider
1894 ha_provider = None
1895 for provider in self.mass.providers:
1896 if provider.domain == "hass" and provider.available:
1897 ha_provider = provider
1898 break
1899
1900 if ha_provider:
1901 # HA provider exists and is available - ensure OAuth provider is registered
1902 if "homeassistant" not in self.login_providers:
1903 ha_provider = cast("HomeAssistantProvider", ha_provider)
1904 ha_url = ha_provider.url
1905 if not ha_url:
1906 # missing URL must never break the login providers endpoint,
1907 # simply leave the HA OAuth provider unregistered
1908 self.logger.debug(
1909 "Home Assistant provider has no URL configured, "
1910 "Home Assistant OAuth login is not available"
1911 )
1912 return
1913 ha_config: HomeAssistantProviderConfig = {"ha_url": ha_url}
1914 self.login_providers["homeassistant"] = HomeAssistantOAuthProvider(
1915 self.mass, "homeassistant", ha_config
1916 )
1917 self.logger.info(
1918 "Home Assistant OAuth provider dynamically enabled (using URL: %s)",
1919 ha_url,
1920 )
1921 # HA provider not available - remove OAuth provider if present
1922 elif "homeassistant" in self.login_providers:
1923 del self.login_providers["homeassistant"]
1924 self.logger.info("Home Assistant OAuth provider removed (HA provider not available)")
1925
1926 async def _has_non_system_users(self) -> bool:
1927 """Check if any non-system users exist."""
1928 user_rows = await self.database.get_rows("users", limit=10)
1929 return any(row["username"] != HOMEASSISTANT_SYSTEM_USER for row in user_rows)
1930
1931 async def _migrate_system_user_role(self) -> None:
1932 """Migrate the Home Assistant system user of pre-existing installs to the service role."""
1933 user_row = await self.database.get_row(
1934 "users", {"username": normalize_username(HOMEASSISTANT_SYSTEM_USER)}
1935 )
1936 if user_row and user_row["role"] != UserRole.SERVICE.value:
1937 await self.database.update(
1938 "users", {"user_id": user_row["user_id"]}, {"role": UserRole.SERVICE.value}
1939 )
1940 self.logger.info(
1941 "Updated Home Assistant system user role to %s", UserRole.SERVICE.value
1942 )
1943
1944 async def _prune_stale_user_filters(self) -> None:
1945 """Drop user access filter entries for providers or players that no longer exist."""
1946 known_providers = set(self.mass.config.get(CONF_PROVIDERS, {}))
1947 known_players = set(self.mass.config.get(CONF_PLAYERS, {}))
1948 # an empty config section means nothing is configured yet, which must not be
1949 # mistaken for everything having been removed
1950 await self._rewrite_user_filters(
1951 keep_provider=(lambda x: x in known_providers) if known_providers else None,
1952 keep_player=(lambda x: x in known_players) if known_players else None,
1953 )
1954
1955 async def _rewrite_user_filters(
1956 self,
1957 keep_provider: Callable[[str], bool] | None,
1958 keep_player: Callable[[str], bool] | None,
1959 ) -> None:
1960 """Rewrite the access filters of all users, dropping the entries that are not kept."""
1961 if keep_provider is None and keep_player is None:
1962 return
1963 # removing a provider wipes the config of its players one by one, so without the lock
1964 # those rewrites would read the same filter and each undo the other's removal
1965 async with self._user_filter_lock:
1966 for row in await self.database.get_rows("users", limit=0):
1967 updates: dict[str, str] = {}
1968 for column, keep_func in (
1969 ("provider_filter", keep_provider),
1970 ("player_filter", keep_player),
1971 ):
1972 if keep_func is None:
1973 continue
1974 current: list[str] = json_loads(row[column])
1975 remaining = [x for x in current if keep_func(x)]
1976 if remaining == current:
1977 continue
1978 updates[column] = json_dumps(remaining)
1979 dropped = ", ".join(x for x in current if x not in remaining)
1980 if remaining:
1981 self.logger.info(
1982 "Removed %s from the %s of user '%s'", dropped, column, row["username"]
1983 )
1984 else:
1985 # An empty filter means unrestricted. A user whose entries are all gone is
1986 # deliberately left unrestricted, the alternative being an account that
1987 # can see nothing at all.
1988 self.logger.warning(
1989 "Removed the last entries (%s) from the %s of user '%s'. This user is "
1990 "no longer restricted, adjust the access settings if needed.",
1991 dropped,
1992 column,
1993 row["username"],
1994 )
1995 if updates:
1996 await self.database.update("users", {"user_id": row["user_id"]}, updates)
1997
1998 async def _migrate_playlog_to_first_user(self, user_id: str) -> None:
1999 """
2000 Migrate all existing playlog entries to the first user.
2001
2002 This is called automatically when the first non-system user is created.
2003 All existing playlog entries (which have NULL userid) will be updated
2004 to belong to this first user.
2005
2006 :param user_id: The user ID of the first user.
2007 """
2008 try:
2009 # Update all playlog entries with NULL userid to this user
2010 await self.mass.music.database.execute(
2011 f"UPDATE {DB_TABLE_PLAYLOG} SET userid = :userid WHERE userid IS NULL",
2012 {"userid": user_id},
2013 )
2014 await self.mass.music.database.commit()
2015 self.logger.info("Migrated existing playlog entries to first user: %s", user_id)
2016 except Exception as err:
2017 self.logger.warning("Failed to migrate playlog entries: %s", err)
2018
2019 async def _update_profile_password(
2020 self,
2021 target_user: User,
2022 password: str,
2023 is_admin_update: bool,
2024 current_user: User,
2025 ) -> None:
2026 """Update user password (helper method)."""
2027 if len(password) < 8:
2028 raise InvalidDataError("Password must be at least 8 characters")
2029
2030 builtin_provider = self.login_providers.get("builtin")
2031 if not builtin_provider or not isinstance(builtin_provider, BuiltinLoginProvider):
2032 raise InvalidDataError("Built-in auth not available")
2033
2034 # Update password (used for both admin resets and user password changes)
2035 await builtin_provider.reset_password(target_user, password)
2036
2037 if is_admin_update:
2038 self.logger.info(
2039 "Password reset for user %s by admin %s",
2040 target_user.username,
2041 current_user.username,
2042 )
2043 else:
2044 self.logger.info("Password changed for user %s", target_user.username)
2045
2046 async def _check_join_code_rate_limit(self, key: str) -> dict[str, Any] | None:
2047 """
2048 Check the join code exchange throttles that apply to the calling client.
2049
2050 :param key: Rate limit key identifying the calling client.
2051 :return: The error result to return to the caller, or None if the attempt may proceed.
2052 """
2053 limiters = (
2054 ("client", self._join_code_rate_limiter, key),
2055 ("server", self._join_code_global_rate_limiter, JOIN_CODE_GLOBAL_RATE_LIMIT_KEY),
2056 )
2057 for scope, limiter, limiter_key in limiters:
2058 allowed, remaining_delay = await limiter.check_rate_limit(limiter_key)
2059 if allowed:
2060 continue
2061 # The attempted code is deliberately absent here: it has not been checked yet,
2062 # so it may well be a valid one. Each failure that filled the bucket already
2063 # logged its own (rejected, and therefore unusable) code.
2064 self.logger.warning(
2065 "Join code exchange throttled by the %s limit "
2066 "(client=%s, client_failures=%d, server_failures=%d). "
2067 "%d seconds remaining.",
2068 scope,
2069 key,
2070 self._join_code_rate_limiter.get_attempt_count(key),
2071 self._join_code_global_rate_limiter.get_attempt_count(
2072 JOIN_CODE_GLOBAL_RATE_LIMIT_KEY
2073 ),
2074 remaining_delay,
2075 )
2076 return {
2077 "success": False,
2078 "error": (
2079 f"Too many failed attempts. Please try again in {remaining_delay} seconds."
2080 ),
2081 }
2082 return None
2083
2084 async def _exchange_join_code(self, code: str) -> str | None:
2085 """
2086 Exchange a join code for a JWT access token.
2087
2088 The token is created for the user associated with the join code.
2089
2090 :param code: The short join code.
2091 :return: JWT token string if valid, None otherwise.
2092 """
2093 now = utc()
2094
2095 cursor = await self.database.execute(
2096 """
2097 UPDATE join_codes
2098 SET use_count = use_count + 1,
2099 last_used_at = :now
2100 WHERE code = :code
2101 AND expires_at > :now
2102 AND (max_uses = 0 OR use_count < max_uses)
2103 RETURNING user_id, device_name
2104 """,
2105 {"now": now.isoformat(), "code": code.upper()},
2106 )
2107 row = await cursor.fetchone()
2108 await self.database.commit()
2109
2110 if not row:
2111 self.logger.warning(
2112 "Join code exchange rejected (client=%s, code=%s)",
2113 get_current_client_id() or JOIN_CODE_ANONYMOUS_RATE_LIMIT_KEY,
2114 _mask_join_code(code),
2115 )
2116 return None
2117
2118 user = await self.get_user(row["user_id"])
2119 if not user:
2120 self.logger.error(
2121 "User not found for join code despite FK constraint (user_id=%s)", row["user_id"]
2122 )
2123 return None
2124
2125 device_name = row["device_name"] or "Short Code Login"
2126 token = await self.create_token(
2127 user,
2128 device_name,
2129 is_long_lived=False,
2130 )
2131
2132 self.logger.info(
2133 "Join code exchanged for token (user=%s)",
2134 user.username,
2135 )
2136 return token
2137
2138 async def _cleanup_expired_join_codes(self) -> None:
2139 """Delete expired and exhausted join codes from the database."""
2140 now = utc()
2141 cursor = await self.database.execute(
2142 """
2143 DELETE FROM join_codes
2144 WHERE expires_at < :now
2145 OR (max_uses > 0 AND use_count >= max_uses)
2146 """,
2147 {"now": now.isoformat()},
2148 )
2149 await self.database.commit()
2150 count = int(cursor.rowcount)
2151 if count > 0:
2152 self.logger.debug("Cleaned up %d expired/exhausted join code(s)", count)
2153
2154 def _schedule_join_code_cleanup(self) -> None:
2155 """Schedule periodic cleanup of expired join codes."""
2156 self.mass.create_task(self._cleanup_expired_join_codes())
2157 self.mass.call_later(86400, self._schedule_join_code_cleanup)
2158
2159 async def _refresh_token_expiration(
2160 self, token_row: Mapping[str, Any], user: User, is_long_lived: bool
2161 ) -> dict[str, str] | None:
2162 """
2163 Build the on-use column updates for a token, enforcing the absolute lifetime cap.
2164
2165 :param token_row: The auth_tokens row for the token being used.
2166 :param user: The user owning the token.
2167 :param is_long_lived: Whether the token is long-lived.
2168 :return: Column updates to apply (empty when the stored activity timestamp is
2169 still fresh, so callers can skip the write), or None if the token exceeded
2170 its max lifetime (in which case the token row is deleted).
2171 """
2172 now = utc()
2173
2174 if not is_long_lived:
2175 created_at = datetime.fromisoformat(token_row["created_at"])
2176 if now > created_at + timedelta(days=TOKEN_ABSOLUTE_MAX_EXPIRATION):
2177 await self.database.delete("auth_tokens", {"token_id": token_row["token_id"]})
2178 return None
2179
2180 # The HTTP API authenticates on every request, so persisting activity per use
2181 # would cost an UPDATE+commit (an fsync) per request. Skip the write while the
2182 # stored timestamp is fresh; last_used_at and the sliding expiration then lag
2183 # by at most this interval, which is negligible against the 30-day idle window.
2184 if last_used_at := token_row["last_used_at"]:
2185 if now - datetime.fromisoformat(last_used_at) < TOKEN_ACTIVITY_PERSIST_INTERVAL:
2186 return {}
2187
2188 updates = {"last_used_at": now.isoformat()}
2189 if not is_long_lived and user.role != UserRole.GUEST:
2190 # Short-lived token: extend expiration on each use (sliding window)
2191 new_expires_at = now + timedelta(days=TOKEN_SHORT_LIVED_EXPIRATION)
2192 updates["expires_at"] = new_expires_at.isoformat()
2193
2194 return updates
2195
2196 async def _can_reuse_ha_integration_token(self, token: str, system_user: User) -> bool:
2197 """Check whether the stored HA integration token is valid and not yet due for rotation."""
2198 token_id = self.jwt_helper.get_token_id(token)
2199 if not token_id:
2200 return False
2201 token_row = await self.database.get_row("auth_tokens", {"token_id": token_id})
2202 if not token_row or token_row["user_id"] != system_user.user_id:
2203 return False
2204 now = utc()
2205 if token_row["expires_at"] and datetime.fromisoformat(token_row["expires_at"]) <= now:
2206 return False
2207 created_at = datetime.fromisoformat(token_row["created_at"])
2208 rotate_after = created_at + timedelta(
2209 days=TOKEN_ABSOLUTE_MAX_EXPIRATION - HA_TOKEN_ROTATION_MARGIN
2210 )
2211 return now < rotate_after
2212
2213 def _notify_user_access_revoked(self, user: User) -> None:
2214 """Dispatch an access withdrawal to subscribers, isolating them from each other."""
2215 for callback in list(self._access_revoked_callbacks):
2216 self.mass.loop.call_soon(callback, user)
2217
2218
2219def _join_code_rate_limit_key() -> tuple[str, bool]:
2220 """
2221 Work out which bucket the calling client's failed join code exchanges belong to.
2222
2223 :return: The rate limit key, and whether that key identifies a single caller
2224 exclusively (a shared key must never be cleared on a successful exchange).
2225 """
2226 if client_id := get_current_client_id():
2227 return client_id, True
2228 if peer_address := get_current_peer_address():
2229 return f"peer:{peer_address}", False
2230 return JOIN_CODE_ANONYMOUS_RATE_LIMIT_KEY, False
2231
2232
2233def _mask_join_code(code: str) -> str:
2234 """
2235 Mask a join code so support logs can correlate attempts without exposing a usable code.
2236
2237 :param code: The join code as supplied by the client.
2238 :return: The code with everything past its prefix replaced by asterisks.
2239 """
2240 normalized = code.upper()
2241 return normalized[:4] + "*" * max(len(normalized) - 4, 0)
2242