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