Compare commits

..
16 Commits
Author SHA1 Message Date
LeoVasanko 4e6f63e9ef Fix auth-host being added to origins even when no origins were wanted. 2026-02-19 01:00:13 +00:00
LeoVasanko d431c75297 Fix request header removal that was causing zstd compressed response when we wanted plain text. 2026-02-19 00:51:11 +00:00
LeoVasanko 6257071efe Cleaned Org Admin styling. 2026-02-19 00:37:44 +00:00
LeoVasanko 7958b6f365 Migrations cleanup by using a MigrationCtx object for meta. 2026-02-19 00:08:38 +00:00
LeoVasanko b3cb540098 Implement database file locking for extra safety. 2026-02-19 00:01:49 +00:00
LeoVasanko 22ba7231b1 Implement read-only database load at startup for CLI to get its settings. Full opening only when server has started. 2026-02-18 23:46:18 +00:00
LeoVasanko 9a9979fb62 Add missing new file. 2026-02-18 23:45:22 +00:00
LeoVasanko 9b7855c0af Faster and simplified hash_secret() that directly produces urlsafe entries. 2026-02-18 23:02:36 +00:00
LeoVasanko dfc4c76d43 Fix static asset serving in devserver mode. 2026-02-18 22:19:18 +00:00
LeoVasanko e1f0fdf664 Fix E2E tests 2026-02-18 22:18:31 +00:00
LeoVasanko f26ac8f33b Simplified My Profile authentication flows, fixed some UX issues with reauth cancelled/accepted leading to incorrect states. 2026-02-18 19:04:02 +00:00
LeoVasanko 880ced3b8c Always load user's theme from API if available, and update the localStorage cache. Previously in various situations the old cached value was being used instead, leading to inconsistent theming or wrong themeselector readout. 2026-02-18 18:35:41 +00:00
LeoVasanko af80b5eefc Make ResetApp Registration layout match the other dialog apps (centered). 2026-02-18 18:01:17 +00:00
LeoVasanko fa1e69d58b Imports to top of file. 2026-02-18 17:47:57 +00:00
LeoVasanko 39000ef831 Remove PASKIA_SITE_URL env, use PASKIA_VITE_URL instead, moving the correct site URL determination to paskia CLI directly. This resolves devserver issues where VITE URL was reported as the external site URL instead of configured auth-host or origins. Main CLI still simplified. Bump deps versions and cleanup pyproject.toml. 2026-02-18 17:21:03 +00:00
LeoVasanko 49119fac81 Fix HostProfile user properties access. 2026-02-18 03:47:08 +00:00
47 changed files with 718 additions and 538 deletions
+2 -2
View File
@@ -216,7 +216,7 @@ test.describe('API Mode - 401 Login Flow', () => {
await clearSessionCookie(page) await clearSessionCookie(page)
// Make API call that triggers 401 (don't await - it blocks until iframe resolves) // Make API call that triggers 401 (don't await - it blocks until iframe resolves)
const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'POST').catch(e => e) const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'GET').catch(e => e)
console.log('✓ Auth iframe appeared on 401') console.log('✓ Auth iframe appeared on 401')
// Verify it's in login mode (not reauth) // Verify it's in login mode (not reauth)
@@ -268,7 +268,7 @@ test.describe('API Mode - 401 Login Flow', () => {
await setupTestHarness(page) await setupTestHarness(page)
// Make API call that triggers 401 // Make API call that triggers 401
const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'POST') const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'GET')
// Wait for auth iframe to appear // Wait for auth iframe to appear
await waitForAuthIframe(page) await waitForAuthIframe(page)
+6 -1
View File
@@ -1,4 +1,4 @@
import { spawn } from 'child_process' import { execSync, spawn } from 'child_process'
import { join, dirname } from 'path' import { join, dirname } from 'path'
import { existsSync, mkdirSync, writeFileSync } from 'fs' import { existsSync, mkdirSync, writeFileSync } from 'fs'
import { fileURLToPath } from 'url' import { fileURLToPath } from 'url'
@@ -31,6 +31,11 @@ export default async function globalSetup() {
mkdirSync(testDataDir, { recursive: true }) mkdirSync(testDataDir, { recursive: true })
} }
// Build the package first
console.log(' Building package with uv build...')
execSync('uv build', { cwd: projectRoot, stdio: 'inherit' })
console.log(' ✅ Build complete\n')
console.log(' Starting server with in-memory database...') console.log(' Starting server with in-memory database...')
if (COLLECT_COVERAGE) { if (COLLECT_COVERAGE) {
console.log(' 📊 Coverage collection enabled for Python backend') console.log(' 📊 Coverage collection enabled for Python backend')
+21 -72
View File
@@ -13,8 +13,8 @@
<script setup> <script setup>
import { computed, onMounted, onUnmounted, ref } from 'vue' import { computed, onMounted, onUnmounted, ref } from 'vue'
import { useAuthStore } from '@/stores/auth' import { useAuthStore } from '@/stores/auth'
import { apiJson, SessionValidator, createAuthIframe, removeAuthIframe } from 'paskia' import { apiJson, SessionValidator } from 'paskia'
import { getAuthIframeUrl } from '@/utils/api' import { updateThemeFromSession } from '@/utils/theme'
import StatusMessage from '@/components/StatusMessage.vue' import StatusMessage from '@/components/StatusMessage.vue'
import ProfileView from '@/components/ProfileView.vue' import ProfileView from '@/components/ProfileView.vue'
import HostProfileView from '@/components/HostProfileView.vue' import HostProfileView from '@/components/HostProfileView.vue'
@@ -48,90 +48,49 @@ const isHostMode = computed(() => {
return currentHost !== configuredHost return currentHost !== configuredHost
}) })
function terminateSession() { function onSessionLost(e) {
store.userInfo = null store.userInfo = null
viewState.value = 'terminal' store.ctx = null
if (e?.name === 'AuthCancelledError') {
viewState.value = 'terminal'
} else {
store.showMessage(e?.message || 'Session lost', 'error', 5000)
viewState.value = 'terminal'
}
} }
const userUuidGetter = () => store.ctx?.user.uuid const userUuidGetter = () => store.ctx?.user.uuid
const sessionValidator = new SessionValidator(userUuidGetter, terminateSession) const sessionValidator = new SessionValidator(userUuidGetter, onSessionLost)
onMounted(() => sessionValidator.start()) onMounted(() => sessionValidator.start())
onUnmounted(() => sessionValidator.stop()) onUnmounted(() => sessionValidator.stop())
async function loadUserInfo() { async function loadUserInfo() {
viewState.value = 'loading'
loadingMessage.value = 'Loading...'
try { try {
// apiJson handles 401/403 with auth.iframe automatically:
// shows overlay iframe, waits for auth, retries the request.
const [validateData, userInfoData] = await Promise.all([ const [validateData, userInfoData] = await Promise.all([
apiJson('/auth/api/validate', { method: 'POST' }), apiJson('/auth/api/validate', { method: 'POST' }),
apiJson('/auth/api/user-info', { method: 'GET' }) apiJson('/auth/api/user-info', { method: 'GET' })
]) ])
store.userInfo = userInfoData store.userInfo = userInfoData
store.ctx = validateData.ctx store.ctx = validateData.ctx
updateThemeFromSession(store.userInfo)
// Verify that the user UUIDs match between user-info and validate responses // Verify that the user UUIDs match between user-info and validate responses
if (store.userInfo.user.uuid !== store.ctx.user.uuid) { if (store.userInfo.user.uuid !== store.ctx.user.uuid) {
console.error('User UUID mismatch between user-info and validate responses') console.error('User UUID mismatch between user-info and validate responses')
window.location.reload() window.location.reload()
return false return
} }
viewState.value = 'profile' viewState.value = 'profile'
return true } catch (e) {
} catch { onSessionLost(e)
store.userInfo = null
store.ctx = null
return false
}
}
async function showAuthIframe() {
const url = await getAuthIframeUrl('login')
createAuthIframe(url)
loadingMessage.value = 'Authentication required...'
}
function handleAuthMessage(event) {
const data = event.data
if (!data?.type) return
switch (data.type) {
case 'auth-success':
// Authentication successful - reload user info
removeAuthIframe()
viewState.value = 'loading'
loadingMessage.value = 'Loading user profile...'
loadUserInfo()
break
case 'auth-error':
// Authentication failed - keep iframe open so user can retry
if (data.cancelled) {
console.log('Authentication cancelled by user')
} else {
store.showMessage(data.message || 'Authentication failed', 'error', 5000)
}
break
case 'auth-cancelled':
// Legacy support - treat as auth-error with cancelled flag
console.log('Authentication cancelled')
break
case 'auth-back':
// User clicked Back - show terminal state
removeAuthIframe()
terminateSession()
break
case 'auth-close-request':
// Legacy support - treat as back
removeAuthIframe()
break
} }
} }
onMounted(async () => { onMounted(async () => {
// Listen for postMessage from auth iframe
window.addEventListener('message', handleAuthMessage)
// Load settings // Load settings
await store.loadSettings() await store.loadSettings()
@@ -145,17 +104,7 @@ onMounted(async () => {
document.title = inHostMode ? `${rpName} · Account summary` : rpName document.title = inHostMode ? `${rpName} · Account summary` : rpName
} }
// Try to load user info // Load user info (apiJson handles auth iframe if needed)
const success = await loadUserInfo() await loadUserInfo()
if (!success) {
// Need authentication - show login iframe
showAuthIframe()
}
})
onUnmounted(() => {
window.removeEventListener('message', handleAuthMessage)
removeAuthIframe()
}) })
</script> </script>
+6 -21
View File
@@ -14,6 +14,7 @@ import AdminDialogs from '@/admin/AdminDialogs.vue'
import { useAuthStore } from '@/stores/auth' import { useAuthStore } from '@/stores/auth'
import { adminUiPath, makeUiHref } from '@/utils/settings' import { adminUiPath, makeUiHref } from '@/utils/settings'
import { apiJson, SessionValidator } from 'paskia' import { apiJson, SessionValidator } from 'paskia'
import { updateThemeFromSession } from '@/utils/theme'
import { uuidv7 } from 'uuidv7' import { uuidv7 } from 'uuidv7'
import { getDirection } from '@/utils/keynav' import { getDirection } from '@/utils/keynav'
import { goBack } from '@/utils/helpers' import { goBack } from '@/utils/helpers'
@@ -197,6 +198,7 @@ function orgUserCount(org) {
async function loadUserInfo() { async function loadUserInfo() {
const data = await apiJson('/auth/api/validate', { method: 'POST' }) const data = await apiJson('/auth/api/validate', { method: 'POST' })
info.value = data info.value = data
updateThemeFromSession(data.ctx)
authenticated.value = true authenticated.value = true
} }
@@ -337,24 +339,9 @@ async function moveUserToRole(userUuid, user, targetRoleUuid) {
} }
} }
function onUserDragStart(e, userUuid, org) { function moveUserToRoleFromDrag(userUuid, newRoleUuid) {
e.dataTransfer.effectAllowed = 'move' const user = selectedOrg.value?.users?.[userUuid]
e.dataTransfer.setData('text/plain', JSON.stringify({ user_uuid: userUuid, org })) if (user) moveUserToRole(userUuid, user, newRoleUuid)
}
function onRoleDragOver(e) {
e.preventDefault()
e.dataTransfer.dropEffect = 'move'
}
function onRoleDrop(e, org, role) {
e.preventDefault()
try {
const data = JSON.parse(e.dataTransfer.getData('text/plain'))
if (data.org !== org.uuid) return // only within same org
const user = org.users[data.user_uuid]
if (user) moveUserToRole(data.user_uuid, user, role.uuid)
} catch (_) { /* ignore */ }
} }
// Role actions // Role actions
@@ -999,10 +986,8 @@ async function submitDialog() {
@create-user-in-role="createUserInRole" @create-user-in-role="createUserInRole"
@open-user="openUser" @open-user="openUser"
@toggle-role-permission="toggleRolePermission" @toggle-role-permission="toggleRolePermission"
@on-role-drag-over="onRoleDragOver" @move-user-to-role="moveUserToRoleFromDrag"
@navigate-out="handlePanelNavigateOut" @navigate-out="handlePanelNavigateOut"
@on-role-drop="onRoleDrop"
@on-user-drag-start="onUserDragStart"
/> />
<AdminOidcDetail <AdminOidcDetail
+1 -1
View File
@@ -3,7 +3,7 @@
<head> <head>
<meta charset="UTF-8"> <meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0"> <meta name="viewport" content="width=device-width, initial-scale=1.0">
<script>{let t=localStorage.getItem('paskia-theme');if(t!=='light'&&t!=='dark')t=new URLSearchParams(location.hash.slice(1)).get('theme');(t==='dark'||t!=='light'&&matchMedia('(prefers-color-scheme:dark)').matches)&&document.documentElement.classList.add('dark');if(window.location.pathname==='/auth/restricted/iframe')document.documentElement.style.background='transparent'}</script> <script>{let t=localStorage.getItem('paskia-theme');if(!t){let p=new URLSearchParams(location.hash.slice(1)).get('theme');if(p==='light'||p==='dark')t=p}(t==='dark'||t!=='light'&&matchMedia('(prefers-color-scheme:dark)').matches)&&document.documentElement.classList.add('dark');if(window.location.pathname==='/auth/restricted/iframe')document.documentElement.style.background='transparent'}</script>
<link rel="stylesheet" href="/src/assets/style.css"> <link rel="stylesheet" href="/src/assets/style.css">
</head> </head>
<body> <body>
+2 -2
View File
@@ -1,9 +1,9 @@
// Early theme for restricted app - first URL param wins, then localStorage // Early theme for restricted app - user preference (localStorage) wins, then URL param
import { applyTheme, getCachedTheme } from '@/utils/theme.js' import { applyTheme, getCachedTheme } from '@/utils/theme.js'
function getTheme() { function getTheme() {
const params = new URLSearchParams(location.hash.slice(1)) const params = new URLSearchParams(location.hash.slice(1))
return params.get('theme') || getCachedTheme() || '' return getCachedTheme() || params.get('theme') || ''
} }
// Apply theme class to document root // Apply theme class to document root
+9 -7
View File
@@ -8,7 +8,7 @@
<main class="view-root"> <main class="view-root">
<div class="surface surface--tight reset-container"> <div class="surface surface--tight reset-container">
<header class="view-header reset-header"> <header class="view-header center">
<h1>🔑 Registration</h1> <h1>🔑 Registration</h1>
<p class="view-lede"> <p class="view-lede">
{{ subtitleMessage }} {{ subtitleMessage }}
@@ -60,6 +60,7 @@ import { computed, onMounted, reactive, ref } from 'vue'
import passkey from '@/utils/passkey' import passkey from '@/utils/passkey'
import { getSettings, uiBasePath } from '@/utils/settings' import { getSettings, uiBasePath } from '@/utils/settings'
import { apiJson, ApiError, getUserFriendlyErrorMessage } from 'paskia' import { apiJson, ApiError, getUserFriendlyErrorMessage } from 'paskia'
import { updateThemeFromSession } from '@/utils/theme'
const status = reactive({ const status = reactive({
show: false, show: false,
@@ -80,7 +81,7 @@ const sessionDescriptor = computed(() => tokenInfo.value?.token_type || 'your en
const subtitleMessage = computed(() => { const subtitleMessage = computed(() => {
if (initializing.value) return 'Preparing your secure enrollment…' if (initializing.value) return 'Preparing your secure enrollment…'
if (!canRegister.value) return 'This authentication link is no longer valid.' if (!canRegister.value) return 'This authentication link is no longer valid.'
return `Finish up ${sessionDescriptor.value}. You may edit the name below if needed, and it will be saved to your passkey.` return `Finish up ${sessionDescriptor.value}. The name entered will be stored on your passkey and on our system.`
}) })
const basePath = computed(() => uiBasePath()) const basePath = computed(() => uiBasePath())
@@ -117,6 +118,7 @@ async function fetchTokenInfo() {
headers: { 'Authorization': `Bearer ${token.value}` }, headers: { 'Authorization': `Bearer ${token.value}` },
}) })
displayName.value = tokenInfo.value.display_name displayName.value = tokenInfo.value.display_name
if (tokenInfo.value.theme) updateThemeFromSession({ user: { theme: tokenInfo.value.theme } })
} catch (error) { } catch (error) {
console.error('Failed to load token info', error) console.error('Failed to load token info', error)
const message = error instanceof ApiError const message = error instanceof ApiError
@@ -201,14 +203,14 @@ onMounted(async () => {
</script> </script>
<style scoped> <style scoped>
main.view-root { min-height: 100vh; align-items: center; justify-content: center; padding: 2rem 1rem; }
.reset-container { .reset-container {
max-width: 560px; max-width: 520px;
margin: 0 auto; margin: 0 auto;
width: 100%; width: 100%;
} display: flex;
flex-direction: column;
.reset-header { gap: 1.75rem;
text-align: center;
} }
.section-body { .section-body {
+2 -1
View File
@@ -15,7 +15,8 @@
"qrcode": "^1.5.4", "qrcode": "^1.5.4",
"sirv": "^3.0.2", "sirv": "^3.0.2",
"uuidv7": "^1.1.0", "uuidv7": "^1.1.0",
"vue": "^3.5.17" "vue": "^3.5.17",
"vuedraggable": "^4.1.0"
}, },
"devDependencies": { "devDependencies": {
"@vitejs/plugin-vue": "^6.0.0", "@vitejs/plugin-vue": "^6.0.0",
+48 -31
View File
@@ -1,5 +1,6 @@
<script setup> <script setup>
import { computed, ref } from 'vue' import { computed, ref } from 'vue'
import draggable from 'vuedraggable'
import { getDirection, navigateButtonRow, focusPreferred } from '@/utils/keynav' import { getDirection, navigateButtonRow, focusPreferred } from '@/utils/keynav'
const props = defineProps({ const props = defineProps({
@@ -8,7 +9,7 @@ const props = defineProps({
navigationDisabled: { type: Boolean, default: false } navigationDisabled: { type: Boolean, default: false }
}) })
const emit = defineEmits(['updateOrg', 'createRole', 'updateRole', 'deleteRole', 'createUserInRole', 'openUser', 'toggleRolePermission', 'onRoleDragOver', 'onRoleDrop', 'onUserDragStart', 'navigateOut']) const emit = defineEmits(['updateOrg', 'createRole', 'updateRole', 'deleteRole', 'createUserInRole', 'openUser', 'toggleRolePermission', 'moveUserToRole', 'navigateOut'])
// Template refs for navigation // Template refs for navigation
const orgTitleRef = ref(null) const orgTitleRef = ref(null)
@@ -50,6 +51,14 @@ function roleUserCount(roleUuid) {
return Object.values(props.selectedOrg.users).filter(u => u.role === roleUuid).length return Object.values(props.selectedOrg.users).filter(u => u.role === roleUuid).length
} }
function onUserChange(evt, targetRoleUuid) {
// Only handle 'added' events (when a user is dropped into this role)
if (evt.added) {
const userUuid = evt.added.element.uuid
emit('moveUserToRole', userUuid, targetRoleUuid)
}
}
function permissionDisplayName(scope) { function permissionDisplayName(scope) {
return props.permissions.find(p => p.scope === scope)?.display_name || scope return props.permissions.find(p => p.scope === scope)?.display_name || scope
} }
@@ -350,8 +359,6 @@ defineExpose({ focusFirstElement })
v-for="(r, roleIndex) in sortedRoles" v-for="(r, roleIndex) in sortedRoles"
:key="r.uuid" :key="r.uuid"
class="role-column" class="role-column"
@dragover="$emit('onRoleDragOver', $event)"
@drop="e => $emit('onRoleDrop', e, selectedOrg, r)"
> >
<div class="role-header" @keydown="e => handleRoleHeaderKeydown(e, roleIndex)"> <div class="role-header" @keydown="e => handleRoleHeaderKeydown(e, roleIndex)">
<strong class="role-name" :title="r.uuid"> <strong class="role-name" :title="r.uuid">
@@ -363,26 +370,32 @@ defineExpose({ focusFirstElement })
<button @click="$emit('createUserInRole', selectedOrg, r)" class="plus-btn" aria-label="Add user" title="Add user"></button> <button @click="$emit('createUserInRole', selectedOrg, r)" class="plus-btn" aria-label="Add user" title="Add user"></button>
</div> </div>
</div> </div>
<template v-if="roleUserCount(r.uuid) > 0"> <div class="user-list-wrapper">
<ul class="user-list" @keydown="handleUserListKeydown"> <draggable
<li :list="roleUsers(r.uuid)"
v-for="u in roleUsers(r.uuid)" group="users"
:key="u.uuid" item-key="uuid"
class="user-chip" tag="ul"
tabindex="0" class="user-list"
draggable="true" @change="evt => onUserChange(evt, r.uuid)"
@dragstart="e => $emit('onUserDragStart', e, u.uuid, selectedOrg.uuid)" @keydown="handleUserListKeydown"
@click="$emit('openUser', u)" >
@keydown.enter="$emit('openUser', u)" <template #item="{ element: u }">
:title="u.uuid" <li
> class="user-chip"
<span class="name">{{ u.display_name }}</span> tabindex="0"
<span class="meta">{{ u.last_seen ? new Date(u.last_seen).toLocaleDateString() : '—' }}</span> @click="$emit('openUser', u)"
</li> @keydown.enter="$emit('openUser', u)"
</ul> :title="u.uuid"
</template> >
<div v-else class="empty-role"> <span class="name">{{ u.display_name }}</span>
<p class="empty-text muted">No members</p> <span class="meta">{{ u.last_seen ? new Date(u.last_seen).toLocaleDateString() : '—' }}</span>
</li>
</template>
</draggable>
<div v-if="roleUserCount(r.uuid) === 0" class="empty-role">
<p class="empty-text muted">No members</p>
</div>
</div> </div>
</div> </div>
</div> </div>
@@ -391,23 +404,27 @@ defineExpose({ focusFirstElement })
<style scoped> <style scoped>
.card.surface { padding: var(--space-lg); } .card.surface { padding: var(--space-lg); }
.org-title { display: flex; align-items: center; gap: var(--space-sm); margin-bottom: var(--space-lg); } .org-title { display: flex; align-items: center; gap: var(--space-sm); margin-bottom: var(--space-lg); font-size: 1.65rem; }
.org-name { font-size: 1.5rem; font-weight: 600; color: var(--color-heading); } .org-name { font-weight: 600; color: var(--color-heading); }
.perm-matrix-grid .role-head { display: flex; align-items: flex-end; justify-content: center; } .perm-matrix-grid .role-head { display: flex; align-items: flex-end; justify-content: center; }
.perm-matrix-grid .role-head span { writing-mode: vertical-rl; transform: rotate(180deg); font-size: 0.65rem; } .perm-matrix-grid .role-head span { writing-mode: vertical-rl; transform: rotate(180deg); font-size: 0.65rem; }
.perm-matrix-grid .add-role-head { cursor: pointer; } .perm-matrix-grid .add-role-head { cursor: pointer; }
.roles-grid { display: flex; gap: var(--space-lg); margin-top: var(--space-lg); } .roles-grid { display: flex; flex-wrap: wrap; gap: var(--space-lg); margin-top: var(--space-lg); justify-content: flex-start; align-items: stretch; }
.role-column { flex: 1; min-width: 200px; border-radius: var(--radius-md); padding: var(--space-md); } .role-column { flex: 0 0 240px; border-radius: var(--radius-md); padding: var(--space-md); display: flex; flex-direction: column; }
.role-header { display: flex; justify-content: space-between; align-items: center; margin-bottom: var(--space-md); } .role-header { display: flex; justify-content: space-between; align-items: center; margin-bottom: var(--space-md); }
.role-name { display: flex; align-items: center; gap: var(--space-xs); font-size: 1.1rem; color: var(--color-heading); } .role-name { display: flex; align-items: center; gap: var(--space-xs); font-size: 1.1rem; color: var(--color-heading); }
.role-actions { display: flex; gap: var(--space-xs); } .role-actions { display: flex; gap: var(--space-xs); }
.plus-btn { background: none; color: var(--color-accent); border: none; border-radius: var(--radius-sm); padding: 0.25rem 0.45rem; font-size: 1.1rem; cursor: pointer; } .plus-btn { background: none; color: var(--color-accent); border: none; border-radius: var(--radius-sm); padding: 0.25rem 0.45rem; font-size: 1.1rem; cursor: pointer; }
.plus-btn:hover { background: rgba(37, 99, 235, 0.18); } .plus-btn:hover { background: rgba(37, 99, 235, 0.18); }
.user-list { list-style: none; padding: 0; margin: 0; display: flex; flex-direction: column; gap: var(--space-xs); } .user-list-wrapper { position: relative; flex: 1; display: flex; flex-direction: column; min-height: 5.5rem; }
.user-chip { background: var(--color-surface); border: 1px solid var(--color-border); border-radius: var(--radius-md); padding: 0.45rem 0.6rem; display: flex; justify-content: space-between; gap: var(--space-sm); cursor: grab; } .user-list { list-style: none; padding: 0; margin: 0; display: flex; flex-direction: column; gap: var(--space-xs); flex: 1; }
.user-chip { background: var(--color-accent-strong); color: white; border: none; border-radius: var(--radius-md); padding: 0.45rem 0.6rem; display: flex; justify-content: space-between; gap: var(--space-sm); cursor: grab; }
.user-chip:focus { outline: 2px solid var(--color-accent); outline-offset: 1px; } .user-chip:focus { outline: 2px solid var(--color-accent); outline-offset: 1px; }
.user-chip .meta { font-size: 0.7rem; color: var(--color-text-muted); } .user-chip .meta { font-size: 0.7rem; color: rgba(255, 255, 255, 0.8); }
.empty-role { border: 1px dashed var(--color-border-strong); border-radius: var(--radius-md); padding: var(--space-sm); display: flex; flex-direction: column; gap: var(--space-xs); align-items: flex-start; } .user-chip.sortable-ghost { opacity: 0.5; }
.user-chip.sortable-chosen { box-shadow: 0 4px 12px rgba(0, 0, 0, 0.2); }
.empty-role { position: absolute; inset: 0; border: 1px dashed var(--color-border-strong); border-radius: var(--radius-md); display: flex; align-items: center; justify-content: center; pointer-events: none; }
.user-list:has(.sortable-ghost) + .empty-role { display: none; }
.empty-text { margin: 0; } .empty-text { margin: 0; }
@media (max-width: 720px) { @media (max-width: 720px) {
-3
View File
@@ -823,9 +823,6 @@ th {
.user-info { .user-info {
display: grid; display: grid;
border-radius: var(--radius-md);
background: var(--color-surface);
padding: 1rem;
} }
.user-details { .user-details {
+3 -3
View File
@@ -10,9 +10,9 @@
<UserBasicInfo <UserBasicInfo
v-if="ctx" v-if="ctx"
:name="ctx.user.display_name" :name="ctx.user.display_name"
:visits="authStore.userInfo?.visits || 0" :visits="authStore.userInfo.user.visits"
:created-at="authStore.userInfo?.created_at" :created-at="authStore.userInfo.user.created_at"
:last-seen="authStore.userInfo?.last_seen" :last-seen="authStore.userInfo.user.last_seen"
:email="ctx.user.email" :email="ctx.user.email"
:telephone="ctx.user.telephone" :telephone="ctx.user.telephone"
:org-display-name="orgDisplayName" :org-display-name="orgDisplayName"
+12 -12
View File
@@ -22,8 +22,8 @@
:created-at="authStore.userInfo.user.created_at" :created-at="authStore.userInfo.user.created_at"
:last-seen="authStore.userInfo.user.last_seen" :last-seen="authStore.userInfo.user.last_seen"
:loading="authStore.isLoading" :loading="authStore.isLoading"
:org-display-name="authStore.ctx?.org.display_name" :org-display-name="authStore.userInfo.org.display_name"
:role-name="authStore.ctx?.role.display_name" :role-name="authStore.userInfo.role.display_name"
update-endpoint="/auth/api/user/info" update-endpoint="/auth/api/user/info"
@saved="authStore.loadUserInfo()" @saved="authStore.loadUserInfo()"
@edit="openEditDialog" @edit="openEditDialog"
@@ -53,7 +53,7 @@
<CredentialList <CredentialList
ref="credentialList" ref="credentialList"
:credentials="credentials" :credentials="credentials"
:aaguid-info="authStore.userInfo?.aaguid_info || {}" :aaguid-info="authStore.userInfo.aaguid_info"
:loading="authStore.isLoading" :loading="authStore.isLoading"
:hovered-credential-uuid="hoveredCredentialUuid" :hovered-credential-uuid="hoveredCredentialUuid"
:hovered-session-credential-uuid="hoveredSession?.credential" :hovered-session-credential-uuid="hoveredSession?.credential"
@@ -184,11 +184,11 @@ const hasActiveModal = computed(() => showEditDialog.value || showRegLink.value)
watch(showEditDialog, (open) => { watch(showEditDialog, (open) => {
if (!open) return if (!open) return
const user = authStore.userInfo?.user const user = authStore.userInfo.user
editName.value = user?.display_name ?? '' editName.value = user.display_name ?? ''
editEmail.value = user?.email ?? '' editEmail.value = user.email ?? ''
editUsername.value = user?.preferred_username ?? '' editUsername.value = user.preferred_username ?? ''
editTelephone.value = user?.telephone ?? '' editTelephone.value = user.telephone ?? ''
editError.value = '' editError.value = ''
}) })
@@ -341,7 +341,7 @@ const handleDelete = async (credential) => {
const rpName = computed(() => authStore.settings?.rp_name || 'this service') const rpName = computed(() => authStore.settings?.rp_name || 'this service')
const paskiaVersion = computed(() => authStore.settings?.version || '') const paskiaVersion = computed(() => authStore.settings?.version || '')
const sessions = computed(() => authStore.userInfo?.sessions || {}) const sessions = computed(() => authStore.userInfo.sessions)
const currentSessionHost = computed(() => { const currentSessionHost = computed(() => {
const currentSession = Object.values(sessions.value).find(session => session.is_current) const currentSession = Object.values(sessions.value).find(session => session.is_current)
return currentSession?.host || 'this host' return currentSession?.host || 'this host'
@@ -365,12 +365,12 @@ const logoutEverywhere = async () => { await authStore.logoutEverywhere() }
const logout = async () => { await authStore.logout() } const logout = async () => { await authStore.logout() }
const openEditDialog = () => { showEditDialog.value = true } const openEditDialog = () => { showEditDialog.value = true }
const isAdmin = computed(() => { const isAdmin = computed(() => {
const perms = authStore.ctx?.permissions const perms = Object.values(authStore.userInfo.permissions).map(p => p.scope)
return perms?.includes('auth:admin') || perms?.includes('auth:org:admin') return perms.includes('auth:admin') || perms.includes('auth:org:admin')
}) })
const hasMultipleSessions = computed(() => Object.keys(sessions.value).length > 1) const hasMultipleSessions = computed(() => Object.keys(sessions.value).length > 1)
const credentials = computed(() => const credentials = computed(() =>
Object.entries(authStore.userInfo?.credentials || {}).map(([uuid, c]) => ({ ...c, credential: uuid })) Object.entries(authStore.userInfo.credentials).map(([uuid, c]) => ({ ...c, credential: uuid }))
) )
const useWideLayout = computed(() => { const useWideLayout = computed(() => {
// Check if any single site has more than 8 sessions // Check if any single site has more than 8 sessions
@@ -37,7 +37,7 @@
<script setup> <script setup>
import { ref, computed, onMounted, onUnmounted, nextTick } from 'vue' import { ref, computed, onMounted, onUnmounted, nextTick } from 'vue'
import QRCodeDisplay from '@/components/QRCodeDisplay.vue' import QRCodeDisplay from '@/components/QRCodeDisplay.vue'
import { apiJson, holdGlobalBackdrop, releaseGlobalBackdrop } from 'paskia' import { apiJson, AuthCancelledError, holdGlobalBackdrop, releaseGlobalBackdrop } from 'paskia'
import { formatDate } from '@/utils/helpers' import { formatDate } from '@/utils/helpers'
import { getDirection } from '@/utils/keynav' import { getDirection } from '@/utils/keynav'
import { useAuthStore } from '@/stores/auth' import { useAuthStore } from '@/stores/auth'
@@ -90,7 +90,9 @@ async function generateLink() {
emit('close') emit('close')
} }
} catch (e) { } catch (e) {
authStore.showMessage(e.message || 'Failed to generate link', 'error') if (!(e instanceof AuthCancelledError)) {
authStore.showMessage(e.message || 'Failed to generate link', 'error')
}
emit('close') emit('close')
} }
} }
@@ -163,6 +163,9 @@ async function startRemoteAuth() {
// PoW challenge // PoW challenge
const powChallenge = await ws.receive_json() const powChallenge = await ws.receive_json()
if (powChallenge.status) {
throw new Error(powChallenge.detail || `Failed to connect: ${powChallenge.status}`)
}
if (powChallenge.pow) { if (powChallenge.pow) {
const challenge = b64dec(powChallenge.pow.challenge) const challenge = b64dec(powChallenge.pow.challenge)
const nonces = await solvePoW(challenge, powChallenge.pow.work) const nonces = await solvePoW(challenge, powChallenge.pow.work)
@@ -61,6 +61,7 @@ import { getSettings, uiBasePath } from '@/utils/settings'
import { fetchJson, getUserFriendlyErrorMessage } from 'paskia' import { fetchJson, getUserFriendlyErrorMessage } from 'paskia'
import RemoteAuthRequest from '@/components/RemoteAuthRequest.vue' import RemoteAuthRequest from '@/components/RemoteAuthRequest.vue'
import { focusDialogButton } from '@/utils/keynav' import { focusDialogButton } from '@/utils/keynav'
import { updateThemeFromSession } from '@/utils/theme'
const props = defineProps({ const props = defineProps({
mode: { mode: {
@@ -147,6 +148,7 @@ async function fetchSettings() {
async function validateSession() { async function validateSession() {
try { try {
session.value = await fetchJson('/auth/api/validate', { method: 'POST' }) session.value = await fetchJson('/auth/api/validate', { method: 'POST' })
updateThemeFromSession(session.value?.ctx)
if (isAuthenticated.value && props.mode !== 'reauth') { if (isAuthenticated.value && props.mode !== 'reauth') {
currentView.value = 'forbidden' currentView.value = 'forbidden'
emit('forbidden', session.value) emit('forbidden', session.value)
+1 -1
View File
@@ -118,7 +118,7 @@ const userLoaded = computed(() => !!props.name)
.user-info-extra { grid-area: extra; padding-left: 1rem; border-left: 1px solid var(--color-border); flex-shrink: 0; } .user-info-extra { grid-area: extra; padding-left: 1rem; border-left: 1px solid var(--color-border); flex-shrink: 0; }
.user-name-row { display: inline-flex; align-items: center; gap: 0.35rem; max-width: 100%; min-width: 0; } .user-name-row { display: inline-flex; align-items: center; gap: 0.35rem; max-width: 100%; min-width: 0; }
.display-name { font-weight: 600; font-size: 1.05em; line-height: 1.2; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; min-width: 0; } .display-name { font-weight: 600; font-size: 1.05em; line-height: 1.2; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; min-width: 0; }
.mini-btn { width: auto; padding: 4px 6px; margin: 0; font-size: 0.75em; line-height: 1; cursor: pointer; } .mini-btn { width: auto; padding: 4px 6px; margin: 0; font-size: 0.75em; line-height: 1; cursor: pointer; background: transparent; }
.mini-btn:hover:not(:disabled) { background: var(--color-accent-soft); color: var(--color-accent); } .mini-btn:hover:not(:disabled) { background: var(--color-accent-soft); color: var(--color-accent); }
.mini-btn:active:not(:disabled) { transform: translateY(1px); } .mini-btn:active:not(:disabled) { transform: translateY(1px); }
.mini-btn:disabled { opacity: 0.5; cursor: not-allowed; } .mini-btn:disabled { opacity: 0.5; cursor: not-allowed; }
+1 -1
View File
@@ -88,7 +88,7 @@ export const useAuthStore = defineStore('auth', {
async loadUserInfo() { async loadUserInfo() {
try { try {
this.userInfo = await apiJson('/auth/api/user-info', { method: 'GET' }) this.userInfo = await apiJson('/auth/api/user-info', { method: 'GET' })
updateThemeFromSession(this.ctx) updateThemeFromSession(this.userInfo)
console.log('User info loaded:', this.userInfo) console.log('User info loaded:', this.userInfo)
} catch (error) { } catch (error) {
// Suppress toast for 401/403 errors - the auth iframe will handle these // Suppress toast for 401/403 errors - the auth iframe will handle these
-32
View File
@@ -1,32 +0,0 @@
// Cache for auth iframe URL by mode
const authIframeUrlCache = {}
/**
* Get the auth iframe URL for a given mode.
* Fetches from /auth/api/forward which returns URL in the auth.iframe field.
* Results are cached per mode.
* @param {string} mode - The auth mode ('login', 'reauth', 'forbidden')
* @returns {Promise<string>} - The URL for the iframe
*/
export async function getAuthIframeUrl(mode = 'login') {
if (authIframeUrlCache[mode]) {
return authIframeUrlCache[mode]
}
// Fetch from forward endpoint - it returns URL in auth.iframe on 401/403
const response = await fetch('/auth/api/forward')
if (response.status === 401 || response.status === 403) {
const data = await response.json()
if (data.auth?.iframe) {
// The iframe field now contains a URL with hash fragment
// If mode differs, update the hash param
let url = data.auth.iframe
if (mode !== data.auth.mode) {
url = url.replace(/mode=[^&]*/, `mode=${mode}`)
}
authIframeUrlCache[mode] = url
return url
}
}
throw new Error('Unable to fetch auth iframe URL')
}
-2
View File
@@ -18,8 +18,6 @@ export {
isAuthIframeOpen, isAuthIframeOpen,
hideAuthIframe, hideAuthIframe,
showAuthIframe, showAuthIframe,
createAuthIframe,
removeAuthIframe,
} from './overlay' } from './overlay'
export { SessionValidator } from './validate' export { SessionValidator } from './validate'
+34 -70
View File
@@ -1,19 +1,14 @@
import argparse import argparse
import asyncio
import json import json
import logging import logging
import os import os
from urllib.parse import urlparse from urllib.parse import urlparse
from fastapi_vue import server from fastapi_vue import server
from fastapi_vue.hostutil import parse_endpoint from fastapi_vue.hostutil import parse_endpoints
from paskia import db
from paskia import globals as _globals
from paskia.bootstrap import bootstrap_if_needed
from paskia.config import PaskiaConfig from paskia.config import PaskiaConfig
from paskia.db.background import flush from paskia.db.jsonl import load_readonly
from paskia.db.structs import Config
from paskia.util import startupbox from paskia.util import startupbox
from paskia.util.hostutil import normalize_origin from paskia.util.hostutil import normalize_origin
@@ -56,6 +51,7 @@ def add_common_options(p: argparse.ArgumentParser) -> None:
"--origin", "--origin",
action="append", action="append",
dest="origins", dest="origins",
default=[],
metavar="URL", metavar="URL",
help="Allowed origin URL(s). May be specified multiple times. If any are specified, only those origins are permitted for WebSocket authentication.", help="Allowed origin URL(s). May be specified multiple times. If any are specified, only those origins are permitted for WebSocket authentication.",
) )
@@ -103,9 +99,10 @@ def main():
if getattr(args, "listen", None) == "": if getattr(args, "listen", None) == "":
args.listen = None args.listen = None
# Init db and load stored config # Read-only load to get stored config (no writes, no global state)
asyncio.run(db.init(rp_id=args.rp_id)) db_path = os.environ.get("PASKIA_DB", f"{args.rp_id}.paskiadb")
stored_config = db.data().config stored_db = load_readonly(db_path, rp_id=args.rp_id)
stored_config = stored_db.config
# Apply defaults from stored config # Apply defaults from stored config
if args.rp_name is None and stored_config.rp_name is not None: if args.rp_name is None and stored_config.rp_name is not None:
@@ -118,42 +115,26 @@ def main():
args.listen = stored_config.listen args.listen = stored_config.listen
# Parse first endpoint for config display and site_url # Parse first endpoint for config display and site_url
first_listen = args.listen[0] if isinstance(args.listen, list) else args.listen ep = next(iter(parse_endpoints(args.listen, DEFAULT_PORT)), {})
endpoints = parse_endpoint(first_listen, DEFAULT_PORT) host, port, uds = ep.get("host"), ep.get("port"), ep.get("uds")
# Extract host/port/uds from first endpoint for config display and site_url # Process and normalize auth_host
ep = endpoints[0] if endpoints else {}
host = ep.get("host")
port = ep.get("port")
uds = ep.get("uds")
# Collect and normalize origins, handle auth_host
origins = [normalize_origin(o) for o in (getattr(args, "origins", None) or [])]
if args.auth_host: if args.auth_host:
# Normalize auth_host with scheme
if "://" not in args.auth_host: if "://" not in args.auth_host:
args.auth_host = f"https://{args.auth_host}" args.auth_host = f"https://{args.auth_host}"
args.auth_host = args.auth_host.rstrip("/")
validate_auth_host(args.auth_host, args.rp_id) validate_auth_host(args.auth_host, args.rp_id)
if args.origins:
args.origins.insert(0, args.auth_host) # Ensure first in origins
# If origins are configured, ensure auth_host is included at top # Normalize, strip trailing slashes, and deduplicate while preserving order
if origins: origins = list({normalize_origin(o).rstrip("/"): ... for o in (args.origins)})
# Insert auth_host at the beginning
origins.insert(0, args.auth_host)
# Remove duplicates while preserving order
seen = set()
origins = [x for x in origins if not (x in seen or seen.add(x))]
# Compute site_url and site_path for reset links # Compute site_url and site_path for reset links
# Priority: PASKIA_SITE_URL (explicit) > auth_host > first origin with localhost > http://localhost:port # Priority: auth_host > first configured origin > PASKIA_VITE_URL (devserver) > http://localhost:port > https://rp_id
explicit_site_url = os.environ.get("PASKIA_SITE_URL") site_path = "/auth/"
if explicit_site_url: if args.auth_host:
# Explicit site URL from devserver or deployment config site_url = args.auth_host
site_url = explicit_site_url.rstrip("/")
site_path = "/" if args.auth_host else "/auth/"
elif args.auth_host:
site_url = args.auth_host.rstrip("/")
site_path = "/" site_path = "/"
elif origins: elif origins:
# Find localhost origin if rp_id is localhost, else use first origin # Find localhost origin if rp_id is localhost, else use first origin
@@ -162,15 +143,13 @@ def main():
if args.rp_id == "localhost" if args.rp_id == "localhost"
else None else None
) )
site_url = (localhost_origin or origins[0]).rstrip("/") site_url = localhost_origin or origins[0]
site_path = "/auth/" elif vite_url := os.environ.get("PASKIA_VITE_URL"):
site_url = vite_url.rstrip("/") # Devserver
elif args.rp_id == "localhost" and port: elif args.rp_id == "localhost" and port:
# Dev mode: use http with port site_url = f"http://localhost:{port}" # Backend directly if we can
site_url = f"http://localhost:{port}"
site_path = "/auth/"
else: else:
site_url = f"https://{args.rp_id}" site_url = f"https://{args.rp_id}" # Assume external reverse proxy
site_path = "/auth/"
# Build runtime configuration # Build runtime configuration
config = PaskiaConfig( config = PaskiaConfig(
@@ -186,6 +165,14 @@ def main():
) )
# Export configuration via single JSON env variable for worker processes # Export configuration via single JSON env variable for worker processes
# Include cli_config and save flag so lifespan can handle bootstrap/persistence
cli_config = {
"rp_id": args.rp_id,
"rp_name": args.rp_name,
"origins": args.origins,
"auth_host": args.auth_host,
"listen": args.listen,
}
config_json = { config_json = {
"rp_id": config.rp_id, "rp_id": config.rp_id,
"rp_name": config.rp_name, "rp_name": config.rp_name,
@@ -193,36 +180,13 @@ def main():
"auth_host": config.auth_host, "auth_host": config.auth_host,
"site_url": config.site_url, "site_url": config.site_url,
"site_path": config.site_path, "site_path": config.site_path,
"save": args.save,
"cli_config": cli_config,
} }
os.environ["PASKIA_CONFIG"] = json.dumps(config_json) os.environ["PASKIA_CONFIG"] = json.dumps(config_json)
startupbox.print_startup_config(config) startupbox.print_startup_config(config)
# Build config to save (for bootstrap or explicit --save)
cli_config = Config(
rp_id=args.rp_id,
rp_name=args.rp_name,
origins=args.origins,
auth_host=args.auth_host,
listen=args.listen,
)
async def startup():
await _globals.init(
rp_id=config.rp_id,
rp_name=config.rp_name,
origins=config.origins,
bootstrap=False,
)
# Pass config to bootstrap - it will be saved within the bootstrap transaction
await bootstrap_if_needed(config=cli_config)
# Also save config if --save was explicitly used (even without bootstrap)
if args.save:
await db.update_config(cli_config)
await flush()
asyncio.run(startup())
dev = {"reload": True, "reload_dirs": ["paskia"]} if DEVMODE else {} dev = {"reload": True, "reload_dirs": ["paskia"]} if DEVMODE else {}
server.run( server.run(
"paskia.fastapi.mainapp:app", "paskia.fastapi.mainapp:app",
+3 -9
View File
@@ -1,7 +1,7 @@
""" """
Database module for WebAuthn passkey authentication. Database module for WebAuthn passkey authentication.
Read: Access data() directly, use build_* to convert to public structs. Read: Access data() directly for structs.
CTX: data().session_ctx(key) returns SessionContext with effective permissions. CTX: data().session_ctx(key) returns SessionContext with effective permissions.
Write: Functions validate and commit, or raise ValueError. Write: Functions validate and commit, or raise ValueError.
@@ -10,7 +10,6 @@ Usage:
# Read (after init) # Read (after init)
user_data = db.data().users[user_uuid] user_data = db.data().users[user_uuid]
user = db.build_user(user_uuid)
# Context # Context
ctx = db.data().session_ctx(session_key) ctx = db.data().session_ctx(session_key)
@@ -27,6 +26,7 @@ from paskia.db.background import (
stop_cleanup, stop_cleanup,
) )
from paskia.db.bootstrap import bootstrap from paskia.db.bootstrap import bootstrap
from paskia.db.jsonl import load_readonly
from paskia.db.lifecycle import cleanup_expired, init from paskia.db.lifecycle import cleanup_expired, init
from paskia.db.operations import ( from paskia.db.operations import (
add_permission_to_org, add_permission_to_org,
@@ -102,18 +102,12 @@ __all__ = [
# Instance # Instance
"data", "data",
"init", "init",
"load_readonly",
# Background # Background
"start_background", "start_background",
"stop_background", "stop_background",
"start_cleanup", "start_cleanup",
"stop_cleanup", "stop_cleanup",
# Builders
"build_credential",
"build_permission",
"build_reset_token",
"build_role",
"build_session",
"build_user",
# Read ops # Read ops
# Write ops # Write ops
"add_permission_to_org", "add_permission_to_org",
+2 -1
View File
@@ -90,7 +90,7 @@ async def start_background():
async def stop_background(): async def stop_background():
"""Stop the background task and flush any pending changes.""" """Stop the background task, flush pending changes, and release the file lock."""
global _background_task global _background_task
if _background_task: if _background_task:
_background_task.cancel() _background_task.cancel()
@@ -99,6 +99,7 @@ async def stop_background():
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
_background_task = None _background_task = None
_ops._store.close()
# Aliases for backwards compatibility # Aliases for backwards compatibility
+1 -2
View File
@@ -7,6 +7,7 @@ from datetime import UTC, datetime
import uuid7 import uuid7
import paskia.db.operations as _ops import paskia.db.operations as _ops
from paskia.authsession import reset_expires
from paskia.db.structs import Config, Org, Permission, ResetToken, Role, User from paskia.db.structs import Config, Org, Permission, ResetToken, Role, User
from paskia.util.crypto import secret_key from paskia.util.crypto import secret_key
@@ -59,8 +60,6 @@ def bootstrap(
# Set reset token expiry (passphrase generated by ResetToken.create) # Set reset token expiry (passphrase generated by ResetToken.create)
if reset_expiry is None: if reset_expiry is None:
from paskia.authsession import reset_expires # noqa: PLC0415
reset_expiry = reset_expires() reset_expiry = reset_expires()
with _ops._db.transaction("bootstrap"): with _ops._db.transaction("bootstrap"):
+247
View File
@@ -0,0 +1,247 @@
"""Cross-platform locked file for the database (no separate .lock files).
Unix: open() + fcntl.flock (advisory, cooperative among processes that flock).
Windows: CreateFileW with FILE_SHARE_READ (OS-enforced, allows readers, blocks writers).
A single file descriptor is opened once for both reading and writing.
The lock is acquired atomically (on Windows) or immediately after open (on Unix),
and the same descriptor is used for the lifetime of the process: first to read
the existing content, then to append new writes.
"""
import logging
import os
import sys
from pathlib import Path
_logger = logging.getLogger(__name__)
def _fatal(msg: str) -> None:
"""Log a fatal error and exit immediately, bypassing exception handlers."""
_logger.critical(msg)
os._exit(1)
if sys.platform == "win32":
import ctypes
from ctypes import wintypes
_kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
_GENERIC_READ = 0x80000000
_GENERIC_WRITE = 0x40000000
_FILE_SHARE_READ = 0x00000001
_OPEN_EXISTING = 3
_OPEN_ALWAYS = 4
_FILE_ATTRIBUTE_NORMAL = 0x80
_FILE_BEGIN = 0
_FILE_END = 2
_ERROR_SHARING_VIOLATION = 32
_INVALID_FILE_SIZE = 0xFFFFFFFF
_kernel32.CreateFileW.restype = wintypes.HANDLE
_kernel32.CreateFileW.argtypes = [
wintypes.LPCWSTR,
wintypes.DWORD,
wintypes.DWORD,
ctypes.c_void_p,
wintypes.DWORD,
wintypes.DWORD,
wintypes.HANDLE,
]
_kernel32.ReadFile.restype = wintypes.BOOL
_kernel32.ReadFile.argtypes = [
wintypes.HANDLE,
ctypes.c_void_p,
wintypes.DWORD,
ctypes.POINTER(wintypes.DWORD),
ctypes.c_void_p,
]
_kernel32.WriteFile.restype = wintypes.BOOL
_kernel32.WriteFile.argtypes = [
wintypes.HANDLE,
ctypes.c_void_p,
wintypes.DWORD,
ctypes.POINTER(wintypes.DWORD),
ctypes.c_void_p,
]
_kernel32.GetFileSize.restype = wintypes.DWORD
_kernel32.GetFileSize.argtypes = [
wintypes.HANDLE,
ctypes.POINTER(wintypes.DWORD),
]
_kernel32.SetFilePointer.restype = wintypes.DWORD
_kernel32.SetFilePointer.argtypes = [
wintypes.HANDLE,
wintypes.LONG,
ctypes.POINTER(wintypes.LONG),
wintypes.DWORD,
]
_kernel32.CloseHandle.restype = wintypes.BOOL
_kernel32.CloseHandle.argtypes = [wintypes.HANDLE]
def _is_invalid_handle(handle) -> bool:
return ctypes.c_void_p(handle).value == ctypes.c_void_p(-1).value
else:
import fcntl
class LockedFile:
"""A file opened with an exclusive write lock.
Usage::
f = LockedFile()
f.open(path) # open + lock (read+write)
content = f.read() # read entire content
f.write(data) # append data (seeks to end first)
f.close() # release lock + close fd
Unix: fcntl.flock (advisory) — read-only callers that don't flock are unaffected.
Windows: CreateFileW with FILE_SHARE_READ — OS blocks other writers.
"""
def __init__(self) -> None:
self._fd: int | None = None # Unix fd or Windows HANDLE
def open(self, path: Path, *, create: bool = False) -> None:
"""Open *path* for read+write with an exclusive lock.
Args:
path: File to open and lock.
create: If True, create the file if it doesn't exist (bootstrap).
Raises:
SystemExit: If the file is locked by another process or not found.
"""
if self._fd is not None:
return # Already open (idempotent)
if sys.platform == "win32":
self._open_win32(path, create)
else:
self._open_unix(path, create)
def open_and_read(self, path: Path) -> bytes:
"""Open *path* with exclusive lock and read all content.
Combined operation for efficient use with asyncio.to_thread().
"""
self.open(path)
return self.read()
def read(self) -> bytes:
"""Read the entire file content from the beginning."""
if self._fd is None:
raise RuntimeError("LockedFile.read() called on a closed file")
if sys.platform == "win32":
return self._read_win32()
else:
return self._read_unix()
def write(self, data: bytes) -> None:
"""Append *data* to the end of the file."""
if self._fd is None:
raise RuntimeError("LockedFile.write() called on a closed file")
if sys.platform == "win32":
self._write_win32(data)
else:
self._write_unix(data)
def close(self) -> None:
"""Release the lock and close the file."""
if self._fd is None:
return
if sys.platform == "win32":
_kernel32.CloseHandle(self._fd)
else:
os.close(self._fd)
self._fd = None
@property
def is_open(self) -> bool:
return self._fd is not None
# -- Unix ----------------------------------------------------------------
def _open_unix(self, path: Path, create: bool) -> None:
flags = os.O_RDWR | (os.O_CREAT if create else 0)
try:
fd = os.open(path, flags, 0o666)
except FileNotFoundError:
_fatal(f"Database file not found: {path.resolve()}")
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
except OSError:
os.close(fd)
_fatal(f"🛑 {path.resolve()}: database already locked by another instance")
self._fd = fd
def _read_unix(self) -> bytes:
os.lseek(self._fd, 0, os.SEEK_SET)
chunks = []
while True:
chunk = os.read(self._fd, 1 << 20) # 1 MiB
if not chunk:
break
chunks.append(chunk)
return b"".join(chunks)
def _write_unix(self, data: bytes) -> None:
os.lseek(self._fd, 0, os.SEEK_END)
os.write(self._fd, data)
# -- Windows -------------------------------------------------------------
def _open_win32(self, path: Path, create: bool) -> None:
disposition = _OPEN_ALWAYS if create else _OPEN_EXISTING
handle = _kernel32.CreateFileW(
str(path),
_GENERIC_READ | _GENERIC_WRITE,
_FILE_SHARE_READ,
None,
disposition,
_FILE_ATTRIBUTE_NORMAL,
None,
)
if _is_invalid_handle(handle):
err = ctypes.get_last_error()
if err == _ERROR_SHARING_VIOLATION:
_fatal(
f"🛑 {path.resolve()}: database already locked by another instance"
)
_fatal(f"Failed to open database {path.resolve()}: Windows error {err}")
self._fd = handle
def _read_win32(self) -> bytes:
_kernel32.SetFilePointer(self._fd, 0, None, _FILE_BEGIN)
size = _kernel32.GetFileSize(self._fd, None)
if size == _INVALID_FILE_SIZE:
raise OSError(
f"GetFileSize failed: Windows error {ctypes.get_last_error()}"
)
if size == 0:
return b""
buf = ctypes.create_string_buffer(size)
bytes_read = wintypes.DWORD()
ok = _kernel32.ReadFile(self._fd, buf, size, ctypes.byref(bytes_read), None)
if not ok:
raise OSError(f"ReadFile failed: Windows error {ctypes.get_last_error()}")
return buf.raw[: bytes_read.value]
def _write_win32(self, data: bytes) -> None:
_kernel32.SetFilePointer(self._fd, 0, None, _FILE_END)
written = wintypes.DWORD()
ok = _kernel32.WriteFile(
self._fd,
data,
len(data),
ctypes.byref(written),
None,
)
if not ok:
raise OSError(f"WriteFile failed: Windows error {ctypes.get_last_error()}")
+98 -56
View File
@@ -2,6 +2,7 @@
JSONL persistence layer for the database. JSONL persistence layer for the database.
""" """
import asyncio
import copy import copy
import logging import logging
import os import os
@@ -13,13 +14,18 @@ from pathlib import Path
from typing import Any from typing import Any
from uuid import UUID from uuid import UUID
import aiofiles
import jsondiff import jsondiff
import msgspec import msgspec
from paskia.db.filelock import LockedFile
from paskia.db.logging import log_change from paskia.db.logging import log_change
from paskia.db.migrations import DBVER, apply_all_migrations from paskia.db.migrations import (
from paskia.db.structs import DB, SessionContext DBVER,
MigrationCtx,
apply_all_migrations,
apply_migrations_readonly,
)
from paskia.db.structs import DB, Config, SessionContext
_logger = logging.getLogger(__name__) _logger = logging.getLogger(__name__)
@@ -27,6 +33,47 @@ _logger = logging.getLogger(__name__)
DB_PATH_DEFAULT = "paskia.jsonl" DB_PATH_DEFAULT = "paskia.jsonl"
def load_readonly(db_path: str, *, rp_id: str = "localhost") -> DB:
"""Replay JSONL and apply migrations to produce a DB, without writing anything.
This is suitable for reading settings before the server starts.
Migrations are applied in-memory only; nothing is queued or flushed.
"""
path = Path(db_path)
if not path.exists():
return DB(config=Config(rp_id=rp_id))
data_dict: dict = {}
version = 0
try:
with open(path, "rb") as f:
content = f.read()
for line_num, line in enumerate(content.split(b"\n"), 1):
line = line.strip()
if not line:
continue
try:
change = msgspec.json.decode(line)
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
version = change.get("v", 0)
except Exception as e:
raise ValueError(f"Error parsing line {line_num}: {e}")
except OSError as e:
raise SystemExit(f"Failed to load database: {e}")
except (ValueError, msgspec.DecodeError) as e:
raise SystemExit(f"Failed to load database: {e}")
if not data_dict:
return DB(config=Config(rp_id=rp_id))
# Apply migrations in-memory (no persistence)
apply_migrations_readonly(data_dict, version, MigrationCtx(rp_id=rp_id))
# Decode to msgspec struct
db = msgspec.json.decode(msgspec.json.encode(data_dict), type=DB)
return db
class _ChangeRecord(msgspec.Struct, omit_defaults=True): class _ChangeRecord(msgspec.Struct, omit_defaults=True):
"""A single change record in the JSONL file.""" """A single change record in the JSONL file."""
@@ -71,54 +118,6 @@ def create_change_record(
# Actions that are allowed to create a new database file # Actions that are allowed to create a new database file
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap"}) _BOOTSTRAP_ACTIONS = frozenset({"bootstrap"})
# Flag to prevent duplicate error messages on fatal flush failure
_flush_failed = False
async def flush_changes(
db_path: Path,
pending_changes: deque[_ChangeRecord],
) -> None:
"""Write all pending changes to disk.
Args:
db_path: Path to the JSONL database file
pending_changes: Queue of pending change records (will be cleared on success)
On failure, logs an error and sends SIGTERM to trigger graceful shutdown.
"""
global _flush_failed
if _flush_failed or not pending_changes:
return
if not db_path.exists():
first_action = pending_changes[0].a
if first_action not in _BOOTSTRAP_ACTIONS:
_logger.error(
"Refusing to create database file with action '%s' - "
"only bootstrap can create a new database",
first_action,
)
_flush_failed = True
os.kill(os.getpid(), signal.SIGTERM)
return
changes_to_write = list(pending_changes)
try:
lines = [_change_encoder.encode(change) for change in changes_to_write]
if not lines:
pending_changes.clear()
return
async with aiofiles.open(db_path, "ab") as f:
await f.write(b"\n".join(lines) + b"\n")
pending_changes.clear()
except OSError as e:
_logger.error("Failed to flush database: %s", e)
_flush_failed = True
os.kill(os.getpid(), signal.SIGTERM)
class JsonlStore: class JsonlStore:
"""JSONL persistence layer for a DB instance.""" """JSONL persistence layer for a DB instance."""
@@ -126,6 +125,8 @@ class JsonlStore:
def __init__(self, db: DB, db_path: str = DB_PATH_DEFAULT): def __init__(self, db: DB, db_path: str = DB_PATH_DEFAULT):
self.db: DB = db self.db: DB = db
self.db_path = Path(db_path) self.db_path = Path(db_path)
self._file = LockedFile()
self._flush_failed = False
self._previous_builtins: dict[str, Any] = {} self._previous_builtins: dict[str, Any] = {}
self._pending_changes: deque[_ChangeRecord] = deque() self._pending_changes: deque[_ChangeRecord] = deque()
self._current_action: str = "system" self._current_action: str = "system"
@@ -144,11 +145,12 @@ class JsonlStore:
if not self.db_path.exists(): if not self.db_path.exists():
return return
# Open with exclusive write lock and read contents — single threadpool call
content = await asyncio.to_thread(self._file.open_and_read, self.db_path)
# Replay change log to reconstruct state # Replay change log to reconstruct state
data_dict: dict = {} data_dict: dict = {}
try: try:
async with aiofiles.open(self.db_path, "rb") as f:
content = await f.read()
for line_num, line in enumerate(content.split(b"\n"), 1): for line_num, line in enumerate(content.split(b"\n"), 1):
line = line.strip() line = line.strip()
if not line: if not line:
@@ -179,7 +181,10 @@ class JsonlStore:
# Apply schema migrations one at a time # Apply schema migrations one at a time
await apply_all_migrations( await apply_all_migrations(
data_dict, self._current_version, persist_migration, rp_id=rp_id data_dict,
self._current_version,
persist_migration,
MigrationCtx(rp_id=rp_id),
) )
# Decode to msgspec struct # Decode to msgspec struct
@@ -289,5 +294,42 @@ class JsonlStore:
self._transaction_snapshot = None self._transaction_snapshot = None
async def flush(self) -> None: async def flush(self) -> None:
"""Write all pending changes to disk.""" """Write all pending changes to disk.
await flush_changes(self.db_path, self._pending_changes)
On failure, logs an error and sends SIGTERM to trigger graceful shutdown.
"""
if self._flush_failed or not self._pending_changes:
return
if not self._file.is_open:
first_action = self._pending_changes[0].a
if first_action not in _BOOTSTRAP_ACTIONS:
_logger.error(
"Refusing to create database file with action '%s' - "
"only bootstrap can create a new database",
first_action,
)
self._flush_failed = True
os.kill(os.getpid(), signal.SIGTERM)
return
# Bootstrap: create and open the file with lock
await asyncio.to_thread(self._file.open, self.db_path, create=True)
changes_to_write = list(self._pending_changes)
try:
lines = [_change_encoder.encode(change) for change in changes_to_write]
if not lines:
self._pending_changes.clear()
return
await asyncio.to_thread(self._file.write, b"\n".join(lines) + b"\n")
self._pending_changes.clear()
except OSError as e:
_logger.error("Failed to flush database: %s", e)
self._flush_failed = True
os.kill(os.getpid(), signal.SIGTERM)
def close(self) -> None:
"""Release the file lock and close the file."""
self._file.close()
+1 -2
View File
@@ -7,6 +7,7 @@ import os
from datetime import UTC, datetime from datetime import UTC, datetime
import paskia.db.operations as _ops import paskia.db.operations as _ops
from paskia import oidc_notify
from paskia.authsession import EXPIRES from paskia.authsession import EXPIRES
_logger = logging.getLogger(__name__) _logger = logging.getLogger(__name__)
@@ -31,8 +32,6 @@ def cleanup_expired() -> int:
limit = now - EXPIRES limit = now - EXPIRES
expired_sessions = [k for k, s in _ops._db.sessions.items() if s.validated < limit] expired_sessions = [k for k, s in _ops._db.sessions.items() if s.validated < limit]
if expired_sessions: if expired_sessions:
from paskia import oidc_notify # noqa: PLC0415
oidc_notify.schedule_notifications(expired_sessions) oidc_notify.schedule_notifications(expired_sessions)
with _ops._db.transaction("expiry"): with _ops._db.transaction("expiry"):
for k in expired_sessions: for k in expired_sessions:
+30 -8
View File
@@ -8,28 +8,36 @@ Each migration should be idempotent and only run when needed.
import base64 import base64
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
import msgspec
from paskia.util.crypto import secret_key from paskia.util.crypto import secret_key
def migrate_v1(d: dict, **kwargs) -> None: class MigrationCtx(msgspec.Struct):
"""Context passed to each migration function."""
rp_id: str
def migrate_v1(d: dict, ctx: MigrationCtx) -> None:
"""Remove Org.created_at fields.""" """Remove Org.created_at fields."""
for org_data in d["orgs"].values(): for org_data in d["orgs"].values():
org_data.pop("created_at", None) org_data.pop("created_at", None)
def migrate_v2(d: dict, *, rp_id: str = "localhost") -> None: def migrate_v2(d: dict, ctx: MigrationCtx) -> None:
"""Add config field if missing.""" """Add config field if missing."""
if "config" not in d: if "config" not in d:
d["config"] = {"rp_id": rp_id} d["config"] = {"rp_id": ctx.rp_id}
def migrate_v3(d: dict, **kwargs) -> None: def migrate_v3(d: dict, ctx: MigrationCtx) -> None:
"""Ensure all users have visits field.""" """Ensure all users have visits field."""
for user_data in d["users"].values(): for user_data in d["users"].values():
user_data.setdefault("visits", 0) user_data.setdefault("visits", 0)
def migrate_v4(d: dict, **kwargs) -> None: def migrate_v4(d: dict, ctx: MigrationCtx) -> None:
"""OpenID Connect support and hardened session keys.""" """OpenID Connect support and hardened session keys."""
# Session keys changed to hashes, drop old sessions # Session keys changed to hashes, drop old sessions
d["sessions"] = {} d["sessions"] = {}
@@ -45,14 +53,28 @@ migrations = sorted(
DBVER = len(migrations) # Used by bootstrap to set initial version DBVER = len(migrations) # Used by bootstrap to set initial version
def apply_migrations_readonly(
data_dict: dict,
current_version: int,
ctx: MigrationCtx,
) -> int:
"""Apply migration functions in-place without persistence.
Returns the new version after all migrations.
"""
while current_version < DBVER:
migrations[current_version](data_dict, ctx)
current_version += 1
return current_version
async def apply_all_migrations( async def apply_all_migrations(
data_dict: dict, data_dict: dict,
current_version: int, current_version: int,
persist: Callable[[str, int, dict], Awaitable[None]], persist: Callable[[str, int, dict], Awaitable[None]],
*, ctx: MigrationCtx,
rp_id: str = "localhost",
) -> None: ) -> None:
while current_version < DBVER: while current_version < DBVER:
migrations[current_version](data_dict, rp_id=rp_id) migrations[current_version](data_dict, ctx)
current_version += 1 current_version += 1
await persist(f"migrate:v{current_version}", current_version, data_dict) await persist(f"migrate:v{current_version}", current_version, data_dict)
+3 -5
View File
@@ -11,9 +11,9 @@ 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.config import SESSION_LIFETIME from paskia.config import SESSION_LIFETIME
from paskia.db.jsonl import ( from paskia.db.jsonl import (
JsonlStore, JsonlStore,
@@ -484,7 +484,6 @@ def delete_session(
""" """
if key not in _db.sessions: if key not in _db.sessions:
raise ValueError("Session not found") raise ValueError("Session not found")
from paskia import oidc_notify # noqa: PLC0415
oidc_notify.schedule_notifications([key]) oidc_notify.schedule_notifications([key])
with _db.transaction(action, ctx): with _db.transaction(action, ctx):
@@ -503,7 +502,6 @@ def delete_sessions_for_user(
user = _db.users.get(user_uuid) user = _db.users.get(user_uuid)
if not user: if not user:
return return
from paskia import oidc_notify # noqa: PLC0415
keys = [s.key for s in user.sessions] keys = [s.key for s in user.sessions]
oidc_notify.schedule_notifications(keys) oidc_notify.schedule_notifications(keys)
@@ -589,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,
@@ -657,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:
+3 -2
View File
@@ -1,7 +1,7 @@
import logging import logging
from uuid import UUID from uuid import UUID
from fastapi import Body, FastAPI, HTTPException, Query, Request, Response from fastapi import Body, FastAPI, HTTPException, Query, Request
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from paskia import aaguid as aaguid_mod from paskia import aaguid as aaguid_mod
@@ -14,6 +14,7 @@ from paskia.db import User as UserDC
from paskia.db.operations import _UNSET from paskia.db.operations import _UNSET
from paskia.db.structs import Client from paskia.db.structs import Client
from paskia.fastapi import authz from paskia.fastapi import authz
from paskia.fastapi.front import frontend
from paskia.fastapi.response import MsgspecResponse from paskia.fastapi.response import MsgspecResponse
from paskia.fastapi.session import AUTH_COOKIE from paskia.fastapi.session import AUTH_COOKIE
from paskia.globals import passkey from paskia.globals import passkey
@@ -78,7 +79,7 @@ async def general_exception_handler(_request, exc: Exception): # pragma: no cov
@app.get("/") @app.get("/")
async def adminapp(request: Request, auth=AUTH_COOKIE): async def adminapp(request: Request, auth=AUTH_COOKIE):
return Response(*await vitedev.read("/auth/admin/index.html")) return await vitedev.handle(request, frontend, "/auth/admin/")
# -------------------- Organizations -------------------- # -------------------- Organizations --------------------
+16 -20
View File
@@ -20,7 +20,7 @@ from paskia.fastapi import authz, session, user
from paskia.fastapi.response import MsgspecResponse from paskia.fastapi.response import MsgspecResponse
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME, get_client_ip from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME, get_client_ip
from paskia.globals import passkey as global_passkey from paskia.globals import passkey as global_passkey
from paskia.util import hostutil, htmlutil, passphrase, userinfo, vitedev from paskia.util import hostutil, htmlutil, passphrase, userinfo
from paskia.util.apistructs import ApiSettings, ApiTokenInfo, ApiValidateResponse from paskia.util.apistructs import ApiSettings, ApiTokenInfo, ApiValidateResponse
bearer_auth = HTTPBearer(auto_error=False) bearer_auth = HTTPBearer(auto_error=False)
@@ -161,25 +161,15 @@ async def forward_authentication(
# Clear cookie only if session is invalid (not for reauth) # Clear cookie only if session is invalid (not for reauth)
if e.clear_session: if e.clear_session:
session.clear_session_cookie(response) session.clear_session_cookie(response)
# Browser request? - return full-page HTML with metadata patched into data attrs
# Check Accept header to decide response format if "text/html" in request.headers.get("accept", ""):
accept = request.headers.get("accept", "") return await htmlutil.patched_html_response(
wants_html = "text/html" in accept request, "/int/forward/", e.status_code, mode=e.mode, **e.metadata
if wants_html:
# Browser request - return full-page HTML with metadata
data_attrs = {"mode": e.mode, **e.metadata}
html = (await vitedev.read("/int/forward/index.html"))[0]
html = htmlutil.patch_html_data_attrs(html, **data_attrs)
return Response(
html, status_code=e.status_code, media_type="text/html; charset=UTF-8"
)
else:
# API request - return JSON with iframe srcdoc HTML
return JSONResponse(
status_code=e.status_code,
content=await authz.auth_error_content(e),
) )
# API request - return JSON with iframe srcdoc HTML
return JSONResponse(
status_code=e.status_code, content=await authz.auth_error_content(e)
)
@app.get("/settings") @app.get("/settings")
@@ -214,7 +204,12 @@ async def api_user_info(
) )
ctx = db.data().session_ctx(auth, request.headers.get("host")) ctx = db.data().session_ctx(auth, request.headers.get("host"))
if not ctx: if not ctx:
raise HTTPException(401, "Session expired") raise authz.AuthException(
status_code=401,
detail="Session expired",
mode="login",
clear_session=True,
)
return MsgspecResponse( return MsgspecResponse(
await userinfo.build_user_info( await userinfo.build_user_info(
@@ -244,6 +239,7 @@ async def token_info(credentials=Depends(bearer_auth)):
ApiTokenInfo( ApiTokenInfo(
token_type=reset_token.token_type, token_type=reset_token.token_type,
display_name=u.display_name, display_name=u.display_name,
theme=u.theme,
) )
) )
+10
View File
@@ -0,0 +1,10 @@
from pathlib import Path
from fastapi_vue import Frontend
# Vue Frontend static files
frontend = Frontend(
Path(__file__).parent.parent / "frontend-build",
cached=["/auth/assets/"],
favicon="/paskia.webp",
)
+21 -17
View File
@@ -6,13 +6,18 @@ from pathlib import Path
from fastapi import FastAPI, HTTPException, Request, Response from fastapi import FastAPI, HTTPException, Request, Response
from fastapi.responses import FileResponse, RedirectResponse from fastapi.responses import FileResponse, RedirectResponse
from fastapi_vue import Frontend
from paskia import authcode, globals from paskia import authcode, db, globals
from paskia.__main__ import DEVMODE from paskia.__main__ import DEVMODE
from paskia.bootstrap import bootstrap_if_needed
from paskia.db import start_background, stop_background from paskia.db import start_background, stop_background
from paskia.db.background import flush
from paskia.db.logging import configure_db_logging from paskia.db.logging import configure_db_logging
from paskia.db.structs import Config
from paskia.fastapi import admin, api, auth_host, oid, ws from paskia.fastapi import admin, api, auth_host, oid, ws
# Import frontend instance
from paskia.fastapi.front import frontend
from paskia.fastapi.logging import AccessLogMiddleware, configure_access_logging from paskia.fastapi.logging import AccessLogMiddleware, configure_access_logging
from paskia.fastapi.session import AUTH_COOKIE from paskia.fastapi.session import AUTH_COOKIE
from paskia.util import hostutil, passphrase, vitedev from paskia.util import hostutil, passphrase, vitedev
@@ -23,14 +28,6 @@ configure_db_logging()
_access_logger = logging.getLogger("paskia.access") _access_logger = logging.getLogger("paskia.access")
# Vue Frontend static files
frontend = Frontend(
Path(__file__).parent.parent / "frontend-build",
cached=["/auth/assets/"],
favicon="/paskia.webp",
)
# Path to examples/index.html when running from source tree # Path to examples/index.html when running from source tree
_EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples" _EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples"
@@ -46,7 +43,6 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
config = json.loads(os.environ["PASKIA_CONFIG"]) config = json.loads(os.environ["PASKIA_CONFIG"])
try: try:
# CLI (__main__) performs bootstrap once; here we skip to avoid duplicate work
await globals.init( await globals.init(
rp_id=config["rp_id"], rp_id=config["rp_id"],
rp_name=config["rp_name"], rp_name=config["rp_name"],
@@ -58,6 +54,15 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
# Re-raise to fail fast # Re-raise to fail fast
raise raise
# Bootstrap and persist config now that the full DB is loaded
cli_config_data = config.get("cli_config")
if cli_config_data:
cli_config = Config(**cli_config_data)
await bootstrap_if_needed(config=cli_config)
if config.get("save"):
await db.update_config(cli_config)
await flush()
# Restore uvicorn info logging (suppressed during startup in dev mode) # Restore uvicorn info logging (suppressed during startup in dev mode)
# Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages # Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages
if app.debug: if app.debug:
@@ -131,9 +136,9 @@ async def openid_configuration(request: Request):
@app.get("/auth/restricted/iframe") @app.get("/auth/restricted/iframe")
@app.get("/auth/restricted/oidc") @app.get("/auth/restricted/oidc")
async def restricted_view(): async def restricted_view(request: Request):
"""Serve the restricted/authentication UI for iframe or OpenID Connect.""" """Serve the restricted/authentication UI for iframe or OpenID Connect."""
return Response(*await vitedev.read("/auth/restricted/index.html")) return await vitedev.handle(request, frontend, "/auth/restricted/")
# Navigable URLs are defined here. We support both / and /auth/ as the base path # Navigable URLs are defined here. We support both / and /auth/ as the base path
@@ -148,7 +153,7 @@ async def frontapp(request: Request, response: Response, auth=AUTH_COOKIE):
The frontend handles mode detection (host mode vs full profile) based on settings. The frontend handles mode detection (host mode vs full profile) based on settings.
Access control is handled via APIs. Access control is handled via APIs.
""" """
return Response(*await vitedev.read("/auth/index.html")) return await vitedev.handle(request, frontend, "/auth/")
@app.get("/admin", include_in_schema=False) @app.get("/admin", include_in_schema=False)
@@ -180,14 +185,13 @@ async def examples_page():
# Frontend static files - must be before /{token} catch-all routes # Frontend static files - must be before /{token} catch-all routes
# (actual routes registered during lifespan after frontend.load())
frontend.route(app, "/") frontend.route(app, "/")
# Note: this catch-all handler must be the last route defined # Note: this catch-all handler must be the last route defined
@app.get("/{token}") @app.get("/{token}")
@app.get("/auth/{token}") @app.get("/auth/{token}")
async def token_link(token: str): async def token_link(request: Request, token: str):
"""Serve the reset app for reset tokens (password reset / device addition). """Serve the reset app for reset tokens (password reset / device addition).
The frontend will validate the token via /auth/api/token-info. The frontend will validate the token via /auth/api/token-info.
@@ -195,4 +199,4 @@ async def token_link(token: str):
if not passphrase.is_well_formed(token): if not passphrase.is_well_formed(token):
raise HTTPException(status_code=404) raise HTTPException(status_code=404)
return Response(*await vitedev.read("/int/reset/index.html")) return await vitedev.handle(request, frontend, "/int/reset/")
+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"],
+2 -1
View File
@@ -171,11 +171,12 @@ class ApiSettings(msgspec.Struct):
version: str version: str
class ApiTokenInfo(msgspec.Struct): class ApiTokenInfo(msgspec.Struct, omit_defaults=True):
"""Token info response struct.""" """Token info response struct."""
token_type: str token_type: str
display_name: str display_name: str
theme: str = ""
class ApiUuidResponse(msgspec.Struct): class ApiUuidResponse(msgspec.Struct):
+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 -1
View File
@@ -11,7 +11,7 @@ __all__ = ["path", "file", "read", "is_dev_mode"]
def _get_dev_server() -> str | None: def _get_dev_server() -> str | None:
"""Get the dev server URL from environment, or None if not in dev mode.""" """Get the dev server URL from environment, or None if not in dev mode."""
return os.environ.get("FASTAPI_VUE_FRONTEND_URL") or None return os.environ.get("PASKIA_VITE_URL") or None
def _resolve_static_dir() -> Path: def _resolve_static_dir() -> Path:
+33
View File
@@ -2,6 +2,39 @@
import re import re
from paskia.fastapi.front import frontend
from paskia.util import vitedev
async def patched_html_response(request, filepath: str, status_code: int, **data_attrs):
"""Fetch HTML from vitedev and patch with data attributes.
Strips caching/compression headers from request to get raw content,
patches the HTML body with data attributes, and strips caching headers
from response.
Args:
request: The FastAPI Request object
filepath: Path to HTML file, e.g. "/int/forward/"
status_code: HTTP status code for the response
**data_attrs: Key-value pairs for data attributes
Returns:
Patched Response object, or original response if not 200.
"""
resp = await vitedev.handle(request, frontend, filepath)
# Pass through non-200 responses
if resp.status_code != 200:
return resp
# Patch HTML with data attrs and strip caching headers from response
resp.body = patch_html_data_attrs(resp.body, **data_attrs)
resp.status_code = status_code
strip_headers = {b"etag", b"last-modified", b"content-length"}
resp.raw_headers = [
(k, v) for k, v in resp.raw_headers if k.lower() not in strip_headers
]
return resp
def patch_html_data_attrs(html: bytes, **data_attrs: str) -> bytes: def patch_html_data_attrs(html: bytes, **data_attrs: str) -> bytes:
"""Patch HTML by adding data attributes to the <html> tag. """Patch HTML by adding data attributes to the <html> tag.
+1 -1
View File
@@ -69,7 +69,7 @@ def print_startup_config(config: "PaskiaConfig") -> None:
lines.append(line(f"Auth Host: {config.auth_host}")) lines.append(line(f"Auth Host: {config.auth_host}"))
# Show frontend URL if in dev mode # Show frontend URL if in dev mode
devmode = os.environ.get("FASTAPI_VUE_FRONTEND_URL") devmode = os.environ.get("PASKIA_VITE_URL")
if devmode: if devmode:
lines.append(line(f"Dev Frontend: {devmode}")) lines.append(line(f"Dev Frontend: {devmode}"))
+4
View File
@@ -5,8 +5,10 @@ from paskia.db import SessionContext
from paskia.util import hostutil from paskia.util import hostutil
from paskia.util.apistructs import ( from paskia.util.apistructs import (
ApiAaguidInfo, ApiAaguidInfo,
ApiOrg,
ApiOrgContext, ApiOrgContext,
ApiPermission, ApiPermission,
ApiRole,
ApiRoleContext, ApiRoleContext,
ApiSessionContext, ApiSessionContext,
ApiUser, ApiUser,
@@ -64,4 +66,6 @@ async def build_user_info(
permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions} permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions}
if ctx if ctx
else {}, else {},
org=ApiOrg.from_db(ctx.org) if ctx else None,
role=ApiRole.from_db(ctx.role) if ctx else None,
) )
+35 -31
View File
@@ -1,39 +1,29 @@
"""Vite dev server proxy for fetching frontend files during development. """Vite dev server proxy for fetching frontend files during development.
In dev mode (FASTAPI_VUE_FRONTEND_URL set), fetches files from Vite. In dev mode (PASKIA_VITE_URL set), fetches files from Vite.
In production, reads from the static build directory. In production, reads from the static build directory.
This complements fastapi_vue.Frontend which handles static file serving This complements fastapi_vue.Frontend which handles static file serving
but doesn't provide server-side fetching of HTML content. but doesn't provide server-side fetching of HTML content.
""" """
import asyncio
import mimetypes import mimetypes
import os import os
from importlib import resources from importlib import resources
from pathlib import Path from pathlib import Path
import httpx import httpx
from fastapi import Response
__all__ = ["read"] __all__ = ["handle"]
def _get_dev_server() -> str | None:
"""Get the dev server URL from environment, or None if not in dev mode."""
return os.environ.get("FASTAPI_VUE_FRONTEND_URL") or None
def _resolve_static_dir() -> Path: def _resolve_static_dir() -> Path:
"""Resolve the static files directory."""
# Try packaged path via importlib.resources (works for wheel/installed). # Try packaged path via importlib.resources (works for wheel/installed).
try: # pragma: no cover - trivial path resolution pkg_dir = resources.files("paskia") / "frontend-build"
pkg_dir = resources.files("paskia") / "frontend-build" fs_path = Path(str(pkg_dir))
fs_path = Path(str(pkg_dir)) if fs_path.is_dir():
if fs_path.is_dir(): return fs_path
return fs_path
except Exception: # pragma: no cover - defensive
pass
# Fallback for editable/development before build. # Fallback for editable/development before build.
return Path(__file__).parent.parent / "frontend-build" return Path(__file__).parent.parent / "frontend-build"
@@ -41,31 +31,45 @@ def _resolve_static_dir() -> Path:
_static_dir: Path = _resolve_static_dir() _static_dir: Path = _resolve_static_dir()
async def read(filepath: str) -> tuple[bytes, int, dict[str, str]]: async def handle(request, frontend, filepath: str):
"""Read file content and return response tuple. """Read file content and return Response.
In dev mode, fetches from the Vite dev server. In dev mode, fetches from the Vite dev server.
In production, reads from the static build directory. In production, uses frontend.handle.
Args: Args:
request: The FastAPI Request object
frontend: The fastapi_vue.Frontend instance
filepath: Path relative to frontend root, e.g. "/auth/index.html" filepath: Path relative to frontend root, e.g. "/auth/index.html"
Returns: Returns:
Tuple of (content, status_code, headers) suitable for FastAPI Response object.
FastAPI Response(*args).
""" """
dev_server = _get_dev_server() if dev_server := os.environ.get("PASKIA_VITE_URL"):
if dev_server:
async with httpx.AsyncClient() as client: async with httpx.AsyncClient() as client:
resp = await client.get(f"{dev_server}{filepath}") resp = await client.get(f"{dev_server}{filepath}")
resp.raise_for_status() resp.raise_for_status()
mime = resp.headers.get("content-type", "application/octet-stream") mime = resp.headers.get("content-type", "application/octet-stream")
# Strip charset suffix if present # Strip charset suffix if present
mime = mime.split(";")[0].strip() mime = mime.split(";")[0].strip()
return resp.content, resp.status_code, {"content-type": mime} return Response(resp.content, resp.status_code, {"content-type": mime})
else:
# Production: read from static build # Read from frontend cache directly to bypass any compression/processing
file_path = _static_dir / filepath.lstrip("/") cached_content = getattr(frontend, "_files", {}).get(filepath)
content = await asyncio.to_thread(file_path.read_bytes) if cached_content is not None:
mime, _ = mimetypes.guess_type(str(file_path)) mime, _ = mimetypes.guess_type(filepath)
return content, 200, {"content-type": mime or "application/octet-stream"} return Response(
cached_content, 200, {"content-type": mime or "application/octet-stream"}
)
# Fallback to frontend.handle for cache negotiation
# Strip accept-encoding to get uncompressed content (needed for HTML patching)
strip_headers = {b"accept-encoding", b"if-none-match", b"if-modified-since"}
request.scope["headers"] = [
(k, v) for k, v in request.scope["headers"] if k.lower() not in strip_headers
]
# Invalidate cached Headers object (it doesn't re-read scope after first access)
if hasattr(request, "_headers"):
del request._headers
return frontend.handle(request, filepath)
+17 -28
View File
@@ -11,20 +11,28 @@ keywords = [ "forward_auth", "auth_request", "FastAPI" ]
authors = [ authors = [
{name = "Leo Vasanko"}, {name = "Leo Vasanko"},
] ]
requires-python = ">=3.11"
dependencies = [ dependencies = [
"fastapi[standard]>=0.104.1", "fastapi[standard]>=0.129.0",
"websockets>=12.0", "websockets>=16.0",
"webauthn>=1.11.1", "webauthn>=2.7.1",
"base64url>=1.0.0", "base64url>=1.1.1",
"uuid7-standard>=1.0.0", "uuid7-standard>=1.1.0",
"pyjwt[crypto]>=2.8.0", "pyjwt[crypto]>=2.11.0",
"jsondiff>=2.2.1", "jsondiff>=2.2.1",
"msgspec>=0.20.0", "msgspec>=0.20.0",
"aiofiles>=25.1.0", "fastapi-vue>=1.1.0",
"fastapi-vue>=0.3.0",
"ua-parser[regex]>=1.0.1", "ua-parser[regex]>=1.0.1",
] ]
requires-python = ">=3.11" [dependency-groups]
dev = [
"coverage>=7.13.4",
"httpx>=0.28.1",
"pytest>=9.0.2",
"pytest-asyncio>=1.3.0",
"pytest-cov>=7.0.0",
"ruff>=0.15.1",
]
[project.urls] [project.urls]
Homepage = "https://git.zi.fi/LeoVasanko/paskia" Homepage = "https://git.zi.fi/LeoVasanko/paskia"
@@ -36,15 +44,6 @@ source = "vcs"
[tool.hatch.build.hooks.vcs] [tool.hatch.build.hooks.vcs]
version-file = "paskia/_version.py" version-file = "paskia/_version.py"
[project.optional-dependencies]
dev = [
"ruff>=0.1.0",
"coverage[toml]>=7.0.0",
"pytest>=8.0.0",
"pytest-asyncio>=0.24.0",
"httpx>=0.27.0",
]
[tool.coverage.run] [tool.coverage.run]
source = ["paskia"] source = ["paskia"]
branch = true branch = true
@@ -75,16 +74,6 @@ select = ["E", "F", "I", "N", "W", "UP", "PLC0415"]
ignore = ["E501"] # Line too long ignore = ["E501"] # Line too long
isort.known-first-party = ["paskia"] isort.known-first-party = ["paskia"]
[dependency-groups]
dev = [
"coverage>=7.12.0",
"httpx>=0.28.1",
"pytest>=9.0.1",
"pytest-asyncio>=1.3.0",
"pytest-cov>=7.0.0",
"ruff>=0.14.8",
]
[project.scripts] [project.scripts]
paskia = "paskia.__main__:main" paskia = "paskia.__main__:main"
+16 -20
View File
@@ -153,29 +153,9 @@ async def run_devserver(args: argparse.Namespace, remaining: list[str]) -> None:
paskia.extend(["--origin", origin]) paskia.extend(["--origin", origin])
paskia.extend(remaining) paskia.extend(remaining)
# Compute origins for Caddy
caddy_origins = []
if args.auth_host:
auth_host = args.auth_host
if "://" not in auth_host:
auth_host = f"https://{auth_host}"
caddy_origins.append(auth_host)
caddy_origins.append(f"https://{args.rp_id}")
if args.origins:
for origin in args.origins:
if "://" not in origin:
origin = f"https://{origin}"
caddy_origins.append(origin)
if not args.auth_host and not args.origins:
caddy_origins.append(f"https://{args.rp_id}")
# Remove duplicates while preserving order
seen = set()
caddy_origins = [x for x in caddy_origins if not (x in seen or seen.add(x))]
# Set environment for subprocesses # Set environment for subprocesses
os.environ["PASKIA_VITE_URL"] = viteurl os.environ["PASKIA_VITE_URL"] = viteurl
os.environ["PASKIA_BACKEND_URL"] = backurl os.environ["PASKIA_BACKEND_URL"] = backurl
os.environ["PASKIA_SITE_URL"] = caddy_origins[0] if args.caddy else viteurl
os.environ["PASKIA_DEV"] = "1" os.environ["PASKIA_DEV"] = "1"
if args.auth_host: if args.auth_host:
os.environ["PASKIA_AUTH_HOST"] = args.auth_host os.environ["PASKIA_AUTH_HOST"] = args.auth_host
@@ -183,6 +163,22 @@ async def run_devserver(args: argparse.Namespace, remaining: list[str]) -> None:
async with ProcessGroup() as pg: async with ProcessGroup() as pg:
# Start Caddy first if requested (needs to bind ports) # Start Caddy first if requested (needs to bind ports)
if args.caddy: if args.caddy:
caddy_origins = []
if args.auth_host:
auth_host = args.auth_host
if "://" not in auth_host:
auth_host = f"https://{auth_host}"
caddy_origins.append(auth_host)
caddy_origins.append(f"https://{args.rp_id}")
if args.origins:
for origin in args.origins:
if "://" not in origin:
origin = f"https://{origin}"
caddy_origins.append(origin)
if not caddy_origins:
caddy_origins.append(f"https://{args.rp_id}")
seen: set = set()
caddy_origins = [x for x in caddy_origins if not (x in seen or seen.add(x))]
caddy_proc = await run_caddy(caddy_origins, viteurl, backurl) caddy_proc = await run_caddy(caddy_origins, viteurl, backurl)
pg._procs.append(caddy_proc) pg._procs.append(caddy_proc)
pg._cmds[caddy_proc.pid] = "caddy" pg._cmds[caddy_proc.pid] = "caddy"
+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 -21
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
@@ -188,25 +187,6 @@ class TestExceptionHandlers:
assert "iframe" in data["auth"] assert "iframe" in data["auth"]
# -------------------- Admin App Root --------------------
class TestAdminAppRoot:
"""Tests for the admin app root endpoint"""
@pytest.mark.asyncio
async def test_admin_app_root_with_auth(
self, client: httpx.AsyncClient, session_token: str
):
"""Admin app root returns HTML when authenticated."""
response = await client.get(
"/auth/api/admin/",
headers={**auth_headers(session_token), "Host": "localhost:4401"},
)
assert response.status_code == 200
assert "text/html" in response.headers.get("content-type", "")
# -------------------- Organization Tests -------------------- # -------------------- Organization Tests --------------------
@@ -1320,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"},
-28
View File
@@ -355,34 +355,6 @@ class TestErrorHandling:
class TestForwardAuthHtmlResponse: class TestForwardAuthHtmlResponse:
"""Tests for forward auth HTML responses""" """Tests for forward auth HTML responses"""
@pytest.mark.asyncio
async def test_forward_401_html_response(self, client: httpx.AsyncClient):
"""Forward auth 401 should return HTML page for browser requests."""
response = await client.get(
"/auth/api/forward",
headers={"Accept": "text/html"},
)
assert response.status_code == 401
assert "text/html" in response.headers.get("content-type", "")
# HTML response should contain the mode data attribute
assert b"data-mode" in response.content or b"mode" in response.content
@pytest.mark.asyncio
async def test_forward_403_html_response(
self, client: httpx.AsyncClient, regular_session_token: str
):
"""Forward auth 403 should return HTML page for browser requests."""
response = await client.get(
"/auth/api/forward?perm=auth:admin",
headers={
**auth_headers(regular_session_token),
"Host": "localhost:4401",
"Accept": "text/html",
},
)
assert response.status_code == 403
assert "text/html" in response.headers.get("content-type", "")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_forward_with_expired_session_clears_cookie( async def test_forward_with_expired_session_clears_cookie(
self, client: httpx.AsyncClient self, client: httpx.AsyncClient