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