Faster and simplified hash_secret() that directly produces urlsafe entries.

This commit is contained in:
2026-02-18 23:02:36 +00:00
parent dfc4c76d43
commit 9b7855c0af
7 changed files with 19 additions and 26 deletions
+2 -3
View File
@@ -11,7 +11,6 @@ import secrets
from datetime import UTC, datetime, timedelta from datetime import UTC, datetime, timedelta
from uuid import UUID from uuid import UUID
import base64url
import uuid7 import uuid7
from paskia import oidc_notify from paskia import oidc_notify
@@ -588,7 +587,7 @@ def login(
session = Session.create( session = Session.create(
user=user_uuid, user=user_uuid,
credential=credential_uuid, credential=credential_uuid,
key=base64url.enc(hash_secret("cookie", token)), key=hash_secret("cookie", token),
host=host, host=host,
ip=ip, ip=ip,
user_agent=user_agent, user_agent=user_agent,
@@ -656,7 +655,7 @@ def create_credential_session(
# Generate token and derive key # Generate token and derive key
token = secrets.token_urlsafe(12) token = secrets.token_urlsafe(12)
key = base64url.enc(hash_secret("cookie", token)) key = hash_secret("cookie", token)
session = Session.create( session = Session.create(
user=user_uuid, user=user_uuid,
+7 -8
View File
@@ -5,7 +5,6 @@ import secrets
from datetime import UTC, datetime from datetime import UTC, datetime
from uuid import UUID from uuid import UUID
import base64url
import msgspec import msgspec
import uuid7 import uuid7
@@ -434,7 +433,7 @@ class Session(msgspec.Struct, dict=True, omit_defaults=True):
"""Create a new Session with the provided key. """Create a new Session with the provided key.
Args: Args:
key: The base64url-encoded hashed session key (derived from secret via hash_secret then base64url.enc) key: The hashed session key (derived from secret via hash_secret)
Returns: Returns:
Session object with key set Session object with key set
@@ -471,7 +470,7 @@ class ResetToken(msgspec.Struct, dict=True):
def __post_init__(self): def __post_init__(self):
if not hasattr(self, "key"): if not hasattr(self, "key"):
self.key: bytes = b"" self.key: str = ""
@property @property
def user(self) -> User: def user(self) -> User:
@@ -487,15 +486,15 @@ class ResetToken(msgspec.Struct, dict=True):
del db.data().reset_tokens[self.key] del db.data().reset_tokens[self.key]
@staticmethod @staticmethod
def hash(passphrase: str) -> bytes: def hash(passphrase: str) -> str:
"""Hash a passphrase to bytes for reset token storage.""" """Hash a passphrase to string for reset token storage."""
if not passphrase_util.is_well_formed(passphrase): if not passphrase_util.is_well_formed(passphrase):
raise ValueError( raise ValueError(
"Trying to reset with a session token in place of a passphrase" "Trying to reset with a session token in place of a passphrase"
if len(passphrase) == 16 if len(passphrase) == 16
else "Invalid passphrase format" else "Invalid passphrase format"
) )
return hashlib.sha512(passphrase.encode()).digest()[:9] return hash_secret("reset", passphrase)
@classmethod @classmethod
def by_passphrase(cls, passphrase: str) -> ResetToken | None: def by_passphrase(cls, passphrase: str) -> ResetToken | None:
@@ -627,7 +626,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
users: dict[UUID, User] = {} users: dict[UUID, User] = {}
credentials: dict[UUID, Credential] = {} credentials: dict[UUID, Credential] = {}
sessions: dict[str, Session] = {} sessions: dict[str, Session] = {}
reset_tokens: dict[bytes, ResetToken] = {} reset_tokens: dict[str, ResetToken] = {}
# OIDC provider data # OIDC provider data
oidc: OIDC = msgspec.field(default_factory=lambda: OIDC()) oidc: OIDC = msgspec.field(default_factory=lambda: OIDC())
@@ -670,7 +669,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
SessionContext if valid, None if session not found, expired, or host mismatch SessionContext if valid, None if session not found, expired, or host mismatch
""" """
key = base64url.enc(hash_secret("cookie", session_secret)) key = hash_secret("cookie", session_secret)
try: try:
s = self.sessions[key] s = self.sessions[key]
except KeyError: except KeyError:
+1 -1
View File
@@ -40,7 +40,7 @@ def _oidc_session_by_token(
token: str, client_uuid: UUID | None = None token: str, client_uuid: UUID | None = None
) -> Session | None: ) -> Session | None:
"""Look up an OIDC session by token (refresh token value).""" """Look up an OIDC session by token (refresh token value)."""
key = base64url.enc(hash_secret("oidc", token)) key = hash_secret("oidc", token)
s = db.data().sessions.get(key) s = db.data().sessions.get(key)
if not s or s.client_uuid is None: if not s or s.client_uuid is None:
return None return None
+1 -2
View File
@@ -3,7 +3,6 @@ from datetime import UTC, datetime
from urllib.parse import urlencode from urllib.parse import urlencode
from uuid import UUID from uuid import UUID
import base64url
from fastapi import FastAPI, WebSocket from fastapi import FastAPI, WebSocket
from paskia import authcode, db from paskia import authcode, db
@@ -218,7 +217,7 @@ async def websocket_authenticate(
session = Session.create( session = Session.create(
user=cred.user_uuid, user=cred.user_uuid,
credential=cred.uuid, credential=cred.uuid,
key=base64url.enc(hash_secret("oidc", token)), key=hash_secret("oidc", token),
host=normalized_host, host=normalized_host,
ip=metadata["ip"], ip=metadata["ip"],
user_agent=metadata["user_agent"], user_agent=metadata["user_agent"],
+6 -8
View File
@@ -1,17 +1,15 @@
import hashlib import hashlib
import base64url
from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
def hash_secret(*data) -> bytes: def hash_secret(*data: str | bytes, length=12) -> str:
"""A custom HMAC that securily combines and hashes the given data (context, secrets). The first argument should be a namespacing string.""" """A custom HMAC that securily combines and hashes the given data. The first argument should be a namespacing string."""
inner = bytearray(len(data).to_bytes(8, "big")) p = [d.encode() if hasattr(d, "encode") else d for d in data]
for d in data: p += [len(x).to_bytes(8, "little") for x in [p, *p]]
if isinstance(d, str): return base64url.enc(hashlib.sha256(b"".join(p)).digest()[:length])
d = d.encode()
inner += hashlib.sha256(d).digest()
return hashlib.sha256(inner).digest()[:12]
def secret_key() -> bytes: def secret_key() -> bytes:
+1 -2
View File
@@ -19,7 +19,6 @@ from collections.abc import AsyncGenerator
from datetime import UTC, datetime, timedelta from datetime import UTC, datetime, timedelta
from uuid import UUID from uuid import UUID
import base64url
import httpx import httpx
import pytest import pytest
import pytest_asyncio import pytest_asyncio
@@ -268,7 +267,7 @@ def create_test_session(
# Generate token and derive key # Generate token and derive key
token = secrets.token_urlsafe(12) token = secrets.token_urlsafe(12)
key = base64url.enc(hash_secret("cookie", token)) key = hash_secret("cookie", token)
session = Session.create( session = Session.create(
user=user_uuid, user=user_uuid,
+1 -2
View File
@@ -16,7 +16,6 @@ import secrets
from datetime import UTC, datetime from datetime import UTC, datetime
from uuid import UUID from uuid import UUID
import base64url
import httpx import httpx
import pytest import pytest
import pytest_asyncio import pytest_asyncio
@@ -1301,7 +1300,7 @@ class TestAdminSessions:
test_user, test_user,
): ):
"""Admin can delete their own current session.""" """Admin can delete their own current session."""
session_db_key = base64url.enc(hash_secret("cookie", session_token)) session_db_key = hash_secret("cookie", session_token)
response = await client.delete( response = await client.delete(
f"/auth/api/admin/users/{test_user.uuid}/sessions/{session_db_key}", f"/auth/api/admin/users/{test_user.uuid}/sessions/{session_db_key}",
headers={**auth_headers(session_token), "Host": "localhost:4401"}, headers={**auth_headers(session_token), "Host": "localhost:4401"},