Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
727625ef4f | ||
|
|
5f7a5ed9b1 | ||
|
|
d64e63527b | ||
|
|
733439b446 | ||
|
|
4e6f63e9ef | ||
|
|
d431c75297 | ||
|
|
6257071efe | ||
|
|
7958b6f365 | ||
|
|
b3cb540098 | ||
|
|
22ba7231b1 | ||
|
|
9a9979fb62 | ||
|
|
9b7855c0af | ||
|
|
dfc4c76d43 | ||
|
|
e1f0fdf664 | ||
|
|
f26ac8f33b | ||
|
|
880ced3b8c | ||
|
|
af80b5eefc | ||
|
|
fa1e69d58b | ||
|
|
39000ef831 | ||
|
|
49119fac81 |
@@ -216,7 +216,7 @@ test.describe('API Mode - 401 Login Flow', () => {
|
||||
await clearSessionCookie(page)
|
||||
|
||||
// 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')
|
||||
|
||||
// Verify it's in login mode (not reauth)
|
||||
@@ -268,7 +268,7 @@ test.describe('API Mode - 401 Login Flow', () => {
|
||||
await setupTestHarness(page)
|
||||
|
||||
// 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
|
||||
await waitForAuthIframe(page)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import { spawn } from 'child_process'
|
||||
import { execSync, spawn } from 'child_process'
|
||||
import { join, dirname } from 'path'
|
||||
import { existsSync, mkdirSync, writeFileSync } from 'fs'
|
||||
import { fileURLToPath } from 'url'
|
||||
@@ -31,6 +31,11 @@ export default async function globalSetup() {
|
||||
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...')
|
||||
if (COLLECT_COVERAGE) {
|
||||
console.log(' 📊 Coverage collection enabled for Python backend')
|
||||
|
||||
+21
-72
@@ -13,8 +13,8 @@
|
||||
<script setup>
|
||||
import { computed, onMounted, onUnmounted, ref } from 'vue'
|
||||
import { useAuthStore } from '@/stores/auth'
|
||||
import { apiJson, SessionValidator, createAuthIframe, removeAuthIframe } from 'paskia'
|
||||
import { getAuthIframeUrl } from '@/utils/api'
|
||||
import { apiJson, SessionValidator } from 'paskia'
|
||||
import { updateThemeFromSession } from '@/utils/theme'
|
||||
import StatusMessage from '@/components/StatusMessage.vue'
|
||||
import ProfileView from '@/components/ProfileView.vue'
|
||||
import HostProfileView from '@/components/HostProfileView.vue'
|
||||
@@ -48,90 +48,49 @@ const isHostMode = computed(() => {
|
||||
return currentHost !== configuredHost
|
||||
})
|
||||
|
||||
function terminateSession() {
|
||||
function onSessionLost(e) {
|
||||
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 sessionValidator = new SessionValidator(userUuidGetter, terminateSession)
|
||||
const sessionValidator = new SessionValidator(userUuidGetter, onSessionLost)
|
||||
|
||||
onMounted(() => sessionValidator.start())
|
||||
onUnmounted(() => sessionValidator.stop())
|
||||
|
||||
async function loadUserInfo() {
|
||||
viewState.value = 'loading'
|
||||
loadingMessage.value = 'Loading...'
|
||||
try {
|
||||
// apiJson handles 401/403 with auth.iframe automatically:
|
||||
// shows overlay iframe, waits for auth, retries the request.
|
||||
const [validateData, userInfoData] = await Promise.all([
|
||||
apiJson('/auth/api/validate', { method: 'POST' }),
|
||||
apiJson('/auth/api/user-info', { method: 'GET' })
|
||||
])
|
||||
store.userInfo = userInfoData
|
||||
store.ctx = validateData.ctx
|
||||
updateThemeFromSession(store.userInfo)
|
||||
// Verify that the user UUIDs match between user-info and validate responses
|
||||
if (store.userInfo.user.uuid !== store.ctx.user.uuid) {
|
||||
console.error('User UUID mismatch between user-info and validate responses')
|
||||
window.location.reload()
|
||||
return false
|
||||
return
|
||||
}
|
||||
viewState.value = 'profile'
|
||||
return true
|
||||
} catch {
|
||||
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
|
||||
} catch (e) {
|
||||
onSessionLost(e)
|
||||
}
|
||||
}
|
||||
|
||||
onMounted(async () => {
|
||||
// Listen for postMessage from auth iframe
|
||||
window.addEventListener('message', handleAuthMessage)
|
||||
|
||||
// Load settings
|
||||
await store.loadSettings()
|
||||
|
||||
@@ -145,17 +104,7 @@ onMounted(async () => {
|
||||
document.title = inHostMode ? `${rpName} · Account summary` : rpName
|
||||
}
|
||||
|
||||
// Try to load user info
|
||||
const success = await loadUserInfo()
|
||||
|
||||
if (!success) {
|
||||
// Need authentication - show login iframe
|
||||
showAuthIframe()
|
||||
}
|
||||
})
|
||||
|
||||
onUnmounted(() => {
|
||||
window.removeEventListener('message', handleAuthMessage)
|
||||
removeAuthIframe()
|
||||
// Load user info (apiJson handles auth iframe if needed)
|
||||
await loadUserInfo()
|
||||
})
|
||||
</script>
|
||||
|
||||
@@ -14,6 +14,7 @@ import AdminDialogs from '@/admin/AdminDialogs.vue'
|
||||
import { useAuthStore } from '@/stores/auth'
|
||||
import { adminUiPath, makeUiHref } from '@/utils/settings'
|
||||
import { apiJson, SessionValidator } from 'paskia'
|
||||
import { updateThemeFromSession } from '@/utils/theme'
|
||||
import { uuidv7 } from 'uuidv7'
|
||||
import { getDirection } from '@/utils/keynav'
|
||||
import { goBack } from '@/utils/helpers'
|
||||
@@ -197,6 +198,7 @@ function orgUserCount(org) {
|
||||
async function loadUserInfo() {
|
||||
const data = await apiJson('/auth/api/validate', { method: 'POST' })
|
||||
info.value = data
|
||||
updateThemeFromSession(data.ctx)
|
||||
authenticated.value = true
|
||||
}
|
||||
|
||||
@@ -337,24 +339,9 @@ async function moveUserToRole(userUuid, user, targetRoleUuid) {
|
||||
}
|
||||
}
|
||||
|
||||
function onUserDragStart(e, userUuid, org) {
|
||||
e.dataTransfer.effectAllowed = 'move'
|
||||
e.dataTransfer.setData('text/plain', JSON.stringify({ user_uuid: userUuid, org }))
|
||||
}
|
||||
|
||||
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 */ }
|
||||
function moveUserToRoleFromDrag(userUuid, newRoleUuid) {
|
||||
const user = selectedOrg.value?.users?.[userUuid]
|
||||
if (user) moveUserToRole(userUuid, user, newRoleUuid)
|
||||
}
|
||||
|
||||
// Role actions
|
||||
@@ -999,10 +986,8 @@ async function submitDialog() {
|
||||
@create-user-in-role="createUserInRole"
|
||||
@open-user="openUser"
|
||||
@toggle-role-permission="toggleRolePermission"
|
||||
@on-role-drag-over="onRoleDragOver"
|
||||
@move-user-to-role="moveUserToRoleFromDrag"
|
||||
@navigate-out="handlePanelNavigateOut"
|
||||
@on-role-drop="onRoleDrop"
|
||||
@on-user-drag-start="onUserDragStart"
|
||||
/>
|
||||
|
||||
<AdminOidcDetail
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<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">
|
||||
</head>
|
||||
<body>
|
||||
|
||||
@@ -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'
|
||||
|
||||
function getTheme() {
|
||||
const params = new URLSearchParams(location.hash.slice(1))
|
||||
return params.get('theme') || getCachedTheme() || ''
|
||||
return getCachedTheme() || params.get('theme') || ''
|
||||
}
|
||||
|
||||
// Apply theme class to document root
|
||||
|
||||
@@ -8,7 +8,7 @@
|
||||
|
||||
<main class="view-root">
|
||||
<div class="surface surface--tight reset-container">
|
||||
<header class="view-header reset-header">
|
||||
<header class="view-header center">
|
||||
<h1>🔑 Registration</h1>
|
||||
<p class="view-lede">
|
||||
{{ subtitleMessage }}
|
||||
@@ -60,6 +60,7 @@ import { computed, onMounted, reactive, ref } from 'vue'
|
||||
import passkey from '@/utils/passkey'
|
||||
import { getSettings, uiBasePath } from '@/utils/settings'
|
||||
import { apiJson, ApiError, getUserFriendlyErrorMessage } from 'paskia'
|
||||
import { updateThemeFromSession } from '@/utils/theme'
|
||||
|
||||
const status = reactive({
|
||||
show: false,
|
||||
@@ -80,7 +81,7 @@ const sessionDescriptor = computed(() => tokenInfo.value?.token_type || 'your en
|
||||
const subtitleMessage = computed(() => {
|
||||
if (initializing.value) return 'Preparing your secure enrollment…'
|
||||
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())
|
||||
@@ -117,6 +118,7 @@ async function fetchTokenInfo() {
|
||||
headers: { 'Authorization': `Bearer ${token.value}` },
|
||||
})
|
||||
displayName.value = tokenInfo.value.display_name
|
||||
if (tokenInfo.value.theme) updateThemeFromSession({ user: { theme: tokenInfo.value.theme } })
|
||||
} catch (error) {
|
||||
console.error('Failed to load token info', error)
|
||||
const message = error instanceof ApiError
|
||||
@@ -201,14 +203,14 @@ onMounted(async () => {
|
||||
</script>
|
||||
|
||||
<style scoped>
|
||||
main.view-root { min-height: 100vh; align-items: center; justify-content: center; padding: 2rem 1rem; }
|
||||
.reset-container {
|
||||
max-width: 560px;
|
||||
max-width: 520px;
|
||||
margin: 0 auto;
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
.reset-header {
|
||||
text-align: center;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 1.75rem;
|
||||
}
|
||||
|
||||
.section-body {
|
||||
|
||||
@@ -15,7 +15,8 @@
|
||||
"qrcode": "^1.5.4",
|
||||
"sirv": "^3.0.2",
|
||||
"uuidv7": "^1.1.0",
|
||||
"vue": "^3.5.17"
|
||||
"vue": "^3.5.17",
|
||||
"vuedraggable": "^4.1.0"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@vitejs/plugin-vue": "^6.0.0",
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
<script setup>
|
||||
import { computed, ref } from 'vue'
|
||||
import draggable from 'vuedraggable'
|
||||
import { getDirection, navigateButtonRow, focusPreferred } from '@/utils/keynav'
|
||||
|
||||
const props = defineProps({
|
||||
@@ -8,7 +9,7 @@ const props = defineProps({
|
||||
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
|
||||
const orgTitleRef = ref(null)
|
||||
@@ -50,6 +51,14 @@ function roleUserCount(roleUuid) {
|
||||
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) {
|
||||
return props.permissions.find(p => p.scope === scope)?.display_name || scope
|
||||
}
|
||||
@@ -350,8 +359,6 @@ defineExpose({ focusFirstElement })
|
||||
v-for="(r, roleIndex) in sortedRoles"
|
||||
:key="r.uuid"
|
||||
class="role-column"
|
||||
@dragover="$emit('onRoleDragOver', $event)"
|
||||
@drop="e => $emit('onRoleDrop', e, selectedOrg, r)"
|
||||
>
|
||||
<div class="role-header" @keydown="e => handleRoleHeaderKeydown(e, roleIndex)">
|
||||
<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>
|
||||
</div>
|
||||
</div>
|
||||
<template v-if="roleUserCount(r.uuid) > 0">
|
||||
<ul class="user-list" @keydown="handleUserListKeydown">
|
||||
<li
|
||||
v-for="u in roleUsers(r.uuid)"
|
||||
:key="u.uuid"
|
||||
class="user-chip"
|
||||
tabindex="0"
|
||||
draggable="true"
|
||||
@dragstart="e => $emit('onUserDragStart', e, u.uuid, selectedOrg.uuid)"
|
||||
@click="$emit('openUser', u)"
|
||||
@keydown.enter="$emit('openUser', u)"
|
||||
:title="u.uuid"
|
||||
>
|
||||
<span class="name">{{ u.display_name }}</span>
|
||||
<span class="meta">{{ u.last_seen ? new Date(u.last_seen).toLocaleDateString() : '—' }}</span>
|
||||
</li>
|
||||
</ul>
|
||||
</template>
|
||||
<div v-else class="empty-role">
|
||||
<p class="empty-text muted">No members</p>
|
||||
<div class="user-list-wrapper">
|
||||
<draggable
|
||||
:list="roleUsers(r.uuid)"
|
||||
group="users"
|
||||
item-key="uuid"
|
||||
tag="ul"
|
||||
class="user-list"
|
||||
@change="evt => onUserChange(evt, r.uuid)"
|
||||
@keydown="handleUserListKeydown"
|
||||
>
|
||||
<template #item="{ element: u }">
|
||||
<li
|
||||
class="user-chip"
|
||||
tabindex="0"
|
||||
@click="$emit('openUser', u)"
|
||||
@keydown.enter="$emit('openUser', u)"
|
||||
:title="u.uuid"
|
||||
>
|
||||
<span class="name">{{ u.display_name }}</span>
|
||||
<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>
|
||||
@@ -391,23 +404,27 @@ defineExpose({ focusFirstElement })
|
||||
|
||||
<style scoped>
|
||||
.card.surface { padding: var(--space-lg); }
|
||||
.org-title { display: flex; align-items: center; gap: var(--space-sm); margin-bottom: var(--space-lg); }
|
||||
.org-name { font-size: 1.5rem; font-weight: 600; color: var(--color-heading); }
|
||||
.org-title { display: flex; align-items: center; gap: var(--space-sm); margin-bottom: var(--space-lg); font-size: 1.65rem; }
|
||||
.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 span { writing-mode: vertical-rl; transform: rotate(180deg); font-size: 0.65rem; }
|
||||
.perm-matrix-grid .add-role-head { cursor: pointer; }
|
||||
.roles-grid { display: flex; gap: var(--space-lg); margin-top: var(--space-lg); }
|
||||
.role-column { flex: 1; min-width: 200px; border-radius: var(--radius-md); padding: var(--space-md); }
|
||||
.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: 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-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); }
|
||||
.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); }
|
||||
.user-list { list-style: none; padding: 0; margin: 0; display: flex; flex-direction: column; gap: var(--space-xs); }
|
||||
.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-wrapper { position: relative; flex: 1; display: flex; flex-direction: column; min-height: 5.5rem; }
|
||||
.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 .meta { font-size: 0.7rem; color: var(--color-text-muted); }
|
||||
.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 .meta { font-size: 0.7rem; color: rgba(255, 255, 255, 0.8); }
|
||||
.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; }
|
||||
|
||||
@media (max-width: 720px) {
|
||||
|
||||
@@ -823,9 +823,6 @@ th {
|
||||
|
||||
.user-info {
|
||||
display: grid;
|
||||
border-radius: var(--radius-md);
|
||||
background: var(--color-surface);
|
||||
padding: 1rem;
|
||||
}
|
||||
|
||||
.user-details {
|
||||
|
||||
@@ -10,9 +10,9 @@
|
||||
<UserBasicInfo
|
||||
v-if="ctx"
|
||||
:name="ctx.user.display_name"
|
||||
:visits="authStore.userInfo?.visits || 0"
|
||||
:created-at="authStore.userInfo?.created_at"
|
||||
:last-seen="authStore.userInfo?.last_seen"
|
||||
:visits="authStore.userInfo.user.visits"
|
||||
:created-at="authStore.userInfo.user.created_at"
|
||||
:last-seen="authStore.userInfo.user.last_seen"
|
||||
:email="ctx.user.email"
|
||||
:telephone="ctx.user.telephone"
|
||||
:org-display-name="orgDisplayName"
|
||||
|
||||
@@ -22,8 +22,8 @@
|
||||
:created-at="authStore.userInfo.user.created_at"
|
||||
:last-seen="authStore.userInfo.user.last_seen"
|
||||
:loading="authStore.isLoading"
|
||||
:org-display-name="authStore.ctx?.org.display_name"
|
||||
:role-name="authStore.ctx?.role.display_name"
|
||||
:org-display-name="authStore.userInfo.org.display_name"
|
||||
:role-name="authStore.userInfo.role.display_name"
|
||||
update-endpoint="/auth/api/user/info"
|
||||
@saved="authStore.loadUserInfo()"
|
||||
@edit="openEditDialog"
|
||||
@@ -53,7 +53,7 @@
|
||||
<CredentialList
|
||||
ref="credentialList"
|
||||
:credentials="credentials"
|
||||
:aaguid-info="authStore.userInfo?.aaguid_info || {}"
|
||||
:aaguid-info="authStore.userInfo.aaguid_info"
|
||||
:loading="authStore.isLoading"
|
||||
:hovered-credential-uuid="hoveredCredentialUuid"
|
||||
:hovered-session-credential-uuid="hoveredSession?.credential"
|
||||
@@ -184,11 +184,11 @@ const hasActiveModal = computed(() => showEditDialog.value || showRegLink.value)
|
||||
|
||||
watch(showEditDialog, (open) => {
|
||||
if (!open) return
|
||||
const user = authStore.userInfo?.user
|
||||
editName.value = user?.display_name ?? ''
|
||||
editEmail.value = user?.email ?? ''
|
||||
editUsername.value = user?.preferred_username ?? ''
|
||||
editTelephone.value = user?.telephone ?? ''
|
||||
const user = authStore.userInfo.user
|
||||
editName.value = user.display_name ?? ''
|
||||
editEmail.value = user.email ?? ''
|
||||
editUsername.value = user.preferred_username ?? ''
|
||||
editTelephone.value = user.telephone ?? ''
|
||||
editError.value = ''
|
||||
})
|
||||
|
||||
@@ -341,7 +341,7 @@ const handleDelete = async (credential) => {
|
||||
|
||||
const rpName = computed(() => authStore.settings?.rp_name || 'this service')
|
||||
const paskiaVersion = computed(() => authStore.settings?.version || '')
|
||||
const sessions = computed(() => authStore.userInfo?.sessions || {})
|
||||
const sessions = computed(() => authStore.userInfo.sessions)
|
||||
const currentSessionHost = computed(() => {
|
||||
const currentSession = Object.values(sessions.value).find(session => session.is_current)
|
||||
return currentSession?.host || 'this host'
|
||||
@@ -365,12 +365,12 @@ const logoutEverywhere = async () => { await authStore.logoutEverywhere() }
|
||||
const logout = async () => { await authStore.logout() }
|
||||
const openEditDialog = () => { showEditDialog.value = true }
|
||||
const isAdmin = computed(() => {
|
||||
const perms = authStore.ctx?.permissions
|
||||
return perms?.includes('auth:admin') || perms?.includes('auth:org:admin')
|
||||
const perms = Object.values(authStore.userInfo.permissions).map(p => p.scope)
|
||||
return perms.includes('auth:admin') || perms.includes('auth:org:admin')
|
||||
})
|
||||
const hasMultipleSessions = computed(() => Object.keys(sessions.value).length > 1)
|
||||
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(() => {
|
||||
// Check if any single site has more than 8 sessions
|
||||
|
||||
@@ -37,7 +37,7 @@
|
||||
<script setup>
|
||||
import { ref, computed, onMounted, onUnmounted, nextTick } from '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 { getDirection } from '@/utils/keynav'
|
||||
import { useAuthStore } from '@/stores/auth'
|
||||
@@ -90,7 +90,9 @@ async function generateLink() {
|
||||
emit('close')
|
||||
}
|
||||
} 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')
|
||||
}
|
||||
}
|
||||
|
||||
@@ -163,6 +163,9 @@ async function startRemoteAuth() {
|
||||
|
||||
// PoW challenge
|
||||
const powChallenge = await ws.receive_json()
|
||||
if (powChallenge.status) {
|
||||
throw new Error(powChallenge.detail || `Failed to connect: ${powChallenge.status}`)
|
||||
}
|
||||
if (powChallenge.pow) {
|
||||
const challenge = b64dec(powChallenge.pow.challenge)
|
||||
const nonces = await solvePoW(challenge, powChallenge.pow.work)
|
||||
|
||||
@@ -61,6 +61,7 @@ import { getSettings, uiBasePath } from '@/utils/settings'
|
||||
import { fetchJson, getUserFriendlyErrorMessage } from 'paskia'
|
||||
import RemoteAuthRequest from '@/components/RemoteAuthRequest.vue'
|
||||
import { focusDialogButton } from '@/utils/keynav'
|
||||
import { updateThemeFromSession } from '@/utils/theme'
|
||||
|
||||
const props = defineProps({
|
||||
mode: {
|
||||
@@ -147,6 +148,7 @@ async function fetchSettings() {
|
||||
async function validateSession() {
|
||||
try {
|
||||
session.value = await fetchJson('/auth/api/validate', { method: 'POST' })
|
||||
updateThemeFromSession(session.value?.ctx)
|
||||
if (isAuthenticated.value && props.mode !== 'reauth') {
|
||||
currentView.value = 'forbidden'
|
||||
emit('forbidden', session.value)
|
||||
|
||||
@@ -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-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; }
|
||||
.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:active:not(:disabled) { transform: translateY(1px); }
|
||||
.mini-btn:disabled { opacity: 0.5; cursor: not-allowed; }
|
||||
|
||||
@@ -88,7 +88,7 @@ export const useAuthStore = defineStore('auth', {
|
||||
async loadUserInfo() {
|
||||
try {
|
||||
this.userInfo = await apiJson('/auth/api/user-info', { method: 'GET' })
|
||||
updateThemeFromSession(this.ctx)
|
||||
updateThemeFromSession(this.userInfo)
|
||||
console.log('User info loaded:', this.userInfo)
|
||||
} catch (error) {
|
||||
// Suppress toast for 401/403 errors - the auth iframe will handle these
|
||||
|
||||
@@ -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')
|
||||
}
|
||||
@@ -18,8 +18,6 @@ export {
|
||||
isAuthIframeOpen,
|
||||
hideAuthIframe,
|
||||
showAuthIframe,
|
||||
createAuthIframe,
|
||||
removeAuthIframe,
|
||||
} from './overlay'
|
||||
|
||||
export { SessionValidator } from './validate'
|
||||
|
||||
+49
-126
@@ -1,21 +1,16 @@
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import msgspec
|
||||
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.db.background import flush
|
||||
from paskia.db.structs import Config
|
||||
from paskia.db.jsonl import load_readonly
|
||||
from paskia.util import startupbox
|
||||
from paskia.util.hostutil import normalize_origin
|
||||
from paskia.util.runtime import RuntimeConfig
|
||||
|
||||
DEFAULT_PORT = 4401
|
||||
DEVMODE = os.getenv("PASKIA_DEV") == "1"
|
||||
@@ -95,138 +90,66 @@ def main():
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Handle clearing options
|
||||
if getattr(args, "auth_host", None) == "":
|
||||
args.auth_host = None
|
||||
if getattr(args, "rp_name", None) == "":
|
||||
args.rp_name = None
|
||||
if getattr(args, "listen", None) == "":
|
||||
args.listen = None
|
||||
# Load stored config (read-only, no writes, no global state)
|
||||
db_path = os.environ.get("PASKIA_DB", f"{args.rp_id}.paskiadb")
|
||||
config = load_readonly(db_path, rp_id=args.rp_id).config
|
||||
|
||||
# Init db and load stored config
|
||||
asyncio.run(db.init(rp_id=args.rp_id))
|
||||
stored_config = db.data().config
|
||||
# Override stored config with CLI args, or clear with empty string
|
||||
if args.rp_name is not None:
|
||||
config.rp_name = args.rp_name or None
|
||||
if args.auth_host is not None:
|
||||
config.auth_host = args.auth_host or None
|
||||
if args.origins is not None:
|
||||
config.origins = None if args.origins == [""] else args.origins
|
||||
if args.listen is not None:
|
||||
config.listen = None if args.listen == [""] else args.listen
|
||||
|
||||
# Apply defaults from stored config
|
||||
if args.rp_name is None and stored_config.rp_name is not None:
|
||||
args.rp_name = stored_config.rp_name
|
||||
if args.origins is None and stored_config.origins is not None:
|
||||
args.origins = stored_config.origins
|
||||
if args.auth_host is None and stored_config.auth_host is not None:
|
||||
args.auth_host = stored_config.auth_host
|
||||
if args.listen is None and stored_config.listen is not None:
|
||||
args.listen = stored_config.listen
|
||||
# Process and normalize auth_host
|
||||
if config.auth_host:
|
||||
if "://" not in config.auth_host:
|
||||
config.auth_host = f"https://{config.auth_host}"
|
||||
config.auth_host = config.auth_host.rstrip("/")
|
||||
validate_auth_host(config.auth_host, config.rp_id)
|
||||
if config.origins:
|
||||
config.origins.insert(0, config.auth_host) # Ensure first in origins
|
||||
|
||||
# Parse first endpoint for config display and site_url
|
||||
first_listen = args.listen[0] if isinstance(args.listen, list) else args.listen
|
||||
endpoints = parse_endpoint(first_listen, DEFAULT_PORT)
|
||||
# Normalize and deduplicate while preserving order
|
||||
if config.origins:
|
||||
config.origins = list({normalize_origin(o): ... for o in config.origins})
|
||||
|
||||
# Extract host/port/uds from first endpoint for config display and site_url
|
||||
ep = endpoints[0] if endpoints else {}
|
||||
host = ep.get("host")
|
||||
# Parse first endpoint for site_url fallback
|
||||
ep = next(iter(parse_endpoints(config.listen, DEFAULT_PORT)), {})
|
||||
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:
|
||||
# Normalize auth_host with scheme
|
||||
if "://" not in args.auth_host:
|
||||
args.auth_host = f"https://{args.auth_host}"
|
||||
|
||||
validate_auth_host(args.auth_host, args.rp_id)
|
||||
|
||||
# If origins are configured, ensure auth_host is included at top
|
||||
if 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
|
||||
# Priority: PASKIA_SITE_URL (explicit) > auth_host > first origin with localhost > http://localhost:port
|
||||
explicit_site_url = os.environ.get("PASKIA_SITE_URL")
|
||||
if explicit_site_url:
|
||||
# Explicit site URL from devserver or deployment config
|
||||
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 = "/"
|
||||
elif origins:
|
||||
# Find localhost origin if rp_id is localhost, else use first origin
|
||||
localhost_origin = (
|
||||
next((o for o in origins if "://localhost" in o), None)
|
||||
if args.rp_id == "localhost"
|
||||
else None
|
||||
)
|
||||
site_url = (localhost_origin or origins[0]).rstrip("/")
|
||||
site_path = "/auth/"
|
||||
elif args.rp_id == "localhost" and port:
|
||||
# Dev mode: use http with port
|
||||
site_url = f"http://localhost:{port}"
|
||||
site_path = "/auth/"
|
||||
# Compute site_url and site_path
|
||||
# Priority: auth_host > origins[0] > PASKIA_VITE_URL > http://localhost:port > https://rp_id
|
||||
site_path = "/auth/"
|
||||
if config.auth_host:
|
||||
site_url, site_path = config.auth_host, "/"
|
||||
elif config.origins:
|
||||
site_url = config.origins[0]
|
||||
elif vite_url := os.environ.get("PASKIA_VITE_URL"):
|
||||
site_url = vite_url.rstrip("/") # Devserver
|
||||
elif config.rp_id == "localhost" and port:
|
||||
site_url = f"http://localhost:{port}" # Backend directly if we can
|
||||
else:
|
||||
site_url = f"https://{args.rp_id}"
|
||||
site_path = "/auth/"
|
||||
site_url = f"https://{config.rp_id}" # Assume external reverse proxy
|
||||
|
||||
# Build runtime configuration
|
||||
config = PaskiaConfig(
|
||||
rp_id=args.rp_id,
|
||||
rp_name=args.rp_name or None,
|
||||
origins=origins or None,
|
||||
auth_host=args.auth_host or None,
|
||||
# Build runtime configuration for the server
|
||||
runtime = RuntimeConfig(
|
||||
config=config,
|
||||
site_url=site_url,
|
||||
site_path=site_path,
|
||||
host=host,
|
||||
port=port,
|
||||
uds=uds,
|
||||
save=args.save,
|
||||
)
|
||||
startupbox.print_startup_config(runtime)
|
||||
os.environ["PASKIA_CONFIG"] = msgspec.json.encode(runtime).decode()
|
||||
|
||||
# Export configuration via single JSON env variable for worker processes
|
||||
config_json = {
|
||||
"rp_id": config.rp_id,
|
||||
"rp_name": config.rp_name,
|
||||
"origins": config.origins,
|
||||
"auth_host": config.auth_host,
|
||||
"site_url": config.site_url,
|
||||
"site_path": config.site_path,
|
||||
}
|
||||
os.environ["PASKIA_CONFIG"] = json.dumps(config_json)
|
||||
|
||||
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())
|
||||
|
||||
# Run the server (spawns processes in dev mode)
|
||||
dev = {"reload": True, "reload_dirs": ["paskia"]} if DEVMODE else {}
|
||||
server.run(
|
||||
"paskia.fastapi.mainapp:app",
|
||||
listen=args.listen,
|
||||
listen=config.listen,
|
||||
default_port=DEFAULT_PORT,
|
||||
log_level="warning",
|
||||
access_log=False,
|
||||
|
||||
@@ -23,6 +23,11 @@ if TYPE_CHECKING:
|
||||
EXPIRES = SESSION_LIFETIME
|
||||
|
||||
|
||||
def session_ctx(auth: str, host: str | None = None):
|
||||
"""Get session context with normalized host."""
|
||||
return db.data().session_ctx(auth, hostutil.normalize_host(host))
|
||||
|
||||
|
||||
def expires() -> datetime:
|
||||
return datetime.now(UTC) + EXPIRES
|
||||
|
||||
@@ -42,7 +47,7 @@ def get_reset(token: str) -> "ResetToken":
|
||||
|
||||
def delete_credential(credential_uuid: UUID, auth: str, host: str | None = None):
|
||||
"""Delete a specific credential for the current user."""
|
||||
ctx = db.data().session_ctx(auth, hostutil.normalize_host(host))
|
||||
ctx = session_ctx(auth, host)
|
||||
if not ctx:
|
||||
raise ValueError("Session expired")
|
||||
db.delete_credential(credential_uuid, ctx.user.uuid)
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
|
||||
# Shared configuration constants for session management.
|
||||
@@ -6,19 +5,3 @@ SESSION_LIFETIME = timedelta(hours=24)
|
||||
|
||||
# Lifetime for reset links created by admins
|
||||
RESET_LIFETIME = timedelta(days=14)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PaskiaConfig:
|
||||
"""Runtime configuration for the Paskia authentication server."""
|
||||
|
||||
rp_id: str
|
||||
rp_name: str | None
|
||||
origins: list[str] | None
|
||||
auth_host: str | None
|
||||
site_url: str # Base URL without trailing path (e.g. https://example.com)
|
||||
site_path: str # Path to auth UI: "/" if auth_host, else "/auth/"
|
||||
# Listen address (one of host:port or uds)
|
||||
host: str | None = None
|
||||
port: int | None = None
|
||||
uds: str | None = None
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""
|
||||
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.
|
||||
Write: Functions validate and commit, or raise ValueError.
|
||||
|
||||
@@ -10,7 +10,6 @@ Usage:
|
||||
|
||||
# Read (after init)
|
||||
user_data = db.data().users[user_uuid]
|
||||
user = db.build_user(user_uuid)
|
||||
|
||||
# Context
|
||||
ctx = db.data().session_ctx(session_key)
|
||||
@@ -27,6 +26,7 @@ from paskia.db.background import (
|
||||
stop_cleanup,
|
||||
)
|
||||
from paskia.db.bootstrap import bootstrap
|
||||
from paskia.db.jsonl import load_readonly
|
||||
from paskia.db.lifecycle import cleanup_expired, init
|
||||
from paskia.db.operations import (
|
||||
add_permission_to_org,
|
||||
@@ -102,18 +102,12 @@ __all__ = [
|
||||
# Instance
|
||||
"data",
|
||||
"init",
|
||||
"load_readonly",
|
||||
# Background
|
||||
"start_background",
|
||||
"stop_background",
|
||||
"start_cleanup",
|
||||
"stop_cleanup",
|
||||
# Builders
|
||||
"build_credential",
|
||||
"build_permission",
|
||||
"build_reset_token",
|
||||
"build_role",
|
||||
"build_session",
|
||||
"build_user",
|
||||
# Read ops
|
||||
# Write ops
|
||||
"add_permission_to_org",
|
||||
|
||||
@@ -21,7 +21,7 @@ _background_task: asyncio.Task | None = None
|
||||
|
||||
async def flush() -> None:
|
||||
"""Write all pending database changes to disk."""
|
||||
store = _ops._store
|
||||
store = _ops._db._store
|
||||
if store is None:
|
||||
_logger.warning("flush() called but _store is None")
|
||||
return
|
||||
@@ -48,6 +48,10 @@ async def _background_loop():
|
||||
cleanup_expired()
|
||||
await flush() # Flush cleanup changes
|
||||
last_cleanup = now
|
||||
|
||||
# Conditionally write a snapshot to speed up future startups
|
||||
if _ops._db._store is not None:
|
||||
_ops._db._store.maybe_snapshot()
|
||||
except asyncio.CancelledError:
|
||||
# Final flush before exit
|
||||
await flush()
|
||||
@@ -90,7 +94,7 @@ async def start_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
|
||||
if _background_task:
|
||||
_background_task.cancel()
|
||||
@@ -99,6 +103,7 @@ async def stop_background():
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
_background_task = None
|
||||
_ops._db._store.close()
|
||||
|
||||
|
||||
# Aliases for backwards compatibility
|
||||
|
||||
@@ -7,6 +7,7 @@ from datetime import UTC, datetime
|
||||
import uuid7
|
||||
|
||||
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.util.crypto import secret_key
|
||||
|
||||
@@ -59,8 +60,6 @@ def bootstrap(
|
||||
|
||||
# Set reset token expiry (passphrase generated by ResetToken.create)
|
||||
if reset_expiry is None:
|
||||
from paskia.authsession import reset_expires # noqa: PLC0415
|
||||
|
||||
reset_expiry = reset_expires()
|
||||
|
||||
with _ops._db.transaction("bootstrap"):
|
||||
|
||||
@@ -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()}")
|
||||
+185
-127
@@ -2,6 +2,7 @@
|
||||
JSONL persistence layer for the database.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import logging
|
||||
import os
|
||||
@@ -13,126 +14,137 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
from uuid import UUID
|
||||
|
||||
import aiofiles
|
||||
import jsondiff
|
||||
import msgspec
|
||||
|
||||
from paskia.db.filelock import LockedFile
|
||||
from paskia.db.logging import log_change
|
||||
from paskia.db.migrations import DBVER, apply_all_migrations
|
||||
from paskia.db.structs import DB, SessionContext
|
||||
from paskia.db.migrations import (
|
||||
DBVER,
|
||||
MigrationCtx,
|
||||
apply_all_migrations,
|
||||
apply_migrations_readonly,
|
||||
)
|
||||
from paskia.db.snapshot import SnapshotState
|
||||
from paskia.db.structs import DB, Config, SessionContext
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
# Default database path
|
||||
DB_PATH_DEFAULT = "paskia.jsonl"
|
||||
|
||||
class ReplayResult(msgspec.Struct, frozen=False):
|
||||
"""Return value of _replay_from_data"""
|
||||
|
||||
state: dict
|
||||
v: int = 0
|
||||
ts: datetime | None = None
|
||||
snapts: datetime | None = None
|
||||
changes: int = 0
|
||||
|
||||
|
||||
class _ChangeRecord(msgspec.Struct, omit_defaults=True):
|
||||
"""A single change record in the JSONL file."""
|
||||
class DatabaseError(Exception):
|
||||
"""Exception raised for database loading errors."""
|
||||
|
||||
ts: datetime
|
||||
pass
|
||||
|
||||
|
||||
def _replay_from_data(data: bytes, db_path: str) -> ReplayResult:
|
||||
"""Replay database state from file data, using the last snapshot if available."""
|
||||
resolved_path = str(Path(db_path).resolve())
|
||||
result = ReplayResult(state={})
|
||||
|
||||
# Find and apply the last snapshot
|
||||
snap, start_offset = SnapshotState.load(data)
|
||||
if snap:
|
||||
result.state = snap.state
|
||||
result.v = snap.v
|
||||
result.snapts = snap.ts
|
||||
|
||||
# Replay change records after the snapshot
|
||||
lines = data[start_offset:].split(b"\n")
|
||||
for line_num, raw in enumerate(lines, start=1): # 1-based line numbering
|
||||
line = raw.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
change = msgspec.json.decode(line, type=ChangeRecord)
|
||||
except msgspec.DecodeError as e:
|
||||
raise DatabaseError(f"{resolved_path}:{line_num}: {e}")
|
||||
result.state = jsondiff.patch(result.state, change.diff, marshal=True)
|
||||
result.v = change.v
|
||||
result.ts = change.ts
|
||||
result.changes += 1
|
||||
|
||||
return result
|
||||
|
||||
|
||||
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))
|
||||
|
||||
try:
|
||||
with open(path, "rb") as f:
|
||||
content = f.read()
|
||||
r = _replay_from_data(content, str(path.resolve()))
|
||||
data_dict = r.state
|
||||
version = r.v
|
||||
except OSError as e:
|
||||
_logger.exception("Failed to load database")
|
||||
raise SystemExit(f"{e}")
|
||||
except (ValueError, msgspec.DecodeError, DatabaseError) as e:
|
||||
raise SystemExit(f"{e}")
|
||||
except Exception as e:
|
||||
_logger.exception("Unexpected error loading database")
|
||||
raise SystemExit(f"{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, kw_only=True):
|
||||
ts: datetime = msgspec.field(default_factory=lambda: datetime.now(UTC))
|
||||
a: str # action - describes the operation (e.g., "migrate", "login", "create_user")
|
||||
v: int # schema version after this change
|
||||
v: int = 0 # schema version after this change
|
||||
u: str | None = None # user UUID who performed the action (None for system)
|
||||
diff: dict = {}
|
||||
|
||||
|
||||
# msgspec encoder for change records
|
||||
_change_encoder = msgspec.json.Encoder()
|
||||
diff: dict
|
||||
|
||||
|
||||
def compute_diff(previous: dict, current: dict) -> dict | None:
|
||||
"""Compute JSON diff between two states.
|
||||
|
||||
Args:
|
||||
previous: Previous state (JSON-compatible dict)
|
||||
current: Current state (JSON-compatible dict)
|
||||
|
||||
Returns:
|
||||
The diff, or None if no changes
|
||||
"""
|
||||
diff = jsondiff.diff(previous, current, marshal=True)
|
||||
return diff if diff else None
|
||||
|
||||
|
||||
def create_change_record(
|
||||
action: str, version: int, diff: dict, user: str | None = None
|
||||
) -> _ChangeRecord:
|
||||
"""Create a change record for persistence."""
|
||||
return _ChangeRecord(
|
||||
ts=datetime.now(UTC),
|
||||
a=action,
|
||||
v=version,
|
||||
u=user,
|
||||
diff=diff,
|
||||
)
|
||||
return jsondiff.diff(previous, current, marshal=True) or None
|
||||
|
||||
|
||||
# Actions that are allowed to create a new database file
|
||||
_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:
|
||||
"""JSONL persistence layer for a DB instance."""
|
||||
|
||||
def __init__(self, db: DB, db_path: str = DB_PATH_DEFAULT):
|
||||
def __init__(self, db: DB, db_path: str):
|
||||
self.db: DB = db
|
||||
self.db_path = Path(db_path)
|
||||
self._previous_builtins: dict[str, Any] = {}
|
||||
self._pending_changes: deque[_ChangeRecord] = deque()
|
||||
self._file = LockedFile()
|
||||
self._flush_failed = False
|
||||
self._statedict: dict[str, Any] = {}
|
||||
self._pending_changes: deque[ChangeRecord] = deque()
|
||||
self._current_action: str = "system"
|
||||
self._current_user: str | None = None
|
||||
self._in_transaction: bool = False
|
||||
self._transaction_snapshot: dict[str, Any] | None = None
|
||||
self._current_version: int = DBVER # Schema version for new databases
|
||||
self._v: int = DBVER # Schema version for new databases
|
||||
self._snapshot = SnapshotState()
|
||||
|
||||
async def load(
|
||||
self, db_path: str | None = None, *, rp_id: str = "localhost"
|
||||
@@ -144,55 +156,52 @@ class JsonlStore:
|
||||
if not self.db_path.exists():
|
||||
return
|
||||
|
||||
# Replay change log to reconstruct state
|
||||
data_dict: dict = {}
|
||||
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):
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
change = msgspec.json.decode(line)
|
||||
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
|
||||
self._current_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}")
|
||||
# Open with exclusive write lock and read contents — single threadpool call
|
||||
content = await asyncio.to_thread(self._file.open_and_read, self.db_path)
|
||||
|
||||
if not data_dict:
|
||||
# Replay change log to reconstruct state (snapshot-accelerated)
|
||||
try:
|
||||
r = _replay_from_data(content, str(self.db_path.resolve()))
|
||||
statedict = r.state
|
||||
self._v = r.v
|
||||
self._snapshot.ts = r.snapts
|
||||
self._snapshot.changes = r.changes
|
||||
except (OSError, ValueError, msgspec.DecodeError, DatabaseError) as e:
|
||||
raise SystemExit(f"{e}")
|
||||
except Exception as e:
|
||||
_logger.exception("Unexpected error loading database")
|
||||
raise SystemExit(f"{e}")
|
||||
|
||||
if not statedict:
|
||||
return
|
||||
|
||||
# Set previous state for diffing (will be updated by _queue_change)
|
||||
self._previous_builtins = copy.deepcopy(data_dict)
|
||||
self._statedict = copy.deepcopy(statedict)
|
||||
|
||||
# Callback to persist each migration
|
||||
async def persist_migration(
|
||||
action: str, new_version: int, current: dict
|
||||
) -> None:
|
||||
self._current_version = new_version
|
||||
self._v = new_version
|
||||
self._queue_change(action, new_version, current)
|
||||
|
||||
# Apply schema migrations one at a time
|
||||
await apply_all_migrations(
|
||||
data_dict, self._current_version, persist_migration, rp_id=rp_id
|
||||
statedict,
|
||||
self._v,
|
||||
persist_migration,
|
||||
MigrationCtx(rp_id=rp_id),
|
||||
)
|
||||
|
||||
# Decode to msgspec struct
|
||||
decoder = msgspec.json.Decoder(DB)
|
||||
self.db = decoder.decode(msgspec.json.encode(data_dict))
|
||||
self.db = decoder.decode(msgspec.json.encode(statedict))
|
||||
self.db._store = self
|
||||
|
||||
# Normalize via msgspec round-trip (handles omit_defaults etc.)
|
||||
# This ensures _previous_builtins matches what msgspec would produce
|
||||
normalized_dict = msgspec.to_builtins(self.db)
|
||||
await persist_migration(
|
||||
"migrate:msgspec", self._current_version, normalized_dict
|
||||
)
|
||||
await persist_migration("migrate:msgspec", self._v, normalized_dict)
|
||||
|
||||
def _queue_change(
|
||||
self, action: str, version: int, current: dict, user: str | None = None
|
||||
@@ -205,10 +214,17 @@ class JsonlStore:
|
||||
current: The current state as a plain dict
|
||||
user: Optional user UUID who performed the action
|
||||
"""
|
||||
diff = compute_diff(self._previous_builtins, current)
|
||||
diff = compute_diff(self._statedict, current)
|
||||
if not diff:
|
||||
return
|
||||
self._pending_changes.append(create_change_record(action, version, diff, user))
|
||||
self._pending_changes.append(
|
||||
ChangeRecord(
|
||||
a=action,
|
||||
v=version,
|
||||
u=user,
|
||||
diff=diff,
|
||||
)
|
||||
)
|
||||
|
||||
# Log the change with user display name if available
|
||||
user_display = None
|
||||
@@ -220,8 +236,8 @@ class JsonlStore:
|
||||
except (ValueError, KeyError):
|
||||
user_display = user
|
||||
|
||||
log_change(action, diff, user_display, self._previous_builtins, self.db)
|
||||
self._previous_builtins = copy.deepcopy(current)
|
||||
log_change(action, diff, user_display, self._statedict, self.db)
|
||||
self._statedict = copy.deepcopy(current)
|
||||
|
||||
@contextmanager
|
||||
def transaction(
|
||||
@@ -243,13 +259,13 @@ class JsonlStore:
|
||||
|
||||
# Check for out-of-transaction modifications
|
||||
current_state = msgspec.to_builtins(self.db)
|
||||
if current_state != self._previous_builtins:
|
||||
if current_state != self._statedict:
|
||||
# Allow bootstrap to create a new database from empty state
|
||||
is_bootstrap = action in _BOOTSTRAP_ACTIONS
|
||||
if is_bootstrap and not self._previous_builtins:
|
||||
if is_bootstrap and not self._statedict:
|
||||
pass # Expected: creating database from scratch
|
||||
else:
|
||||
diff = compute_diff(self._previous_builtins, current_state)
|
||||
diff = compute_diff(self._statedict, current_state)
|
||||
diff_json = msgspec.json.encode(diff).decode()
|
||||
_logger.critical(
|
||||
"Database state modified outside of transaction! "
|
||||
@@ -270,7 +286,7 @@ class JsonlStore:
|
||||
yield
|
||||
current = msgspec.to_builtins(self.db)
|
||||
self._queue_change(
|
||||
self._current_action, self._current_version, current, self._current_user
|
||||
self._current_action, self._v, current, self._current_user
|
||||
)
|
||||
except Exception:
|
||||
# Rollback on error: restore from snapshot
|
||||
@@ -289,5 +305,47 @@ class JsonlStore:
|
||||
self._transaction_snapshot = None
|
||||
|
||||
async def flush(self) -> None:
|
||||
"""Write all pending changes to disk."""
|
||||
await flush_changes(self.db_path, self._pending_changes)
|
||||
"""Write all pending changes to disk.
|
||||
|
||||
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 = [msgspec.json.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._snapshot.record_lines(len(lines))
|
||||
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 maybe_snapshot(self) -> None:
|
||||
"""Write a snapshot if conditions are met."""
|
||||
self._snapshot.maybe_write(self._file, self._v, self._statedict)
|
||||
|
||||
def close(self) -> None:
|
||||
"""Release the file lock and close the file."""
|
||||
self._file.close()
|
||||
|
||||
+11
-9
@@ -7,21 +7,25 @@ import os
|
||||
from datetime import UTC, datetime
|
||||
|
||||
import paskia.db.operations as _ops
|
||||
from paskia import oidc_notify
|
||||
from paskia.authsession import EXPIRES
|
||||
from paskia.db.jsonl import JsonlStore
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def init(rp_id: str = "localhost", *args, **kwargs):
|
||||
async def init(rp_id: str, *args, **kwargs):
|
||||
"""Load database from JSONL file."""
|
||||
if _ops._initialized:
|
||||
if _ops._db._store:
|
||||
_logger.debug("Database already initialized, skipping reload")
|
||||
return
|
||||
default_path = f"{rp_id}.paskiadb"
|
||||
db_path = os.environ.get("PASKIA_DB", default_path)
|
||||
await _ops._store.load(db_path, rp_id=rp_id)
|
||||
_ops._db = _ops._store.db
|
||||
_ops._initialized = True
|
||||
db_path = os.environ.get("PASKIA_DB", f"{rp_id}.paskiadb")
|
||||
store = JsonlStore(_ops._db, db_path)
|
||||
await store.load(db_path, rp_id=rp_id)
|
||||
_ops._db = store.db
|
||||
_ops._db._store = store
|
||||
# Request a snapshot after successful startup
|
||||
store._snapshot.request_force()
|
||||
|
||||
|
||||
def cleanup_expired() -> int:
|
||||
@@ -31,8 +35,6 @@ def cleanup_expired() -> int:
|
||||
limit = now - EXPIRES
|
||||
expired_sessions = [k for k, s in _ops._db.sessions.items() if s.validated < limit]
|
||||
if expired_sessions:
|
||||
from paskia import oidc_notify # noqa: PLC0415
|
||||
|
||||
oidc_notify.schedule_notifications(expired_sessions)
|
||||
with _ops._db.transaction("expiry"):
|
||||
for k in expired_sessions:
|
||||
|
||||
+10
-10
@@ -35,9 +35,9 @@ _UNSAFE_CHARS = re.compile(
|
||||
|
||||
# ANSI color codes (matching FastAPI logging style)
|
||||
_RESET = "\033[0m"
|
||||
_DIM = "\033[2m"
|
||||
_PATH_PREFIX = "\033[1;30m" # Dark grey for path prefix (like host in access log)
|
||||
_PATH_FINAL = "\033[0m" # Default for final element (like path in access log)
|
||||
_SEP = "\033[38;5;242m" # Dark grey for separators (like host/timing in access log)
|
||||
_PATH_PREFIX = "\033[38;5;242m" # Dark grey for path prefix (like host in access log)
|
||||
_PATH_FINAL = "\033[38;5;250m" # Default for final element (like path in access log)
|
||||
_DELETE = "\033[1;31m" # Red for deletions
|
||||
_ADD = "\033[0;32m" # Green for additions
|
||||
_ACTION = "\033[1;34m" # Bold blue for action name
|
||||
@@ -317,7 +317,7 @@ def _format_change_lines(
|
||||
# Helper to format a value, checking for censored paths
|
||||
def fmt_value(v: Any, child_path: list[str]) -> str:
|
||||
if child_path[-2:] == ["oidc", "key"]:
|
||||
return f"{_DIM}<hidden>{_RESET}"
|
||||
return f"{_SEP}<hidden>{_RESET}"
|
||||
return _format_value(v, resolver=resolver)
|
||||
|
||||
# Helper to format path with UUID replacement
|
||||
@@ -342,12 +342,12 @@ def _format_change_lines(
|
||||
lines = []
|
||||
# First line: path with green final element and grey =
|
||||
if len(formatted_path) == 1:
|
||||
lines.append(f" {_ADD}{formatted_path[0]}{_RESET} {_DIM}={_RESET}")
|
||||
lines.append(f" {_ADD}{formatted_path[0]}{_RESET} {_SEP}={_RESET}")
|
||||
else:
|
||||
prefix = ".".join(formatted_path[:-1])
|
||||
final = formatted_path[-1]
|
||||
lines.append(
|
||||
f" {_PATH_PREFIX}{prefix}.{_RESET}{_ADD}{final}{_RESET} {_DIM}={_RESET}"
|
||||
f" {_PATH_PREFIX}{prefix}.{_RESET}{_ADD}{final}{_RESET} {_SEP}={_RESET}"
|
||||
)
|
||||
# Child lines: indented key: value, with aligned values
|
||||
# Format keys (may contain UUIDs)
|
||||
@@ -360,24 +360,24 @@ def _format_change_lines(
|
||||
field_width = max(max_key_len, 12) # minimum 12 chars
|
||||
for k_display, v_str in formatted_items:
|
||||
padding = " " * (field_width - len(k_display))
|
||||
lines.append(f" {k_display}{_DIM}:{_RESET}{padding} {v_str}")
|
||||
lines.append(f" {k_display}{_SEP}:{_RESET}{padding} {v_str}")
|
||||
return lines
|
||||
else:
|
||||
value_str = fmt_value(value, path)
|
||||
if len(formatted_path) == 1:
|
||||
return [
|
||||
f" {_ADD}{formatted_path[0]}{_RESET} {_DIM}={_RESET} {value_str}"
|
||||
f" {_ADD}{formatted_path[0]}{_RESET} {_SEP}={_RESET} {value_str}"
|
||||
]
|
||||
prefix = ".".join(formatted_path[:-1])
|
||||
final = formatted_path[-1]
|
||||
return [
|
||||
f" {_PATH_PREFIX}{prefix}.{_RESET}{_ADD}{final}{_RESET} {_DIM}={_RESET} {value_str}"
|
||||
f" {_PATH_PREFIX}{prefix}.{_RESET}{_ADD}{final}{_RESET} {_SEP}={_RESET} {value_str}"
|
||||
]
|
||||
|
||||
# update: Existing item being updated - normal path colors
|
||||
value_str = fmt_value(value, path)
|
||||
path_str = _format_path(path, resolver=resolver)
|
||||
return [f" {path_str} {_DIM}={_RESET} {value_str}"]
|
||||
return [f" {path_str} {_SEP}={_RESET} {value_str}"]
|
||||
|
||||
|
||||
def format_diff(
|
||||
|
||||
+30
-8
@@ -8,28 +8,36 @@ Each migration should be idempotent and only run when needed.
|
||||
import base64
|
||||
from collections.abc import Awaitable, Callable
|
||||
|
||||
import msgspec
|
||||
|
||||
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."""
|
||||
for org_data in d["orgs"].values():
|
||||
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."""
|
||||
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."""
|
||||
for user_data in d["users"].values():
|
||||
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."""
|
||||
# Session keys changed to hashes, drop old sessions
|
||||
d["sessions"] = {}
|
||||
@@ -45,14 +53,28 @@ migrations = sorted(
|
||||
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(
|
||||
data_dict: dict,
|
||||
current_version: int,
|
||||
persist: Callable[[str, int, dict], Awaitable[None]],
|
||||
*,
|
||||
rp_id: str = "localhost",
|
||||
ctx: MigrationCtx,
|
||||
) -> None:
|
||||
while current_version < DBVER:
|
||||
migrations[current_version](data_dict, rp_id=rp_id)
|
||||
migrations[current_version](data_dict, ctx)
|
||||
current_version += 1
|
||||
await persist(f"migrate:v{current_version}", current_version, data_dict)
|
||||
|
||||
+3
-11
@@ -11,13 +11,10 @@ import secrets
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import UUID
|
||||
|
||||
import base64url
|
||||
import uuid7
|
||||
|
||||
from paskia import oidc_notify
|
||||
from paskia.config import SESSION_LIFETIME
|
||||
from paskia.db.jsonl import (
|
||||
JsonlStore,
|
||||
)
|
||||
from paskia.db.structs import (
|
||||
DB,
|
||||
Client,
|
||||
@@ -41,9 +38,6 @@ _UNSET = object()
|
||||
|
||||
# Global database instance (empty until init() loads data)
|
||||
_db = DB(config=Config(rp_id="uninitialized.invalid"))
|
||||
_store = JsonlStore(_db)
|
||||
_db._store = _store
|
||||
_initialized = False
|
||||
|
||||
|
||||
def is_username_taken(username: str, exclude_uuid: UUID | None = None) -> bool:
|
||||
@@ -484,7 +478,6 @@ def delete_session(
|
||||
"""
|
||||
if key not in _db.sessions:
|
||||
raise ValueError("Session not found")
|
||||
from paskia import oidc_notify # noqa: PLC0415
|
||||
|
||||
oidc_notify.schedule_notifications([key])
|
||||
with _db.transaction(action, ctx):
|
||||
@@ -503,7 +496,6 @@ def delete_sessions_for_user(
|
||||
user = _db.users.get(user_uuid)
|
||||
if not user:
|
||||
return
|
||||
from paskia import oidc_notify # noqa: PLC0415
|
||||
|
||||
keys = [s.key for s in user.sessions]
|
||||
oidc_notify.schedule_notifications(keys)
|
||||
@@ -589,7 +581,7 @@ def login(
|
||||
session = Session.create(
|
||||
user=user_uuid,
|
||||
credential=credential_uuid,
|
||||
key=base64url.enc(hash_secret("cookie", token)),
|
||||
key=hash_secret("cookie", token),
|
||||
host=host,
|
||||
ip=ip,
|
||||
user_agent=user_agent,
|
||||
@@ -657,7 +649,7 @@ def create_credential_session(
|
||||
|
||||
# Generate token and derive key
|
||||
token = secrets.token_urlsafe(12)
|
||||
key = base64url.enc(hash_secret("cookie", token))
|
||||
key = hash_secret("cookie", token)
|
||||
|
||||
session = Session.create(
|
||||
user=user_uuid,
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
"""
|
||||
Snapshot handling for JSONL database persistence.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
import msgspec
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
LINEPREFIX = b"SNAPSHOT "
|
||||
MINDIFFS = 100
|
||||
|
||||
|
||||
class Snapshot(msgspec.Struct):
|
||||
"""Snapshot data structure for database persistence."""
|
||||
|
||||
ts: datetime
|
||||
v: int
|
||||
state: dict[str, Any]
|
||||
|
||||
|
||||
class SnapshotState:
|
||||
"""Tracks snapshot timing and line counts for a database file."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.ts: datetime | None = None
|
||||
self.changes: int = 0
|
||||
self._force_pending: bool = False
|
||||
|
||||
def request_force(self) -> None:
|
||||
"""Request a forced snapshot on the next maybe_write call."""
|
||||
self._force_pending = True
|
||||
|
||||
def record_lines(self, count: int) -> None:
|
||||
self.changes += count
|
||||
|
||||
def maybe_write(self, file, version: int, state: dict) -> None:
|
||||
"""Write a snapshot if conditions are met (enough changes, and Sunday UTC or forced)."""
|
||||
if self.changes < MINDIFFS:
|
||||
return
|
||||
force = self._force_pending
|
||||
now = datetime.now(UTC)
|
||||
if not force and now.weekday() != 6: # 6 = Sunday
|
||||
return
|
||||
sunday_midnight = now.replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
if not force and self.ts is not None and self.ts >= sunday_midnight:
|
||||
return
|
||||
if not file.is_open:
|
||||
return
|
||||
try:
|
||||
self._write(file, version, state, now)
|
||||
self._force_pending = False
|
||||
except Exception as exc:
|
||||
_logger.error("snapshot: failed to write snapshot: %r", exc)
|
||||
|
||||
def _write(self, file, version: int, state: dict, now: datetime) -> None:
|
||||
"""Write a snapshot and update internal state."""
|
||||
data = msgspec.json.encode(Snapshot(ts=now, v=version, state=state))
|
||||
file.write(LINEPREFIX + data + b"\n")
|
||||
self.changes = 0
|
||||
self.ts = now
|
||||
|
||||
@staticmethod
|
||||
def load(data: bytes) -> tuple[Snapshot | None, int]:
|
||||
"""Find and parse the last snapshot in file data.
|
||||
|
||||
Returns (snapshot, replay_offset) where replay_offset is the byte
|
||||
position to start replaying change records from. If no valid snapshot
|
||||
is found, returns (None, 0).
|
||||
"""
|
||||
marker = b"\n" + LINEPREFIX
|
||||
pos = data.rfind(marker)
|
||||
if pos != -1:
|
||||
pos += 1 # skip the newline
|
||||
elif data.startswith(LINEPREFIX):
|
||||
pos = 0
|
||||
else:
|
||||
return None, 0
|
||||
|
||||
end = data.find(b"\n", pos)
|
||||
if end == -1:
|
||||
raise ValueError("Incomplete snapshot line at end of file")
|
||||
|
||||
snap = msgspec.json.decode(data[pos + len(LINEPREFIX) : end], type=Snapshot)
|
||||
return snap, end + 1
|
||||
+11
-15
@@ -5,12 +5,10 @@ import secrets
|
||||
from datetime import UTC, datetime
|
||||
from uuid import UUID
|
||||
|
||||
import base64url
|
||||
import msgspec
|
||||
import uuid7
|
||||
|
||||
from paskia import db
|
||||
from paskia.util import hostutil
|
||||
from paskia.util import passphrase as passphrase_util
|
||||
from paskia.util.crypto import hash_secret
|
||||
|
||||
@@ -434,7 +432,7 @@ class Session(msgspec.Struct, dict=True, omit_defaults=True):
|
||||
"""Create a new Session with the provided key.
|
||||
|
||||
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:
|
||||
Session object with key set
|
||||
@@ -471,7 +469,7 @@ class ResetToken(msgspec.Struct, dict=True):
|
||||
|
||||
def __post_init__(self):
|
||||
if not hasattr(self, "key"):
|
||||
self.key: bytes = b""
|
||||
self.key: str = ""
|
||||
|
||||
@property
|
||||
def user(self) -> User:
|
||||
@@ -487,15 +485,15 @@ class ResetToken(msgspec.Struct, dict=True):
|
||||
del db.data().reset_tokens[self.key]
|
||||
|
||||
@staticmethod
|
||||
def hash(passphrase: str) -> bytes:
|
||||
"""Hash a passphrase to bytes for reset token storage."""
|
||||
def hash(passphrase: str) -> str:
|
||||
"""Hash a passphrase to string for reset token storage."""
|
||||
if not passphrase_util.is_well_formed(passphrase):
|
||||
raise ValueError(
|
||||
"Trying to reset with a session token in place of a passphrase"
|
||||
if len(passphrase) == 16
|
||||
else "Invalid passphrase format"
|
||||
)
|
||||
return hashlib.sha512(passphrase.encode()).digest()[:9]
|
||||
return hash_secret("reset", passphrase)
|
||||
|
||||
@classmethod
|
||||
def by_passphrase(cls, passphrase: str) -> ResetToken | None:
|
||||
@@ -602,14 +600,14 @@ class OIDC(msgspec.Struct, dict=True):
|
||||
key: bytes | None = None
|
||||
|
||||
|
||||
class Config(msgspec.Struct, frozen=True, dict=True, omit_defaults=True):
|
||||
class Config(msgspec.Struct, omit_defaults=True):
|
||||
"""Stored configuration for the instance."""
|
||||
|
||||
rp_id: str
|
||||
rp_name: str | None = None
|
||||
origins: list[str] | None = None
|
||||
auth_host: str | None = None
|
||||
listen: str | None = None
|
||||
origins: list[str] | None = None
|
||||
listen: list[str] | None = None
|
||||
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
@@ -627,7 +625,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
users: dict[UUID, User] = {}
|
||||
credentials: dict[UUID, Credential] = {}
|
||||
sessions: dict[str, Session] = {}
|
||||
reset_tokens: dict[bytes, ResetToken] = {}
|
||||
reset_tokens: dict[str, ResetToken] = {}
|
||||
# OIDC provider data
|
||||
oidc: OIDC = msgspec.field(default_factory=lambda: OIDC())
|
||||
|
||||
@@ -670,7 +668,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
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:
|
||||
s = self.sessions[key]
|
||||
except KeyError:
|
||||
@@ -680,10 +678,8 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||
if s.client_uuid is not None:
|
||||
return None
|
||||
|
||||
# Normalize host for comparison (stored hosts are already normalized)
|
||||
normalized_input = hostutil.normalize_host(host)
|
||||
|
||||
# Validate host matches (sessions are always created with a host)
|
||||
normalized_input = host
|
||||
if s.host != normalized_input:
|
||||
# Session bound to different host
|
||||
return None
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import logging
|
||||
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 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.structs import Client
|
||||
from paskia.fastapi import authz
|
||||
from paskia.fastapi.front import frontend
|
||||
from paskia.fastapi.response import MsgspecResponse
|
||||
from paskia.fastapi.session import AUTH_COOKIE
|
||||
from paskia.globals import passkey
|
||||
@@ -78,7 +79,7 @@ async def general_exception_handler(_request, exc: Exception): # pragma: no cov
|
||||
|
||||
@app.get("/")
|
||||
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 --------------------
|
||||
|
||||
+20
-24
@@ -15,12 +15,12 @@ from fastapi.security import HTTPBearer
|
||||
|
||||
from paskia import authcode, db
|
||||
from paskia._version import __version__
|
||||
from paskia.authsession import EXPIRES, get_reset
|
||||
from paskia.authsession import EXPIRES, get_reset, session_ctx
|
||||
from paskia.fastapi import authz, session, user
|
||||
from paskia.fastapi.response import MsgspecResponse
|
||||
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME, get_client_ip
|
||||
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
|
||||
|
||||
bearer_auth = HTTPBearer(auto_error=False)
|
||||
@@ -161,25 +161,15 @@ async def forward_authentication(
|
||||
# Clear cookie only if session is invalid (not for reauth)
|
||||
if e.clear_session:
|
||||
session.clear_session_cookie(response)
|
||||
|
||||
# Check Accept header to decide response format
|
||||
accept = request.headers.get("accept", "")
|
||||
wants_html = "text/html" in accept
|
||||
|
||||
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),
|
||||
# Browser request? - return full-page HTML with metadata patched into data attrs
|
||||
if "text/html" in request.headers.get("accept", ""):
|
||||
return await htmlutil.patched_html_response(
|
||||
request, "/int/forward/", e.status_code, mode=e.mode, **e.metadata
|
||||
)
|
||||
# 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")
|
||||
@@ -212,9 +202,14 @@ async def api_user_info(
|
||||
detail="Authentication required",
|
||||
mode="login",
|
||||
)
|
||||
ctx = db.data().session_ctx(auth, request.headers.get("host"))
|
||||
ctx = session_ctx(auth, request.headers.get("host"))
|
||||
if not ctx:
|
||||
raise HTTPException(401, "Session expired")
|
||||
raise authz.AuthException(
|
||||
status_code=401,
|
||||
detail="Session expired",
|
||||
mode="login",
|
||||
clear_session=True,
|
||||
)
|
||||
|
||||
return MsgspecResponse(
|
||||
await userinfo.build_user_info(
|
||||
@@ -244,6 +239,7 @@ async def token_info(credentials=Depends(bearer_auth)):
|
||||
ApiTokenInfo(
|
||||
token_type=reset_token.token_type,
|
||||
display_name=u.display_name,
|
||||
theme=u.theme,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -253,7 +249,7 @@ async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
||||
if not auth:
|
||||
return {"message": "Already logged out"}
|
||||
host = request.headers.get("host")
|
||||
ctx = db.data().session_ctx(auth, host)
|
||||
ctx = session_ctx(auth, host)
|
||||
if not ctx:
|
||||
return {"message": "Already logged out"}
|
||||
with suppress(Exception):
|
||||
@@ -286,7 +282,7 @@ async def api_set_session(
|
||||
secret = a.session_key
|
||||
|
||||
# Verify the session exists
|
||||
ctx = db.data().session_ctx(secret, host)
|
||||
ctx = session_ctx(secret, host)
|
||||
if not ctx:
|
||||
raise HTTPException(401, f"Session not found on {host}")
|
||||
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
+14
-39
@@ -115,25 +115,16 @@ def format_access_log(
|
||||
client: str, status: int, method: str, host: str, path: str, duration_ms: float
|
||||
) -> str:
|
||||
"""Format access log line with colors and aligned fields."""
|
||||
use_color = sys.stderr.isatty()
|
||||
|
||||
# Format components with fixed widths for alignment
|
||||
ip = format_client_ip(client).ljust(19) # IPv6 network max 19 chars
|
||||
timing = f"{duration_ms:.0f}ms"
|
||||
method_padded = method.ljust(7) # Longest method is OPTIONS (7)
|
||||
|
||||
if use_color:
|
||||
status_str = f"{status_color(status)}{status}{_RESET}"
|
||||
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||
method_str = f"{method_color(method)}{method_padded}{_RESET}"
|
||||
host_str = f"{_HOST}{host}{_RESET}"
|
||||
path_str = f"{_PATH}{path}{_RESET}"
|
||||
else:
|
||||
status_str = str(status)
|
||||
timing_str = timing
|
||||
method_str = method_padded
|
||||
host_str = host
|
||||
path_str = path
|
||||
status_str = f"{status_color(status)}{status}{_RESET}"
|
||||
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||
method_str = f"{method_color(method)}{method_padded}{_RESET}"
|
||||
host_str = f"{_HOST}{host}{_RESET}"
|
||||
path_str = f"{_PATH}{path}{_RESET}"
|
||||
|
||||
# Format: "IP STATUS METHOD host path TIMING"
|
||||
return f"{ip} {status_str} {method_str} {host_str}{path_str} {timing_str}"
|
||||
@@ -153,7 +144,6 @@ def _next_ws_id() -> int:
|
||||
|
||||
def log_ws_open(ws) -> int:
|
||||
"""Log WebSocket connection open. Returns connection ID for use in close."""
|
||||
use_color = sys.stderr.isatty()
|
||||
ws_id = _next_ws_id()
|
||||
|
||||
client = ws.client.host if ws.client else "-"
|
||||
@@ -169,19 +159,11 @@ def log_ws_open(ws) -> int:
|
||||
origin_host = origin.split("://", 1)[-1] if origin else None
|
||||
show_origin = origin_host and origin_host != host
|
||||
|
||||
if use_color:
|
||||
# 🔌 aligned with status (takes ~2 char width), ID aligned with method
|
||||
prefix = f"🔌 {_WS_OPEN}{id_str}{_RESET}"
|
||||
host_str = f"{_HOST}{host}{_RESET}"
|
||||
path_str = f"{_PATH}{path}{_RESET}"
|
||||
origin_str = (
|
||||
f" {_RESET}from {_HOST}{origin_host}{_RESET}" if show_origin else ""
|
||||
)
|
||||
else:
|
||||
prefix = f"WS+ {id_str}"
|
||||
host_str = host
|
||||
path_str = path
|
||||
origin_str = f" from {origin_host}" if show_origin else ""
|
||||
# 🔌 aligned with status (takes ~2 char width), ID aligned with method
|
||||
prefix = f"🔌 {_WS_OPEN}{id_str}{_RESET}"
|
||||
host_str = f"{_HOST}{host}{_RESET}"
|
||||
path_str = f"{_PATH}{path}{_RESET}"
|
||||
origin_str = f" {_RESET}from {_HOST}{origin_host}{_RESET}" if show_origin else ""
|
||||
|
||||
logger.info(f"{ip} {prefix} {host_str}{path_str}{origin_str}")
|
||||
return ws_id
|
||||
@@ -209,8 +191,6 @@ WS_CLOSE_CODES = {
|
||||
|
||||
def log_ws_close(ws_id: int, close_code: int | None, duration: float) -> None:
|
||||
"""Log WebSocket connection close with duration and status."""
|
||||
use_color = sys.stderr.isatty()
|
||||
|
||||
id_str = f"{ws_id:02d}".ljust(7) # Align with method field (7 chars)
|
||||
timing = f"{duration * 1000:.0f}ms"
|
||||
|
||||
@@ -220,15 +200,10 @@ def log_ws_close(ws_id: int, close_code: int | None, duration: float) -> None:
|
||||
else:
|
||||
status = WS_CLOSE_CODES.get(close_code, f"code {close_code}")
|
||||
|
||||
if use_color:
|
||||
# 🔌 aligned with status, ID aligned with method
|
||||
prefix = f"🔌 {_WS_CLOSE}{id_str}{_RESET}"
|
||||
status_str = f"{_WS_STATUS}{status}{_RESET}"
|
||||
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||
else:
|
||||
prefix = f"WS- {id_str}"
|
||||
status_str = status
|
||||
timing_str = timing
|
||||
# 🔌 aligned with status, ID aligned with method
|
||||
prefix = f"🔌 {_WS_CLOSE}{id_str}{_RESET}"
|
||||
status_str = f"{_WS_STATUS}{status}{_RESET}"
|
||||
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||
|
||||
logger.info(f"{' ' * 19} {prefix} {status_str} {timing_str}")
|
||||
|
||||
|
||||
+23
-22
@@ -1,21 +1,26 @@
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from contextlib import asynccontextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import msgspec
|
||||
from fastapi import FastAPI, HTTPException, Request, Response
|
||||
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.bootstrap import bootstrap_if_needed
|
||||
from paskia.db import start_background, stop_background
|
||||
from paskia.db.background import flush
|
||||
from paskia.db.logging import configure_db_logging
|
||||
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.session import AUTH_COOKIE
|
||||
from paskia.util import hostutil, passphrase, vitedev
|
||||
from paskia.util.runtime import RuntimeConfig
|
||||
|
||||
# Configure custom logging
|
||||
configure_access_logging()
|
||||
@@ -23,14 +28,6 @@ configure_db_logging()
|
||||
|
||||
_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
|
||||
_EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples"
|
||||
|
||||
@@ -43,14 +40,13 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
||||
so that uvicorn reload / multiprocess workers inherit the settings.
|
||||
All keys are guaranteed to exist; values are already normalized by __main__.py.
|
||||
"""
|
||||
config = json.loads(os.environ["PASKIA_CONFIG"])
|
||||
runtime = msgspec.json.decode(os.environ["PASKIA_CONFIG"], type=RuntimeConfig)
|
||||
|
||||
try:
|
||||
# CLI (__main__) performs bootstrap once; here we skip to avoid duplicate work
|
||||
await globals.init(
|
||||
rp_id=config["rp_id"],
|
||||
rp_name=config["rp_name"],
|
||||
origins=config["origins"],
|
||||
rp_id=runtime.config.rp_id,
|
||||
rp_name=runtime.config.rp_name,
|
||||
origins=runtime.config.origins,
|
||||
bootstrap=False,
|
||||
)
|
||||
except ValueError as e:
|
||||
@@ -58,6 +54,12 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
||||
# Re-raise to fail fast
|
||||
raise
|
||||
|
||||
# Bootstrap and persist config now that the full DB is loaded
|
||||
await bootstrap_if_needed(config=runtime.config)
|
||||
if runtime.save:
|
||||
await db.update_config(runtime.config)
|
||||
await flush()
|
||||
|
||||
# Restore uvicorn info logging (suppressed during startup in dev mode)
|
||||
# Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages
|
||||
if app.debug:
|
||||
@@ -131,9 +133,9 @@ async def openid_configuration(request: Request):
|
||||
|
||||
@app.get("/auth/restricted/iframe")
|
||||
@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."""
|
||||
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
|
||||
@@ -148,7 +150,7 @@ async def frontapp(request: Request, response: Response, auth=AUTH_COOKIE):
|
||||
The frontend handles mode detection (host mode vs full profile) based on settings.
|
||||
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)
|
||||
@@ -180,14 +182,13 @@ async def examples_page():
|
||||
|
||||
|
||||
# Frontend static files - must be before /{token} catch-all routes
|
||||
# (actual routes registered during lifespan after frontend.load())
|
||||
frontend.route(app, "/")
|
||||
|
||||
|
||||
# Note: this catch-all handler must be the last route defined
|
||||
@app.get("/{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).
|
||||
|
||||
The frontend will validate the token via /auth/api/token-info.
|
||||
@@ -195,4 +196,4 @@ async def token_link(token: str):
|
||||
if not passphrase.is_well_formed(token):
|
||||
raise HTTPException(status_code=404)
|
||||
|
||||
return Response(*await vitedev.read("/int/reset/index.html"))
|
||||
return await vitedev.handle(request, frontend, "/int/reset/")
|
||||
|
||||
@@ -40,7 +40,7 @@ def _oidc_session_by_token(
|
||||
token: str, client_uuid: UUID | None = None
|
||||
) -> Session | None:
|
||||
"""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)
|
||||
if not s or s.client_uuid is None:
|
||||
return None
|
||||
@@ -259,8 +259,10 @@ async def _handle_refresh_token(
|
||||
The refresh_token is the session secret. On refresh:
|
||||
- Validates session exists and belongs to client
|
||||
- Extends session expiry (24h sliding window)
|
||||
- Records current IP and user_agent
|
||||
- Issues new access_token and id_token
|
||||
|
||||
Note: ip and user_agent are NOT updated because the refresh request
|
||||
comes from the OIDC client's backend, not the end user's browser.
|
||||
"""
|
||||
if not refresh_token_value:
|
||||
return JSONResponse(
|
||||
@@ -287,17 +289,13 @@ async def _handle_refresh_token(
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
# Refresh the session - extend expiry and record IP/user_agent
|
||||
# Refresh the session - extend expiry only
|
||||
# Don't update ip/user_agent: the refresh request comes from the OIDC
|
||||
# client's backend, not the end user's browser.
|
||||
now = datetime.now(UTC)
|
||||
ip = request.headers.get("x-forwarded-for", "").split(",")[0].strip()
|
||||
if not ip:
|
||||
ip = request.client.host if request.client else ""
|
||||
user_agent = request.headers.get("user-agent", "")
|
||||
|
||||
db.update_session(
|
||||
session.key,
|
||||
ip=ip,
|
||||
user_agent=user_agent,
|
||||
validated=now,
|
||||
)
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ from paskia import db
|
||||
from paskia.authsession import (
|
||||
delete_credential,
|
||||
expires,
|
||||
session_ctx,
|
||||
)
|
||||
from paskia.fastapi import authz, session
|
||||
from paskia.fastapi.response import MsgspecResponse
|
||||
@@ -45,7 +46,7 @@ async def user_update_display_name(
|
||||
status_code=401, detail="Authentication Required", mode="login"
|
||||
)
|
||||
host = request.headers.get("host")
|
||||
ctx = db.data().session_ctx(auth, host)
|
||||
ctx = session_ctx(auth, host)
|
||||
if not ctx:
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Session expired", mode="login"
|
||||
@@ -74,7 +75,7 @@ async def user_update_info(
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Authentication Required", mode="login"
|
||||
)
|
||||
ctx = db.data().session_ctx(auth, request.headers.get("host"))
|
||||
ctx = session_ctx(auth, request.headers.get("host"))
|
||||
if not ctx:
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Session expired", mode="login"
|
||||
@@ -112,7 +113,7 @@ async def user_update_theme(
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Authentication Required", mode="login"
|
||||
)
|
||||
ctx = db.data().session_ctx(auth, request.headers.get("host"))
|
||||
ctx = session_ctx(auth, request.headers.get("host"))
|
||||
if not ctx:
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Session expired", mode="login"
|
||||
@@ -129,7 +130,7 @@ async def api_logout_all(request: Request, response: Response, auth=AUTH_COOKIE)
|
||||
if not auth:
|
||||
return {"message": "Already logged out"}
|
||||
host = request.headers.get("host")
|
||||
ctx = db.data().session_ctx(auth, host)
|
||||
ctx = session_ctx(auth, host)
|
||||
if not ctx:
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Session expired", mode="login"
|
||||
@@ -151,7 +152,7 @@ async def api_delete_session(
|
||||
status_code=401, detail="Authentication Required", mode="login"
|
||||
)
|
||||
host = request.headers.get("host")
|
||||
ctx = db.data().session_ctx(auth, host)
|
||||
ctx = session_ctx(auth, host)
|
||||
if not ctx:
|
||||
raise authz.AuthException(
|
||||
status_code=401, detail="Session expired", mode="login"
|
||||
|
||||
@@ -3,12 +3,11 @@ from datetime import UTC, datetime
|
||||
from urllib.parse import urlencode
|
||||
from uuid import UUID
|
||||
|
||||
import base64url
|
||||
from fastapi import FastAPI, WebSocket
|
||||
|
||||
from paskia import authcode, db
|
||||
from paskia.authcode import CookieCode, OIDCCode
|
||||
from paskia.authsession import get_reset
|
||||
from paskia.authsession import get_reset, session_ctx
|
||||
from paskia.db.structs import Session
|
||||
from paskia.fastapi import authz, remote
|
||||
from paskia.fastapi.session import AUTH_COOKIE, infodict
|
||||
@@ -196,7 +195,7 @@ async def websocket_authenticate(
|
||||
# If there's an existing session, restrict to that user's credentials (reauth)
|
||||
session_user_uuid = None
|
||||
if auth:
|
||||
existing_ctx = db.data().session_ctx(auth, host)
|
||||
existing_ctx = session_ctx(auth, host)
|
||||
if existing_ctx:
|
||||
session_user_uuid = existing_ctx.user.uuid
|
||||
|
||||
@@ -218,7 +217,7 @@ async def websocket_authenticate(
|
||||
session = Session.create(
|
||||
user=cred.user_uuid,
|
||||
credential=cred.uuid,
|
||||
key=base64url.enc(hash_secret("oidc", token)),
|
||||
key=hash_secret("oidc", token),
|
||||
host=normalized_host,
|
||||
ip=metadata["ip"],
|
||||
user_agent=metadata["user_agent"],
|
||||
|
||||
@@ -7,6 +7,7 @@ from uuid import UUID
|
||||
from fastapi import WebSocket
|
||||
|
||||
from paskia import db
|
||||
from paskia.authsession import session_ctx
|
||||
from paskia.db import Credential, SessionContext
|
||||
from paskia.fastapi.session import infodict
|
||||
from paskia.fastapi.wsutil import validate_origin
|
||||
@@ -90,7 +91,7 @@ async def authenticate_and_login(
|
||||
# Get credential IDs if restricting to a user's credentials
|
||||
credential_ids = None
|
||||
if auth:
|
||||
existing_ctx = db.data().session_ctx(auth, host)
|
||||
existing_ctx = session_ctx(auth, host)
|
||||
if existing_ctx:
|
||||
credential_ids = existing_ctx.user.credential_ids or None
|
||||
|
||||
@@ -107,7 +108,7 @@ async def authenticate_and_login(
|
||||
)
|
||||
|
||||
# Fetch and return the full session context
|
||||
ctx = db.data().session_ctx(secret, normalized_host)
|
||||
ctx = session_ctx(secret, host)
|
||||
if not ctx:
|
||||
raise ValueError("Failed to create session context")
|
||||
return ctx, secret
|
||||
|
||||
@@ -13,7 +13,7 @@ import httpx
|
||||
|
||||
from paskia import db
|
||||
from paskia.util import oidjwt
|
||||
from paskia.util.hostutil import _load_config
|
||||
from paskia.util.runtime import _load_config
|
||||
|
||||
_logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -171,11 +171,12 @@ class ApiSettings(msgspec.Struct):
|
||||
version: str
|
||||
|
||||
|
||||
class ApiTokenInfo(msgspec.Struct):
|
||||
class ApiTokenInfo(msgspec.Struct, omit_defaults=True):
|
||||
"""Token info response struct."""
|
||||
|
||||
token_type: str
|
||||
display_name: str
|
||||
theme: str = ""
|
||||
|
||||
|
||||
class ApiUuidResponse(msgspec.Struct):
|
||||
|
||||
@@ -1,17 +1,15 @@
|
||||
import hashlib
|
||||
|
||||
import base64url
|
||||
from cryptography.hazmat.primitives import serialization
|
||||
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
|
||||
|
||||
|
||||
def hash_secret(*data) -> bytes:
|
||||
"""A custom HMAC that securily combines and hashes the given data (context, secrets). The first argument should be a namespacing string."""
|
||||
inner = bytearray(len(data).to_bytes(8, "big"))
|
||||
for d in data:
|
||||
if isinstance(d, str):
|
||||
d = d.encode()
|
||||
inner += hashlib.sha256(d).digest()
|
||||
return hashlib.sha256(inner).digest()[:12]
|
||||
def hash_secret(*data: str | bytes, length=12) -> str:
|
||||
"""A custom HMAC that securily combines and hashes the given data. The first argument should be a namespacing string."""
|
||||
p = [d.encode() if hasattr(d, "encode") else d for d in data]
|
||||
p += [len(x).to_bytes(8, "little") for x in [p, *p]]
|
||||
return base64url.enc(hashlib.sha256(b"".join(p)).digest()[:length])
|
||||
|
||||
|
||||
def secret_key() -> bytes:
|
||||
|
||||
@@ -11,7 +11,7 @@ __all__ = ["path", "file", "read", "is_dev_mode"]
|
||||
|
||||
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
|
||||
return os.environ.get("PASKIA_VITE_URL") or None
|
||||
|
||||
|
||||
def _resolve_static_dir() -> Path:
|
||||
|
||||
+26
-16
@@ -1,27 +1,23 @@
|
||||
"""Utilities for determining the auth UI host and base URLs."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from urllib.parse import urlparse, urlsplit
|
||||
|
||||
from paskia.util.runtime import _load_config
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _load_config() -> dict:
|
||||
"""Load PASKIA_CONFIG JSON."""
|
||||
config_json = os.getenv("PASKIA_CONFIG")
|
||||
if not config_json:
|
||||
return {}
|
||||
return json.loads(config_json)
|
||||
|
||||
def _cfg():
|
||||
return _load_config()
|
||||
|
||||
|
||||
def is_root_mode() -> bool:
|
||||
return _load_config().get("auth_host") is not None
|
||||
cfg = _cfg()
|
||||
return cfg is not None and cfg.config.auth_host is not None
|
||||
|
||||
|
||||
def dedicated_auth_host() -> str | None:
|
||||
"""Return configured auth_host netloc, or None."""
|
||||
auth_host = _load_config().get("auth_host")
|
||||
cfg = _cfg()
|
||||
auth_host = cfg.config.auth_host if cfg else None
|
||||
if not auth_host:
|
||||
return None
|
||||
|
||||
@@ -35,8 +31,10 @@ def ui_base_path() -> str:
|
||||
|
||||
def auth_site_url() -> str:
|
||||
"""Return the base URL for the auth site UI (computed at startup)."""
|
||||
cfg = _load_config()
|
||||
return cfg.get("site_url", "https://localhost") + cfg.get("site_path", "/auth/")
|
||||
cfg = _cfg()
|
||||
if cfg:
|
||||
return cfg.site_url + cfg.site_path
|
||||
return "https://localhost/auth/"
|
||||
|
||||
|
||||
def reset_link_url(token: str) -> str:
|
||||
@@ -45,10 +43,10 @@ def reset_link_url(token: str) -> str:
|
||||
|
||||
|
||||
def normalize_origin(origin: str) -> str:
|
||||
"""Normalize an origin URL by adding https:// if no scheme is present."""
|
||||
"""Normalize an origin URL by adding https:// if no scheme is present, removing trailing slashes."""
|
||||
if "://" not in origin:
|
||||
return f"https://{origin}"
|
||||
return origin
|
||||
return origin.rstrip("/")
|
||||
|
||||
|
||||
def reload_config() -> None:
|
||||
@@ -74,3 +72,15 @@ def normalize_host(raw_host: str | None) -> str | None:
|
||||
# Strip port from host:port
|
||||
netloc = netloc.rsplit(":", 1)[0]
|
||||
return netloc.lower() or None
|
||||
|
||||
|
||||
def format_endpoint(ep: dict) -> str:
|
||||
"""Format an endpoint dict to a listen string (e.g. 'unix:/path' or 'host:port')."""
|
||||
if uds := ep.get("uds"):
|
||||
return f"unix:{uds}"
|
||||
host = ep["host"]
|
||||
port = ep["port"]
|
||||
# Bracket IPv6 addresses
|
||||
if ":" in host:
|
||||
host = f"[{host}]"
|
||||
return f"{host}:{port}"
|
||||
|
||||
@@ -2,6 +2,39 @@
|
||||
|
||||
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:
|
||||
"""Patch HTML by adding data attributes to the <html> tag.
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from collections.abc import Sequence
|
||||
from fnmatch import fnmatchcase
|
||||
|
||||
from paskia import db
|
||||
from paskia.authsession import session_ctx
|
||||
from paskia.util.hostutil import normalize_host
|
||||
|
||||
__all__ = ["has_any", "has_all", "session_context"]
|
||||
@@ -40,4 +40,4 @@ async def session_context(auth: str | None, host: str | None = None):
|
||||
if not auth:
|
||||
return None
|
||||
normalized_host = normalize_host(host) if host else None
|
||||
return db.data().session_ctx(auth, normalized_host)
|
||||
return session_ctx(auth, normalized_host)
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
"""Runtime configuration utilities."""
|
||||
|
||||
import os
|
||||
from functools import lru_cache
|
||||
|
||||
import msgspec
|
||||
|
||||
from paskia.db.structs import Config
|
||||
|
||||
|
||||
class RuntimeConfig(msgspec.Struct):
|
||||
"""Runtime configuration for the Paskia authentication server.
|
||||
|
||||
Wraps the db Config (CLI/stored settings) with computed runtime fields.
|
||||
Serialized to PASKIA_CONFIG env var as JSON via msgspec.
|
||||
"""
|
||||
|
||||
config: Config # CLI/stored configuration to persist
|
||||
site_url: str # Base URL without trailing path (e.g. https://example.com)
|
||||
site_path: str # Path to auth UI: "/" if auth_host, else "/auth/"
|
||||
save: bool = False # Whether to persist config to database
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _load_config() -> "RuntimeConfig | None":
|
||||
"""Load RuntimeConfig from PASKIA_CONFIG env var."""
|
||||
config_json = os.getenv("PASKIA_CONFIG")
|
||||
if not config_json:
|
||||
return None
|
||||
|
||||
return msgspec.json.decode(config_json.encode(), type=RuntimeConfig)
|
||||
+28
-24
@@ -1,14 +1,19 @@
|
||||
"""Startup configuration box formatting utilities."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from sys import stderr
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from fastapi_vue.hostutil import parse_endpoints
|
||||
|
||||
from paskia._version import __version__
|
||||
from paskia.util.hostutil import format_endpoint
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from paskia.config import PaskiaConfig
|
||||
from paskia.util.runtime import RuntimeConfig
|
||||
|
||||
BOX_WIDTH = 60 # Inner width (excluding box chars)
|
||||
|
||||
@@ -42,7 +47,7 @@ def bottom() -> str:
|
||||
return "┗" + "━" * (BOX_WIDTH + 2) + "┛\n"
|
||||
|
||||
|
||||
def print_startup_config(config: "PaskiaConfig") -> None:
|
||||
def print_startup_config(runtime: RuntimeConfig) -> None:
|
||||
"""Print server configuration on startup."""
|
||||
# Key graphic with yellow shading (bright for highlights, dark for body)
|
||||
y = YELLOW # Dark yellow for main body
|
||||
@@ -57,41 +62,40 @@ def print_startup_config(config: "PaskiaConfig") -> None:
|
||||
lines.append(
|
||||
line(
|
||||
f"{b}█{y} {b}█{y}▀▀▀▀{b}█{y}▀▀{b}█{y}▀▀{b}█{r} {w}"
|
||||
+ config.site_url
|
||||
+ config.site_path
|
||||
+ runtime.site_url
|
||||
+ runtime.site_path
|
||||
+ r
|
||||
)
|
||||
)
|
||||
lines.append(line(f" {y}▀▀▀▀▀{r}"))
|
||||
|
||||
# Format auth host section
|
||||
if config.auth_host:
|
||||
lines.append(line(f"Auth Host: {config.auth_host}"))
|
||||
if runtime.config.auth_host:
|
||||
lines.append(line(f"Auth Host: {runtime.config.auth_host}"))
|
||||
|
||||
from paskia.__main__ import DEFAULT_PORT as P # noqa: PLC0415 - circular
|
||||
from paskia.__main__ import DEVMODE # noqa: PLC0415 - circular
|
||||
|
||||
# Show frontend URL if in dev mode
|
||||
devmode = os.environ.get("FASTAPI_VUE_FRONTEND_URL")
|
||||
if devmode:
|
||||
lines.append(line(f"Dev Frontend: {devmode}"))
|
||||
if DEVMODE:
|
||||
lines.append(line(f"Dev Frontend: {os.environ.get('PASKIA_VITE_URL')}"))
|
||||
|
||||
# Format listen address with scheme
|
||||
if config.uds:
|
||||
listen = f"unix:{config.uds}"
|
||||
elif config.host:
|
||||
listen = f"http://{config.host}:{config.port}"
|
||||
else:
|
||||
listen = f"http://0.0.0.0:{config.port} + [::]:{config.port}"
|
||||
lines.append(line(f"Backend: {listen}"))
|
||||
# Format listen endpoints (dev mode only uses the first endpoint)
|
||||
|
||||
endpoints = list(parse_endpoints(runtime.config.listen, P))
|
||||
if DEVMODE:
|
||||
endpoints = endpoints[:1] # server.run reload=True uses only one
|
||||
parts = [format_endpoint(ep) for ep in endpoints]
|
||||
lines.append(line(f"Backend: {' '.join(parts)}"))
|
||||
|
||||
# Relying Party line (omit name if same as id)
|
||||
rp_id = config.rp_id
|
||||
rp_name = config.rp_name
|
||||
if rp_name and rp_name != rp_id:
|
||||
lines.append(line(f"Relying Party: {rp_id} ({rp_name})"))
|
||||
else:
|
||||
lines.append(line(f"Relying Party: {rp_id}"))
|
||||
rp_id = runtime.config.rp_id
|
||||
rp_name = runtime.config.rp_name
|
||||
suffix = f" ({rp_name})" if rp_name and rp_name != rp_id else ""
|
||||
lines.append(line(f"Relying Party: {rp_id}{suffix}"))
|
||||
|
||||
# Format origins section
|
||||
allowed = config.origins
|
||||
allowed = runtime.config.origins
|
||||
if allowed:
|
||||
lines.append(line("Permitted Origins:"))
|
||||
for origin in sorted(allowed):
|
||||
|
||||
@@ -5,8 +5,10 @@ from paskia.db import SessionContext
|
||||
from paskia.util import hostutil
|
||||
from paskia.util.apistructs import (
|
||||
ApiAaguidInfo,
|
||||
ApiOrg,
|
||||
ApiOrgContext,
|
||||
ApiPermission,
|
||||
ApiRole,
|
||||
ApiRoleContext,
|
||||
ApiSessionContext,
|
||||
ApiUser,
|
||||
@@ -64,4 +66,6 @@ async def build_user_info(
|
||||
permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions}
|
||||
if ctx
|
||||
else {},
|
||||
org=ApiOrg.from_db(ctx.org) if ctx else None,
|
||||
role=ApiRole.from_db(ctx.role) if ctx else None,
|
||||
)
|
||||
|
||||
+35
-31
@@ -1,39 +1,29 @@
|
||||
"""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.
|
||||
|
||||
This complements fastapi_vue.Frontend which handles static file serving
|
||||
but doesn't provide server-side fetching of HTML content.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import mimetypes
|
||||
import os
|
||||
from importlib import resources
|
||||
from pathlib import Path
|
||||
|
||||
import httpx
|
||||
from fastapi import Response
|
||||
|
||||
__all__ = ["read"]
|
||||
|
||||
|
||||
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
|
||||
__all__ = ["handle"]
|
||||
|
||||
|
||||
def _resolve_static_dir() -> Path:
|
||||
"""Resolve the static files directory."""
|
||||
|
||||
# Try packaged path via importlib.resources (works for wheel/installed).
|
||||
try: # pragma: no cover - trivial path resolution
|
||||
pkg_dir = resources.files("paskia") / "frontend-build"
|
||||
fs_path = Path(str(pkg_dir))
|
||||
if fs_path.is_dir():
|
||||
return fs_path
|
||||
except Exception: # pragma: no cover - defensive
|
||||
pass
|
||||
pkg_dir = resources.files("paskia") / "frontend-build"
|
||||
fs_path = Path(str(pkg_dir))
|
||||
if fs_path.is_dir():
|
||||
return fs_path
|
||||
# Fallback for editable/development before build.
|
||||
return Path(__file__).parent.parent / "frontend-build"
|
||||
|
||||
@@ -41,31 +31,45 @@ def _resolve_static_dir() -> Path:
|
||||
_static_dir: Path = _resolve_static_dir()
|
||||
|
||||
|
||||
async def read(filepath: str) -> tuple[bytes, int, dict[str, str]]:
|
||||
"""Read file content and return response tuple.
|
||||
async def handle(request, frontend, filepath: str):
|
||||
"""Read file content and return Response.
|
||||
|
||||
In dev mode, fetches from the Vite dev server.
|
||||
In production, reads from the static build directory.
|
||||
In production, uses frontend.handle.
|
||||
|
||||
Args:
|
||||
request: The FastAPI Request object
|
||||
frontend: The fastapi_vue.Frontend instance
|
||||
filepath: Path relative to frontend root, e.g. "/auth/index.html"
|
||||
|
||||
Returns:
|
||||
Tuple of (content, status_code, headers) suitable for
|
||||
FastAPI Response(*args).
|
||||
FastAPI Response object.
|
||||
"""
|
||||
dev_server = _get_dev_server()
|
||||
if dev_server:
|
||||
if dev_server := os.environ.get("PASKIA_VITE_URL"):
|
||||
async with httpx.AsyncClient() as client:
|
||||
resp = await client.get(f"{dev_server}{filepath}")
|
||||
resp.raise_for_status()
|
||||
mime = resp.headers.get("content-type", "application/octet-stream")
|
||||
# Strip charset suffix if present
|
||||
mime = mime.split(";")[0].strip()
|
||||
return resp.content, resp.status_code, {"content-type": mime}
|
||||
else:
|
||||
# Production: read from static build
|
||||
file_path = _static_dir / filepath.lstrip("/")
|
||||
content = await asyncio.to_thread(file_path.read_bytes)
|
||||
mime, _ = mimetypes.guess_type(str(file_path))
|
||||
return content, 200, {"content-type": mime or "application/octet-stream"}
|
||||
return Response(resp.content, resp.status_code, {"content-type": mime})
|
||||
|
||||
# Read from frontend cache directly to bypass any compression/processing
|
||||
cached_content = getattr(frontend, "_files", {}).get(filepath)
|
||||
if cached_content is not None:
|
||||
mime, _ = mimetypes.guess_type(filepath)
|
||||
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
@@ -11,20 +11,28 @@ keywords = [ "forward_auth", "auth_request", "FastAPI" ]
|
||||
authors = [
|
||||
{name = "Leo Vasanko"},
|
||||
]
|
||||
requires-python = ">=3.11"
|
||||
dependencies = [
|
||||
"fastapi[standard]>=0.104.1",
|
||||
"websockets>=12.0",
|
||||
"webauthn>=1.11.1",
|
||||
"base64url>=1.0.0",
|
||||
"uuid7-standard>=1.0.0",
|
||||
"pyjwt[crypto]>=2.8.0",
|
||||
"fastapi[standard]>=0.129.0",
|
||||
"websockets>=16.0",
|
||||
"webauthn>=2.7.1",
|
||||
"base64url>=1.1.1",
|
||||
"uuid7-standard>=1.1.0",
|
||||
"pyjwt[crypto]>=2.11.0",
|
||||
"jsondiff>=2.2.1",
|
||||
"msgspec>=0.20.0",
|
||||
"aiofiles>=25.1.0",
|
||||
"fastapi-vue>=0.3.0",
|
||||
"fastapi-vue>=1.1.0",
|
||||
"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]
|
||||
Homepage = "https://git.zi.fi/LeoVasanko/paskia"
|
||||
@@ -36,15 +44,6 @@ source = "vcs"
|
||||
[tool.hatch.build.hooks.vcs]
|
||||
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]
|
||||
source = ["paskia"]
|
||||
branch = true
|
||||
@@ -75,16 +74,6 @@ select = ["E", "F", "I", "N", "W", "UP", "PLC0415"]
|
||||
ignore = ["E501"] # Line too long
|
||||
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]
|
||||
paskia = "paskia.__main__:main"
|
||||
|
||||
|
||||
+16
-20
@@ -153,29 +153,9 @@ async def run_devserver(args: argparse.Namespace, remaining: list[str]) -> None:
|
||||
paskia.extend(["--origin", origin])
|
||||
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
|
||||
os.environ["PASKIA_VITE_URL"] = viteurl
|
||||
os.environ["PASKIA_BACKEND_URL"] = backurl
|
||||
os.environ["PASKIA_SITE_URL"] = caddy_origins[0] if args.caddy else viteurl
|
||||
os.environ["PASKIA_DEV"] = "1"
|
||||
if 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:
|
||||
# Start Caddy first if requested (needs to bind ports)
|
||||
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)
|
||||
pg._procs.append(caddy_proc)
|
||||
pg._cmds[caddy_proc.pid] = "caddy"
|
||||
|
||||
+1
-2
@@ -19,7 +19,6 @@ from collections.abc import AsyncGenerator
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from uuid import UUID
|
||||
|
||||
import base64url
|
||||
import httpx
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
@@ -268,7 +267,7 @@ def create_test_session(
|
||||
|
||||
# Generate token and derive key
|
||||
token = secrets.token_urlsafe(12)
|
||||
key = base64url.enc(hash_secret("cookie", token))
|
||||
key = hash_secret("cookie", token)
|
||||
|
||||
session = Session.create(
|
||||
user=user_uuid,
|
||||
|
||||
+1
-21
@@ -16,7 +16,6 @@ import secrets
|
||||
from datetime import UTC, datetime
|
||||
from uuid import UUID
|
||||
|
||||
import base64url
|
||||
import httpx
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
@@ -188,25 +187,6 @@ class TestExceptionHandlers:
|
||||
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 --------------------
|
||||
|
||||
|
||||
@@ -1320,7 +1300,7 @@ class TestAdminSessions:
|
||||
test_user,
|
||||
):
|
||||
"""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(
|
||||
f"/auth/api/admin/users/{test_user.uuid}/sessions/{session_db_key}",
|
||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||
|
||||
@@ -355,34 +355,6 @@ class TestErrorHandling:
|
||||
class TestForwardAuthHtmlResponse:
|
||||
"""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
|
||||
async def test_forward_with_expired_session_clears_cookie(
|
||||
self, client: httpx.AsyncClient
|
||||
|
||||
Reference in New Issue
Block a user