This commit is contained in:
@@ -0,0 +1,89 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hmac
|
||||
import secrets
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from threading import Lock
|
||||
|
||||
from host_agent.local_account import LocalAccountStore
|
||||
|
||||
SESSION_TOKEN_BYTES = 32
|
||||
CSRF_TOKEN_BYTES = 32
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SessionState:
|
||||
username: str
|
||||
csrf_token: str
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class SessionManager:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
ttl_seconds: float,
|
||||
now: Callable[[], datetime] | None = None,
|
||||
) -> None:
|
||||
self._ttl = timedelta(seconds=ttl_seconds)
|
||||
self._now = now or (lambda: datetime.now(UTC))
|
||||
self._lock = Lock()
|
||||
self._sessions: dict[str, SessionState] = {}
|
||||
|
||||
def create_session(self, username: str) -> tuple[str, str]:
|
||||
session_token = secrets.token_urlsafe(SESSION_TOKEN_BYTES)
|
||||
csrf_token = secrets.token_urlsafe(CSRF_TOKEN_BYTES)
|
||||
state = SessionState(
|
||||
username=username,
|
||||
csrf_token=csrf_token,
|
||||
expires_at=self._now() + self._ttl,
|
||||
)
|
||||
with self._lock:
|
||||
self._sessions[session_token] = state
|
||||
return session_token, csrf_token
|
||||
|
||||
def validate(self, session_token: str) -> SessionState | None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
state = self._sessions.get(session_token)
|
||||
if state is None:
|
||||
return None
|
||||
if state.expires_at <= now:
|
||||
del self._sessions[session_token]
|
||||
return None
|
||||
renewed = SessionState(
|
||||
username=state.username,
|
||||
csrf_token=state.csrf_token,
|
||||
expires_at=now + self._ttl,
|
||||
)
|
||||
self._sessions[session_token] = renewed
|
||||
return renewed
|
||||
|
||||
def validate_csrf(self, session_token: str, csrf_token: str) -> bool:
|
||||
state = self.validate(session_token)
|
||||
if state is None:
|
||||
return False
|
||||
return hmac.compare_digest(state.csrf_token, csrf_token)
|
||||
|
||||
def invalidate(self, session_token: str) -> None:
|
||||
with self._lock:
|
||||
self._sessions.pop(session_token, None)
|
||||
|
||||
|
||||
def attempt_login(store: LocalAccountStore, *, username: str, password: str) -> bool:
|
||||
account = store.load()
|
||||
if account is None or account.username != username:
|
||||
return False
|
||||
return store.verify(account, password)
|
||||
|
||||
|
||||
def change_password(
|
||||
store: LocalAccountStore, *, current_password: str, new_password: str
|
||||
) -> bool:
|
||||
account = store.load()
|
||||
if account is None or not store.verify(account, current_password):
|
||||
return False
|
||||
store.create(account.username, new_password)
|
||||
return True
|
||||
Reference in New Issue
Block a user