/
/
/
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 # the lock the automatic rewrites take as well, so a player or provider that is
1252 # being removed cannot overwrite the filters an admin just saved
1253 async with self._user_filter_lock:
1254 await self.database.update("users", {"user_id": target_user.user_id}, updates)
1255 self.webserver.update_active_user_filters(
1256 target_user.user_id,
1257 player_filter=player_filter,
1258 provider_filter=provider_filter,
1259 )
1260 # Refresh target user to get updated filters
1261 refreshed_user = await self.get_user(target_user.user_id)
1262 if not refreshed_user:
1263 raise InvalidDataError("Failed to refresh user after filter update")
1264 return refreshed_user
1265 return target_user
1266
1267 async def remove_from_user_filters(
1268 self,
1269 provider_instance_ids: Collection[str] = (),
1270 player_ids: Collection[str] = (),
1271 ) -> None:
1272 """
1273 Remove the given providers and/or players from the access filters of all users.
1274
1275 Call this when a provider or player is permanently removed, so no user is left with
1276 an access filter that points at something that no longer exists.
1277
1278 :param provider_instance_ids: Instance IDs of the removed providers.
1279 :param player_ids: IDs of the removed players.
1280 """
1281 await self._rewrite_user_filters(
1282 keep_provider=(lambda x: x not in provider_instance_ids)
1283 if provider_instance_ids
1284 else None,
1285 keep_player=(lambda x: x not in player_ids) if player_ids else None,
1286 )
1287
1288 async def replace_player_in_user_filters(
1289 self,
1290 old_player_id: str,
1291 new_player_id: str,
1292 removed_player_ids: Collection[str] = (),
1293 ) -> None:
1294 """
1295 Point the access filters of all users at the replacement of a removed player.
1296
1297 Call this when a player is automatically replaced by another one, so a user that
1298 is restricted to the old player follows the replacement instead of silently
1299 ending up with access to every player.
1300
1301 :param old_player_id: ID of the player that is replaced.
1302 :param new_player_id: ID of the player that takes its place, must not be one of
1303 the removed players.
1304 :param removed_player_ids: IDs of all players whose config is removed, which
1305 normally includes the replaced player itself.
1306 """
1307 await self._rewrite_user_filters(
1308 keep_provider=None,
1309 keep_player=(lambda x: x not in removed_player_ids) if removed_player_ids else None,
1310 map_player=lambda x: new_player_id if x == old_player_id else x,
1311 )
1312
1313 @api_command("auth/user/update")
1314 async def update_user_profile(
1315 self,
1316 user_id: str | None = None,
1317 username: str | None = None,
1318 display_name: str | None = None,
1319 avatar_url: str | None = None,
1320 password: str | None = None,
1321 role: str | None = None,
1322 preferences: dict[str, Any] | None = None,
1323 player_filter: list[str] | None = None,
1324 provider_filter: list[str] | None = None,
1325 ) -> User:
1326 """
1327 Update user profile information.
1328
1329 Users can update their own profile. Admins can update any user including role and password.
1330
1331 :param user_id: User ID to update (optional, defaults to current user).
1332 :param username: New username (optional).
1333 :param display_name: New display name (optional).
1334 :param avatar_url: New avatar URL (optional).
1335 :param password: New password (optional, minimum 8 characters).
1336 :param role: New role - "admin" or "user" (optional, set by admin only).
1337 :param preferences: User preferences dict (completely replaces existing, optional).
1338 :param player_filter: List of player IDs user has access to (set by admin only, optional).
1339 :param provider_filter: List of provider instance IDs user has access to (set by admin only, optional).
1340 :return: Updated user object.
1341 """
1342 current_user_obj = get_current_user()
1343 if not current_user_obj:
1344 raise AuthenticationRequired("Not authenticated")
1345
1346 # Determine target user
1347 may_manage_users = has_scope(current_user_obj, Scope.USERS_MANAGE)
1348 if user_id and user_id != current_user_obj.user_id:
1349 # Updating another user - requires the users.manage scope
1350 if not may_manage_users:
1351 raise InsufficientPermissions(
1352 "The users.manage scope is required to update other users"
1353 )
1354 target_user = await self.get_user(user_id)
1355 if not target_user:
1356 raise InvalidDataError("User not found")
1357 else:
1358 # Updating own profile
1359 target_user = current_user_obj
1360
1361 # Update role (requires the users.manage scope)
1362 if role:
1363 if not may_manage_users:
1364 raise InsufficientPermissions(
1365 "The users.manage scope is required to update user roles"
1366 )
1367
1368 try:
1369 new_role = UserRole(role)
1370 except ValueError as err:
1371 raise InvalidDataError("Invalid role. Must be 'admin' or 'user'") from err
1372
1373 success = await self.update_user_role(target_user.user_id, new_role, current_user_obj)
1374 if not success:
1375 raise InvalidDataError("Failed to update role")
1376
1377 # Refresh target user to get updated role
1378 refreshed_user = await self.get_user(target_user.user_id)
1379 if not refreshed_user:
1380 raise InvalidDataError("Failed to refresh user after role update")
1381 target_user = refreshed_user
1382
1383 # Update basic profile fields
1384 if username or display_name or avatar_url:
1385 updated_user = await self.update_user(
1386 target_user,
1387 username=username,
1388 display_name=display_name,
1389 avatar_url=avatar_url,
1390 )
1391 if not updated_user:
1392 raise InvalidDataError("Failed to update user profile")
1393 target_user = updated_user
1394
1395 # Update preferences if provided
1396 if preferences is not None:
1397 target_user = await self.update_user_preferences(target_user, preferences)
1398
1399 # Update player_filter and provider_filter (requires the users.manage scope)
1400 if player_filter is not None or provider_filter is not None:
1401 if not may_manage_users:
1402 raise InsufficientPermissions(
1403 "The users.manage scope is required to update player/provider filters"
1404 )
1405 target_user = await self.update_user_filters(
1406 target_user, player_filter, provider_filter
1407 )
1408
1409 # Update password if provided
1410 if password:
1411 await self._update_profile_password(
1412 target_user, password, may_manage_users, current_user_obj
1413 )
1414
1415 return target_user
1416
1417 @api_command("auth/logout")
1418 async def logout(self) -> None:
1419 """Logout current user by revoking the current token."""
1420 user = get_current_user()
1421 if not user:
1422 raise AuthenticationRequired("Not authenticated")
1423
1424 # Get current token from context
1425 token = get_current_token()
1426 if not token:
1427 raise InvalidDataError("No token in context")
1428
1429 # Find and revoke the token
1430 token_hash = hashlib.sha256(token.encode()).hexdigest()
1431 token_row = await self.database.get_row("auth_tokens", {"token_hash": token_hash})
1432 if token_row:
1433 await self.database.delete("auth_tokens", {"token_id": token_row["token_id"]})
1434
1435 # Disconnect any WebSocket connections using this token
1436 self.webserver.disconnect_websockets_for_token(token_row["token_id"])
1437
1438 self.logger.info("User '%s' logged out", user.username)
1439
1440 @api_command("auth/user/providers")
1441 async def get_my_providers(self) -> list[dict[str, Any]]:
1442 """
1443 Get current user's linked authentication providers.
1444
1445 :return: List of provider links.
1446 """
1447 user = get_current_user()
1448 if not user:
1449 return []
1450
1451 # Get provider links from database
1452 rows = await self.database.get_rows("user_auth_providers", {"user_id": user.user_id})
1453 providers = [UserAuthProvider.from_dict(dict(row)) for row in rows]
1454 return [p.to_dict() for p in providers]
1455
1456 @api_command("auth/user/unlink_provider", required_scope=Scope.USERS_MANAGE)
1457 async def unlink_provider(self, user_id: str, provider_type: str) -> bool:
1458 """
1459 Unlink authentication provider from user (admin only).
1460
1461 :param user_id: The user ID.
1462 :param provider_type: Provider type to unlink.
1463 :return: True if successful.
1464 """
1465 await self.database.delete(
1466 "user_auth_providers", {"user_id": user_id, "provider_type": provider_type}
1467 )
1468 await self.database.commit()
1469
1470 self.logger.info(
1471 "Auth provider '%s' unlinked from user (user_id=%s)",
1472 provider_type,
1473 user_id,
1474 )
1475 return True
1476
1477 # ==================== Join Code Methods ====================
1478
1479 async def generate_join_code(
1480 self,
1481 user: User,
1482 expires_in_hours: int = JOIN_CODE_DEFAULT_EXPIRY_HOURS,
1483 max_uses: int = 1,
1484 device_name: str = "Short Code Login",
1485 ) -> tuple[str, datetime]:
1486 """
1487 Generate a short join code for link/QR-based login.
1488
1489 This creates a short alphanumeric code that can be exchanged for a JWT token.
1490 Used for features like the party provider guest access, device pairing,
1491 or other short-code authentication flows.
1492
1493 :param user: The guest user that tokens created from this code will belong to.
1494 :param expires_in_hours: Hours until code expires (default: 8).
1495 :param max_uses: Maximum number of uses (0 = unlimited).
1496 :param device_name: Device name for tokens created with this code.
1497 :return: Tuple of (code, expires_at datetime).
1498 """
1499 if expires_in_hours <= 0:
1500 raise ValueError("expires_in_hours must be positive")
1501 if max_uses < 0:
1502 raise ValueError("max_uses must be non-negative (0 = unlimited)")
1503 if user.role != UserRole.GUEST:
1504 raise ValueError("Join codes can only be generated for guest accounts")
1505
1506 now = utc()
1507 expires_at = now + timedelta(hours=expires_in_hours)
1508
1509 for _ in range(3): # Try up to 3 times to avoid code collisions
1510 code = "".join(secrets.choice(JOIN_CODE_CHARSET) for _ in range(JOIN_CODE_LENGTH))
1511 code_data = {
1512 "code_id": secrets.token_urlsafe(32),
1513 "code": code,
1514 "user_id": user.user_id,
1515 "created_at": now.isoformat(),
1516 "expires_at": expires_at.isoformat(),
1517 "max_uses": max_uses,
1518 "use_count": 0,
1519 "device_name": device_name,
1520 }
1521 try:
1522 await self.database.insert("join_codes", code_data)
1523 await self.database.commit()
1524 self.logger.info(
1525 "Join code generated for user %s (expires: %s, max_uses: %s)",
1526 user.username,
1527 expires_at,
1528 max_uses,
1529 )
1530 return code, expires_at
1531 except IntegrityError:
1532 self.logger.warning("Join code collision, retrying...")
1533 continue
1534
1535 raise RuntimeError("Failed to generate a unique join code after 3 attempts")
1536
1537 async def revoke_join_codes(self, user: User) -> int:
1538 """
1539 Revoke all join codes for a user.
1540
1541 :param user: The user whose join codes should be revoked.
1542 :return: Number of codes revoked.
1543 """
1544 cursor = await self.database.execute(
1545 "DELETE FROM join_codes WHERE user_id = :user_id",
1546 {"user_id": user.user_id},
1547 )
1548 await self.database.commit()
1549
1550 count = int(cursor.rowcount)
1551 if count > 0:
1552 self.logger.info("Revoked %d join code(s) for user %s", count, user.username)
1553 return count
1554
1555 async def get_active_join_code(self, user: User) -> str | None:
1556 """
1557 Get the most recently created, non-expired join code for a user.
1558
1559 :param user: The user to look up codes for.
1560 :return: The join code string if found, None otherwise.
1561 """
1562 now = utc()
1563 cursor = await self.database.execute(
1564 """
1565 SELECT code FROM join_codes
1566 WHERE user_id = :user_id
1567 AND expires_at > :now
1568 AND (max_uses = 0 OR use_count < max_uses)
1569 ORDER BY created_at DESC
1570 LIMIT 1
1571 """,
1572 {"user_id": user.user_id, "now": now.isoformat()},
1573 )
1574 row = await cursor.fetchone()
1575 return str(row["code"]) if row else None
1576
1577 async def get_join_code_expiry(self, code: str, user: User | None = None) -> datetime | None:
1578 """
1579 Get the expiry datetime for an active join code.
1580
1581 :param code: The join code to look up.
1582 :param user: Optional user that must own the join code.
1583 :return: The expiry datetime if the code is active, None otherwise.
1584 """
1585 query = """
1586 SELECT expires_at FROM join_codes
1587 WHERE code = :code
1588 AND expires_at > :now
1589 AND (max_uses = 0 OR use_count < max_uses)
1590 """
1591 params: dict[str, Any] = {"code": code.upper(), "now": utc().isoformat()}
1592 if user is not None:
1593 query += "AND user_id = :user_id "
1594 params["user_id"] = user.user_id
1595 cursor = await self.database.execute(query + "LIMIT 1", params)
1596 row = await cursor.fetchone()
1597 return datetime.fromisoformat(str(row["expires_at"])) if row else None
1598
1599 @api_command("auth/join_code/exchange", authenticated=False)
1600 async def exchange_join_code(self, code: str) -> dict[str, Any]:
1601 """
1602 Exchange a join code for an access token (public API).
1603
1604 This is the public API endpoint for short-code authentication.
1605 Clients call this with a code (e.g., from QR scan or link) to receive a JWT token.
1606
1607 :param code: The short join code.
1608 :return: Authentication result with access token if successful.
1609 """
1610 rate_limit_key, key_is_exclusive = _join_code_rate_limit_key()
1611 async with self._join_code_exchange_lock:
1612 if throttled := await self._check_join_code_rate_limit(rate_limit_key):
1613 return throttled
1614
1615 token = await self._exchange_join_code(code)
1616
1617 if not token:
1618 await self._join_code_rate_limiter.record_failed_attempt(rate_limit_key)
1619 await self._join_code_global_rate_limiter.record_failed_attempt(
1620 JOIN_CODE_GLOBAL_RATE_LIMIT_KEY
1621 )
1622 return {
1623 "success": False,
1624 "error": "Invalid or expired join code",
1625 }
1626
1627 # A bucket is only cleared when it belongs to one caller alone, so presenting a
1628 # valid code never lifts the throttle for anyone else.
1629 if key_is_exclusive:
1630 await self._join_code_rate_limiter.clear_attempts(rate_limit_key)
1631
1632 # Decode token to get user info
1633 try:
1634 payload = self.jwt_helper.decode_token(token)
1635 return {
1636 "success": True,
1637 "access_token": token,
1638 "user": {
1639 "user_id": payload.get("sub"),
1640 "username": payload.get("username"),
1641 "role": payload.get("role"),
1642 },
1643 }
1644 except pyjwt.InvalidTokenError:
1645 return {
1646 "success": False,
1647 "error": "Failed to create access token",
1648 }
1649
1650 @api_command("auth/join_codes", required_scope=Scope.USERS_MANAGE)
1651 async def list_join_codes(self, user_id: str | None = None) -> list[dict[str, Any]]:
1652 """
1653 List join codes, optionally filtered by user (admin only).
1654
1655 :param user_id: Optional user ID to filter codes for.
1656 :return: List of join code records.
1657 """
1658 filter_args = {"user_id": user_id} if user_id else None
1659 rows = await self.database.get_rows("join_codes", filter_args, limit=100)
1660 return [dict(row) for row in rows]
1661
1662 @api_command("auth/join_code/revoke", required_scope=Scope.USERS_MANAGE)
1663 async def revoke_join_code(self, code_id: str) -> None:
1664 """
1665 Revoke a specific join code (admin only).
1666
1667 :param code_id: The code ID to revoke.
1668 """
1669 code_row = await self.database.get_row("join_codes", {"code_id": code_id})
1670 if not code_row:
1671 raise InvalidDataError("Join code not found")
1672
1673 await self.database.delete("join_codes", {"code_id": code_id})
1674 await self.database.commit()
1675 self.logger.info("Join code revoked (code_id=%s)", code_id)
1676
1677 async def _setup_database(self) -> None:
1678 """Set up database schema and handle migrations."""
1679 # Always create tables if they don't exist
1680 await self._create_database_tables()
1681
1682 # Check current schema version
1683 try:
1684 if db_row := await self.database.get_row("settings", {"key": "schema_version"}):
1685 prev_version = int(db_row["value"])
1686 else:
1687 prev_version = DB_SCHEMA_VERSION
1688 except KeyError, ValueError, Exception:
1689 # settings table doesn't exist yet or other error
1690 prev_version = 0
1691
1692 # Perform migration if needed
1693 if prev_version < DB_SCHEMA_VERSION:
1694 self.logger.warning(
1695 "Performing database migration from schema version %s to %s",
1696 prev_version,
1697 DB_SCHEMA_VERSION,
1698 )
1699 await self._migrate_database(prev_version)
1700
1701 # Store current schema version
1702 await self.database.insert_or_replace(
1703 "settings",
1704 {"key": "schema_version", "value": str(DB_SCHEMA_VERSION), "type": "int"},
1705 )
1706
1707 # Create indexes
1708 await self._create_database_indexes()
1709 await self.database.commit()
1710
1711 async def _create_database_tables(self) -> None:
1712 """Create database tables."""
1713 # Settings table (for schema version and other settings)
1714 await self.database.execute(
1715 """
1716 CREATE TABLE IF NOT EXISTS settings (
1717 key TEXT PRIMARY KEY,
1718 value TEXT,
1719 type TEXT
1720 )
1721 """
1722 )
1723 # Users table
1724 await self.database.execute(
1725 """
1726 CREATE TABLE IF NOT EXISTS users (
1727 user_id TEXT PRIMARY KEY,
1728 username TEXT NOT NULL UNIQUE,
1729 role TEXT NOT NULL,
1730 enabled INTEGER NOT NULL DEFAULT 1,
1731 created_at TEXT NOT NULL,
1732 display_name TEXT,
1733 avatar_url TEXT,
1734 preferences json NOT NULL DEFAULT '{}',
1735 player_filter json NOT NULL DEFAULT '[]',
1736 provider_filter json NOT NULL DEFAULT '[]'
1737 )
1738 """
1739 )
1740 # User auth provider links (many-to-many)
1741 await self.database.execute(
1742 """
1743 CREATE TABLE IF NOT EXISTS user_auth_providers (
1744 link_id TEXT PRIMARY KEY,
1745 user_id TEXT NOT NULL,
1746 provider_type TEXT NOT NULL,
1747 provider_user_id TEXT NOT NULL,
1748 created_at TEXT NOT NULL,
1749 UNIQUE(provider_type, provider_user_id),
1750 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1751 )
1752 """
1753 )
1754 # Auth tokens table
1755 await self.database.execute(
1756 """
1757 CREATE TABLE IF NOT EXISTS auth_tokens (
1758 token_id TEXT PRIMARY KEY,
1759 user_id TEXT NOT NULL,
1760 token_hash TEXT NOT NULL UNIQUE,
1761 name TEXT NOT NULL,
1762 created_at TEXT NOT NULL,
1763 expires_at TEXT,
1764 last_used_at TEXT,
1765 is_long_lived INTEGER NOT NULL DEFAULT 0,
1766 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1767 )
1768 """
1769 )
1770 # Join codes table (for short code to JWT exchange, used by providers like party)
1771 await self.database.execute(
1772 """
1773 CREATE TABLE IF NOT EXISTS join_codes (
1774 code_id TEXT PRIMARY KEY,
1775 code TEXT NOT NULL UNIQUE,
1776 user_id TEXT NOT NULL,
1777 created_at TEXT NOT NULL,
1778 expires_at TEXT NOT NULL,
1779 max_uses INTEGER DEFAULT 0,
1780 use_count INTEGER DEFAULT 0,
1781 last_used_at TEXT,
1782 device_name TEXT,
1783 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1784 )
1785 """
1786 )
1787 await self.database.commit()
1788
1789 async def _create_database_indexes(self) -> None:
1790 """Create database indexes."""
1791 await self.database.execute(
1792 "CREATE INDEX IF NOT EXISTS idx_user_auth_providers_user "
1793 "ON user_auth_providers(user_id)"
1794 )
1795 await self.database.execute(
1796 "CREATE INDEX IF NOT EXISTS idx_user_auth_providers_provider "
1797 "ON user_auth_providers(provider_type, provider_user_id)"
1798 )
1799 await self.database.execute(
1800 "CREATE INDEX IF NOT EXISTS idx_tokens_user ON auth_tokens(user_id)"
1801 )
1802 await self.database.execute(
1803 "CREATE INDEX IF NOT EXISTS idx_tokens_hash ON auth_tokens(token_hash)"
1804 )
1805 await self.database.execute(
1806 "CREATE INDEX IF NOT EXISTS idx_join_codes_user ON join_codes(user_id)"
1807 )
1808
1809 async def _migrate_database(self, from_version: int) -> None:
1810 """
1811 Perform database migration.
1812
1813 :param from_version: The schema version to migrate from.
1814 """
1815 self.logger.info(
1816 "Migrating auth database from version %s to %s", from_version, DB_SCHEMA_VERSION
1817 )
1818 # Migration to version 2: Recreate tables due to password salt breaking change
1819 if from_version < 2:
1820 # Drop all auth-related tables
1821 await self.database.execute("DROP TABLE IF EXISTS auth_tokens")
1822 await self.database.execute("DROP TABLE IF EXISTS user_auth_providers")
1823 await self.database.execute("DROP TABLE IF EXISTS users")
1824 await self.database.commit()
1825
1826 # Recreate tables with current schema
1827 await self._create_database_tables()
1828
1829 # Migration to version 3: Add player_filter and provider_filter columns
1830 if from_version < 3:
1831 with contextlib.suppress(OperationalError):
1832 # Column(s) may already exist
1833 await self.database.execute(
1834 "ALTER TABLE users ADD COLUMN player_filter json NOT NULL DEFAULT '[]'"
1835 )
1836 await self.database.execute(
1837 "ALTER TABLE users ADD COLUMN provider_filter json NOT NULL DEFAULT '[]'"
1838 )
1839 await self.database.commit()
1840
1841 # Migration to version 4: Make usernames case-insensitive by converting to lowercase
1842 if from_version < 4:
1843 await self.database.execute("UPDATE users SET username = LOWER(username)")
1844 await self.database.commit()
1845
1846 # Migration to version 5: Add join codes table
1847 if from_version < 5:
1848 await self.database.execute(
1849 """
1850 CREATE TABLE IF NOT EXISTS join_codes (
1851 code_id TEXT PRIMARY KEY,
1852 code TEXT NOT NULL UNIQUE,
1853 user_id TEXT NOT NULL,
1854 created_at TEXT NOT NULL,
1855 expires_at TEXT NOT NULL,
1856 max_uses INTEGER DEFAULT 0,
1857 use_count INTEGER DEFAULT 0,
1858 last_used_at TEXT,
1859 device_name TEXT,
1860 FOREIGN KEY (user_id) REFERENCES users(user_id) ON DELETE CASCADE
1861 )
1862 """
1863 )
1864 await self.database.commit()
1865
1866 async def _get_or_create_jwt_secret(self) -> str:
1867 """
1868 Get or create JWT secret key from database.
1869
1870 :return: JWT secret key for signing tokens.
1871 """
1872 # Try to get existing secret
1873 if secret_row := await self.database.get_row("settings", {"key": "jwt_secret"}):
1874 return str(secret_row["value"])
1875
1876 # Generate new secret
1877 jwt_secret = JWTHelper.generate_secret_key()
1878
1879 # Store in database
1880 await self.database.insert_or_replace(
1881 "settings",
1882 {"key": "jwt_secret", "value": jwt_secret, "type": "string"},
1883 )
1884 await self.database.commit()
1885
1886 self.logger.info("Generated new JWT secret key")
1887 return jwt_secret
1888
1889 async def _setup_login_providers(self) -> None:
1890 """Set up available login providers based on configuration."""
1891 # Always enable built-in provider
1892 self.login_providers["builtin"] = BuiltinLoginProvider(self.mass, "builtin", {})
1893
1894 # Home Assistant OAuth provider
1895 # Automatically enabled if HA provider (plugin) is configured
1896 ha_provider = None
1897 for provider in self.mass.providers:
1898 if provider.domain == "hass" and provider.available:
1899 ha_provider = provider
1900 break
1901
1902 if ha_provider:
1903 ha_provider = cast("HomeAssistantProvider", ha_provider)
1904 ha_url = ha_provider.url
1905 if not ha_url:
1906 self.logger.warning(
1907 "Home Assistant provider has no URL configured, "
1908 "Home Assistant OAuth login is not available"
1909 )
1910 return
1911 ha_config: HomeAssistantProviderConfig = {"ha_url": ha_url}
1912 self.login_providers["homeassistant"] = HomeAssistantOAuthProvider(
1913 self.mass, "homeassistant", ha_config
1914 )
1915 self.logger.info(
1916 "Home Assistant OAuth provider enabled (using URL from HA provider: %s)",
1917 ha_url,
1918 )
1919
1920 async def _sync_ha_oauth_provider(self) -> None:
1921 """
1922 Sync HA OAuth provider with HA provider availability (dynamic check).
1923
1924 Adds the provider if HA is available, removes it if HA is not available.
1925 """
1926 # Find HA provider
1927 ha_provider = None
1928 for provider in self.mass.providers:
1929 if provider.domain == "hass" and provider.available:
1930 ha_provider = provider
1931 break
1932
1933 if ha_provider:
1934 # HA provider exists and is available - ensure OAuth provider is registered
1935 if "homeassistant" not in self.login_providers:
1936 ha_provider = cast("HomeAssistantProvider", ha_provider)
1937 ha_url = ha_provider.url
1938 if not ha_url:
1939 # missing URL must never break the login providers endpoint,
1940 # simply leave the HA OAuth provider unregistered
1941 self.logger.debug(
1942 "Home Assistant provider has no URL configured, "
1943 "Home Assistant OAuth login is not available"
1944 )
1945 return
1946 ha_config: HomeAssistantProviderConfig = {"ha_url": ha_url}
1947 self.login_providers["homeassistant"] = HomeAssistantOAuthProvider(
1948 self.mass, "homeassistant", ha_config
1949 )
1950 self.logger.info(
1951 "Home Assistant OAuth provider dynamically enabled (using URL: %s)",
1952 ha_url,
1953 )
1954 # HA provider not available - remove OAuth provider if present
1955 elif "homeassistant" in self.login_providers:
1956 del self.login_providers["homeassistant"]
1957 self.logger.info("Home Assistant OAuth provider removed (HA provider not available)")
1958
1959 async def _has_non_system_users(self) -> bool:
1960 """Check if any non-system users exist."""
1961 user_rows = await self.database.get_rows("users", limit=10)
1962 return any(row["username"] != HOMEASSISTANT_SYSTEM_USER for row in user_rows)
1963
1964 async def _migrate_system_user_role(self) -> None:
1965 """Migrate the Home Assistant system user of pre-existing installs to the service role."""
1966 user_row = await self.database.get_row(
1967 "users", {"username": normalize_username(HOMEASSISTANT_SYSTEM_USER)}
1968 )
1969 if user_row and user_row["role"] != UserRole.SERVICE.value:
1970 await self.database.update(
1971 "users", {"user_id": user_row["user_id"]}, {"role": UserRole.SERVICE.value}
1972 )
1973 self.logger.info(
1974 "Updated Home Assistant system user role to %s", UserRole.SERVICE.value
1975 )
1976
1977 async def _prune_stale_user_filters(self) -> None:
1978 """Drop user access filter entries for providers or players that no longer exist."""
1979 known_providers = set(self.mass.config.get(CONF_PROVIDERS, {}))
1980 known_players = set(self.mass.config.get(CONF_PLAYERS, {}))
1981
1982 # one-off: the connected-player plugins collapsed their instances into a single
1983 # instance keyed by the bare domain; filter entries naming a collapsed instance
1984 # follow it instead of being pruned (which would lift the user's restriction).
1985 # TODO: remove after 2.12 release
1986 def _map_collapsed_plugin(entry: str) -> str:
1987 for domain in ("spotify_connect", "airplay_receiver"):
1988 if entry.startswith(f"{domain}--") and domain in known_providers:
1989 return domain
1990 return entry
1991
1992 # an empty config section means nothing is configured yet, which must not be
1993 # mistaken for everything having been removed
1994 await self._rewrite_user_filters(
1995 keep_provider=(lambda x: x in known_providers) if known_providers else None,
1996 keep_player=(lambda x: x in known_players) if known_players else None,
1997 map_provider=_map_collapsed_plugin if known_providers else None,
1998 )
1999
2000 async def _rewrite_user_filters(
2001 self,
2002 keep_provider: Callable[[str], bool] | None,
2003 keep_player: Callable[[str], bool] | None,
2004 map_player: Callable[[str], str] | None = None,
2005 map_provider: Callable[[str], str] | None = None,
2006 ) -> None:
2007 """
2008 Rewrite the access filters of all users.
2009
2010 :param keep_provider: Returns False for the provider entries that must be dropped.
2011 :param keep_player: Returns False for the player entries that must be dropped.
2012 :param map_player: Maps a player entry onto its replacement, applied before keep_player.
2013 :param map_provider: Maps a provider entry onto its replacement, applied before
2014 keep_provider.
2015 """
2016 if keep_provider is None and keep_player is None and map_player is None:
2017 return
2018 # removing a provider wipes the config of its players one by one, so without the lock
2019 # those rewrites would read the same filter and each undo the other's removal
2020 async with self._user_filter_lock:
2021 for row in await self.database.get_rows("users", limit=0):
2022 changed: dict[str, list[str]] = {}
2023 for column, keep_func, map_func in (
2024 ("provider_filter", keep_provider, map_provider),
2025 ("player_filter", keep_player, map_player),
2026 ):
2027 if keep_func is None and map_func is None:
2028 continue
2029 current: list[str] = json_loads(row[column])
2030 remaining: list[str] = []
2031 dropped: list[str] = []
2032 for entry in current:
2033 mapped = map_func(entry) if map_func else entry
2034 if keep_func and not keep_func(mapped):
2035 dropped.append(entry)
2036 elif mapped not in remaining:
2037 remaining.append(mapped)
2038 if remaining == current:
2039 continue
2040 changed[column] = remaining
2041 if not dropped:
2042 self.logger.info(
2043 "Updated the %s of user '%s' to %s",
2044 column,
2045 row["username"],
2046 ", ".join(remaining),
2047 )
2048 elif remaining:
2049 self.logger.info(
2050 "Removed %s from the %s of user '%s'",
2051 ", ".join(dropped),
2052 column,
2053 row["username"],
2054 )
2055 else:
2056 # An empty filter means unrestricted. A user whose entries are all gone is
2057 # deliberately left unrestricted, the alternative being an account that
2058 # can see nothing at all.
2059 self.logger.warning(
2060 "Removed the last entries (%s) from the %s of user '%s'. This user is "
2061 "no longer restricted, adjust the access settings if needed.",
2062 ", ".join(dropped),
2063 column,
2064 row["username"],
2065 )
2066 if changed:
2067 await self.database.update(
2068 "users",
2069 {"user_id": row["user_id"]},
2070 {column: json_dumps(value) for column, value in changed.items()},
2071 )
2072 # a session holds its own copy of the User object, so the live ones have to
2073 # follow or they keep applying the filter that was just rewritten
2074 self.webserver.update_active_user_filters(
2075 row["user_id"],
2076 player_filter=changed.get("player_filter"),
2077 provider_filter=changed.get("provider_filter"),
2078 )
2079
2080 async def _migrate_playlog_to_first_user(self, user_id: str) -> None:
2081 """
2082 Migrate all existing playlog entries to the first user.
2083
2084 This is called automatically when the first non-system user is created.
2085 All existing playlog entries (which have NULL userid) will be updated
2086 to belong to this first user.
2087
2088 :param user_id: The user ID of the first user.
2089 """
2090 try:
2091 # Update all playlog entries with NULL userid to this user
2092 await self.mass.music.database.execute(
2093 f"UPDATE {DB_TABLE_PLAYLOG} SET userid = :userid WHERE userid IS NULL",
2094 {"userid": user_id},
2095 )
2096 await self.mass.music.database.commit()
2097 self.logger.info("Migrated existing playlog entries to first user: %s", user_id)
2098 except Exception as err:
2099 self.logger.warning("Failed to migrate playlog entries: %s", err)
2100
2101 async def _update_profile_password(
2102 self,
2103 target_user: User,
2104 password: str,
2105 is_admin_update: bool,
2106 current_user: User,
2107 ) -> None:
2108 """Update user password (helper method)."""
2109 if len(password) < 8:
2110 raise InvalidDataError("Password must be at least 8 characters")
2111
2112 builtin_provider = self.login_providers.get("builtin")
2113 if not builtin_provider or not isinstance(builtin_provider, BuiltinLoginProvider):
2114 raise InvalidDataError("Built-in auth not available")
2115
2116 # Update password (used for both admin resets and user password changes)
2117 await builtin_provider.reset_password(target_user, password)
2118
2119 if is_admin_update:
2120 self.logger.info(
2121 "Password reset for user %s by admin %s",
2122 target_user.username,
2123 current_user.username,
2124 )
2125 else:
2126 self.logger.info("Password changed for user %s", target_user.username)
2127
2128 async def _check_join_code_rate_limit(self, key: str) -> dict[str, Any] | None:
2129 """
2130 Check the join code exchange throttles that apply to the calling client.
2131
2132 :param key: Rate limit key identifying the calling client.
2133 :return: The error result to return to the caller, or None if the attempt may proceed.
2134 """
2135 limiters = (
2136 ("client", self._join_code_rate_limiter, key),
2137 ("server", self._join_code_global_rate_limiter, JOIN_CODE_GLOBAL_RATE_LIMIT_KEY),
2138 )
2139 for scope, limiter, limiter_key in limiters:
2140 allowed, remaining_delay = await limiter.check_rate_limit(limiter_key)
2141 if allowed:
2142 continue
2143 # The attempted code is deliberately absent here: it has not been checked yet,
2144 # so it may well be a valid one. Each failure that filled the bucket already
2145 # logged its own (rejected, and therefore unusable) code.
2146 self.logger.warning(
2147 "Join code exchange throttled by the %s limit "
2148 "(client=%s, client_failures=%d, server_failures=%d). "
2149 "%d seconds remaining.",
2150 scope,
2151 key,
2152 self._join_code_rate_limiter.get_attempt_count(key),
2153 self._join_code_global_rate_limiter.get_attempt_count(
2154 JOIN_CODE_GLOBAL_RATE_LIMIT_KEY
2155 ),
2156 remaining_delay,
2157 )
2158 return {
2159 "success": False,
2160 "error": (
2161 f"Too many failed attempts. Please try again in {remaining_delay} seconds."
2162 ),
2163 }
2164 return None
2165
2166 async def _exchange_join_code(self, code: str) -> str | None:
2167 """
2168 Exchange a join code for a JWT access token.
2169
2170 The token is created for the user associated with the join code.
2171
2172 :param code: The short join code.
2173 :return: JWT token string if valid, None otherwise.
2174 """
2175 now = utc()
2176
2177 cursor = await self.database.execute(
2178 """
2179 UPDATE join_codes
2180 SET use_count = use_count + 1,
2181 last_used_at = :now
2182 WHERE code = :code
2183 AND expires_at > :now
2184 AND (max_uses = 0 OR use_count < max_uses)
2185 RETURNING user_id, device_name
2186 """,
2187 {"now": now.isoformat(), "code": code.upper()},
2188 )
2189 row = await cursor.fetchone()
2190 await self.database.commit()
2191
2192 if not row:
2193 self.logger.warning(
2194 "Join code exchange rejected (client=%s, code=%s)",
2195 get_current_client_id() or JOIN_CODE_ANONYMOUS_RATE_LIMIT_KEY,
2196 _mask_join_code(code),
2197 )
2198 return None
2199
2200 user = await self.get_user(row["user_id"])
2201 if not user:
2202 self.logger.error(
2203 "User not found for join code despite FK constraint (user_id=%s)", row["user_id"]
2204 )
2205 return None
2206
2207 device_name = row["device_name"] or "Short Code Login"
2208 token = await self.create_token(
2209 user,
2210 device_name,
2211 is_long_lived=False,
2212 )
2213
2214 self.logger.info(
2215 "Join code exchanged for token (user=%s)",
2216 user.username,
2217 )
2218 return token
2219
2220 async def _cleanup_expired_join_codes(self) -> None:
2221 """Delete expired and exhausted join codes from the database."""
2222 now = utc()
2223 cursor = await self.database.execute(
2224 """
2225 DELETE FROM join_codes
2226 WHERE expires_at < :now
2227 OR (max_uses > 0 AND use_count >= max_uses)
2228 """,
2229 {"now": now.isoformat()},
2230 )
2231 await self.database.commit()
2232 count = int(cursor.rowcount)
2233 if count > 0:
2234 self.logger.debug("Cleaned up %d expired/exhausted join code(s)", count)
2235
2236 def _schedule_join_code_cleanup(self) -> None:
2237 """Schedule periodic cleanup of expired join codes."""
2238 self.mass.create_task(self._cleanup_expired_join_codes())
2239 self.mass.call_later(86400, self._schedule_join_code_cleanup)
2240
2241 async def _refresh_token_expiration(
2242 self, token_row: Mapping[str, Any], user: User, is_long_lived: bool
2243 ) -> dict[str, str] | None:
2244 """
2245 Build the on-use column updates for a token, enforcing the absolute lifetime cap.
2246
2247 :param token_row: The auth_tokens row for the token being used.
2248 :param user: The user owning the token.
2249 :param is_long_lived: Whether the token is long-lived.
2250 :return: Column updates to apply (empty when the stored activity timestamp is
2251 still fresh, so callers can skip the write), or None if the token exceeded
2252 its max lifetime (in which case the token row is deleted).
2253 """
2254 now = utc()
2255
2256 if not is_long_lived:
2257 created_at = datetime.fromisoformat(token_row["created_at"])
2258 if now > created_at + timedelta(days=TOKEN_ABSOLUTE_MAX_EXPIRATION):
2259 await self.database.delete("auth_tokens", {"token_id": token_row["token_id"]})
2260 return None
2261
2262 # The HTTP API authenticates on every request, so persisting activity per use
2263 # would cost an UPDATE+commit (an fsync) per request. Skip the write while the
2264 # stored timestamp is fresh; last_used_at and the sliding expiration then lag
2265 # by at most this interval, which is negligible against the 30-day idle window.
2266 if last_used_at := token_row["last_used_at"]:
2267 if now - datetime.fromisoformat(last_used_at) < TOKEN_ACTIVITY_PERSIST_INTERVAL:
2268 return {}
2269
2270 updates = {"last_used_at": now.isoformat()}
2271 if not is_long_lived and user.role != UserRole.GUEST:
2272 # Short-lived token: extend expiration on each use (sliding window)
2273 new_expires_at = now + timedelta(days=TOKEN_SHORT_LIVED_EXPIRATION)
2274 updates["expires_at"] = new_expires_at.isoformat()
2275
2276 return updates
2277
2278 async def _can_reuse_ha_integration_token(self, token: str, system_user: User) -> bool:
2279 """Check whether the stored HA integration token is valid and not yet due for rotation."""
2280 token_id = self.jwt_helper.get_token_id(token)
2281 if not token_id:
2282 return False
2283 token_row = await self.database.get_row("auth_tokens", {"token_id": token_id})
2284 if not token_row or token_row["user_id"] != system_user.user_id:
2285 return False
2286 now = utc()
2287 if token_row["expires_at"] and datetime.fromisoformat(token_row["expires_at"]) <= now:
2288 return False
2289 created_at = datetime.fromisoformat(token_row["created_at"])
2290 rotate_after = created_at + timedelta(
2291 days=TOKEN_ABSOLUTE_MAX_EXPIRATION - HA_TOKEN_ROTATION_MARGIN
2292 )
2293 return now < rotate_after
2294
2295 def _notify_user_access_revoked(self, user: User) -> None:
2296 """Dispatch an access withdrawal to subscribers, isolating them from each other."""
2297 for callback in list(self._access_revoked_callbacks):
2298 self.mass.loop.call_soon(callback, user)
2299
2300
2301def _join_code_rate_limit_key() -> tuple[str, bool]:
2302 """
2303 Work out which bucket the calling client's failed join code exchanges belong to.
2304
2305 :return: The rate limit key, and whether that key identifies a single caller
2306 exclusively (a shared key must never be cleared on a successful exchange).
2307 """
2308 if client_id := get_current_client_id():
2309 return client_id, True
2310 if peer_address := get_current_peer_address():
2311 return f"peer:{peer_address}", False
2312 return JOIN_CODE_ANONYMOUS_RATE_LIMIT_KEY, False
2313
2314
2315def _mask_join_code(code: str) -> str:
2316 """
2317 Mask a join code so support logs can correlate attempts without exposing a usable code.
2318
2319 :param code: The join code as supplied by the client.
2320 :return: The code with everything past its prefix replaced by asterisks.
2321 """
2322 normalized = code.upper()
2323 return normalized[:4] + "*" * max(len(normalized) - 4, 0)
2324