Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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)
|
await clearSessionCookie(page)
|
||||||
|
|
||||||
// Make API call that triggers 401 (don't await - it blocks until iframe resolves)
|
// Make API call that triggers 401 (don't await - it blocks until iframe resolves)
|
||||||
const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'POST').catch(e => e)
|
const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'GET').catch(e => e)
|
||||||
console.log('✓ Auth iframe appeared on 401')
|
console.log('✓ Auth iframe appeared on 401')
|
||||||
|
|
||||||
// Verify it's in login mode (not reauth)
|
// Verify it's in login mode (not reauth)
|
||||||
@@ -268,7 +268,7 @@ test.describe('API Mode - 401 Login Flow', () => {
|
|||||||
await setupTestHarness(page)
|
await setupTestHarness(page)
|
||||||
|
|
||||||
// Make API call that triggers 401
|
// Make API call that triggers 401
|
||||||
const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'POST')
|
const apiCallPromise = makeApiCall(page, '/auth/api/user-info', 'GET')
|
||||||
|
|
||||||
// Wait for auth iframe to appear
|
// Wait for auth iframe to appear
|
||||||
await waitForAuthIframe(page)
|
await waitForAuthIframe(page)
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
import { spawn } from 'child_process'
|
import { execSync, spawn } from 'child_process'
|
||||||
import { join, dirname } from 'path'
|
import { join, dirname } from 'path'
|
||||||
import { existsSync, mkdirSync, writeFileSync } from 'fs'
|
import { existsSync, mkdirSync, writeFileSync } from 'fs'
|
||||||
import { fileURLToPath } from 'url'
|
import { fileURLToPath } from 'url'
|
||||||
@@ -31,6 +31,11 @@ export default async function globalSetup() {
|
|||||||
mkdirSync(testDataDir, { recursive: true })
|
mkdirSync(testDataDir, { recursive: true })
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Build the package first
|
||||||
|
console.log(' Building package with uv build...')
|
||||||
|
execSync('uv build', { cwd: projectRoot, stdio: 'inherit' })
|
||||||
|
console.log(' ✅ Build complete\n')
|
||||||
|
|
||||||
console.log(' Starting server with in-memory database...')
|
console.log(' Starting server with in-memory database...')
|
||||||
if (COLLECT_COVERAGE) {
|
if (COLLECT_COVERAGE) {
|
||||||
console.log(' 📊 Coverage collection enabled for Python backend')
|
console.log(' 📊 Coverage collection enabled for Python backend')
|
||||||
|
|||||||
+21
-72
@@ -13,8 +13,8 @@
|
|||||||
<script setup>
|
<script setup>
|
||||||
import { computed, onMounted, onUnmounted, ref } from 'vue'
|
import { computed, onMounted, onUnmounted, ref } from 'vue'
|
||||||
import { useAuthStore } from '@/stores/auth'
|
import { useAuthStore } from '@/stores/auth'
|
||||||
import { apiJson, SessionValidator, createAuthIframe, removeAuthIframe } from 'paskia'
|
import { apiJson, SessionValidator } from 'paskia'
|
||||||
import { getAuthIframeUrl } from '@/utils/api'
|
import { updateThemeFromSession } from '@/utils/theme'
|
||||||
import StatusMessage from '@/components/StatusMessage.vue'
|
import StatusMessage from '@/components/StatusMessage.vue'
|
||||||
import ProfileView from '@/components/ProfileView.vue'
|
import ProfileView from '@/components/ProfileView.vue'
|
||||||
import HostProfileView from '@/components/HostProfileView.vue'
|
import HostProfileView from '@/components/HostProfileView.vue'
|
||||||
@@ -48,90 +48,49 @@ const isHostMode = computed(() => {
|
|||||||
return currentHost !== configuredHost
|
return currentHost !== configuredHost
|
||||||
})
|
})
|
||||||
|
|
||||||
function terminateSession() {
|
function onSessionLost(e) {
|
||||||
store.userInfo = null
|
store.userInfo = null
|
||||||
viewState.value = 'terminal'
|
store.ctx = null
|
||||||
|
if (e?.name === 'AuthCancelledError') {
|
||||||
|
viewState.value = 'terminal'
|
||||||
|
} else {
|
||||||
|
store.showMessage(e?.message || 'Session lost', 'error', 5000)
|
||||||
|
viewState.value = 'terminal'
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
const userUuidGetter = () => store.ctx?.user.uuid
|
const userUuidGetter = () => store.ctx?.user.uuid
|
||||||
const sessionValidator = new SessionValidator(userUuidGetter, terminateSession)
|
const sessionValidator = new SessionValidator(userUuidGetter, onSessionLost)
|
||||||
|
|
||||||
onMounted(() => sessionValidator.start())
|
onMounted(() => sessionValidator.start())
|
||||||
onUnmounted(() => sessionValidator.stop())
|
onUnmounted(() => sessionValidator.stop())
|
||||||
|
|
||||||
async function loadUserInfo() {
|
async function loadUserInfo() {
|
||||||
|
viewState.value = 'loading'
|
||||||
|
loadingMessage.value = 'Loading...'
|
||||||
try {
|
try {
|
||||||
|
// apiJson handles 401/403 with auth.iframe automatically:
|
||||||
|
// shows overlay iframe, waits for auth, retries the request.
|
||||||
const [validateData, userInfoData] = await Promise.all([
|
const [validateData, userInfoData] = await Promise.all([
|
||||||
apiJson('/auth/api/validate', { method: 'POST' }),
|
apiJson('/auth/api/validate', { method: 'POST' }),
|
||||||
apiJson('/auth/api/user-info', { method: 'GET' })
|
apiJson('/auth/api/user-info', { method: 'GET' })
|
||||||
])
|
])
|
||||||
store.userInfo = userInfoData
|
store.userInfo = userInfoData
|
||||||
store.ctx = validateData.ctx
|
store.ctx = validateData.ctx
|
||||||
|
updateThemeFromSession(store.userInfo)
|
||||||
// Verify that the user UUIDs match between user-info and validate responses
|
// Verify that the user UUIDs match between user-info and validate responses
|
||||||
if (store.userInfo.user.uuid !== store.ctx.user.uuid) {
|
if (store.userInfo.user.uuid !== store.ctx.user.uuid) {
|
||||||
console.error('User UUID mismatch between user-info and validate responses')
|
console.error('User UUID mismatch between user-info and validate responses')
|
||||||
window.location.reload()
|
window.location.reload()
|
||||||
return false
|
return
|
||||||
}
|
}
|
||||||
viewState.value = 'profile'
|
viewState.value = 'profile'
|
||||||
return true
|
} catch (e) {
|
||||||
} catch {
|
onSessionLost(e)
|
||||||
store.userInfo = null
|
|
||||||
store.ctx = null
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
async function showAuthIframe() {
|
|
||||||
const url = await getAuthIframeUrl('login')
|
|
||||||
createAuthIframe(url)
|
|
||||||
loadingMessage.value = 'Authentication required...'
|
|
||||||
}
|
|
||||||
|
|
||||||
function handleAuthMessage(event) {
|
|
||||||
const data = event.data
|
|
||||||
if (!data?.type) return
|
|
||||||
|
|
||||||
switch (data.type) {
|
|
||||||
case 'auth-success':
|
|
||||||
// Authentication successful - reload user info
|
|
||||||
removeAuthIframe()
|
|
||||||
viewState.value = 'loading'
|
|
||||||
loadingMessage.value = 'Loading user profile...'
|
|
||||||
loadUserInfo()
|
|
||||||
break
|
|
||||||
|
|
||||||
case 'auth-error':
|
|
||||||
// Authentication failed - keep iframe open so user can retry
|
|
||||||
if (data.cancelled) {
|
|
||||||
console.log('Authentication cancelled by user')
|
|
||||||
} else {
|
|
||||||
store.showMessage(data.message || 'Authentication failed', 'error', 5000)
|
|
||||||
}
|
|
||||||
break
|
|
||||||
|
|
||||||
case 'auth-cancelled':
|
|
||||||
// Legacy support - treat as auth-error with cancelled flag
|
|
||||||
console.log('Authentication cancelled')
|
|
||||||
break
|
|
||||||
|
|
||||||
case 'auth-back':
|
|
||||||
// User clicked Back - show terminal state
|
|
||||||
removeAuthIframe()
|
|
||||||
terminateSession()
|
|
||||||
break
|
|
||||||
|
|
||||||
case 'auth-close-request':
|
|
||||||
// Legacy support - treat as back
|
|
||||||
removeAuthIframe()
|
|
||||||
break
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
onMounted(async () => {
|
onMounted(async () => {
|
||||||
// Listen for postMessage from auth iframe
|
|
||||||
window.addEventListener('message', handleAuthMessage)
|
|
||||||
|
|
||||||
// Load settings
|
// Load settings
|
||||||
await store.loadSettings()
|
await store.loadSettings()
|
||||||
|
|
||||||
@@ -145,17 +104,7 @@ onMounted(async () => {
|
|||||||
document.title = inHostMode ? `${rpName} · Account summary` : rpName
|
document.title = inHostMode ? `${rpName} · Account summary` : rpName
|
||||||
}
|
}
|
||||||
|
|
||||||
// Try to load user info
|
// Load user info (apiJson handles auth iframe if needed)
|
||||||
const success = await loadUserInfo()
|
await loadUserInfo()
|
||||||
|
|
||||||
if (!success) {
|
|
||||||
// Need authentication - show login iframe
|
|
||||||
showAuthIframe()
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
onUnmounted(() => {
|
|
||||||
window.removeEventListener('message', handleAuthMessage)
|
|
||||||
removeAuthIframe()
|
|
||||||
})
|
})
|
||||||
</script>
|
</script>
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ import AdminDialogs from '@/admin/AdminDialogs.vue'
|
|||||||
import { useAuthStore } from '@/stores/auth'
|
import { useAuthStore } from '@/stores/auth'
|
||||||
import { adminUiPath, makeUiHref } from '@/utils/settings'
|
import { adminUiPath, makeUiHref } from '@/utils/settings'
|
||||||
import { apiJson, SessionValidator } from 'paskia'
|
import { apiJson, SessionValidator } from 'paskia'
|
||||||
|
import { updateThemeFromSession } from '@/utils/theme'
|
||||||
import { uuidv7 } from 'uuidv7'
|
import { uuidv7 } from 'uuidv7'
|
||||||
import { getDirection } from '@/utils/keynav'
|
import { getDirection } from '@/utils/keynav'
|
||||||
import { goBack } from '@/utils/helpers'
|
import { goBack } from '@/utils/helpers'
|
||||||
@@ -197,6 +198,7 @@ function orgUserCount(org) {
|
|||||||
async function loadUserInfo() {
|
async function loadUserInfo() {
|
||||||
const data = await apiJson('/auth/api/validate', { method: 'POST' })
|
const data = await apiJson('/auth/api/validate', { method: 'POST' })
|
||||||
info.value = data
|
info.value = data
|
||||||
|
updateThemeFromSession(data.ctx)
|
||||||
authenticated.value = true
|
authenticated.value = true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -337,24 +339,9 @@ async function moveUserToRole(userUuid, user, targetRoleUuid) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
function onUserDragStart(e, userUuid, org) {
|
function moveUserToRoleFromDrag(userUuid, newRoleUuid) {
|
||||||
e.dataTransfer.effectAllowed = 'move'
|
const user = selectedOrg.value?.users?.[userUuid]
|
||||||
e.dataTransfer.setData('text/plain', JSON.stringify({ user_uuid: userUuid, org }))
|
if (user) moveUserToRole(userUuid, user, newRoleUuid)
|
||||||
}
|
|
||||||
|
|
||||||
function onRoleDragOver(e) {
|
|
||||||
e.preventDefault()
|
|
||||||
e.dataTransfer.dropEffect = 'move'
|
|
||||||
}
|
|
||||||
|
|
||||||
function onRoleDrop(e, org, role) {
|
|
||||||
e.preventDefault()
|
|
||||||
try {
|
|
||||||
const data = JSON.parse(e.dataTransfer.getData('text/plain'))
|
|
||||||
if (data.org !== org.uuid) return // only within same org
|
|
||||||
const user = org.users[data.user_uuid]
|
|
||||||
if (user) moveUserToRole(data.user_uuid, user, role.uuid)
|
|
||||||
} catch (_) { /* ignore */ }
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Role actions
|
// Role actions
|
||||||
@@ -999,10 +986,8 @@ async function submitDialog() {
|
|||||||
@create-user-in-role="createUserInRole"
|
@create-user-in-role="createUserInRole"
|
||||||
@open-user="openUser"
|
@open-user="openUser"
|
||||||
@toggle-role-permission="toggleRolePermission"
|
@toggle-role-permission="toggleRolePermission"
|
||||||
@on-role-drag-over="onRoleDragOver"
|
@move-user-to-role="moveUserToRoleFromDrag"
|
||||||
@navigate-out="handlePanelNavigateOut"
|
@navigate-out="handlePanelNavigateOut"
|
||||||
@on-role-drop="onRoleDrop"
|
|
||||||
@on-user-drag-start="onUserDragStart"
|
|
||||||
/>
|
/>
|
||||||
|
|
||||||
<AdminOidcDetail
|
<AdminOidcDetail
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
<head>
|
<head>
|
||||||
<meta charset="UTF-8">
|
<meta charset="UTF-8">
|
||||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||||
<script>{let t=localStorage.getItem('paskia-theme');if(t!=='light'&&t!=='dark')t=new URLSearchParams(location.hash.slice(1)).get('theme');(t==='dark'||t!=='light'&&matchMedia('(prefers-color-scheme:dark)').matches)&&document.documentElement.classList.add('dark');if(window.location.pathname==='/auth/restricted/iframe')document.documentElement.style.background='transparent'}</script>
|
<script>{let t=localStorage.getItem('paskia-theme');if(!t){let p=new URLSearchParams(location.hash.slice(1)).get('theme');if(p==='light'||p==='dark')t=p}(t==='dark'||t!=='light'&&matchMedia('(prefers-color-scheme:dark)').matches)&&document.documentElement.classList.add('dark');if(window.location.pathname==='/auth/restricted/iframe')document.documentElement.style.background='transparent'}</script>
|
||||||
<link rel="stylesheet" href="/src/assets/style.css">
|
<link rel="stylesheet" href="/src/assets/style.css">
|
||||||
</head>
|
</head>
|
||||||
<body>
|
<body>
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
// Early theme for restricted app - first URL param wins, then localStorage
|
// Early theme for restricted app - user preference (localStorage) wins, then URL param
|
||||||
import { applyTheme, getCachedTheme } from '@/utils/theme.js'
|
import { applyTheme, getCachedTheme } from '@/utils/theme.js'
|
||||||
|
|
||||||
function getTheme() {
|
function getTheme() {
|
||||||
const params = new URLSearchParams(location.hash.slice(1))
|
const params = new URLSearchParams(location.hash.slice(1))
|
||||||
return params.get('theme') || getCachedTheme() || ''
|
return getCachedTheme() || params.get('theme') || ''
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply theme class to document root
|
// Apply theme class to document root
|
||||||
|
|||||||
@@ -8,7 +8,7 @@
|
|||||||
|
|
||||||
<main class="view-root">
|
<main class="view-root">
|
||||||
<div class="surface surface--tight reset-container">
|
<div class="surface surface--tight reset-container">
|
||||||
<header class="view-header reset-header">
|
<header class="view-header center">
|
||||||
<h1>🔑 Registration</h1>
|
<h1>🔑 Registration</h1>
|
||||||
<p class="view-lede">
|
<p class="view-lede">
|
||||||
{{ subtitleMessage }}
|
{{ subtitleMessage }}
|
||||||
@@ -60,6 +60,7 @@ import { computed, onMounted, reactive, ref } from 'vue'
|
|||||||
import passkey from '@/utils/passkey'
|
import passkey from '@/utils/passkey'
|
||||||
import { getSettings, uiBasePath } from '@/utils/settings'
|
import { getSettings, uiBasePath } from '@/utils/settings'
|
||||||
import { apiJson, ApiError, getUserFriendlyErrorMessage } from 'paskia'
|
import { apiJson, ApiError, getUserFriendlyErrorMessage } from 'paskia'
|
||||||
|
import { updateThemeFromSession } from '@/utils/theme'
|
||||||
|
|
||||||
const status = reactive({
|
const status = reactive({
|
||||||
show: false,
|
show: false,
|
||||||
@@ -80,7 +81,7 @@ const sessionDescriptor = computed(() => tokenInfo.value?.token_type || 'your en
|
|||||||
const subtitleMessage = computed(() => {
|
const subtitleMessage = computed(() => {
|
||||||
if (initializing.value) return 'Preparing your secure enrollment…'
|
if (initializing.value) return 'Preparing your secure enrollment…'
|
||||||
if (!canRegister.value) return 'This authentication link is no longer valid.'
|
if (!canRegister.value) return 'This authentication link is no longer valid.'
|
||||||
return `Finish up ${sessionDescriptor.value}. You may edit the name below if needed, and it will be saved to your passkey.`
|
return `Finish up ${sessionDescriptor.value}. The name entered will be stored on your passkey and on our system.`
|
||||||
})
|
})
|
||||||
|
|
||||||
const basePath = computed(() => uiBasePath())
|
const basePath = computed(() => uiBasePath())
|
||||||
@@ -117,6 +118,7 @@ async function fetchTokenInfo() {
|
|||||||
headers: { 'Authorization': `Bearer ${token.value}` },
|
headers: { 'Authorization': `Bearer ${token.value}` },
|
||||||
})
|
})
|
||||||
displayName.value = tokenInfo.value.display_name
|
displayName.value = tokenInfo.value.display_name
|
||||||
|
if (tokenInfo.value.theme) updateThemeFromSession({ user: { theme: tokenInfo.value.theme } })
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Failed to load token info', error)
|
console.error('Failed to load token info', error)
|
||||||
const message = error instanceof ApiError
|
const message = error instanceof ApiError
|
||||||
@@ -201,14 +203,14 @@ onMounted(async () => {
|
|||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style scoped>
|
||||||
|
main.view-root { min-height: 100vh; align-items: center; justify-content: center; padding: 2rem 1rem; }
|
||||||
.reset-container {
|
.reset-container {
|
||||||
max-width: 560px;
|
max-width: 520px;
|
||||||
margin: 0 auto;
|
margin: 0 auto;
|
||||||
width: 100%;
|
width: 100%;
|
||||||
}
|
display: flex;
|
||||||
|
flex-direction: column;
|
||||||
.reset-header {
|
gap: 1.75rem;
|
||||||
text-align: center;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
.section-body {
|
.section-body {
|
||||||
|
|||||||
@@ -15,7 +15,8 @@
|
|||||||
"qrcode": "^1.5.4",
|
"qrcode": "^1.5.4",
|
||||||
"sirv": "^3.0.2",
|
"sirv": "^3.0.2",
|
||||||
"uuidv7": "^1.1.0",
|
"uuidv7": "^1.1.0",
|
||||||
"vue": "^3.5.17"
|
"vue": "^3.5.17",
|
||||||
|
"vuedraggable": "^4.1.0"
|
||||||
},
|
},
|
||||||
"devDependencies": {
|
"devDependencies": {
|
||||||
"@vitejs/plugin-vue": "^6.0.0",
|
"@vitejs/plugin-vue": "^6.0.0",
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
<script setup>
|
<script setup>
|
||||||
import { computed, ref } from 'vue'
|
import { computed, ref } from 'vue'
|
||||||
|
import draggable from 'vuedraggable'
|
||||||
import { getDirection, navigateButtonRow, focusPreferred } from '@/utils/keynav'
|
import { getDirection, navigateButtonRow, focusPreferred } from '@/utils/keynav'
|
||||||
|
|
||||||
const props = defineProps({
|
const props = defineProps({
|
||||||
@@ -8,7 +9,7 @@ const props = defineProps({
|
|||||||
navigationDisabled: { type: Boolean, default: false }
|
navigationDisabled: { type: Boolean, default: false }
|
||||||
})
|
})
|
||||||
|
|
||||||
const emit = defineEmits(['updateOrg', 'createRole', 'updateRole', 'deleteRole', 'createUserInRole', 'openUser', 'toggleRolePermission', 'onRoleDragOver', 'onRoleDrop', 'onUserDragStart', 'navigateOut'])
|
const emit = defineEmits(['updateOrg', 'createRole', 'updateRole', 'deleteRole', 'createUserInRole', 'openUser', 'toggleRolePermission', 'moveUserToRole', 'navigateOut'])
|
||||||
|
|
||||||
// Template refs for navigation
|
// Template refs for navigation
|
||||||
const orgTitleRef = ref(null)
|
const orgTitleRef = ref(null)
|
||||||
@@ -50,6 +51,14 @@ function roleUserCount(roleUuid) {
|
|||||||
return Object.values(props.selectedOrg.users).filter(u => u.role === roleUuid).length
|
return Object.values(props.selectedOrg.users).filter(u => u.role === roleUuid).length
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function onUserChange(evt, targetRoleUuid) {
|
||||||
|
// Only handle 'added' events (when a user is dropped into this role)
|
||||||
|
if (evt.added) {
|
||||||
|
const userUuid = evt.added.element.uuid
|
||||||
|
emit('moveUserToRole', userUuid, targetRoleUuid)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
function permissionDisplayName(scope) {
|
function permissionDisplayName(scope) {
|
||||||
return props.permissions.find(p => p.scope === scope)?.display_name || scope
|
return props.permissions.find(p => p.scope === scope)?.display_name || scope
|
||||||
}
|
}
|
||||||
@@ -350,8 +359,6 @@ defineExpose({ focusFirstElement })
|
|||||||
v-for="(r, roleIndex) in sortedRoles"
|
v-for="(r, roleIndex) in sortedRoles"
|
||||||
:key="r.uuid"
|
:key="r.uuid"
|
||||||
class="role-column"
|
class="role-column"
|
||||||
@dragover="$emit('onRoleDragOver', $event)"
|
|
||||||
@drop="e => $emit('onRoleDrop', e, selectedOrg, r)"
|
|
||||||
>
|
>
|
||||||
<div class="role-header" @keydown="e => handleRoleHeaderKeydown(e, roleIndex)">
|
<div class="role-header" @keydown="e => handleRoleHeaderKeydown(e, roleIndex)">
|
||||||
<strong class="role-name" :title="r.uuid">
|
<strong class="role-name" :title="r.uuid">
|
||||||
@@ -363,26 +370,32 @@ defineExpose({ focusFirstElement })
|
|||||||
<button @click="$emit('createUserInRole', selectedOrg, r)" class="plus-btn" aria-label="Add user" title="Add user">➕</button>
|
<button @click="$emit('createUserInRole', selectedOrg, r)" class="plus-btn" aria-label="Add user" title="Add user">➕</button>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<template v-if="roleUserCount(r.uuid) > 0">
|
<div class="user-list-wrapper">
|
||||||
<ul class="user-list" @keydown="handleUserListKeydown">
|
<draggable
|
||||||
<li
|
:list="roleUsers(r.uuid)"
|
||||||
v-for="u in roleUsers(r.uuid)"
|
group="users"
|
||||||
:key="u.uuid"
|
item-key="uuid"
|
||||||
class="user-chip"
|
tag="ul"
|
||||||
tabindex="0"
|
class="user-list"
|
||||||
draggable="true"
|
@change="evt => onUserChange(evt, r.uuid)"
|
||||||
@dragstart="e => $emit('onUserDragStart', e, u.uuid, selectedOrg.uuid)"
|
@keydown="handleUserListKeydown"
|
||||||
@click="$emit('openUser', u)"
|
>
|
||||||
@keydown.enter="$emit('openUser', u)"
|
<template #item="{ element: u }">
|
||||||
:title="u.uuid"
|
<li
|
||||||
>
|
class="user-chip"
|
||||||
<span class="name">{{ u.display_name }}</span>
|
tabindex="0"
|
||||||
<span class="meta">{{ u.last_seen ? new Date(u.last_seen).toLocaleDateString() : '—' }}</span>
|
@click="$emit('openUser', u)"
|
||||||
</li>
|
@keydown.enter="$emit('openUser', u)"
|
||||||
</ul>
|
:title="u.uuid"
|
||||||
</template>
|
>
|
||||||
<div v-else class="empty-role">
|
<span class="name">{{ u.display_name }}</span>
|
||||||
<p class="empty-text muted">No members</p>
|
<span class="meta">{{ u.last_seen ? new Date(u.last_seen).toLocaleDateString() : '—' }}</span>
|
||||||
|
</li>
|
||||||
|
</template>
|
||||||
|
</draggable>
|
||||||
|
<div v-if="roleUserCount(r.uuid) === 0" class="empty-role">
|
||||||
|
<p class="empty-text muted">No members</p>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -391,23 +404,27 @@ defineExpose({ focusFirstElement })
|
|||||||
|
|
||||||
<style scoped>
|
<style scoped>
|
||||||
.card.surface { padding: var(--space-lg); }
|
.card.surface { padding: var(--space-lg); }
|
||||||
.org-title { display: flex; align-items: center; gap: var(--space-sm); margin-bottom: var(--space-lg); }
|
.org-title { display: flex; align-items: center; gap: var(--space-sm); margin-bottom: var(--space-lg); font-size: 1.65rem; }
|
||||||
.org-name { font-size: 1.5rem; font-weight: 600; color: var(--color-heading); }
|
.org-name { font-weight: 600; color: var(--color-heading); }
|
||||||
.perm-matrix-grid .role-head { display: flex; align-items: flex-end; justify-content: center; }
|
.perm-matrix-grid .role-head { display: flex; align-items: flex-end; justify-content: center; }
|
||||||
.perm-matrix-grid .role-head span { writing-mode: vertical-rl; transform: rotate(180deg); font-size: 0.65rem; }
|
.perm-matrix-grid .role-head span { writing-mode: vertical-rl; transform: rotate(180deg); font-size: 0.65rem; }
|
||||||
.perm-matrix-grid .add-role-head { cursor: pointer; }
|
.perm-matrix-grid .add-role-head { cursor: pointer; }
|
||||||
.roles-grid { display: flex; gap: var(--space-lg); margin-top: var(--space-lg); }
|
.roles-grid { display: flex; flex-wrap: wrap; gap: var(--space-lg); margin-top: var(--space-lg); justify-content: flex-start; align-items: stretch; }
|
||||||
.role-column { flex: 1; min-width: 200px; border-radius: var(--radius-md); padding: var(--space-md); }
|
.role-column { flex: 0 0 240px; border-radius: var(--radius-md); padding: var(--space-md); display: flex; flex-direction: column; }
|
||||||
.role-header { display: flex; justify-content: space-between; align-items: center; margin-bottom: var(--space-md); }
|
.role-header { display: flex; justify-content: space-between; align-items: center; margin-bottom: var(--space-md); }
|
||||||
.role-name { display: flex; align-items: center; gap: var(--space-xs); font-size: 1.1rem; color: var(--color-heading); }
|
.role-name { display: flex; align-items: center; gap: var(--space-xs); font-size: 1.1rem; color: var(--color-heading); }
|
||||||
.role-actions { display: flex; gap: var(--space-xs); }
|
.role-actions { display: flex; gap: var(--space-xs); }
|
||||||
.plus-btn { background: none; color: var(--color-accent); border: none; border-radius: var(--radius-sm); padding: 0.25rem 0.45rem; font-size: 1.1rem; cursor: pointer; }
|
.plus-btn { background: none; color: var(--color-accent); border: none; border-radius: var(--radius-sm); padding: 0.25rem 0.45rem; font-size: 1.1rem; cursor: pointer; }
|
||||||
.plus-btn:hover { background: rgba(37, 99, 235, 0.18); }
|
.plus-btn:hover { background: rgba(37, 99, 235, 0.18); }
|
||||||
.user-list { list-style: none; padding: 0; margin: 0; display: flex; flex-direction: column; gap: var(--space-xs); }
|
.user-list-wrapper { position: relative; flex: 1; display: flex; flex-direction: column; min-height: 5.5rem; }
|
||||||
.user-chip { background: var(--color-surface); border: 1px solid var(--color-border); border-radius: var(--radius-md); padding: 0.45rem 0.6rem; display: flex; justify-content: space-between; gap: var(--space-sm); cursor: grab; }
|
.user-list { list-style: none; padding: 0; margin: 0; display: flex; flex-direction: column; gap: var(--space-xs); flex: 1; }
|
||||||
|
.user-chip { background: var(--color-accent-strong); color: white; border: none; border-radius: var(--radius-md); padding: 0.45rem 0.6rem; display: flex; justify-content: space-between; gap: var(--space-sm); cursor: grab; }
|
||||||
.user-chip:focus { outline: 2px solid var(--color-accent); outline-offset: 1px; }
|
.user-chip:focus { outline: 2px solid var(--color-accent); outline-offset: 1px; }
|
||||||
.user-chip .meta { font-size: 0.7rem; color: var(--color-text-muted); }
|
.user-chip .meta { font-size: 0.7rem; color: rgba(255, 255, 255, 0.8); }
|
||||||
.empty-role { border: 1px dashed var(--color-border-strong); border-radius: var(--radius-md); padding: var(--space-sm); display: flex; flex-direction: column; gap: var(--space-xs); align-items: flex-start; }
|
.user-chip.sortable-ghost { opacity: 0.5; }
|
||||||
|
.user-chip.sortable-chosen { box-shadow: 0 4px 12px rgba(0, 0, 0, 0.2); }
|
||||||
|
.empty-role { position: absolute; inset: 0; border: 1px dashed var(--color-border-strong); border-radius: var(--radius-md); display: flex; align-items: center; justify-content: center; pointer-events: none; }
|
||||||
|
.user-list:has(.sortable-ghost) + .empty-role { display: none; }
|
||||||
.empty-text { margin: 0; }
|
.empty-text { margin: 0; }
|
||||||
|
|
||||||
@media (max-width: 720px) {
|
@media (max-width: 720px) {
|
||||||
|
|||||||
@@ -823,9 +823,6 @@ th {
|
|||||||
|
|
||||||
.user-info {
|
.user-info {
|
||||||
display: grid;
|
display: grid;
|
||||||
border-radius: var(--radius-md);
|
|
||||||
background: var(--color-surface);
|
|
||||||
padding: 1rem;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
.user-details {
|
.user-details {
|
||||||
|
|||||||
@@ -10,9 +10,9 @@
|
|||||||
<UserBasicInfo
|
<UserBasicInfo
|
||||||
v-if="ctx"
|
v-if="ctx"
|
||||||
:name="ctx.user.display_name"
|
:name="ctx.user.display_name"
|
||||||
:visits="authStore.userInfo?.visits || 0"
|
:visits="authStore.userInfo.user.visits"
|
||||||
:created-at="authStore.userInfo?.created_at"
|
:created-at="authStore.userInfo.user.created_at"
|
||||||
:last-seen="authStore.userInfo?.last_seen"
|
:last-seen="authStore.userInfo.user.last_seen"
|
||||||
:email="ctx.user.email"
|
:email="ctx.user.email"
|
||||||
:telephone="ctx.user.telephone"
|
:telephone="ctx.user.telephone"
|
||||||
:org-display-name="orgDisplayName"
|
:org-display-name="orgDisplayName"
|
||||||
|
|||||||
@@ -22,8 +22,8 @@
|
|||||||
:created-at="authStore.userInfo.user.created_at"
|
:created-at="authStore.userInfo.user.created_at"
|
||||||
:last-seen="authStore.userInfo.user.last_seen"
|
:last-seen="authStore.userInfo.user.last_seen"
|
||||||
:loading="authStore.isLoading"
|
:loading="authStore.isLoading"
|
||||||
:org-display-name="authStore.ctx?.org.display_name"
|
:org-display-name="authStore.userInfo.org.display_name"
|
||||||
:role-name="authStore.ctx?.role.display_name"
|
:role-name="authStore.userInfo.role.display_name"
|
||||||
update-endpoint="/auth/api/user/info"
|
update-endpoint="/auth/api/user/info"
|
||||||
@saved="authStore.loadUserInfo()"
|
@saved="authStore.loadUserInfo()"
|
||||||
@edit="openEditDialog"
|
@edit="openEditDialog"
|
||||||
@@ -53,7 +53,7 @@
|
|||||||
<CredentialList
|
<CredentialList
|
||||||
ref="credentialList"
|
ref="credentialList"
|
||||||
:credentials="credentials"
|
:credentials="credentials"
|
||||||
:aaguid-info="authStore.userInfo?.aaguid_info || {}"
|
:aaguid-info="authStore.userInfo.aaguid_info"
|
||||||
:loading="authStore.isLoading"
|
:loading="authStore.isLoading"
|
||||||
:hovered-credential-uuid="hoveredCredentialUuid"
|
:hovered-credential-uuid="hoveredCredentialUuid"
|
||||||
:hovered-session-credential-uuid="hoveredSession?.credential"
|
:hovered-session-credential-uuid="hoveredSession?.credential"
|
||||||
@@ -184,11 +184,11 @@ const hasActiveModal = computed(() => showEditDialog.value || showRegLink.value)
|
|||||||
|
|
||||||
watch(showEditDialog, (open) => {
|
watch(showEditDialog, (open) => {
|
||||||
if (!open) return
|
if (!open) return
|
||||||
const user = authStore.userInfo?.user
|
const user = authStore.userInfo.user
|
||||||
editName.value = user?.display_name ?? ''
|
editName.value = user.display_name ?? ''
|
||||||
editEmail.value = user?.email ?? ''
|
editEmail.value = user.email ?? ''
|
||||||
editUsername.value = user?.preferred_username ?? ''
|
editUsername.value = user.preferred_username ?? ''
|
||||||
editTelephone.value = user?.telephone ?? ''
|
editTelephone.value = user.telephone ?? ''
|
||||||
editError.value = ''
|
editError.value = ''
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -341,7 +341,7 @@ const handleDelete = async (credential) => {
|
|||||||
|
|
||||||
const rpName = computed(() => authStore.settings?.rp_name || 'this service')
|
const rpName = computed(() => authStore.settings?.rp_name || 'this service')
|
||||||
const paskiaVersion = computed(() => authStore.settings?.version || '')
|
const paskiaVersion = computed(() => authStore.settings?.version || '')
|
||||||
const sessions = computed(() => authStore.userInfo?.sessions || {})
|
const sessions = computed(() => authStore.userInfo.sessions)
|
||||||
const currentSessionHost = computed(() => {
|
const currentSessionHost = computed(() => {
|
||||||
const currentSession = Object.values(sessions.value).find(session => session.is_current)
|
const currentSession = Object.values(sessions.value).find(session => session.is_current)
|
||||||
return currentSession?.host || 'this host'
|
return currentSession?.host || 'this host'
|
||||||
@@ -365,12 +365,12 @@ const logoutEverywhere = async () => { await authStore.logoutEverywhere() }
|
|||||||
const logout = async () => { await authStore.logout() }
|
const logout = async () => { await authStore.logout() }
|
||||||
const openEditDialog = () => { showEditDialog.value = true }
|
const openEditDialog = () => { showEditDialog.value = true }
|
||||||
const isAdmin = computed(() => {
|
const isAdmin = computed(() => {
|
||||||
const perms = authStore.ctx?.permissions
|
const perms = Object.values(authStore.userInfo.permissions).map(p => p.scope)
|
||||||
return perms?.includes('auth:admin') || perms?.includes('auth:org:admin')
|
return perms.includes('auth:admin') || perms.includes('auth:org:admin')
|
||||||
})
|
})
|
||||||
const hasMultipleSessions = computed(() => Object.keys(sessions.value).length > 1)
|
const hasMultipleSessions = computed(() => Object.keys(sessions.value).length > 1)
|
||||||
const credentials = computed(() =>
|
const credentials = computed(() =>
|
||||||
Object.entries(authStore.userInfo?.credentials || {}).map(([uuid, c]) => ({ ...c, credential: uuid }))
|
Object.entries(authStore.userInfo.credentials).map(([uuid, c]) => ({ ...c, credential: uuid }))
|
||||||
)
|
)
|
||||||
const useWideLayout = computed(() => {
|
const useWideLayout = computed(() => {
|
||||||
// Check if any single site has more than 8 sessions
|
// Check if any single site has more than 8 sessions
|
||||||
|
|||||||
@@ -37,7 +37,7 @@
|
|||||||
<script setup>
|
<script setup>
|
||||||
import { ref, computed, onMounted, onUnmounted, nextTick } from 'vue'
|
import { ref, computed, onMounted, onUnmounted, nextTick } from 'vue'
|
||||||
import QRCodeDisplay from '@/components/QRCodeDisplay.vue'
|
import QRCodeDisplay from '@/components/QRCodeDisplay.vue'
|
||||||
import { apiJson, holdGlobalBackdrop, releaseGlobalBackdrop } from 'paskia'
|
import { apiJson, AuthCancelledError, holdGlobalBackdrop, releaseGlobalBackdrop } from 'paskia'
|
||||||
import { formatDate } from '@/utils/helpers'
|
import { formatDate } from '@/utils/helpers'
|
||||||
import { getDirection } from '@/utils/keynav'
|
import { getDirection } from '@/utils/keynav'
|
||||||
import { useAuthStore } from '@/stores/auth'
|
import { useAuthStore } from '@/stores/auth'
|
||||||
@@ -90,7 +90,9 @@ async function generateLink() {
|
|||||||
emit('close')
|
emit('close')
|
||||||
}
|
}
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
authStore.showMessage(e.message || 'Failed to generate link', 'error')
|
if (!(e instanceof AuthCancelledError)) {
|
||||||
|
authStore.showMessage(e.message || 'Failed to generate link', 'error')
|
||||||
|
}
|
||||||
emit('close')
|
emit('close')
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -163,6 +163,9 @@ async function startRemoteAuth() {
|
|||||||
|
|
||||||
// PoW challenge
|
// PoW challenge
|
||||||
const powChallenge = await ws.receive_json()
|
const powChallenge = await ws.receive_json()
|
||||||
|
if (powChallenge.status) {
|
||||||
|
throw new Error(powChallenge.detail || `Failed to connect: ${powChallenge.status}`)
|
||||||
|
}
|
||||||
if (powChallenge.pow) {
|
if (powChallenge.pow) {
|
||||||
const challenge = b64dec(powChallenge.pow.challenge)
|
const challenge = b64dec(powChallenge.pow.challenge)
|
||||||
const nonces = await solvePoW(challenge, powChallenge.pow.work)
|
const nonces = await solvePoW(challenge, powChallenge.pow.work)
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ import { getSettings, uiBasePath } from '@/utils/settings'
|
|||||||
import { fetchJson, getUserFriendlyErrorMessage } from 'paskia'
|
import { fetchJson, getUserFriendlyErrorMessage } from 'paskia'
|
||||||
import RemoteAuthRequest from '@/components/RemoteAuthRequest.vue'
|
import RemoteAuthRequest from '@/components/RemoteAuthRequest.vue'
|
||||||
import { focusDialogButton } from '@/utils/keynav'
|
import { focusDialogButton } from '@/utils/keynav'
|
||||||
|
import { updateThemeFromSession } from '@/utils/theme'
|
||||||
|
|
||||||
const props = defineProps({
|
const props = defineProps({
|
||||||
mode: {
|
mode: {
|
||||||
@@ -147,6 +148,7 @@ async function fetchSettings() {
|
|||||||
async function validateSession() {
|
async function validateSession() {
|
||||||
try {
|
try {
|
||||||
session.value = await fetchJson('/auth/api/validate', { method: 'POST' })
|
session.value = await fetchJson('/auth/api/validate', { method: 'POST' })
|
||||||
|
updateThemeFromSession(session.value?.ctx)
|
||||||
if (isAuthenticated.value && props.mode !== 'reauth') {
|
if (isAuthenticated.value && props.mode !== 'reauth') {
|
||||||
currentView.value = 'forbidden'
|
currentView.value = 'forbidden'
|
||||||
emit('forbidden', session.value)
|
emit('forbidden', session.value)
|
||||||
|
|||||||
@@ -118,7 +118,7 @@ const userLoaded = computed(() => !!props.name)
|
|||||||
.user-info-extra { grid-area: extra; padding-left: 1rem; border-left: 1px solid var(--color-border); flex-shrink: 0; }
|
.user-info-extra { grid-area: extra; padding-left: 1rem; border-left: 1px solid var(--color-border); flex-shrink: 0; }
|
||||||
.user-name-row { display: inline-flex; align-items: center; gap: 0.35rem; max-width: 100%; min-width: 0; }
|
.user-name-row { display: inline-flex; align-items: center; gap: 0.35rem; max-width: 100%; min-width: 0; }
|
||||||
.display-name { font-weight: 600; font-size: 1.05em; line-height: 1.2; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; min-width: 0; }
|
.display-name { font-weight: 600; font-size: 1.05em; line-height: 1.2; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; min-width: 0; }
|
||||||
.mini-btn { width: auto; padding: 4px 6px; margin: 0; font-size: 0.75em; line-height: 1; cursor: pointer; }
|
.mini-btn { width: auto; padding: 4px 6px; margin: 0; font-size: 0.75em; line-height: 1; cursor: pointer; background: transparent; }
|
||||||
.mini-btn:hover:not(:disabled) { background: var(--color-accent-soft); color: var(--color-accent); }
|
.mini-btn:hover:not(:disabled) { background: var(--color-accent-soft); color: var(--color-accent); }
|
||||||
.mini-btn:active:not(:disabled) { transform: translateY(1px); }
|
.mini-btn:active:not(:disabled) { transform: translateY(1px); }
|
||||||
.mini-btn:disabled { opacity: 0.5; cursor: not-allowed; }
|
.mini-btn:disabled { opacity: 0.5; cursor: not-allowed; }
|
||||||
|
|||||||
@@ -88,7 +88,7 @@ export const useAuthStore = defineStore('auth', {
|
|||||||
async loadUserInfo() {
|
async loadUserInfo() {
|
||||||
try {
|
try {
|
||||||
this.userInfo = await apiJson('/auth/api/user-info', { method: 'GET' })
|
this.userInfo = await apiJson('/auth/api/user-info', { method: 'GET' })
|
||||||
updateThemeFromSession(this.ctx)
|
updateThemeFromSession(this.userInfo)
|
||||||
console.log('User info loaded:', this.userInfo)
|
console.log('User info loaded:', this.userInfo)
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
// Suppress toast for 401/403 errors - the auth iframe will handle these
|
// Suppress toast for 401/403 errors - the auth iframe will handle these
|
||||||
|
|||||||
@@ -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,
|
isAuthIframeOpen,
|
||||||
hideAuthIframe,
|
hideAuthIframe,
|
||||||
showAuthIframe,
|
showAuthIframe,
|
||||||
createAuthIframe,
|
|
||||||
removeAuthIframe,
|
|
||||||
} from './overlay'
|
} from './overlay'
|
||||||
|
|
||||||
export { SessionValidator } from './validate'
|
export { SessionValidator } from './validate'
|
||||||
|
|||||||
+49
-126
@@ -1,21 +1,16 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import asyncio
|
|
||||||
import json
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
|
import msgspec
|
||||||
from fastapi_vue import server
|
from fastapi_vue import server
|
||||||
from fastapi_vue.hostutil import parse_endpoint
|
from fastapi_vue.hostutil import parse_endpoints
|
||||||
|
|
||||||
from paskia import db
|
from paskia.db.jsonl import load_readonly
|
||||||
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.util import startupbox
|
from paskia.util import startupbox
|
||||||
from paskia.util.hostutil import normalize_origin
|
from paskia.util.hostutil import normalize_origin
|
||||||
|
from paskia.util.runtime import RuntimeConfig
|
||||||
|
|
||||||
DEFAULT_PORT = 4401
|
DEFAULT_PORT = 4401
|
||||||
DEVMODE = os.getenv("PASKIA_DEV") == "1"
|
DEVMODE = os.getenv("PASKIA_DEV") == "1"
|
||||||
@@ -95,138 +90,66 @@ def main():
|
|||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
# Handle clearing options
|
# Load stored config (read-only, no writes, no global state)
|
||||||
if getattr(args, "auth_host", None) == "":
|
db_path = os.environ.get("PASKIA_DB", f"{args.rp_id}.paskiadb")
|
||||||
args.auth_host = None
|
config = load_readonly(db_path, rp_id=args.rp_id).config
|
||||||
if getattr(args, "rp_name", None) == "":
|
|
||||||
args.rp_name = None
|
|
||||||
if getattr(args, "listen", None) == "":
|
|
||||||
args.listen = None
|
|
||||||
|
|
||||||
# Init db and load stored config
|
# Override stored config with CLI args, or clear with empty string
|
||||||
asyncio.run(db.init(rp_id=args.rp_id))
|
if args.rp_name is not None:
|
||||||
stored_config = db.data().config
|
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
|
# Process and normalize auth_host
|
||||||
if args.rp_name is None and stored_config.rp_name is not None:
|
if config.auth_host:
|
||||||
args.rp_name = stored_config.rp_name
|
if "://" not in config.auth_host:
|
||||||
if args.origins is None and stored_config.origins is not None:
|
config.auth_host = f"https://{config.auth_host}"
|
||||||
args.origins = stored_config.origins
|
config.auth_host = config.auth_host.rstrip("/")
|
||||||
if args.auth_host is None and stored_config.auth_host is not None:
|
validate_auth_host(config.auth_host, config.rp_id)
|
||||||
args.auth_host = stored_config.auth_host
|
if config.origins:
|
||||||
if args.listen is None and stored_config.listen is not None:
|
config.origins.insert(0, config.auth_host) # Ensure first in origins
|
||||||
args.listen = stored_config.listen
|
|
||||||
|
|
||||||
# Parse first endpoint for config display and site_url
|
# Normalize and deduplicate while preserving order
|
||||||
first_listen = args.listen[0] if isinstance(args.listen, list) else args.listen
|
if config.origins:
|
||||||
endpoints = parse_endpoint(first_listen, DEFAULT_PORT)
|
config.origins = list({normalize_origin(o): ... for o in config.origins})
|
||||||
|
|
||||||
# Extract host/port/uds from first endpoint for config display and site_url
|
# Parse first endpoint for site_url fallback
|
||||||
ep = endpoints[0] if endpoints else {}
|
ep = next(iter(parse_endpoints(config.listen, DEFAULT_PORT)), {})
|
||||||
host = ep.get("host")
|
|
||||||
port = ep.get("port")
|
port = ep.get("port")
|
||||||
uds = ep.get("uds")
|
|
||||||
|
|
||||||
# Collect and normalize origins, handle auth_host
|
# Compute site_url and site_path
|
||||||
origins = [normalize_origin(o) for o in (getattr(args, "origins", None) or [])]
|
# Priority: auth_host > origins[0] > PASKIA_VITE_URL > http://localhost:port > https://rp_id
|
||||||
if args.auth_host:
|
site_path = "/auth/"
|
||||||
# Normalize auth_host with scheme
|
if config.auth_host:
|
||||||
if "://" not in args.auth_host:
|
site_url, site_path = config.auth_host, "/"
|
||||||
args.auth_host = f"https://{args.auth_host}"
|
elif config.origins:
|
||||||
|
site_url = config.origins[0]
|
||||||
validate_auth_host(args.auth_host, args.rp_id)
|
elif vite_url := os.environ.get("PASKIA_VITE_URL"):
|
||||||
|
site_url = vite_url.rstrip("/") # Devserver
|
||||||
# If origins are configured, ensure auth_host is included at top
|
elif config.rp_id == "localhost" and port:
|
||||||
if origins:
|
site_url = f"http://localhost:{port}" # Backend directly if we can
|
||||||
# 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/"
|
|
||||||
else:
|
else:
|
||||||
site_url = f"https://{args.rp_id}"
|
site_url = f"https://{config.rp_id}" # Assume external reverse proxy
|
||||||
site_path = "/auth/"
|
|
||||||
|
|
||||||
# Build runtime configuration
|
# Build runtime configuration for the server
|
||||||
config = PaskiaConfig(
|
runtime = RuntimeConfig(
|
||||||
rp_id=args.rp_id,
|
config=config,
|
||||||
rp_name=args.rp_name or None,
|
|
||||||
origins=origins or None,
|
|
||||||
auth_host=args.auth_host or None,
|
|
||||||
site_url=site_url,
|
site_url=site_url,
|
||||||
site_path=site_path,
|
site_path=site_path,
|
||||||
host=host,
|
save=args.save,
|
||||||
port=port,
|
|
||||||
uds=uds,
|
|
||||||
)
|
)
|
||||||
|
startupbox.print_startup_config(runtime)
|
||||||
|
os.environ["PASKIA_CONFIG"] = msgspec.json.encode(runtime).decode()
|
||||||
|
|
||||||
# Export configuration via single JSON env variable for worker processes
|
# Run the server (spawns processes in dev mode)
|
||||||
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())
|
|
||||||
|
|
||||||
dev = {"reload": True, "reload_dirs": ["paskia"]} if DEVMODE else {}
|
dev = {"reload": True, "reload_dirs": ["paskia"]} if DEVMODE else {}
|
||||||
server.run(
|
server.run(
|
||||||
"paskia.fastapi.mainapp:app",
|
"paskia.fastapi.mainapp:app",
|
||||||
listen=args.listen,
|
listen=config.listen,
|
||||||
default_port=DEFAULT_PORT,
|
default_port=DEFAULT_PORT,
|
||||||
log_level="warning",
|
log_level="warning",
|
||||||
access_log=False,
|
access_log=False,
|
||||||
|
|||||||
@@ -23,6 +23,11 @@ if TYPE_CHECKING:
|
|||||||
EXPIRES = SESSION_LIFETIME
|
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:
|
def expires() -> datetime:
|
||||||
return datetime.now(UTC) + EXPIRES
|
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):
|
def delete_credential(credential_uuid: UUID, auth: str, host: str | None = None):
|
||||||
"""Delete a specific credential for the current user."""
|
"""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:
|
if not ctx:
|
||||||
raise ValueError("Session expired")
|
raise ValueError("Session expired")
|
||||||
db.delete_credential(credential_uuid, ctx.user.uuid)
|
db.delete_credential(credential_uuid, ctx.user.uuid)
|
||||||
|
|||||||
@@ -1,4 +1,3 @@
|
|||||||
from dataclasses import dataclass
|
|
||||||
from datetime import timedelta
|
from datetime import timedelta
|
||||||
|
|
||||||
# Shared configuration constants for session management.
|
# Shared configuration constants for session management.
|
||||||
@@ -6,19 +5,3 @@ SESSION_LIFETIME = timedelta(hours=24)
|
|||||||
|
|
||||||
# Lifetime for reset links created by admins
|
# Lifetime for reset links created by admins
|
||||||
RESET_LIFETIME = timedelta(days=14)
|
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.
|
Database module for WebAuthn passkey authentication.
|
||||||
|
|
||||||
Read: Access data() directly, use build_* to convert to public structs.
|
Read: Access data() directly for structs.
|
||||||
CTX: data().session_ctx(key) returns SessionContext with effective permissions.
|
CTX: data().session_ctx(key) returns SessionContext with effective permissions.
|
||||||
Write: Functions validate and commit, or raise ValueError.
|
Write: Functions validate and commit, or raise ValueError.
|
||||||
|
|
||||||
@@ -10,7 +10,6 @@ Usage:
|
|||||||
|
|
||||||
# Read (after init)
|
# Read (after init)
|
||||||
user_data = db.data().users[user_uuid]
|
user_data = db.data().users[user_uuid]
|
||||||
user = db.build_user(user_uuid)
|
|
||||||
|
|
||||||
# Context
|
# Context
|
||||||
ctx = db.data().session_ctx(session_key)
|
ctx = db.data().session_ctx(session_key)
|
||||||
@@ -27,6 +26,7 @@ from paskia.db.background import (
|
|||||||
stop_cleanup,
|
stop_cleanup,
|
||||||
)
|
)
|
||||||
from paskia.db.bootstrap import bootstrap
|
from paskia.db.bootstrap import bootstrap
|
||||||
|
from paskia.db.jsonl import load_readonly
|
||||||
from paskia.db.lifecycle import cleanup_expired, init
|
from paskia.db.lifecycle import cleanup_expired, init
|
||||||
from paskia.db.operations import (
|
from paskia.db.operations import (
|
||||||
add_permission_to_org,
|
add_permission_to_org,
|
||||||
@@ -102,18 +102,12 @@ __all__ = [
|
|||||||
# Instance
|
# Instance
|
||||||
"data",
|
"data",
|
||||||
"init",
|
"init",
|
||||||
|
"load_readonly",
|
||||||
# Background
|
# Background
|
||||||
"start_background",
|
"start_background",
|
||||||
"stop_background",
|
"stop_background",
|
||||||
"start_cleanup",
|
"start_cleanup",
|
||||||
"stop_cleanup",
|
"stop_cleanup",
|
||||||
# Builders
|
|
||||||
"build_credential",
|
|
||||||
"build_permission",
|
|
||||||
"build_reset_token",
|
|
||||||
"build_role",
|
|
||||||
"build_session",
|
|
||||||
"build_user",
|
|
||||||
# Read ops
|
# Read ops
|
||||||
# Write ops
|
# Write ops
|
||||||
"add_permission_to_org",
|
"add_permission_to_org",
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ _background_task: asyncio.Task | None = None
|
|||||||
|
|
||||||
async def flush() -> None:
|
async def flush() -> None:
|
||||||
"""Write all pending database changes to disk."""
|
"""Write all pending database changes to disk."""
|
||||||
store = _ops._store
|
store = _ops._db._store
|
||||||
if store is None:
|
if store is None:
|
||||||
_logger.warning("flush() called but _store is None")
|
_logger.warning("flush() called but _store is None")
|
||||||
return
|
return
|
||||||
@@ -48,6 +48,10 @@ async def _background_loop():
|
|||||||
cleanup_expired()
|
cleanup_expired()
|
||||||
await flush() # Flush cleanup changes
|
await flush() # Flush cleanup changes
|
||||||
last_cleanup = now
|
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:
|
except asyncio.CancelledError:
|
||||||
# Final flush before exit
|
# Final flush before exit
|
||||||
await flush()
|
await flush()
|
||||||
@@ -90,7 +94,7 @@ async def start_background():
|
|||||||
|
|
||||||
|
|
||||||
async def stop_background():
|
async def stop_background():
|
||||||
"""Stop the background task and flush any pending changes."""
|
"""Stop the background task, flush pending changes, and release the file lock."""
|
||||||
global _background_task
|
global _background_task
|
||||||
if _background_task:
|
if _background_task:
|
||||||
_background_task.cancel()
|
_background_task.cancel()
|
||||||
@@ -99,6 +103,7 @@ async def stop_background():
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
pass
|
pass
|
||||||
_background_task = None
|
_background_task = None
|
||||||
|
_ops._db._store.close()
|
||||||
|
|
||||||
|
|
||||||
# Aliases for backwards compatibility
|
# Aliases for backwards compatibility
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from datetime import UTC, datetime
|
|||||||
import uuid7
|
import uuid7
|
||||||
|
|
||||||
import paskia.db.operations as _ops
|
import paskia.db.operations as _ops
|
||||||
|
from paskia.authsession import reset_expires
|
||||||
from paskia.db.structs import Config, Org, Permission, ResetToken, Role, User
|
from paskia.db.structs import Config, Org, Permission, ResetToken, Role, User
|
||||||
from paskia.util.crypto import secret_key
|
from paskia.util.crypto import secret_key
|
||||||
|
|
||||||
@@ -59,8 +60,6 @@ def bootstrap(
|
|||||||
|
|
||||||
# Set reset token expiry (passphrase generated by ResetToken.create)
|
# Set reset token expiry (passphrase generated by ResetToken.create)
|
||||||
if reset_expiry is None:
|
if reset_expiry is None:
|
||||||
from paskia.authsession import reset_expires # noqa: PLC0415
|
|
||||||
|
|
||||||
reset_expiry = reset_expires()
|
reset_expiry = reset_expires()
|
||||||
|
|
||||||
with _ops._db.transaction("bootstrap"):
|
with _ops._db.transaction("bootstrap"):
|
||||||
|
|||||||
@@ -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.
|
JSONL persistence layer for the database.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import copy
|
import copy
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
@@ -13,126 +14,137 @@ from pathlib import Path
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import aiofiles
|
|
||||||
import jsondiff
|
import jsondiff
|
||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
|
from paskia.db.filelock import LockedFile
|
||||||
from paskia.db.logging import log_change
|
from paskia.db.logging import log_change
|
||||||
from paskia.db.migrations import DBVER, apply_all_migrations
|
from paskia.db.migrations import (
|
||||||
from paskia.db.structs import DB, SessionContext
|
DBVER,
|
||||||
|
MigrationCtx,
|
||||||
|
apply_all_migrations,
|
||||||
|
apply_migrations_readonly,
|
||||||
|
)
|
||||||
|
from paskia.db.snapshot import SnapshotState
|
||||||
|
from paskia.db.structs import DB, Config, SessionContext
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_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):
|
class DatabaseError(Exception):
|
||||||
"""A single change record in the JSONL file."""
|
"""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")
|
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)
|
u: str | None = None # user UUID who performed the action (None for system)
|
||||||
diff: dict = {}
|
diff: dict
|
||||||
|
|
||||||
|
|
||||||
# msgspec encoder for change records
|
|
||||||
_change_encoder = msgspec.json.Encoder()
|
|
||||||
|
|
||||||
|
|
||||||
def compute_diff(previous: dict, current: dict) -> dict | None:
|
def compute_diff(previous: dict, current: dict) -> dict | None:
|
||||||
"""Compute JSON diff between two states.
|
return jsondiff.diff(previous, current, marshal=True) or None
|
||||||
|
|
||||||
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,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# Actions that are allowed to create a new database file
|
# Actions that are allowed to create a new database file
|
||||||
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap"})
|
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap"})
|
||||||
|
|
||||||
# Flag to prevent duplicate error messages on fatal flush failure
|
|
||||||
_flush_failed = False
|
|
||||||
|
|
||||||
|
|
||||||
async def flush_changes(
|
|
||||||
db_path: Path,
|
|
||||||
pending_changes: deque[_ChangeRecord],
|
|
||||||
) -> None:
|
|
||||||
"""Write all pending changes to disk.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
db_path: Path to the JSONL database file
|
|
||||||
pending_changes: Queue of pending change records (will be cleared on success)
|
|
||||||
|
|
||||||
On failure, logs an error and sends SIGTERM to trigger graceful shutdown.
|
|
||||||
"""
|
|
||||||
global _flush_failed
|
|
||||||
if _flush_failed or not pending_changes:
|
|
||||||
return
|
|
||||||
|
|
||||||
if not db_path.exists():
|
|
||||||
first_action = pending_changes[0].a
|
|
||||||
if first_action not in _BOOTSTRAP_ACTIONS:
|
|
||||||
_logger.error(
|
|
||||||
"Refusing to create database file with action '%s' - "
|
|
||||||
"only bootstrap can create a new database",
|
|
||||||
first_action,
|
|
||||||
)
|
|
||||||
_flush_failed = True
|
|
||||||
os.kill(os.getpid(), signal.SIGTERM)
|
|
||||||
return
|
|
||||||
|
|
||||||
changes_to_write = list(pending_changes)
|
|
||||||
|
|
||||||
try:
|
|
||||||
lines = [_change_encoder.encode(change) for change in changes_to_write]
|
|
||||||
if not lines:
|
|
||||||
pending_changes.clear()
|
|
||||||
return
|
|
||||||
|
|
||||||
async with aiofiles.open(db_path, "ab") as f:
|
|
||||||
await f.write(b"\n".join(lines) + b"\n")
|
|
||||||
pending_changes.clear()
|
|
||||||
except OSError as e:
|
|
||||||
_logger.error("Failed to flush database: %s", e)
|
|
||||||
_flush_failed = True
|
|
||||||
os.kill(os.getpid(), signal.SIGTERM)
|
|
||||||
|
|
||||||
|
|
||||||
class JsonlStore:
|
class JsonlStore:
|
||||||
"""JSONL persistence layer for a DB instance."""
|
"""JSONL persistence layer for a DB instance."""
|
||||||
|
|
||||||
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: DB = db
|
||||||
self.db_path = Path(db_path)
|
self.db_path = Path(db_path)
|
||||||
self._previous_builtins: dict[str, Any] = {}
|
self._file = LockedFile()
|
||||||
self._pending_changes: deque[_ChangeRecord] = deque()
|
self._flush_failed = False
|
||||||
|
self._statedict: dict[str, Any] = {}
|
||||||
|
self._pending_changes: deque[ChangeRecord] = deque()
|
||||||
self._current_action: str = "system"
|
self._current_action: str = "system"
|
||||||
self._current_user: str | None = None
|
self._current_user: str | None = None
|
||||||
self._in_transaction: bool = False
|
self._in_transaction: bool = False
|
||||||
self._transaction_snapshot: dict[str, Any] | None = None
|
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(
|
async def load(
|
||||||
self, db_path: str | None = None, *, rp_id: str = "localhost"
|
self, db_path: str | None = None, *, rp_id: str = "localhost"
|
||||||
@@ -144,55 +156,52 @@ class JsonlStore:
|
|||||||
if not self.db_path.exists():
|
if not self.db_path.exists():
|
||||||
return
|
return
|
||||||
|
|
||||||
# Replay change log to reconstruct state
|
# Open with exclusive write lock and read contents — single threadpool call
|
||||||
data_dict: dict = {}
|
content = await asyncio.to_thread(self._file.open_and_read, self.db_path)
|
||||||
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}")
|
|
||||||
|
|
||||||
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
|
return
|
||||||
|
|
||||||
# Set previous state for diffing (will be updated by _queue_change)
|
# 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
|
# Callback to persist each migration
|
||||||
async def persist_migration(
|
async def persist_migration(
|
||||||
action: str, new_version: int, current: dict
|
action: str, new_version: int, current: dict
|
||||||
) -> None:
|
) -> None:
|
||||||
self._current_version = new_version
|
self._v = new_version
|
||||||
self._queue_change(action, new_version, current)
|
self._queue_change(action, new_version, current)
|
||||||
|
|
||||||
# Apply schema migrations one at a time
|
# Apply schema migrations one at a time
|
||||||
await apply_all_migrations(
|
await apply_all_migrations(
|
||||||
data_dict, self._current_version, persist_migration, rp_id=rp_id
|
statedict,
|
||||||
|
self._v,
|
||||||
|
persist_migration,
|
||||||
|
MigrationCtx(rp_id=rp_id),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Decode to msgspec struct
|
# Decode to msgspec struct
|
||||||
decoder = msgspec.json.Decoder(DB)
|
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
|
self.db._store = self
|
||||||
|
|
||||||
# Normalize via msgspec round-trip (handles omit_defaults etc.)
|
# Normalize via msgspec round-trip (handles omit_defaults etc.)
|
||||||
# This ensures _previous_builtins matches what msgspec would produce
|
# This ensures _previous_builtins matches what msgspec would produce
|
||||||
normalized_dict = msgspec.to_builtins(self.db)
|
normalized_dict = msgspec.to_builtins(self.db)
|
||||||
await persist_migration(
|
await persist_migration("migrate:msgspec", self._v, normalized_dict)
|
||||||
"migrate:msgspec", self._current_version, normalized_dict
|
|
||||||
)
|
|
||||||
|
|
||||||
def _queue_change(
|
def _queue_change(
|
||||||
self, action: str, version: int, current: dict, user: str | None = None
|
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
|
current: The current state as a plain dict
|
||||||
user: Optional user UUID who performed the action
|
user: Optional user UUID who performed the action
|
||||||
"""
|
"""
|
||||||
diff = compute_diff(self._previous_builtins, current)
|
diff = compute_diff(self._statedict, current)
|
||||||
if not diff:
|
if not diff:
|
||||||
return
|
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
|
# Log the change with user display name if available
|
||||||
user_display = None
|
user_display = None
|
||||||
@@ -220,8 +236,8 @@ class JsonlStore:
|
|||||||
except (ValueError, KeyError):
|
except (ValueError, KeyError):
|
||||||
user_display = user
|
user_display = user
|
||||||
|
|
||||||
log_change(action, diff, user_display, self._previous_builtins, self.db)
|
log_change(action, diff, user_display, self._statedict, self.db)
|
||||||
self._previous_builtins = copy.deepcopy(current)
|
self._statedict = copy.deepcopy(current)
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def transaction(
|
def transaction(
|
||||||
@@ -243,13 +259,13 @@ class JsonlStore:
|
|||||||
|
|
||||||
# Check for out-of-transaction modifications
|
# Check for out-of-transaction modifications
|
||||||
current_state = msgspec.to_builtins(self.db)
|
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
|
# Allow bootstrap to create a new database from empty state
|
||||||
is_bootstrap = action in _BOOTSTRAP_ACTIONS
|
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
|
pass # Expected: creating database from scratch
|
||||||
else:
|
else:
|
||||||
diff = compute_diff(self._previous_builtins, current_state)
|
diff = compute_diff(self._statedict, current_state)
|
||||||
diff_json = msgspec.json.encode(diff).decode()
|
diff_json = msgspec.json.encode(diff).decode()
|
||||||
_logger.critical(
|
_logger.critical(
|
||||||
"Database state modified outside of transaction! "
|
"Database state modified outside of transaction! "
|
||||||
@@ -270,7 +286,7 @@ class JsonlStore:
|
|||||||
yield
|
yield
|
||||||
current = msgspec.to_builtins(self.db)
|
current = msgspec.to_builtins(self.db)
|
||||||
self._queue_change(
|
self._queue_change(
|
||||||
self._current_action, self._current_version, current, self._current_user
|
self._current_action, self._v, current, self._current_user
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception:
|
||||||
# Rollback on error: restore from snapshot
|
# Rollback on error: restore from snapshot
|
||||||
@@ -289,5 +305,47 @@ class JsonlStore:
|
|||||||
self._transaction_snapshot = None
|
self._transaction_snapshot = None
|
||||||
|
|
||||||
async def flush(self) -> None:
|
async def flush(self) -> None:
|
||||||
"""Write all pending changes to disk."""
|
"""Write all pending changes to disk.
|
||||||
await flush_changes(self.db_path, self._pending_changes)
|
|
||||||
|
On failure, logs an error and sends SIGTERM to trigger graceful shutdown.
|
||||||
|
"""
|
||||||
|
if self._flush_failed or not self._pending_changes:
|
||||||
|
return
|
||||||
|
|
||||||
|
if not self._file.is_open:
|
||||||
|
first_action = self._pending_changes[0].a
|
||||||
|
if first_action not in _BOOTSTRAP_ACTIONS:
|
||||||
|
_logger.error(
|
||||||
|
"Refusing to create database file with action '%s' - "
|
||||||
|
"only bootstrap can create a new database",
|
||||||
|
first_action,
|
||||||
|
)
|
||||||
|
self._flush_failed = True
|
||||||
|
os.kill(os.getpid(), signal.SIGTERM)
|
||||||
|
return
|
||||||
|
# Bootstrap: create and open the file with lock
|
||||||
|
await asyncio.to_thread(self._file.open, self.db_path, create=True)
|
||||||
|
|
||||||
|
changes_to_write = list(self._pending_changes)
|
||||||
|
|
||||||
|
try:
|
||||||
|
lines = [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
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
import paskia.db.operations as _ops
|
import paskia.db.operations as _ops
|
||||||
|
from paskia import oidc_notify
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
|
from paskia.db.jsonl import JsonlStore
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_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."""
|
"""Load database from JSONL file."""
|
||||||
if _ops._initialized:
|
if _ops._db._store:
|
||||||
_logger.debug("Database already initialized, skipping reload")
|
_logger.debug("Database already initialized, skipping reload")
|
||||||
return
|
return
|
||||||
default_path = f"{rp_id}.paskiadb"
|
db_path = os.environ.get("PASKIA_DB", f"{rp_id}.paskiadb")
|
||||||
db_path = os.environ.get("PASKIA_DB", default_path)
|
store = JsonlStore(_ops._db, db_path)
|
||||||
await _ops._store.load(db_path, rp_id=rp_id)
|
await store.load(db_path, rp_id=rp_id)
|
||||||
_ops._db = _ops._store.db
|
_ops._db = store.db
|
||||||
_ops._initialized = True
|
_ops._db._store = store
|
||||||
|
# Request a snapshot after successful startup
|
||||||
|
store._snapshot.request_force()
|
||||||
|
|
||||||
|
|
||||||
def cleanup_expired() -> int:
|
def cleanup_expired() -> int:
|
||||||
@@ -31,8 +35,6 @@ def cleanup_expired() -> int:
|
|||||||
limit = now - EXPIRES
|
limit = now - EXPIRES
|
||||||
expired_sessions = [k for k, s in _ops._db.sessions.items() if s.validated < limit]
|
expired_sessions = [k for k, s in _ops._db.sessions.items() if s.validated < limit]
|
||||||
if expired_sessions:
|
if expired_sessions:
|
||||||
from paskia import oidc_notify # noqa: PLC0415
|
|
||||||
|
|
||||||
oidc_notify.schedule_notifications(expired_sessions)
|
oidc_notify.schedule_notifications(expired_sessions)
|
||||||
with _ops._db.transaction("expiry"):
|
with _ops._db.transaction("expiry"):
|
||||||
for k in expired_sessions:
|
for k in expired_sessions:
|
||||||
|
|||||||
+10
-10
@@ -35,9 +35,9 @@ _UNSAFE_CHARS = re.compile(
|
|||||||
|
|
||||||
# ANSI color codes (matching FastAPI logging style)
|
# ANSI color codes (matching FastAPI logging style)
|
||||||
_RESET = "\033[0m"
|
_RESET = "\033[0m"
|
||||||
_DIM = "\033[2m"
|
_SEP = "\033[38;5;242m" # Dark grey for separators (like host/timing in access log)
|
||||||
_PATH_PREFIX = "\033[1;30m" # Dark grey for path prefix (like host in access log)
|
_PATH_PREFIX = "\033[38;5;242m" # Dark grey for path prefix (like host in access log)
|
||||||
_PATH_FINAL = "\033[0m" # Default for final element (like path 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
|
_DELETE = "\033[1;31m" # Red for deletions
|
||||||
_ADD = "\033[0;32m" # Green for additions
|
_ADD = "\033[0;32m" # Green for additions
|
||||||
_ACTION = "\033[1;34m" # Bold blue for action name
|
_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
|
# Helper to format a value, checking for censored paths
|
||||||
def fmt_value(v: Any, child_path: list[str]) -> str:
|
def fmt_value(v: Any, child_path: list[str]) -> str:
|
||||||
if child_path[-2:] == ["oidc", "key"]:
|
if child_path[-2:] == ["oidc", "key"]:
|
||||||
return f"{_DIM}<hidden>{_RESET}"
|
return f"{_SEP}<hidden>{_RESET}"
|
||||||
return _format_value(v, resolver=resolver)
|
return _format_value(v, resolver=resolver)
|
||||||
|
|
||||||
# Helper to format path with UUID replacement
|
# Helper to format path with UUID replacement
|
||||||
@@ -342,12 +342,12 @@ def _format_change_lines(
|
|||||||
lines = []
|
lines = []
|
||||||
# First line: path with green final element and grey =
|
# First line: path with green final element and grey =
|
||||||
if len(formatted_path) == 1:
|
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:
|
else:
|
||||||
prefix = ".".join(formatted_path[:-1])
|
prefix = ".".join(formatted_path[:-1])
|
||||||
final = formatted_path[-1]
|
final = formatted_path[-1]
|
||||||
lines.append(
|
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
|
# Child lines: indented key: value, with aligned values
|
||||||
# Format keys (may contain UUIDs)
|
# Format keys (may contain UUIDs)
|
||||||
@@ -360,24 +360,24 @@ def _format_change_lines(
|
|||||||
field_width = max(max_key_len, 12) # minimum 12 chars
|
field_width = max(max_key_len, 12) # minimum 12 chars
|
||||||
for k_display, v_str in formatted_items:
|
for k_display, v_str in formatted_items:
|
||||||
padding = " " * (field_width - len(k_display))
|
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
|
return lines
|
||||||
else:
|
else:
|
||||||
value_str = fmt_value(value, path)
|
value_str = fmt_value(value, path)
|
||||||
if len(formatted_path) == 1:
|
if len(formatted_path) == 1:
|
||||||
return [
|
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])
|
prefix = ".".join(formatted_path[:-1])
|
||||||
final = formatted_path[-1]
|
final = formatted_path[-1]
|
||||||
return [
|
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
|
# update: Existing item being updated - normal path colors
|
||||||
value_str = fmt_value(value, path)
|
value_str = fmt_value(value, path)
|
||||||
path_str = _format_path(path, resolver=resolver)
|
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(
|
def format_diff(
|
||||||
|
|||||||
+30
-8
@@ -8,28 +8,36 @@ Each migration should be idempotent and only run when needed.
|
|||||||
import base64
|
import base64
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
|
|
||||||
|
import msgspec
|
||||||
|
|
||||||
from paskia.util.crypto import secret_key
|
from paskia.util.crypto import secret_key
|
||||||
|
|
||||||
|
|
||||||
def migrate_v1(d: dict, **kwargs) -> None:
|
class MigrationCtx(msgspec.Struct):
|
||||||
|
"""Context passed to each migration function."""
|
||||||
|
|
||||||
|
rp_id: str
|
||||||
|
|
||||||
|
|
||||||
|
def migrate_v1(d: dict, ctx: MigrationCtx) -> None:
|
||||||
"""Remove Org.created_at fields."""
|
"""Remove Org.created_at fields."""
|
||||||
for org_data in d["orgs"].values():
|
for org_data in d["orgs"].values():
|
||||||
org_data.pop("created_at", None)
|
org_data.pop("created_at", None)
|
||||||
|
|
||||||
|
|
||||||
def migrate_v2(d: dict, *, rp_id: str = "localhost") -> None:
|
def migrate_v2(d: dict, ctx: MigrationCtx) -> None:
|
||||||
"""Add config field if missing."""
|
"""Add config field if missing."""
|
||||||
if "config" not in d:
|
if "config" not in d:
|
||||||
d["config"] = {"rp_id": rp_id}
|
d["config"] = {"rp_id": ctx.rp_id}
|
||||||
|
|
||||||
|
|
||||||
def migrate_v3(d: dict, **kwargs) -> None:
|
def migrate_v3(d: dict, ctx: MigrationCtx) -> None:
|
||||||
"""Ensure all users have visits field."""
|
"""Ensure all users have visits field."""
|
||||||
for user_data in d["users"].values():
|
for user_data in d["users"].values():
|
||||||
user_data.setdefault("visits", 0)
|
user_data.setdefault("visits", 0)
|
||||||
|
|
||||||
|
|
||||||
def migrate_v4(d: dict, **kwargs) -> None:
|
def migrate_v4(d: dict, ctx: MigrationCtx) -> None:
|
||||||
"""OpenID Connect support and hardened session keys."""
|
"""OpenID Connect support and hardened session keys."""
|
||||||
# Session keys changed to hashes, drop old sessions
|
# Session keys changed to hashes, drop old sessions
|
||||||
d["sessions"] = {}
|
d["sessions"] = {}
|
||||||
@@ -45,14 +53,28 @@ migrations = sorted(
|
|||||||
DBVER = len(migrations) # Used by bootstrap to set initial version
|
DBVER = len(migrations) # Used by bootstrap to set initial version
|
||||||
|
|
||||||
|
|
||||||
|
def apply_migrations_readonly(
|
||||||
|
data_dict: dict,
|
||||||
|
current_version: int,
|
||||||
|
ctx: MigrationCtx,
|
||||||
|
) -> int:
|
||||||
|
"""Apply migration functions in-place without persistence.
|
||||||
|
|
||||||
|
Returns the new version after all migrations.
|
||||||
|
"""
|
||||||
|
while current_version < DBVER:
|
||||||
|
migrations[current_version](data_dict, ctx)
|
||||||
|
current_version += 1
|
||||||
|
return current_version
|
||||||
|
|
||||||
|
|
||||||
async def apply_all_migrations(
|
async def apply_all_migrations(
|
||||||
data_dict: dict,
|
data_dict: dict,
|
||||||
current_version: int,
|
current_version: int,
|
||||||
persist: Callable[[str, int, dict], Awaitable[None]],
|
persist: Callable[[str, int, dict], Awaitable[None]],
|
||||||
*,
|
ctx: MigrationCtx,
|
||||||
rp_id: str = "localhost",
|
|
||||||
) -> None:
|
) -> None:
|
||||||
while current_version < DBVER:
|
while current_version < DBVER:
|
||||||
migrations[current_version](data_dict, rp_id=rp_id)
|
migrations[current_version](data_dict, ctx)
|
||||||
current_version += 1
|
current_version += 1
|
||||||
await persist(f"migrate:v{current_version}", current_version, data_dict)
|
await persist(f"migrate:v{current_version}", current_version, data_dict)
|
||||||
|
|||||||
+3
-11
@@ -11,13 +11,10 @@ import secrets
|
|||||||
from datetime import UTC, datetime, timedelta
|
from datetime import UTC, datetime, timedelta
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
|
||||||
import uuid7
|
import uuid7
|
||||||
|
|
||||||
|
from paskia import oidc_notify
|
||||||
from paskia.config import SESSION_LIFETIME
|
from paskia.config import SESSION_LIFETIME
|
||||||
from paskia.db.jsonl import (
|
|
||||||
JsonlStore,
|
|
||||||
)
|
|
||||||
from paskia.db.structs import (
|
from paskia.db.structs import (
|
||||||
DB,
|
DB,
|
||||||
Client,
|
Client,
|
||||||
@@ -41,9 +38,6 @@ _UNSET = object()
|
|||||||
|
|
||||||
# Global database instance (empty until init() loads data)
|
# Global database instance (empty until init() loads data)
|
||||||
_db = DB(config=Config(rp_id="uninitialized.invalid"))
|
_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:
|
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:
|
if key not in _db.sessions:
|
||||||
raise ValueError("Session not found")
|
raise ValueError("Session not found")
|
||||||
from paskia import oidc_notify # noqa: PLC0415
|
|
||||||
|
|
||||||
oidc_notify.schedule_notifications([key])
|
oidc_notify.schedule_notifications([key])
|
||||||
with _db.transaction(action, ctx):
|
with _db.transaction(action, ctx):
|
||||||
@@ -503,7 +496,6 @@ def delete_sessions_for_user(
|
|||||||
user = _db.users.get(user_uuid)
|
user = _db.users.get(user_uuid)
|
||||||
if not user:
|
if not user:
|
||||||
return
|
return
|
||||||
from paskia import oidc_notify # noqa: PLC0415
|
|
||||||
|
|
||||||
keys = [s.key for s in user.sessions]
|
keys = [s.key for s in user.sessions]
|
||||||
oidc_notify.schedule_notifications(keys)
|
oidc_notify.schedule_notifications(keys)
|
||||||
@@ -589,7 +581,7 @@ def login(
|
|||||||
session = Session.create(
|
session = Session.create(
|
||||||
user=user_uuid,
|
user=user_uuid,
|
||||||
credential=credential_uuid,
|
credential=credential_uuid,
|
||||||
key=base64url.enc(hash_secret("cookie", token)),
|
key=hash_secret("cookie", token),
|
||||||
host=host,
|
host=host,
|
||||||
ip=ip,
|
ip=ip,
|
||||||
user_agent=user_agent,
|
user_agent=user_agent,
|
||||||
@@ -657,7 +649,7 @@ def create_credential_session(
|
|||||||
|
|
||||||
# Generate token and derive key
|
# Generate token and derive key
|
||||||
token = secrets.token_urlsafe(12)
|
token = secrets.token_urlsafe(12)
|
||||||
key = base64url.enc(hash_secret("cookie", token))
|
key = hash_secret("cookie", token)
|
||||||
|
|
||||||
session = Session.create(
|
session = Session.create(
|
||||||
user=user_uuid,
|
user=user_uuid,
|
||||||
|
|||||||
@@ -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 datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
|
||||||
import msgspec
|
import msgspec
|
||||||
import uuid7
|
import uuid7
|
||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.util import hostutil
|
|
||||||
from paskia.util import passphrase as passphrase_util
|
from paskia.util import passphrase as passphrase_util
|
||||||
from paskia.util.crypto import hash_secret
|
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.
|
"""Create a new Session with the provided key.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
key: The base64url-encoded hashed session key (derived from secret via hash_secret then base64url.enc)
|
key: The hashed session key (derived from secret via hash_secret)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Session object with key set
|
Session object with key set
|
||||||
@@ -471,7 +469,7 @@ class ResetToken(msgspec.Struct, dict=True):
|
|||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
if not hasattr(self, "key"):
|
if not hasattr(self, "key"):
|
||||||
self.key: bytes = b""
|
self.key: str = ""
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def user(self) -> User:
|
def user(self) -> User:
|
||||||
@@ -487,15 +485,15 @@ class ResetToken(msgspec.Struct, dict=True):
|
|||||||
del db.data().reset_tokens[self.key]
|
del db.data().reset_tokens[self.key]
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def hash(passphrase: str) -> bytes:
|
def hash(passphrase: str) -> str:
|
||||||
"""Hash a passphrase to bytes for reset token storage."""
|
"""Hash a passphrase to string for reset token storage."""
|
||||||
if not passphrase_util.is_well_formed(passphrase):
|
if not passphrase_util.is_well_formed(passphrase):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Trying to reset with a session token in place of a passphrase"
|
"Trying to reset with a session token in place of a passphrase"
|
||||||
if len(passphrase) == 16
|
if len(passphrase) == 16
|
||||||
else "Invalid passphrase format"
|
else "Invalid passphrase format"
|
||||||
)
|
)
|
||||||
return hashlib.sha512(passphrase.encode()).digest()[:9]
|
return hash_secret("reset", passphrase)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def by_passphrase(cls, passphrase: str) -> ResetToken | None:
|
def by_passphrase(cls, passphrase: str) -> ResetToken | None:
|
||||||
@@ -602,14 +600,14 @@ class OIDC(msgspec.Struct, dict=True):
|
|||||||
key: bytes | None = None
|
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."""
|
"""Stored configuration for the instance."""
|
||||||
|
|
||||||
rp_id: str
|
rp_id: str
|
||||||
rp_name: str | None = None
|
rp_name: str | None = None
|
||||||
origins: list[str] | None = None
|
|
||||||
auth_host: 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] = {}
|
users: dict[UUID, User] = {}
|
||||||
credentials: dict[UUID, Credential] = {}
|
credentials: dict[UUID, Credential] = {}
|
||||||
sessions: dict[str, Session] = {}
|
sessions: dict[str, Session] = {}
|
||||||
reset_tokens: dict[bytes, ResetToken] = {}
|
reset_tokens: dict[str, ResetToken] = {}
|
||||||
# OIDC provider data
|
# OIDC provider data
|
||||||
oidc: OIDC = msgspec.field(default_factory=lambda: OIDC())
|
oidc: OIDC = msgspec.field(default_factory=lambda: OIDC())
|
||||||
|
|
||||||
@@ -670,7 +668,7 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
|||||||
SessionContext if valid, None if session not found, expired, or host mismatch
|
SessionContext if valid, None if session not found, expired, or host mismatch
|
||||||
"""
|
"""
|
||||||
|
|
||||||
key = base64url.enc(hash_secret("cookie", session_secret))
|
key = hash_secret("cookie", session_secret)
|
||||||
try:
|
try:
|
||||||
s = self.sessions[key]
|
s = self.sessions[key]
|
||||||
except KeyError:
|
except KeyError:
|
||||||
@@ -680,10 +678,8 @@ class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
|||||||
if s.client_uuid is not None:
|
if s.client_uuid is not None:
|
||||||
return 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)
|
# Validate host matches (sessions are always created with a host)
|
||||||
|
normalized_input = host
|
||||||
if s.host != normalized_input:
|
if s.host != normalized_input:
|
||||||
# Session bound to different host
|
# Session bound to different host
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import logging
|
import logging
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import Body, FastAPI, HTTPException, Query, Request, Response
|
from fastapi import Body, FastAPI, HTTPException, Query, Request
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
from paskia import aaguid as aaguid_mod
|
from paskia import aaguid as aaguid_mod
|
||||||
@@ -14,6 +14,7 @@ from paskia.db import User as UserDC
|
|||||||
from paskia.db.operations import _UNSET
|
from paskia.db.operations import _UNSET
|
||||||
from paskia.db.structs import Client
|
from paskia.db.structs import Client
|
||||||
from paskia.fastapi import authz
|
from paskia.fastapi import authz
|
||||||
|
from paskia.fastapi.front import frontend
|
||||||
from paskia.fastapi.response import MsgspecResponse
|
from paskia.fastapi.response import MsgspecResponse
|
||||||
from paskia.fastapi.session import AUTH_COOKIE
|
from paskia.fastapi.session import AUTH_COOKIE
|
||||||
from paskia.globals import passkey
|
from paskia.globals import passkey
|
||||||
@@ -78,7 +79,7 @@ async def general_exception_handler(_request, exc: Exception): # pragma: no cov
|
|||||||
|
|
||||||
@app.get("/")
|
@app.get("/")
|
||||||
async def adminapp(request: Request, auth=AUTH_COOKIE):
|
async def adminapp(request: Request, auth=AUTH_COOKIE):
|
||||||
return Response(*await vitedev.read("/auth/admin/index.html"))
|
return await vitedev.handle(request, frontend, "/auth/admin/")
|
||||||
|
|
||||||
|
|
||||||
# -------------------- Organizations --------------------
|
# -------------------- Organizations --------------------
|
||||||
|
|||||||
+20
-24
@@ -15,12 +15,12 @@ from fastapi.security import HTTPBearer
|
|||||||
|
|
||||||
from paskia import authcode, db
|
from paskia import authcode, db
|
||||||
from paskia._version import __version__
|
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 import authz, session, user
|
||||||
from paskia.fastapi.response import MsgspecResponse
|
from paskia.fastapi.response import MsgspecResponse
|
||||||
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME, get_client_ip
|
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME, get_client_ip
|
||||||
from paskia.globals import passkey as global_passkey
|
from paskia.globals import passkey as global_passkey
|
||||||
from paskia.util import hostutil, htmlutil, passphrase, userinfo, vitedev
|
from paskia.util import hostutil, htmlutil, passphrase, userinfo
|
||||||
from paskia.util.apistructs import ApiSettings, ApiTokenInfo, ApiValidateResponse
|
from paskia.util.apistructs import ApiSettings, ApiTokenInfo, ApiValidateResponse
|
||||||
|
|
||||||
bearer_auth = HTTPBearer(auto_error=False)
|
bearer_auth = HTTPBearer(auto_error=False)
|
||||||
@@ -161,25 +161,15 @@ async def forward_authentication(
|
|||||||
# Clear cookie only if session is invalid (not for reauth)
|
# Clear cookie only if session is invalid (not for reauth)
|
||||||
if e.clear_session:
|
if e.clear_session:
|
||||||
session.clear_session_cookie(response)
|
session.clear_session_cookie(response)
|
||||||
|
# Browser request? - return full-page HTML with metadata patched into data attrs
|
||||||
# Check Accept header to decide response format
|
if "text/html" in request.headers.get("accept", ""):
|
||||||
accept = request.headers.get("accept", "")
|
return await htmlutil.patched_html_response(
|
||||||
wants_html = "text/html" in accept
|
request, "/int/forward/", e.status_code, mode=e.mode, **e.metadata
|
||||||
|
|
||||||
if wants_html:
|
|
||||||
# Browser request - return full-page HTML with metadata
|
|
||||||
data_attrs = {"mode": e.mode, **e.metadata}
|
|
||||||
html = (await vitedev.read("/int/forward/index.html"))[0]
|
|
||||||
html = htmlutil.patch_html_data_attrs(html, **data_attrs)
|
|
||||||
return Response(
|
|
||||||
html, status_code=e.status_code, media_type="text/html; charset=UTF-8"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
# API request - return JSON with iframe srcdoc HTML
|
|
||||||
return JSONResponse(
|
|
||||||
status_code=e.status_code,
|
|
||||||
content=await authz.auth_error_content(e),
|
|
||||||
)
|
)
|
||||||
|
# API request - return JSON with iframe srcdoc HTML
|
||||||
|
return JSONResponse(
|
||||||
|
status_code=e.status_code, content=await authz.auth_error_content(e)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@app.get("/settings")
|
@app.get("/settings")
|
||||||
@@ -212,9 +202,14 @@ async def api_user_info(
|
|||||||
detail="Authentication required",
|
detail="Authentication required",
|
||||||
mode="login",
|
mode="login",
|
||||||
)
|
)
|
||||||
ctx = db.data().session_ctx(auth, request.headers.get("host"))
|
ctx = session_ctx(auth, request.headers.get("host"))
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise HTTPException(401, "Session expired")
|
raise authz.AuthException(
|
||||||
|
status_code=401,
|
||||||
|
detail="Session expired",
|
||||||
|
mode="login",
|
||||||
|
clear_session=True,
|
||||||
|
)
|
||||||
|
|
||||||
return MsgspecResponse(
|
return MsgspecResponse(
|
||||||
await userinfo.build_user_info(
|
await userinfo.build_user_info(
|
||||||
@@ -244,6 +239,7 @@ async def token_info(credentials=Depends(bearer_auth)):
|
|||||||
ApiTokenInfo(
|
ApiTokenInfo(
|
||||||
token_type=reset_token.token_type,
|
token_type=reset_token.token_type,
|
||||||
display_name=u.display_name,
|
display_name=u.display_name,
|
||||||
|
theme=u.theme,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -253,7 +249,7 @@ async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
|||||||
if not auth:
|
if not auth:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.data().session_ctx(auth, host)
|
ctx = session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
with suppress(Exception):
|
with suppress(Exception):
|
||||||
@@ -286,7 +282,7 @@ async def api_set_session(
|
|||||||
secret = a.session_key
|
secret = a.session_key
|
||||||
|
|
||||||
# Verify the session exists
|
# Verify the session exists
|
||||||
ctx = db.data().session_ctx(secret, host)
|
ctx = session_ctx(secret, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise HTTPException(401, f"Session not found on {host}")
|
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
|
client: str, status: int, method: str, host: str, path: str, duration_ms: float
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Format access log line with colors and aligned fields."""
|
"""Format access log line with colors and aligned fields."""
|
||||||
use_color = sys.stderr.isatty()
|
|
||||||
|
|
||||||
# Format components with fixed widths for alignment
|
# Format components with fixed widths for alignment
|
||||||
ip = format_client_ip(client).ljust(19) # IPv6 network max 19 chars
|
ip = format_client_ip(client).ljust(19) # IPv6 network max 19 chars
|
||||||
timing = f"{duration_ms:.0f}ms"
|
timing = f"{duration_ms:.0f}ms"
|
||||||
method_padded = method.ljust(7) # Longest method is OPTIONS (7)
|
method_padded = method.ljust(7) # Longest method is OPTIONS (7)
|
||||||
|
|
||||||
if use_color:
|
status_str = f"{status_color(status)}{status}{_RESET}"
|
||||||
status_str = f"{status_color(status)}{status}{_RESET}"
|
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||||
timing_str = f"{_TIMING}{timing}{_RESET}"
|
method_str = f"{method_color(method)}{method_padded}{_RESET}"
|
||||||
method_str = f"{method_color(method)}{method_padded}{_RESET}"
|
host_str = f"{_HOST}{host}{_RESET}"
|
||||||
host_str = f"{_HOST}{host}{_RESET}"
|
path_str = f"{_PATH}{path}{_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
|
|
||||||
|
|
||||||
# Format: "IP STATUS METHOD host path TIMING"
|
# Format: "IP STATUS METHOD host path TIMING"
|
||||||
return f"{ip} {status_str} {method_str} {host_str}{path_str} {timing_str}"
|
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:
|
def log_ws_open(ws) -> int:
|
||||||
"""Log WebSocket connection open. Returns connection ID for use in close."""
|
"""Log WebSocket connection open. Returns connection ID for use in close."""
|
||||||
use_color = sys.stderr.isatty()
|
|
||||||
ws_id = _next_ws_id()
|
ws_id = _next_ws_id()
|
||||||
|
|
||||||
client = ws.client.host if ws.client else "-"
|
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
|
origin_host = origin.split("://", 1)[-1] if origin else None
|
||||||
show_origin = origin_host and origin_host != host
|
show_origin = origin_host and origin_host != host
|
||||||
|
|
||||||
if use_color:
|
# 🔌 aligned with status (takes ~2 char width), ID aligned with method
|
||||||
# 🔌 aligned with status (takes ~2 char width), ID aligned with method
|
prefix = f"🔌 {_WS_OPEN}{id_str}{_RESET}"
|
||||||
prefix = f"🔌 {_WS_OPEN}{id_str}{_RESET}"
|
host_str = f"{_HOST}{host}{_RESET}"
|
||||||
host_str = f"{_HOST}{host}{_RESET}"
|
path_str = f"{_PATH}{path}{_RESET}"
|
||||||
path_str = f"{_PATH}{path}{_RESET}"
|
origin_str = f" {_RESET}from {_HOST}{origin_host}{_RESET}" if show_origin else ""
|
||||||
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 ""
|
|
||||||
|
|
||||||
logger.info(f"{ip} {prefix} {host_str}{path_str}{origin_str}")
|
logger.info(f"{ip} {prefix} {host_str}{path_str}{origin_str}")
|
||||||
return ws_id
|
return ws_id
|
||||||
@@ -209,8 +191,6 @@ WS_CLOSE_CODES = {
|
|||||||
|
|
||||||
def log_ws_close(ws_id: int, close_code: int | None, duration: float) -> None:
|
def log_ws_close(ws_id: int, close_code: int | None, duration: float) -> None:
|
||||||
"""Log WebSocket connection close with duration and status."""
|
"""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)
|
id_str = f"{ws_id:02d}".ljust(7) # Align with method field (7 chars)
|
||||||
timing = f"{duration * 1000:.0f}ms"
|
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:
|
else:
|
||||||
status = WS_CLOSE_CODES.get(close_code, f"code {close_code}")
|
status = WS_CLOSE_CODES.get(close_code, f"code {close_code}")
|
||||||
|
|
||||||
if use_color:
|
# 🔌 aligned with status, ID aligned with method
|
||||||
# 🔌 aligned with status, ID aligned with method
|
prefix = f"🔌 {_WS_CLOSE}{id_str}{_RESET}"
|
||||||
prefix = f"🔌 {_WS_CLOSE}{id_str}{_RESET}"
|
status_str = f"{_WS_STATUS}{status}{_RESET}"
|
||||||
status_str = f"{_WS_STATUS}{status}{_RESET}"
|
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||||
timing_str = f"{_TIMING}{timing}{_RESET}"
|
|
||||||
else:
|
|
||||||
prefix = f"WS- {id_str}"
|
|
||||||
status_str = status
|
|
||||||
timing_str = timing
|
|
||||||
|
|
||||||
logger.info(f"{' ' * 19} {prefix} {status_str} {timing_str}")
|
logger.info(f"{' ' * 19} {prefix} {status_str} {timing_str}")
|
||||||
|
|
||||||
|
|||||||
+23
-22
@@ -1,21 +1,26 @@
|
|||||||
import json
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
import msgspec
|
||||||
from fastapi import FastAPI, HTTPException, Request, Response
|
from fastapi import FastAPI, HTTPException, Request, Response
|
||||||
from fastapi.responses import FileResponse, RedirectResponse
|
from fastapi.responses import FileResponse, RedirectResponse
|
||||||
from fastapi_vue import Frontend
|
|
||||||
|
|
||||||
from paskia import authcode, globals
|
from paskia import authcode, db, globals
|
||||||
from paskia.__main__ import DEVMODE
|
from paskia.__main__ import DEVMODE
|
||||||
|
from paskia.bootstrap import bootstrap_if_needed
|
||||||
from paskia.db import start_background, stop_background
|
from paskia.db import start_background, stop_background
|
||||||
|
from paskia.db.background import flush
|
||||||
from paskia.db.logging import configure_db_logging
|
from paskia.db.logging import configure_db_logging
|
||||||
from paskia.fastapi import admin, api, auth_host, oid, ws
|
from paskia.fastapi import admin, api, auth_host, oid, ws
|
||||||
|
|
||||||
|
# Import frontend instance
|
||||||
|
from paskia.fastapi.front import frontend
|
||||||
from paskia.fastapi.logging import AccessLogMiddleware, configure_access_logging
|
from paskia.fastapi.logging import AccessLogMiddleware, configure_access_logging
|
||||||
from paskia.fastapi.session import AUTH_COOKIE
|
from paskia.fastapi.session import AUTH_COOKIE
|
||||||
from paskia.util import hostutil, passphrase, vitedev
|
from paskia.util import hostutil, passphrase, vitedev
|
||||||
|
from paskia.util.runtime import RuntimeConfig
|
||||||
|
|
||||||
# Configure custom logging
|
# Configure custom logging
|
||||||
configure_access_logging()
|
configure_access_logging()
|
||||||
@@ -23,14 +28,6 @@ configure_db_logging()
|
|||||||
|
|
||||||
_access_logger = logging.getLogger("paskia.access")
|
_access_logger = logging.getLogger("paskia.access")
|
||||||
|
|
||||||
# Vue Frontend static files
|
|
||||||
frontend = Frontend(
|
|
||||||
Path(__file__).parent.parent / "frontend-build",
|
|
||||||
cached=["/auth/assets/"],
|
|
||||||
favicon="/paskia.webp",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# Path to examples/index.html when running from source tree
|
# Path to examples/index.html when running from source tree
|
||||||
_EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples"
|
_EXAMPLES_DIR = Path(__file__).parent.parent.parent / "examples"
|
||||||
|
|
||||||
@@ -43,14 +40,13 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
|||||||
so that uvicorn reload / multiprocess workers inherit the settings.
|
so that uvicorn reload / multiprocess workers inherit the settings.
|
||||||
All keys are guaranteed to exist; values are already normalized by __main__.py.
|
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:
|
try:
|
||||||
# CLI (__main__) performs bootstrap once; here we skip to avoid duplicate work
|
|
||||||
await globals.init(
|
await globals.init(
|
||||||
rp_id=config["rp_id"],
|
rp_id=runtime.config.rp_id,
|
||||||
rp_name=config["rp_name"],
|
rp_name=runtime.config.rp_name,
|
||||||
origins=config["origins"],
|
origins=runtime.config.origins,
|
||||||
bootstrap=False,
|
bootstrap=False,
|
||||||
)
|
)
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
@@ -58,6 +54,12 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
|||||||
# Re-raise to fail fast
|
# Re-raise to fail fast
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
# Bootstrap and persist config now that the full DB is loaded
|
||||||
|
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)
|
# Restore uvicorn info logging (suppressed during startup in dev mode)
|
||||||
# Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages
|
# Keep uvicorn.error at WARNING to suppress WebSocket "connection open/closed" messages
|
||||||
if app.debug:
|
if app.debug:
|
||||||
@@ -131,9 +133,9 @@ async def openid_configuration(request: Request):
|
|||||||
|
|
||||||
@app.get("/auth/restricted/iframe")
|
@app.get("/auth/restricted/iframe")
|
||||||
@app.get("/auth/restricted/oidc")
|
@app.get("/auth/restricted/oidc")
|
||||||
async def restricted_view():
|
async def restricted_view(request: Request):
|
||||||
"""Serve the restricted/authentication UI for iframe or OpenID Connect."""
|
"""Serve the restricted/authentication UI for iframe or OpenID Connect."""
|
||||||
return Response(*await vitedev.read("/auth/restricted/index.html"))
|
return await vitedev.handle(request, frontend, "/auth/restricted/")
|
||||||
|
|
||||||
|
|
||||||
# Navigable URLs are defined here. We support both / and /auth/ as the base path
|
# Navigable URLs are defined here. We support both / and /auth/ as the base path
|
||||||
@@ -148,7 +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.
|
The frontend handles mode detection (host mode vs full profile) based on settings.
|
||||||
Access control is handled via APIs.
|
Access control is handled via APIs.
|
||||||
"""
|
"""
|
||||||
return Response(*await vitedev.read("/auth/index.html"))
|
return await vitedev.handle(request, frontend, "/auth/")
|
||||||
|
|
||||||
|
|
||||||
@app.get("/admin", include_in_schema=False)
|
@app.get("/admin", include_in_schema=False)
|
||||||
@@ -180,14 +182,13 @@ async def examples_page():
|
|||||||
|
|
||||||
|
|
||||||
# Frontend static files - must be before /{token} catch-all routes
|
# Frontend static files - must be before /{token} catch-all routes
|
||||||
# (actual routes registered during lifespan after frontend.load())
|
|
||||||
frontend.route(app, "/")
|
frontend.route(app, "/")
|
||||||
|
|
||||||
|
|
||||||
# Note: this catch-all handler must be the last route defined
|
# Note: this catch-all handler must be the last route defined
|
||||||
@app.get("/{token}")
|
@app.get("/{token}")
|
||||||
@app.get("/auth/{token}")
|
@app.get("/auth/{token}")
|
||||||
async def token_link(token: str):
|
async def token_link(request: Request, token: str):
|
||||||
"""Serve the reset app for reset tokens (password reset / device addition).
|
"""Serve the reset app for reset tokens (password reset / device addition).
|
||||||
|
|
||||||
The frontend will validate the token via /auth/api/token-info.
|
The frontend will validate the token via /auth/api/token-info.
|
||||||
@@ -195,4 +196,4 @@ async def token_link(token: str):
|
|||||||
if not passphrase.is_well_formed(token):
|
if not passphrase.is_well_formed(token):
|
||||||
raise HTTPException(status_code=404)
|
raise HTTPException(status_code=404)
|
||||||
|
|
||||||
return Response(*await vitedev.read("/int/reset/index.html"))
|
return await vitedev.handle(request, frontend, "/int/reset/")
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ def _oidc_session_by_token(
|
|||||||
token: str, client_uuid: UUID | None = None
|
token: str, client_uuid: UUID | None = None
|
||||||
) -> Session | None:
|
) -> Session | None:
|
||||||
"""Look up an OIDC session by token (refresh token value)."""
|
"""Look up an OIDC session by token (refresh token value)."""
|
||||||
key = base64url.enc(hash_secret("oidc", token))
|
key = hash_secret("oidc", token)
|
||||||
s = db.data().sessions.get(key)
|
s = db.data().sessions.get(key)
|
||||||
if not s or s.client_uuid is None:
|
if not s or s.client_uuid is None:
|
||||||
return None
|
return None
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from paskia import db
|
|||||||
from paskia.authsession import (
|
from paskia.authsession import (
|
||||||
delete_credential,
|
delete_credential,
|
||||||
expires,
|
expires,
|
||||||
|
session_ctx,
|
||||||
)
|
)
|
||||||
from paskia.fastapi import authz, session
|
from paskia.fastapi import authz, session
|
||||||
from paskia.fastapi.response import MsgspecResponse
|
from paskia.fastapi.response import MsgspecResponse
|
||||||
@@ -45,7 +46,7 @@ async def user_update_display_name(
|
|||||||
status_code=401, detail="Authentication Required", mode="login"
|
status_code=401, detail="Authentication Required", mode="login"
|
||||||
)
|
)
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.data().session_ctx(auth, host)
|
ctx = session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
@@ -74,7 +75,7 @@ async def user_update_info(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Authentication Required", mode="login"
|
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:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
@@ -112,7 +113,7 @@ async def user_update_theme(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Authentication Required", mode="login"
|
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:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
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:
|
if not auth:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.data().session_ctx(auth, host)
|
ctx = session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
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"
|
status_code=401, detail="Authentication Required", mode="login"
|
||||||
)
|
)
|
||||||
host = request.headers.get("host")
|
host = request.headers.get("host")
|
||||||
ctx = db.data().session_ctx(auth, host)
|
ctx = session_ctx(auth, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
|
|||||||
@@ -3,12 +3,11 @@ from datetime import UTC, datetime
|
|||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
|
||||||
from fastapi import FastAPI, WebSocket
|
from fastapi import FastAPI, WebSocket
|
||||||
|
|
||||||
from paskia import authcode, db
|
from paskia import authcode, db
|
||||||
from paskia.authcode import CookieCode, OIDCCode
|
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.db.structs import Session
|
||||||
from paskia.fastapi import authz, remote
|
from paskia.fastapi import authz, remote
|
||||||
from paskia.fastapi.session import AUTH_COOKIE, infodict
|
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)
|
# If there's an existing session, restrict to that user's credentials (reauth)
|
||||||
session_user_uuid = None
|
session_user_uuid = None
|
||||||
if auth:
|
if auth:
|
||||||
existing_ctx = db.data().session_ctx(auth, host)
|
existing_ctx = session_ctx(auth, host)
|
||||||
if existing_ctx:
|
if existing_ctx:
|
||||||
session_user_uuid = existing_ctx.user.uuid
|
session_user_uuid = existing_ctx.user.uuid
|
||||||
|
|
||||||
@@ -218,7 +217,7 @@ async def websocket_authenticate(
|
|||||||
session = Session.create(
|
session = Session.create(
|
||||||
user=cred.user_uuid,
|
user=cred.user_uuid,
|
||||||
credential=cred.uuid,
|
credential=cred.uuid,
|
||||||
key=base64url.enc(hash_secret("oidc", token)),
|
key=hash_secret("oidc", token),
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
ip=metadata["ip"],
|
ip=metadata["ip"],
|
||||||
user_agent=metadata["user_agent"],
|
user_agent=metadata["user_agent"],
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from uuid import UUID
|
|||||||
from fastapi import WebSocket
|
from fastapi import WebSocket
|
||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
|
from paskia.authsession import session_ctx
|
||||||
from paskia.db import Credential, SessionContext
|
from paskia.db import Credential, SessionContext
|
||||||
from paskia.fastapi.session import infodict
|
from paskia.fastapi.session import infodict
|
||||||
from paskia.fastapi.wsutil import validate_origin
|
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
|
# Get credential IDs if restricting to a user's credentials
|
||||||
credential_ids = None
|
credential_ids = None
|
||||||
if auth:
|
if auth:
|
||||||
existing_ctx = db.data().session_ctx(auth, host)
|
existing_ctx = session_ctx(auth, host)
|
||||||
if existing_ctx:
|
if existing_ctx:
|
||||||
credential_ids = existing_ctx.user.credential_ids or None
|
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
|
# Fetch and return the full session context
|
||||||
ctx = db.data().session_ctx(secret, normalized_host)
|
ctx = session_ctx(secret, host)
|
||||||
if not ctx:
|
if not ctx:
|
||||||
raise ValueError("Failed to create session context")
|
raise ValueError("Failed to create session context")
|
||||||
return ctx, secret
|
return ctx, secret
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ import httpx
|
|||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.util import oidjwt
|
from paskia.util import oidjwt
|
||||||
from paskia.util.hostutil import _load_config
|
from paskia.util.runtime import _load_config
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -171,11 +171,12 @@ class ApiSettings(msgspec.Struct):
|
|||||||
version: str
|
version: str
|
||||||
|
|
||||||
|
|
||||||
class ApiTokenInfo(msgspec.Struct):
|
class ApiTokenInfo(msgspec.Struct, omit_defaults=True):
|
||||||
"""Token info response struct."""
|
"""Token info response struct."""
|
||||||
|
|
||||||
token_type: str
|
token_type: str
|
||||||
display_name: str
|
display_name: str
|
||||||
|
theme: str = ""
|
||||||
|
|
||||||
|
|
||||||
class ApiUuidResponse(msgspec.Struct):
|
class ApiUuidResponse(msgspec.Struct):
|
||||||
|
|||||||
@@ -1,17 +1,15 @@
|
|||||||
import hashlib
|
import hashlib
|
||||||
|
|
||||||
|
import base64url
|
||||||
from cryptography.hazmat.primitives import serialization
|
from cryptography.hazmat.primitives import serialization
|
||||||
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
|
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
|
||||||
|
|
||||||
|
|
||||||
def hash_secret(*data) -> bytes:
|
def hash_secret(*data: str | bytes, length=12) -> str:
|
||||||
"""A custom HMAC that securily combines and hashes the given data (context, secrets). The first argument should be a namespacing string."""
|
"""A custom HMAC that securily combines and hashes the given data. The first argument should be a namespacing string."""
|
||||||
inner = bytearray(len(data).to_bytes(8, "big"))
|
p = [d.encode() if hasattr(d, "encode") else d for d in data]
|
||||||
for d in data:
|
p += [len(x).to_bytes(8, "little") for x in [p, *p]]
|
||||||
if isinstance(d, str):
|
return base64url.enc(hashlib.sha256(b"".join(p)).digest()[:length])
|
||||||
d = d.encode()
|
|
||||||
inner += hashlib.sha256(d).digest()
|
|
||||||
return hashlib.sha256(inner).digest()[:12]
|
|
||||||
|
|
||||||
|
|
||||||
def secret_key() -> bytes:
|
def secret_key() -> bytes:
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ __all__ = ["path", "file", "read", "is_dev_mode"]
|
|||||||
|
|
||||||
def _get_dev_server() -> str | None:
|
def _get_dev_server() -> str | None:
|
||||||
"""Get the dev server URL from environment, or None if not in dev mode."""
|
"""Get the dev server URL from environment, or None if not in dev mode."""
|
||||||
return os.environ.get("FASTAPI_VUE_FRONTEND_URL") or None
|
return os.environ.get("PASKIA_VITE_URL") or None
|
||||||
|
|
||||||
|
|
||||||
def _resolve_static_dir() -> Path:
|
def _resolve_static_dir() -> Path:
|
||||||
|
|||||||
+26
-16
@@ -1,27 +1,23 @@
|
|||||||
"""Utilities for determining the auth UI host and base URLs."""
|
"""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 urllib.parse import urlparse, urlsplit
|
||||||
|
|
||||||
|
from paskia.util.runtime import _load_config
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
|
||||||
def _load_config() -> dict:
|
def _cfg():
|
||||||
"""Load PASKIA_CONFIG JSON."""
|
return _load_config()
|
||||||
config_json = os.getenv("PASKIA_CONFIG")
|
|
||||||
if not config_json:
|
|
||||||
return {}
|
|
||||||
return json.loads(config_json)
|
|
||||||
|
|
||||||
|
|
||||||
def is_root_mode() -> bool:
|
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:
|
def dedicated_auth_host() -> str | None:
|
||||||
"""Return configured auth_host netloc, or 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:
|
if not auth_host:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
@@ -35,8 +31,10 @@ def ui_base_path() -> str:
|
|||||||
|
|
||||||
def auth_site_url() -> str:
|
def auth_site_url() -> str:
|
||||||
"""Return the base URL for the auth site UI (computed at startup)."""
|
"""Return the base URL for the auth site UI (computed at startup)."""
|
||||||
cfg = _load_config()
|
cfg = _cfg()
|
||||||
return cfg.get("site_url", "https://localhost") + cfg.get("site_path", "/auth/")
|
if cfg:
|
||||||
|
return cfg.site_url + cfg.site_path
|
||||||
|
return "https://localhost/auth/"
|
||||||
|
|
||||||
|
|
||||||
def reset_link_url(token: str) -> str:
|
def reset_link_url(token: str) -> str:
|
||||||
@@ -45,10 +43,10 @@ def reset_link_url(token: str) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def normalize_origin(origin: 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:
|
if "://" not in origin:
|
||||||
return f"https://{origin}"
|
return f"https://{origin}"
|
||||||
return origin
|
return origin.rstrip("/")
|
||||||
|
|
||||||
|
|
||||||
def reload_config() -> None:
|
def reload_config() -> None:
|
||||||
@@ -74,3 +72,15 @@ def normalize_host(raw_host: str | None) -> str | None:
|
|||||||
# Strip port from host:port
|
# Strip port from host:port
|
||||||
netloc = netloc.rsplit(":", 1)[0]
|
netloc = netloc.rsplit(":", 1)[0]
|
||||||
return netloc.lower() or None
|
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
|
import re
|
||||||
|
|
||||||
|
from paskia.fastapi.front import frontend
|
||||||
|
from paskia.util import vitedev
|
||||||
|
|
||||||
|
|
||||||
|
async def patched_html_response(request, filepath: str, status_code: int, **data_attrs):
|
||||||
|
"""Fetch HTML from vitedev and patch with data attributes.
|
||||||
|
|
||||||
|
Strips caching/compression headers from request to get raw content,
|
||||||
|
patches the HTML body with data attributes, and strips caching headers
|
||||||
|
from response.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
request: The FastAPI Request object
|
||||||
|
filepath: Path to HTML file, e.g. "/int/forward/"
|
||||||
|
status_code: HTTP status code for the response
|
||||||
|
**data_attrs: Key-value pairs for data attributes
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Patched Response object, or original response if not 200.
|
||||||
|
"""
|
||||||
|
resp = await vitedev.handle(request, frontend, filepath)
|
||||||
|
# Pass through non-200 responses
|
||||||
|
if resp.status_code != 200:
|
||||||
|
return resp
|
||||||
|
# Patch HTML with data attrs and strip caching headers from response
|
||||||
|
resp.body = patch_html_data_attrs(resp.body, **data_attrs)
|
||||||
|
resp.status_code = status_code
|
||||||
|
strip_headers = {b"etag", b"last-modified", b"content-length"}
|
||||||
|
resp.raw_headers = [
|
||||||
|
(k, v) for k, v in resp.raw_headers if k.lower() not in strip_headers
|
||||||
|
]
|
||||||
|
return resp
|
||||||
|
|
||||||
|
|
||||||
def patch_html_data_attrs(html: bytes, **data_attrs: str) -> bytes:
|
def patch_html_data_attrs(html: bytes, **data_attrs: str) -> bytes:
|
||||||
"""Patch HTML by adding data attributes to the <html> tag.
|
"""Patch HTML by adding data attributes to the <html> tag.
|
||||||
|
|||||||
@@ -3,7 +3,7 @@
|
|||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from fnmatch import fnmatchcase
|
from fnmatch import fnmatchcase
|
||||||
|
|
||||||
from paskia import db
|
from paskia.authsession import session_ctx
|
||||||
from paskia.util.hostutil import normalize_host
|
from paskia.util.hostutil import normalize_host
|
||||||
|
|
||||||
__all__ = ["has_any", "has_all", "session_context"]
|
__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:
|
if not auth:
|
||||||
return None
|
return None
|
||||||
normalized_host = normalize_host(host) if host else 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."""
|
"""Startup configuration box formatting utilities."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from sys import stderr
|
from sys import stderr
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
|
from fastapi_vue.hostutil import parse_endpoints
|
||||||
|
|
||||||
from paskia._version import __version__
|
from paskia._version import __version__
|
||||||
|
from paskia.util.hostutil import format_endpoint
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from paskia.config import PaskiaConfig
|
from paskia.util.runtime import RuntimeConfig
|
||||||
|
|
||||||
BOX_WIDTH = 60 # Inner width (excluding box chars)
|
BOX_WIDTH = 60 # Inner width (excluding box chars)
|
||||||
|
|
||||||
@@ -42,7 +47,7 @@ def bottom() -> str:
|
|||||||
return "┗" + "━" * (BOX_WIDTH + 2) + "┛\n"
|
return "┗" + "━" * (BOX_WIDTH + 2) + "┛\n"
|
||||||
|
|
||||||
|
|
||||||
def print_startup_config(config: "PaskiaConfig") -> None:
|
def print_startup_config(runtime: RuntimeConfig) -> None:
|
||||||
"""Print server configuration on startup."""
|
"""Print server configuration on startup."""
|
||||||
# Key graphic with yellow shading (bright for highlights, dark for body)
|
# Key graphic with yellow shading (bright for highlights, dark for body)
|
||||||
y = YELLOW # Dark yellow for main body
|
y = YELLOW # Dark yellow for main body
|
||||||
@@ -57,41 +62,40 @@ def print_startup_config(config: "PaskiaConfig") -> None:
|
|||||||
lines.append(
|
lines.append(
|
||||||
line(
|
line(
|
||||||
f"{b}█{y} {b}█{y}▀▀▀▀{b}█{y}▀▀{b}█{y}▀▀{b}█{r} {w}"
|
f"{b}█{y} {b}█{y}▀▀▀▀{b}█{y}▀▀{b}█{y}▀▀{b}█{r} {w}"
|
||||||
+ config.site_url
|
+ runtime.site_url
|
||||||
+ config.site_path
|
+ runtime.site_path
|
||||||
+ r
|
+ r
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
lines.append(line(f" {y}▀▀▀▀▀{r}"))
|
lines.append(line(f" {y}▀▀▀▀▀{r}"))
|
||||||
|
|
||||||
# Format auth host section
|
# Format auth host section
|
||||||
if config.auth_host:
|
if runtime.config.auth_host:
|
||||||
lines.append(line(f"Auth Host: {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
|
# Show frontend URL if in dev mode
|
||||||
devmode = os.environ.get("FASTAPI_VUE_FRONTEND_URL")
|
if DEVMODE:
|
||||||
if devmode:
|
lines.append(line(f"Dev Frontend: {os.environ.get('PASKIA_VITE_URL')}"))
|
||||||
lines.append(line(f"Dev Frontend: {devmode}"))
|
|
||||||
|
|
||||||
# Format listen address with scheme
|
# Format listen endpoints (dev mode only uses the first endpoint)
|
||||||
if config.uds:
|
|
||||||
listen = f"unix:{config.uds}"
|
endpoints = list(parse_endpoints(runtime.config.listen, P))
|
||||||
elif config.host:
|
if DEVMODE:
|
||||||
listen = f"http://{config.host}:{config.port}"
|
endpoints = endpoints[:1] # server.run reload=True uses only one
|
||||||
else:
|
parts = [format_endpoint(ep) for ep in endpoints]
|
||||||
listen = f"http://0.0.0.0:{config.port} + [::]:{config.port}"
|
lines.append(line(f"Backend: {' '.join(parts)}"))
|
||||||
lines.append(line(f"Backend: {listen}"))
|
|
||||||
|
|
||||||
# Relying Party line (omit name if same as id)
|
# Relying Party line (omit name if same as id)
|
||||||
rp_id = config.rp_id
|
rp_id = runtime.config.rp_id
|
||||||
rp_name = config.rp_name
|
rp_name = runtime.config.rp_name
|
||||||
if rp_name and rp_name != rp_id:
|
suffix = f" ({rp_name})" if rp_name and rp_name != rp_id else ""
|
||||||
lines.append(line(f"Relying Party: {rp_id} ({rp_name})"))
|
lines.append(line(f"Relying Party: {rp_id}{suffix}"))
|
||||||
else:
|
|
||||||
lines.append(line(f"Relying Party: {rp_id}"))
|
|
||||||
|
|
||||||
# Format origins section
|
# Format origins section
|
||||||
allowed = config.origins
|
allowed = runtime.config.origins
|
||||||
if allowed:
|
if allowed:
|
||||||
lines.append(line("Permitted Origins:"))
|
lines.append(line("Permitted Origins:"))
|
||||||
for origin in sorted(allowed):
|
for origin in sorted(allowed):
|
||||||
|
|||||||
@@ -5,8 +5,10 @@ from paskia.db import SessionContext
|
|||||||
from paskia.util import hostutil
|
from paskia.util import hostutil
|
||||||
from paskia.util.apistructs import (
|
from paskia.util.apistructs import (
|
||||||
ApiAaguidInfo,
|
ApiAaguidInfo,
|
||||||
|
ApiOrg,
|
||||||
ApiOrgContext,
|
ApiOrgContext,
|
||||||
ApiPermission,
|
ApiPermission,
|
||||||
|
ApiRole,
|
||||||
ApiRoleContext,
|
ApiRoleContext,
|
||||||
ApiSessionContext,
|
ApiSessionContext,
|
||||||
ApiUser,
|
ApiUser,
|
||||||
@@ -64,4 +66,6 @@ async def build_user_info(
|
|||||||
permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions}
|
permissions={p.uuid: ApiPermission.from_db(p) for p in ctx.permissions}
|
||||||
if ctx
|
if ctx
|
||||||
else {},
|
else {},
|
||||||
|
org=ApiOrg.from_db(ctx.org) if ctx else None,
|
||||||
|
role=ApiRole.from_db(ctx.role) if ctx else None,
|
||||||
)
|
)
|
||||||
|
|||||||
+35
-31
@@ -1,39 +1,29 @@
|
|||||||
"""Vite dev server proxy for fetching frontend files during development.
|
"""Vite dev server proxy for fetching frontend files during development.
|
||||||
|
|
||||||
In dev mode (FASTAPI_VUE_FRONTEND_URL set), fetches files from Vite.
|
In dev mode (PASKIA_VITE_URL set), fetches files from Vite.
|
||||||
In production, reads from the static build directory.
|
In production, reads from the static build directory.
|
||||||
|
|
||||||
This complements fastapi_vue.Frontend which handles static file serving
|
This complements fastapi_vue.Frontend which handles static file serving
|
||||||
but doesn't provide server-side fetching of HTML content.
|
but doesn't provide server-side fetching of HTML content.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import os
|
import os
|
||||||
from importlib import resources
|
from importlib import resources
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
from fastapi import Response
|
||||||
|
|
||||||
__all__ = ["read"]
|
__all__ = ["handle"]
|
||||||
|
|
||||||
|
|
||||||
def _get_dev_server() -> str | None:
|
|
||||||
"""Get the dev server URL from environment, or None if not in dev mode."""
|
|
||||||
return os.environ.get("FASTAPI_VUE_FRONTEND_URL") or None
|
|
||||||
|
|
||||||
|
|
||||||
def _resolve_static_dir() -> Path:
|
def _resolve_static_dir() -> Path:
|
||||||
"""Resolve the static files directory."""
|
|
||||||
|
|
||||||
# Try packaged path via importlib.resources (works for wheel/installed).
|
# Try packaged path via importlib.resources (works for wheel/installed).
|
||||||
try: # pragma: no cover - trivial path resolution
|
pkg_dir = resources.files("paskia") / "frontend-build"
|
||||||
pkg_dir = resources.files("paskia") / "frontend-build"
|
fs_path = Path(str(pkg_dir))
|
||||||
fs_path = Path(str(pkg_dir))
|
if fs_path.is_dir():
|
||||||
if fs_path.is_dir():
|
return fs_path
|
||||||
return fs_path
|
|
||||||
except Exception: # pragma: no cover - defensive
|
|
||||||
pass
|
|
||||||
# Fallback for editable/development before build.
|
# Fallback for editable/development before build.
|
||||||
return Path(__file__).parent.parent / "frontend-build"
|
return Path(__file__).parent.parent / "frontend-build"
|
||||||
|
|
||||||
@@ -41,31 +31,45 @@ def _resolve_static_dir() -> Path:
|
|||||||
_static_dir: Path = _resolve_static_dir()
|
_static_dir: Path = _resolve_static_dir()
|
||||||
|
|
||||||
|
|
||||||
async def read(filepath: str) -> tuple[bytes, int, dict[str, str]]:
|
async def handle(request, frontend, filepath: str):
|
||||||
"""Read file content and return response tuple.
|
"""Read file content and return Response.
|
||||||
|
|
||||||
In dev mode, fetches from the Vite dev server.
|
In dev mode, fetches from the Vite dev server.
|
||||||
In production, reads from the static build directory.
|
In production, uses frontend.handle.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
|
request: The FastAPI Request object
|
||||||
|
frontend: The fastapi_vue.Frontend instance
|
||||||
filepath: Path relative to frontend root, e.g. "/auth/index.html"
|
filepath: Path relative to frontend root, e.g. "/auth/index.html"
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (content, status_code, headers) suitable for
|
FastAPI Response object.
|
||||||
FastAPI Response(*args).
|
|
||||||
"""
|
"""
|
||||||
dev_server = _get_dev_server()
|
if dev_server := os.environ.get("PASKIA_VITE_URL"):
|
||||||
if dev_server:
|
|
||||||
async with httpx.AsyncClient() as client:
|
async with httpx.AsyncClient() as client:
|
||||||
resp = await client.get(f"{dev_server}{filepath}")
|
resp = await client.get(f"{dev_server}{filepath}")
|
||||||
resp.raise_for_status()
|
resp.raise_for_status()
|
||||||
mime = resp.headers.get("content-type", "application/octet-stream")
|
mime = resp.headers.get("content-type", "application/octet-stream")
|
||||||
# Strip charset suffix if present
|
# Strip charset suffix if present
|
||||||
mime = mime.split(";")[0].strip()
|
mime = mime.split(";")[0].strip()
|
||||||
return resp.content, resp.status_code, {"content-type": mime}
|
return Response(resp.content, resp.status_code, {"content-type": mime})
|
||||||
else:
|
|
||||||
# Production: read from static build
|
# Read from frontend cache directly to bypass any compression/processing
|
||||||
file_path = _static_dir / filepath.lstrip("/")
|
cached_content = getattr(frontend, "_files", {}).get(filepath)
|
||||||
content = await asyncio.to_thread(file_path.read_bytes)
|
if cached_content is not None:
|
||||||
mime, _ = mimetypes.guess_type(str(file_path))
|
mime, _ = mimetypes.guess_type(filepath)
|
||||||
return content, 200, {"content-type": mime or "application/octet-stream"}
|
return Response(
|
||||||
|
cached_content, 200, {"content-type": mime or "application/octet-stream"}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Fallback to frontend.handle for cache negotiation
|
||||||
|
# Strip accept-encoding to get uncompressed content (needed for HTML patching)
|
||||||
|
strip_headers = {b"accept-encoding", b"if-none-match", b"if-modified-since"}
|
||||||
|
request.scope["headers"] = [
|
||||||
|
(k, v) for k, v in request.scope["headers"] if k.lower() not in strip_headers
|
||||||
|
]
|
||||||
|
# Invalidate cached Headers object (it doesn't re-read scope after first access)
|
||||||
|
if hasattr(request, "_headers"):
|
||||||
|
del request._headers
|
||||||
|
|
||||||
|
return frontend.handle(request, filepath)
|
||||||
|
|||||||
+17
-28
@@ -11,20 +11,28 @@ keywords = [ "forward_auth", "auth_request", "FastAPI" ]
|
|||||||
authors = [
|
authors = [
|
||||||
{name = "Leo Vasanko"},
|
{name = "Leo Vasanko"},
|
||||||
]
|
]
|
||||||
|
requires-python = ">=3.11"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"fastapi[standard]>=0.104.1",
|
"fastapi[standard]>=0.129.0",
|
||||||
"websockets>=12.0",
|
"websockets>=16.0",
|
||||||
"webauthn>=1.11.1",
|
"webauthn>=2.7.1",
|
||||||
"base64url>=1.0.0",
|
"base64url>=1.1.1",
|
||||||
"uuid7-standard>=1.0.0",
|
"uuid7-standard>=1.1.0",
|
||||||
"pyjwt[crypto]>=2.8.0",
|
"pyjwt[crypto]>=2.11.0",
|
||||||
"jsondiff>=2.2.1",
|
"jsondiff>=2.2.1",
|
||||||
"msgspec>=0.20.0",
|
"msgspec>=0.20.0",
|
||||||
"aiofiles>=25.1.0",
|
"fastapi-vue>=1.1.0",
|
||||||
"fastapi-vue>=0.3.0",
|
|
||||||
"ua-parser[regex]>=1.0.1",
|
"ua-parser[regex]>=1.0.1",
|
||||||
]
|
]
|
||||||
requires-python = ">=3.11"
|
[dependency-groups]
|
||||||
|
dev = [
|
||||||
|
"coverage>=7.13.4",
|
||||||
|
"httpx>=0.28.1",
|
||||||
|
"pytest>=9.0.2",
|
||||||
|
"pytest-asyncio>=1.3.0",
|
||||||
|
"pytest-cov>=7.0.0",
|
||||||
|
"ruff>=0.15.1",
|
||||||
|
]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
Homepage = "https://git.zi.fi/LeoVasanko/paskia"
|
Homepage = "https://git.zi.fi/LeoVasanko/paskia"
|
||||||
@@ -36,15 +44,6 @@ source = "vcs"
|
|||||||
[tool.hatch.build.hooks.vcs]
|
[tool.hatch.build.hooks.vcs]
|
||||||
version-file = "paskia/_version.py"
|
version-file = "paskia/_version.py"
|
||||||
|
|
||||||
[project.optional-dependencies]
|
|
||||||
dev = [
|
|
||||||
"ruff>=0.1.0",
|
|
||||||
"coverage[toml]>=7.0.0",
|
|
||||||
"pytest>=8.0.0",
|
|
||||||
"pytest-asyncio>=0.24.0",
|
|
||||||
"httpx>=0.27.0",
|
|
||||||
]
|
|
||||||
|
|
||||||
[tool.coverage.run]
|
[tool.coverage.run]
|
||||||
source = ["paskia"]
|
source = ["paskia"]
|
||||||
branch = true
|
branch = true
|
||||||
@@ -75,16 +74,6 @@ select = ["E", "F", "I", "N", "W", "UP", "PLC0415"]
|
|||||||
ignore = ["E501"] # Line too long
|
ignore = ["E501"] # Line too long
|
||||||
isort.known-first-party = ["paskia"]
|
isort.known-first-party = ["paskia"]
|
||||||
|
|
||||||
[dependency-groups]
|
|
||||||
dev = [
|
|
||||||
"coverage>=7.12.0",
|
|
||||||
"httpx>=0.28.1",
|
|
||||||
"pytest>=9.0.1",
|
|
||||||
"pytest-asyncio>=1.3.0",
|
|
||||||
"pytest-cov>=7.0.0",
|
|
||||||
"ruff>=0.14.8",
|
|
||||||
]
|
|
||||||
|
|
||||||
[project.scripts]
|
[project.scripts]
|
||||||
paskia = "paskia.__main__:main"
|
paskia = "paskia.__main__:main"
|
||||||
|
|
||||||
|
|||||||
+16
-20
@@ -153,29 +153,9 @@ async def run_devserver(args: argparse.Namespace, remaining: list[str]) -> None:
|
|||||||
paskia.extend(["--origin", origin])
|
paskia.extend(["--origin", origin])
|
||||||
paskia.extend(remaining)
|
paskia.extend(remaining)
|
||||||
|
|
||||||
# Compute origins for Caddy
|
|
||||||
caddy_origins = []
|
|
||||||
if args.auth_host:
|
|
||||||
auth_host = args.auth_host
|
|
||||||
if "://" not in auth_host:
|
|
||||||
auth_host = f"https://{auth_host}"
|
|
||||||
caddy_origins.append(auth_host)
|
|
||||||
caddy_origins.append(f"https://{args.rp_id}")
|
|
||||||
if args.origins:
|
|
||||||
for origin in args.origins:
|
|
||||||
if "://" not in origin:
|
|
||||||
origin = f"https://{origin}"
|
|
||||||
caddy_origins.append(origin)
|
|
||||||
if not args.auth_host and not args.origins:
|
|
||||||
caddy_origins.append(f"https://{args.rp_id}")
|
|
||||||
# Remove duplicates while preserving order
|
|
||||||
seen = set()
|
|
||||||
caddy_origins = [x for x in caddy_origins if not (x in seen or seen.add(x))]
|
|
||||||
|
|
||||||
# Set environment for subprocesses
|
# Set environment for subprocesses
|
||||||
os.environ["PASKIA_VITE_URL"] = viteurl
|
os.environ["PASKIA_VITE_URL"] = viteurl
|
||||||
os.environ["PASKIA_BACKEND_URL"] = backurl
|
os.environ["PASKIA_BACKEND_URL"] = backurl
|
||||||
os.environ["PASKIA_SITE_URL"] = caddy_origins[0] if args.caddy else viteurl
|
|
||||||
os.environ["PASKIA_DEV"] = "1"
|
os.environ["PASKIA_DEV"] = "1"
|
||||||
if args.auth_host:
|
if args.auth_host:
|
||||||
os.environ["PASKIA_AUTH_HOST"] = args.auth_host
|
os.environ["PASKIA_AUTH_HOST"] = args.auth_host
|
||||||
@@ -183,6 +163,22 @@ async def run_devserver(args: argparse.Namespace, remaining: list[str]) -> None:
|
|||||||
async with ProcessGroup() as pg:
|
async with ProcessGroup() as pg:
|
||||||
# Start Caddy first if requested (needs to bind ports)
|
# Start Caddy first if requested (needs to bind ports)
|
||||||
if args.caddy:
|
if args.caddy:
|
||||||
|
caddy_origins = []
|
||||||
|
if args.auth_host:
|
||||||
|
auth_host = args.auth_host
|
||||||
|
if "://" not in auth_host:
|
||||||
|
auth_host = f"https://{auth_host}"
|
||||||
|
caddy_origins.append(auth_host)
|
||||||
|
caddy_origins.append(f"https://{args.rp_id}")
|
||||||
|
if args.origins:
|
||||||
|
for origin in args.origins:
|
||||||
|
if "://" not in origin:
|
||||||
|
origin = f"https://{origin}"
|
||||||
|
caddy_origins.append(origin)
|
||||||
|
if not caddy_origins:
|
||||||
|
caddy_origins.append(f"https://{args.rp_id}")
|
||||||
|
seen: set = set()
|
||||||
|
caddy_origins = [x for x in caddy_origins if not (x in seen or seen.add(x))]
|
||||||
caddy_proc = await run_caddy(caddy_origins, viteurl, backurl)
|
caddy_proc = await run_caddy(caddy_origins, viteurl, backurl)
|
||||||
pg._procs.append(caddy_proc)
|
pg._procs.append(caddy_proc)
|
||||||
pg._cmds[caddy_proc.pid] = "caddy"
|
pg._cmds[caddy_proc.pid] = "caddy"
|
||||||
|
|||||||
+1
-2
@@ -19,7 +19,6 @@ from collections.abc import AsyncGenerator
|
|||||||
from datetime import UTC, datetime, timedelta
|
from datetime import UTC, datetime, timedelta
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
import pytest_asyncio
|
import pytest_asyncio
|
||||||
@@ -268,7 +267,7 @@ def create_test_session(
|
|||||||
|
|
||||||
# Generate token and derive key
|
# Generate token and derive key
|
||||||
token = secrets.token_urlsafe(12)
|
token = secrets.token_urlsafe(12)
|
||||||
key = base64url.enc(hash_secret("cookie", token))
|
key = hash_secret("cookie", token)
|
||||||
|
|
||||||
session = Session.create(
|
session = Session.create(
|
||||||
user=user_uuid,
|
user=user_uuid,
|
||||||
|
|||||||
+1
-21
@@ -16,7 +16,6 @@ import secrets
|
|||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
import pytest_asyncio
|
import pytest_asyncio
|
||||||
@@ -188,25 +187,6 @@ class TestExceptionHandlers:
|
|||||||
assert "iframe" in data["auth"]
|
assert "iframe" in data["auth"]
|
||||||
|
|
||||||
|
|
||||||
# -------------------- Admin App Root --------------------
|
|
||||||
|
|
||||||
|
|
||||||
class TestAdminAppRoot:
|
|
||||||
"""Tests for the admin app root endpoint"""
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_admin_app_root_with_auth(
|
|
||||||
self, client: httpx.AsyncClient, session_token: str
|
|
||||||
):
|
|
||||||
"""Admin app root returns HTML when authenticated."""
|
|
||||||
response = await client.get(
|
|
||||||
"/auth/api/admin/",
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
|
||||||
)
|
|
||||||
assert response.status_code == 200
|
|
||||||
assert "text/html" in response.headers.get("content-type", "")
|
|
||||||
|
|
||||||
|
|
||||||
# -------------------- Organization Tests --------------------
|
# -------------------- Organization Tests --------------------
|
||||||
|
|
||||||
|
|
||||||
@@ -1320,7 +1300,7 @@ class TestAdminSessions:
|
|||||||
test_user,
|
test_user,
|
||||||
):
|
):
|
||||||
"""Admin can delete their own current session."""
|
"""Admin can delete their own current session."""
|
||||||
session_db_key = base64url.enc(hash_secret("cookie", session_token))
|
session_db_key = hash_secret("cookie", session_token)
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/users/{test_user.uuid}/sessions/{session_db_key}",
|
f"/auth/api/admin/users/{test_user.uuid}/sessions/{session_db_key}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
|
|||||||
@@ -355,34 +355,6 @@ class TestErrorHandling:
|
|||||||
class TestForwardAuthHtmlResponse:
|
class TestForwardAuthHtmlResponse:
|
||||||
"""Tests for forward auth HTML responses"""
|
"""Tests for forward auth HTML responses"""
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_forward_401_html_response(self, client: httpx.AsyncClient):
|
|
||||||
"""Forward auth 401 should return HTML page for browser requests."""
|
|
||||||
response = await client.get(
|
|
||||||
"/auth/api/forward",
|
|
||||||
headers={"Accept": "text/html"},
|
|
||||||
)
|
|
||||||
assert response.status_code == 401
|
|
||||||
assert "text/html" in response.headers.get("content-type", "")
|
|
||||||
# HTML response should contain the mode data attribute
|
|
||||||
assert b"data-mode" in response.content or b"mode" in response.content
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_forward_403_html_response(
|
|
||||||
self, client: httpx.AsyncClient, regular_session_token: str
|
|
||||||
):
|
|
||||||
"""Forward auth 403 should return HTML page for browser requests."""
|
|
||||||
response = await client.get(
|
|
||||||
"/auth/api/forward?perm=auth:admin",
|
|
||||||
headers={
|
|
||||||
**auth_headers(regular_session_token),
|
|
||||||
"Host": "localhost:4401",
|
|
||||||
"Accept": "text/html",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
assert response.status_code == 403
|
|
||||||
assert "text/html" in response.headers.get("content-type", "")
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_forward_with_expired_session_clears_cookie(
|
async def test_forward_with_expired_session_clears_cookie(
|
||||||
self, client: httpx.AsyncClient
|
self, client: httpx.AsyncClient
|
||||||
|
|||||||
Reference in New Issue
Block a user