Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7c4418e631 | ||
|
|
1648c8641f | ||
|
|
63eb088dbd | ||
|
|
e88cc004dd | ||
|
|
76921e8b31 | ||
|
|
c1b0aab296 | ||
|
|
8f89bb6d4b | ||
|
|
d16d1ed1c2 | ||
|
|
2cfca81672 | ||
|
|
ce300ebdaf | ||
|
|
7329223784 | ||
|
|
c8d659b5ca | ||
|
|
29ea6426fe | ||
|
|
88a170a37b | ||
|
|
38d240d86d | ||
|
|
53362b8061 | ||
|
|
d3d5f5a3c8 | ||
|
|
d156fb9221 | ||
|
|
f868bc59d1 | ||
|
|
0022986d4e | ||
|
|
b08cca754f | ||
|
|
aa58f08bc5 | ||
|
|
1062b5d6c8 | ||
|
|
3d49cbf2d6 | ||
|
|
13c49aebfd | ||
|
|
abec77d561 | ||
|
|
9b505ff553 | ||
|
|
ddd70e6130 | ||
|
|
cf1124c251 | ||
|
|
7504aaf7e0 | ||
|
|
e8247a2c7f | ||
|
|
968964c4c9 | ||
|
|
6aa1a08e39 | ||
|
|
31f40d874c | ||
|
|
7530d7a710 | ||
|
|
90d5f0e45f | ||
|
|
f0d1b86d6b | ||
|
|
dbdd1dbd3c | ||
|
|
8f862fb4d1 | ||
|
|
cfb917da46 | ||
|
|
3a8e7d1f4f | ||
|
|
86966526c4 | ||
|
|
3196aa7688 | ||
|
|
2fadaea19c | ||
|
|
cb84a81a06 | ||
|
|
9bdca1f43a | ||
|
|
0f29544bdb | ||
|
|
4ddaa9fdf4 | ||
|
|
7e568dbd10 | ||
|
|
fbc6108b7a |
@@ -97,14 +97,14 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
// Verify registration result
|
// Verify registration result
|
||||||
expect(result.session_token).toBeDefined()
|
expect(result.session_token).toBeDefined()
|
||||||
expect(result.session_token).toHaveLength(16)
|
expect(result.session_token).toHaveLength(16)
|
||||||
expect(result.user_uuid).toBeDefined()
|
expect(result.user).toBeDefined()
|
||||||
expect(result.credential_uuid).toBeDefined()
|
expect(result.credential).toBeDefined()
|
||||||
expect(result.message).toContain('successfully')
|
expect(result.message).toContain('successfully')
|
||||||
|
|
||||||
// Store for subsequent tests
|
// Store for subsequent tests
|
||||||
sessionToken = result.session_token
|
sessionToken = result.session_token
|
||||||
userUuid = result.user_uuid
|
userUuid = result.user
|
||||||
credentialUuid = result.credential_uuid
|
credentialUuid = result.credential
|
||||||
|
|
||||||
// Save session token for other test groups to use
|
// Save session token for other test groups to use
|
||||||
saveSessionToken(sessionToken)
|
saveSessionToken(sessionToken)
|
||||||
@@ -138,9 +138,9 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
const validation = await validateSession(page, baseUrl, sessionToken)
|
const validation = await validateSession(page, baseUrl, sessionToken)
|
||||||
|
|
||||||
expect(validation.valid).toBe(true)
|
expect(validation.valid).toBe(true)
|
||||||
expect(validation.user_uuid).toBe(userUuid)
|
expect(validation.ctx.user.uuid).toBe(userUuid)
|
||||||
|
|
||||||
console.log(`✓ Session validated for user: ${validation.user_uuid}`)
|
console.log(`✓ Session validated for user: ${validation.ctx.user.uuid}`)
|
||||||
})
|
})
|
||||||
|
|
||||||
test('should retrieve user info', async ({ page }) => {
|
test('should retrieve user info', async ({ page }) => {
|
||||||
@@ -148,8 +148,8 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
|
|
||||||
const userInfo = await getUserInfo(page, baseUrl, sessionToken)
|
const userInfo = await getUserInfo(page, baseUrl, sessionToken)
|
||||||
|
|
||||||
expect(userInfo.user.user_uuid).toBe(userUuid)
|
expect(userInfo.ctx.user.uuid).toBe(userUuid)
|
||||||
expect(userInfo.user.user_name).toBe('Admin User')
|
expect(userInfo.ctx.user.display_name).toBe('Admin User')
|
||||||
expect(userInfo.credentials).toBeDefined()
|
expect(userInfo.credentials).toBeDefined()
|
||||||
expect(userInfo.credentials.length).toBeGreaterThanOrEqual(1)
|
expect(userInfo.credentials.length).toBeGreaterThanOrEqual(1)
|
||||||
|
|
||||||
@@ -169,7 +169,7 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
await page.screenshot({ path: 'test-results/profile-view.png' })
|
await page.screenshot({ path: 'test-results/profile-view.png' })
|
||||||
console.log('✓ Screenshot saved: test-results/profile-view.png')
|
console.log('✓ Screenshot saved: test-results/profile-view.png')
|
||||||
|
|
||||||
console.log(`✓ User info retrieved: ${userInfo.user.user_name}`)
|
console.log(`✓ User info retrieved: ${userInfo.ctx.user.display_name}`)
|
||||||
console.log(`✓ Credentials count: ${userInfo.credentials.length}`)
|
console.log(`✓ Credentials count: ${userInfo.credentials.length}`)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -190,7 +190,7 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
displayName: 'Admin User (test device)'
|
displayName: 'Admin User (test device)'
|
||||||
})
|
})
|
||||||
|
|
||||||
console.log(`✓ Added test credential: ${regResult.credential_uuid}`)
|
console.log(`✓ Added test credential: ${regResult.credential}`)
|
||||||
|
|
||||||
// Now logout and authenticate with the fresh credential
|
// Now logout and authenticate with the fresh credential
|
||||||
await logout(page, baseUrl, regResult.session_token)
|
await logout(page, baseUrl, regResult.session_token)
|
||||||
@@ -201,7 +201,7 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
|
|
||||||
expect(result.session_token).toBeDefined()
|
expect(result.session_token).toBeDefined()
|
||||||
expect(result.session_token).toHaveLength(16)
|
expect(result.session_token).toHaveLength(16)
|
||||||
expect(result.user_uuid).toBe(userUuid)
|
expect(result.user).toBe(userUuid)
|
||||||
|
|
||||||
// Update session token for subsequent tests
|
// Update session token for subsequent tests
|
||||||
sessionToken = result.session_token
|
sessionToken = result.session_token
|
||||||
@@ -209,7 +209,7 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
// Save session token for other test groups to use
|
// Save session token for other test groups to use
|
||||||
saveSessionToken(sessionToken)
|
saveSessionToken(sessionToken)
|
||||||
|
|
||||||
console.log(`✓ Authenticated as user: ${result.user_uuid}`)
|
console.log(`✓ Authenticated as user: ${result.user}`)
|
||||||
console.log(`✓ New session token: ${sessionToken.substring(0, 4)}...`)
|
console.log(`✓ New session token: ${sessionToken.substring(0, 4)}...`)
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -219,7 +219,7 @@ test.describe('Passkey Authentication E2E', () => {
|
|||||||
const validation = await validateSession(page, baseUrl, sessionToken)
|
const validation = await validateSession(page, baseUrl, sessionToken)
|
||||||
|
|
||||||
expect(validation.valid).toBe(true)
|
expect(validation.valid).toBe(true)
|
||||||
expect(validation.user_uuid).toBe(userUuid)
|
expect(validation.ctx.user.uuid).toBe(userUuid)
|
||||||
|
|
||||||
console.log(`✓ New session validated`)
|
console.log(`✓ New session validated`)
|
||||||
})
|
})
|
||||||
@@ -291,8 +291,8 @@ test.describe('Device Addition Dialog', () => {
|
|||||||
// Wait for the profile view to load
|
// Wait for the profile view to load
|
||||||
await page.waitForSelector('[data-view="profile"]', { timeout: 5000 })
|
await page.waitForSelector('[data-view="profile"]', { timeout: 5000 })
|
||||||
|
|
||||||
// Click the "Add Another Device" button
|
// Click the "Another Device" button
|
||||||
const addDeviceButton = page.getByRole('button', { name: 'Add Another Device' })
|
const addDeviceButton = page.getByRole('button', { name: 'Another Device' })
|
||||||
await expect(addDeviceButton).toBeVisible()
|
await expect(addDeviceButton).toBeVisible()
|
||||||
await addDeviceButton.click()
|
await addDeviceButton.click()
|
||||||
|
|
||||||
@@ -301,7 +301,7 @@ test.describe('Device Addition Dialog', () => {
|
|||||||
await expect(dialog).toBeVisible({ timeout: 5000 })
|
await expect(dialog).toBeVisible({ timeout: 5000 })
|
||||||
|
|
||||||
// Verify dialog contains expected elements
|
// Verify dialog contains expected elements
|
||||||
await expect(dialog.locator('h2')).toContainText('Device Registration Link')
|
await expect(dialog.locator('h2')).toContainText('Add Another Device')
|
||||||
|
|
||||||
// Wait for QR code to be generated (canvas should have content)
|
// Wait for QR code to be generated (canvas should have content)
|
||||||
const qrCanvas = dialog.locator('.qr-code')
|
const qrCanvas = dialog.locator('.qr-code')
|
||||||
@@ -318,16 +318,16 @@ test.describe('Device Addition Dialog', () => {
|
|||||||
expect(linkHref).toContain('http://localhost:4404/auth/')
|
expect(linkHref).toContain('http://localhost:4404/auth/')
|
||||||
console.log(`✓ Device link displayed: ${linkText} (href: ${linkHref})`)
|
console.log(`✓ Device link displayed: ${linkText} (href: ${linkHref})`)
|
||||||
|
|
||||||
// Verify expiration warning is shown
|
// Verify help text is shown
|
||||||
await expect(dialog.locator('.reg-help')).toContainText('Expires')
|
await expect(dialog.locator('.reg-help')).toContainText('Scan this QR code')
|
||||||
|
|
||||||
// Take screenshot of the dialog
|
// Take screenshot of the dialog
|
||||||
await dialog.screenshot({ path: 'test-results/device-addition-dialog.png' })
|
await dialog.screenshot({ path: 'test-results/device-addition-dialog.png' })
|
||||||
console.log(`✓ Screenshot saved: test-results/device-addition-dialog.png`)
|
console.log(`✓ Screenshot saved: test-results/device-addition-dialog.png`)
|
||||||
|
|
||||||
// Verify Copy Link button exists
|
// Verify the QR link element is clickable (copy functionality is built into clicking it)
|
||||||
const copyButton = dialog.getByRole('button', { name: 'Copy Link' })
|
const qrLink = dialog.locator('a.qr-link')
|
||||||
await expect(copyButton).toBeVisible()
|
await expect(qrLink).toBeVisible()
|
||||||
|
|
||||||
// Close the dialog (use the text button, not the icon button)
|
// Close the dialog (use the text button, not the icon button)
|
||||||
const closeButton = dialog.locator('button.btn-secondary', { hasText: 'Close' })
|
const closeButton = dialog.locator('button.btn-secondary', { hasText: 'Close' })
|
||||||
@@ -357,12 +357,12 @@ test.describe('Device Addition Dialog', () => {
|
|||||||
await page.waitForSelector('[data-view="profile"]', { timeout: 5000 })
|
await page.waitForSelector('[data-view="profile"]', { timeout: 5000 })
|
||||||
|
|
||||||
// Open the dialog
|
// Open the dialog
|
||||||
await page.getByRole('button', { name: 'Add Another Device' }).click()
|
await page.getByRole('button', { name: 'Another Device' }).click()
|
||||||
const dialog = page.locator('.device-dialog')
|
const dialog = page.locator('.device-dialog')
|
||||||
await expect(dialog).toBeVisible({ timeout: 5000 })
|
await expect(dialog).toBeVisible({ timeout: 5000 })
|
||||||
|
|
||||||
// Extract the reset token from the displayed URL
|
// Extract the reset token from the displayed URL
|
||||||
const linkText = dialog.locator('.qr-link p')
|
const linkText = dialog.locator('.qr-link .link-text')
|
||||||
const linkContent = await linkText.textContent()
|
const linkContent = await linkText.textContent()
|
||||||
|
|
||||||
// URL format: localhost/auth/word1.word2.word3.word4.word5
|
// URL format: localhost/auth/word1.word2.word3.word4.word5
|
||||||
@@ -405,7 +405,7 @@ test.describe('Device Addition Dialog', () => {
|
|||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
test.describe('ProfileView - Add New Passkey', () => {
|
test.describe('ProfileView - Register New', () => {
|
||||||
const baseUrl = process.env.BASE_URL || 'http://localhost:4404'
|
const baseUrl = process.env.BASE_URL || 'http://localhost:4404'
|
||||||
|
|
||||||
test('should show credentials list in profile', async ({ page }) => {
|
test('should show credentials list in profile', async ({ page }) => {
|
||||||
@@ -427,7 +427,7 @@ test.describe('ProfileView - Add New Passkey', () => {
|
|||||||
console.log(`✓ Profile shows ${credentialItems} credential(s) in list`)
|
console.log(`✓ Profile shows ${credentialItems} credential(s) in list`)
|
||||||
})
|
})
|
||||||
|
|
||||||
test('should add a new passkey using Add New Passkey button', async ({ page }) => {
|
test('should add a new passkey using Register New button', async ({ page }) => {
|
||||||
const sessionToken = getSavedSessionToken()
|
const sessionToken = getSavedSessionToken()
|
||||||
test.skip(!sessionToken, 'Requires saved session token')
|
test.skip(!sessionToken, 'Requires saved session token')
|
||||||
|
|
||||||
@@ -444,8 +444,8 @@ test.describe('ProfileView - Add New Passkey', () => {
|
|||||||
const initialCredentialCount = await page.locator('.credential-item').count()
|
const initialCredentialCount = await page.locator('.credential-item').count()
|
||||||
console.log(`Initial credential count: ${initialCredentialCount}`)
|
console.log(`Initial credential count: ${initialCredentialCount}`)
|
||||||
|
|
||||||
// Click "Add New Passkey" button
|
// Click "Register New" button
|
||||||
const addPasskeyBtn = page.locator('button:has-text("Add New Passkey")')
|
const addPasskeyBtn = page.locator('button:has-text("Register New")')
|
||||||
await expect(addPasskeyBtn).toBeVisible()
|
await expect(addPasskeyBtn).toBeVisible()
|
||||||
await addPasskeyBtn.click()
|
await addPasskeyBtn.click()
|
||||||
|
|
||||||
@@ -490,7 +490,7 @@ test.describe('ProfileView - Add New Passkey', () => {
|
|||||||
|
|
||||||
// Try to add a passkey - with excludeCredentials the authenticator should
|
// Try to add a passkey - with excludeCredentials the authenticator should
|
||||||
// prevent re-registration of the same credential
|
// prevent re-registration of the same credential
|
||||||
const addPasskeyBtn = page.locator('button:has-text("Add New Passkey")')
|
const addPasskeyBtn = page.locator('button:has-text("Register New")')
|
||||||
await expect(addPasskeyBtn).toBeVisible()
|
await expect(addPasskeyBtn).toBeVisible()
|
||||||
await addPasskeyBtn.click()
|
await addPasskeyBtn.click()
|
||||||
|
|
||||||
@@ -541,8 +541,8 @@ test.describe('ProfileView - Multi-Authenticator', () => {
|
|||||||
await page.waitForSelector('.credential-list', { timeout: 10000 })
|
await page.waitForSelector('.credential-list', { timeout: 10000 })
|
||||||
const initialCredentialCount = await page.locator('.credential-item').count()
|
const initialCredentialCount = await page.locator('.credential-item').count()
|
||||||
|
|
||||||
// Click "Add New Passkey" button
|
// Click "Register New" button
|
||||||
const addPasskeyBtn = page.locator('button:has-text("Add New Passkey")')
|
const addPasskeyBtn = page.locator('button:has-text("Register New")')
|
||||||
await expect(addPasskeyBtn).toBeVisible()
|
await expect(addPasskeyBtn).toBeVisible()
|
||||||
await addPasskeyBtn.click()
|
await addPasskeyBtn.click()
|
||||||
|
|
||||||
|
|||||||
@@ -84,7 +84,7 @@ async function makeApiCall(page: Page, url: string, method = 'GET'): Promise<{ s
|
|||||||
// Wait a tick for the page's handler to retry, then make our own call
|
// Wait a tick for the page's handler to retry, then make our own call
|
||||||
setTimeout(async () => {
|
setTimeout(async () => {
|
||||||
try {
|
try {
|
||||||
const response = await fetch(url, { method, credentials: 'include' });
|
const response = await fetch(url, { method });
|
||||||
if (response.status === 204) {
|
if (response.status === 204) {
|
||||||
resolve({ status: 204 });
|
resolve({ status: 204 });
|
||||||
} else if (response.ok) {
|
} else if (response.ok) {
|
||||||
@@ -111,7 +111,7 @@ async function makeApiCall(page: Page, url: string, method = 'GET'): Promise<{ s
|
|||||||
setTimeout(async () => {
|
setTimeout(async () => {
|
||||||
if (resolved) return;
|
if (resolved) return;
|
||||||
try {
|
try {
|
||||||
const response = await fetch(url, { method, credentials: 'include' });
|
const response = await fetch(url, { method });
|
||||||
// Only resolve if this is a success or non-auth error
|
// Only resolve if this is a success or non-auth error
|
||||||
if (response.status !== 401 && response.status !== 403) {
|
if (response.status !== 401 && response.status !== 403) {
|
||||||
if (resolved) return;
|
if (resolved) return;
|
||||||
@@ -242,7 +242,7 @@ test.describe('API Mode - 401 Login Flow', () => {
|
|||||||
resetToken: deviceToken,
|
resetToken: deviceToken,
|
||||||
displayName: 'API Test Device',
|
displayName: 'API Test Device',
|
||||||
})
|
})
|
||||||
console.log(`✓ Registered credential: ${regResult.credential_uuid}`)
|
console.log(`✓ Registered credential: ${regResult.credential}`)
|
||||||
|
|
||||||
// Logout to clear session (but keep the passkey in virtual authenticator)
|
// Logout to clear session (but keep the passkey in virtual authenticator)
|
||||||
await logout(page, baseUrl, regResult.session_token)
|
await logout(page, baseUrl, regResult.session_token)
|
||||||
@@ -268,7 +268,7 @@ test.describe('API Mode - 401 Login Flow', () => {
|
|||||||
// Wait for API call to complete and verify result
|
// Wait for API call to complete and verify result
|
||||||
const result = await apiCallPromise
|
const result = await apiCallPromise
|
||||||
expect(result.status).toBe(200)
|
expect(result.status).toBe(200)
|
||||||
expect(result.data.user).toBeDefined()
|
expect(result.data.ctx).toBeDefined()
|
||||||
console.log('✓ API call succeeded after authentication')
|
console.log('✓ API call succeeded after authentication')
|
||||||
|
|
||||||
// Save the session for other tests
|
// Save the session for other tests
|
||||||
|
|||||||
+39
-5
@@ -12,17 +12,51 @@ const stateFile = join(__dirname, '..', '..', 'test-data', 'test-state.json')
|
|||||||
*/
|
*/
|
||||||
|
|
||||||
export interface RegistrationResult {
|
export interface RegistrationResult {
|
||||||
user_uuid: string
|
user: string
|
||||||
credential_uuid: string
|
credential: string
|
||||||
session_token: string
|
session_token: string
|
||||||
message: string
|
message: string
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface AuthenticationResult {
|
export interface AuthenticationResult {
|
||||||
user_uuid: string
|
user: string
|
||||||
session_token: string
|
session_token: string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
export interface SessionContext {
|
||||||
|
user: { uuid: string; display_name: string }
|
||||||
|
org: { uuid: string; display_name: string }
|
||||||
|
role: { uuid: string; display_name: string }
|
||||||
|
permissions: string[]
|
||||||
|
}
|
||||||
|
|
||||||
|
export interface UserInfo {
|
||||||
|
ctx: SessionContext
|
||||||
|
created_at: string
|
||||||
|
last_seen: string
|
||||||
|
visits: number
|
||||||
|
credentials: Array<{
|
||||||
|
credential: string
|
||||||
|
aaguid: string
|
||||||
|
created_at: string
|
||||||
|
last_used: string | null
|
||||||
|
last_verified: string | null
|
||||||
|
sign_count: number
|
||||||
|
is_current_session: boolean
|
||||||
|
}>
|
||||||
|
aaguid_info: Record<string, { name: string; icon_light?: string; icon_dark?: string }>
|
||||||
|
sessions: Array<{
|
||||||
|
id: string
|
||||||
|
credential: string
|
||||||
|
host: string
|
||||||
|
ip: string
|
||||||
|
user_agent: string
|
||||||
|
last_renewed: string
|
||||||
|
is_current: boolean
|
||||||
|
is_current_host: boolean
|
||||||
|
}>
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Get the bootstrap reset token from the test state file.
|
* Get the bootstrap reset token from the test state file.
|
||||||
*/
|
*/
|
||||||
@@ -376,7 +410,7 @@ export async function validateSession(
|
|||||||
page: Page,
|
page: Page,
|
||||||
baseUrl: string,
|
baseUrl: string,
|
||||||
sessionToken: string
|
sessionToken: string
|
||||||
): Promise<{ valid: boolean; user_uuid: string; renewed: boolean }> {
|
): Promise<{ valid: boolean; ctx: SessionContext; renewed: boolean }> {
|
||||||
const cookieName = getSessionCookieName()
|
const cookieName = getSessionCookieName()
|
||||||
const response = await page.request.post(`${baseUrl}/auth/api/validate`, {
|
const response = await page.request.post(`${baseUrl}/auth/api/validate`, {
|
||||||
headers: {
|
headers: {
|
||||||
@@ -393,7 +427,7 @@ export async function getUserInfo(
|
|||||||
page: Page,
|
page: Page,
|
||||||
baseUrl: string,
|
baseUrl: string,
|
||||||
sessionToken: string
|
sessionToken: string
|
||||||
): Promise<any> {
|
): Promise<UserInfo> {
|
||||||
const cookieName = getSessionCookieName()
|
const cookieName = getSessionCookieName()
|
||||||
const response = await page.request.post(`${baseUrl}/auth/api/user-info`, {
|
const response = await page.request.post(`${baseUrl}/auth/api/user-info`, {
|
||||||
headers: {
|
headers: {
|
||||||
|
|||||||
@@ -42,21 +42,23 @@ export default async function globalSetup() {
|
|||||||
const serverArgs = COLLECT_COVERAGE
|
const serverArgs = COLLECT_COVERAGE
|
||||||
? [
|
? [
|
||||||
'run', 'coverage', 'run', '--parallel-mode',
|
'run', 'coverage', 'run', '--parallel-mode',
|
||||||
'-m', 'paskia.fastapi', 'serve', 'localhost:4404',
|
'-m', 'paskia.fastapi', 'localhost:4404',
|
||||||
'--rp-id', 'localhost'
|
'--rp-id', 'localhost'
|
||||||
]
|
]
|
||||||
: [
|
: [
|
||||||
'run', 'paskia', 'serve', 'localhost:4404',
|
'run', 'paskia', 'localhost:4404',
|
||||||
'--rp-id', 'localhost'
|
'--rp-id', 'localhost'
|
||||||
]
|
]
|
||||||
|
|
||||||
|
// Use a temporary jsonl file for test database
|
||||||
|
const testDbFile = join(testDataDir, 'test-db.jsonl')
|
||||||
|
|
||||||
// Start the server using Node's spawn
|
// Start the server using Node's spawn
|
||||||
// Use in-memory SQLite for faster tests
|
|
||||||
const serverProcess = spawn('uv', serverArgs, {
|
const serverProcess = spawn('uv', serverArgs, {
|
||||||
cwd: projectRoot,
|
cwd: projectRoot,
|
||||||
env: {
|
env: {
|
||||||
...process.env,
|
...process.env,
|
||||||
PASKIA_DB: 'sqlite+aiosqlite:///:memory:',
|
PASKIA_DB: testDbFile,
|
||||||
COVERAGE_FILE: join(projectRoot, '.coverage'),
|
COVERAGE_FILE: join(projectRoot, '.coverage'),
|
||||||
},
|
},
|
||||||
stdio: ['ignore', 'pipe', 'pipe'],
|
stdio: ['ignore', 'pipe', 'pipe'],
|
||||||
|
|||||||
@@ -59,18 +59,11 @@ export default async function globalTeardown() {
|
|||||||
rmSync(stateFile, { force: true })
|
rmSync(stateFile, { force: true })
|
||||||
}
|
}
|
||||||
|
|
||||||
// Optionally clean up test database (keep it for debugging by default)
|
// Clean up test database
|
||||||
if (process.env.CLEANUP_TEST_DB === 'true') {
|
const testDbFile = join(testDataDir, 'test-db.jsonl')
|
||||||
const dbPath = join(testDataDir, 'test.sqlite')
|
if (existsSync(testDbFile)) {
|
||||||
if (existsSync(dbPath)) {
|
console.log(' Removing test database...')
|
||||||
console.log(' Removing test database...')
|
rmSync(testDbFile)
|
||||||
rmSync(dbPath)
|
|
||||||
}
|
|
||||||
// Remove wal/shm files too
|
|
||||||
for (const ext of ['-wal', '-shm']) {
|
|
||||||
const file = dbPath + ext
|
|
||||||
if (existsSync(file)) rmSync(file)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// Generate Python coverage report if coverage was collected
|
// Generate Python coverage report if coverage was collected
|
||||||
|
|||||||
+2
-2
@@ -96,7 +96,7 @@
|
|||||||
async function apiCall(url, method = 'GET') {
|
async function apiCall(url, method = 'GET') {
|
||||||
log(`${method} ${url}...`);
|
log(`${method} ${url}...`);
|
||||||
|
|
||||||
const response = await fetch(url, { method, credentials: 'include' });
|
const response = await fetch(url, { method });
|
||||||
|
|
||||||
// Server returns 401 (login/reauth) or 403 (missing permissions)
|
// Server returns 401 (login/reauth) or 403 (missing permissions)
|
||||||
// with a JSON body containing the iframe URL for authentication
|
// with a JSON body containing the iframe URL for authentication
|
||||||
@@ -131,7 +131,7 @@
|
|||||||
}
|
}
|
||||||
|
|
||||||
async function logout() {
|
async function logout() {
|
||||||
await fetch('/auth/api/logout', { method: 'POST', credentials: 'include' });
|
await fetch('/auth/api/logout', { method: 'POST' });
|
||||||
log('Logged out');
|
log('Logged out');
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+21
-37
@@ -2,10 +2,10 @@
|
|||||||
<div class="app-shell">
|
<div class="app-shell">
|
||||||
<StatusMessage />
|
<StatusMessage />
|
||||||
<main class="app-main">
|
<main class="app-main">
|
||||||
<HostProfileView v-if="authenticated && isHostMode" :initializing="loading" />
|
<HostProfileView v-if="viewState === 'profile' && isHostMode" />
|
||||||
<ProfileView v-else-if="authenticated" />
|
<ProfileView v-else-if="viewState === 'profile'" />
|
||||||
<LoadingView v-else-if="loading" :message="loadingMessage" />
|
<LoadingView v-else-if="viewState === 'loading'" :message="loadingMessage" />
|
||||||
<AuthRequiredMessage v-else-if="showBackMessage" @reload="reloadPage" />
|
<AccessDenied v-else-if="viewState === 'terminal'" />
|
||||||
</main>
|
</main>
|
||||||
</div>
|
</div>
|
||||||
</template>
|
</template>
|
||||||
@@ -18,13 +18,11 @@ 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'
|
||||||
import LoadingView from '@/components/LoadingView.vue'
|
import LoadingView from '@/components/LoadingView.vue'
|
||||||
import AuthRequiredMessage from '@/components/AccessDenied.vue'
|
import AccessDenied from '@/components/AccessDenied.vue'
|
||||||
|
|
||||||
const store = useAuthStore()
|
const store = useAuthStore()
|
||||||
const loading = ref(true)
|
const viewState = ref('loading') // 'loading' | 'profile' | 'terminal'
|
||||||
const loadingMessage = ref('Loading...')
|
const loadingMessage = ref('Loading...')
|
||||||
const authenticated = ref(false)
|
|
||||||
const showBackMessage = ref(false)
|
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Normalize a host string for comparison (lowercase, strip default ports).
|
* Normalize a host string for comparison (lowercase, strip default ports).
|
||||||
@@ -51,14 +49,19 @@ const isHostMode = computed(() => {
|
|||||||
let validationTimer = null
|
let validationTimer = null
|
||||||
let authIframe = null
|
let authIframe = null
|
||||||
|
|
||||||
|
function terminateSession() {
|
||||||
|
store.userInfo = null
|
||||||
|
viewState.value = 'terminal'
|
||||||
|
}
|
||||||
|
|
||||||
async function loadUserInfo() {
|
async function loadUserInfo() {
|
||||||
try {
|
try {
|
||||||
store.userInfo = await apiJson('/auth/api/user-info', { method: 'POST' })
|
store.userInfo = await apiJson('/auth/api/user-info', { method: 'POST' })
|
||||||
authenticated.value = true
|
viewState.value = 'profile'
|
||||||
loading.value = false
|
|
||||||
startSessionValidation()
|
startSessionValidation()
|
||||||
return true
|
return true
|
||||||
} catch (e) {
|
} catch {
|
||||||
|
store.userInfo = null
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -85,10 +88,6 @@ function hideAuthIframe() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
function reloadPage() {
|
|
||||||
window.location.reload()
|
|
||||||
}
|
|
||||||
|
|
||||||
function handleAuthMessage(event) {
|
function handleAuthMessage(event) {
|
||||||
const data = event.data
|
const data = event.data
|
||||||
if (!data?.type) return
|
if (!data?.type) return
|
||||||
@@ -97,7 +96,7 @@ function handleAuthMessage(event) {
|
|||||||
case 'auth-success':
|
case 'auth-success':
|
||||||
// Authentication successful - reload user info
|
// Authentication successful - reload user info
|
||||||
hideAuthIframe()
|
hideAuthIframe()
|
||||||
loading.value = true
|
viewState.value = 'loading'
|
||||||
loadingMessage.value = 'Loading user profile...'
|
loadingMessage.value = 'Loading user profile...'
|
||||||
loadUserInfo()
|
loadUserInfo()
|
||||||
break
|
break
|
||||||
@@ -117,11 +116,9 @@ function handleAuthMessage(event) {
|
|||||||
break
|
break
|
||||||
|
|
||||||
case 'auth-back':
|
case 'auth-back':
|
||||||
// User clicked Back - show message with reload option
|
// User clicked Back - show terminal state
|
||||||
hideAuthIframe()
|
hideAuthIframe()
|
||||||
loading.value = false
|
terminateSession()
|
||||||
showBackMessage.value = true
|
|
||||||
store.showMessage('Authentication cancelled', 'info', 3000)
|
|
||||||
break
|
break
|
||||||
|
|
||||||
case 'auth-close-request':
|
case 'auth-close-request':
|
||||||
@@ -133,23 +130,10 @@ function handleAuthMessage(event) {
|
|||||||
|
|
||||||
async function validateSession() {
|
async function validateSession() {
|
||||||
try {
|
try {
|
||||||
await apiJson('/auth/api/validate', {
|
await apiJson('/auth/api/validate', { method: 'POST' })
|
||||||
method: 'POST',
|
} catch {
|
||||||
credentials: 'include'
|
stopSessionValidation()
|
||||||
})
|
terminateSession()
|
||||||
// If successful, session was renewed automatically
|
|
||||||
} catch (error) {
|
|
||||||
if (error.status === 401) {
|
|
||||||
// Session expired - need to re-authenticate
|
|
||||||
console.log('Session expired, requiring re-authentication')
|
|
||||||
authenticated.value = false
|
|
||||||
loading.value = true
|
|
||||||
stopSessionValidation()
|
|
||||||
showAuthIframe()
|
|
||||||
} else {
|
|
||||||
console.error('Session validation error:', error)
|
|
||||||
// Don't treat network errors as session expiry
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ import CredentialList from '@/components/CredentialList.vue'
|
|||||||
import UserBasicInfo from '@/components/UserBasicInfo.vue'
|
import UserBasicInfo from '@/components/UserBasicInfo.vue'
|
||||||
import StatusMessage from '@/components/StatusMessage.vue'
|
import StatusMessage from '@/components/StatusMessage.vue'
|
||||||
import LoadingView from '@/components/LoadingView.vue'
|
import LoadingView from '@/components/LoadingView.vue'
|
||||||
import AuthRequiredMessage from '@/components/AccessDenied.vue'
|
import AccessDenied from '@/components/AccessDenied.vue'
|
||||||
import AdminOverview from '@/admin/AdminOverview.vue'
|
import AdminOverview from '@/admin/AdminOverview.vue'
|
||||||
import AdminOrgDetail from '@/admin/AdminOrgDetail.vue'
|
import AdminOrgDetail from '@/admin/AdminOrgDetail.vue'
|
||||||
import AdminUserDetail from '@/admin/AdminUserDetail.vue'
|
import AdminUserDetail from '@/admin/AdminUserDetail.vue'
|
||||||
@@ -48,8 +48,8 @@ const adminUserDetailRef = ref(null)
|
|||||||
const hasActiveModal = computed(() => dialog.value.type !== null || showRegModal.value)
|
const hasActiveModal = computed(() => dialog.value.type !== null || showRegModal.value)
|
||||||
|
|
||||||
// Derive admin status from permissions
|
// Derive admin status from permissions
|
||||||
const isGlobalAdmin = computed(() => info.value?.permissions?.includes('auth:admin') ?? false)
|
const isMasterAdmin = computed(() => info.value?.ctx.permissions.includes('auth:admin'))
|
||||||
const isOrgAdmin = computed(() => info.value?.permissions?.includes('auth:org:admin') ?? false)
|
const isOrgAdmin = computed(() => info.value?.ctx.permissions.includes('auth:org:admin'))
|
||||||
|
|
||||||
function sanitizeRenameId() { if (renameIdValue.value) renameIdValue.value = renameIdValue.value.replace(safeIdRegex, '') }
|
function sanitizeRenameId() { if (renameIdValue.value) renameIdValue.value = renameIdValue.value.replace(safeIdRegex, '') }
|
||||||
|
|
||||||
@@ -130,7 +130,7 @@ function parseHash() {
|
|||||||
async function loadOrgs() {
|
async function loadOrgs() {
|
||||||
const data = await apiJson('/auth/api/admin/orgs')
|
const data = await apiJson('/auth/api/admin/orgs')
|
||||||
orgs.value = data.map(o => {
|
orgs.value = data.map(o => {
|
||||||
const roles = o.roles.map(r => ({ ...r, org_uuid: o.uuid, users: [] }))
|
const roles = o.roles.map(r => ({ ...r, org: o.uuid, users: [] }))
|
||||||
const roleMap = Object.fromEntries(roles.map(r => [r.display_name, r]))
|
const roleMap = Object.fromEntries(roles.map(r => [r.display_name, r]))
|
||||||
for (const u of o.users || []) {
|
for (const u of o.users || []) {
|
||||||
if (roleMap[u.role]) roleMap[u.role].users.push(u)
|
if (roleMap[u.role]) roleMap[u.role].users.push(u)
|
||||||
@@ -144,10 +144,19 @@ async function loadPermissions() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async function loadUserInfo() {
|
async function loadUserInfo() {
|
||||||
info.value = await apiJson('/auth/api/user-info', { method: 'POST' })
|
const data = await apiJson('/auth/api/validate', { method: 'POST' })
|
||||||
|
info.value = data
|
||||||
authenticated.value = true
|
authenticated.value = true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function clearSensitiveState() {
|
||||||
|
info.value = null
|
||||||
|
orgs.value = []
|
||||||
|
permissions.value = []
|
||||||
|
userDetail.value = null
|
||||||
|
authenticated.value = false
|
||||||
|
}
|
||||||
|
|
||||||
async function load() {
|
async function load() {
|
||||||
loading.value = true
|
loading.value = true
|
||||||
loadingMessage.value = 'Loading...'
|
loadingMessage.value = 'Loading...'
|
||||||
@@ -158,7 +167,7 @@ async function load() {
|
|||||||
// If we get here, user has admin access - now fetch user info for display
|
// If we get here, user has admin access - now fetch user info for display
|
||||||
await loadUserInfo()
|
await loadUserInfo()
|
||||||
|
|
||||||
if (!isGlobalAdmin.value && isOrgAdmin.value && orgs.value.length === 1) {
|
if (!isMasterAdmin.value && isOrgAdmin.value && orgs.value.length === 1) {
|
||||||
if (!window.location.hash || window.location.hash === '#overview') {
|
if (!window.location.hash || window.location.hash === '#overview') {
|
||||||
currentOrgId.value = orgs.value[0].uuid
|
currentOrgId.value = orgs.value[0].uuid
|
||||||
window.location.hash = `#org/${currentOrgId.value}`
|
window.location.hash = `#org/${currentOrgId.value}`
|
||||||
@@ -168,6 +177,7 @@ async function load() {
|
|||||||
}
|
}
|
||||||
} else parseHash()
|
} else parseHash()
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
|
clearSensitiveState()
|
||||||
if (e.name === 'AuthCancelledError') {
|
if (e.name === 'AuthCancelledError') {
|
||||||
showBackMessage.value = true
|
showBackMessage.value = true
|
||||||
} else {
|
} else {
|
||||||
@@ -191,8 +201,6 @@ async function performOrgDeletion(orgUuid) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
function deleteOrg(org) {
|
function deleteOrg(org) {
|
||||||
if (!isGlobalAdmin.value) { authStore.showMessage('Global admin only'); return }
|
|
||||||
|
|
||||||
const userCount = org.roles.reduce((acc, r) => acc + r.users.length, 0)
|
const userCount = org.roles.reduce((acc, r) => acc + r.users.length, 0)
|
||||||
|
|
||||||
if (userCount === 0) {
|
if (userCount === 0) {
|
||||||
@@ -234,9 +242,9 @@ async function moveUserToRole(org, user, targetRoleDisplayName) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
function onUserDragStart(e, user, org_uuid) {
|
function onUserDragStart(e, user, org) {
|
||||||
e.dataTransfer.effectAllowed = 'move'
|
e.dataTransfer.effectAllowed = 'move'
|
||||||
e.dataTransfer.setData('text/plain', JSON.stringify({ user_uuid: user.uuid, org_uuid }))
|
e.dataTransfer.setData('text/plain', JSON.stringify({ user_uuid: user.uuid, org }))
|
||||||
}
|
}
|
||||||
|
|
||||||
function onRoleDragOver(e) {
|
function onRoleDragOver(e) {
|
||||||
@@ -248,7 +256,7 @@ function onRoleDrop(e, org, role) {
|
|||||||
e.preventDefault()
|
e.preventDefault()
|
||||||
try {
|
try {
|
||||||
const data = JSON.parse(e.dataTransfer.getData('text/plain'))
|
const data = JSON.parse(e.dataTransfer.getData('text/plain'))
|
||||||
if (data.org_uuid !== org.uuid) return // only within same org
|
if (data.org !== org.uuid) return // only within same org
|
||||||
const user = org.roles.flatMap(r => r.users).find(u => u.uuid === data.user_uuid)
|
const user = org.roles.flatMap(r => r.users).find(u => u.uuid === data.user_uuid)
|
||||||
if (user) moveUserToRole(org, user, role.display_name)
|
if (user) moveUserToRole(org, user, role.display_name)
|
||||||
} catch (_) { /* ignore */ }
|
} catch (_) { /* ignore */ }
|
||||||
@@ -261,7 +269,7 @@ function updateRole(role) { openDialog('role-update', { role, name: role.display
|
|||||||
|
|
||||||
function deleteRole(role) {
|
function deleteRole(role) {
|
||||||
// UI only allows deleting empty roles, so no confirmation needed
|
// UI only allows deleting empty roles, so no confirmation needed
|
||||||
apiJson(`/auth/api/admin/orgs/${role.org_uuid}/roles/${role.uuid}`, { method: 'DELETE' })
|
apiJson(`/auth/api/admin/orgs/${role.org}/roles/${role.uuid}`, { method: 'DELETE' })
|
||||||
.then(() => {
|
.then(() => {
|
||||||
authStore.showMessage(`Role "${role.display_name}" deleted.`, 'success', 2500)
|
authStore.showMessage(`Role "${role.display_name}" deleted.`, 'success', 2500)
|
||||||
loadOrgs()
|
loadOrgs()
|
||||||
@@ -281,7 +289,7 @@ async function toggleRolePermission(role, pid, checked) {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
const method = checked ? 'POST' : 'DELETE'
|
const method = checked ? 'POST' : 'DELETE'
|
||||||
await apiJson(`/auth/api/admin/orgs/${role.org_uuid}/roles/${role.uuid}/permissions/${pid}`, {
|
await apiJson(`/auth/api/admin/orgs/${role.org}/roles/${role.uuid}/permissions/${pid}`, {
|
||||||
method
|
method
|
||||||
})
|
})
|
||||||
await loadOrgs()
|
await loadOrgs()
|
||||||
@@ -292,8 +300,8 @@ async function toggleRolePermission(role, pid, checked) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Permission actions
|
// Permission actions
|
||||||
async function performPermissionDeletion(permissionScope) {
|
async function performPermissionDeletion(permissionUuid) {
|
||||||
const params = new URLSearchParams({ permission_id: permissionScope })
|
const params = new URLSearchParams({ permission_uuid: permissionUuid })
|
||||||
await apiJson(`/auth/api/admin/permission?${params.toString()}`, { method: 'DELETE' })
|
await apiJson(`/auth/api/admin/permission?${params.toString()}`, { method: 'DELETE' })
|
||||||
await loadPermissions()
|
await loadPermissions()
|
||||||
}
|
}
|
||||||
@@ -313,7 +321,7 @@ function deletePermission(p) {
|
|||||||
|
|
||||||
if (roleCount === 0) {
|
if (roleCount === 0) {
|
||||||
// No roles have this permission, safe to delete directly
|
// No roles have this permission, safe to delete directly
|
||||||
performPermissionDeletion(p.scope)
|
performPermissionDeletion(p.uuid)
|
||||||
.then(() => {
|
.then(() => {
|
||||||
authStore.showMessage(`Permission "${p.display_name}" deleted.`, 'success', 2500)
|
authStore.showMessage(`Permission "${p.display_name}" deleted.`, 'success', 2500)
|
||||||
})
|
})
|
||||||
@@ -329,14 +337,10 @@ function deletePermission(p) {
|
|||||||
const affects = parts.join(', ')
|
const affects = parts.join(', ')
|
||||||
|
|
||||||
openDialog('confirm', { message: `Delete permission "${p.display_name}" (${affects})?`, action: async () => {
|
openDialog('confirm', { message: `Delete permission "${p.display_name}" (${affects})?`, action: async () => {
|
||||||
await performPermissionDeletion(p.scope)
|
await performPermissionDeletion(p.uuid)
|
||||||
} })
|
} })
|
||||||
}
|
}
|
||||||
|
|
||||||
function reloadPage() {
|
|
||||||
window.location.reload()
|
|
||||||
}
|
|
||||||
|
|
||||||
const selectedOrg = computed(() => orgs.value.find(o => o.uuid === currentOrgId.value) || null)
|
const selectedOrg = computed(() => orgs.value.find(o => o.uuid === currentOrgId.value) || null)
|
||||||
|
|
||||||
function openOrg(o) {
|
function openOrg(o) {
|
||||||
@@ -356,7 +360,7 @@ const selectedUser = computed(() => {
|
|||||||
for (const o of orgs.value) {
|
for (const o of orgs.value) {
|
||||||
for (const r of o.roles) {
|
for (const r of o.roles) {
|
||||||
const u = r.users.find(x => x.uuid === currentUserId.value)
|
const u = r.users.find(x => x.uuid === currentUserId.value)
|
||||||
if (u) return { ...u, org_uuid: o.uuid, role_display_name: r.display_name }
|
if (u) return { ...u, org: o.uuid, role_display_name: r.display_name }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return null
|
return null
|
||||||
@@ -377,14 +381,14 @@ const breadcrumbEntries = computed(() => {
|
|||||||
// Determine organization for user view if selectedOrg not explicitly chosen.
|
// Determine organization for user view if selectedOrg not explicitly chosen.
|
||||||
let orgForUser = null
|
let orgForUser = null
|
||||||
if (selectedUser.value) {
|
if (selectedUser.value) {
|
||||||
orgForUser = orgs.value.find(o => o.uuid === selectedUser.value.org_uuid) || null
|
orgForUser = orgs.value.find(o => o.uuid === selectedUser.value.org) || null
|
||||||
}
|
}
|
||||||
const orgToShow = selectedOrg.value || orgForUser
|
const orgToShow = selectedOrg.value || orgForUser
|
||||||
if (orgToShow) {
|
if (orgToShow) {
|
||||||
entries.push({ label: orgToShow.display_name, href: `#org/${orgToShow.uuid}` })
|
entries.push({ label: orgToShow.display_name, href: `#org/${orgToShow.uuid}` })
|
||||||
}
|
}
|
||||||
if (selectedUser.value) {
|
if (selectedUser.value) {
|
||||||
entries.push({ label: selectedUser.value.display_name || 'User', href: `#user/${selectedUser.value.uuid}` })
|
entries.push({ label: selectedUser.value.display_name, href: `#user/${selectedUser.value.uuid}` })
|
||||||
}
|
}
|
||||||
return entries
|
return entries
|
||||||
})
|
})
|
||||||
@@ -392,7 +396,7 @@ const breadcrumbEntries = computed(() => {
|
|||||||
watch(selectedUser, async (u) => {
|
watch(selectedUser, async (u) => {
|
||||||
if (!u) { userDetail.value = null; return }
|
if (!u) { userDetail.value = null; return }
|
||||||
try {
|
try {
|
||||||
userDetail.value = await apiJson(`/auth/api/admin/orgs/${u.org_uuid}/users/${u.uuid}`)
|
userDetail.value = await apiJson(`/auth/api/admin/orgs/${u.org}/users/${u.uuid}`)
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
userDetail.value = { error: e.message }
|
userDetail.value = { error: e.message }
|
||||||
}
|
}
|
||||||
@@ -413,7 +417,7 @@ async function toggleOrgPermission(org, permId, checked) {
|
|||||||
const prev = [...org.permissions]
|
const prev = [...org.permissions]
|
||||||
org.permissions = next
|
org.permissions = next
|
||||||
try {
|
try {
|
||||||
const params = new URLSearchParams({ permission_id: permId })
|
const params = new URLSearchParams({ permission_uuid: permId })
|
||||||
await apiJson(`/auth/api/admin/orgs/${org.uuid}/permission?${params.toString()}`, { method: checked ? 'POST' : 'DELETE' })
|
await apiJson(`/auth/api/admin/orgs/${org.uuid}/permission?${params.toString()}`, { method: checked ? 'POST' : 'DELETE' })
|
||||||
await loadOrgs()
|
await loadOrgs()
|
||||||
} catch (e) {
|
} catch (e) {
|
||||||
@@ -528,7 +532,7 @@ async function refreshUserDetail() {
|
|||||||
await loadOrgs()
|
await loadOrgs()
|
||||||
if (selectedUser.value) {
|
if (selectedUser.value) {
|
||||||
try {
|
try {
|
||||||
userDetail.value = await apiJson(`/auth/api/admin/orgs/${selectedUser.value.org_uuid}/users/${selectedUser.value.uuid}`)
|
userDetail.value = await apiJson(`/auth/api/admin/orgs/${selectedUser.value.org}/users/${selectedUser.value.uuid}`)
|
||||||
} catch (e) { authStore.showMessage(e.message || 'Failed to reload user', 'error') }
|
} catch (e) { authStore.showMessage(e.message || 'Failed to reload user', 'error') }
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -590,7 +594,7 @@ async function submitDialog() {
|
|||||||
|
|
||||||
// Close dialog immediately, then perform async operation
|
// Close dialog immediately, then perform async operation
|
||||||
closeDialog()
|
closeDialog()
|
||||||
apiJson(`/auth/api/admin/orgs/${role.org_uuid}/roles/${role.uuid}`, { method: 'PATCH', body: { display_name: name } })
|
apiJson(`/auth/api/admin/orgs/${role.org}/roles/${role.uuid}`, { method: 'PATCH', body: { display_name: name } })
|
||||||
.then(() => {
|
.then(() => {
|
||||||
authStore.showMessage(`Role renamed to "${name}".`, 'success', 2500)
|
authStore.showMessage(`Role renamed to "${name}".`, 'success', 2500)
|
||||||
loadOrgs()
|
loadOrgs()
|
||||||
@@ -618,7 +622,7 @@ async function submitDialog() {
|
|||||||
|
|
||||||
// Close dialog immediately, then perform async operation
|
// Close dialog immediately, then perform async operation
|
||||||
closeDialog()
|
closeDialog()
|
||||||
apiJson(`/auth/api/admin/orgs/${user.org_uuid}/users/${user.uuid}/display-name`, { method: 'PATCH', body: { display_name: name } })
|
apiJson(`/auth/api/admin/orgs/${user.org}/users/${user.uuid}/display-name`, { method: 'PATCH', body: { display_name: name } })
|
||||||
.then(() => {
|
.then(() => {
|
||||||
authStore.showMessage(`User renamed to "${name}".`, 'success', 2500)
|
authStore.showMessage(`User renamed to "${name}".`, 'success', 2500)
|
||||||
onUserNameSaved()
|
onUserNameSaved()
|
||||||
@@ -629,31 +633,28 @@ async function submitDialog() {
|
|||||||
return // Don't call closeDialog() again
|
return // Don't call closeDialog() again
|
||||||
} else if (t === 'perm-display') {
|
} else if (t === 'perm-display') {
|
||||||
const { permission } = dialog.value.data
|
const { permission } = dialog.value.data
|
||||||
const newId = dialog.value.data.scope?.trim()
|
const newScope = dialog.value.data.scope?.trim()
|
||||||
const newDisplay = dialog.value.data.display_name?.trim()
|
const newDisplay = dialog.value.data.display_name?.trim()
|
||||||
const newDomain = dialog.value.data.domain?.trim() || ''
|
const newDomain = dialog.value.data.domain?.trim() || ''
|
||||||
if (!newDisplay) throw new Error('Display name required')
|
if (!newDisplay) throw new Error('Display name required')
|
||||||
if (!newId) throw new Error('Scope required')
|
if (!newScope) throw new Error('Scope required')
|
||||||
|
|
||||||
// Close dialog immediately, then perform async operation
|
// Close dialog immediately, then perform async operation
|
||||||
closeDialog()
|
closeDialog()
|
||||||
|
|
||||||
const oldDomain = permission.domain || ''
|
const oldDomain = permission.domain || ''
|
||||||
let apiCall;
|
// Check if anything changed
|
||||||
if (newId !== permission.scope) {
|
if (newScope === permission.scope && newDisplay === permission.display_name && newDomain === oldDomain) {
|
||||||
// Scope changed, use rename endpoint (also update domain)
|
return // No changes
|
||||||
apiCall = apiJson('/auth/api/admin/permission/rename', { method: 'POST', body: { old_scope: permission.scope, new_scope: newId, display_name: newDisplay, domain: newDomain } })
|
|
||||||
} else if (newDisplay !== permission.display_name || newDomain !== oldDomain) {
|
|
||||||
// Display name or domain changed
|
|
||||||
const params = new URLSearchParams({ permission_id: permission.scope, display_name: newDisplay })
|
|
||||||
if (newDomain) params.set('domain', newDomain)
|
|
||||||
apiCall = apiJson(`/auth/api/admin/permission?${params.toString()}`, { method: 'PATCH' })
|
|
||||||
} else {
|
|
||||||
// No changes
|
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
apiCall
|
// Always use PATCH with permission_uuid
|
||||||
|
const params = new URLSearchParams({ permission_uuid: permission.uuid })
|
||||||
|
if (newScope !== permission.scope) params.set('scope', newScope)
|
||||||
|
if (newDisplay !== permission.display_name) params.set('display_name', newDisplay)
|
||||||
|
if (newDomain !== oldDomain) params.set('domain', newDomain || '')
|
||||||
|
|
||||||
|
apiJson(`/auth/api/admin/permission?${params.toString()}`, { method: 'PATCH' })
|
||||||
.then(() => {
|
.then(() => {
|
||||||
authStore.showMessage(`Permission "${newDisplay}" updated.`, 'success', 2500)
|
authStore.showMessage(`Permission "${newDisplay}" updated.`, 'success', 2500)
|
||||||
loadPermissions()
|
loadPermissions()
|
||||||
@@ -703,23 +704,19 @@ async function submitDialog() {
|
|||||||
<StatusMessage />
|
<StatusMessage />
|
||||||
<main class="app-main">
|
<main class="app-main">
|
||||||
<LoadingView v-if="loading" :message="loadingMessage" />
|
<LoadingView v-if="loading" :message="loadingMessage" />
|
||||||
<AuthRequiredMessage
|
<AccessDenied v-else-if="showBackMessage" />
|
||||||
v-else-if="showBackMessage"
|
<AccessDenied
|
||||||
@reload="reloadPage"
|
v-else-if="error"
|
||||||
|
icon="⚠️"
|
||||||
|
title="Error"
|
||||||
|
:message="error"
|
||||||
/>
|
/>
|
||||||
<!-- Access denied: authenticated but not admin, or error occurred -->
|
<AccessDenied
|
||||||
<div v-else-if="error || (authenticated && !isGlobalAdmin && !isOrgAdmin)" class="access-denied-container">
|
v-else-if="authenticated && !isMasterAdmin && !isOrgAdmin"
|
||||||
<div class="access-denied-content">
|
icon="⛔"
|
||||||
<h2>⛔ Access Denied</h2>
|
message="You do not have admin permissions for this application."
|
||||||
<p v-if="error" class="error-detail">{{ error }}</p>
|
/>
|
||||||
<p v-else class="error-detail">You do not have admin permissions for this application.</p>
|
<section v-else-if="authenticated && (isMasterAdmin || isOrgAdmin)" class="view-root view-root--wide view-admin">
|
||||||
<div class="button-row">
|
|
||||||
<button class="btn-secondary" @click="goBack">Back</button>
|
|
||||||
<button class="btn-primary" @click="reloadPage">Reload Page</button>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
</div>
|
|
||||||
<section v-else-if="authenticated && (isGlobalAdmin || isOrgAdmin)" class="view-root view-root--wide view-admin">
|
|
||||||
<header class="view-header">
|
<header class="view-header">
|
||||||
<h1>{{ pageHeading }}</h1>
|
<h1>{{ pageHeading }}</h1>
|
||||||
<Breadcrumbs ref="breadcrumbsRef" :entries="breadcrumbEntries" @keydown="handleBreadcrumbKeydown" />
|
<Breadcrumbs ref="breadcrumbsRef" :entries="breadcrumbEntries" @keydown="handleBreadcrumbKeydown" />
|
||||||
@@ -729,7 +726,7 @@ async function submitDialog() {
|
|||||||
<div class="section-body admin-section-body">
|
<div class="section-body admin-section-body">
|
||||||
<div class="admin-panels">
|
<div class="admin-panels">
|
||||||
<AdminOverview
|
<AdminOverview
|
||||||
v-if="!selectedUser && !selectedOrg && (isGlobalAdmin || isOrgAdmin)"
|
v-if="!selectedUser && !selectedOrg && (isMasterAdmin || isOrgAdmin)"
|
||||||
ref="adminOverviewRef"
|
ref="adminOverviewRef"
|
||||||
:info="info"
|
:info="info"
|
||||||
:orgs="orgs"
|
:orgs="orgs"
|
||||||
@@ -805,9 +802,4 @@ async function submitDialog() {
|
|||||||
.admin-section { margin-top: var(--space-xl); }
|
.admin-section { margin-top: var(--space-xl); }
|
||||||
.admin-section-body { display: flex; flex-direction: column; gap: var(--space-xl); }
|
.admin-section-body { display: flex; flex-direction: column; gap: var(--space-xl); }
|
||||||
.admin-panels { display: flex; flex-direction: column; gap: var(--space-xl); }
|
.admin-panels { display: flex; flex-direction: column; gap: var(--space-xl); }
|
||||||
.access-denied-container { display: flex; flex-direction: column; align-items: center; justify-content: center; min-height: 60vh; padding: 2rem; }
|
|
||||||
.access-denied-content { text-align: center; max-width: 480px; }
|
|
||||||
.access-denied-content h2 { margin: 0 0 1rem; color: var(--color-heading); font-size: 1.5rem; }
|
|
||||||
.access-denied-content .error-detail { margin: 0 0 1.5rem; color: var(--color-text-muted); }
|
|
||||||
.access-denied-content .button-row { display: flex; gap: 0.75rem; justify-content: center; }
|
|
||||||
</style>
|
</style>
|
||||||
|
|||||||
@@ -71,12 +71,12 @@ const initializing = ref(true)
|
|||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
const token = ref('')
|
const token = ref('')
|
||||||
const settings = ref(null)
|
const settings = ref(null)
|
||||||
const userInfo = ref(null)
|
const tokenInfo = ref(null)
|
||||||
const displayName = ref('')
|
const displayName = ref('')
|
||||||
const errorMessage = ref('')
|
const errorMessage = ref('')
|
||||||
let statusTimer = null
|
let statusTimer = null
|
||||||
|
|
||||||
const sessionDescriptor = computed(() => userInfo.value?.session_type || 'your enrollment')
|
const sessionDescriptor = computed(() => tokenInfo.value?.token_type || 'your enrollment')
|
||||||
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.'
|
||||||
@@ -85,7 +85,7 @@ const subtitleMessage = computed(() => {
|
|||||||
|
|
||||||
const basePath = computed(() => uiBasePath())
|
const basePath = computed(() => uiBasePath())
|
||||||
|
|
||||||
const canRegister = computed(() => !!(token.value && userInfo.value))
|
const canRegister = computed(() => !!(token.value && tokenInfo.value))
|
||||||
|
|
||||||
function showMessage(message, type = 'info', duration = 3000) {
|
function showMessage(message, type = 'info', duration = 3000) {
|
||||||
status.show = true
|
status.show = true
|
||||||
@@ -109,15 +109,16 @@ async function fetchSettings() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async function fetchUserInfo() {
|
async function fetchTokenInfo() {
|
||||||
if (!token.value) return
|
if (!token.value) return
|
||||||
try {
|
try {
|
||||||
userInfo.value = await apiJson(`/auth/api/user-info?reset=${encodeURIComponent(token.value)}`, {
|
tokenInfo.value = await apiJson('/auth/api/token-info', {
|
||||||
method: 'POST'
|
method: 'GET',
|
||||||
|
headers: { 'Authorization': `Bearer ${token.value}` },
|
||||||
})
|
})
|
||||||
displayName.value = userInfo.value?.user?.user_name || ''
|
displayName.value = tokenInfo.value.display_name
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Failed to load user info', error)
|
console.error('Failed to load token info', error)
|
||||||
const message = error instanceof ApiError
|
const message = error instanceof ApiError
|
||||||
? (error.data?.detail || 'The authentication link is invalid or expired.')
|
? (error.data?.detail || 'The authentication link is invalid or expired.')
|
||||||
: getUserFriendlyErrorMessage(error)
|
: getUserFriendlyErrorMessage(error)
|
||||||
@@ -196,7 +197,7 @@ onMounted(async () => {
|
|||||||
initializing.value = false
|
initializing.value = false
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
await fetchUserInfo()
|
await fetchTokenInfo()
|
||||||
initializing.value = false
|
initializing.value = false
|
||||||
})
|
})
|
||||||
</script>
|
</script>
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ const rpId = computed(() => props.settings?.rp_id || 'the configured domain')
|
|||||||
<template v-else-if="dialog.type==='role-update'">Edit Role</template>
|
<template v-else-if="dialog.type==='role-update'">Edit Role</template>
|
||||||
<template v-else-if="dialog.type==='user-create'">Add User To Role</template>
|
<template v-else-if="dialog.type==='user-create'">Add User To Role</template>
|
||||||
<template v-else-if="dialog.type==='user-update-name'">Edit User Name</template>
|
<template v-else-if="dialog.type==='user-update-name'">Edit User Name</template>
|
||||||
<template v-else-if="dialog.type==='perm-create' || dialog.type==='perm-display'">{{ dialog.type === 'perm-create' ? 'Create Permission' : 'Edit Permission Display' }}</template>
|
<template v-else-if="dialog.type==='perm-create' || dialog.type==='perm-display'">{{ dialog.type === 'perm-create' ? 'Create Permission' : 'Edit Permission' }}</template>
|
||||||
<template v-else-if="dialog.type==='confirm'">Confirm</template>
|
<template v-else-if="dialog.type==='confirm'">Confirm</template>
|
||||||
</h3>
|
</h3>
|
||||||
<form @submit.prevent="$emit('submitDialog')" class="modal-form">
|
<form @submit.prevent="$emit('submitDialog')" class="modal-form">
|
||||||
|
|||||||
@@ -26,9 +26,9 @@ const sortedOrgs = computed(() => [...props.orgs].sort((a,b)=> {
|
|||||||
}))
|
}))
|
||||||
const sortedPermissions = computed(() => [...props.permissions].sort((a,b)=> a.scope.localeCompare(b.scope)))
|
const sortedPermissions = computed(() => [...props.permissions].sort((a,b)=> a.scope.localeCompare(b.scope)))
|
||||||
|
|
||||||
// Derive admin status from permissions
|
// Derive admin status from permissions (info contains ctx from validate response)
|
||||||
const isGlobalAdmin = computed(() => props.info?.permissions?.includes('auth:admin') ?? false)
|
const isMasterAdmin = computed(() => props.info?.ctx.permissions.includes('auth:admin'))
|
||||||
const isOrgAdmin = computed(() => props.info?.permissions?.includes('auth:org:admin') ?? false)
|
const isOrgAdmin = computed(() => props.info?.ctx.permissions.includes('auth:org:admin'))
|
||||||
|
|
||||||
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
|
||||||
@@ -93,7 +93,7 @@ function handleTableKeydown(event, tableType) {
|
|||||||
} else if (direction === 'down' && currentIndex === rows.length - 1) {
|
} else if (direction === 'down' && currentIndex === rows.length - 1) {
|
||||||
// At bottom of org table, navigate to permissions section
|
// At bottom of org table, navigate to permissions section
|
||||||
event.preventDefault()
|
event.preventDefault()
|
||||||
if (tableType === 'org' && isGlobalAdmin.value) {
|
if (tableType === 'org' && isMasterAdmin.value) {
|
||||||
// Navigate to permissions matrix or actions
|
// Navigate to permissions matrix or actions
|
||||||
if (permMatrixRef.value) {
|
if (permMatrixRef.value) {
|
||||||
const firstCheckbox = permMatrixRef.value.querySelector('input[type="checkbox"]')
|
const firstCheckbox = permMatrixRef.value.querySelector('input[type="checkbox"]')
|
||||||
@@ -236,7 +236,7 @@ function handlePermActionsKeydown(event) {
|
|||||||
|
|
||||||
// Focus helper for external navigation
|
// Focus helper for external navigation
|
||||||
function focusFirstElement() {
|
function focusFirstElement() {
|
||||||
if (isGlobalAdmin.value) {
|
if (isMasterAdmin.value) {
|
||||||
focusPreferred(orgActionsRef.value, { itemSelector: 'button' })
|
focusPreferred(orgActionsRef.value, { itemSelector: 'button' })
|
||||||
} else {
|
} else {
|
||||||
const firstFocusable = orgTableRef.value?.querySelector('tbody tr a, tbody tr button:not([disabled])')
|
const firstFocusable = orgTableRef.value?.querySelector('tbody tr a, tbody tr button:not([disabled])')
|
||||||
@@ -249,9 +249,9 @@ defineExpose({ focusFirstElement })
|
|||||||
|
|
||||||
<template>
|
<template>
|
||||||
<div class="permissions-section" ref="orgSection">
|
<div class="permissions-section" ref="orgSection">
|
||||||
<h2>{{ isGlobalAdmin ? 'Organizations' : 'Your Organizations' }}</h2>
|
<h2>{{ isMasterAdmin ? 'Organizations' : 'Your Organizations' }}</h2>
|
||||||
<div class="actions" ref="orgActionsRef" @keydown="handleOrgActionsKeydown">
|
<div class="actions" ref="orgActionsRef" @keydown="handleOrgActionsKeydown">
|
||||||
<button v-if="isGlobalAdmin" @click="$emit('createOrg')">+ Create Org</button>
|
<button v-if="isMasterAdmin" @click="$emit('createOrg')">+ Create Org</button>
|
||||||
</div>
|
</div>
|
||||||
<table class="org-table" ref="orgTableRef" @keydown="e => handleTableKeydown(e, 'org')">
|
<table class="org-table" ref="orgTableRef" @keydown="e => handleTableKeydown(e, 'org')">
|
||||||
<thead>
|
<thead>
|
||||||
@@ -259,18 +259,18 @@ defineExpose({ focusFirstElement })
|
|||||||
<th>Name</th>
|
<th>Name</th>
|
||||||
<th>Roles</th>
|
<th>Roles</th>
|
||||||
<th>Members</th>
|
<th>Members</th>
|
||||||
<th v-if="isGlobalAdmin">Actions</th>
|
<th v-if="isMasterAdmin">Actions</th>
|
||||||
</tr>
|
</tr>
|
||||||
</thead>
|
</thead>
|
||||||
<tbody>
|
<tbody>
|
||||||
<tr v-for="o in sortedOrgs" :key="o.uuid">
|
<tr v-for="o in sortedOrgs" :key="o.uuid">
|
||||||
<td>
|
<td>
|
||||||
<a href="#org/{{o.uuid}}" @click.prevent="$emit('openOrg', o)">{{ o.display_name }}</a>
|
<a href="#org/{{o.uuid}}" @click.prevent="$emit('openOrg', o)">{{ o.display_name }}</a>
|
||||||
<button v-if="isGlobalAdmin || isOrgAdmin" @click="$emit('updateOrg', o)" class="icon-btn edit-org-btn" aria-label="Rename organization" title="Rename organization">✏️</button>
|
<button v-if="isMasterAdmin || isOrgAdmin" @click="$emit('updateOrg', o)" class="icon-btn edit-org-btn" aria-label="Rename organization" title="Rename organization">✏️</button>
|
||||||
</td>
|
</td>
|
||||||
<td class="role-names">{{ getRoleNames(o) }}</td>
|
<td class="role-names">{{ getRoleNames(o) }}</td>
|
||||||
<td class="center">{{ o.roles.reduce((acc,r)=>acc + r.users.length,0) }}</td>
|
<td class="center">{{ o.roles.reduce((acc,r)=>acc + r.users.length,0) }}</td>
|
||||||
<td v-if="isGlobalAdmin" class="center">
|
<td v-if="isMasterAdmin" class="center">
|
||||||
<button @click="$emit('deleteOrg', o)" class="icon-btn delete-icon" aria-label="Delete organization" title="Delete organization">❌</button>
|
<button @click="$emit('deleteOrg', o)" class="icon-btn delete-icon" aria-label="Delete organization" title="Delete organization">❌</button>
|
||||||
</td>
|
</td>
|
||||||
</tr>
|
</tr>
|
||||||
@@ -278,7 +278,7 @@ defineExpose({ focusFirstElement })
|
|||||||
</table>
|
</table>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
<div v-if="isGlobalAdmin" class="permissions-section">
|
<div v-if="isMasterAdmin" class="permissions-section">
|
||||||
<h2>Permissions</h2>
|
<h2>Permissions</h2>
|
||||||
<div class="matrix-wrapper" ref="permMatrixRef" @keydown="handleMatrixKeydown">
|
<div class="matrix-wrapper" ref="permMatrixRef" @keydown="handleMatrixKeydown">
|
||||||
<div class="matrix-scroll">
|
<div class="matrix-scroll">
|
||||||
@@ -317,7 +317,7 @@ defineExpose({ focusFirstElement })
|
|||||||
<p class="matrix-hint muted">Toggle which permissions each organization can grant to its members.</p>
|
<p class="matrix-hint muted">Toggle which permissions each organization can grant to its members.</p>
|
||||||
</div>
|
</div>
|
||||||
<div class="actions" ref="permActionsRef" @keydown="handlePermActionsKeydown">
|
<div class="actions" ref="permActionsRef" @keydown="handlePermActionsKeydown">
|
||||||
<button v-if="isGlobalAdmin" @click="$emit('openDialog', 'perm-create', { display_name: '', scope: '', domain: '' })">+ Create Permission</button>
|
<button v-if="isMasterAdmin" @click="$emit('openDialog', 'perm-create', { display_name: '', scope: '', domain: '' })">+ Create Permission</button>
|
||||||
</div>
|
</div>
|
||||||
<table class="org-table" ref="permTableRef" @keydown="e => handleTableKeydown(e, 'perm')">
|
<table class="org-table" ref="permTableRef" @keydown="e => handleTableKeydown(e, 'perm')">
|
||||||
<thead>
|
<thead>
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ function handleEditName() {
|
|||||||
|
|
||||||
async function handleDelete(credential) {
|
async function handleDelete(credential) {
|
||||||
try {
|
try {
|
||||||
const data = await apiJson(`/auth/api/admin/orgs/${props.selectedUser.org_uuid}/users/${props.selectedUser.uuid}/credentials/${credential.credential_uuid}`, { method: 'DELETE' })
|
const data = await apiJson(`/auth/api/admin/orgs/${props.selectedUser.org}/users/${props.selectedUser.uuid}/credentials/${credential.credential}`, { method: 'DELETE' })
|
||||||
if (data.status === 'ok') {
|
if (data.status === 'ok') {
|
||||||
emit('onUserNameSaved') // Reuse to refresh user detail
|
emit('onUserNameSaved') // Reuse to refresh user detail
|
||||||
} else {
|
} else {
|
||||||
@@ -61,7 +61,7 @@ async function handleTerminateSession(session) {
|
|||||||
if (!sessionId) return
|
if (!sessionId) return
|
||||||
terminatingSessions.value = { ...terminatingSessions.value, [sessionId]: true }
|
terminatingSessions.value = { ...terminatingSessions.value, [sessionId]: true }
|
||||||
try {
|
try {
|
||||||
const data = await apiJson(`/auth/api/admin/orgs/${props.selectedUser.org_uuid}/users/${props.selectedUser.uuid}/sessions/${sessionId}`, { method: 'DELETE' })
|
const data = await apiJson(`/auth/api/admin/orgs/${props.selectedUser.org}/users/${props.selectedUser.uuid}/sessions/${sessionId}`, { method: 'DELETE' })
|
||||||
if (data.status === 'ok') {
|
if (data.status === 'ok') {
|
||||||
if (data.current_session_terminated) {
|
if (data.current_session_terminated) {
|
||||||
sessionStorage.clear()
|
sessionStorage.clear()
|
||||||
@@ -183,7 +183,7 @@ defineExpose({ focusFirstElement })
|
|||||||
:loading="loading"
|
:loading="loading"
|
||||||
:org-display-name="userDetail.org.display_name"
|
:org-display-name="userDetail.org.display_name"
|
||||||
:role-name="userDetail.role"
|
:role-name="userDetail.role"
|
||||||
:update-endpoint="`/auth/api/admin/orgs/${selectedUser.org_uuid}/users/${selectedUser.uuid}/display-name`"
|
:update-endpoint="`/auth/api/admin/orgs/${selectedUser.org}/users/${selectedUser.uuid}/display-name`"
|
||||||
@saved="$emit('onUserNameSaved')"
|
@saved="$emit('onUserNameSaved')"
|
||||||
@edit-name="handleEditName"
|
@edit-name="handleEditName"
|
||||||
/>
|
/>
|
||||||
@@ -212,7 +212,7 @@ defineExpose({ focusFirstElement })
|
|||||||
:aaguid-info="userDetail.aaguid_info"
|
:aaguid-info="userDetail.aaguid_info"
|
||||||
:allow-delete="true"
|
:allow-delete="true"
|
||||||
:hovered-credential-uuid="hoveredCredentialUuid"
|
:hovered-credential-uuid="hoveredCredentialUuid"
|
||||||
:hovered-session-credential-uuid="hoveredSession?.credential_uuid"
|
:hovered-session-credential-uuid="hoveredSession?.credential"
|
||||||
:navigation-disabled="hasActiveModal"
|
:navigation-disabled="hasActiveModal"
|
||||||
@delete="handleDelete"
|
@delete="handleDelete"
|
||||||
@credential-hover="hoveredCredentialUuid = $event"
|
@credential-hover="hoveredCredentialUuid = $event"
|
||||||
@@ -238,7 +238,7 @@ defineExpose({ focusFirstElement })
|
|||||||
</div>
|
</div>
|
||||||
<RegistrationLinkModal
|
<RegistrationLinkModal
|
||||||
v-if="showRegModal"
|
v-if="showRegModal"
|
||||||
:endpoint="`/auth/api/admin/orgs/${selectedUser.org_uuid}/users/${selectedUser.uuid}/create-link`"
|
:endpoint="`/auth/api/admin/orgs/${selectedUser.org}/users/${selectedUser.uuid}/create-link`"
|
||||||
:user-name="userDetail?.display_name || selectedUser.display_name"
|
:user-name="userDetail?.display_name || selectedUser.display_name"
|
||||||
@close="$emit('closeRegModal')"
|
@close="$emit('closeRegModal')"
|
||||||
@copied="onLinkCopied"
|
@copied="onLinkCopied"
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
<template>
|
<template>
|
||||||
<div class="message-container">
|
<div class="message-container">
|
||||||
<div class="message-content">
|
<div class="message-content">
|
||||||
<h2>🔒 Access Denied</h2>
|
<h2>{{ icon }} {{ title }}</h2>
|
||||||
|
<p v-if="message" class="error-detail">{{ message }}</p>
|
||||||
<div class="button-row">
|
<div class="button-row">
|
||||||
<button class="btn-secondary" @click="goBack">Back</button>
|
<button class="btn-secondary" @click="goBack">Back</button>
|
||||||
<button class="btn-primary" @click="$emit('reload')">Reload Page</button>
|
<button class="btn-primary" @click="reload">Reload Page</button>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -13,7 +14,15 @@
|
|||||||
<script setup>
|
<script setup>
|
||||||
import { goBack } from '@/utils/helpers'
|
import { goBack } from '@/utils/helpers'
|
||||||
|
|
||||||
defineEmits(['reload'])
|
const props = defineProps({
|
||||||
|
title: { type: String, default: 'Access Denied' },
|
||||||
|
icon: { type: String, default: '🔒' },
|
||||||
|
message: { type: String, default: null },
|
||||||
|
})
|
||||||
|
|
||||||
|
function reload() {
|
||||||
|
window.location.reload()
|
||||||
|
}
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
<style scoped>
|
<style scoped>
|
||||||
@@ -32,10 +41,15 @@ defineEmits(['reload'])
|
|||||||
}
|
}
|
||||||
|
|
||||||
.message-content h2 {
|
.message-content h2 {
|
||||||
margin: 0 0 1.5rem;
|
margin: 0 0 1rem;
|
||||||
color: var(--color-heading);
|
color: var(--color-heading);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.message-content .error-detail {
|
||||||
|
margin: 0 0 1.5rem;
|
||||||
|
color: var(--color-text-muted);
|
||||||
|
}
|
||||||
|
|
||||||
.message-content .button-row {
|
.message-content .button-row {
|
||||||
display: flex;
|
display: flex;
|
||||||
gap: 0.75rem;
|
gap: 0.75rem;
|
||||||
|
|||||||
@@ -5,16 +5,16 @@
|
|||||||
<template v-else>
|
<template v-else>
|
||||||
<div
|
<div
|
||||||
v-for="credential in credentials"
|
v-for="credential in credentials"
|
||||||
:key="credential.credential_uuid"
|
:key="credential.credential"
|
||||||
:class="['credential-item', {
|
:class="['credential-item', {
|
||||||
'current-session': credential.is_current_session && !hoveredCredentialUuid && !hoveredSessionCredentialUuid,
|
'current-session': credential.is_current_session && !hoveredCredentialUuid && !hoveredSessionCredentialUuid,
|
||||||
'is-hovered': hoveredCredentialUuid === credential.credential_uuid,
|
'is-hovered': hoveredCredentialUuid === credential.credential,
|
||||||
'is-linked-session': hoveredSessionCredentialUuid === credential.credential_uuid
|
'is-linked-session': hoveredSessionCredentialUuid === credential.credential
|
||||||
}]"
|
}]"
|
||||||
tabindex="-1"
|
tabindex="-1"
|
||||||
@mousedown.prevent
|
@mousedown.prevent
|
||||||
@click.capture="handleCardClick"
|
@click.capture="handleCardClick"
|
||||||
@focusin="handleCredentialFocus(credential.credential_uuid)"
|
@focusin="handleCredentialFocus(credential.credential)"
|
||||||
@focusout="handleCredentialBlur($event)"
|
@focusout="handleCredentialBlur($event)"
|
||||||
@keydown="handleItemKeydown($event, credential)"
|
@keydown="handleItemKeydown($event, credential)"
|
||||||
>
|
>
|
||||||
@@ -33,8 +33,8 @@
|
|||||||
<h4 class="item-title">{{ getCredentialAuthName(credential) }}</h4>
|
<h4 class="item-title">{{ getCredentialAuthName(credential) }}</h4>
|
||||||
<div class="item-actions">
|
<div class="item-actions">
|
||||||
<span v-if="credential.is_current_session && !hoveredCredentialUuid && !hoveredSessionCredentialUuid" class="badge badge-current">Current</span>
|
<span v-if="credential.is_current_session && !hoveredCredentialUuid && !hoveredSessionCredentialUuid" class="badge badge-current">Current</span>
|
||||||
<span v-else-if="hoveredCredentialUuid === credential.credential_uuid" class="badge badge-current">Selected</span>
|
<span v-else-if="hoveredCredentialUuid === credential.credential" class="badge badge-current">Selected</span>
|
||||||
<span v-else-if="hoveredSessionCredentialUuid === credential.credential_uuid" class="badge badge-current">Linked</span>
|
<span v-else-if="hoveredSessionCredentialUuid === credential.credential" class="badge badge-current">Linked</span>
|
||||||
<button
|
<button
|
||||||
v-if="allowDelete"
|
v-if="allowDelete"
|
||||||
@click="$emit('delete', credential)"
|
@click="$emit('delete', credential)"
|
||||||
|
|||||||
@@ -8,11 +8,11 @@
|
|||||||
<section class="section-block" ref="userInfoSection">
|
<section class="section-block" ref="userInfoSection">
|
||||||
<div class="section-body">
|
<div class="section-body">
|
||||||
<UserBasicInfo
|
<UserBasicInfo
|
||||||
v-if="user"
|
v-if="ctx"
|
||||||
:name="user.user_name"
|
:name="ctx.user.display_name"
|
||||||
:visits="user.visits || 0"
|
:visits="authStore.userInfo?.visits || 0"
|
||||||
:created-at="user.created_at"
|
:created-at="authStore.userInfo?.created_at"
|
||||||
:last-seen="user.last_seen"
|
:last-seen="authStore.userInfo?.last_seen"
|
||||||
:org-display-name="orgDisplayName"
|
:org-display-name="orgDisplayName"
|
||||||
:role-name="roleDisplayName"
|
:role-name="roleDisplayName"
|
||||||
:can-edit="false"
|
:can-edit="false"
|
||||||
@@ -78,9 +78,9 @@ const currentHost = window.location.host
|
|||||||
const userInfoSection = ref(null)
|
const userInfoSection = ref(null)
|
||||||
const buttonRow = ref(null)
|
const buttonRow = ref(null)
|
||||||
|
|
||||||
const user = computed(() => authStore.userInfo?.user || null)
|
const ctx = computed(() => authStore.userInfo?.ctx || null)
|
||||||
const orgDisplayName = computed(() => authStore.userInfo?.org?.display_name || '')
|
const orgDisplayName = computed(() => ctx.value?.org.display_name ?? '')
|
||||||
const roleDisplayName = computed(() => authStore.userInfo?.role?.display_name || '')
|
const roleDisplayName = computed(() => ctx.value?.role.display_name ?? '')
|
||||||
|
|
||||||
const headingTitle = computed(() => {
|
const headingTitle = computed(() => {
|
||||||
const service = authStore.settings?.rp_name
|
const service = authStore.settings?.rp_name
|
||||||
|
|||||||
@@ -8,12 +8,12 @@
|
|||||||
|
|
||||||
<section class="section-block" ref="userInfoSection">
|
<section class="section-block" ref="userInfoSection">
|
||||||
<UserBasicInfo
|
<UserBasicInfo
|
||||||
v-if="authStore.userInfo?.user"
|
v-if="authStore.userInfo?.ctx"
|
||||||
ref="userBasicInfo"
|
ref="userBasicInfo"
|
||||||
:name="authStore.userInfo.user.user_name"
|
:name="authStore.userInfo.ctx.user.display_name"
|
||||||
:visits="authStore.userInfo.user.visits || 0"
|
:visits="authStore.userInfo.visits"
|
||||||
:created-at="authStore.userInfo.user.created_at"
|
:created-at="authStore.userInfo.created_at"
|
||||||
:last-seen="authStore.userInfo.user.last_seen"
|
:last-seen="authStore.userInfo.last_seen"
|
||||||
:loading="authStore.isLoading"
|
:loading="authStore.isLoading"
|
||||||
update-endpoint="/auth/api/user/display-name"
|
update-endpoint="/auth/api/user/display-name"
|
||||||
@saved="authStore.loadUserInfo()"
|
@saved="authStore.loadUserInfo()"
|
||||||
@@ -47,7 +47,7 @@
|
|||||||
: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_uuid"
|
:hovered-session-credential-uuid="hoveredSession?.credential"
|
||||||
:navigation-disabled="hasActiveModal"
|
:navigation-disabled="hasActiveModal"
|
||||||
allow-delete
|
allow-delete
|
||||||
@delete="handleDelete"
|
@delete="handleDelete"
|
||||||
@@ -151,7 +151,7 @@ const userInfoSection = ref(null)
|
|||||||
// Check if any modal/dialog is open (blocks arrow key navigation)
|
// Check if any modal/dialog is open (blocks arrow key navigation)
|
||||||
const hasActiveModal = computed(() => showNameDialog.value || showRegLink.value)
|
const hasActiveModal = computed(() => showNameDialog.value || showRegLink.value)
|
||||||
|
|
||||||
watch(showNameDialog, (newVal) => { if (newVal) newName.value = authStore.userInfo?.user?.user_name || '' })
|
watch(showNameDialog, (newVal) => { if (newVal) newName.value = authStore.userInfo?.ctx.user.display_name ?? '' })
|
||||||
|
|
||||||
onMounted(() => {
|
onMounted(() => {
|
||||||
updateInterval.value = setInterval(() => { if (authStore.userInfo) authStore.userInfo = { ...authStore.userInfo } }, 60000)
|
updateInterval.value = setInterval(() => { if (authStore.userInfo) authStore.userInfo = { ...authStore.userInfo } }, 60000)
|
||||||
@@ -292,7 +292,7 @@ const handleLogoutButtonKeydown = (event) => {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const handleDelete = async (credential) => {
|
const handleDelete = async (credential) => {
|
||||||
const credentialId = credential?.credential_uuid
|
const credentialId = credential?.credential
|
||||||
if (!credentialId) return
|
if (!credentialId) return
|
||||||
try {
|
try {
|
||||||
await authStore.deleteCredential(credentialId)
|
await authStore.deleteCredential(credentialId)
|
||||||
@@ -323,9 +323,9 @@ const terminateSession = async (session) => {
|
|||||||
|
|
||||||
const logoutEverywhere = async () => { await authStore.logoutEverywhere() }
|
const logoutEverywhere = async () => { await authStore.logoutEverywhere() }
|
||||||
const logout = async () => { await authStore.logout() }
|
const logout = async () => { await authStore.logout() }
|
||||||
const openNameDialog = () => { newName.value = authStore.userInfo?.user?.user_name || ''; showNameDialog.value = true }
|
const openNameDialog = () => { newName.value = authStore.userInfo?.ctx.user.display_name ?? ''; showNameDialog.value = true }
|
||||||
const isAdmin = computed(() => {
|
const isAdmin = computed(() => {
|
||||||
const perms = authStore.userInfo?.permissions ?? []
|
const perms = authStore.userInfo?.ctx.permissions
|
||||||
return perms.includes('auth:admin') || perms.includes('auth:org:admin')
|
return perms.includes('auth:admin') || perms.includes('auth:org:admin')
|
||||||
})
|
})
|
||||||
const hasMultipleSessions = computed(() => sessions.value.length > 1)
|
const hasMultipleSessions = computed(() => sessions.value.length > 1)
|
||||||
|
|||||||
@@ -76,13 +76,13 @@ const status = reactive({ show: false, message: '', type: 'info' })
|
|||||||
const initializing = ref(true)
|
const initializing = ref(true)
|
||||||
const loading = ref(false)
|
const loading = ref(false)
|
||||||
const settings = ref(null)
|
const settings = ref(null)
|
||||||
const userInfo = ref(null)
|
const session = ref(null)
|
||||||
const currentView = ref('initial') // 'initial', 'login', 'forbidden'
|
const currentView = ref('initial') // 'initial', 'login', 'forbidden'
|
||||||
const authView = ref('local') // 'local' or 'remote'
|
const authView = ref('local') // 'local' or 'remote'
|
||||||
const buttonRow = ref(null)
|
const buttonRow = ref(null)
|
||||||
let statusTimer = null
|
let statusTimer = null
|
||||||
|
|
||||||
const isAuthenticated = computed(() => !!userInfo.value?.authenticated)
|
const isAuthenticated = computed(() => !!session.value)
|
||||||
|
|
||||||
const canAuthenticate = computed(() => {
|
const canAuthenticate = computed(() => {
|
||||||
if (initializing.value) return false
|
if (initializing.value) return false
|
||||||
@@ -115,7 +115,7 @@ const headerMessage = computed(() => {
|
|||||||
return 'Please sign in with your passkey.'
|
return 'Please sign in with your passkey.'
|
||||||
})
|
})
|
||||||
|
|
||||||
const userDisplayName = computed(() => userInfo.value?.user?.user_name || 'User')
|
const userDisplayName = computed(() => session.value?.ctx.user.display_name || 'User')
|
||||||
|
|
||||||
function showMessage(message, type = 'info', duration = 3000) {
|
function showMessage(message, type = 'info', duration = 3000) {
|
||||||
status.show = true
|
status.show = true
|
||||||
@@ -140,22 +140,21 @@ async function fetchSettings() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async function fetchUserInfo() {
|
async function validateSession() {
|
||||||
try {
|
try {
|
||||||
userInfo.value = await fetchJson('/auth/api/user-info', { method: 'POST' })
|
session.value = await fetchJson('/auth/api/validate', { method: 'POST' })
|
||||||
if (isAuthenticated.value && props.mode !== 'reauth') {
|
if (isAuthenticated.value && props.mode !== 'reauth') {
|
||||||
currentView.value = 'forbidden'
|
currentView.value = 'forbidden'
|
||||||
emit('forbidden', userInfo.value)
|
emit('forbidden', session.value)
|
||||||
} else {
|
} else {
|
||||||
currentView.value = 'login'
|
currentView.value = 'login'
|
||||||
}
|
}
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
console.error('Failed to load user info', error)
|
session.value = null
|
||||||
|
currentView.value = 'login'
|
||||||
if (error.status !== 401 && error.status !== 403) {
|
if (error.status !== 401 && error.status !== 403) {
|
||||||
showMessage(getUserFriendlyErrorMessage(error), 'error', 4000)
|
showMessage(getUserFriendlyErrorMessage(error), 'error', 4000)
|
||||||
}
|
}
|
||||||
userInfo.value = null
|
|
||||||
currentView.value = 'login'
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,7 +187,7 @@ async function logoutUser() {
|
|||||||
loading.value = true
|
loading.value = true
|
||||||
try {
|
try {
|
||||||
await fetchJson('/auth/api/logout', { method: 'POST' })
|
await fetchJson('/auth/api/logout', { method: 'POST' })
|
||||||
userInfo.value = null
|
session.value = null
|
||||||
currentView.value = 'login'
|
currentView.value = 'login'
|
||||||
showMessage('Logged out. You can sign in with a different account.', 'info', 3000)
|
showMessage('Logged out. You can sign in with a different account.', 'info', 3000)
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
@@ -266,7 +265,7 @@ watch(initializing, (newVal) => {
|
|||||||
|
|
||||||
onMounted(async () => {
|
onMounted(async () => {
|
||||||
await fetchSettings()
|
await fetchSettings()
|
||||||
await fetchUserInfo()
|
await validateSession()
|
||||||
initializing.value = false
|
initializing.value = false
|
||||||
|
|
||||||
// Add click handler for inline links
|
// Add click handler for inline links
|
||||||
@@ -280,7 +279,7 @@ onUnmounted(() => {
|
|||||||
defineExpose({
|
defineExpose({
|
||||||
showMessage,
|
showMessage,
|
||||||
isAuthenticated,
|
isAuthenticated,
|
||||||
userInfo
|
session
|
||||||
})
|
})
|
||||||
</script>
|
</script>
|
||||||
|
|
||||||
|
|||||||
@@ -20,7 +20,7 @@
|
|||||||
:class="['session-item', {
|
:class="['session-item', {
|
||||||
'is-current': session.is_current && !hoveredIp && !hoveredCredentialUuid,
|
'is-current': session.is_current && !hoveredIp && !hoveredCredentialUuid,
|
||||||
'is-hovered': hoveredSession?.id === session.id,
|
'is-hovered': hoveredSession?.id === session.id,
|
||||||
'is-linked-credential': hoveredCredentialUuid === session.credential_uuid
|
'is-linked-credential': hoveredCredentialUuid === session.credential
|
||||||
}]"
|
}]"
|
||||||
tabindex="-1"
|
tabindex="-1"
|
||||||
@mousedown.prevent
|
@mousedown.prevent
|
||||||
@@ -34,7 +34,7 @@
|
|||||||
<div class="item-actions">
|
<div class="item-actions">
|
||||||
<span v-if="session.is_current && !hoveredIp && !hoveredCredentialUuid" class="badge badge-current">Current</span>
|
<span v-if="session.is_current && !hoveredIp && !hoveredCredentialUuid" class="badge badge-current">Current</span>
|
||||||
<span v-else-if="hoveredSession?.id === session.id" class="badge badge-current">Selected</span>
|
<span v-else-if="hoveredSession?.id === session.id" class="badge badge-current">Selected</span>
|
||||||
<span v-else-if="hoveredCredentialUuid === session.credential_uuid" class="badge badge-current">Linked</span>
|
<span v-else-if="hoveredCredentialUuid === session.credential" class="badge badge-current">Linked</span>
|
||||||
<span v-else-if="!hoveredCredentialUuid && isSameHost(session.ip)" class="badge">Same IP</span>
|
<span v-else-if="!hoveredCredentialUuid && isSameHost(session.ip)" class="badge">Same IP</span>
|
||||||
<button
|
<button
|
||||||
@click="$emit('terminate', session)"
|
@click="$emit('terminate', session)"
|
||||||
|
|||||||
@@ -88,7 +88,7 @@ export async function getAuthIframeUrl(mode = 'login') {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Fetch from forward endpoint - it returns URL in auth.iframe on 401/403
|
// Fetch from forward endpoint - it returns URL in auth.iframe on 401/403
|
||||||
const response = await fetch('/auth/api/forward', { credentials: 'include' })
|
const response = await fetch('/auth/api/forward')
|
||||||
if (response.status === 401 || response.status === 403) {
|
if (response.status === 401 || response.status === 403) {
|
||||||
const data = await response.json()
|
const data = await response.json()
|
||||||
if (data.auth?.iframe) {
|
if (data.auth?.iframe) {
|
||||||
@@ -321,7 +321,6 @@ export async function apiJson(url, options = {}) {
|
|||||||
*/
|
*/
|
||||||
export async function fetchJson(url, options = {}) {
|
export async function fetchJson(url, options = {}) {
|
||||||
const fetchOptions = {
|
const fetchOptions = {
|
||||||
credentials: 'include',
|
|
||||||
...options,
|
...options,
|
||||||
headers: {
|
headers: {
|
||||||
'Accept': 'application/json',
|
'Accept': 'application/json',
|
||||||
|
|||||||
@@ -1,55 +0,0 @@
|
|||||||
# auto-upgrade@fastapi-vue-setup - remove this if you modify this file
|
|
||||||
import argparse
|
|
||||||
import asyncio
|
|
||||||
import os
|
|
||||||
|
|
||||||
import uvicorn
|
|
||||||
from fastapi_vue.hostutil import parse_endpoint
|
|
||||||
from uvicorn import Config, Server
|
|
||||||
|
|
||||||
from .APP_MODULE import APP_VAR
|
|
||||||
|
|
||||||
DEFAULT_PORT = 5080
|
|
||||||
|
|
||||||
|
|
||||||
def run_server(endpoints: list[dict], *, proxy="", devmode=False):
|
|
||||||
conf: dict[str, object] = {"app": "MODULE_NAME.APP_MODULE:APP_VAR"}
|
|
||||||
if proxy:
|
|
||||||
conf["proxy_headers"] = True
|
|
||||||
conf["forwarded_allow_ips"] = proxy
|
|
||||||
if devmode:
|
|
||||||
conf["reload"] = True
|
|
||||||
conf["reload_dirs"] = ["MODULE_NAME"]
|
|
||||||
APP_VAR.debug = True
|
|
||||||
|
|
||||||
if len(endpoints) > 1:
|
|
||||||
# Run separate servers for multiple endpoints
|
|
||||||
async def serve_all():
|
|
||||||
async with asyncio.TaskGroup() as tg:
|
|
||||||
for ep in endpoints:
|
|
||||||
tg.create_task(Server(Config(**conf, **ep)).serve())
|
|
||||||
|
|
||||||
asyncio.run(serve_all())
|
|
||||||
else:
|
|
||||||
uvicorn.run(**conf, **endpoints[0])
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
|
||||||
parser = argparse.ArgumentParser(description="Run the MODULE_NAME server.")
|
|
||||||
parser.add_argument(
|
|
||||||
"endpoint",
|
|
||||||
nargs="?",
|
|
||||||
help=(
|
|
||||||
f"Endpoint (default: localhost:{DEFAULT_PORT}). "
|
|
||||||
"Forms: host:port | :port | [ipv6]:port | ip | host | unix:/path.sock"
|
|
||||||
),
|
|
||||||
)
|
|
||||||
args = parser.parse_args()
|
|
||||||
proxy = os.getenv("FORWARDED_ALLOW_IPS", "127.0.0.1,::1")
|
|
||||||
devmode = bool(os.getenv("FASTAPI_VUE_FRONTEND_URL"))
|
|
||||||
endpoints = parse_endpoint(args.endpoint, DEFAULT_PORT)
|
|
||||||
run_server(endpoints, proxy=proxy, devmode=devmode)
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -10,6 +10,7 @@ This module provides functionality to:
|
|||||||
import json
|
import json
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
from importlib.resources import files
|
from importlib.resources import files
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
__ALL__ = ["AAGUID", "filter"]
|
__ALL__ = ["AAGUID", "filter"]
|
||||||
|
|
||||||
@@ -18,15 +19,15 @@ AAGUID_FILE = files("paskia") / "aaguid" / "combined_aaguid.json"
|
|||||||
AAGUID: dict[str, dict] = json.loads(AAGUID_FILE.read_text(encoding="utf-8"))
|
AAGUID: dict[str, dict] = json.loads(AAGUID_FILE.read_text(encoding="utf-8"))
|
||||||
|
|
||||||
|
|
||||||
def filter(aaguids: Iterable[str]) -> dict[str, dict]:
|
def filter(aaguids: Iterable[UUID]) -> dict[str, dict]:
|
||||||
"""
|
"""
|
||||||
Get AAGUID information only for the provided set of AAGUIDs.
|
Get AAGUID information only for the provided set of AAGUIDs.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
aaguids: Set of AAGUID strings that the user has credentials for
|
aaguids: Iterable of AAGUIDs (UUIDs) that the user has credentials for
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary mapping AAGUID to authenticator information for only
|
Dictionary mapping AAGUID string to authenticator information for only
|
||||||
the AAGUIDs that the user has and that we have data for
|
the AAGUIDs that the user has and that we have data for
|
||||||
"""
|
"""
|
||||||
return {aaguid: AAGUID[aaguid] for aaguid in aaguids if aaguid in AAGUID}
|
return {(s := str(a)): AAGUID[s] for a in aaguids if (s := str(a)) in AAGUID}
|
||||||
|
|||||||
+15
-43
@@ -8,68 +8,40 @@ independent of any web framework:
|
|||||||
- Credential management
|
- Credential management
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.config import SESSION_LIFETIME
|
from paskia.config import RESET_LIFETIME, SESSION_LIFETIME
|
||||||
from paskia.db import ResetToken, Session
|
|
||||||
from paskia.util import hostutil
|
from paskia.util import hostutil
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from paskia.db import ResetToken
|
||||||
|
|
||||||
EXPIRES = SESSION_LIFETIME
|
EXPIRES = SESSION_LIFETIME
|
||||||
|
|
||||||
|
|
||||||
def expires() -> datetime:
|
def expires() -> datetime:
|
||||||
return datetime.now(timezone.utc) + EXPIRES
|
return datetime.now(UTC) + EXPIRES
|
||||||
|
|
||||||
|
|
||||||
def reset_expires() -> datetime:
|
def reset_expires() -> datetime:
|
||||||
from .config import RESET_LIFETIME
|
return datetime.now(UTC) + RESET_LIFETIME
|
||||||
|
|
||||||
return datetime.now(timezone.utc) + RESET_LIFETIME
|
|
||||||
|
|
||||||
|
|
||||||
async def get_reset(token: str) -> ResetToken:
|
def get_reset(token: str) -> "ResetToken":
|
||||||
"""Validate a credential reset token."""
|
"""Validate a credential reset token."""
|
||||||
|
|
||||||
record = db.get_reset_token(token)
|
record = db.get_reset_token(token)
|
||||||
if record:
|
if record:
|
||||||
return record
|
return record
|
||||||
raise ValueError("This authentication link is no longer valid.")
|
raise ValueError("This authentication link is no longer valid.")
|
||||||
|
|
||||||
|
|
||||||
async def get_session(token: str, host: str | None = None) -> Session:
|
def delete_credential(credential_uuid: UUID, auth: str, host: str | None = None):
|
||||||
"""Validate a session token and return session data if valid."""
|
|
||||||
host = hostutil.normalize_host(host)
|
|
||||||
if not host:
|
|
||||||
raise ValueError("Invalid host")
|
|
||||||
session = db.get_session(token)
|
|
||||||
if session:
|
|
||||||
if session.host is None:
|
|
||||||
# First time binding: store exact host:port (or IPv6 form) now.
|
|
||||||
db.set_session_host(session.key, host)
|
|
||||||
session.host = host
|
|
||||||
elif session.host != host:
|
|
||||||
raise ValueError("Session host mismatch")
|
|
||||||
return session
|
|
||||||
raise ValueError("Your session has expired. Please sign in again!")
|
|
||||||
|
|
||||||
|
|
||||||
async def refresh_session_token(token: str, *, ip: str, user_agent: str):
|
|
||||||
"""Refresh a session extending its expiry."""
|
|
||||||
session_record = db.get_session(token)
|
|
||||||
if not session_record:
|
|
||||||
raise ValueError("Session not found or expired")
|
|
||||||
updated = db.update_session(
|
|
||||||
token,
|
|
||||||
ip=ip,
|
|
||||||
user_agent=user_agent,
|
|
||||||
expiry=expires(),
|
|
||||||
)
|
|
||||||
if not updated:
|
|
||||||
raise ValueError("Session not found or expired")
|
|
||||||
|
|
||||||
|
|
||||||
async 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."""
|
||||||
s = await get_session(auth, host=host)
|
ctx = db.data().session_ctx(auth, hostutil.normalize_host(host))
|
||||||
db.delete_credential(credential_uuid, s.user_uuid)
|
if not ctx:
|
||||||
|
raise ValueError("Session expired")
|
||||||
|
db.delete_credential(credential_uuid, ctx.user.uuid)
|
||||||
|
|||||||
+30
-102
@@ -8,26 +8,11 @@ generating a reset link for initial admin setup.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from datetime import datetime, timezone
|
|
||||||
|
|
||||||
import uuid7
|
from paskia import authsession, db, globals
|
||||||
|
|
||||||
from paskia import authsession, db
|
|
||||||
from paskia.db import Org, Permission, Role, User
|
|
||||||
from paskia.util import hostutil, passphrase
|
from paskia.util import hostutil, passphrase
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
def _init_logger() -> logging.Logger:
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
if not logger.handlers and not logging.getLogger().handlers:
|
|
||||||
h = logging.StreamHandler()
|
|
||||||
h.setFormatter(logging.Formatter("%(message)s"))
|
|
||||||
logger.addHandler(h)
|
|
||||||
logger.setLevel(logging.INFO)
|
|
||||||
return logger
|
|
||||||
|
|
||||||
|
|
||||||
logger = _init_logger()
|
|
||||||
|
|
||||||
# Shared log message template for admin reset links
|
# Shared log message template for admin reset links
|
||||||
ADMIN_RESET_MESSAGE = """\
|
ADMIN_RESET_MESSAGE = """\
|
||||||
@@ -38,84 +23,25 @@ ADMIN_RESET_MESSAGE = """\
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
async def _create_and_log_admin_reset_link(user_uuid, message, session_type) -> str:
|
def _log_reset_link(message: str, passphrase: str) -> str:
|
||||||
"""Create an admin reset link and log it with the provided message."""
|
"""Log a reset link message and return the URL."""
|
||||||
token = passphrase.generate()
|
reset_link = hostutil.reset_link_url(passphrase)
|
||||||
expiry = authsession.reset_expires()
|
|
||||||
db.create_reset_token(
|
|
||||||
user_uuid=user_uuid,
|
|
||||||
passphrase=token,
|
|
||||||
expiry=expiry,
|
|
||||||
token_type=session_type,
|
|
||||||
)
|
|
||||||
reset_link = hostutil.reset_link_url(token)
|
|
||||||
logger.info(ADMIN_RESET_MESSAGE, message, reset_link)
|
logger.info(ADMIN_RESET_MESSAGE, message, reset_link)
|
||||||
return reset_link
|
return reset_link
|
||||||
|
|
||||||
|
|
||||||
async def bootstrap_system() -> dict:
|
async def bootstrap_system() -> None:
|
||||||
"""
|
"""
|
||||||
Bootstrap the entire system with default data.
|
Bootstrap the entire system with default data.
|
||||||
|
|
||||||
Returns:
|
Uses db.bootstrap() which performs all operations in a single transaction.
|
||||||
dict: Contains information about created entities and reset link
|
The transaction log will show a single "bootstrap" action with all changes.
|
||||||
"""
|
"""
|
||||||
# Create permission first - will fail if already exists
|
# Call the single-transaction bootstrap function
|
||||||
perm0 = Permission(
|
reset_passphrase = db.bootstrap()
|
||||||
uuid=uuid7.create(), scope="auth:admin", display_name="Master Admin"
|
|
||||||
)
|
|
||||||
db.create_permission(perm0)
|
|
||||||
|
|
||||||
# Create org admin permission - allows managing users within an org
|
# Log the reset link (this is separate from the transaction log)
|
||||||
perm_org_admin = Permission(
|
_log_reset_link("✅ Bootstrap completed!", reset_passphrase)
|
||||||
uuid=uuid7.create(), scope="auth:org:admin", display_name="Org Admin"
|
|
||||||
)
|
|
||||||
db.create_permission(perm_org_admin)
|
|
||||||
|
|
||||||
org = Org(uuid7.create(), "Organization")
|
|
||||||
db.create_organization(org)
|
|
||||||
|
|
||||||
# Allow this org to grant global admin and org admin permissions
|
|
||||||
db.add_permission_to_organization(str(org.uuid), perm0.scope)
|
|
||||||
db.add_permission_to_organization(str(org.uuid), perm_org_admin.scope)
|
|
||||||
|
|
||||||
# Create an Administration role granting both org and global admin
|
|
||||||
role = Role(
|
|
||||||
uuid7.create(),
|
|
||||||
org.uuid,
|
|
||||||
"Administration",
|
|
||||||
permissions=[perm0.scope, perm_org_admin.scope],
|
|
||||||
)
|
|
||||||
db.create_role(role)
|
|
||||||
|
|
||||||
user = User(
|
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Admin",
|
|
||||||
role_uuid=role.uuid,
|
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
visits=0,
|
|
||||||
)
|
|
||||||
db.create_user(user)
|
|
||||||
|
|
||||||
# Generate reset link and log it
|
|
||||||
reset_link = await _create_and_log_admin_reset_link(
|
|
||||||
user.uuid, "✅ Bootstrap completed!", "admin bootstrap"
|
|
||||||
)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"user": user,
|
|
||||||
"org": org,
|
|
||||||
"role": role,
|
|
||||||
"permissions": [
|
|
||||||
perm0,
|
|
||||||
*[
|
|
||||||
db.get_permission_by_scope(p)
|
|
||||||
for p in org.permissions
|
|
||||||
if db.get_permission_by_scope(p)
|
|
||||||
],
|
|
||||||
],
|
|
||||||
"reset_link": reset_link,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
async def check_admin_credentials() -> bool:
|
async def check_admin_credentials() -> bool:
|
||||||
@@ -127,13 +53,15 @@ async def check_admin_credentials() -> bool:
|
|||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
# Get permission organizations to find admin users
|
# Get permission organizations to find admin users
|
||||||
permission_orgs = db.get_permission_organizations("auth:admin")
|
p = next(
|
||||||
|
(p for p in db.data().permissions.values() if p.scope == "auth:admin"), None
|
||||||
if not permission_orgs:
|
)
|
||||||
|
if not p or not p.orgs:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# Get users from the first organization with admin permission
|
# Get users from the first organization with admin permission
|
||||||
org_users = db.get_organization_users(str(permission_orgs[0].uuid))
|
first_org_uuid = next(iter(p.orgs))
|
||||||
|
org_users = db.get_organization_users(first_org_uuid)
|
||||||
admin_users = [user for user, role in org_users if role == "Administration"]
|
admin_users = [user for user, role in org_users if role == "Administration"]
|
||||||
|
|
||||||
if not admin_users:
|
if not admin_users:
|
||||||
@@ -141,15 +69,19 @@ async def check_admin_credentials() -> bool:
|
|||||||
|
|
||||||
# Check first admin user for credentials
|
# Check first admin user for credentials
|
||||||
admin_user = admin_users[0]
|
admin_user = admin_users[0]
|
||||||
credentials = db.get_credentials_by_user_uuid(admin_user.uuid)
|
|
||||||
|
|
||||||
if not credentials:
|
if not db.get_user_credential_ids(admin_user.uuid):
|
||||||
# Admin exists but has no credentials, create reset link
|
# Admin exists but has no credentials, create reset link
|
||||||
await _create_and_log_admin_reset_link(
|
|
||||||
admin_user.uuid,
|
token = passphrase.generate()
|
||||||
"⚠️ Admin user has no credentials!",
|
expiry = authsession.reset_expires()
|
||||||
"admin registration",
|
db.create_reset_token(
|
||||||
|
user_uuid=admin_user.uuid,
|
||||||
|
passphrase=token,
|
||||||
|
expiry=expiry,
|
||||||
|
token_type="admin registration",
|
||||||
)
|
)
|
||||||
|
_log_reset_link("⚠️ Admin user has no credentials!", token)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
return False
|
return False
|
||||||
@@ -165,16 +97,12 @@ async def bootstrap_if_needed() -> bool:
|
|||||||
Returns:
|
Returns:
|
||||||
bool: True if bootstrapping was performed, False if system was already set up
|
bool: True if bootstrapping was performed, False if system was already set up
|
||||||
"""
|
"""
|
||||||
try:
|
# Check if the admin permission exists - if it does, system is already bootstrapped
|
||||||
# Check if the admin permission exists - if it does, system is already bootstrapped
|
if any(p.scope == "auth:admin" for p in db.data().permissions.values()):
|
||||||
db.get_permission("auth:admin")
|
|
||||||
# Permission exists, system is already bootstrapped
|
# Permission exists, system is already bootstrapped
|
||||||
# Check if admin needs credentials (only for already-bootstrapped systems)
|
# Check if admin needs credentials (only for already-bootstrapped systems)
|
||||||
await check_admin_credentials()
|
await check_admin_credentials()
|
||||||
return False
|
return False
|
||||||
except Exception:
|
|
||||||
# Permission doesn't exist, need to bootstrap
|
|
||||||
pass
|
|
||||||
|
|
||||||
# No admin permission found, need to bootstrap
|
# No admin permission found, need to bootstrap
|
||||||
# Bootstrap creates the admin user AND the reset link, so no need to check credentials after
|
# Bootstrap creates the admin user AND the reset link, so no need to check credentials after
|
||||||
|
|||||||
+27
-55
@@ -1,24 +1,25 @@
|
|||||||
"""
|
"""
|
||||||
Database module for WebAuthn passkey authentication.
|
Database module for WebAuthn passkey authentication.
|
||||||
|
|
||||||
Read: Access _db._data directly, use build_* to convert to public structs.
|
Read: Access data() directly, use build_* to convert to public structs.
|
||||||
CTX: get_session_context(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.
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
from paskia import db
|
from paskia import db
|
||||||
|
|
||||||
# Read (after init)
|
# Read (after init)
|
||||||
user_data = db._db._data.users[user_uuid]
|
user_data = db.data().users[user_uuid]
|
||||||
user = db.build_user(user_uuid)
|
user = db.build_user(user_uuid)
|
||||||
|
|
||||||
# Context
|
# Context
|
||||||
ctx = db.get_session_context(session_key)
|
ctx = db.data().session_ctx(session_key)
|
||||||
|
|
||||||
# Write
|
# Write
|
||||||
db.create_user(user)
|
db.create_user(user)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import paskia.db.operations as operations
|
||||||
from paskia.db.background import (
|
from paskia.db.background import (
|
||||||
start_background,
|
start_background,
|
||||||
start_cleanup,
|
start_cleanup,
|
||||||
@@ -26,59 +27,37 @@ from paskia.db.background import (
|
|||||||
stop_cleanup,
|
stop_cleanup,
|
||||||
)
|
)
|
||||||
from paskia.db.operations import (
|
from paskia.db.operations import (
|
||||||
DB,
|
add_permission_to_org,
|
||||||
_db,
|
|
||||||
add_permission_to_organization,
|
|
||||||
add_permission_to_role,
|
add_permission_to_role,
|
||||||
build_credential,
|
bootstrap,
|
||||||
build_org,
|
|
||||||
build_permission,
|
|
||||||
build_reset_token,
|
|
||||||
build_role,
|
|
||||||
build_session,
|
|
||||||
build_user,
|
|
||||||
cleanup_expired,
|
cleanup_expired,
|
||||||
create_credential,
|
create_credential,
|
||||||
create_credential_session,
|
create_credential_session,
|
||||||
create_organization,
|
create_org,
|
||||||
create_permission,
|
create_permission,
|
||||||
create_reset_token,
|
create_reset_token,
|
||||||
create_role,
|
create_role,
|
||||||
create_session,
|
create_session,
|
||||||
create_user,
|
create_user,
|
||||||
delete_credential,
|
delete_credential,
|
||||||
delete_organization,
|
delete_org,
|
||||||
delete_permission,
|
delete_permission,
|
||||||
delete_reset_token,
|
delete_reset_token,
|
||||||
delete_role,
|
delete_role,
|
||||||
delete_session,
|
delete_session,
|
||||||
delete_sessions_for_user,
|
delete_sessions_for_user,
|
||||||
delete_user,
|
delete_user,
|
||||||
get_credential_by_id,
|
|
||||||
get_credentials_by_user_uuid,
|
|
||||||
get_organization,
|
|
||||||
get_organization_users,
|
get_organization_users,
|
||||||
get_permission,
|
|
||||||
get_permission_by_scope,
|
|
||||||
get_permission_organizations,
|
|
||||||
get_reset_token,
|
get_reset_token,
|
||||||
get_role,
|
get_user_credential_ids,
|
||||||
get_roles_by_organization,
|
|
||||||
get_session,
|
|
||||||
get_session_context,
|
|
||||||
get_user_by_uuid,
|
|
||||||
get_user_organization,
|
get_user_organization,
|
||||||
init,
|
init,
|
||||||
list_organizations,
|
|
||||||
list_permissions,
|
|
||||||
list_sessions_for_user,
|
|
||||||
login,
|
login,
|
||||||
remove_permission_from_organization,
|
remove_permission_from_org,
|
||||||
remove_permission_from_role,
|
remove_permission_from_role,
|
||||||
rename_permission,
|
|
||||||
set_session_host,
|
set_session_host,
|
||||||
update_credential_sign_count,
|
update_credential_sign_count,
|
||||||
update_organization_name,
|
update_org_name,
|
||||||
update_permission,
|
update_permission,
|
||||||
update_role_name,
|
update_role_name,
|
||||||
update_session,
|
update_session,
|
||||||
@@ -87,6 +66,7 @@ from paskia.db.operations import (
|
|||||||
update_user_role_in_organization,
|
update_user_role_in_organization,
|
||||||
)
|
)
|
||||||
from paskia.db.structs import (
|
from paskia.db.structs import (
|
||||||
|
DB,
|
||||||
Credential,
|
Credential,
|
||||||
Org,
|
Org,
|
||||||
Permission,
|
Permission,
|
||||||
@@ -97,6 +77,12 @@ from paskia.db.structs import (
|
|||||||
User,
|
User,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def data() -> DB:
|
||||||
|
"""Get the database instance for direct read access."""
|
||||||
|
return operations._db
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
# Types
|
# Types
|
||||||
"Credential",
|
"Credential",
|
||||||
@@ -109,7 +95,7 @@ __all__ = [
|
|||||||
"SessionContext",
|
"SessionContext",
|
||||||
"User",
|
"User",
|
||||||
# Instance
|
# Instance
|
||||||
"_db",
|
"data",
|
||||||
"init",
|
"init",
|
||||||
# Background
|
# Background
|
||||||
"start_background",
|
"start_background",
|
||||||
@@ -118,44 +104,31 @@ __all__ = [
|
|||||||
"stop_cleanup",
|
"stop_cleanup",
|
||||||
# Builders
|
# Builders
|
||||||
"build_credential",
|
"build_credential",
|
||||||
"build_org",
|
|
||||||
"build_permission",
|
"build_permission",
|
||||||
"build_reset_token",
|
"build_reset_token",
|
||||||
"build_role",
|
"build_role",
|
||||||
"build_session",
|
"build_session",
|
||||||
"build_user",
|
"build_user",
|
||||||
# Read ops
|
# Read ops
|
||||||
"get_credential_by_id",
|
|
||||||
"get_credentials_by_user_uuid",
|
|
||||||
"get_organization",
|
|
||||||
"get_organization_users",
|
"get_organization_users",
|
||||||
"get_permission",
|
|
||||||
"get_permission_by_scope",
|
|
||||||
"get_permission_organizations",
|
|
||||||
"get_reset_token",
|
"get_reset_token",
|
||||||
"get_role",
|
"get_user_credential_ids",
|
||||||
"get_roles_by_organization",
|
|
||||||
"get_session",
|
|
||||||
"get_session_context",
|
|
||||||
"get_user_by_uuid",
|
|
||||||
"get_user_organization",
|
"get_user_organization",
|
||||||
"list_organizations",
|
|
||||||
"list_permissions",
|
|
||||||
"list_sessions_for_user",
|
|
||||||
# Write ops
|
# Write ops
|
||||||
"add_permission_to_organization",
|
"add_permission_to_org",
|
||||||
"add_permission_to_role",
|
"add_permission_to_role",
|
||||||
|
"bootstrap",
|
||||||
"cleanup_expired",
|
"cleanup_expired",
|
||||||
"create_credential",
|
"create_credential",
|
||||||
"create_credential_session",
|
"create_credential_session",
|
||||||
"create_organization",
|
"create_org",
|
||||||
"create_permission",
|
"create_permission",
|
||||||
"create_reset_token",
|
"create_reset_token",
|
||||||
"create_role",
|
"create_role",
|
||||||
"create_session",
|
"create_session",
|
||||||
"create_user",
|
"create_user",
|
||||||
"delete_credential",
|
"delete_credential",
|
||||||
"delete_organization",
|
"delete_org",
|
||||||
"delete_permission",
|
"delete_permission",
|
||||||
"delete_reset_token",
|
"delete_reset_token",
|
||||||
"delete_role",
|
"delete_role",
|
||||||
@@ -163,12 +136,11 @@ __all__ = [
|
|||||||
"delete_sessions_for_user",
|
"delete_sessions_for_user",
|
||||||
"delete_user",
|
"delete_user",
|
||||||
"login",
|
"login",
|
||||||
"remove_permission_from_organization",
|
"remove_permission_from_org",
|
||||||
"remove_permission_from_role",
|
"remove_permission_from_role",
|
||||||
"rename_permission",
|
|
||||||
"set_session_host",
|
"set_session_host",
|
||||||
"update_credential_sign_count",
|
"update_credential_sign_count",
|
||||||
"update_organization_name",
|
"update_org_name",
|
||||||
"update_permission",
|
"update_permission",
|
||||||
"update_role_name",
|
"update_role_name",
|
||||||
"update_session",
|
"update_session",
|
||||||
|
|||||||
+20
-40
@@ -6,62 +6,34 @@ Periodically flushes pending changes to disk and cleans up expired items.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from paskia.db.jsonl import flush_changes
|
from paskia.db.operations import _store, cleanup_expired
|
||||||
|
|
||||||
# Flush changes to disk every N seconds
|
FLUSH_INTERVAL = 0.1 # Flush to disk
|
||||||
FLUSH_INTERVAL = 1
|
CLEANUP_INTERVAL = 1 # Expired item cleanup
|
||||||
# Cleanup expired items every N seconds (cheap when nothing to remove)
|
|
||||||
CLEANUP_INTERVAL = 1
|
|
||||||
|
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
_background_task: asyncio.Task | None = None
|
_background_task: asyncio.Task | None = None
|
||||||
|
|
||||||
|
|
||||||
def cleanup() -> None:
|
|
||||||
"""Remove expired sessions and reset tokens from the database."""
|
|
||||||
from paskia.db.operations import _db
|
|
||||||
|
|
||||||
if _db is None or _db._data is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
with _db.transaction("expiry"):
|
|
||||||
current_time = datetime.now(timezone.utc)
|
|
||||||
|
|
||||||
# Clean expired sessions
|
|
||||||
to_delete_sessions = [
|
|
||||||
k for k, s in _db._data.sessions.items() if s.expiry < current_time
|
|
||||||
]
|
|
||||||
for k in to_delete_sessions:
|
|
||||||
del _db._data.sessions[k]
|
|
||||||
|
|
||||||
# Clean expired reset tokens
|
|
||||||
to_delete_tokens = [
|
|
||||||
k for k, t in _db._data.reset_tokens.items() if t.expiry < current_time
|
|
||||||
]
|
|
||||||
for k in to_delete_tokens:
|
|
||||||
del _db._data.reset_tokens[k]
|
|
||||||
|
|
||||||
|
|
||||||
async def flush() -> None:
|
async def flush() -> None:
|
||||||
"""Write all pending database changes to disk."""
|
"""Write all pending database changes to disk."""
|
||||||
from paskia.db.operations import _db
|
|
||||||
|
|
||||||
if _db is None:
|
if _store is None:
|
||||||
_logger.warning("flush() called but _db is None")
|
_logger.warning("flush() called but _store is None")
|
||||||
return
|
return
|
||||||
await flush_changes(_db.db_path, _db._pending_changes)
|
await _store.flush()
|
||||||
|
|
||||||
|
|
||||||
async def _background_loop():
|
async def _background_loop():
|
||||||
"""Background task that periodically flushes changes and cleans up."""
|
"""Background task that periodically flushes changes and cleans up."""
|
||||||
# Run cleanup immediately on startup to clear old expired items
|
# Run cleanup immediately on startup to clear old expired items
|
||||||
cleanup()
|
cleanup_expired()
|
||||||
await flush()
|
await flush()
|
||||||
|
|
||||||
last_cleanup = datetime.now(timezone.utc)
|
last_cleanup = datetime.now(UTC)
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
@@ -69,10 +41,10 @@ async def _background_loop():
|
|||||||
# Flush pending changes to disk
|
# Flush pending changes to disk
|
||||||
await flush()
|
await flush()
|
||||||
|
|
||||||
# Run cleanup less frequently
|
# Run cleanup periodically
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if (now - last_cleanup).total_seconds() >= CLEANUP_INTERVAL:
|
if (now - last_cleanup).total_seconds() >= CLEANUP_INTERVAL:
|
||||||
cleanup()
|
cleanup_expired()
|
||||||
await flush() # Flush cleanup changes
|
await flush() # Flush cleanup changes
|
||||||
last_cleanup = now
|
last_cleanup = now
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
@@ -101,6 +73,14 @@ async def start_background():
|
|||||||
if loop is not task_loop:
|
if loop is not task_loop:
|
||||||
_logger.debug("Background task in different event loop, restarting")
|
_logger.debug("Background task in different event loop, restarting")
|
||||||
_background_task = None
|
_background_task = None
|
||||||
|
else:
|
||||||
|
# Task is running in the same event loop - this is an error
|
||||||
|
raise RuntimeError(
|
||||||
|
"Background task is already running. "
|
||||||
|
"start_background() must not be called multiple times in the same event loop."
|
||||||
|
)
|
||||||
|
except RuntimeError:
|
||||||
|
raise # Re-raise RuntimeError from above
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
_logger.debug("Error checking background task loop: %s, restarting", e)
|
_logger.debug("Error checking background task loop: %s, restarting", e)
|
||||||
_background_task = None
|
_background_task = None
|
||||||
|
|||||||
+196
-46
@@ -1,19 +1,24 @@
|
|||||||
"""
|
"""
|
||||||
JSONL persistence layer for the database.
|
JSONL persistence layer for the database.
|
||||||
|
|
||||||
Handles file I/O, JSON diffs, and persistence. Works with plain JSON/dict data.
|
|
||||||
Uses aiofiles for async I/O operations.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import copy
|
||||||
import logging
|
import logging
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from datetime import datetime, timezone
|
from contextlib import contextmanager
|
||||||
|
from datetime import UTC, datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
import aiofiles
|
import aiofiles
|
||||||
import jsondiff
|
import jsondiff
|
||||||
import msgspec
|
import msgspec
|
||||||
|
|
||||||
|
from paskia.db.logging import log_change
|
||||||
|
from paskia.db.migrations import DBVER, apply_all_migrations
|
||||||
|
from paskia.db.structs import DB, SessionContext
|
||||||
|
|
||||||
_logger = logging.getLogger(__name__)
|
_logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Default database path
|
# Default database path
|
||||||
@@ -25,6 +30,7 @@ class _ChangeRecord(msgspec.Struct, omit_defaults=True):
|
|||||||
|
|
||||||
ts: datetime
|
ts: datetime
|
||||||
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
|
||||||
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 = {}
|
||||||
|
|
||||||
@@ -33,43 +39,6 @@ class _ChangeRecord(msgspec.Struct, omit_defaults=True):
|
|||||||
_change_encoder = msgspec.json.Encoder()
|
_change_encoder = msgspec.json.Encoder()
|
||||||
|
|
||||||
|
|
||||||
async def load_jsonl(db_path: Path) -> dict:
|
|
||||||
"""Load data from disk by applying change log.
|
|
||||||
|
|
||||||
Replays all changes from JSONL file using plain dicts (to handle
|
|
||||||
schema evolution).
|
|
||||||
|
|
||||||
Args:
|
|
||||||
db_path: Path to the JSONL database file
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
The final state after applying all changes
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
ValueError: If file doesn't exist or cannot be loaded
|
|
||||||
"""
|
|
||||||
if not db_path.exists():
|
|
||||||
raise ValueError(f"Database file not found: {db_path}")
|
|
||||||
data_dict: dict = {}
|
|
||||||
try:
|
|
||||||
# Read entire file at once and split into lines
|
|
||||||
async with aiofiles.open(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)
|
|
||||||
# Apply the diff to current state (marshal=True for $-prefixed keys)
|
|
||||||
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
|
|
||||||
except Exception as e:
|
|
||||||
raise ValueError(f"Error parsing line {line_num}: {e}")
|
|
||||||
except (OSError, ValueError, msgspec.DecodeError) as e:
|
|
||||||
raise ValueError(f"Failed to load database: {e}")
|
|
||||||
return data_dict
|
|
||||||
|
|
||||||
|
|
||||||
def compute_diff(previous: dict, current: dict) -> dict | None:
|
def compute_diff(previous: dict, current: dict) -> dict | None:
|
||||||
"""Compute JSON diff between two states.
|
"""Compute JSON diff between two states.
|
||||||
|
|
||||||
@@ -85,17 +54,22 @@ def compute_diff(previous: dict, current: dict) -> dict | None:
|
|||||||
|
|
||||||
|
|
||||||
def create_change_record(
|
def create_change_record(
|
||||||
action: str, diff: dict, user: str | None = None
|
action: str, version: int, diff: dict, user: str | None = None
|
||||||
) -> _ChangeRecord:
|
) -> _ChangeRecord:
|
||||||
"""Create a change record for persistence."""
|
"""Create a change record for persistence."""
|
||||||
return _ChangeRecord(
|
return _ChangeRecord(
|
||||||
ts=datetime.now(timezone.utc),
|
ts=datetime.now(UTC),
|
||||||
a=action,
|
a=action,
|
||||||
|
v=version,
|
||||||
u=user,
|
u=user,
|
||||||
diff=diff,
|
diff=diff,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Actions that are allowed to create a new database file
|
||||||
|
_BOOTSTRAP_ACTIONS = frozenset({"bootstrap", "migrate:sql"})
|
||||||
|
|
||||||
|
|
||||||
async def flush_changes(
|
async def flush_changes(
|
||||||
db_path: Path,
|
db_path: Path,
|
||||||
pending_changes: deque[_ChangeRecord],
|
pending_changes: deque[_ChangeRecord],
|
||||||
@@ -112,15 +86,25 @@ async def flush_changes(
|
|||||||
if not pending_changes:
|
if not pending_changes:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# Collect all pending changes
|
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 or migrate can create a new database",
|
||||||
|
first_action,
|
||||||
|
)
|
||||||
|
pending_changes.clear()
|
||||||
|
return False
|
||||||
|
|
||||||
changes_to_write = list(pending_changes)
|
changes_to_write = list(pending_changes)
|
||||||
pending_changes.clear()
|
pending_changes.clear()
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Build lines to append (keep as bytes, join with \n)
|
|
||||||
lines = [_change_encoder.encode(change) for change in changes_to_write]
|
lines = [_change_encoder.encode(change) for change in changes_to_write]
|
||||||
|
if not lines:
|
||||||
|
return True
|
||||||
|
|
||||||
# Append all lines in a single write (binary mode for Windows compatibility)
|
|
||||||
async with aiofiles.open(db_path, "ab") as f:
|
async with aiofiles.open(db_path, "ab") as f:
|
||||||
await f.write(b"\n".join(lines) + b"\n")
|
await f.write(b"\n".join(lines) + b"\n")
|
||||||
return True
|
return True
|
||||||
@@ -130,3 +114,169 @@ async def flush_changes(
|
|||||||
for change in reversed(changes_to_write):
|
for change in reversed(changes_to_write):
|
||||||
pending_changes.appendleft(change)
|
pending_changes.appendleft(change)
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
class JsonlStore:
|
||||||
|
"""JSONL persistence layer for a DB instance."""
|
||||||
|
|
||||||
|
def __init__(self, db: DB, db_path: str = DB_PATH_DEFAULT):
|
||||||
|
self.db: DB = db
|
||||||
|
self.db_path = Path(db_path)
|
||||||
|
self._previous_builtins: dict[str, Any] = {}
|
||||||
|
self._pending_changes: deque[_ChangeRecord] = deque()
|
||||||
|
self._current_action: str = "system"
|
||||||
|
self._current_user: str | None = None
|
||||||
|
self._in_transaction: bool = False
|
||||||
|
self._transaction_snapshot: dict[str, Any] | None = None
|
||||||
|
self._current_version: int = DBVER # Schema version for new databases
|
||||||
|
|
||||||
|
async def load(self, db_path: str | None = None) -> None:
|
||||||
|
"""Load data from JSONL change log."""
|
||||||
|
if db_path is not None:
|
||||||
|
self.db_path = Path(db_path)
|
||||||
|
if not self.db_path.exists():
|
||||||
|
return
|
||||||
|
|
||||||
|
# Replay change log to reconstruct state
|
||||||
|
data_dict: dict = {}
|
||||||
|
try:
|
||||||
|
async with aiofiles.open(self.db_path, "rb") as f:
|
||||||
|
content = await f.read()
|
||||||
|
for line_num, line in enumerate(content.split(b"\n"), 1):
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
change = msgspec.json.decode(line)
|
||||||
|
data_dict = jsondiff.patch(data_dict, change["diff"], marshal=True)
|
||||||
|
self._current_version = change.get("v", 0)
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"Error parsing line {line_num}: {e}")
|
||||||
|
except (OSError, ValueError, msgspec.DecodeError) as e:
|
||||||
|
raise ValueError(f"Failed to load database: {e}")
|
||||||
|
|
||||||
|
if not data_dict:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Set previous state for diffing (will be updated by _queue_change)
|
||||||
|
self._previous_builtins = copy.deepcopy(data_dict)
|
||||||
|
|
||||||
|
# Callback to persist each migration
|
||||||
|
async def persist_migration(
|
||||||
|
action: str, new_version: int, current: dict
|
||||||
|
) -> None:
|
||||||
|
self._current_version = new_version
|
||||||
|
self._queue_change(action, new_version, current)
|
||||||
|
|
||||||
|
# Apply schema migrations one at a time
|
||||||
|
await apply_all_migrations(data_dict, self._current_version, persist_migration)
|
||||||
|
|
||||||
|
# Decode to msgspec struct
|
||||||
|
decoder = msgspec.json.Decoder(DB)
|
||||||
|
self.db = decoder.decode(msgspec.json.encode(data_dict))
|
||||||
|
self.db._store = self
|
||||||
|
|
||||||
|
# Normalize via msgspec round-trip (handles omit_defaults etc.)
|
||||||
|
# This ensures _previous_builtins matches what msgspec would produce
|
||||||
|
normalized_dict = msgspec.to_builtins(self.db)
|
||||||
|
await persist_migration(
|
||||||
|
"migrate:msgspec", self._current_version, normalized_dict
|
||||||
|
)
|
||||||
|
|
||||||
|
def _queue_change(
|
||||||
|
self, action: str, version: int, current: dict, user: str | None = None
|
||||||
|
) -> None:
|
||||||
|
"""Queue a change record and log it.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
action: The action name for the change record
|
||||||
|
version: The schema version for the change record
|
||||||
|
current: The current state as a plain dict
|
||||||
|
user: Optional user UUID who performed the action
|
||||||
|
"""
|
||||||
|
diff = compute_diff(self._previous_builtins, current)
|
||||||
|
if not diff:
|
||||||
|
return
|
||||||
|
self._pending_changes.append(create_change_record(action, version, diff, user))
|
||||||
|
self._previous_builtins = copy.deepcopy(current)
|
||||||
|
|
||||||
|
# Log the change with user display name if available
|
||||||
|
user_display = None
|
||||||
|
if user:
|
||||||
|
try:
|
||||||
|
user_uuid = UUID(user)
|
||||||
|
if user_uuid in self.db.users:
|
||||||
|
user_display = self.db.users[user_uuid].display_name
|
||||||
|
except (ValueError, KeyError):
|
||||||
|
user_display = user
|
||||||
|
|
||||||
|
log_change(action, diff, user_display)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def transaction(
|
||||||
|
self,
|
||||||
|
action: str,
|
||||||
|
ctx: SessionContext | None = None,
|
||||||
|
*,
|
||||||
|
user: str | None = None,
|
||||||
|
):
|
||||||
|
"""Wrap writes in transaction. Queues change on successful exit.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
action: Describes the operation (e.g., "Created user", "Login")
|
||||||
|
ctx: Session context of user performing the action (None for system operations)
|
||||||
|
user: User UUID string (alternative to ctx when full context unavailable)
|
||||||
|
"""
|
||||||
|
if self._in_transaction:
|
||||||
|
raise RuntimeError("Nested transactions are not supported")
|
||||||
|
|
||||||
|
# Check for out-of-transaction modifications
|
||||||
|
current_state = msgspec.to_builtins(self.db)
|
||||||
|
if current_state != self._previous_builtins:
|
||||||
|
# Allow bootstrap/migrate to create a new database from empty state
|
||||||
|
is_bootstrap = action in _BOOTSTRAP_ACTIONS or action.startswith("migrate:")
|
||||||
|
if is_bootstrap and not self._previous_builtins:
|
||||||
|
pass # Expected: creating database from scratch
|
||||||
|
else:
|
||||||
|
diff = compute_diff(self._previous_builtins, current_state)
|
||||||
|
diff_json = msgspec.json.encode(diff).decode()
|
||||||
|
_logger.critical(
|
||||||
|
"Database state modified outside of transaction! "
|
||||||
|
"This indicates a bug where DB changes occurred without a transaction wrapper.\n"
|
||||||
|
f"Changes detected:\n{diff_json}"
|
||||||
|
)
|
||||||
|
raise SystemExit(1)
|
||||||
|
|
||||||
|
old_action = self._current_action
|
||||||
|
old_user = self._current_user
|
||||||
|
self._current_action = action
|
||||||
|
# Prefer ctx.user.uuid if ctx provided, otherwise use user param
|
||||||
|
self._current_user = str(ctx.user.uuid) if ctx else user
|
||||||
|
self._in_transaction = True
|
||||||
|
self._transaction_snapshot = current_state
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
current = msgspec.to_builtins(self.db)
|
||||||
|
self._queue_change(
|
||||||
|
self._current_action, self._current_version, current, self._current_user
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
# Rollback on error: restore from snapshot
|
||||||
|
_logger.warning("Transaction '%s' failed, rolling back changes", action)
|
||||||
|
if self._transaction_snapshot is not None:
|
||||||
|
decoder = msgspec.json.Decoder(DB)
|
||||||
|
self.db = decoder.decode(
|
||||||
|
msgspec.json.encode(self._transaction_snapshot)
|
||||||
|
)
|
||||||
|
self.db._store = self
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
self._current_action = old_action
|
||||||
|
self._current_user = old_user
|
||||||
|
self._in_transaction = False
|
||||||
|
self._transaction_snapshot = None
|
||||||
|
|
||||||
|
async def flush(self) -> bool:
|
||||||
|
"""Write all pending changes to disk."""
|
||||||
|
return await flush_changes(self.db_path, self._pending_changes)
|
||||||
|
|||||||
@@ -0,0 +1,233 @@
|
|||||||
|
"""
|
||||||
|
Database change logging with pretty-printed diffs.
|
||||||
|
|
||||||
|
Provides a logger for JSONL database changes that formats diffs
|
||||||
|
in a human-readable path.notation style with color coding.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
import sys
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
logger = logging.getLogger("paskia.db")
|
||||||
|
|
||||||
|
# Pattern to match control characters and bidirectional overrides
|
||||||
|
_UNSAFE_CHARS = re.compile(
|
||||||
|
r"[\x00-\x1f\x7f-\x9f" # C0 and C1 control characters
|
||||||
|
r"\u200e\u200f" # LRM, RLM
|
||||||
|
r"\u202a-\u202e" # LRE, RLE, PDF, LRO, RLO
|
||||||
|
r"\u2066-\u2069" # LRI, RLI, FSI, PDI
|
||||||
|
r"]"
|
||||||
|
)
|
||||||
|
|
||||||
|
# ANSI color codes (matching FastAPI logging style)
|
||||||
|
_RESET = "\033[0m"
|
||||||
|
_DIM = "\033[2m"
|
||||||
|
_PATH_PREFIX = "\033[1;30m" # Dark grey for path prefix (like host in access log)
|
||||||
|
_PATH_FINAL = "\033[0m" # Default for final element (like path in access log)
|
||||||
|
_REPLACE = "\033[0;33m" # Yellow for replacements
|
||||||
|
_DELETE = "\033[0;31m" # Red for deletions
|
||||||
|
_ADD = "\033[0;32m" # Green for additions
|
||||||
|
_ACTION = "\033[1;34m" # Bold blue for action name
|
||||||
|
_USER = "\033[0;34m" # Blue for user display
|
||||||
|
|
||||||
|
|
||||||
|
def _use_color() -> bool:
|
||||||
|
"""Check if we should use color output."""
|
||||||
|
return sys.stderr.isatty()
|
||||||
|
|
||||||
|
|
||||||
|
def _format_value(value: Any, use_color: bool, max_len: int = 60) -> str:
|
||||||
|
"""Format a value for display, truncating if needed."""
|
||||||
|
if value is None:
|
||||||
|
return "null"
|
||||||
|
|
||||||
|
if isinstance(value, bool):
|
||||||
|
return "true" if value else "false"
|
||||||
|
|
||||||
|
if isinstance(value, (int, float)):
|
||||||
|
return str(value)
|
||||||
|
|
||||||
|
if isinstance(value, str):
|
||||||
|
# Filter out control characters and bidirectional overrides
|
||||||
|
value = _UNSAFE_CHARS.sub("", value)
|
||||||
|
# Truncate long strings
|
||||||
|
if len(value) > max_len:
|
||||||
|
return value[: max_len - 3] + "..."
|
||||||
|
return value
|
||||||
|
|
||||||
|
if isinstance(value, dict):
|
||||||
|
if not value:
|
||||||
|
return "{}"
|
||||||
|
# For small dicts, show inline
|
||||||
|
if len(value) == 1:
|
||||||
|
k, v = next(iter(value.items()))
|
||||||
|
return "{" + f"{k}: {_format_value(v, use_color, max_len=30)}" + "}"
|
||||||
|
return f"{{...{len(value)} keys}}"
|
||||||
|
|
||||||
|
if isinstance(value, list):
|
||||||
|
if not value:
|
||||||
|
return "[]"
|
||||||
|
if len(value) == 1:
|
||||||
|
return "[" + _format_value(value[0], use_color, max_len=30) + "]"
|
||||||
|
return f"[...{len(value)} items]"
|
||||||
|
|
||||||
|
# Fallback for other types
|
||||||
|
text = str(value)
|
||||||
|
if len(text) > max_len:
|
||||||
|
text = text[: max_len - 3] + "..."
|
||||||
|
return text
|
||||||
|
|
||||||
|
|
||||||
|
def _format_path(path: list[str], use_color: bool) -> str:
|
||||||
|
"""Format a path as dot notation with prefix in dark grey, final in default."""
|
||||||
|
if not path:
|
||||||
|
return ""
|
||||||
|
if not use_color:
|
||||||
|
return ".".join(path)
|
||||||
|
if len(path) == 1:
|
||||||
|
return f"{_PATH_FINAL}{path[0]}{_RESET}"
|
||||||
|
prefix = ".".join(path[:-1])
|
||||||
|
final = path[-1]
|
||||||
|
return f"{_PATH_PREFIX}{prefix}.{_RESET}{_PATH_FINAL}{final}{_RESET}"
|
||||||
|
|
||||||
|
|
||||||
|
def _collect_changes(
|
||||||
|
diff: dict, path: list[str], changes: list[tuple[str, list[str], Any, Any | None]]
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
Recursively collect changes from a diff into a flat list.
|
||||||
|
|
||||||
|
Each change is a tuple of (change_type, path, new_value, old_value).
|
||||||
|
change_type is one of: 'set', 'replace', 'delete'
|
||||||
|
"""
|
||||||
|
if not isinstance(diff, dict):
|
||||||
|
# Leaf value - this is a set operation
|
||||||
|
changes.append(("set", path, diff, None))
|
||||||
|
return
|
||||||
|
|
||||||
|
for key, value in diff.items():
|
||||||
|
if key == "$delete":
|
||||||
|
# $delete contains a list of keys to delete
|
||||||
|
if isinstance(value, list):
|
||||||
|
for deleted_key in value:
|
||||||
|
changes.append(("delete", path + [str(deleted_key)], None, None))
|
||||||
|
else:
|
||||||
|
changes.append(("delete", path + [str(value)], None, None))
|
||||||
|
|
||||||
|
elif key == "$replace":
|
||||||
|
# $replace contains the new value for this path
|
||||||
|
if isinstance(value, dict):
|
||||||
|
# Replacing with a dict - show each key as a replacement
|
||||||
|
for rkey, rval in value.items():
|
||||||
|
changes.append(("replace", path + [str(rkey)], rval, None))
|
||||||
|
if not value:
|
||||||
|
# Empty replacement - clearing the collection
|
||||||
|
changes.append(("replace", path, {}, None))
|
||||||
|
else:
|
||||||
|
changes.append(("replace", path, value, None))
|
||||||
|
|
||||||
|
elif key.startswith("$"):
|
||||||
|
# Other special operations (future-proofing)
|
||||||
|
changes.append(("set", path, {key: value}, None))
|
||||||
|
|
||||||
|
else:
|
||||||
|
# Regular nested key
|
||||||
|
_collect_changes(value, path + [str(key)], changes)
|
||||||
|
|
||||||
|
|
||||||
|
def _format_change_line(
|
||||||
|
change_type: str, path: list[str], value: Any, use_color: bool
|
||||||
|
) -> str:
|
||||||
|
"""Format a single change as a one-line string."""
|
||||||
|
path_str = _format_path(path, use_color)
|
||||||
|
value_str = _format_value(value, use_color)
|
||||||
|
|
||||||
|
if change_type == "delete":
|
||||||
|
if use_color:
|
||||||
|
return f" ❌ {path_str}"
|
||||||
|
return f" - {path_str}"
|
||||||
|
|
||||||
|
if change_type == "replace":
|
||||||
|
if use_color:
|
||||||
|
return f" {_REPLACE}⟳{_RESET} {path_str} {_DIM}={_RESET} {value_str}"
|
||||||
|
return f" ~ {path_str} = {value_str}"
|
||||||
|
|
||||||
|
# Default: set/add
|
||||||
|
if use_color:
|
||||||
|
return f" {_ADD}+{_RESET} {path_str} {_DIM}={_RESET} {value_str}"
|
||||||
|
return f" + {path_str} = {value_str}"
|
||||||
|
|
||||||
|
|
||||||
|
def format_diff(diff: dict) -> list[str]:
|
||||||
|
"""
|
||||||
|
Format a JSON diff as human-readable lines.
|
||||||
|
|
||||||
|
Returns a list of formatted lines (without newlines).
|
||||||
|
Single changes return one line, multiple changes return multiple lines.
|
||||||
|
"""
|
||||||
|
use_color = _use_color()
|
||||||
|
changes: list[tuple[str, list[str], Any, Any | None]] = []
|
||||||
|
_collect_changes(diff, [], changes)
|
||||||
|
|
||||||
|
if not changes:
|
||||||
|
return []
|
||||||
|
|
||||||
|
# Format each change
|
||||||
|
lines = []
|
||||||
|
for change_type, path, value, _ in changes:
|
||||||
|
lines.append(_format_change_line(change_type, path, value, use_color))
|
||||||
|
|
||||||
|
return lines
|
||||||
|
|
||||||
|
|
||||||
|
def format_action_header(action: str, user_display: str | None = None) -> str:
|
||||||
|
"""Format the action header line."""
|
||||||
|
use_color = _use_color()
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
action_str = f"{_ACTION}{action}{_RESET}"
|
||||||
|
if user_display:
|
||||||
|
user_str = f"{_USER}{user_display}{_RESET}"
|
||||||
|
return f"{action_str} by {user_str}"
|
||||||
|
return action_str
|
||||||
|
else:
|
||||||
|
if user_display:
|
||||||
|
return f"{action} by {user_display}"
|
||||||
|
return action
|
||||||
|
|
||||||
|
|
||||||
|
def log_change(action: str, diff: dict, user_display: str | None = None) -> None:
|
||||||
|
"""
|
||||||
|
Log a database change with pretty-printed diff.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
action: The action name (e.g., "login", "admin:delete_user")
|
||||||
|
diff: The JSON diff dict
|
||||||
|
user_display: Optional display name of the user who performed the action
|
||||||
|
"""
|
||||||
|
header = format_action_header(action, user_display)
|
||||||
|
diff_lines = format_diff(diff)
|
||||||
|
|
||||||
|
if not diff_lines:
|
||||||
|
logger.info(header)
|
||||||
|
return
|
||||||
|
|
||||||
|
if len(diff_lines) == 1:
|
||||||
|
# Single change - combine on one line
|
||||||
|
logger.info(f"{header}{diff_lines[0]}")
|
||||||
|
else:
|
||||||
|
# Multiple changes - header on its own line, then changes
|
||||||
|
logger.info(header)
|
||||||
|
for line in diff_lines:
|
||||||
|
logger.info(line)
|
||||||
|
|
||||||
|
|
||||||
|
def configure_db_logging() -> None:
|
||||||
|
"""Configure the database logger to output to stderr without prefix."""
|
||||||
|
handler = logging.StreamHandler(sys.stderr)
|
||||||
|
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||||
|
logger.addHandler(handler)
|
||||||
|
logger.setLevel(logging.INFO)
|
||||||
|
logger.propagate = False
|
||||||
@@ -0,0 +1,33 @@
|
|||||||
|
"""
|
||||||
|
Database schema migrations.
|
||||||
|
|
||||||
|
Migrations are applied during database load based on the version field.
|
||||||
|
Each migration should be idempotent and only run when needed.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
|
||||||
|
|
||||||
|
def migrate_v1(d: dict) -> None:
|
||||||
|
"""Remove Org.created_at fields."""
|
||||||
|
for org_data in d["orgs"].values():
|
||||||
|
org_data.pop("created_at", None)
|
||||||
|
|
||||||
|
|
||||||
|
migrations = sorted(
|
||||||
|
[f for n, f in globals().items() if n.startswith("migrate_v")],
|
||||||
|
key=lambda f: int(f.__name__.removeprefix("migrate_v")),
|
||||||
|
)
|
||||||
|
|
||||||
|
DBVER = len(migrations) # Used by bootstrap and migrate:sql to set initial version
|
||||||
|
|
||||||
|
|
||||||
|
async def apply_all_migrations(
|
||||||
|
data_dict: dict,
|
||||||
|
current_version: int,
|
||||||
|
persist: Callable[[str, int, dict], Awaitable[None]],
|
||||||
|
) -> None:
|
||||||
|
while current_version < DBVER:
|
||||||
|
migrations[current_version](data_dict)
|
||||||
|
current_version += 1
|
||||||
|
await persist(f"migrate:v{current_version}", current_version, data_dict)
|
||||||
+411
-827
File diff suppressed because it is too large
Load Diff
+395
-81
@@ -1,43 +1,223 @@
|
|||||||
from datetime import datetime
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import secrets
|
||||||
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import msgspec
|
import msgspec
|
||||||
|
import uuid7
|
||||||
|
|
||||||
|
from paskia import db
|
||||||
|
from paskia.util.hostutil import normalize_host
|
||||||
|
|
||||||
|
# Sentinel for uuid fields before they are set by create() or DB post init
|
||||||
|
_UUID_UNSET = UUID(int=0)
|
||||||
|
|
||||||
|
|
||||||
class Permission(msgspec.Struct, omit_defaults=True):
|
class Permission(msgspec.Struct, dict=True, omit_defaults=True):
|
||||||
uuid: UUID # UUID primary key
|
"""Permission data structure.
|
||||||
|
|
||||||
|
Mutable fields: scope, display_name, domain, orgs
|
||||||
|
Immutable fields: None (all fields can be updated via update_permission)
|
||||||
|
uuid is generated at creation.
|
||||||
|
"""
|
||||||
|
|
||||||
scope: str # Permission scope identifier (e.g. "auth:admin", "myapp:write")
|
scope: str # Permission scope identifier (e.g. "auth:admin", "myapp:write")
|
||||||
display_name: str
|
display_name: str
|
||||||
domain: str | None = None # If set, scopes permission to this domain
|
domain: str | None = None # If set, scopes permission to this domain
|
||||||
|
orgs: dict[UUID, bool] = {} # org_uuid -> True (which orgs can grant this)
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def org_set(self) -> set[UUID]:
|
||||||
|
"""Get orgs that can grant this permission as a set."""
|
||||||
|
return set(self.orgs.keys())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def orgs_list(self) -> list[Org]:
|
||||||
|
"""Get list of Org objects that can grant this permission."""
|
||||||
|
return [
|
||||||
|
db.data().orgs[org_uuid]
|
||||||
|
for org_uuid in self.orgs.keys()
|
||||||
|
if org_uuid in db.data().orgs
|
||||||
|
]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
scope: str,
|
||||||
|
display_name: str,
|
||||||
|
domain: str | None = None,
|
||||||
|
) -> Permission:
|
||||||
|
"""Create a new Permission with auto-generated uuid7."""
|
||||||
|
perm = cls(
|
||||||
|
scope=scope,
|
||||||
|
display_name=display_name,
|
||||||
|
domain=domain,
|
||||||
|
)
|
||||||
|
perm.uuid = uuid7.create()
|
||||||
|
return perm
|
||||||
|
|
||||||
|
|
||||||
class Role(msgspec.Struct):
|
class Org(msgspec.Struct, dict=True):
|
||||||
uuid: UUID
|
"""Organization data structure."""
|
||||||
org_uuid: UUID
|
|
||||||
display_name: str
|
display_name: str
|
||||||
permissions: list[str] = [] # permission UUIDs this role grants
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def roles(self) -> list[Role]:
|
||||||
|
"""Get all roles that belong to this organization."""
|
||||||
|
return [r for r in db.data().roles.values() if r.org_uuid == self.uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def permissions(self) -> list[Permission]:
|
||||||
|
"""Get all permissions that this organization can grant."""
|
||||||
|
return [p for p in db.data().permissions.values() if self.uuid in p.orgs]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(cls, display_name: str) -> Org:
|
||||||
|
"""Create a new Org with auto-generated uuid7."""
|
||||||
|
org = cls(display_name=display_name)
|
||||||
|
org.uuid = uuid7.create()
|
||||||
|
return org
|
||||||
|
|
||||||
|
|
||||||
class Org(msgspec.Struct):
|
class Role(msgspec.Struct, dict=True, omit_defaults=True):
|
||||||
uuid: UUID
|
"""Role data structure.
|
||||||
|
|
||||||
|
Mutable fields: display_name, permissions
|
||||||
|
Immutable fields: org_uuid (set at creation, never modified)
|
||||||
|
uuid is generated at creation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
org_uuid: UUID = msgspec.field(name="org")
|
||||||
display_name: str
|
display_name: str
|
||||||
permissions: list[str] = [] # permission UUIDs this org can grant
|
permissions: dict[UUID, bool] = {} # permission_uuid -> True
|
||||||
roles: list[Role] = [] # roles belonging to this org
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def permission_set(self) -> set[UUID]:
|
||||||
|
"""Get permissions as a set of UUIDs."""
|
||||||
|
return set(self.permissions.keys())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def permissions_list(self) -> list[Permission]:
|
||||||
|
"""Get list of Permission objects for this role."""
|
||||||
|
return [
|
||||||
|
db.data().permissions[perm_uuid]
|
||||||
|
for perm_uuid in self.permissions.keys()
|
||||||
|
if perm_uuid in db.data().permissions
|
||||||
|
]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def org(self) -> Org:
|
||||||
|
"""Get the organization object this role belongs to."""
|
||||||
|
return db.data().orgs[self.org_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def users(self) -> list[User]:
|
||||||
|
"""Get all users that have this role."""
|
||||||
|
return [u for u in db.data().users.values() if u.role_uuid == self.uuid]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
org: UUID | Org,
|
||||||
|
display_name: str,
|
||||||
|
permissions: set[UUID] | None = None,
|
||||||
|
) -> Role:
|
||||||
|
"""Create a new Role with auto-generated uuid7."""
|
||||||
|
org_uuid = org if isinstance(org, UUID) else org.uuid
|
||||||
|
role = cls(
|
||||||
|
org_uuid=org_uuid,
|
||||||
|
display_name=display_name,
|
||||||
|
permissions={p: True for p in (permissions or set())},
|
||||||
|
)
|
||||||
|
role.uuid = uuid7.create()
|
||||||
|
return role
|
||||||
|
|
||||||
|
|
||||||
class User(msgspec.Struct):
|
class User(msgspec.Struct, dict=True):
|
||||||
uuid: UUID
|
"""User data structure.
|
||||||
|
|
||||||
|
Mutable fields: display_name, role_uuid, last_seen, visits
|
||||||
|
Immutable fields: created_at (set at creation, never modified)
|
||||||
|
uuid is derived from created_at using uuid7.
|
||||||
|
"""
|
||||||
|
|
||||||
display_name: str
|
display_name: str
|
||||||
role_uuid: UUID
|
role_uuid: UUID = msgspec.field(name="role")
|
||||||
created_at: datetime | None = None
|
created_at: datetime
|
||||||
last_seen: datetime | None = None
|
last_seen: datetime | None = None
|
||||||
visits: int = 0
|
visits: int = 0
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
|
@property
|
||||||
|
def role(self) -> Role:
|
||||||
|
"""Get the role object this user has."""
|
||||||
|
return db.data().roles[self.role_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def org(self) -> Org:
|
||||||
|
"""Get the organization this user belongs to (via role)."""
|
||||||
|
return self.role.org
|
||||||
|
|
||||||
|
@property
|
||||||
|
def credentials(self) -> list[Credential]:
|
||||||
|
"""Get all credentials for this user."""
|
||||||
|
return [c for c in db.data().credentials.values() if c.user_uuid == self.uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def sessions(self) -> list[Session]:
|
||||||
|
"""Get all sessions for this user."""
|
||||||
|
return [s for s in db.data().sessions.values() if s.user_uuid == self.uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def reset_tokens(self) -> list[ResetToken]:
|
||||||
|
"""Get all reset tokens for this user."""
|
||||||
|
return [t for t in db.data().reset_tokens.values() if t.user_uuid == self.uuid]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
display_name: str,
|
||||||
|
role: UUID | Role,
|
||||||
|
created_at: datetime | None = None,
|
||||||
|
) -> User:
|
||||||
|
"""Create a new User with auto-generated uuid7."""
|
||||||
|
role_uuid = role if isinstance(role, UUID) else role.uuid
|
||||||
|
user = cls(
|
||||||
|
display_name=display_name,
|
||||||
|
role_uuid=role_uuid,
|
||||||
|
created_at=created_at or datetime.now(UTC),
|
||||||
|
)
|
||||||
|
user.uuid = uuid7.create(user.created_at)
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
class Credential(msgspec.Struct, dict=True):
|
||||||
|
"""Credential (passkey) data structure.
|
||||||
|
|
||||||
|
Mutable fields: sign_count, last_used, last_verified
|
||||||
|
Immutable fields: credential_id, user, aaguid, public_key, created_at
|
||||||
|
uuid is derived from created_at using uuid7.
|
||||||
|
"""
|
||||||
|
|
||||||
class Credential(msgspec.Struct):
|
|
||||||
uuid: UUID
|
|
||||||
credential_id: bytes # Long binary ID from the authenticator
|
credential_id: bytes # Long binary ID from the authenticator
|
||||||
user_uuid: UUID
|
user_uuid: UUID = msgspec.field(name="user")
|
||||||
aaguid: UUID
|
aaguid: UUID
|
||||||
public_key: bytes
|
public_key: bytes
|
||||||
sign_count: int
|
sign_count: int
|
||||||
@@ -45,16 +225,78 @@ class Credential(msgspec.Struct):
|
|||||||
last_used: datetime | None = None
|
last_used: datetime | None = None
|
||||||
last_verified: datetime | None = None
|
last_verified: datetime | None = None
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "uuid"):
|
||||||
|
self.uuid: UUID = _UUID_UNSET
|
||||||
|
|
||||||
class Session(msgspec.Struct):
|
@property
|
||||||
key: str
|
def user(self) -> User:
|
||||||
user_uuid: UUID
|
"""Get the User object for this credential."""
|
||||||
credential_uuid: UUID
|
return db.data().users[self.user_uuid]
|
||||||
host: str | None
|
|
||||||
ip: str | None
|
@property
|
||||||
user_agent: str | None
|
def sessions(self) -> list[Session]:
|
||||||
|
"""Get all sessions using this credential."""
|
||||||
|
return [
|
||||||
|
s for s in db.data().sessions.values() if s.credential_uuid == self.uuid
|
||||||
|
]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
credential_id: bytes,
|
||||||
|
user: UUID | User,
|
||||||
|
aaguid: UUID,
|
||||||
|
public_key: bytes,
|
||||||
|
sign_count: int,
|
||||||
|
created_at: datetime | None = None,
|
||||||
|
) -> Credential:
|
||||||
|
"""Create a new Credential with auto-generated uuid7."""
|
||||||
|
user_uuid = user if isinstance(user, UUID) else user.uuid
|
||||||
|
now = created_at or datetime.now(UTC)
|
||||||
|
cred = cls(
|
||||||
|
credential_id=credential_id,
|
||||||
|
user_uuid=user_uuid,
|
||||||
|
aaguid=aaguid,
|
||||||
|
public_key=public_key,
|
||||||
|
sign_count=sign_count,
|
||||||
|
created_at=now,
|
||||||
|
last_used=now,
|
||||||
|
last_verified=now,
|
||||||
|
)
|
||||||
|
cred.uuid = uuid7.create(now)
|
||||||
|
return cred
|
||||||
|
|
||||||
|
|
||||||
|
class Session(msgspec.Struct, dict=True):
|
||||||
|
"""Session data structure.
|
||||||
|
|
||||||
|
Mutable fields: expiry (updated on session refresh)
|
||||||
|
Immutable fields: user_uuid, credential_uuid, host, ip, user_agent
|
||||||
|
key is stored in the dict key, not in the struct.
|
||||||
|
"""
|
||||||
|
|
||||||
|
user_uuid: UUID = msgspec.field(name="user")
|
||||||
|
credential_uuid: UUID = msgspec.field(name="credential")
|
||||||
|
host: str
|
||||||
|
ip: str
|
||||||
|
user_agent: str
|
||||||
expiry: datetime
|
expiry: datetime
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "key"):
|
||||||
|
self.key: str = ""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def user(self) -> User:
|
||||||
|
"""Get the User object for this session."""
|
||||||
|
return db.data().users[self.user_uuid]
|
||||||
|
|
||||||
|
@property
|
||||||
|
def credential(self) -> Credential:
|
||||||
|
"""Get the Credential object for this session."""
|
||||||
|
return db.data().credentials[self.credential_uuid]
|
||||||
|
|
||||||
def metadata(self) -> dict:
|
def metadata(self) -> dict:
|
||||||
"""Return session metadata for backwards compatibility."""
|
"""Return session metadata for backwards compatibility."""
|
||||||
return {
|
return {
|
||||||
@@ -63,86 +305,158 @@ class Session(msgspec.Struct):
|
|||||||
"expiry": self.expiry.isoformat(),
|
"expiry": self.expiry.isoformat(),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create(
|
||||||
|
cls,
|
||||||
|
user: UUID | User,
|
||||||
|
credential: UUID | Credential,
|
||||||
|
host: str,
|
||||||
|
ip: str,
|
||||||
|
user_agent: str,
|
||||||
|
expiry: datetime,
|
||||||
|
) -> Session:
|
||||||
|
"""Create a new Session with auto-generated key."""
|
||||||
|
user_uuid = user if isinstance(user, UUID) else user.uuid
|
||||||
|
credential_uuid = (
|
||||||
|
credential if isinstance(credential, UUID) else credential.uuid
|
||||||
|
)
|
||||||
|
session = cls(
|
||||||
|
user_uuid=user_uuid,
|
||||||
|
credential_uuid=credential_uuid,
|
||||||
|
host=host,
|
||||||
|
ip=ip,
|
||||||
|
user_agent=user_agent,
|
||||||
|
expiry=expiry,
|
||||||
|
)
|
||||||
|
session.key = secrets.token_urlsafe(12)
|
||||||
|
return session
|
||||||
|
|
||||||
class ResetToken(msgspec.Struct):
|
|
||||||
key: bytes
|
class ResetToken(msgspec.Struct, dict=True):
|
||||||
user_uuid: UUID
|
"""Reset/device-addition token data structure.
|
||||||
|
|
||||||
|
Immutable fields: All fields (tokens are created and deleted, never modified)
|
||||||
|
key is stored in the dict key, not in the struct.
|
||||||
|
"""
|
||||||
|
|
||||||
|
user_uuid: UUID = msgspec.field(name="user")
|
||||||
expiry: datetime
|
expiry: datetime
|
||||||
token_type: str
|
token_type: str
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
if not hasattr(self, "key"):
|
||||||
|
self.key: bytes = b""
|
||||||
|
|
||||||
|
@property
|
||||||
|
def user(self) -> User:
|
||||||
|
"""Get the User object for this reset token."""
|
||||||
|
return db.data().users[self.user_uuid]
|
||||||
|
|
||||||
|
|
||||||
class SessionContext(msgspec.Struct):
|
class SessionContext(msgspec.Struct):
|
||||||
session: Session
|
session: Session
|
||||||
user: User
|
user: User
|
||||||
org: Org
|
org: Org
|
||||||
role: Role
|
role: Role
|
||||||
credential: Credential | None = None
|
credential: Credential
|
||||||
permissions: list[Permission] | None = None
|
permissions: list[Permission] = []
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
# Internal storage types (different structure for efficient storage)
|
# Database storage structure
|
||||||
# -------------------------------------------------------------------------
|
# -------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class _PermissionData(msgspec.Struct, omit_defaults=True):
|
class DB(msgspec.Struct, dict=True, omit_defaults=False):
|
||||||
scope: str # Permission scope identifier
|
"""In-memory database. Access fields directly for reads."""
|
||||||
display_name: str
|
|
||||||
domain: str | None = None
|
|
||||||
orgs: dict[UUID, bool] = {} # org_uuid -> True (which orgs can grant this)
|
|
||||||
|
|
||||||
|
permissions: dict[UUID, Permission] = {}
|
||||||
|
orgs: dict[UUID, Org] = {}
|
||||||
|
roles: dict[UUID, Role] = {}
|
||||||
|
users: dict[UUID, User] = {}
|
||||||
|
credentials: dict[UUID, Credential] = {}
|
||||||
|
sessions: dict[str, Session] = {}
|
||||||
|
reset_tokens: dict[bytes, ResetToken] = {}
|
||||||
|
|
||||||
class _OrgData(msgspec.Struct):
|
def __post_init__(self):
|
||||||
display_name: str
|
# Store reference for persistence (not serialized)
|
||||||
created_at: datetime | None = None
|
self._store = None
|
||||||
|
# Set the key fields on all stored objects
|
||||||
|
for uuid, perm in self.permissions.items():
|
||||||
|
perm.uuid = uuid
|
||||||
|
for uuid, org in self.orgs.items():
|
||||||
|
org.uuid = uuid
|
||||||
|
for uuid, role in self.roles.items():
|
||||||
|
role.uuid = uuid
|
||||||
|
for uuid, user in self.users.items():
|
||||||
|
user.uuid = uuid
|
||||||
|
for uuid, cred in self.credentials.items():
|
||||||
|
cred.uuid = uuid
|
||||||
|
for key, session in self.sessions.items():
|
||||||
|
session.key = key
|
||||||
|
for key, token in self.reset_tokens.items():
|
||||||
|
token.key = key
|
||||||
|
|
||||||
|
def transaction(self, action, ctx=None, *, user=None):
|
||||||
|
"""Wrap writes in transaction. Delegates to JsonlStore."""
|
||||||
|
return self._store.transaction(action, ctx, user=user)
|
||||||
|
|
||||||
class _RoleData(msgspec.Struct):
|
def session_ctx(
|
||||||
org: UUID
|
self, session_key: str, host: str | None = None
|
||||||
display_name: str
|
) -> SessionContext | None:
|
||||||
permissions: dict[UUID, bool] = {} # permission_uuid -> True
|
"""Get full session context with effective permissions.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session_key: The session key string
|
||||||
|
host: Optional host for binding/validation and domain-scoped permissions
|
||||||
|
|
||||||
class _UserData(msgspec.Struct):
|
Returns:
|
||||||
display_name: str
|
SessionContext if valid, None if session not found, expired, or host mismatch
|
||||||
role: UUID
|
"""
|
||||||
created_at: datetime
|
try:
|
||||||
last_seen: datetime | None
|
s = self.sessions[session_key]
|
||||||
visits: int
|
except KeyError:
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Validate host matches (sessions are always created with a host)
|
||||||
|
if s.host != host:
|
||||||
|
# Session bound to different host
|
||||||
|
return None
|
||||||
|
|
||||||
class _CredentialData(msgspec.Struct):
|
try:
|
||||||
credential_id: bytes
|
user = s.user
|
||||||
user: UUID
|
role = user.role
|
||||||
aaguid: UUID
|
org = role.org
|
||||||
public_key: bytes
|
credential = s.credential
|
||||||
sign_count: int
|
except KeyError:
|
||||||
created_at: datetime
|
return None
|
||||||
last_used: datetime | None
|
|
||||||
last_verified: datetime | None
|
|
||||||
|
|
||||||
|
# Effective permissions: role's permissions that the org can grant
|
||||||
|
# Also filter by domain if host is provided
|
||||||
|
org_perm_uuids = {p.uuid for p in org.permissions}
|
||||||
|
normalized_host = normalize_host(host)
|
||||||
|
host_without_port = (
|
||||||
|
normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
||||||
|
)
|
||||||
|
|
||||||
class _SessionData(msgspec.Struct):
|
effective_perms = []
|
||||||
user: UUID
|
for perm_uuid in role.permission_set:
|
||||||
credential: UUID
|
if perm_uuid not in org_perm_uuids:
|
||||||
host: str | None
|
continue
|
||||||
ip: str | None
|
try:
|
||||||
user_agent: str | None
|
p = self.permissions[perm_uuid]
|
||||||
expiry: datetime
|
except KeyError:
|
||||||
|
continue
|
||||||
|
# Check domain restriction
|
||||||
|
if p.domain is not None and p.domain != host_without_port:
|
||||||
|
continue
|
||||||
|
effective_perms.append(p)
|
||||||
|
|
||||||
|
return SessionContext(
|
||||||
class _ResetTokenData(msgspec.Struct):
|
session=s,
|
||||||
user: UUID
|
user=user,
|
||||||
expiry: datetime
|
org=org,
|
||||||
token_type: str
|
role=role,
|
||||||
|
credential=credential,
|
||||||
|
permissions=effective_perms,
|
||||||
class _DatabaseData(msgspec.Struct, omit_defaults=True):
|
)
|
||||||
permissions: dict[UUID, _PermissionData]
|
|
||||||
orgs: dict[UUID, _OrgData]
|
|
||||||
roles: dict[UUID, _RoleData]
|
|
||||||
users: dict[UUID, _UserData]
|
|
||||||
credentials: dict[UUID, _CredentialData]
|
|
||||||
sessions: dict[str, _SessionData]
|
|
||||||
reset_tokens: dict[bytes, _ResetTokenData]
|
|
||||||
v: int = 0
|
|
||||||
|
|||||||
+24
-27
@@ -5,13 +5,13 @@ import logging
|
|||||||
import os
|
import os
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
|
|
||||||
import uvicorn
|
|
||||||
from fastapi_vue.hostutil import parse_endpoint
|
from fastapi_vue.hostutil import parse_endpoint
|
||||||
from uvicorn import Config, Server
|
from uvicorn import Config, Server
|
||||||
|
|
||||||
from paskia import globals as _globals
|
from paskia import globals as _globals
|
||||||
from paskia.bootstrap import bootstrap_if_needed
|
from paskia.bootstrap import bootstrap_if_needed
|
||||||
from paskia.config import PaskiaConfig
|
from paskia.config import PaskiaConfig
|
||||||
|
from paskia.db.background import flush
|
||||||
from paskia.fastapi import app as fastapi_app
|
from paskia.fastapi import app as fastapi_app
|
||||||
from paskia.fastapi import reset as reset_cmd
|
from paskia.fastapi import reset as reset_cmd
|
||||||
from paskia.util import startupbox
|
from paskia.util import startupbox
|
||||||
@@ -183,32 +183,13 @@ def main():
|
|||||||
}
|
}
|
||||||
os.environ["PASKIA_CONFIG"] = json.dumps(config_json)
|
os.environ["PASKIA_CONFIG"] = json.dumps(config_json)
|
||||||
|
|
||||||
# Initialize globals (without bootstrap yet)
|
|
||||||
asyncio.run(
|
|
||||||
_globals.init(
|
|
||||||
rp_id=config.rp_id,
|
|
||||||
rp_name=config.rp_name,
|
|
||||||
origins=config.origins,
|
|
||||||
bootstrap=False,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Print startup configuration
|
|
||||||
startupbox.print_startup_config(config)
|
startupbox.print_startup_config(config)
|
||||||
|
|
||||||
# Bootstrap after startup box is printed
|
|
||||||
asyncio.run(bootstrap_if_needed())
|
|
||||||
|
|
||||||
# Handle reset command (no server start)
|
|
||||||
if is_reset:
|
|
||||||
exit_code = reset_cmd.run(args.reset_query)
|
|
||||||
raise SystemExit(exit_code)
|
|
||||||
|
|
||||||
# Dev mode: enable reload when FASTAPI_VUE_FRONTEND_URL is set
|
|
||||||
devmode = bool(os.environ.get("FASTAPI_VUE_FRONTEND_URL"))
|
devmode = bool(os.environ.get("FASTAPI_VUE_FRONTEND_URL"))
|
||||||
|
|
||||||
run_kwargs: dict = {
|
run_kwargs: dict = {
|
||||||
"log_level": "info",
|
"log_level": "info",
|
||||||
|
"access_log": False, # We use custom AccessLogMiddleware instead
|
||||||
}
|
}
|
||||||
|
|
||||||
if devmode:
|
if devmode:
|
||||||
@@ -221,18 +202,34 @@ def main():
|
|||||||
# Suppress uvicorn startup messages in dev mode
|
# Suppress uvicorn startup messages in dev mode
|
||||||
run_kwargs["log_level"] = "warning"
|
run_kwargs["log_level"] = "warning"
|
||||||
|
|
||||||
if len(endpoints) > 1:
|
async def async_main():
|
||||||
# Run separate servers for multiple endpoints (e.g. IPv4 + IPv6)
|
await _globals.init(
|
||||||
async def serve_all():
|
rp_id=config.rp_id,
|
||||||
|
rp_name=config.rp_name,
|
||||||
|
origins=config.origins,
|
||||||
|
bootstrap=False,
|
||||||
|
)
|
||||||
|
await bootstrap_if_needed()
|
||||||
|
await flush()
|
||||||
|
|
||||||
|
if is_reset:
|
||||||
|
exit_code = reset_cmd.run(args.reset_query)
|
||||||
|
raise SystemExit(exit_code)
|
||||||
|
|
||||||
|
if len(endpoints) > 1:
|
||||||
async with asyncio.TaskGroup() as tg:
|
async with asyncio.TaskGroup() as tg:
|
||||||
for ep in endpoints:
|
for ep in endpoints:
|
||||||
tg.create_task(
|
tg.create_task(
|
||||||
Server(Config(app=fastapi_app, **run_kwargs, **ep)).serve()
|
Server(Config(app=fastapi_app, **run_kwargs, **ep)).serve()
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
server = Server(Config(app=fastapi_app, **run_kwargs, **endpoints[0]))
|
||||||
|
await server.serve()
|
||||||
|
|
||||||
asyncio.run(serve_all())
|
try:
|
||||||
else:
|
asyncio.run(async_main())
|
||||||
uvicorn.run("paskia.fastapi:app", **run_kwargs, **endpoints[0])
|
except KeyboardInterrupt:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
|
|||||||
+149
-331
@@ -1,58 +1,45 @@
|
|||||||
import logging
|
import logging
|
||||||
from datetime import timezone
|
from uuid import UUID
|
||||||
from uuid import UUID, uuid4
|
|
||||||
|
|
||||||
from fastapi import Body, FastAPI, HTTPException, Request, Response
|
from fastapi import Body, FastAPI, HTTPException, Query, Request, Response
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
|
|
||||||
|
from paskia import aaguid as aaguid_mod
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.authsession import EXPIRES, reset_expires
|
from paskia.authsession import EXPIRES, reset_expires
|
||||||
|
from paskia.db import Org as OrgDC
|
||||||
|
from paskia.db import Permission as PermDC
|
||||||
|
from paskia.db import Role as RoleDC
|
||||||
|
from paskia.db import User as UserDC
|
||||||
from paskia.fastapi import authz
|
from paskia.fastapi import authz
|
||||||
|
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.util import (
|
from paskia.util import (
|
||||||
hostutil,
|
hostutil,
|
||||||
passphrase,
|
passphrase,
|
||||||
permutil,
|
permutil,
|
||||||
querysafe,
|
querysafe,
|
||||||
useragent,
|
|
||||||
vitedev,
|
vitedev,
|
||||||
)
|
)
|
||||||
|
from paskia.util.apistructs import ApiPermission, ApiSession, format_datetime
|
||||||
|
from paskia.util.hostutil import normalize_host
|
||||||
|
|
||||||
app = FastAPI()
|
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
|
||||||
|
|
||||||
def is_global_admin(ctx) -> bool:
|
def master_admin(ctx) -> bool:
|
||||||
"""Check if user has global admin permission."""
|
return any(p.scope == "auth:admin" for p in ctx.permissions)
|
||||||
effective_scopes = (
|
|
||||||
{p.scope for p in (ctx.permissions or [])}
|
|
||||||
if ctx.permissions
|
def org_admin(ctx, org_uuid: UUID) -> bool:
|
||||||
else set(ctx.role.permissions or [])
|
return ctx.org.uuid == org_uuid and any(
|
||||||
|
p.scope == "auth:org:admin" for p in ctx.permissions
|
||||||
)
|
)
|
||||||
return "auth:admin" in effective_scopes
|
|
||||||
|
|
||||||
|
|
||||||
def is_org_admin(ctx, org_uuid: UUID | None = None) -> bool:
|
|
||||||
"""Check if user has org admin permission.
|
|
||||||
|
|
||||||
If org_uuid is provided, checks if user is admin of that specific org.
|
|
||||||
If org_uuid is None, checks if user is admin of their own org.
|
|
||||||
"""
|
|
||||||
effective_scopes = (
|
|
||||||
{p.scope for p in (ctx.permissions or [])}
|
|
||||||
if ctx.permissions
|
|
||||||
else set(ctx.role.permissions or [])
|
|
||||||
)
|
|
||||||
if "auth:org:admin" not in effective_scopes:
|
|
||||||
return False
|
|
||||||
if org_uuid is None:
|
|
||||||
return True
|
|
||||||
# User must belong to the target org (via their role)
|
|
||||||
return ctx.org.uuid == org_uuid
|
|
||||||
|
|
||||||
|
|
||||||
def can_manage_org(ctx, org_uuid: UUID) -> bool:
|
def can_manage_org(ctx, org_uuid: UUID) -> bool:
|
||||||
"""Check if user can manage the specified organization."""
|
return master_admin(ctx) or org_admin(ctx, org_uuid)
|
||||||
return is_global_admin(ctx) or is_org_admin(ctx, org_uuid)
|
|
||||||
|
|
||||||
|
|
||||||
@app.exception_handler(ValueError)
|
@app.exception_handler(ValueError)
|
||||||
@@ -91,39 +78,39 @@ async def admin_list_orgs(request: Request, auth=AUTH_COOKIE):
|
|||||||
match=permutil.has_any,
|
match=permutil.has_any,
|
||||||
host=request.headers.get("host"),
|
host=request.headers.get("host"),
|
||||||
)
|
)
|
||||||
orgs = db.list_organizations()
|
orgs = list(db.data().orgs.values())
|
||||||
if not is_global_admin(ctx):
|
if not master_admin(ctx):
|
||||||
# Org admins can only see their own organization
|
# Org admins can only see their own organization
|
||||||
orgs = [o for o in orgs if o.uuid == ctx.org.uuid]
|
orgs = [o for o in orgs if o.uuid == ctx.org.uuid]
|
||||||
|
|
||||||
def role_to_dict(r):
|
def org_to_dict(o):
|
||||||
|
users = db.get_organization_users(o.uuid)
|
||||||
return {
|
return {
|
||||||
"uuid": str(r.uuid),
|
"uuid": o.uuid,
|
||||||
"org_uuid": str(r.org_uuid),
|
|
||||||
"display_name": r.display_name,
|
|
||||||
"permissions": r.permissions,
|
|
||||||
}
|
|
||||||
|
|
||||||
async def org_to_dict(o):
|
|
||||||
users = db.get_organization_users(str(o.uuid))
|
|
||||||
return {
|
|
||||||
"uuid": str(o.uuid),
|
|
||||||
"display_name": o.display_name,
|
"display_name": o.display_name,
|
||||||
"permissions": o.permissions,
|
"permissions": {p.uuid for p in o.permissions},
|
||||||
"roles": [role_to_dict(r) for r in o.roles],
|
"roles": [
|
||||||
|
{
|
||||||
|
"uuid": r.uuid,
|
||||||
|
"org": r.org_uuid,
|
||||||
|
"display_name": r.display_name,
|
||||||
|
"permissions": list(r.permissions.keys()),
|
||||||
|
}
|
||||||
|
for r in o.roles
|
||||||
|
],
|
||||||
"users": [
|
"users": [
|
||||||
{
|
{
|
||||||
"uuid": str(u.uuid),
|
"uuid": u.uuid,
|
||||||
"display_name": u.display_name,
|
"display_name": u.display_name,
|
||||||
"role": role_name,
|
"role": role_name,
|
||||||
"visits": u.visits,
|
"visits": u.visits,
|
||||||
"last_seen": u.last_seen.isoformat() if u.last_seen else None,
|
"last_seen": u.last_seen,
|
||||||
}
|
}
|
||||||
for (u, role_name) in users
|
for (u, role_name) in users
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
|
|
||||||
return [await org_to_dict(o) for o in orgs]
|
return MsgspecResponse([org_to_dict(o) for o in orgs])
|
||||||
|
|
||||||
|
|
||||||
@app.post("/orgs")
|
@app.post("/orgs")
|
||||||
@@ -133,15 +120,16 @@ async def admin_create_org(
|
|||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
||||||
)
|
)
|
||||||
from ..db import Org as OrgDC # local import to avoid cycles
|
|
||||||
|
|
||||||
org_uuid = uuid4()
|
|
||||||
display_name = payload.get("display_name") or "New Organization"
|
display_name = payload.get("display_name") or "New Organization"
|
||||||
permissions = payload.get("permissions") or []
|
permissions = payload.get("permissions") or []
|
||||||
org = OrgDC(uuid=org_uuid, display_name=display_name, permissions=permissions)
|
org = OrgDC.create(display_name=display_name)
|
||||||
db.create_organization(org, ctx=ctx)
|
db.create_org(org, ctx=ctx)
|
||||||
|
# Grant requested permissions to the new org
|
||||||
|
for perm in permissions:
|
||||||
|
db.add_permission_to_org(str(org.uuid), perm)
|
||||||
|
|
||||||
return {"uuid": str(org_uuid)}
|
return {"uuid": str(org.uuid)}
|
||||||
|
|
||||||
|
|
||||||
@app.patch("/orgs/{org_uuid}")
|
@app.patch("/orgs/{org_uuid}")
|
||||||
@@ -166,7 +154,7 @@ async def admin_update_org_name(
|
|||||||
if not display_name:
|
if not display_name:
|
||||||
raise ValueError("display_name is required")
|
raise ValueError("display_name is required")
|
||||||
|
|
||||||
db.update_organization_name(org_uuid, display_name, ctx=ctx)
|
db.update_org_name(org_uuid, display_name, ctx=ctx)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
@@ -188,7 +176,7 @@ async def admin_delete_org(org_uuid: UUID, request: Request, auth=AUTH_COOKIE):
|
|||||||
|
|
||||||
# Delete organization-specific permissions
|
# Delete organization-specific permissions
|
||||||
org_perm_pattern = f"org:{str(org_uuid).lower()}"
|
org_perm_pattern = f"org:{str(org_uuid).lower()}"
|
||||||
all_permissions = db.list_permissions()
|
all_permissions = list(db.data().permissions.values())
|
||||||
for perm in all_permissions:
|
for perm in all_permissions:
|
||||||
perm_scope_lower = perm.scope.lower()
|
perm_scope_lower = perm.scope.lower()
|
||||||
# Check if permission contains "org:{uuid}" separated by colons or at boundaries
|
# Check if permission contains "org:{uuid}" separated by colons or at boundaries
|
||||||
@@ -198,39 +186,43 @@ async def admin_delete_org(org_uuid: UUID, request: Request, auth=AUTH_COOKIE):
|
|||||||
or perm_scope_lower.endswith(f":{org_perm_pattern}")
|
or perm_scope_lower.endswith(f":{org_perm_pattern}")
|
||||||
or perm_scope_lower == org_perm_pattern
|
or perm_scope_lower == org_perm_pattern
|
||||||
):
|
):
|
||||||
db.delete_permission(str(perm.uuid), ctx=ctx)
|
db.delete_permission(perm.uuid, ctx=ctx)
|
||||||
|
|
||||||
db.delete_organization(org_uuid, ctx=ctx)
|
db.delete_org(org_uuid, ctx=ctx)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
@app.post("/orgs/{org_uuid}/permission")
|
@app.post("/orgs/{org_uuid}/permission")
|
||||||
async def admin_add_org_permission(
|
async def admin_add_org_permission(
|
||||||
org_uuid: UUID,
|
org_uuid: UUID,
|
||||||
permission_id: str,
|
|
||||||
request: Request,
|
request: Request,
|
||||||
|
permission_uuid: UUID = Query(...),
|
||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
):
|
):
|
||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
||||||
)
|
)
|
||||||
db.add_permission_to_organization(str(org_uuid), permission_id, ctx=ctx)
|
|
||||||
|
db.add_permission_to_org(org_uuid, permission_uuid, ctx=ctx)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
@app.delete("/orgs/{org_uuid}/permission")
|
@app.delete("/orgs/{org_uuid}/permission")
|
||||||
async def admin_remove_org_permission(
|
async def admin_remove_org_permission(
|
||||||
org_uuid: UUID,
|
org_uuid: UUID,
|
||||||
permission_id: str,
|
|
||||||
request: Request,
|
request: Request,
|
||||||
|
permission_uuid: UUID = Query(...),
|
||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
):
|
):
|
||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
||||||
)
|
)
|
||||||
|
|
||||||
|
db.remove_permission_from_org(org_uuid, permission_uuid, ctx=ctx)
|
||||||
|
|
||||||
# Guard rail: prevent removing auth:admin from your own org if it would lock you out
|
# Guard rail: prevent removing auth:admin from your own org if it would lock you out
|
||||||
if permission_id == "auth:admin" and ctx.org.uuid == org_uuid:
|
perm = db.data().permissions.get(permission_uuid)
|
||||||
|
if perm and perm.scope == "auth:admin" and ctx.org.uuid == org_uuid:
|
||||||
# Check if any other org grants auth:admin that we're a member of
|
# Check if any other org grants auth:admin that we're a member of
|
||||||
# (we only know our current org, so this effectively means we can't remove it from our own org)
|
# (we only know our current org, so this effectively means we can't remove it from our own org)
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -238,7 +230,7 @@ async def admin_remove_org_permission(
|
|||||||
"This would lock you out of admin access."
|
"This would lock you out of admin access."
|
||||||
)
|
)
|
||||||
|
|
||||||
db.remove_permission_from_organization(str(org_uuid), permission_id, ctx=ctx)
|
db.remove_permission_from_org(org_uuid, permission_uuid, ctx=ctx)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
@@ -262,33 +254,31 @@ async def admin_create_role(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
from ..db import Role as RoleDC
|
|
||||||
|
|
||||||
role_uuid = uuid4()
|
|
||||||
display_name = payload.get("display_name") or "New Role"
|
display_name = payload.get("display_name") or "New Role"
|
||||||
perms = payload.get("permissions") or []
|
perms = payload.get("permissions") or []
|
||||||
org = db.get_organization(str(org_uuid))
|
if org_uuid not in db.data().orgs:
|
||||||
grantable = set(org.permissions or [])
|
raise HTTPException(status_code=404, detail="Organization not found")
|
||||||
|
org = db.data().orgs[org_uuid]
|
||||||
|
grantable = {p.uuid for p in org.permissions}
|
||||||
|
|
||||||
# Normalize permission IDs to UUIDs
|
# Normalize permission IDs to UUIDs
|
||||||
permission_uuids = []
|
permission_uuids: set[UUID] = set()
|
||||||
for pid in perms:
|
for pid in perms:
|
||||||
perm = db.get_permission(pid)
|
perm = db.data().permissions.get(UUID(pid))
|
||||||
if not perm:
|
if not perm:
|
||||||
raise ValueError(f"Permission {pid} not found")
|
raise ValueError(f"Permission {pid} not found")
|
||||||
perm_uuid_str = str(perm.uuid)
|
if perm.uuid not in grantable:
|
||||||
if perm_uuid_str not in grantable:
|
|
||||||
raise ValueError(f"Permission not grantable by org: {pid}")
|
raise ValueError(f"Permission not grantable by org: {pid}")
|
||||||
permission_uuids.append(perm_uuid_str)
|
permission_uuids.add(perm.uuid)
|
||||||
|
|
||||||
role = RoleDC(
|
role = RoleDC.create(
|
||||||
uuid=role_uuid,
|
org=org_uuid,
|
||||||
org_uuid=org_uuid,
|
|
||||||
display_name=display_name,
|
display_name=display_name,
|
||||||
permissions=permission_uuids,
|
permissions=permission_uuids,
|
||||||
)
|
)
|
||||||
db.create_role(role, ctx=ctx)
|
db.create_role(role, ctx=ctx)
|
||||||
return {"uuid": str(role_uuid)}
|
return {"uuid": str(role.uuid)}
|
||||||
|
|
||||||
|
|
||||||
@app.patch("/orgs/{org_uuid}/roles/{role_uuid}")
|
@app.patch("/orgs/{org_uuid}/roles/{role_uuid}")
|
||||||
@@ -310,8 +300,8 @@ async def admin_update_role_name(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
role = db.get_role(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if role.org_uuid != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
display_name = payload.get("display_name")
|
display_name = payload.get("display_name")
|
||||||
@@ -342,16 +332,15 @@ async def admin_add_role_permission(
|
|||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
|
|
||||||
role = db.get_role(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if role.org_uuid != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
# Verify permission exists and org can grant it
|
# Verify permission exists and org can grant it
|
||||||
perm = db.get_permission(permission_uuid)
|
perm = db.data().permissions.get(permission_uuid)
|
||||||
if not perm:
|
if not perm:
|
||||||
raise HTTPException(status_code=404, detail="Permission not found")
|
raise HTTPException(status_code=404, detail="Permission not found")
|
||||||
org = db.get_organization(str(org_uuid))
|
if org_uuid not in perm.orgs:
|
||||||
if str(permission_uuid) not in org.permissions:
|
|
||||||
raise ValueError("Permission not grantable by organization")
|
raise ValueError("Permission not grantable by organization")
|
||||||
|
|
||||||
db.add_permission_to_role(role_uuid, permission_uuid, ctx=ctx)
|
db.add_permission_to_role(role_uuid, permission_uuid, ctx=ctx)
|
||||||
@@ -378,21 +367,19 @@ async def admin_remove_role_permission(
|
|||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
|
|
||||||
role = db.get_role(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if role.org_uuid != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
# Sanity check: prevent admin from removing their own access
|
# Sanity check: prevent admin from removing their own access
|
||||||
# Find auth:admin and auth:org:admin permission UUIDs
|
perm = db.data().permissions.get(permission_uuid)
|
||||||
perm_uuid_str = str(permission_uuid)
|
|
||||||
perm = db.get_permission(permission_uuid)
|
|
||||||
if ctx.org.uuid == org_uuid and ctx.role.uuid == role_uuid:
|
if ctx.org.uuid == org_uuid and ctx.role.uuid == role_uuid:
|
||||||
if perm and perm.scope in ["auth:admin", "auth:org:admin"]:
|
if perm and perm.scope in ["auth:admin", "auth:org:admin"]:
|
||||||
# Check if removing this permission would leave no admin access
|
# Check if removing this permission would leave no admin access
|
||||||
remaining_perms = set(role.permissions) - {perm_uuid_str}
|
remaining_perms = role.permission_set - {permission_uuid}
|
||||||
has_admin = False
|
has_admin = False
|
||||||
for rp_uuid in remaining_perms:
|
for rp_uuid in remaining_perms:
|
||||||
rp = db.get_permission(rp_uuid)
|
rp = db.data().permissions.get(rp_uuid)
|
||||||
if rp and rp.scope in ["auth:admin", "auth:org:admin"]:
|
if rp and rp.scope in ["auth:admin", "auth:org:admin"]:
|
||||||
has_admin = True
|
has_admin = True
|
||||||
break
|
break
|
||||||
@@ -421,8 +408,8 @@ async def admin_delete_role(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
role = db.get_role(role_uuid)
|
role = db.data().roles.get(role_uuid)
|
||||||
if role.org_uuid != org_uuid:
|
if not role or role.org_uuid != org_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Role not found in organization")
|
raise HTTPException(status_code=404, detail="Role not found in organization")
|
||||||
|
|
||||||
# Sanity check: prevent admin from deleting their own role
|
# Sanity check: prevent admin from deleting their own role
|
||||||
@@ -457,22 +444,20 @@ async def admin_create_user(
|
|||||||
role_name = payload.get("role")
|
role_name = payload.get("role")
|
||||||
if not display_name or not role_name:
|
if not display_name or not role_name:
|
||||||
raise ValueError("display_name and role are required")
|
raise ValueError("display_name and role are required")
|
||||||
from ..db import User as UserDC
|
|
||||||
|
|
||||||
roles = db.get_roles_by_organization(str(org_uuid))
|
org = db.data().orgs[org_uuid]
|
||||||
role_obj = next((r for r in roles if r.display_name == role_name), None)
|
role_obj = next(
|
||||||
|
(r for r in org.roles if r.display_name == role_name),
|
||||||
|
None,
|
||||||
|
)
|
||||||
if not role_obj:
|
if not role_obj:
|
||||||
raise ValueError("Role not found in organization")
|
raise ValueError("Role not found in organization")
|
||||||
user_uuid = uuid4()
|
user = UserDC.create(
|
||||||
user = UserDC(
|
|
||||||
uuid=user_uuid,
|
|
||||||
display_name=display_name,
|
display_name=display_name,
|
||||||
role_uuid=role_obj.uuid,
|
role=role_obj.uuid,
|
||||||
visits=0,
|
|
||||||
created_at=None,
|
|
||||||
)
|
)
|
||||||
db.create_user(user, ctx=ctx)
|
db.create_user(user, ctx=ctx)
|
||||||
return {"uuid": str(user_uuid)}
|
return {"uuid": str(user.uuid)}
|
||||||
|
|
||||||
|
|
||||||
@app.patch("/orgs/{org_uuid}/users/{user_uuid}/role")
|
@app.patch("/orgs/{org_uuid}/users/{user_uuid}/role")
|
||||||
@@ -502,7 +487,7 @@ async def admin_update_user_role(
|
|||||||
raise ValueError("User not found")
|
raise ValueError("User not found")
|
||||||
if user_org.uuid != org_uuid:
|
if user_org.uuid != org_uuid:
|
||||||
raise ValueError("User does not belong to this organization")
|
raise ValueError("User does not belong to this organization")
|
||||||
roles = db.get_roles_by_organization(str(org_uuid))
|
roles = user_org.roles
|
||||||
if not any(r.display_name == new_role for r in roles):
|
if not any(r.display_name == new_role for r in roles):
|
||||||
raise ValueError("Role not found in organization")
|
raise ValueError("Role not found in organization")
|
||||||
|
|
||||||
@@ -513,7 +498,7 @@ async def admin_update_user_role(
|
|||||||
# Check if any permission in the new role is an admin permission
|
# Check if any permission in the new role is an admin permission
|
||||||
has_admin_access = False
|
has_admin_access = False
|
||||||
for perm_uuid in new_role_obj.permissions:
|
for perm_uuid in new_role_obj.permissions:
|
||||||
perm = db.get_permission(perm_uuid)
|
perm = db.data().permissions.get(perm_uuid)
|
||||||
if perm and perm.scope in ["auth:admin", "auth:org:admin"]:
|
if perm and perm.scope in ["auth:admin", "auth:org:admin"]:
|
||||||
has_admin_access = True
|
has_admin_access = True
|
||||||
break
|
break
|
||||||
@@ -552,8 +537,8 @@ async def admin_create_user_registration_link(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Check if user has existing credentials
|
# Check if user has existing credentials
|
||||||
credentials = db.get_credentials_by_user_uuid(user_uuid)
|
has_credentials = db.get_user_credential_ids(user_uuid)
|
||||||
token_type = "user registration" if not credentials else "account recovery"
|
token_type = "user registration" if not has_credentials else "account recovery"
|
||||||
|
|
||||||
token = passphrase.generate()
|
token = passphrase.generate()
|
||||||
expiry = reset_expires()
|
expiry = reset_expires()
|
||||||
@@ -567,11 +552,7 @@ async def admin_create_user_registration_link(
|
|||||||
url = hostutil.reset_link_url(token)
|
url = hostutil.reset_link_url(token)
|
||||||
return {
|
return {
|
||||||
"url": url,
|
"url": url,
|
||||||
"expires": (
|
"expires": format_datetime(expiry),
|
||||||
expiry.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if expiry.tzinfo
|
|
||||||
else expiry.replace(tzinfo=timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -598,122 +579,40 @@ async def admin_get_user_detail(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
user = db.get_user_by_uuid(user_uuid)
|
user = db.data().users.get(user_uuid)
|
||||||
user_creds = db.get_credentials_by_user_uuid(user_uuid)
|
normalized_host = hostutil.normalize_host(request.headers.get("host"))
|
||||||
creds: list[dict] = []
|
|
||||||
aaguids: set[str] = set()
|
return MsgspecResponse(
|
||||||
for c in user_creds:
|
{
|
||||||
aaguid_str = str(c.aaguid)
|
"display_name": user.display_name,
|
||||||
aaguids.add(aaguid_str)
|
"org": {"display_name": user_org.display_name},
|
||||||
creds.append(
|
"role": role_name,
|
||||||
{
|
"visits": user.visits,
|
||||||
"credential_uuid": str(c.uuid),
|
"created_at": user.created_at,
|
||||||
"aaguid": aaguid_str,
|
"last_seen": user.last_seen,
|
||||||
"created_at": (
|
"credentials": [
|
||||||
c.created_at.astimezone(timezone.utc)
|
{
|
||||||
.isoformat()
|
"credential": c.uuid,
|
||||||
.replace("+00:00", "Z")
|
"aaguid": c.aaguid,
|
||||||
if c.created_at.tzinfo
|
"created_at": c.created_at,
|
||||||
else c.created_at.replace(tzinfo=timezone.utc)
|
"last_used": c.last_used,
|
||||||
.isoformat()
|
"last_verified": c.last_verified,
|
||||||
.replace("+00:00", "Z")
|
"sign_count": c.sign_count,
|
||||||
),
|
}
|
||||||
"last_used": (
|
for c in user.credentials
|
||||||
c.last_used.astimezone(timezone.utc)
|
],
|
||||||
.isoformat()
|
"aaguid_info": aaguid_mod.filter(c.aaguid for c in user.credentials),
|
||||||
.replace("+00:00", "Z")
|
"sessions": [
|
||||||
if c.last_used and c.last_used.tzinfo
|
ApiSession.from_db(
|
||||||
else (
|
s,
|
||||||
c.last_used.replace(tzinfo=timezone.utc)
|
current_key=auth,
|
||||||
.isoformat()
|
normalized_host=normalized_host,
|
||||||
.replace("+00:00", "Z")
|
expires_delta=EXPIRES,
|
||||||
if c.last_used
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
),
|
|
||||||
"last_verified": (
|
|
||||||
c.last_verified.astimezone(timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if c.last_verified and c.last_verified.tzinfo
|
|
||||||
else (
|
|
||||||
c.last_verified.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if c.last_verified
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
if c.last_verified
|
for s in user.sessions
|
||||||
else None,
|
],
|
||||||
"sign_count": c.sign_count,
|
}
|
||||||
}
|
)
|
||||||
)
|
|
||||||
from .. import aaguid as aaguid_mod
|
|
||||||
|
|
||||||
aaguid_info = aaguid_mod.filter(aaguids)
|
|
||||||
|
|
||||||
# Get sessions for the user
|
|
||||||
normalized_request_host = hostutil.normalize_host(request.headers.get("host"))
|
|
||||||
session_records = db.list_sessions_for_user(user_uuid)
|
|
||||||
current_session_key = auth
|
|
||||||
sessions_payload: list[dict] = []
|
|
||||||
for entry in session_records:
|
|
||||||
renewed = entry.expiry - EXPIRES
|
|
||||||
sessions_payload.append(
|
|
||||||
{
|
|
||||||
"id": entry.key,
|
|
||||||
"credential_uuid": str(entry.credential_uuid),
|
|
||||||
"host": entry.host,
|
|
||||||
"ip": entry.ip,
|
|
||||||
"user_agent": useragent.compact_user_agent(entry.user_agent),
|
|
||||||
"last_renewed": (
|
|
||||||
renewed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if renewed.tzinfo
|
|
||||||
else renewed.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
),
|
|
||||||
"is_current": entry.key == current_session_key,
|
|
||||||
"is_current_host": bool(
|
|
||||||
normalized_request_host
|
|
||||||
and entry.host
|
|
||||||
and entry.host == normalized_request_host
|
|
||||||
),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"display_name": user.display_name,
|
|
||||||
"org": {"display_name": user_org.display_name},
|
|
||||||
"role": role_name,
|
|
||||||
"visits": user.visits,
|
|
||||||
"created_at": (
|
|
||||||
user.created_at.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if user.created_at and user.created_at.tzinfo
|
|
||||||
else (
|
|
||||||
user.created_at.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if user.created_at
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
),
|
|
||||||
"last_seen": (
|
|
||||||
user.last_seen.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
|
||||||
if user.last_seen and user.last_seen.tzinfo
|
|
||||||
else (
|
|
||||||
user.last_seen.replace(tzinfo=timezone.utc)
|
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if user.last_seen
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
),
|
|
||||||
"credentials": creds,
|
|
||||||
"aaguid_info": aaguid_info,
|
|
||||||
"sessions": sessions_payload,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.patch("/orgs/{org_uuid}/users/{user_uuid}/display-name")
|
@app.patch("/orgs/{org_uuid}/users/{user_uuid}/display-name")
|
||||||
@@ -803,7 +702,7 @@ async def admin_delete_user_session(
|
|||||||
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
status_code=403, detail="Insufficient permissions", mode="forbidden"
|
||||||
)
|
)
|
||||||
|
|
||||||
target_session = db.get_session(session_id)
|
target_session = db.data().sessions.get(session_id)
|
||||||
if not target_session or target_session.user_uuid != user_uuid:
|
if not target_session or target_session.user_uuid != user_uuid:
|
||||||
raise HTTPException(status_code=404, detail="Session not found")
|
raise HTTPException(status_code=404, detail="Session not found")
|
||||||
|
|
||||||
@@ -817,19 +716,10 @@ async def admin_delete_user_session(
|
|||||||
# -------------------- Permissions (global) --------------------
|
# -------------------- Permissions (global) --------------------
|
||||||
|
|
||||||
|
|
||||||
def _perm_to_dict(p):
|
|
||||||
"""Convert Permission to dict, omitting domain if None."""
|
|
||||||
d = {"uuid": str(p.uuid), "scope": p.scope, "display_name": p.display_name}
|
|
||||||
if p.domain is not None:
|
|
||||||
d["domain"] = p.domain
|
|
||||||
return d
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_permission_domain(domain: str | None) -> None:
|
def _validate_permission_domain(domain: str | None) -> None:
|
||||||
"""Validate that domain is rp_id or a subdomain of it."""
|
"""Validate that domain is rp_id or a subdomain of it."""
|
||||||
if domain is None:
|
if domain is None:
|
||||||
return
|
return
|
||||||
from paskia.globals import passkey
|
|
||||||
|
|
||||||
rp_id = passkey.instance.rp_id
|
rp_id = passkey.instance.rp_id
|
||||||
if domain == rp_id or domain.endswith(f".{rp_id}"):
|
if domain == rp_id or domain.endswith(f".{rp_id}"):
|
||||||
@@ -845,13 +735,12 @@ def _check_admin_lockout(
|
|||||||
Raises ValueError if this change would result in no auth:admin permissions
|
Raises ValueError if this change would result in no auth:admin permissions
|
||||||
being accessible from the current host.
|
being accessible from the current host.
|
||||||
"""
|
"""
|
||||||
from paskia.util.hostutil import normalize_host
|
|
||||||
|
|
||||||
normalized_host = normalize_host(current_host)
|
normalized_host = normalize_host(current_host)
|
||||||
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
||||||
|
|
||||||
# Get all auth:admin permissions
|
# Get all auth:admin permissions
|
||||||
all_perms = db.list_permissions()
|
all_perms = list(db.data().permissions.values())
|
||||||
admin_perms = [p for p in all_perms if p.scope == "auth:admin"]
|
admin_perms = [p for p in all_perms if p.scope == "auth:admin"]
|
||||||
|
|
||||||
# Check if at least one auth:admin would remain accessible
|
# Check if at least one auth:admin would remain accessible
|
||||||
@@ -880,13 +769,12 @@ def _check_admin_lockout_on_delete(perm_uuid: str, current_host: str | None) ->
|
|||||||
Raises ValueError if this deletion would result in no auth:admin permissions
|
Raises ValueError if this deletion would result in no auth:admin permissions
|
||||||
being accessible from the current host.
|
being accessible from the current host.
|
||||||
"""
|
"""
|
||||||
from paskia.util.hostutil import normalize_host
|
|
||||||
|
|
||||||
normalized_host = normalize_host(current_host)
|
normalized_host = normalize_host(current_host)
|
||||||
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
host_without_port = normalized_host.rsplit(":", 1)[0] if normalized_host else None
|
||||||
|
|
||||||
# Get all auth:admin permissions
|
# Get all auth:admin permissions
|
||||||
all_perms = db.list_permissions()
|
all_perms = list(db.data().permissions.values())
|
||||||
admin_perms = [p for p in all_perms if p.scope == "auth:admin"]
|
admin_perms = [p for p in all_perms if p.scope == "auth:admin"]
|
||||||
|
|
||||||
# Check if at least one auth:admin would remain accessible after deletion
|
# Check if at least one auth:admin would remain accessible after deletion
|
||||||
@@ -918,16 +806,8 @@ async def admin_list_permissions(request: Request, auth=AUTH_COOKIE):
|
|||||||
match=permutil.has_any,
|
match=permutil.has_any,
|
||||||
host=request.headers.get("host"),
|
host=request.headers.get("host"),
|
||||||
)
|
)
|
||||||
perms = db.list_permissions()
|
perms = db.data().permissions.values() if master_admin(ctx) else ctx.org.permissions
|
||||||
|
return MsgspecResponse([ApiPermission.from_db(p) for p in perms])
|
||||||
# Global admins see all permissions
|
|
||||||
if is_global_admin(ctx):
|
|
||||||
return [_perm_to_dict(p) for p in perms]
|
|
||||||
|
|
||||||
# Org admins only see permissions their org can grant (by UUID)
|
|
||||||
grantable = set(ctx.org.permissions or [])
|
|
||||||
filtered_perms = [p for p in perms if str(p.uuid) in grantable]
|
|
||||||
return [_perm_to_dict(p) for p in filtered_perms]
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/permissions")
|
@app.post("/permissions")
|
||||||
@@ -943,9 +823,6 @@ async def admin_create_permission(
|
|||||||
match=permutil.has_all,
|
match=permutil.has_all,
|
||||||
max_age="5m",
|
max_age="5m",
|
||||||
)
|
)
|
||||||
import uuid7
|
|
||||||
|
|
||||||
from ..db import Permission as PermDC
|
|
||||||
|
|
||||||
scope = payload.get("scope") or payload.get(
|
scope = payload.get("scope") or payload.get(
|
||||||
"id"
|
"id"
|
||||||
@@ -957,9 +834,7 @@ async def admin_create_permission(
|
|||||||
querysafe.assert_safe(scope, field="scope")
|
querysafe.assert_safe(scope, field="scope")
|
||||||
_validate_permission_domain(domain)
|
_validate_permission_domain(domain)
|
||||||
db.create_permission(
|
db.create_permission(
|
||||||
PermDC(
|
PermDC.create(scope=scope, display_name=display_name, domain=domain),
|
||||||
uuid=uuid7.create(), scope=scope, display_name=display_name, domain=domain
|
|
||||||
),
|
|
||||||
ctx=ctx,
|
ctx=ctx,
|
||||||
)
|
)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
@@ -969,29 +844,27 @@ async def admin_create_permission(
|
|||||||
async def admin_update_permission(
|
async def admin_update_permission(
|
||||||
request: Request,
|
request: Request,
|
||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
permission_uuid: str | None = None,
|
permission_uuid: UUID = Query(...),
|
||||||
permission_id: str | None = None, # Backwards compat - treated as scope
|
display_name: str | None = Query(None),
|
||||||
display_name: str | None = None,
|
scope: str | None = Query(None),
|
||||||
scope: str | None = None,
|
domain: str | None = Query(None),
|
||||||
domain: str | None = None,
|
|
||||||
):
|
):
|
||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
||||||
)
|
)
|
||||||
|
|
||||||
# permission_uuid or permission_id (scope) to identify the permission
|
|
||||||
perm_identifier = permission_uuid or permission_id
|
|
||||||
if not perm_identifier:
|
|
||||||
raise ValueError("permission_uuid or permission_id required")
|
|
||||||
|
|
||||||
# Get existing permission
|
# Get existing permission
|
||||||
perm = db.get_permission(perm_identifier)
|
perm = db.data().permissions.get(permission_uuid)
|
||||||
|
|
||||||
# Update fields that were provided
|
# Update fields that were provided
|
||||||
new_scope = scope if scope is not None else perm.scope
|
new_scope = scope if scope is not None else perm.scope
|
||||||
new_display_name = display_name if display_name is not None else perm.display_name
|
new_display_name = display_name if display_name is not None else perm.display_name
|
||||||
domain_value = domain if domain else None
|
domain_value = domain if domain else None
|
||||||
|
|
||||||
|
# Sanity check: prevent changing the auth:admin permission scope
|
||||||
|
if perm.scope == "auth:admin" and new_scope != "auth:admin":
|
||||||
|
raise ValueError("Cannot rename the master admin permission")
|
||||||
|
|
||||||
if not new_display_name:
|
if not new_display_name:
|
||||||
raise ValueError("display_name is required")
|
raise ValueError("display_name is required")
|
||||||
querysafe.assert_safe(new_scope, field="scope")
|
querysafe.assert_safe(new_scope, field="scope")
|
||||||
@@ -1001,71 +874,21 @@ async def admin_update_permission(
|
|||||||
if perm.scope == "auth:admin" or new_scope == "auth:admin":
|
if perm.scope == "auth:admin" or new_scope == "auth:admin":
|
||||||
_check_admin_lockout(str(perm.uuid), domain_value, request.headers.get("host"))
|
_check_admin_lockout(str(perm.uuid), domain_value, request.headers.get("host"))
|
||||||
|
|
||||||
from ..db import Permission as PermDC
|
|
||||||
|
|
||||||
db.update_permission(
|
db.update_permission(
|
||||||
PermDC(
|
uuid=perm.uuid,
|
||||||
uuid=perm.uuid,
|
scope=new_scope,
|
||||||
scope=new_scope,
|
display_name=new_display_name,
|
||||||
display_name=new_display_name,
|
domain=domain_value,
|
||||||
domain=domain_value,
|
|
||||||
),
|
|
||||||
ctx=ctx,
|
ctx=ctx,
|
||||||
)
|
)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
@app.post("/permission/rename")
|
|
||||||
async def admin_rename_permission(
|
|
||||||
request: Request,
|
|
||||||
payload: dict = Body(...),
|
|
||||||
auth=AUTH_COOKIE,
|
|
||||||
):
|
|
||||||
ctx = await authz.verify(
|
|
||||||
auth, ["auth:admin"], host=request.headers.get("host"), match=permutil.has_all
|
|
||||||
)
|
|
||||||
old_scope = payload.get("old_scope") or payload.get("old_id") # Support both
|
|
||||||
new_scope = payload.get("new_scope") or payload.get("new_id") # Support both
|
|
||||||
display_name = payload.get("display_name")
|
|
||||||
domain = payload.get(
|
|
||||||
"domain"
|
|
||||||
) # Can be None (not provided), empty string (clear), or value
|
|
||||||
if not old_scope or not new_scope:
|
|
||||||
raise ValueError("old_scope and new_scope required")
|
|
||||||
|
|
||||||
# Sanity check: prevent renaming critical permissions
|
|
||||||
if old_scope == "auth:admin":
|
|
||||||
raise ValueError("Cannot rename the master admin permission")
|
|
||||||
|
|
||||||
querysafe.assert_safe(old_scope, field="old_scope")
|
|
||||||
querysafe.assert_safe(new_scope, field="new_scope")
|
|
||||||
|
|
||||||
# Get existing permission to preserve values not being changed
|
|
||||||
perm = db.get_permission(old_scope)
|
|
||||||
if display_name is None:
|
|
||||||
display_name = perm.display_name
|
|
||||||
# domain=None means "not provided, keep existing", domain="" means "clear it"
|
|
||||||
if domain is None:
|
|
||||||
domain_value = perm.domain
|
|
||||||
else:
|
|
||||||
domain_value = domain if domain else None
|
|
||||||
_validate_permission_domain(domain_value)
|
|
||||||
|
|
||||||
# Safety check: prevent admin lockout when setting domain on auth:admin
|
|
||||||
if perm.scope == "auth:admin" or new_scope == "auth:admin":
|
|
||||||
_check_admin_lockout(str(perm.uuid), domain_value, request.headers.get("host"))
|
|
||||||
|
|
||||||
# All current backends support rename_permission
|
|
||||||
db.rename_permission(old_scope, new_scope, display_name, domain_value, ctx=ctx)
|
|
||||||
return {"status": "ok"}
|
|
||||||
|
|
||||||
|
|
||||||
@app.delete("/permission")
|
@app.delete("/permission")
|
||||||
async def admin_delete_permission(
|
async def admin_delete_permission(
|
||||||
request: Request,
|
request: Request,
|
||||||
|
permission_uuid: UUID = Query(...),
|
||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
permission_uuid: str | None = None,
|
|
||||||
permission_id: str | None = None, # Backwards compat - treated as scope
|
|
||||||
):
|
):
|
||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth,
|
auth,
|
||||||
@@ -1075,17 +898,12 @@ async def admin_delete_permission(
|
|||||||
max_age="5m",
|
max_age="5m",
|
||||||
)
|
)
|
||||||
|
|
||||||
perm_identifier = permission_uuid or permission_id
|
|
||||||
if not perm_identifier:
|
|
||||||
raise ValueError("permission_uuid or permission_id required")
|
|
||||||
querysafe.assert_safe(perm_identifier, field="permission_id")
|
|
||||||
|
|
||||||
# Get the permission to check its scope
|
# Get the permission to check its scope
|
||||||
perm = db.get_permission(perm_identifier)
|
perm = db.data().permissions.get(permission_uuid)
|
||||||
|
|
||||||
# Sanity check: prevent deleting critical permissions if it would lock out admin
|
# Sanity check: prevent deleting critical permissions if it would lock out admin
|
||||||
if perm.scope == "auth:admin":
|
if perm.scope == "auth:admin":
|
||||||
_check_admin_lockout_on_delete(str(perm.uuid), request.headers.get("host"))
|
_check_admin_lockout_on_delete(str(perm.uuid), request.headers.get("host"))
|
||||||
|
|
||||||
db.delete_permission(str(perm.uuid), ctx=ctx)
|
db.delete_permission(permission_uuid, ctx=ctx)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|||||||
+68
-110
@@ -1,6 +1,6 @@
|
|||||||
import logging
|
import logging
|
||||||
from contextlib import suppress
|
from contextlib import suppress
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import UTC, datetime, timedelta
|
||||||
|
|
||||||
from fastapi import (
|
from fastapi import (
|
||||||
Depends,
|
Depends,
|
||||||
@@ -14,20 +14,16 @@ from fastapi.responses import JSONResponse
|
|||||||
from fastapi.security import HTTPBearer
|
from fastapi.security import HTTPBearer
|
||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.authsession import (
|
from paskia.authsession import EXPIRES, expires, get_reset
|
||||||
EXPIRES,
|
|
||||||
get_reset,
|
|
||||||
get_session,
|
|
||||||
refresh_session_token,
|
|
||||||
)
|
|
||||||
from paskia.fastapi import authz, session, user
|
from paskia.fastapi import authz, session, user
|
||||||
|
from paskia.fastapi.response import MsgspecResponse
|
||||||
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME
|
from paskia.fastapi.session import AUTH_COOKIE, AUTH_COOKIE_NAME
|
||||||
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, vitedev
|
||||||
|
|
||||||
bearer_auth = HTTPBearer(auto_error=True)
|
bearer_auth = HTTPBearer(auto_error=True)
|
||||||
|
|
||||||
app = FastAPI()
|
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
|
||||||
app.mount("/user", user.app)
|
app.mount("/user", user.app)
|
||||||
|
|
||||||
@@ -78,12 +74,7 @@ async def validate_token(
|
|||||||
max_age: str | None = Query(None),
|
max_age: str | None = Query(None),
|
||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
):
|
):
|
||||||
"""Validate the current session and extend its expiry.
|
"""Validate session and return context. Refreshes session expiry."""
|
||||||
|
|
||||||
Always refreshes the session (sliding expiration) and re-sets the cookie with a
|
|
||||||
renewed max-age. This keeps active users logged in without needing a separate
|
|
||||||
refresh endpoint.
|
|
||||||
"""
|
|
||||||
try:
|
try:
|
||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth,
|
auth,
|
||||||
@@ -96,26 +87,23 @@ async def validate_token(
|
|||||||
raise
|
raise
|
||||||
renewed = False
|
renewed = False
|
||||||
if auth:
|
if auth:
|
||||||
consumed = EXPIRES - (ctx.session.expiry - datetime.now(timezone.utc))
|
consumed = EXPIRES - (ctx.session.expiry - datetime.now(UTC))
|
||||||
if not timedelta(0) < consumed < _REFRESH_INTERVAL:
|
if not timedelta(0) < consumed < _REFRESH_INTERVAL:
|
||||||
try:
|
db.update_session(
|
||||||
await refresh_session_token(
|
auth,
|
||||||
auth,
|
ip=request.client.host if request.client else "",
|
||||||
ip=request.client.host if request.client else "",
|
user_agent=request.headers.get("user-agent") or "",
|
||||||
user_agent=request.headers.get("user-agent") or "",
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
session.set_session_cookie(response, auth)
|
session.set_session_cookie(response, auth)
|
||||||
renewed = True
|
renewed = True
|
||||||
except ValueError:
|
return MsgspecResponse(
|
||||||
# Session disappeared, e.g. due to concurrent logout; global handler will clear
|
{
|
||||||
raise authz.AuthException(
|
"valid": True,
|
||||||
status_code=401, detail="Session expired", mode="login"
|
"renewed": renewed,
|
||||||
)
|
"ctx": userinfo.build_session_context(ctx),
|
||||||
return {
|
}
|
||||||
"valid": True,
|
)
|
||||||
"user_uuid": str(ctx.session.user_uuid),
|
|
||||||
"renewed": renewed,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@app.get("/forward")
|
@app.get("/forward")
|
||||||
@@ -144,9 +132,10 @@ async def forward_authentication(
|
|||||||
ctx = await authz.verify(
|
ctx = await authz.verify(
|
||||||
auth, perm, host=request.headers.get("host"), max_age=max_age
|
auth, perm, host=request.headers.get("host"), max_age=max_age
|
||||||
)
|
)
|
||||||
role_permissions = set(ctx.role.permissions or [])
|
# Build permission scopes for Remote-Groups header
|
||||||
if ctx.permissions:
|
role_permissions = (
|
||||||
role_permissions.update(permission.scope for permission in ctx.permissions)
|
{p.scope for p in ctx.permissions} if ctx.permissions else set()
|
||||||
|
)
|
||||||
|
|
||||||
remote_headers: dict[str, str] = {
|
remote_headers: dict[str, str] = {
|
||||||
"Remote-User": str(ctx.user.uuid),
|
"Remote-User": str(ctx.user.uuid),
|
||||||
@@ -157,15 +146,13 @@ async def forward_authentication(
|
|||||||
"Remote-Role": str(ctx.role.uuid),
|
"Remote-Role": str(ctx.role.uuid),
|
||||||
"Remote-Role-Name": ctx.role.display_name,
|
"Remote-Role-Name": ctx.role.display_name,
|
||||||
"Remote-Session-Expires": (
|
"Remote-Session-Expires": (
|
||||||
ctx.session.expiry.astimezone(timezone.utc)
|
ctx.session.expiry.astimezone(UTC).isoformat().replace("+00:00", "Z")
|
||||||
.isoformat()
|
|
||||||
.replace("+00:00", "Z")
|
|
||||||
if ctx.session.expiry.tzinfo
|
if ctx.session.expiry.tzinfo
|
||||||
else ctx.session.expiry.replace(tzinfo=timezone.utc)
|
else ctx.session.expiry.replace(tzinfo=UTC)
|
||||||
.isoformat()
|
.isoformat()
|
||||||
.replace("+00:00", "Z")
|
.replace("+00:00", "Z")
|
||||||
),
|
),
|
||||||
"Remote-Credential": str(ctx.session.credential_uuid),
|
"Remote-Credential": str(ctx.session.credential),
|
||||||
}
|
}
|
||||||
return Response(status_code=204, headers=remote_headers)
|
return Response(status_code=204, headers=remote_headers)
|
||||||
except authz.AuthException as e:
|
except authz.AuthException as e:
|
||||||
@@ -207,92 +194,61 @@ async def get_settings():
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@app.get("/token-info")
|
|
||||||
async def api_token_info(token: str):
|
|
||||||
"""Get information about a reset token.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
- type: "reset"
|
|
||||||
- user_name: display name of the user
|
|
||||||
- token_type: type of reset token
|
|
||||||
"""
|
|
||||||
if not passphrase.is_well_formed(token):
|
|
||||||
raise HTTPException(status_code=404, detail="Invalid token")
|
|
||||||
|
|
||||||
# Check if this is a reset token
|
|
||||||
try:
|
|
||||||
reset_token = await get_reset(token)
|
|
||||||
user = db.get_user_by_uuid(reset_token.user_uuid)
|
|
||||||
return {
|
|
||||||
"type": "reset",
|
|
||||||
"user_name": user.display_name,
|
|
||||||
"token_type": reset_token.token_type,
|
|
||||||
}
|
|
||||||
except (ValueError, Exception):
|
|
||||||
raise HTTPException(status_code=404, detail="Token not found or expired")
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/user-info")
|
@app.post("/user-info")
|
||||||
async def api_user_info(
|
async def api_user_info(
|
||||||
request: Request,
|
request: Request,
|
||||||
response: Response,
|
response: Response,
|
||||||
reset: str | None = None,
|
|
||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
):
|
):
|
||||||
"""Get user information including credentials, sessions, and permissions.
|
"""Get full user profile including credentials and sessions."""
|
||||||
|
if auth is None:
|
||||||
|
raise authz.AuthException(
|
||||||
|
status_code=401,
|
||||||
|
detail="Authentication required",
|
||||||
|
mode="login",
|
||||||
|
)
|
||||||
|
ctx = db.data().session_ctx(auth, request.headers.get("host"))
|
||||||
|
if not ctx:
|
||||||
|
raise HTTPException(401, "Session expired")
|
||||||
|
|
||||||
Can be called with either:
|
return MsgspecResponse(
|
||||||
- A session cookie (auth) for authenticated users
|
await userinfo.build_user_info(
|
||||||
- A reset token for users in password reset flow
|
user_uuid=ctx.user.uuid,
|
||||||
"""
|
auth=auth,
|
||||||
authenticated = False
|
session_record=ctx.session,
|
||||||
session_record = None
|
request_host=request.headers.get("host"),
|
||||||
reset_token = None
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@app.get("/token-info")
|
||||||
|
async def token_info(credentials=Depends(bearer_auth)):
|
||||||
|
"""Get reset/device-add token info. Pass token via Bearer header."""
|
||||||
|
token = credentials.credentials
|
||||||
|
if not passphrase.is_well_formed(token):
|
||||||
|
raise HTTPException(400, "Invalid token format")
|
||||||
try:
|
try:
|
||||||
if reset:
|
reset_token = get_reset(token)
|
||||||
if not passphrase.is_well_formed(reset):
|
|
||||||
raise ValueError("Invalid reset token")
|
|
||||||
reset_token = await get_reset(reset)
|
|
||||||
target_user_uuid = reset_token.user_uuid
|
|
||||||
else:
|
|
||||||
if auth is None:
|
|
||||||
raise authz.AuthException(
|
|
||||||
status_code=401,
|
|
||||||
detail="Authentication required",
|
|
||||||
mode="login",
|
|
||||||
)
|
|
||||||
session_record = await get_session(auth, host=request.headers.get("host"))
|
|
||||||
authenticated = True
|
|
||||||
target_user_uuid = session_record.user_uuid
|
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
raise HTTPException(401, str(e))
|
raise HTTPException(401, str(e))
|
||||||
|
|
||||||
# Return minimal response for reset tokens
|
u = reset_token.user
|
||||||
if not authenticated and reset_token:
|
return {
|
||||||
return await userinfo.format_reset_user_info(target_user_uuid, reset_token)
|
"token_type": reset_token.token_type,
|
||||||
|
"display_name": u.display_name,
|
||||||
# Return full user info for authenticated users
|
}
|
||||||
assert auth is not None
|
|
||||||
assert session_record is not None
|
|
||||||
|
|
||||||
return await userinfo.format_user_info(
|
|
||||||
user_uuid=target_user_uuid,
|
|
||||||
auth=auth,
|
|
||||||
session_record=session_record,
|
|
||||||
request_host=request.headers.get("host"),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@app.post("/logout")
|
@app.post("/logout")
|
||||||
async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
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"}
|
||||||
try:
|
host = request.headers.get("host")
|
||||||
_s = await get_session(auth, host=request.headers.get("host"))
|
ctx = db.data().session_ctx(auth, host)
|
||||||
except ValueError:
|
if not ctx:
|
||||||
return {"message": "Already logged out"}
|
return {"message": "Already logged out"}
|
||||||
with suppress(Exception):
|
with suppress(Exception):
|
||||||
db.delete_session(auth)
|
db.delete_session(auth, ctx=ctx)
|
||||||
session.clear_session_cookie(response)
|
session.clear_session_cookie(response)
|
||||||
return {"message": "Logged out successfully"}
|
return {"message": "Logged out successfully"}
|
||||||
|
|
||||||
@@ -301,9 +257,11 @@ async def api_logout(request: Request, response: Response, auth=AUTH_COOKIE):
|
|||||||
async def api_set_session(
|
async def api_set_session(
|
||||||
request: Request, response: Response, auth=Depends(bearer_auth)
|
request: Request, response: Response, auth=Depends(bearer_auth)
|
||||||
):
|
):
|
||||||
user = await get_session(auth.credentials, host=request.headers.get("host"))
|
ctx = db.data().session_ctx(auth.credentials, request.headers.get("host"))
|
||||||
|
if not ctx:
|
||||||
|
raise HTTPException(401, "Session expired")
|
||||||
session.set_session_cookie(response, auth.credentials)
|
session.set_session_cookie(response, auth.credentials)
|
||||||
return {
|
return {
|
||||||
"message": "Session cookie set successfully",
|
"message": "Session cookie set successfully",
|
||||||
"user_uuid": str(user.user_uuid),
|
"user": str(ctx.user.uuid),
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,218 @@
|
|||||||
|
"""Custom access logging middleware for FastAPI/Uvicorn."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from ipaddress import IPv6Address
|
||||||
|
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
from starlette.requests import Request
|
||||||
|
from starlette.responses import Response
|
||||||
|
|
||||||
|
logger = logging.getLogger("paskia.access")
|
||||||
|
|
||||||
|
_RESET = "\033[0m"
|
||||||
|
_STATUS_INFO = "\033[32m" # 1xx (green)
|
||||||
|
_STATUS_OK = "\033[92m" # 2xx (bright green)
|
||||||
|
_STATUS_REDIRECT = "\033[32m" # 3xx (green)
|
||||||
|
_STATUS_CLIENT_ERR = "\033[0;31m" # 4xx (red)
|
||||||
|
_STATUS_SERVER_ERR = "\033[1;31m" # 5xx (bright red)
|
||||||
|
_METHOD_READ = "\033[0;34m" # GET, HEAD, OPTIONS (blue)
|
||||||
|
_METHOD_WRITE = "\033[1;34m" # POST, PUT, DELETE, PATCH (bright blue)
|
||||||
|
_HOST = "\033[1;30m" # hostname (dark grey)
|
||||||
|
_PATH = "\033[0m" # path (default)
|
||||||
|
_TIMING = "\033[2m" # timing (dim)
|
||||||
|
_WS_OPEN = "\033[1;33m" # WebSocket connect (bright yellow)
|
||||||
|
_WS_CLOSE = "\033[0;33m" # WebSocket disconnect (yellow)
|
||||||
|
_WS_STATUS = "\033[1;30m" # WebSocket close status (dark grey)
|
||||||
|
|
||||||
|
|
||||||
|
def format_ipv6_network(ip: str) -> str:
|
||||||
|
"""Format IPv6 address to show only network part (first 64 bits)."""
|
||||||
|
try:
|
||||||
|
addr = IPv6Address(ip)
|
||||||
|
# Get the integer representation and mask to first 64 bits
|
||||||
|
network_int = int(addr) >> 64
|
||||||
|
# Format as IPv6 with trailing ::
|
||||||
|
# Split into 4 groups of 16 bits
|
||||||
|
groups = []
|
||||||
|
for _ in range(4):
|
||||||
|
groups.insert(0, format(network_int & 0xFFFF, "x"))
|
||||||
|
network_int >>= 16
|
||||||
|
# Compress consecutive zero groups
|
||||||
|
result = ":".join(groups) + "::"
|
||||||
|
# Simplify leading zeros in groups and compress
|
||||||
|
return str(IPv6Address(result + "0"))
|
||||||
|
except Exception:
|
||||||
|
return ip
|
||||||
|
|
||||||
|
|
||||||
|
def format_client_ip(ip: str) -> str:
|
||||||
|
"""Format client IP, compressing IPv6 to network part only."""
|
||||||
|
if not ip or ip == "-":
|
||||||
|
return "-"
|
||||||
|
if ":" in ip:
|
||||||
|
return format_ipv6_network(ip)
|
||||||
|
return ip
|
||||||
|
|
||||||
|
|
||||||
|
def status_color(status: int) -> str:
|
||||||
|
"""Return color code based on HTTP status."""
|
||||||
|
if status < 200:
|
||||||
|
return _STATUS_INFO
|
||||||
|
if status < 300:
|
||||||
|
return _STATUS_OK
|
||||||
|
if status < 400:
|
||||||
|
return _STATUS_REDIRECT
|
||||||
|
if status < 500:
|
||||||
|
return _STATUS_CLIENT_ERR
|
||||||
|
return _STATUS_SERVER_ERR
|
||||||
|
|
||||||
|
|
||||||
|
def method_color(method: str) -> str:
|
||||||
|
"""Return color code based on HTTP method."""
|
||||||
|
if method in ("GET", "HEAD", "OPTIONS"):
|
||||||
|
return _METHOD_READ
|
||||||
|
return _METHOD_WRITE
|
||||||
|
|
||||||
|
|
||||||
|
def format_access_log(
|
||||||
|
client: str, status: int, method: str, host: str, path: str, duration_ms: float
|
||||||
|
) -> str:
|
||||||
|
"""Format access log line with colors and aligned fields."""
|
||||||
|
use_color = sys.stderr.isatty()
|
||||||
|
|
||||||
|
# Format components with fixed widths for alignment
|
||||||
|
ip = format_client_ip(client).ljust(15) # IPv4 max 15 chars
|
||||||
|
timing = f"{duration_ms:.0f}ms"
|
||||||
|
method_padded = method.ljust(7) # Longest method is OPTIONS (7)
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
status_str = f"{status_color(status)}{status}{_RESET}"
|
||||||
|
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||||
|
method_str = f"{method_color(method)}{method_padded}{_RESET}"
|
||||||
|
host_str = f"{_HOST}{host}{_RESET}"
|
||||||
|
path_str = f"{_PATH}{path}{_RESET}"
|
||||||
|
else:
|
||||||
|
status_str = str(status)
|
||||||
|
timing_str = timing
|
||||||
|
method_str = method_padded
|
||||||
|
host_str = host
|
||||||
|
path_str = path
|
||||||
|
|
||||||
|
# Format: "IP STATUS METHOD host path TIMING"
|
||||||
|
return f"{ip} {status_str} {method_str} {host_str}{path_str} {timing_str}"
|
||||||
|
|
||||||
|
|
||||||
|
# WebSocket connection counter (mod 100)
|
||||||
|
_ws_counter = 0
|
||||||
|
|
||||||
|
|
||||||
|
def _next_ws_id() -> int:
|
||||||
|
"""Get next WebSocket connection ID (0-99)."""
|
||||||
|
global _ws_counter
|
||||||
|
ws_id = _ws_counter
|
||||||
|
_ws_counter = (_ws_counter + 1) % 100
|
||||||
|
return ws_id
|
||||||
|
|
||||||
|
|
||||||
|
def log_ws_open(client: str, host: str, path: str) -> int:
|
||||||
|
"""Log WebSocket connection open. Returns connection ID for use in close."""
|
||||||
|
use_color = sys.stderr.isatty()
|
||||||
|
ws_id = _next_ws_id()
|
||||||
|
|
||||||
|
ip = format_client_ip(client).ljust(15)
|
||||||
|
id_str = f"{ws_id:02d}".ljust(7) # Align with method field (7 chars)
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
# 🔌 aligned with status (takes ~2 char width), ID aligned with method
|
||||||
|
prefix = f"🔌 {_WS_OPEN}{id_str}{_RESET}"
|
||||||
|
host_str = f"{_HOST}{host}{_RESET}"
|
||||||
|
path_str = f"{_PATH}{path}{_RESET}"
|
||||||
|
else:
|
||||||
|
prefix = f"WS+ {id_str}"
|
||||||
|
host_str = host
|
||||||
|
path_str = path
|
||||||
|
|
||||||
|
logger.info(f"{ip} {prefix} {host_str}{path_str}")
|
||||||
|
return ws_id
|
||||||
|
|
||||||
|
|
||||||
|
# WebSocket close codes to human-readable status
|
||||||
|
WS_CLOSE_CODES = {
|
||||||
|
1000: "ok",
|
||||||
|
1001: "going away",
|
||||||
|
1002: "protocol error",
|
||||||
|
1003: "unsupported",
|
||||||
|
1005: "no status",
|
||||||
|
1006: "abnormal",
|
||||||
|
1007: "invalid data",
|
||||||
|
1008: "policy violation",
|
||||||
|
1009: "too large",
|
||||||
|
1010: "extension required",
|
||||||
|
1011: "server error",
|
||||||
|
1012: "restarting",
|
||||||
|
1013: "try again",
|
||||||
|
1014: "bad gateway",
|
||||||
|
1015: "tls error",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def log_ws_close(
|
||||||
|
client: str, ws_id: int, close_code: int | None, duration_ms: float
|
||||||
|
) -> None:
|
||||||
|
"""Log WebSocket connection close with duration and status."""
|
||||||
|
use_color = sys.stderr.isatty()
|
||||||
|
|
||||||
|
ip = format_client_ip(client).ljust(15)
|
||||||
|
id_str = f"{ws_id:02d}".ljust(7) # Align with method field (7 chars)
|
||||||
|
timing = f"{duration_ms:.0f}ms"
|
||||||
|
|
||||||
|
# Convert close code to status text
|
||||||
|
if close_code is None:
|
||||||
|
status = "closed"
|
||||||
|
else:
|
||||||
|
status = WS_CLOSE_CODES.get(close_code, f"code {close_code}")
|
||||||
|
|
||||||
|
if use_color:
|
||||||
|
# 🔌 aligned with status, ID aligned with method
|
||||||
|
prefix = f"🔌 {_WS_CLOSE}{id_str}{_RESET}"
|
||||||
|
status_str = f"{_WS_STATUS}{status}{_RESET}"
|
||||||
|
timing_str = f"{_TIMING}{timing}{_RESET}"
|
||||||
|
else:
|
||||||
|
prefix = f"WS- {id_str}"
|
||||||
|
status_str = status
|
||||||
|
timing_str = timing
|
||||||
|
|
||||||
|
logger.info(f"{ip} {prefix} {status_str} {timing_str}")
|
||||||
|
|
||||||
|
|
||||||
|
class AccessLogMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Middleware that logs HTTP requests with custom format."""
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next) -> Response:
|
||||||
|
start = time.perf_counter()
|
||||||
|
response = await call_next(request)
|
||||||
|
duration_ms = (time.perf_counter() - start) * 1000
|
||||||
|
|
||||||
|
client = request.client.host if request.client else "-"
|
||||||
|
host = request.headers.get("host", "-")
|
||||||
|
method = request.method
|
||||||
|
path = request.url.path
|
||||||
|
if request.url.query:
|
||||||
|
path = f"{path}?{request.url.query}"
|
||||||
|
status = response.status_code
|
||||||
|
|
||||||
|
line = format_access_log(client, status, method, host, path, duration_ms)
|
||||||
|
logger.info(line)
|
||||||
|
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
def configure_access_logging():
|
||||||
|
"""Configure the access logger to output to stderr."""
|
||||||
|
handler = logging.StreamHandler(sys.stderr)
|
||||||
|
handler.setFormatter(logging.Formatter("%(message)s"))
|
||||||
|
logger.addHandler(handler)
|
||||||
|
logger.setLevel(logging.INFO)
|
||||||
|
logger.propagate = False
|
||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
@@ -7,13 +8,21 @@ 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 fastapi_vue import Frontend
|
||||||
|
|
||||||
|
from paskia import globals
|
||||||
|
from paskia.db import start_background, stop_background
|
||||||
|
from paskia.db.logging import configure_db_logging
|
||||||
from paskia.fastapi import admin, api, auth_host, ws
|
from paskia.fastapi import admin, api, auth_host, ws
|
||||||
|
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
|
||||||
|
|
||||||
|
# Configure custom logging
|
||||||
|
configure_access_logging()
|
||||||
|
configure_db_logging()
|
||||||
|
|
||||||
# Vue Frontend static files
|
# Vue Frontend static files
|
||||||
frontend = Frontend(
|
frontend = Frontend(
|
||||||
Path(__file__).with_name("frontend-build"),
|
Path(__file__).parent.parent / "frontend-build",
|
||||||
cached=["/auth/assets/"],
|
cached=["/auth/assets/"],
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -30,10 +39,6 @@ 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.
|
||||||
"""
|
"""
|
||||||
import json
|
|
||||||
|
|
||||||
from paskia import globals
|
|
||||||
|
|
||||||
config = json.loads(os.environ["PASKIA_CONFIG"])
|
config = json.loads(os.environ["PASKIA_CONFIG"])
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -49,16 +54,28 @@ async def lifespan(app: FastAPI): # pragma: no cover - startup path
|
|||||||
# Re-raise to fail fast
|
# Re-raise to fail fast
|
||||||
raise
|
raise
|
||||||
|
|
||||||
# Restore info level logging after startup (suppressed during uvicorn init 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
|
||||||
if frontend.devmode:
|
if frontend.devmode:
|
||||||
logging.getLogger("uvicorn").setLevel(logging.INFO)
|
logging.getLogger("uvicorn").setLevel(logging.INFO)
|
||||||
logging.getLogger("uvicorn.access").setLevel(logging.INFO)
|
logging.getLogger("uvicorn.error").setLevel(logging.WARNING)
|
||||||
|
|
||||||
await frontend.load()
|
await frontend.load()
|
||||||
|
await start_background()
|
||||||
yield
|
yield
|
||||||
|
await stop_background()
|
||||||
|
|
||||||
|
|
||||||
app = FastAPI(lifespan=lifespan, redirect_slashes=False)
|
app = FastAPI(
|
||||||
|
lifespan=lifespan,
|
||||||
|
redirect_slashes=False,
|
||||||
|
docs_url=None,
|
||||||
|
redoc_url=None,
|
||||||
|
openapi_url=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Custom access logging (uvicorn's access_log is disabled)
|
||||||
|
app.add_middleware(AccessLogMiddleware)
|
||||||
|
|
||||||
# Apply redirections to auth-host if configured (deny access to restricted endpoints, remove /auth/)
|
# Apply redirections to auth-host if configured (deny access to restricted endpoints, remove /auth/)
|
||||||
app.middleware("http")(auth_host.redirect_middleware)
|
app.middleware("http")(auth_host.redirect_middleware)
|
||||||
|
|||||||
+16
-39
@@ -16,13 +16,14 @@ import base64url
|
|||||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||||
|
|
||||||
from paskia import db, remoteauth
|
from paskia import db, remoteauth
|
||||||
|
from paskia.authsession import expires
|
||||||
from paskia.fastapi.session import infodict
|
from paskia.fastapi.session import infodict
|
||||||
|
from paskia.fastapi.wschat import authenticate_chat
|
||||||
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
|
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
|
||||||
from paskia.globals import passkey
|
from paskia.util import hostutil, passphrase, pow, useragent
|
||||||
from paskia.util import passphrase, pow
|
|
||||||
|
|
||||||
# Create a FastAPI subapp for remote auth WebSocket endpoints
|
# Create a FastAPI subapp for remote auth WebSocket endpoints
|
||||||
app = FastAPI()
|
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
|
||||||
|
|
||||||
@app.websocket("/request")
|
@app.websocket("/request")
|
||||||
@@ -179,7 +180,7 @@ async def websocket_remote_auth_request(ws: WebSocket):
|
|||||||
):
|
):
|
||||||
response = {
|
response = {
|
||||||
"status": "authenticated",
|
"status": "authenticated",
|
||||||
"user_uuid": str(result_data["user_uuid"]),
|
"user": str(result_data["user_uuid"]),
|
||||||
}
|
}
|
||||||
if result_data.get("session_token"):
|
if result_data.get("session_token"):
|
||||||
response["session_token"] = result_data["session_token"]
|
response["session_token"] = result_data["session_token"]
|
||||||
@@ -268,7 +269,6 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
6. Client sends WebAuthn response
|
6. Client sends WebAuthn response
|
||||||
7. Server sends {status: "success", message: "..."}
|
7. Server sends {status: "success", message: "..."}
|
||||||
"""
|
"""
|
||||||
from paskia.util import useragent
|
|
||||||
|
|
||||||
origin = validate_origin(ws)
|
origin = validate_origin(ws)
|
||||||
|
|
||||||
@@ -289,7 +289,6 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
)
|
)
|
||||||
|
|
||||||
request = None
|
request = None
|
||||||
webauthn_challenge = None
|
|
||||||
explicitly_denied = False
|
explicitly_denied = False
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -311,43 +310,21 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
|
|
||||||
# Handle authenticate request (no PoW needed - already validated during lookup)
|
# Handle authenticate request (no PoW needed - already validated during lookup)
|
||||||
if msg.get("authenticate") and request is not None:
|
if msg.get("authenticate") and request is not None:
|
||||||
# Generate authentication options
|
cred, new_sign_count = await authenticate_chat(ws, origin)
|
||||||
options, webauthn_challenge = passkey.instance.auth_generate_options(
|
|
||||||
credential_ids=None
|
|
||||||
)
|
|
||||||
await ws.send_json({"optionsJSON": options})
|
|
||||||
|
|
||||||
# Wait for WebAuthn response
|
|
||||||
credential = passkey.instance.auth_parse(await ws.receive_json())
|
|
||||||
|
|
||||||
# Fetch and verify credential
|
|
||||||
try:
|
|
||||||
stored_cred = db.get_credential_by_id(credential.raw_id)
|
|
||||||
except ValueError:
|
|
||||||
raise ValueError(
|
|
||||||
f"This passkey is no longer registered with {passkey.instance.rp_name}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Verify the credential
|
|
||||||
passkey.instance.auth_verify(
|
|
||||||
credential, webauthn_challenge, stored_cred, origin
|
|
||||||
)
|
|
||||||
|
|
||||||
# Create a session for the REQUESTING device
|
# Create a session for the REQUESTING device
|
||||||
assert stored_cred.uuid is not None
|
assert cred.uuid is not None
|
||||||
|
|
||||||
session_token = None
|
session_token = None
|
||||||
reset_token = None
|
reset_token = None
|
||||||
|
|
||||||
if request.action == "register":
|
if request.action == "register":
|
||||||
# For registration, create a reset token for device addition
|
# For registration, create a reset token for device addition
|
||||||
from paskia.authsession import expires
|
|
||||||
from paskia.util import hostutil
|
|
||||||
|
|
||||||
token_str = passphrase.generate()
|
token_str = passphrase.generate()
|
||||||
expiry = expires()
|
expiry = expires()
|
||||||
db.create_reset_token(
|
db.create_reset_token(
|
||||||
user_uuid=stored_cred.user_uuid,
|
user_uuid=cred.user_uuid,
|
||||||
passphrase=token_str,
|
passphrase=token_str,
|
||||||
expiry=expiry,
|
expiry=expiry,
|
||||||
token_type="device addition",
|
token_type="device addition",
|
||||||
@@ -356,8 +333,9 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
# Also create a session so the device is logged in
|
# Also create a session so the device is logged in
|
||||||
normalized_host = hostutil.normalize_host(request.host)
|
normalized_host = hostutil.normalize_host(request.host)
|
||||||
session_token = db.login(
|
session_token = db.login(
|
||||||
user_uuid=stored_cred.user_uuid,
|
user_uuid=cred.user_uuid,
|
||||||
credential=stored_cred,
|
credential_uuid=cred.uuid,
|
||||||
|
sign_count=new_sign_count,
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
ip=request.ip,
|
ip=request.ip,
|
||||||
user_agent=request.user_agent,
|
user_agent=request.user_agent,
|
||||||
@@ -365,13 +343,12 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Default login action
|
# Default login action
|
||||||
from paskia.authsession import expires
|
|
||||||
from paskia.util import hostutil
|
|
||||||
|
|
||||||
normalized_host = hostutil.normalize_host(request.host)
|
normalized_host = hostutil.normalize_host(request.host)
|
||||||
session_token = db.login(
|
session_token = db.login(
|
||||||
user_uuid=stored_cred.user_uuid,
|
user_uuid=cred.user_uuid,
|
||||||
credential=stored_cred,
|
credential_uuid=cred.uuid,
|
||||||
|
sign_count=new_sign_count,
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
ip=request.ip,
|
ip=request.ip,
|
||||||
user_agent=request.user_agent,
|
user_agent=request.user_agent,
|
||||||
@@ -382,8 +359,8 @@ async def websocket_remote_auth_permit(ws: WebSocket):
|
|||||||
completed = await remoteauth.instance.complete_request(
|
completed = await remoteauth.instance.complete_request(
|
||||||
token=request.key,
|
token=request.key,
|
||||||
session_token=session_token,
|
session_token=session_token,
|
||||||
user_uuid=stored_cred.user_uuid,
|
user_uuid=cred.user_uuid,
|
||||||
credential_uuid=stored_cred.uuid,
|
credential_uuid=cred.uuid,
|
||||||
reset_token=reset_token,
|
reset_token=reset_token,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
+27
-19
@@ -10,13 +10,11 @@ display name. If multiple users match, they are listed and the command
|
|||||||
aborts. A new one-time reset link is always created.
|
aborts. A new one-time reset link is always created.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from paskia import authsession as _authsession
|
from paskia import authsession as _authsession
|
||||||
from paskia import db as _db
|
from paskia import db
|
||||||
from paskia.util import hostutil, passphrase
|
from paskia.util import hostutil, passphrase
|
||||||
|
|
||||||
|
|
||||||
@@ -26,23 +24,30 @@ async def _resolve_targets(query: str | None):
|
|||||||
targets: list[tuple] = []
|
targets: list[tuple] = []
|
||||||
try:
|
try:
|
||||||
q_uuid = UUID(query)
|
q_uuid = UUID(query)
|
||||||
perm_orgs = _db.get_permission_organizations("auth:admin")
|
p = next(
|
||||||
for o in perm_orgs:
|
(p for p in db.data().permissions.values() if p.scope == "auth:admin"),
|
||||||
users = _db.get_organization_users(str(o.uuid))
|
None,
|
||||||
for u, role_name in users:
|
)
|
||||||
if u.uuid == q_uuid:
|
if p:
|
||||||
return [(u, role_name)]
|
for org_uuid in p.orgs:
|
||||||
|
users = db.get_organization_users(org_uuid)
|
||||||
|
for u, role_name in users:
|
||||||
|
if u.uuid == q_uuid:
|
||||||
|
return [(u, role_name)]
|
||||||
# UUID not found among admin orgs -> fall back to substring search (rare case)
|
# UUID not found among admin orgs -> fall back to substring search (rare case)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
pass
|
pass
|
||||||
# Substring search
|
# Substring search
|
||||||
needle = query.lower()
|
needle = query.lower()
|
||||||
perm_orgs = _db.get_permission_organizations("auth:admin")
|
p = next(
|
||||||
for o in perm_orgs:
|
(p for p in db.data().permissions.values() if p.scope == "auth:admin"), None
|
||||||
users = _db.get_organization_users(str(o.uuid))
|
)
|
||||||
for u, role_name in users:
|
if p:
|
||||||
if needle in (u.display_name or "").lower():
|
for org_uuid in p.orgs:
|
||||||
targets.append((u, role_name))
|
users = db.get_organization_users(org_uuid)
|
||||||
|
for u, role_name in users:
|
||||||
|
if needle in (u.display_name or "").lower():
|
||||||
|
targets.append((u, role_name))
|
||||||
# De-duplicate
|
# De-duplicate
|
||||||
seen = set()
|
seen = set()
|
||||||
deduped = []
|
deduped = []
|
||||||
@@ -52,10 +57,13 @@ async def _resolve_targets(query: str | None):
|
|||||||
deduped.append((u, role_name))
|
deduped.append((u, role_name))
|
||||||
return deduped
|
return deduped
|
||||||
# No query -> master admin
|
# No query -> master admin
|
||||||
perm_orgs = _db.get_permission_organizations("auth:admin")
|
p = next(
|
||||||
if not perm_orgs:
|
(p for p in db.data().permissions.values() if p.scope == "auth:admin"), None
|
||||||
|
)
|
||||||
|
if not p or not p.orgs:
|
||||||
return []
|
return []
|
||||||
users = _db.get_organization_users(str(perm_orgs[0].uuid))
|
first_org_uuid = next(iter(p.orgs))
|
||||||
|
users = db.get_organization_users(first_org_uuid)
|
||||||
admin_users = [pair for pair in users if pair[1] == "Administration"]
|
admin_users = [pair for pair in users if pair[1] == "Administration"]
|
||||||
return admin_users[:1]
|
return admin_users[:1]
|
||||||
|
|
||||||
@@ -63,7 +71,7 @@ async def _resolve_targets(query: str | None):
|
|||||||
async def _create_reset(user, role_name: str):
|
async def _create_reset(user, role_name: str):
|
||||||
token = passphrase.generate()
|
token = passphrase.generate()
|
||||||
expiry = _authsession.reset_expires()
|
expiry = _authsession.reset_expires()
|
||||||
_db.create_reset_token(
|
db.create_reset_token(
|
||||||
passphrase=token,
|
passphrase=token,
|
||||||
user_uuid=user.uuid,
|
user_uuid=user.uuid,
|
||||||
expiry=expiry,
|
expiry=expiry,
|
||||||
|
|||||||
@@ -0,0 +1,22 @@
|
|||||||
|
"""FastAPI response utilities for msgspec.Struct serialization."""
|
||||||
|
|
||||||
|
import msgspec
|
||||||
|
from fastapi import Response
|
||||||
|
|
||||||
|
|
||||||
|
class MsgspecResponse(Response):
|
||||||
|
"""Response that uses msgspec for JSON encoding.
|
||||||
|
|
||||||
|
Use this for returning msgspec.Struct, dict, or list with proper serialization.
|
||||||
|
"""
|
||||||
|
|
||||||
|
media_type = "application/json"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
content: msgspec.Struct | dict | list,
|
||||||
|
status_code: int = 200,
|
||||||
|
headers: dict | None = None,
|
||||||
|
):
|
||||||
|
body = msgspec.json.encode(content)
|
||||||
|
super().__init__(content=body, status_code=status_code, headers=headers)
|
||||||
@@ -19,8 +19,8 @@ AUTH_COOKIE = Cookie(None, alias=AUTH_COOKIE_NAME)
|
|||||||
def infodict(request: Request | WebSocket, type: str) -> dict:
|
def infodict(request: Request | WebSocket, type: str) -> dict:
|
||||||
"""Extract client information from request."""
|
"""Extract client information from request."""
|
||||||
return {
|
return {
|
||||||
"ip": request.client.host if request.client else None,
|
"ip": request.client.host if request.client else "",
|
||||||
"user_agent": request.headers.get("user-agent", "")[:500] or None,
|
"user_agent": request.headers.get("user-agent", "")[:500],
|
||||||
"session_type": type,
|
"session_type": type,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+24
-30
@@ -1,4 +1,4 @@
|
|||||||
from datetime import timezone
|
from datetime import UTC
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from fastapi import (
|
from fastapi import (
|
||||||
@@ -14,13 +14,12 @@ from paskia import db
|
|||||||
from paskia.authsession import (
|
from paskia.authsession import (
|
||||||
delete_credential,
|
delete_credential,
|
||||||
expires,
|
expires,
|
||||||
get_session,
|
|
||||||
)
|
)
|
||||||
from paskia.fastapi import authz, session
|
from paskia.fastapi import authz, session
|
||||||
from paskia.fastapi.session import AUTH_COOKIE
|
from paskia.fastapi.session import AUTH_COOKIE
|
||||||
from paskia.util import hostutil, passphrase
|
from paskia.util import hostutil, passphrase
|
||||||
|
|
||||||
app = FastAPI()
|
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
|
||||||
|
|
||||||
@app.exception_handler(authz.AuthException)
|
@app.exception_handler(authz.AuthException)
|
||||||
@@ -43,18 +42,18 @@ async def user_update_display_name(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Authentication Required", mode="login"
|
status_code=401, detail="Authentication Required", mode="login"
|
||||||
)
|
)
|
||||||
try:
|
host = request.headers.get("host")
|
||||||
s = await get_session(auth, host=request.headers.get("host"))
|
ctx = db.data().session_ctx(auth, host)
|
||||||
except ValueError as e:
|
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"
|
||||||
) from e
|
)
|
||||||
new_name = (payload.get("display_name") or "").strip()
|
new_name = (payload.get("display_name") or "").strip()
|
||||||
if not new_name:
|
if not new_name:
|
||||||
raise HTTPException(status_code=400, detail="display_name required")
|
raise HTTPException(status_code=400, detail="display_name required")
|
||||||
if len(new_name) > 64:
|
if len(new_name) > 64:
|
||||||
raise HTTPException(status_code=400, detail="display_name too long")
|
raise HTTPException(status_code=400, detail="display_name too long")
|
||||||
db.update_user_display_name(s.user_uuid, new_name)
|
db.update_user_display_name(ctx.user.uuid, new_name, ctx=ctx)
|
||||||
return {"status": "ok"}
|
return {"status": "ok"}
|
||||||
|
|
||||||
|
|
||||||
@@ -62,13 +61,13 @@ async def user_update_display_name(
|
|||||||
async def api_logout_all(request: Request, response: Response, auth=AUTH_COOKIE):
|
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"}
|
||||||
try:
|
host = request.headers.get("host")
|
||||||
s = await get_session(auth, host=request.headers.get("host"))
|
ctx = db.data().session_ctx(auth, host)
|
||||||
except ValueError:
|
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"
|
||||||
)
|
)
|
||||||
db.delete_sessions_for_user(s.user_uuid)
|
db.delete_sessions_for_user(ctx.user.uuid, ctx=ctx)
|
||||||
session.clear_session_cookie(response)
|
session.clear_session_cookie(response)
|
||||||
return {"message": "Logged out from all hosts"}
|
return {"message": "Logged out from all hosts"}
|
||||||
|
|
||||||
@@ -84,18 +83,18 @@ async def api_delete_session(
|
|||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Authentication Required", mode="login"
|
status_code=401, detail="Authentication Required", mode="login"
|
||||||
)
|
)
|
||||||
try:
|
host = request.headers.get("host")
|
||||||
current_session = await get_session(auth, host=request.headers.get("host"))
|
ctx = db.data().session_ctx(auth, host)
|
||||||
except ValueError as exc:
|
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"
|
||||||
) from exc
|
)
|
||||||
|
|
||||||
target_session = db.get_session(session_id)
|
target_session = db.data().sessions.get(session_id)
|
||||||
if not target_session or target_session.user_uuid != current_session.user_uuid:
|
if not target_session or target_session.user_uuid != ctx.user.uuid:
|
||||||
raise HTTPException(status_code=404, detail="Session not found")
|
raise HTTPException(status_code=404, detail="Session not found")
|
||||||
|
|
||||||
db.delete_session(session_id)
|
db.delete_session(session_id, ctx=ctx)
|
||||||
current_terminated = session_id == auth
|
current_terminated = session_id == auth
|
||||||
if current_terminated:
|
if current_terminated:
|
||||||
session.clear_session_cookie(response) # explicit because 200
|
session.clear_session_cookie(response) # explicit because 200
|
||||||
@@ -112,7 +111,7 @@ async def api_delete_credential(
|
|||||||
# Require recent authentication for sensitive operation
|
# Require recent authentication for sensitive operation
|
||||||
await authz.verify(auth, [], host=request.headers.get("host"), max_age="5m")
|
await authz.verify(auth, [], host=request.headers.get("host"), max_age="5m")
|
||||||
try:
|
try:
|
||||||
await delete_credential(uuid, auth, host=request.headers.get("host"))
|
delete_credential(uuid, auth, host=request.headers.get("host"))
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
raise authz.AuthException(
|
raise authz.AuthException(
|
||||||
status_code=401, detail="Session expired", mode="login"
|
status_code=401, detail="Session expired", mode="login"
|
||||||
@@ -127,28 +126,23 @@ async def api_create_link(
|
|||||||
auth=AUTH_COOKIE,
|
auth=AUTH_COOKIE,
|
||||||
):
|
):
|
||||||
# Require recent authentication for sensitive operation
|
# Require recent authentication for sensitive operation
|
||||||
await authz.verify(auth, [], host=request.headers.get("host"), max_age="5m")
|
ctx = await authz.verify(auth, [], host=request.headers.get("host"), max_age="5m")
|
||||||
try:
|
|
||||||
s = await get_session(auth, host=request.headers.get("host"))
|
|
||||||
except ValueError as e:
|
|
||||||
raise authz.AuthException(
|
|
||||||
status_code=401, detail="Session expired", mode="login"
|
|
||||||
) from e
|
|
||||||
token = passphrase.generate()
|
token = passphrase.generate()
|
||||||
expiry = expires()
|
expiry = expires()
|
||||||
db.create_reset_token(
|
db.create_reset_token(
|
||||||
user_uuid=s.user_uuid,
|
user_uuid=ctx.user.uuid,
|
||||||
passphrase=token,
|
passphrase=token,
|
||||||
expiry=expiry,
|
expiry=expiry,
|
||||||
token_type="device addition",
|
token_type="device addition",
|
||||||
|
ctx=ctx,
|
||||||
)
|
)
|
||||||
url = hostutil.reset_link_url(token)
|
url = hostutil.reset_link_url(token)
|
||||||
return {
|
return {
|
||||||
"message": "Registration link generated successfully",
|
"message": "Registration link generated successfully",
|
||||||
"url": url,
|
"url": url,
|
||||||
"expires": (
|
"expires": (
|
||||||
expiry.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
expiry.astimezone(UTC).isoformat().replace("+00:00", "Z")
|
||||||
if expiry.tzinfo
|
if expiry.tzinfo
|
||||||
else expiry.replace(tzinfo=timezone.utc).isoformat().replace("+00:00", "Z")
|
else expiry.replace(tzinfo=UTC).isoformat().replace("+00:00", "Z")
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
|||||||
+24
-59
@@ -1,40 +1,21 @@
|
|||||||
from uuid import UUID
|
|
||||||
|
|
||||||
from fastapi import FastAPI, WebSocket
|
from fastapi import FastAPI, WebSocket
|
||||||
|
|
||||||
from paskia import db
|
from paskia import db
|
||||||
from paskia.authsession import expires, get_reset, get_session
|
from paskia.authsession import expires, get_reset
|
||||||
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
|
||||||
|
from paskia.fastapi.wschat import authenticate_chat, register_chat
|
||||||
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
|
from paskia.fastapi.wsutil import validate_origin, websocket_error_handler
|
||||||
from paskia.globals import passkey
|
from paskia.globals import passkey
|
||||||
from paskia.util import hostutil, passphrase
|
from paskia.util import hostutil, passphrase
|
||||||
|
|
||||||
# Create a FastAPI subapp for WebSocket endpoints
|
# Create a FastAPI subapp for WebSocket endpoints
|
||||||
app = FastAPI()
|
app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None)
|
||||||
|
|
||||||
# Mount the remote auth WebSocket endpoints
|
# Mount the remote auth WebSocket endpoints
|
||||||
app.mount("/remote-auth", remote.app)
|
app.mount("/remote-auth", remote.app)
|
||||||
|
|
||||||
|
|
||||||
async def register_chat(
|
|
||||||
ws: WebSocket,
|
|
||||||
user_uuid: UUID,
|
|
||||||
user_name: str,
|
|
||||||
origin: str,
|
|
||||||
credential_ids: list[bytes] | None = None,
|
|
||||||
):
|
|
||||||
"""Generate registration options and send them to the client."""
|
|
||||||
options, challenge = passkey.instance.reg_generate_options(
|
|
||||||
user_id=user_uuid,
|
|
||||||
user_name=user_name,
|
|
||||||
credential_ids=credential_ids,
|
|
||||||
)
|
|
||||||
await ws.send_json({"optionsJSON": options})
|
|
||||||
response = await ws.receive_json()
|
|
||||||
return passkey.instance.reg_verify(response, challenge, user_uuid, origin=origin)
|
|
||||||
|
|
||||||
|
|
||||||
@app.websocket("/register")
|
@app.websocket("/register")
|
||||||
@websocket_error_handler
|
@websocket_error_handler
|
||||||
async def websocket_register_add(
|
async def websocket_register_add(
|
||||||
@@ -56,7 +37,7 @@ async def websocket_register_add(
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"The reset link for {passkey.instance.rp_name} is invalid or has expired"
|
f"The reset link for {passkey.instance.rp_name} is invalid or has expired"
|
||||||
)
|
)
|
||||||
s = await get_reset(reset)
|
s = get_reset(reset)
|
||||||
user_uuid = s.user_uuid
|
user_uuid = s.user_uuid
|
||||||
else:
|
else:
|
||||||
# Require recent authentication for adding a new passkey
|
# Require recent authentication for adding a new passkey
|
||||||
@@ -65,16 +46,16 @@ async def websocket_register_add(
|
|||||||
s = ctx.session
|
s = ctx.session
|
||||||
|
|
||||||
# Get user information and determine effective user_name for this registration
|
# Get user information and determine effective user_name for this registration
|
||||||
user = db.get_user_by_uuid(user_uuid)
|
user = db.data().users.get(user_uuid)
|
||||||
user_name = user.display_name
|
user_name = user.display_name
|
||||||
if name is not None:
|
if name is not None:
|
||||||
stripped = name.strip()
|
stripped = name.strip()
|
||||||
if stripped:
|
if stripped:
|
||||||
user_name = stripped
|
user_name = stripped
|
||||||
challenge_ids = db.get_credentials_by_user_uuid(user_uuid)
|
credential_ids = db.get_user_credential_ids(user_uuid) or None
|
||||||
|
|
||||||
# WebAuthn registration
|
# WebAuthn registration
|
||||||
credential = await register_chat(ws, user_uuid, user_name, origin, challenge_ids)
|
credential = await register_chat(ws, user_uuid, user_name, origin, credential_ids)
|
||||||
|
|
||||||
# Create a new session and store everything in database
|
# Create a new session and store everything in database
|
||||||
metadata = infodict(ws, "authenticated")
|
metadata = infodict(ws, "authenticated")
|
||||||
@@ -84,16 +65,16 @@ async def websocket_register_add(
|
|||||||
reset_key=(s.key if reset is not None else None),
|
reset_key=(s.key if reset is not None else None),
|
||||||
display_name=user_name,
|
display_name=user_name,
|
||||||
host=host,
|
host=host,
|
||||||
ip=metadata.get("ip"),
|
ip=metadata["ip"],
|
||||||
user_agent=metadata.get("user_agent"),
|
user_agent=metadata["user_agent"],
|
||||||
)
|
)
|
||||||
auth = token
|
auth = token
|
||||||
|
|
||||||
assert isinstance(auth, str) and len(auth) == 16
|
assert isinstance(auth, str) and len(auth) == 16
|
||||||
await ws.send_json(
|
await ws.send_json(
|
||||||
{
|
{
|
||||||
"user_uuid": str(user.uuid),
|
"user": str(user.uuid),
|
||||||
"credential_uuid": str(credential.uuid),
|
"credential": str(credential.uuid),
|
||||||
"session_token": auth,
|
"session_token": auth,
|
||||||
"message": "New credential added successfully",
|
"message": "New credential added successfully",
|
||||||
}
|
}
|
||||||
@@ -110,36 +91,19 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
|
|||||||
session_user_uuid = None
|
session_user_uuid = None
|
||||||
credential_ids = None
|
credential_ids = None
|
||||||
if auth:
|
if auth:
|
||||||
try:
|
ctx = db.data().session_ctx(auth, host)
|
||||||
session = await get_session(auth, host=host)
|
if ctx:
|
||||||
session_user_uuid = session.user_uuid
|
session_user_uuid = ctx.user.uuid
|
||||||
credential_ids = db.get_credentials_by_user_uuid(session_user_uuid)
|
credential_ids = db.get_user_credential_ids(session_user_uuid) or None
|
||||||
except ValueError:
|
|
||||||
pass # Invalid/expired session - allow normal authentication
|
|
||||||
|
|
||||||
options, challenge = passkey.instance.auth_generate_options(
|
cred, new_sign_count = await authenticate_chat(ws, origin, credential_ids)
|
||||||
credential_ids=credential_ids
|
|
||||||
)
|
|
||||||
await ws.send_json({"optionsJSON": options})
|
|
||||||
# Wait for the client to use his authenticator to authenticate
|
|
||||||
credential = passkey.instance.auth_parse(await ws.receive_json())
|
|
||||||
# Fetch from the database by credential ID
|
|
||||||
try:
|
|
||||||
stored_cred = db.get_credential_by_id(credential.raw_id)
|
|
||||||
except ValueError:
|
|
||||||
raise ValueError(
|
|
||||||
f"This passkey is no longer registered with {passkey.instance.rp_name}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# If reauth mode, verify the credential belongs to the session's user
|
# If reauth mode, verify the credential belongs to the session's user
|
||||||
if session_user_uuid and stored_cred.user_uuid != session_user_uuid:
|
if session_user_uuid and cred.user_uuid != session_user_uuid:
|
||||||
raise ValueError("This passkey belongs to a different account")
|
raise ValueError("This passkey belongs to a different account")
|
||||||
|
|
||||||
# Verify the credential matches the stored data
|
|
||||||
passkey.instance.auth_verify(credential, challenge, stored_cred, origin)
|
|
||||||
|
|
||||||
# Create session and update user/credential in a single transaction
|
# Create session and update user/credential in a single transaction
|
||||||
assert stored_cred.uuid is not None
|
assert cred.uuid is not None
|
||||||
metadata = infodict(ws, "auth")
|
metadata = infodict(ws, "auth")
|
||||||
normalized_host = hostutil.normalize_host(host)
|
normalized_host = hostutil.normalize_host(host)
|
||||||
if not normalized_host:
|
if not normalized_host:
|
||||||
@@ -150,17 +114,18 @@ async def websocket_authenticate(ws: WebSocket, auth=AUTH_COOKIE):
|
|||||||
raise ValueError(f"Host must be the same as or a subdomain of {rp_id}")
|
raise ValueError(f"Host must be the same as or a subdomain of {rp_id}")
|
||||||
|
|
||||||
token = db.login(
|
token = db.login(
|
||||||
user_uuid=stored_cred.user_uuid,
|
user_uuid=cred.user_uuid,
|
||||||
credential=stored_cred,
|
credential_uuid=cred.uuid,
|
||||||
|
sign_count=new_sign_count,
|
||||||
host=normalized_host,
|
host=normalized_host,
|
||||||
ip=metadata.get("ip") or "",
|
ip=metadata["ip"],
|
||||||
user_agent=metadata.get("user_agent") or "",
|
user_agent=metadata["user_agent"],
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
|
|
||||||
await ws.send_json(
|
await ws.send_json(
|
||||||
{
|
{
|
||||||
"user_uuid": str(stored_cred.user_uuid),
|
"user": str(cred.user_uuid),
|
||||||
"session_token": token,
|
"session_token": token,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -0,0 +1,62 @@
|
|||||||
|
"""
|
||||||
|
WebSocket chat functions for WebAuthn registration and authentication flows.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
from fastapi import WebSocket
|
||||||
|
|
||||||
|
from paskia import db
|
||||||
|
from paskia.db import Credential
|
||||||
|
from paskia.globals import passkey
|
||||||
|
|
||||||
|
|
||||||
|
async def register_chat(
|
||||||
|
ws: WebSocket,
|
||||||
|
user_uuid: UUID,
|
||||||
|
user_name: str,
|
||||||
|
origin: str,
|
||||||
|
credential_ids: list[bytes] | None = None,
|
||||||
|
):
|
||||||
|
"""Run WebAuthn registration flow and return the verified credential."""
|
||||||
|
options, challenge = passkey.instance.reg_generate_options(
|
||||||
|
user_id=user_uuid,
|
||||||
|
user_name=user_name,
|
||||||
|
credential_ids=credential_ids,
|
||||||
|
)
|
||||||
|
await ws.send_json({"optionsJSON": options})
|
||||||
|
response = await ws.receive_json()
|
||||||
|
return passkey.instance.reg_verify(response, challenge, user_uuid, origin=origin)
|
||||||
|
|
||||||
|
|
||||||
|
async def authenticate_chat(
|
||||||
|
ws: WebSocket,
|
||||||
|
origin: str,
|
||||||
|
credential_ids: list[bytes] | None = None,
|
||||||
|
) -> tuple[Credential, int]:
|
||||||
|
"""Run WebAuthn authentication flow and return the credential and new sign count.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
tuple of (credential, new_sign_count) where new_sign_count comes from WebAuthn verification
|
||||||
|
"""
|
||||||
|
options, challenge = passkey.instance.auth_generate_options(
|
||||||
|
credential_ids=credential_ids
|
||||||
|
)
|
||||||
|
await ws.send_json({"optionsJSON": options})
|
||||||
|
authcred = passkey.instance.auth_parse(await ws.receive_json())
|
||||||
|
|
||||||
|
cred = next(
|
||||||
|
(
|
||||||
|
c
|
||||||
|
for c in db.data().credentials.values()
|
||||||
|
if c.credential_id == authcred.raw_id
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not cred:
|
||||||
|
raise ValueError(
|
||||||
|
f"This passkey is no longer registered with {passkey.instance.rp_name}"
|
||||||
|
)
|
||||||
|
|
||||||
|
verification = passkey.instance.auth_verify(authcred, challenge, cred, origin)
|
||||||
|
return cred, verification.new_sign_count
|
||||||
@@ -3,6 +3,7 @@ Shared WebSocket utilities for FastAPI endpoints.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import time
|
||||||
from functools import wraps
|
from functools import wraps
|
||||||
|
|
||||||
import base64url
|
import base64url
|
||||||
@@ -10,6 +11,7 @@ from fastapi import WebSocket, WebSocketDisconnect
|
|||||||
from webauthn.helpers.exceptions import InvalidAuthenticationResponse
|
from webauthn.helpers.exceptions import InvalidAuthenticationResponse
|
||||||
|
|
||||||
from paskia.fastapi import authz
|
from paskia.fastapi import authz
|
||||||
|
from paskia.fastapi.logging import log_ws_close, log_ws_open
|
||||||
from paskia.globals import passkey
|
from paskia.globals import passkey
|
||||||
from paskia.util import pow
|
from paskia.util import pow
|
||||||
|
|
||||||
@@ -19,11 +21,19 @@ def websocket_error_handler(func):
|
|||||||
|
|
||||||
@wraps(func)
|
@wraps(func)
|
||||||
async def wrapper(ws: WebSocket, *args, **kwargs):
|
async def wrapper(ws: WebSocket, *args, **kwargs):
|
||||||
|
client = ws.client.host if ws.client else "-"
|
||||||
|
host = ws.headers.get("host", "-")
|
||||||
|
path = ws.url.path
|
||||||
|
|
||||||
|
start = time.perf_counter()
|
||||||
|
ws_id = log_ws_open(client, host, path)
|
||||||
|
close_code = None
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await ws.accept()
|
await ws.accept()
|
||||||
return await func(ws, *args, **kwargs)
|
return await func(ws, *args, **kwargs)
|
||||||
except WebSocketDisconnect:
|
except WebSocketDisconnect as e:
|
||||||
pass
|
close_code = e.code
|
||||||
except authz.AuthException as e:
|
except authz.AuthException as e:
|
||||||
await ws.send_json(
|
await ws.send_json(
|
||||||
{
|
{
|
||||||
@@ -36,6 +46,9 @@ def websocket_error_handler(func):
|
|||||||
except Exception:
|
except Exception:
|
||||||
logging.exception("Internal Server Error")
|
logging.exception("Internal Server Error")
|
||||||
await ws.send_json({"status": 500, "detail": "Internal Server Error"})
|
await ws.send_json({"status": 500, "detail": "Internal Server Error"})
|
||||||
|
finally:
|
||||||
|
duration_ms = (time.perf_counter() - start) * 1000
|
||||||
|
log_ws_close(client, ws_id, close_code, duration_ms)
|
||||||
|
|
||||||
return wrapper
|
return wrapper
|
||||||
|
|
||||||
|
|||||||
+2
-2
@@ -1,5 +1,7 @@
|
|||||||
from typing import Generic, TypeVar
|
from typing import Generic, TypeVar
|
||||||
|
|
||||||
|
from paskia import db, remoteauth
|
||||||
|
from paskia.bootstrap import bootstrap_if_needed
|
||||||
from paskia.sansio import Passkey
|
from paskia.sansio import Passkey
|
||||||
|
|
||||||
T = TypeVar("T")
|
T = TypeVar("T")
|
||||||
@@ -42,7 +44,6 @@ async def init(
|
|||||||
Set PASKIA_DB environment variable to specify the JSONL database file path.
|
Set PASKIA_DB environment variable to specify the JSONL database file path.
|
||||||
Default: paskia.jsonl
|
Default: paskia.jsonl
|
||||||
"""
|
"""
|
||||||
from . import db, remoteauth
|
|
||||||
|
|
||||||
# Initialize passkey instance with provided parameters
|
# Initialize passkey instance with provided parameters
|
||||||
passkey.instance = Passkey(
|
passkey.instance = Passkey(
|
||||||
@@ -59,7 +60,6 @@ async def init(
|
|||||||
|
|
||||||
if bootstrap:
|
if bootstrap:
|
||||||
# Bootstrap system if needed
|
# Bootstrap system if needed
|
||||||
from .bootstrap import bootstrap_if_needed
|
|
||||||
|
|
||||||
await bootstrap_if_needed()
|
await bootstrap_if_needed()
|
||||||
|
|
||||||
|
|||||||
+67
-60
@@ -11,13 +11,28 @@ Or via the CLI entry point (if installed):
|
|||||||
paskia-migrate --sql sqlite+aiosqlite:///paskia.sqlite --json paskia.jsonl
|
paskia-migrate --sql sqlite+aiosqlite:///paskia.sqlite --json paskia.jsonl
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
import asyncio
|
import asyncio
|
||||||
from datetime import datetime, timezone
|
import re
|
||||||
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import base64url
|
import base64url
|
||||||
|
import uuid7
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
|
from paskia.db.jsonl import JsonlStore
|
||||||
|
from paskia.db.structs import (
|
||||||
|
DB,
|
||||||
|
Credential,
|
||||||
|
Org,
|
||||||
|
Permission,
|
||||||
|
ResetToken,
|
||||||
|
Role,
|
||||||
|
Session,
|
||||||
|
User,
|
||||||
|
)
|
||||||
|
|
||||||
from .sql import (
|
from .sql import (
|
||||||
DB as SQLDB,
|
DB as SQLDB,
|
||||||
@@ -47,30 +62,14 @@ async def migrate_from_sql(
|
|||||||
sql_db_path: SQLAlchemy connection string for the source SQL database
|
sql_db_path: SQLAlchemy connection string for the source SQL database
|
||||||
json_db_path: Path for the destination JSONL file
|
json_db_path: Path for the destination JSONL file
|
||||||
"""
|
"""
|
||||||
# Import here to avoid circular imports and to not require JSON db at import time
|
|
||||||
import re
|
|
||||||
|
|
||||||
import uuid7
|
|
||||||
from sqlalchemy import select
|
|
||||||
|
|
||||||
from paskia.db.operations import DB as JSONDB
|
|
||||||
from paskia.db.structs import (
|
|
||||||
_CredentialData,
|
|
||||||
_OrgData,
|
|
||||||
_PermissionData,
|
|
||||||
_ResetTokenData,
|
|
||||||
_RoleData,
|
|
||||||
_SessionData,
|
|
||||||
_UserData,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Initialize source SQL database
|
# Initialize source SQL database
|
||||||
sql_db = SQLDB(sql_db_path)
|
sql_db = SQLDB(sql_db_path)
|
||||||
await sql_db.init_db()
|
await sql_db.init_db()
|
||||||
|
|
||||||
# Initialize destination JSON database (fresh, don't load existing)
|
# Initialize destination JSON database (fresh, don't load existing)
|
||||||
json_db = JSONDB(json_db_path)
|
db = DB()
|
||||||
# Don't call json_db.load() - we want a fresh database, not to load existing
|
store = JsonlStore(db, json_db_path)
|
||||||
|
db._store = store
|
||||||
|
|
||||||
print(f"Migrating from {sql_db_path} to {json_db_path}...")
|
print(f"Migrating from {sql_db_path} to {json_db_path}...")
|
||||||
|
|
||||||
@@ -90,11 +89,13 @@ async def migrate_from_sql(
|
|||||||
# Migrate permissions with UUID keys and scope field
|
# Migrate permissions with UUID keys and scope field
|
||||||
# Always create exactly one common auth:org:admin permission for all org admin needs
|
# Always create exactly one common auth:org:admin permission for all org admin needs
|
||||||
org_admin_perm_uuid: UUID = uuid7.create()
|
org_admin_perm_uuid: UUID = uuid7.create()
|
||||||
json_db._data.permissions[org_admin_perm_uuid] = _PermissionData(
|
org_admin_perm = Permission(
|
||||||
scope="auth:org:admin",
|
scope="auth:org:admin",
|
||||||
display_name="Org Admin",
|
display_name="Org Admin",
|
||||||
orgs={},
|
orgs={},
|
||||||
)
|
)
|
||||||
|
org_admin_perm.uuid = org_admin_perm_uuid
|
||||||
|
db.permissions[org_admin_perm_uuid] = org_admin_perm
|
||||||
|
|
||||||
# Mapping from old permission ID to new permission UUID
|
# Mapping from old permission ID to new permission UUID
|
||||||
perm_id_to_uuid: dict[str, UUID] = {}
|
perm_id_to_uuid: dict[str, UUID] = {}
|
||||||
@@ -113,11 +114,13 @@ async def migrate_from_sql(
|
|||||||
|
|
||||||
# Regular permission - create with UUID key
|
# Regular permission - create with UUID key
|
||||||
perm_uuid: UUID = uuid7.create()
|
perm_uuid: UUID = uuid7.create()
|
||||||
json_db._data.permissions[perm_uuid] = _PermissionData(
|
new_perm = Permission(
|
||||||
scope=perm.id, # Old ID becomes the scope
|
scope=perm.id, # Old ID becomes the scope
|
||||||
display_name=perm.display_name,
|
display_name=perm.display_name,
|
||||||
orgs={},
|
orgs={},
|
||||||
)
|
)
|
||||||
|
new_perm.uuid = perm_uuid
|
||||||
|
db.permissions[perm_uuid] = new_perm
|
||||||
perm_id_to_uuid[perm.id] = perm_uuid
|
perm_id_to_uuid[perm.id] = perm_uuid
|
||||||
print(
|
print(
|
||||||
f" Migrated {len(permissions)} permissions (with {len(org_admin_uuids)} org-specific admins consolidated to auth:org:admin)"
|
f" Migrated {len(permissions)} permissions (with {len(org_admin_uuids)} org-specific admins consolidated to auth:org:admin)"
|
||||||
@@ -127,16 +130,16 @@ async def migrate_from_sql(
|
|||||||
orgs = await sql_db.list_organizations()
|
orgs = await sql_db.list_organizations()
|
||||||
for org in orgs:
|
for org in orgs:
|
||||||
org_key: UUID = org.uuid
|
org_key: UUID = org.uuid
|
||||||
json_db._data.orgs[org_key] = _OrgData(
|
new_org = Org(display_name=org.display_name)
|
||||||
display_name=org.display_name,
|
new_org.uuid = org_key
|
||||||
)
|
db.orgs[org_key] = new_org
|
||||||
# Update permissions to allow this org to grant them (by UUID)
|
# Update permissions to allow this org to grant them (by UUID)
|
||||||
for old_perm_id in org.permissions:
|
for old_perm_id in org.permissions:
|
||||||
perm_uuid = perm_id_to_uuid.get(old_perm_id)
|
perm_uuid = perm_id_to_uuid.get(old_perm_id)
|
||||||
if perm_uuid and perm_uuid in json_db._data.permissions:
|
if perm_uuid and perm_uuid in db.permissions:
|
||||||
json_db._data.permissions[perm_uuid].orgs[org_key] = True
|
db.permissions[perm_uuid].orgs[org_key] = True
|
||||||
# Ensure every org can grant auth:org:admin
|
# Ensure every org can grant auth:org:admin
|
||||||
json_db._data.permissions[org_admin_perm_uuid].orgs[org_key] = True
|
db.permissions[org_admin_perm_uuid].orgs[org_key] = True
|
||||||
print(f" Migrated {len(orgs)} organizations")
|
print(f" Migrated {len(orgs)} organizations")
|
||||||
|
|
||||||
# Migrate roles - convert old permission IDs to UUIDs
|
# Migrate roles - convert old permission IDs to UUIDs
|
||||||
@@ -150,11 +153,13 @@ async def migrate_from_sql(
|
|||||||
perm_uuid = perm_id_to_uuid.get(old_perm_id)
|
perm_uuid = perm_id_to_uuid.get(old_perm_id)
|
||||||
if perm_uuid:
|
if perm_uuid:
|
||||||
new_permissions[perm_uuid] = True
|
new_permissions[perm_uuid] = True
|
||||||
json_db._data.roles[role_key] = _RoleData(
|
new_role = Role(
|
||||||
org=role.org_uuid,
|
org_uuid=role.org_uuid,
|
||||||
display_name=role.display_name,
|
display_name=role.display_name,
|
||||||
permissions=new_permissions,
|
permissions=new_permissions,
|
||||||
)
|
)
|
||||||
|
new_role.uuid = role_key
|
||||||
|
db.roles[role_key] = new_role
|
||||||
role_count += 1
|
role_count += 1
|
||||||
print(f" Migrated {role_count} roles")
|
print(f" Migrated {role_count} roles")
|
||||||
|
|
||||||
@@ -163,15 +168,17 @@ async def migrate_from_sql(
|
|||||||
result = await session.execute(select(UserModel))
|
result = await session.execute(select(UserModel))
|
||||||
user_models = result.scalars().all()
|
user_models = result.scalars().all()
|
||||||
for um in user_models:
|
for um in user_models:
|
||||||
user = um.as_dataclass()
|
legacy_user = um.as_dataclass()
|
||||||
user_key: UUID = user.uuid
|
user_key: UUID = legacy_user.uuid
|
||||||
json_db._data.users[user_key] = _UserData(
|
new_user = User(
|
||||||
display_name=user.display_name,
|
display_name=legacy_user.display_name,
|
||||||
role=user.role_uuid,
|
role_uuid=legacy_user.role_uuid,
|
||||||
created_at=user.created_at or datetime.now(timezone.utc),
|
created_at=legacy_user.created_at or datetime.now(UTC),
|
||||||
last_seen=user.last_seen,
|
last_seen=legacy_user.last_seen,
|
||||||
visits=user.visits,
|
visits=legacy_user.visits,
|
||||||
)
|
)
|
||||||
|
new_user.uuid = user_key
|
||||||
|
db.users[user_key] = new_user
|
||||||
print(f" Migrated {len(user_models)} users")
|
print(f" Migrated {len(user_models)} users")
|
||||||
|
|
||||||
# Migrate credentials
|
# Migrate credentials
|
||||||
@@ -179,18 +186,20 @@ async def migrate_from_sql(
|
|||||||
result = await session.execute(select(CredentialModel))
|
result = await session.execute(select(CredentialModel))
|
||||||
cred_models = result.scalars().all()
|
cred_models = result.scalars().all()
|
||||||
for cm in cred_models:
|
for cm in cred_models:
|
||||||
cred = cm.as_dataclass()
|
legacy_cred = cm.as_dataclass()
|
||||||
cred_key: UUID = cred.uuid
|
cred_key: UUID = legacy_cred.uuid
|
||||||
json_db._data.credentials[cred_key] = _CredentialData(
|
new_cred = Credential(
|
||||||
credential_id=cred.credential_id,
|
credential_id=legacy_cred.credential_id,
|
||||||
user=cred.user_uuid,
|
user_uuid=legacy_cred.user_uuid,
|
||||||
aaguid=cred.aaguid,
|
aaguid=legacy_cred.aaguid,
|
||||||
public_key=cred.public_key,
|
public_key=legacy_cred.public_key,
|
||||||
sign_count=cred.sign_count,
|
sign_count=legacy_cred.sign_count,
|
||||||
created_at=cred.created_at,
|
created_at=legacy_cred.created_at,
|
||||||
last_used=cred.last_used,
|
last_used=legacy_cred.last_used,
|
||||||
last_verified=cred.last_verified,
|
last_verified=legacy_cred.last_verified,
|
||||||
)
|
)
|
||||||
|
new_cred.uuid = cred_key
|
||||||
|
db.credentials[cred_key] = new_cred
|
||||||
print(f" Migrated {len(cred_models)} credentials")
|
print(f" Migrated {len(cred_models)} credentials")
|
||||||
|
|
||||||
# Migrate sessions
|
# Migrate sessions
|
||||||
@@ -207,9 +216,9 @@ async def migrate_from_sql(
|
|||||||
else:
|
else:
|
||||||
# Already in new format or unknown - try to use as-is
|
# Already in new format or unknown - try to use as-is
|
||||||
session_key = base64url.enc(old_key[:12])
|
session_key = base64url.enc(old_key[:12])
|
||||||
json_db._data.sessions[session_key] = _SessionData(
|
db.sessions[session_key] = Session(
|
||||||
user=sess.user_uuid,
|
user_uuid=sess.user_uuid,
|
||||||
credential=sess.credential_uuid,
|
credential_uuid=sess.credential_uuid,
|
||||||
host=sess.host,
|
host=sess.host,
|
||||||
ip=sess.ip,
|
ip=sess.ip,
|
||||||
user_agent=sess.user_agent,
|
user_agent=sess.user_agent,
|
||||||
@@ -231,26 +240,24 @@ async def migrate_from_sql(
|
|||||||
else:
|
else:
|
||||||
# Already in new format or unknown - truncate to 9 bytes
|
# Already in new format or unknown - truncate to 9 bytes
|
||||||
token_key = old_key[:9]
|
token_key = old_key[:9]
|
||||||
json_db._data.reset_tokens[token_key] = _ResetTokenData(
|
db.reset_tokens[token_key] = ResetToken(
|
||||||
user=token.user_uuid,
|
user_uuid=token.user_uuid,
|
||||||
expiry=token.expiry,
|
expiry=token.expiry,
|
||||||
token_type=token.token_type,
|
token_type=token.token_type,
|
||||||
)
|
)
|
||||||
print(f" Migrated {len(token_models)} reset tokens")
|
print(f" Migrated {len(token_models)} reset tokens")
|
||||||
|
|
||||||
# Queue and flush all changes with actor "migrate"
|
# Queue and flush all changes using the transaction mechanism
|
||||||
json_db._current_actor = "migrate"
|
with db.transaction("migrate:sql"):
|
||||||
json_db._queue_change()
|
pass # All data already added to _data, transaction commits on exit
|
||||||
from paskia.db.jsonl import flush_changes
|
|
||||||
|
|
||||||
await flush_changes(json_db.db_path, json_db._pending_changes)
|
await store.flush()
|
||||||
|
|
||||||
print("Migration complete!")
|
print("Migration complete!")
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
"""CLI entry point for migration."""
|
"""CLI entry point for migration."""
|
||||||
import argparse
|
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Migrate Paskia database from SQL to JSON"
|
description="Migrate Paskia database from SQL to JSON"
|
||||||
|
|||||||
+100
-43
@@ -9,7 +9,7 @@ DO NOT use this module for new code. Use paskia.db instead.
|
|||||||
|
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from sqlalchemy import (
|
from sqlalchemy import (
|
||||||
@@ -25,13 +25,81 @@ from sqlalchemy.dialects.sqlite import BLOB
|
|||||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||||
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column
|
||||||
|
|
||||||
from paskia.db import (
|
|
||||||
Credential,
|
# Legacy User class for SQL schema (uses 'role_uuid' not 'role')
|
||||||
Org,
|
@dataclass
|
||||||
ResetToken,
|
class _LegacyUser:
|
||||||
Role,
|
"""User as stored in the old SQL schema with role_uuid field."""
|
||||||
User,
|
|
||||||
)
|
uuid: UUID
|
||||||
|
display_name: str
|
||||||
|
role_uuid: UUID
|
||||||
|
created_at: datetime | None = None
|
||||||
|
last_seen: datetime | None = None
|
||||||
|
visits: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy Credential class for SQL schema (uses 'user_uuid' not 'user')
|
||||||
|
@dataclass
|
||||||
|
class _LegacyCredential:
|
||||||
|
"""Credential as stored in the old SQL schema with user_uuid field."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
credential_id: bytes
|
||||||
|
user_uuid: UUID
|
||||||
|
aaguid: UUID
|
||||||
|
public_key: bytes
|
||||||
|
sign_count: int
|
||||||
|
created_at: datetime
|
||||||
|
last_used: datetime | None = None
|
||||||
|
last_verified: datetime | None = None
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy Role class for SQL schema (uses 'org_uuid' not 'org')
|
||||||
|
@dataclass
|
||||||
|
class _LegacyRole:
|
||||||
|
"""Role as stored in the old SQL schema with org_uuid field."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
org_uuid: UUID
|
||||||
|
display_name: str
|
||||||
|
permissions: list[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy Org class for SQL schema (has mutable permissions/roles lists)
|
||||||
|
@dataclass
|
||||||
|
class _LegacyOrg:
|
||||||
|
"""Org as stored in the old SQL schema with mutable permissions/roles."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
display_name: str
|
||||||
|
permissions: list[str] | None = None
|
||||||
|
roles: list[_LegacyRole] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy Session class for SQL schema (uses 'key' as field, 'user_uuid', 'credential_uuid')
|
||||||
|
@dataclass
|
||||||
|
class _LegacySession:
|
||||||
|
"""Session as stored in the old SQL schema."""
|
||||||
|
|
||||||
|
key: bytes
|
||||||
|
user_uuid: UUID
|
||||||
|
credential_uuid: UUID
|
||||||
|
host: str
|
||||||
|
ip: str
|
||||||
|
user_agent: str
|
||||||
|
renewed: datetime
|
||||||
|
|
||||||
|
|
||||||
|
# Legacy ResetToken class for SQL schema (uses 'key' as field, 'user_uuid')
|
||||||
|
@dataclass
|
||||||
|
class _LegacyResetToken:
|
||||||
|
"""ResetToken as stored in the old SQL schema."""
|
||||||
|
|
||||||
|
key: bytes
|
||||||
|
user_uuid: UUID
|
||||||
|
token_type: str
|
||||||
|
expiry: datetime
|
||||||
|
|
||||||
|
|
||||||
# Local Permission class for SQL schema (uses 'id' not 'uuid' + 'scope')
|
# Local Permission class for SQL schema (uses 'id' not 'uuid' + 'scope')
|
||||||
@@ -46,26 +114,12 @@ class SqlPermission:
|
|||||||
DB_PATH_DEFAULT = "sqlite+aiosqlite:///paskia.sqlite"
|
DB_PATH_DEFAULT = "sqlite+aiosqlite:///paskia.sqlite"
|
||||||
|
|
||||||
|
|
||||||
# Local Session class for SQL schema (uses 'renewed' not 'expiry')
|
|
||||||
@dataclass
|
|
||||||
class _SqlSession:
|
|
||||||
"""Session as stored in the old SQL schema with renewed timestamp."""
|
|
||||||
|
|
||||||
key: bytes
|
|
||||||
user_uuid: UUID
|
|
||||||
credential_uuid: UUID
|
|
||||||
host: str
|
|
||||||
ip: str
|
|
||||||
user_agent: str
|
|
||||||
renewed: datetime
|
|
||||||
|
|
||||||
|
|
||||||
def _normalize_dt(value: datetime | None) -> datetime | None:
|
def _normalize_dt(value: datetime | None) -> datetime | None:
|
||||||
if value is None:
|
if value is None:
|
||||||
return None
|
return None
|
||||||
if value.tzinfo is None:
|
if value.tzinfo is None:
|
||||||
return value.replace(tzinfo=timezone.utc)
|
return value.replace(tzinfo=UTC)
|
||||||
return value.astimezone(timezone.utc)
|
return value.astimezone(UTC)
|
||||||
|
|
||||||
|
|
||||||
class Base(DeclarativeBase):
|
class Base(DeclarativeBase):
|
||||||
@@ -80,10 +134,13 @@ class OrgModel(Base):
|
|||||||
|
|
||||||
def as_dataclass(self):
|
def as_dataclass(self):
|
||||||
# Base Org without permissions/roles (filled by data accessors)
|
# Base Org without permissions/roles (filled by data accessors)
|
||||||
return Org(UUID(bytes=self.uuid), self.display_name)
|
return _LegacyOrg(
|
||||||
|
uuid=UUID(bytes=self.uuid),
|
||||||
|
display_name=self.display_name,
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dataclass(org: Org):
|
def from_dataclass(org: _LegacyOrg):
|
||||||
return OrgModel(uuid=org.uuid.bytes, display_name=org.display_name)
|
return OrgModel(uuid=org.uuid.bytes, display_name=org.display_name)
|
||||||
|
|
||||||
|
|
||||||
@@ -98,14 +155,14 @@ class RoleModel(Base):
|
|||||||
|
|
||||||
def as_dataclass(self):
|
def as_dataclass(self):
|
||||||
# Base Role without permissions (filled by data accessors)
|
# Base Role without permissions (filled by data accessors)
|
||||||
return Role(
|
return _LegacyRole(
|
||||||
uuid=UUID(bytes=self.uuid),
|
uuid=UUID(bytes=self.uuid),
|
||||||
org_uuid=UUID(bytes=self.org_uuid),
|
org_uuid=UUID(bytes=self.org_uuid),
|
||||||
display_name=self.display_name,
|
display_name=self.display_name,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dataclass(role: Role):
|
def from_dataclass(role: _LegacyRole):
|
||||||
return RoleModel(
|
return RoleModel(
|
||||||
uuid=role.uuid.bytes,
|
uuid=role.uuid.bytes,
|
||||||
org_uuid=role.org_uuid.bytes,
|
org_uuid=role.org_uuid.bytes,
|
||||||
@@ -122,15 +179,15 @@ class UserModel(Base):
|
|||||||
LargeBinary(16), ForeignKey("roles.uuid", ondelete="CASCADE"), nullable=False
|
LargeBinary(16), ForeignKey("roles.uuid", ondelete="CASCADE"), nullable=False
|
||||||
)
|
)
|
||||||
created_at: Mapped[datetime] = mapped_column(
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc)
|
DateTime(timezone=True), default=lambda: datetime.now(UTC)
|
||||||
)
|
)
|
||||||
last_seen: Mapped[datetime | None] = mapped_column(
|
last_seen: Mapped[datetime | None] = mapped_column(
|
||||||
DateTime(timezone=True), nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
)
|
)
|
||||||
visits: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
visits: Mapped[int] = mapped_column(Integer, nullable=False, default=0)
|
||||||
|
|
||||||
def as_dataclass(self) -> User:
|
def as_dataclass(self) -> "_LegacyUser":
|
||||||
return User(
|
return _LegacyUser(
|
||||||
uuid=UUID(bytes=self.uuid),
|
uuid=UUID(bytes=self.uuid),
|
||||||
display_name=self.display_name,
|
display_name=self.display_name,
|
||||||
role_uuid=UUID(bytes=self.role_uuid),
|
role_uuid=UUID(bytes=self.role_uuid),
|
||||||
@@ -140,12 +197,12 @@ class UserModel(Base):
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dataclass(user: User):
|
def from_dataclass(user: "_LegacyUser"):
|
||||||
return UserModel(
|
return UserModel(
|
||||||
uuid=user.uuid.bytes,
|
uuid=user.uuid.bytes,
|
||||||
display_name=user.display_name,
|
display_name=user.display_name,
|
||||||
role_uuid=user.role_uuid.bytes,
|
role_uuid=user.role_uuid.bytes,
|
||||||
created_at=user.created_at or datetime.now(timezone.utc),
|
created_at=user.created_at or datetime.now(UTC),
|
||||||
last_seen=user.last_seen,
|
last_seen=user.last_seen,
|
||||||
visits=user.visits,
|
visits=user.visits,
|
||||||
)
|
)
|
||||||
@@ -165,7 +222,7 @@ class CredentialModel(Base):
|
|||||||
public_key: Mapped[bytes] = mapped_column(BLOB, nullable=False)
|
public_key: Mapped[bytes] = mapped_column(BLOB, nullable=False)
|
||||||
sign_count: Mapped[int] = mapped_column(Integer, nullable=False)
|
sign_count: Mapped[int] = mapped_column(Integer, nullable=False)
|
||||||
created_at: Mapped[datetime] = mapped_column(
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc)
|
DateTime(timezone=True), default=lambda: datetime.now(UTC)
|
||||||
)
|
)
|
||||||
last_used: Mapped[datetime | None] = mapped_column(
|
last_used: Mapped[datetime | None] = mapped_column(
|
||||||
DateTime(timezone=True), nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
@@ -175,7 +232,7 @@ class CredentialModel(Base):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def as_dataclass(self):
|
def as_dataclass(self):
|
||||||
return Credential(
|
return _LegacyCredential(
|
||||||
uuid=UUID(bytes=self.uuid),
|
uuid=UUID(bytes=self.uuid),
|
||||||
credential_id=self.credential_id,
|
credential_id=self.credential_id,
|
||||||
user_uuid=UUID(bytes=self.user_uuid),
|
user_uuid=UUID(bytes=self.user_uuid),
|
||||||
@@ -205,12 +262,12 @@ class SessionModel(Base):
|
|||||||
user_agent: Mapped[str] = mapped_column(String(512), nullable=False)
|
user_agent: Mapped[str] = mapped_column(String(512), nullable=False)
|
||||||
renewed: Mapped[datetime] = mapped_column(
|
renewed: Mapped[datetime] = mapped_column(
|
||||||
DateTime(timezone=True),
|
DateTime(timezone=True),
|
||||||
default=lambda: datetime.now(timezone.utc),
|
default=lambda: datetime.now(UTC),
|
||||||
nullable=False,
|
nullable=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
def as_dataclass(self):
|
def as_dataclass(self):
|
||||||
return _SqlSession(
|
return _LegacySession(
|
||||||
key=self.key,
|
key=self.key,
|
||||||
user_uuid=UUID(bytes=self.user_uuid),
|
user_uuid=UUID(bytes=self.user_uuid),
|
||||||
credential_uuid=UUID(bytes=self.credential_uuid),
|
credential_uuid=UUID(bytes=self.credential_uuid),
|
||||||
@@ -221,7 +278,7 @@ class SessionModel(Base):
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def from_dataclass(session: _SqlSession):
|
def from_dataclass(session: _LegacySession):
|
||||||
return SessionModel(
|
return SessionModel(
|
||||||
key=session.key,
|
key=session.key,
|
||||||
user_uuid=session.user_uuid.bytes,
|
user_uuid=session.user_uuid.bytes,
|
||||||
@@ -243,8 +300,8 @@ class ResetTokenModel(Base):
|
|||||||
token_type: Mapped[str] = mapped_column(String, nullable=False)
|
token_type: Mapped[str] = mapped_column(String, nullable=False)
|
||||||
expiry: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
expiry: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||||
|
|
||||||
def as_dataclass(self) -> ResetToken:
|
def as_dataclass(self) -> _LegacyResetToken:
|
||||||
return ResetToken(
|
return _LegacyResetToken(
|
||||||
key=self.key,
|
key=self.key,
|
||||||
user_uuid=UUID(bytes=self.user_uuid),
|
user_uuid=UUID(bytes=self.user_uuid),
|
||||||
token_type=self.token_type,
|
token_type=self.token_type,
|
||||||
@@ -338,7 +395,7 @@ class DB:
|
|||||||
result = await session.execute(select(PermissionModel))
|
result = await session.execute(select(PermissionModel))
|
||||||
return [p.as_dataclass() for p in result.scalars().all()]
|
return [p.as_dataclass() for p in result.scalars().all()]
|
||||||
|
|
||||||
async def list_organizations(self) -> list[Org]:
|
async def list_organizations(self) -> list[_LegacyOrg]:
|
||||||
async with self.session() as session:
|
async with self.session() as session:
|
||||||
# Load all orgs
|
# Load all orgs
|
||||||
orgs_result = await session.execute(select(OrgModel))
|
orgs_result = await session.execute(select(OrgModel))
|
||||||
@@ -365,13 +422,13 @@ class DB:
|
|||||||
perms_by_role.setdefault(rp.role_uuid, []).append(rp.permission_id)
|
perms_by_role.setdefault(rp.role_uuid, []).append(rp.permission_id)
|
||||||
|
|
||||||
# Build org dataclasses with roles and permission IDs
|
# Build org dataclasses with roles and permission IDs
|
||||||
roles_by_org: dict[bytes, list[Role]] = {}
|
roles_by_org: dict[bytes, list[_LegacyRole]] = {}
|
||||||
for rm in role_models:
|
for rm in role_models:
|
||||||
r_dc = rm.as_dataclass()
|
r_dc = rm.as_dataclass()
|
||||||
r_dc.permissions = perms_by_role.get(rm.uuid, [])
|
r_dc.permissions = perms_by_role.get(rm.uuid, [])
|
||||||
roles_by_org.setdefault(rm.org_uuid, []).append(r_dc)
|
roles_by_org.setdefault(rm.org_uuid, []).append(r_dc)
|
||||||
|
|
||||||
orgs: list[Org] = []
|
orgs: list[_LegacyOrg] = []
|
||||||
for om in org_models:
|
for om in org_models:
|
||||||
o_dc = om.as_dataclass()
|
o_dc = om.as_dataclass()
|
||||||
o_dc.permissions = perms_by_org.get(om.uuid, [])
|
o_dc.permissions = perms_by_org.get(om.uuid, [])
|
||||||
|
|||||||
@@ -19,12 +19,12 @@ The first 3 words of the token serve as the pairing code for manual entry.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
|
from collections.abc import Callable
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import UTC, datetime, timedelta
|
||||||
from typing import Callable
|
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
from paskia.util import passphrase
|
from paskia.util import passphrase, pow
|
||||||
|
|
||||||
# Remote auth requests expire after this duration
|
# Remote auth requests expire after this duration
|
||||||
REMOTE_AUTH_LIFETIME = timedelta(minutes=5)
|
REMOTE_AUTH_LIFETIME = timedelta(minutes=5)
|
||||||
@@ -94,7 +94,7 @@ class RemoteAuthManager:
|
|||||||
|
|
||||||
async def _cleanup_expired(self):
|
async def _cleanup_expired(self):
|
||||||
"""Remove expired requests and notify waiting clients."""
|
"""Remove expired requests and notify waiting clients."""
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
expired_keys = []
|
expired_keys = []
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
for key, req in self._requests.items():
|
for key, req in self._requests.items():
|
||||||
@@ -123,7 +123,7 @@ class RemoteAuthManager:
|
|||||||
Returns:
|
Returns:
|
||||||
(code, expiry) - The 3-word passphrase code and expiration time
|
(code, expiry) - The 3-word passphrase code and expiration time
|
||||||
"""
|
"""
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
expiry = now + REMOTE_AUTH_LIFETIME
|
expiry = now + REMOTE_AUTH_LIFETIME
|
||||||
|
|
||||||
async with self._lock:
|
async with self._lock:
|
||||||
@@ -160,7 +160,7 @@ class RemoteAuthManager:
|
|||||||
req = self._requests.get(normalized)
|
req = self._requests.get(normalized)
|
||||||
if req is None:
|
if req is None:
|
||||||
return None
|
return None
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
||||||
# Expired
|
# Expired
|
||||||
del self._requests[normalized]
|
del self._requests[normalized]
|
||||||
@@ -319,7 +319,6 @@ class RemoteAuthManager:
|
|||||||
Returns:
|
Returns:
|
||||||
PoW work units (pow.NORMAL or pow.HARD)
|
PoW work units (pow.NORMAL or pow.HARD)
|
||||||
"""
|
"""
|
||||||
from paskia.util import pow
|
|
||||||
|
|
||||||
count = self.get_connection_count()
|
count = self.get_connection_count()
|
||||||
return pow.HARD if count >= 10 else pow.NORMAL
|
return pow.HARD if count >= 10 else pow.NORMAL
|
||||||
@@ -332,7 +331,7 @@ class RemoteAuthManager:
|
|||||||
req = self._requests.get(token)
|
req = self._requests.get(token)
|
||||||
if req is None:
|
if req is None:
|
||||||
return None
|
return None
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(UTC)
|
||||||
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
if now > req.created_at + REMOTE_AUTH_LIFETIME:
|
||||||
del self._requests[token]
|
del self._requests[token]
|
||||||
return None
|
return None
|
||||||
|
|||||||
+6
-12
@@ -8,11 +8,9 @@ This module provides a unified interface for WebAuthn operations including:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
from datetime import datetime, timezone
|
|
||||||
from urllib.parse import urlparse
|
from urllib.parse import urlparse
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import uuid7
|
|
||||||
from webauthn import (
|
from webauthn import (
|
||||||
generate_authentication_options,
|
generate_authentication_options,
|
||||||
generate_registration_options,
|
generate_registration_options,
|
||||||
@@ -176,14 +174,12 @@ class Passkey:
|
|||||||
expected_origin=origin,
|
expected_origin=origin,
|
||||||
expected_rp_id=self.rp_id,
|
expected_rp_id=self.rp_id,
|
||||||
)
|
)
|
||||||
return Credential(
|
return Credential.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
credential_id=credential.raw_id,
|
credential_id=credential.raw_id,
|
||||||
user_uuid=user_uuid,
|
user=user_uuid,
|
||||||
aaguid=UUID(registration.aaguid),
|
aaguid=UUID(registration.aaguid),
|
||||||
public_key=registration.credential_public_key,
|
public_key=registration.credential_public_key,
|
||||||
sign_count=registration.sign_count,
|
sign_count=registration.sign_count,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
### Authentication Methods ###
|
### Authentication Methods ###
|
||||||
@@ -234,8 +230,11 @@ class Passkey:
|
|||||||
Args:
|
Args:
|
||||||
credential: The authentication credential response from the client
|
credential: The authentication credential response from the client
|
||||||
expected_challenge: The earlier generated challenge bytes
|
expected_challenge: The earlier generated challenge bytes
|
||||||
stored_cred: The server stored credential record (modified by this function)
|
stored_cred: The server stored credential record (NOT modified)
|
||||||
origin: The origin URL (required, must be pre-validated)
|
origin: The origin URL (required, must be pre-validated)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
VerifiedAuthentication with new_sign_count and user_verified status
|
||||||
"""
|
"""
|
||||||
# Verify the authentication response
|
# Verify the authentication response
|
||||||
verification = verify_authentication_response(
|
verification = verify_authentication_response(
|
||||||
@@ -246,11 +245,6 @@ class Passkey:
|
|||||||
credential_public_key=stored_cred.public_key,
|
credential_public_key=stored_cred.public_key,
|
||||||
credential_current_sign_count=stored_cred.sign_count,
|
credential_current_sign_count=stored_cred.sign_count,
|
||||||
)
|
)
|
||||||
stored_cred.sign_count = verification.new_sign_count
|
|
||||||
now = datetime.now(timezone.utc)
|
|
||||||
stored_cred.last_used = now
|
|
||||||
if verification.user_verified:
|
|
||||||
stored_cred.last_verified = now
|
|
||||||
return verification
|
return verification
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,110 @@
|
|||||||
|
"""API response utilities using msgspec for JSON serialization.
|
||||||
|
|
||||||
|
msgspec handles UUID and datetime conversion automatically.
|
||||||
|
API structs inherit from db structs with kw_only=True to add uuid/key fields.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from datetime import UTC, datetime
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
import msgspec
|
||||||
|
|
||||||
|
from paskia.db.structs import Org, Permission, Role, User
|
||||||
|
from paskia.util import useragent
|
||||||
|
|
||||||
|
|
||||||
|
def _utc_datetime(dt: datetime | None) -> datetime | None:
|
||||||
|
"""Convert datetime to UTC, handling both aware and naive datetimes."""
|
||||||
|
if dt is None:
|
||||||
|
return None
|
||||||
|
if dt.tzinfo:
|
||||||
|
return dt.astimezone(UTC)
|
||||||
|
return dt.replace(tzinfo=UTC)
|
||||||
|
|
||||||
|
|
||||||
|
def format_datetime(dt: datetime | None) -> str | None:
|
||||||
|
"""Format a datetime to ISO 8601 string with Z suffix for UTC."""
|
||||||
|
if dt is None:
|
||||||
|
return None
|
||||||
|
utc_dt = _utc_datetime(dt)
|
||||||
|
return utc_dt.isoformat().replace("+00:00", "Z") if utc_dt else None
|
||||||
|
|
||||||
|
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
# API structs - inherit from db structs, add uuid for serialization
|
||||||
|
# -------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
class ApiUser(User, kw_only=True):
|
||||||
|
"""User with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, u: User) -> "ApiUser":
|
||||||
|
return cls(uuid=u.uuid, **msgspec.structs.asdict(u))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiOrg(Org, kw_only=True):
|
||||||
|
"""Org with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, o: Org) -> "ApiOrg":
|
||||||
|
return cls(uuid=o.uuid, **msgspec.structs.asdict(o))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiRole(Role, kw_only=True):
|
||||||
|
"""Role with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, r: Role) -> "ApiRole":
|
||||||
|
return cls(uuid=r.uuid, **msgspec.structs.asdict(r))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiPermission(Permission, kw_only=True):
|
||||||
|
"""Permission with uuid serialized."""
|
||||||
|
|
||||||
|
uuid: UUID
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(cls, p: Permission) -> "ApiPermission":
|
||||||
|
return cls(uuid=p.uuid, **msgspec.structs.asdict(p))
|
||||||
|
|
||||||
|
|
||||||
|
class ApiSession(msgspec.Struct):
|
||||||
|
"""Session for API responses with computed fields."""
|
||||||
|
|
||||||
|
id: str
|
||||||
|
credential_uuid: UUID = msgspec.field(name="credential")
|
||||||
|
host: str
|
||||||
|
ip: str
|
||||||
|
user_agent: str
|
||||||
|
last_renewed: datetime
|
||||||
|
is_current: bool = False
|
||||||
|
is_current_host: bool = False
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_db(
|
||||||
|
cls,
|
||||||
|
s, # Session
|
||||||
|
*,
|
||||||
|
current_key: str,
|
||||||
|
normalized_host: str | None,
|
||||||
|
expires_delta, # timedelta
|
||||||
|
) -> "ApiSession":
|
||||||
|
return cls(
|
||||||
|
id=s.key,
|
||||||
|
credential_uuid=s.credential_uuid,
|
||||||
|
host=s.host,
|
||||||
|
ip=s.ip,
|
||||||
|
user_agent=useragent.compact_user_agent(s.user_agent),
|
||||||
|
last_renewed=s.expiry - expires_delta,
|
||||||
|
is_current=s.key == current_key,
|
||||||
|
is_current_host=bool(
|
||||||
|
normalized_host and s.host and s.host == normalized_host
|
||||||
|
),
|
||||||
|
)
|
||||||
@@ -3,7 +3,7 @@
|
|||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from urllib.parse import urlsplit
|
from urllib.parse import urlparse, urlsplit
|
||||||
|
|
||||||
|
|
||||||
@lru_cache(maxsize=1)
|
@lru_cache(maxsize=1)
|
||||||
@@ -24,7 +24,6 @@ def dedicated_auth_host() -> str | None:
|
|||||||
auth_host = _load_config().get("auth_host")
|
auth_host = _load_config().get("auth_host")
|
||||||
if not auth_host:
|
if not auth_host:
|
||||||
return None
|
return None
|
||||||
from urllib.parse import urlparse
|
|
||||||
|
|
||||||
parsed = urlparse(auth_host if "://" in auth_host else f"//{auth_host}")
|
parsed = urlparse(auth_host if "://" in auth_host else f"//{auth_host}")
|
||||||
return parsed.netloc or parsed.path or None
|
return parsed.netloc or parsed.path or None
|
||||||
|
|||||||
@@ -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.get_session_context(auth, normalized_host)
|
return db.data().session_ctx(auth, normalized_host)
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
"""Utility functions for session validation and checking."""
|
"""Utility functions for session validation and checking."""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
from paskia.db import SessionContext
|
from paskia.db import SessionContext
|
||||||
@@ -34,5 +34,5 @@ def check_session_age(ctx: SessionContext, max_age: str | None) -> bool:
|
|||||||
else:
|
else:
|
||||||
auth_time = ctx.session.expiry - EXPIRES
|
auth_time = ctx.session.expiry - EXPIRES
|
||||||
|
|
||||||
time_since_auth = datetime.now(timezone.utc) - auth_time
|
time_since_auth = datetime.now(UTC) - auth_time
|
||||||
return time_since_auth <= max_age_delta
|
return time_since_auth <= max_age_delta
|
||||||
|
|||||||
+40
-125
@@ -1,145 +1,60 @@
|
|||||||
"""User information formatting and retrieval logic."""
|
"""User information formatting and retrieval logic."""
|
||||||
|
|
||||||
from datetime import timezone
|
|
||||||
|
|
||||||
from paskia import aaguid, db
|
from paskia import aaguid, db
|
||||||
from paskia.authsession import EXPIRES
|
from paskia.authsession import EXPIRES
|
||||||
from paskia.util import hostutil, permutil, useragent
|
from paskia.db import SessionContext
|
||||||
|
from paskia.util import hostutil, permutil
|
||||||
|
from paskia.util.apistructs import ApiSession
|
||||||
|
|
||||||
|
|
||||||
def _format_datetime(dt):
|
def build_session_context(ctx: SessionContext) -> dict:
|
||||||
"""Format a datetime object to ISO 8601 string with UTC timezone."""
|
"""Build session context dict from SessionContext."""
|
||||||
if dt is None:
|
return {
|
||||||
return None
|
"user": {"uuid": ctx.user.uuid, "display_name": ctx.user.display_name},
|
||||||
if dt.tzinfo:
|
"org": {"uuid": ctx.org.uuid, "display_name": ctx.org.display_name},
|
||||||
return dt.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
"role": {"uuid": ctx.role.uuid, "display_name": ctx.role.display_name},
|
||||||
else:
|
"permissions": [p.scope for p in ctx.permissions],
|
||||||
return dt.replace(tzinfo=timezone.utc).isoformat().replace("+00:00", "Z")
|
}
|
||||||
|
|
||||||
|
|
||||||
async def format_user_info(
|
async def build_user_info(
|
||||||
*,
|
*,
|
||||||
user_uuid,
|
user_uuid,
|
||||||
auth: str,
|
auth: str,
|
||||||
session_record,
|
session_record,
|
||||||
request_host: str | None,
|
request_host: str | None,
|
||||||
) -> dict:
|
) -> dict:
|
||||||
"""Format complete user information for authenticated users.
|
"""Build user info dict for authenticated users."""
|
||||||
|
|
||||||
Args:
|
|
||||||
user_uuid: UUID of the user to fetch information for
|
|
||||||
auth: Authentication token
|
|
||||||
session_record: Current session record
|
|
||||||
request_host: Host header from the request
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dictionary containing formatted user information including:
|
|
||||||
- User details
|
|
||||||
- Organization and role information
|
|
||||||
- Credentials list
|
|
||||||
- Sessions list
|
|
||||||
- Permissions
|
|
||||||
"""
|
|
||||||
u = db.get_user_by_uuid(user_uuid)
|
|
||||||
ctx = await permutil.session_context(auth, request_host)
|
ctx = await permutil.session_context(auth, request_host)
|
||||||
|
user = db.data().users[user_uuid]
|
||||||
|
normalized_host = hostutil.normalize_host(request_host)
|
||||||
|
|
||||||
# Fetch and format credentials
|
credentials = sorted(user.credentials, key=lambda c: c.created_at)
|
||||||
user_credentials = db.get_credentials_by_user_uuid(user_uuid)
|
return {
|
||||||
credentials: list[dict] = []
|
"ctx": build_session_context(ctx),
|
||||||
user_aaguids: set[str] = set()
|
"created_at": ctx.user.created_at,
|
||||||
|
"last_seen": ctx.user.last_seen,
|
||||||
for c in user_credentials:
|
"visits": ctx.user.visits,
|
||||||
aaguid_str = str(c.aaguid)
|
"credentials": [
|
||||||
user_aaguids.add(aaguid_str)
|
|
||||||
credentials.append(
|
|
||||||
{
|
{
|
||||||
"credential_uuid": str(c.uuid),
|
"credential": c.uuid,
|
||||||
"aaguid": aaguid_str,
|
"aaguid": c.aaguid,
|
||||||
"created_at": _format_datetime(c.created_at),
|
"created_at": c.created_at,
|
||||||
"last_used": _format_datetime(c.last_used),
|
"last_used": c.last_used,
|
||||||
"last_verified": _format_datetime(c.last_verified),
|
"last_verified": c.last_verified,
|
||||||
"sign_count": c.sign_count,
|
"sign_count": c.sign_count,
|
||||||
"is_current_session": session_record.credential_uuid == c.uuid,
|
"is_current_session": session_record.credential == c.uuid,
|
||||||
}
|
}
|
||||||
)
|
for c in credentials
|
||||||
|
],
|
||||||
credentials.sort(key=lambda cred: cred["created_at"])
|
"aaguid_info": aaguid.filter(c.aaguid for c in credentials),
|
||||||
aaguid_info = aaguid.filter(user_aaguids)
|
"sessions": [
|
||||||
|
ApiSession.from_db(
|
||||||
# Format role and org information
|
s,
|
||||||
role_info = None
|
current_key=auth,
|
||||||
org_info = None
|
normalized_host=normalized_host,
|
||||||
effective_permissions: list[str] = []
|
expires_delta=EXPIRES,
|
||||||
|
)
|
||||||
if ctx:
|
for s in user.sessions
|
||||||
role_info = {
|
],
|
||||||
"uuid": str(ctx.role.uuid),
|
|
||||||
"display_name": ctx.role.display_name,
|
|
||||||
"permissions": ctx.role.permissions,
|
|
||||||
}
|
|
||||||
org_info = {
|
|
||||||
"uuid": str(ctx.org.uuid),
|
|
||||||
"display_name": ctx.org.display_name,
|
|
||||||
"permissions": ctx.org.permissions,
|
|
||||||
}
|
|
||||||
effective_permissions = [p.scope for p in (ctx.permissions or [])]
|
|
||||||
|
|
||||||
# Format sessions
|
|
||||||
normalized_request_host = hostutil.normalize_host(request_host)
|
|
||||||
session_records = db.list_sessions_for_user(user_uuid)
|
|
||||||
current_session_key = auth
|
|
||||||
sessions_payload: list[dict] = []
|
|
||||||
|
|
||||||
for entry in session_records:
|
|
||||||
sessions_payload.append(
|
|
||||||
{
|
|
||||||
"id": entry.key,
|
|
||||||
"credential_uuid": str(entry.credential_uuid),
|
|
||||||
"host": entry.host,
|
|
||||||
"ip": entry.ip,
|
|
||||||
"user_agent": useragent.compact_user_agent(entry.user_agent),
|
|
||||||
"last_renewed": _format_datetime(entry.expiry - EXPIRES),
|
|
||||||
"is_current": entry.key == current_session_key,
|
|
||||||
"is_current_host": bool(
|
|
||||||
normalized_request_host
|
|
||||||
and entry.host
|
|
||||||
and entry.host == normalized_request_host
|
|
||||||
),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"authenticated": True,
|
|
||||||
"user": {
|
|
||||||
"user_uuid": str(u.uuid),
|
|
||||||
"user_name": u.display_name,
|
|
||||||
"created_at": _format_datetime(u.created_at),
|
|
||||||
"last_seen": _format_datetime(u.last_seen),
|
|
||||||
"visits": u.visits,
|
|
||||||
},
|
|
||||||
"org": org_info,
|
|
||||||
"role": role_info,
|
|
||||||
"permissions": effective_permissions,
|
|
||||||
"credentials": credentials,
|
|
||||||
"aaguid_info": aaguid_info,
|
|
||||||
"sessions": sessions_payload,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
async def format_reset_user_info(user_uuid, reset_token) -> dict:
|
|
||||||
"""Format minimal user information for reset token requests.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
user_uuid: UUID of the user
|
|
||||||
reset_token: Reset token record
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Dictionary with minimal user info for password reset flow
|
|
||||||
"""
|
|
||||||
u = db.get_user_by_uuid(user_uuid)
|
|
||||||
|
|
||||||
return {
|
|
||||||
"authenticated": False,
|
|
||||||
"session_type": reset_token.token_type,
|
|
||||||
"user": {"user_uuid": str(u.uuid), "user_name": u.display_name},
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ but doesn't provide server-side fetching of HTML content.
|
|||||||
import asyncio
|
import asyncio
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import os
|
import os
|
||||||
|
from importlib import resources
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -24,7 +25,6 @@ def _get_dev_server() -> str | None:
|
|||||||
|
|
||||||
def _resolve_static_dir() -> Path:
|
def _resolve_static_dir() -> Path:
|
||||||
"""Resolve the static files directory."""
|
"""Resolve the static files directory."""
|
||||||
from importlib import resources
|
|
||||||
|
|
||||||
# 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
|
try: # pragma: no cover - trivial path resolution
|
||||||
|
|||||||
+1
-5
@@ -74,12 +74,8 @@ filterwarnings = [
|
|||||||
"ignore::DeprecationWarning",
|
"ignore::DeprecationWarning",
|
||||||
]
|
]
|
||||||
|
|
||||||
[tool.ruff]
|
|
||||||
target-version = "py39"
|
|
||||||
line-length = 88
|
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
select = ["E", "F", "I", "N", "W", "UP"]
|
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"]
|
||||||
|
|
||||||
|
|||||||
+51
-101
@@ -13,34 +13,33 @@ import asyncio
|
|||||||
import os
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
from collections.abc import AsyncGenerator
|
from collections.abc import AsyncGenerator
|
||||||
from datetime import datetime, timezone
|
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
import pytest_asyncio
|
import pytest_asyncio
|
||||||
import uuid7
|
|
||||||
|
|
||||||
|
import paskia.db.operations as ops_db
|
||||||
from paskia import globals as paskia_globals
|
from paskia import globals as paskia_globals
|
||||||
from paskia.authsession import expires
|
from paskia.authsession import expires, reset_expires
|
||||||
from paskia.db import (
|
from paskia.db import (
|
||||||
Credential,
|
Credential,
|
||||||
Org,
|
Org,
|
||||||
Permission,
|
Permission,
|
||||||
Role,
|
Role,
|
||||||
User,
|
User,
|
||||||
add_permission_to_organization,
|
|
||||||
create_credential,
|
create_credential,
|
||||||
create_organization,
|
|
||||||
create_permission,
|
|
||||||
create_reset_token,
|
create_reset_token,
|
||||||
create_role,
|
create_role,
|
||||||
create_session,
|
create_session,
|
||||||
create_user,
|
create_user,
|
||||||
)
|
)
|
||||||
from paskia.db.operations import DB, _create_token
|
from paskia.db.jsonl import JsonlStore
|
||||||
|
from paskia.db.operations import DB
|
||||||
|
from paskia.fastapi.mainapp import app
|
||||||
from paskia.fastapi.session import AUTH_COOKIE_NAME
|
from paskia.fastapi.session import AUTH_COOKIE_NAME
|
||||||
from paskia.sansio import Passkey
|
from paskia.sansio import Passkey
|
||||||
|
from paskia.util.passphrase import generate
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(scope="session")
|
@pytest.fixture(scope="session")
|
||||||
@@ -55,17 +54,27 @@ def event_loop():
|
|||||||
async def test_db() -> AsyncGenerator[DB, None]:
|
async def test_db() -> AsyncGenerator[DB, None]:
|
||||||
"""Create an in-memory JSON database for testing.
|
"""Create an in-memory JSON database for testing.
|
||||||
|
|
||||||
Uses a temp file that gets cleaned up after each test.
|
Uses bootstrap() to properly initialize the database with:
|
||||||
|
- auth:admin and auth:org:admin permissions
|
||||||
|
- A default organization with Administration role
|
||||||
|
- An admin user with the Administration role
|
||||||
"""
|
"""
|
||||||
import paskia.db.operations as ops_db
|
|
||||||
|
|
||||||
with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f:
|
with tempfile.NamedTemporaryFile(suffix=".jsonl", delete=True) as f:
|
||||||
db = DB(f.name)
|
db = DB()
|
||||||
await db.load()
|
store = JsonlStore(db, f.name)
|
||||||
|
db._store = store
|
||||||
|
await store.load()
|
||||||
ops_db._db = db
|
ops_db._db = db
|
||||||
|
ops_db._store = store
|
||||||
|
# Bootstrap creates the initial permissions, org, role, and admin user
|
||||||
|
ops_db.bootstrap(
|
||||||
|
org_name="Test Organization",
|
||||||
|
admin_name="Test Admin",
|
||||||
|
)
|
||||||
yield db
|
yield db
|
||||||
# Clean up
|
|
||||||
ops_db._db = None
|
ops_db._db = None
|
||||||
|
ops_db._store = None
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -81,98 +90,56 @@ async def passkey_instance() -> Passkey:
|
|||||||
paskia_globals.passkey._instance = None
|
paskia_globals.passkey._instance = None
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
|
||||||
async def test_org(test_db: DB, admin_permission: Permission) -> Org:
|
|
||||||
"""Create a test organization with admin permission."""
|
|
||||||
org = Org(
|
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Test Organization",
|
|
||||||
permissions=[str(admin_permission.uuid)], # Org can grant this permission
|
|
||||||
)
|
|
||||||
create_organization(org)
|
|
||||||
return org
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def admin_permission(test_db: DB) -> Permission:
|
async def admin_permission(test_db: DB) -> Permission:
|
||||||
"""Create the auth:admin permission."""
|
"""Get the auth:admin permission created by bootstrap."""
|
||||||
import uuid7
|
return next(p for p in test_db.permissions.values() if p.scope == "auth:admin")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="auth:admin", display_name="Master Admin"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
|
||||||
return perm
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def org_admin_permission(test_db: DB, test_org: Org) -> Permission:
|
async def org_admin_permission(test_db: DB) -> Permission:
|
||||||
"""Create the auth:org:admin permission."""
|
"""Get the auth:org:admin permission created by bootstrap."""
|
||||||
import uuid7
|
return next(p for p in test_db.permissions.values() if p.scope == "auth:org:admin")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="auth:org:admin", display_name="Organization Admin"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
|
||||||
# Make it grantable by the org
|
|
||||||
add_permission_to_organization(str(test_org.uuid), "auth:org:admin")
|
|
||||||
return perm
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def test_role(
|
async def test_org(test_db: DB) -> Org:
|
||||||
test_db: DB,
|
"""Get the test organization created by bootstrap."""
|
||||||
test_org: Org,
|
# Bootstrap creates exactly one org
|
||||||
admin_permission: Permission,
|
return next(iter(test_db.orgs.values()))
|
||||||
org_admin_permission: Permission,
|
|
||||||
) -> Role:
|
|
||||||
"""Create a test role with admin permission."""
|
@pytest_asyncio.fixture(scope="function")
|
||||||
role = Role(
|
async def test_role(test_db: DB) -> Role:
|
||||||
uuid=uuid7.create(),
|
"""Get the Administration role created by bootstrap."""
|
||||||
org_uuid=test_org.uuid,
|
# Bootstrap creates exactly one role (Administration)
|
||||||
display_name="Test Admin Role",
|
return next(iter(test_db.roles.values()))
|
||||||
permissions=[str(admin_permission.uuid), str(org_admin_permission.uuid)],
|
|
||||||
)
|
|
||||||
create_role(role)
|
|
||||||
return role
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def user_role(test_db: DB, test_org: Org) -> Role:
|
async def user_role(test_db: DB, test_org: Org) -> Role:
|
||||||
"""Create a test role without admin permission (regular user)."""
|
"""Create a test role without admin permission (regular user)."""
|
||||||
role = Role(
|
role = Role.create(
|
||||||
uuid=uuid7.create(),
|
org=test_org.uuid,
|
||||||
org_uuid=test_org.uuid,
|
|
||||||
display_name="User Role",
|
display_name="User Role",
|
||||||
permissions=[],
|
|
||||||
)
|
)
|
||||||
create_role(role)
|
create_role(role)
|
||||||
return role
|
return role
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def test_user(test_db: DB, test_role: Role) -> User:
|
async def test_user(test_db: DB) -> User:
|
||||||
"""Create a test user with admin role."""
|
"""Get the admin user created by bootstrap."""
|
||||||
user = User(
|
# Bootstrap creates exactly one user (admin)
|
||||||
uuid=uuid7.create(),
|
return next(iter(test_db.users.values()))
|
||||||
display_name="Test Admin",
|
|
||||||
role_uuid=test_role.uuid,
|
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
visits=0,
|
|
||||||
)
|
|
||||||
create_user(user)
|
|
||||||
return user
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def regular_user(test_db: DB, user_role: Role) -> User:
|
async def regular_user(test_db: DB, user_role: Role) -> User:
|
||||||
"""Create a regular test user without admin permissions."""
|
"""Create a regular test user without admin permissions."""
|
||||||
user = User(
|
user = User.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Regular User",
|
display_name="Regular User",
|
||||||
role_uuid=user_role.uuid,
|
role=user_role.uuid,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
visits=0,
|
|
||||||
)
|
)
|
||||||
create_user(user)
|
create_user(user)
|
||||||
return user
|
return user
|
||||||
@@ -181,16 +148,12 @@ async def regular_user(test_db: DB, user_role: Role) -> User:
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def test_credential(test_db: DB, test_user: User) -> Credential:
|
async def test_credential(test_db: DB, test_user: User) -> Credential:
|
||||||
"""Create a test credential for the admin user."""
|
"""Create a test credential for the admin user."""
|
||||||
credential = Credential(
|
credential = Credential.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
credential_id=os.urandom(32),
|
credential_id=os.urandom(32),
|
||||||
user_uuid=test_user.uuid,
|
user=test_user.uuid,
|
||||||
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
||||||
public_key=os.urandom(64),
|
public_key=os.urandom(64),
|
||||||
sign_count=0,
|
sign_count=0,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
last_used=None,
|
|
||||||
last_verified=None,
|
|
||||||
)
|
)
|
||||||
create_credential(credential)
|
create_credential(credential)
|
||||||
return credential
|
return credential
|
||||||
@@ -199,16 +162,12 @@ async def test_credential(test_db: DB, test_user: User) -> Credential:
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def regular_credential(test_db: DB, regular_user: User) -> Credential:
|
async def regular_credential(test_db: DB, regular_user: User) -> Credential:
|
||||||
"""Create a test credential for the regular user."""
|
"""Create a test credential for the regular user."""
|
||||||
credential = Credential(
|
credential = Credential.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
credential_id=os.urandom(32),
|
credential_id=os.urandom(32),
|
||||||
user_uuid=regular_user.uuid,
|
user=regular_user.uuid,
|
||||||
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
||||||
public_key=os.urandom(64),
|
public_key=os.urandom(64),
|
||||||
sign_count=0,
|
sign_count=0,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
last_used=None,
|
|
||||||
last_verified=None,
|
|
||||||
)
|
)
|
||||||
create_credential(credential)
|
create_credential(credential)
|
||||||
return credential
|
return credential
|
||||||
@@ -219,17 +178,14 @@ async def session_token(
|
|||||||
test_db: DB, test_user: User, test_credential: Credential
|
test_db: DB, test_user: User, test_credential: Credential
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Create a session for the admin user and return the token."""
|
"""Create a session for the admin user and return the token."""
|
||||||
token = _create_token()
|
return create_session(
|
||||||
create_session(
|
|
||||||
user_uuid=test_user.uuid,
|
user_uuid=test_user.uuid,
|
||||||
credential_uuid=test_credential.uuid,
|
credential_uuid=test_credential.uuid,
|
||||||
key=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -237,24 +193,19 @@ async def regular_session_token(
|
|||||||
test_db: DB, regular_user: User, regular_credential: Credential
|
test_db: DB, regular_user: User, regular_credential: Credential
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Create a session for a regular user and return the token."""
|
"""Create a session for a regular user and return the token."""
|
||||||
token = _create_token()
|
return create_session(
|
||||||
create_session(
|
|
||||||
user_uuid=regular_user.uuid,
|
user_uuid=regular_user.uuid,
|
||||||
credential_uuid=regular_credential.uuid,
|
credential_uuid=regular_credential.uuid,
|
||||||
key=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def reset_token(test_db: DB, test_user: User, test_credential: Credential) -> str:
|
async def reset_token(test_db: DB, test_user: User, test_credential: Credential) -> str:
|
||||||
"""Create a reset token for the test user."""
|
"""Create a reset token for the test user."""
|
||||||
from paskia.authsession import reset_expires
|
|
||||||
from paskia.util.passphrase import generate
|
|
||||||
|
|
||||||
token = generate()
|
token = generate()
|
||||||
create_reset_token(
|
create_reset_token(
|
||||||
@@ -276,7 +227,6 @@ async def client(
|
|||||||
initialized first.
|
initialized first.
|
||||||
"""
|
"""
|
||||||
# Import app after globals are set
|
# Import app after globals are set
|
||||||
from paskia.fastapi.mainapp import app
|
|
||||||
|
|
||||||
transport = httpx.ASGITransport(app=app)
|
transport = httpx.ASGITransport(app=app)
|
||||||
async with httpx.AsyncClient(
|
async with httpx.AsyncClient(
|
||||||
|
|||||||
+110
-173
@@ -11,7 +11,9 @@ These tests cover:
|
|||||||
- Credential management
|
- Credential management
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
import os
|
||||||
|
import secrets
|
||||||
|
from datetime import UTC, datetime
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -19,6 +21,7 @@ import pytest
|
|||||||
import pytest_asyncio
|
import pytest_asyncio
|
||||||
import uuid7
|
import uuid7
|
||||||
|
|
||||||
|
from paskia import db
|
||||||
from paskia.authsession import expires
|
from paskia.authsession import expires
|
||||||
from paskia.db import (
|
from paskia.db import (
|
||||||
Credential,
|
Credential,
|
||||||
@@ -26,15 +29,15 @@ from paskia.db import (
|
|||||||
Permission,
|
Permission,
|
||||||
Role,
|
Role,
|
||||||
User,
|
User,
|
||||||
add_permission_to_organization,
|
add_permission_to_org,
|
||||||
create_credential,
|
create_credential,
|
||||||
create_organization,
|
create_org,
|
||||||
create_permission,
|
create_permission,
|
||||||
create_role,
|
create_role,
|
||||||
create_session,
|
create_session,
|
||||||
create_user,
|
create_user,
|
||||||
)
|
)
|
||||||
from paskia.db.operations import DB, _create_token
|
from paskia.db.operations import DB
|
||||||
from tests.conftest import auth_headers
|
from tests.conftest import auth_headers
|
||||||
|
|
||||||
# -------------------- Additional Fixtures --------------------
|
# -------------------- Additional Fixtures --------------------
|
||||||
@@ -43,12 +46,10 @@ from tests.conftest import auth_headers
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def second_org(test_db: DB) -> Org:
|
async def second_org(test_db: DB) -> Org:
|
||||||
"""Create a second organization for deletion tests."""
|
"""Create a second organization for deletion tests."""
|
||||||
org = Org(
|
org = Org.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Second Organization",
|
display_name="Second Organization",
|
||||||
permissions=[],
|
|
||||||
)
|
)
|
||||||
create_organization(org)
|
create_org(org)
|
||||||
return org
|
return org
|
||||||
|
|
||||||
|
|
||||||
@@ -57,11 +58,10 @@ async def second_org_role(
|
|||||||
test_db: DB, second_org: Org, admin_permission: Permission
|
test_db: DB, second_org: Org, admin_permission: Permission
|
||||||
) -> Role:
|
) -> Role:
|
||||||
"""Create a role in the second org with admin permission."""
|
"""Create a role in the second org with admin permission."""
|
||||||
role = Role(
|
role = Role.create(
|
||||||
uuid=uuid7.create(),
|
org=second_org.uuid,
|
||||||
org_uuid=second_org.uuid,
|
|
||||||
display_name="Second Org Admin Role",
|
display_name="Second Org Admin Role",
|
||||||
permissions=[str(admin_permission.uuid)],
|
permissions={admin_permission.uuid},
|
||||||
)
|
)
|
||||||
create_role(role)
|
create_role(role)
|
||||||
return role
|
return role
|
||||||
@@ -70,12 +70,9 @@ async def second_org_role(
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def second_org_user(test_db: DB, second_org_role: Role) -> User:
|
async def second_org_user(test_db: DB, second_org_role: Role) -> User:
|
||||||
"""Create a user in the second org."""
|
"""Create a user in the second org."""
|
||||||
user = User(
|
user = User.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Second Org User",
|
display_name="Second Org User",
|
||||||
role_uuid=second_org_role.uuid,
|
role=second_org_role.uuid,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
visits=0,
|
|
||||||
)
|
)
|
||||||
create_user(user)
|
create_user(user)
|
||||||
return user
|
return user
|
||||||
@@ -84,18 +81,13 @@ async def second_org_user(test_db: DB, second_org_role: Role) -> User:
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def second_org_credential(test_db: DB, second_org_user: User) -> Credential:
|
async def second_org_credential(test_db: DB, second_org_user: User) -> Credential:
|
||||||
"""Create a credential for the second org user."""
|
"""Create a credential for the second org user."""
|
||||||
import os
|
|
||||||
|
|
||||||
credential = Credential(
|
credential = Credential.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
credential_id=os.urandom(32),
|
credential_id=os.urandom(32),
|
||||||
user_uuid=second_org_user.uuid,
|
user=second_org_user.uuid,
|
||||||
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
||||||
public_key=os.urandom(64),
|
public_key=os.urandom(64),
|
||||||
sign_count=0,
|
sign_count=0,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
last_used=datetime.now(timezone.utc),
|
|
||||||
last_verified=datetime.now(timezone.utc),
|
|
||||||
)
|
)
|
||||||
create_credential(credential)
|
create_credential(credential)
|
||||||
return credential
|
return credential
|
||||||
@@ -106,17 +98,14 @@ async def second_org_session_token(
|
|||||||
test_db: DB, second_org_user: User, second_org_credential: Credential
|
test_db: DB, second_org_user: User, second_org_credential: Credential
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Create a session for the second org admin user."""
|
"""Create a session for the second org admin user."""
|
||||||
token = _create_token()
|
return create_session(
|
||||||
create_session(
|
|
||||||
user_uuid=second_org_user.uuid,
|
user_uuid=second_org_user.uuid,
|
||||||
credential_uuid=second_org_credential.uuid,
|
credential_uuid=second_org_credential.uuid,
|
||||||
key=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
@@ -124,11 +113,10 @@ async def org_admin_role(
|
|||||||
test_db: DB, test_org: Org, org_admin_permission: Permission
|
test_db: DB, test_org: Org, org_admin_permission: Permission
|
||||||
) -> Role:
|
) -> Role:
|
||||||
"""Create a role with org admin permission only (no global admin)."""
|
"""Create a role with org admin permission only (no global admin)."""
|
||||||
role = Role(
|
role = Role.create(
|
||||||
uuid=uuid7.create(),
|
org=test_org.uuid,
|
||||||
org_uuid=test_org.uuid,
|
|
||||||
display_name="Org Admin Role",
|
display_name="Org Admin Role",
|
||||||
permissions=[str(org_admin_permission.uuid)],
|
permissions={org_admin_permission.uuid},
|
||||||
)
|
)
|
||||||
create_role(role)
|
create_role(role)
|
||||||
return role
|
return role
|
||||||
@@ -137,14 +125,12 @@ async def org_admin_role(
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def org_admin_user(test_db: DB, org_admin_role: Role) -> User:
|
async def org_admin_user(test_db: DB, org_admin_role: Role) -> User:
|
||||||
"""Create a user with org admin permission only."""
|
"""Create a user with org admin permission only."""
|
||||||
user = User(
|
user = User.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Org Admin User",
|
display_name="Org Admin User",
|
||||||
role_uuid=org_admin_role.uuid,
|
role=org_admin_role.uuid,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
visits=5,
|
|
||||||
last_seen=datetime.now(timezone.utc),
|
|
||||||
)
|
)
|
||||||
|
user.visits = 5
|
||||||
|
user.last_seen = datetime.now(UTC)
|
||||||
create_user(user)
|
create_user(user)
|
||||||
return user
|
return user
|
||||||
|
|
||||||
@@ -152,18 +138,13 @@ async def org_admin_user(test_db: DB, org_admin_role: Role) -> User:
|
|||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def org_admin_credential(test_db: DB, org_admin_user: User) -> Credential:
|
async def org_admin_credential(test_db: DB, org_admin_user: User) -> Credential:
|
||||||
"""Create a credential for the org admin user."""
|
"""Create a credential for the org admin user."""
|
||||||
import os
|
|
||||||
|
|
||||||
credential = Credential(
|
credential = Credential.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
credential_id=os.urandom(32),
|
credential_id=os.urandom(32),
|
||||||
user_uuid=org_admin_user.uuid,
|
user=org_admin_user.uuid,
|
||||||
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
aaguid=UUID("00000000-0000-0000-0000-000000000000"),
|
||||||
public_key=os.urandom(64),
|
public_key=os.urandom(64),
|
||||||
sign_count=0,
|
sign_count=0,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
last_used=datetime.now(timezone.utc),
|
|
||||||
last_verified=None,
|
|
||||||
)
|
)
|
||||||
create_credential(credential)
|
create_credential(credential)
|
||||||
return credential
|
return credential
|
||||||
@@ -174,30 +155,23 @@ async def org_admin_session_token(
|
|||||||
test_db: DB, org_admin_user: User, org_admin_credential: Credential
|
test_db: DB, org_admin_user: User, org_admin_credential: Credential
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Create a session for the org admin user."""
|
"""Create a session for the org admin user."""
|
||||||
token = _create_token()
|
return create_session(
|
||||||
create_session(
|
|
||||||
user_uuid=org_admin_user.uuid,
|
user_uuid=org_admin_user.uuid,
|
||||||
credential_uuid=org_admin_credential.uuid,
|
credential_uuid=org_admin_credential.uuid,
|
||||||
key=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
expiry=expires(),
|
expiry=expires(),
|
||||||
)
|
)
|
||||||
return token
|
|
||||||
|
|
||||||
|
|
||||||
@pytest_asyncio.fixture(scope="function")
|
@pytest_asyncio.fixture(scope="function")
|
||||||
async def grantable_permission(test_db: DB, test_org: Org) -> Permission:
|
async def grantable_permission(test_db: DB, test_org: Org) -> Permission:
|
||||||
"""Create a permission and add it to org's grantable permissions."""
|
"""Create a permission and add it to org's grantable permissions."""
|
||||||
import uuid7
|
perm = Permission.create(scope="test:grantable:perm", display_name="Grantable Perm")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="test:grantable:perm", display_name="Grantable Perm"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
create_permission(perm)
|
||||||
# Add to org's grantable permissions
|
# Add to org's grantable permissions
|
||||||
add_permission_to_organization(str(test_org.uuid), perm.scope)
|
add_permission_to_org(test_org.uuid, perm.uuid)
|
||||||
return perm
|
return perm
|
||||||
|
|
||||||
|
|
||||||
@@ -390,10 +364,13 @@ class TestAdminOrganizations:
|
|||||||
):
|
):
|
||||||
"""Org admin cannot remove their org admin permission from org's permissions."""
|
"""Org admin cannot remove their org admin permission from org's permissions."""
|
||||||
# The auth:org:admin perm is already created and added by org_admin_permission fixture
|
# The auth:org:admin perm is already created and added by org_admin_permission fixture
|
||||||
|
org_admin_perm = next(
|
||||||
|
p for p in db.data().permissions.values() if p.scope == "auth:org:admin"
|
||||||
|
)
|
||||||
|
|
||||||
# Try to remove org admin perm (this is validated server-side in the remove endpoint)
|
# Try to remove org admin perm (this is validated server-side in the remove endpoint)
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=auth:org:admin",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={org_admin_perm.uuid}",
|
||||||
headers={**auth_headers(org_admin_session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(org_admin_session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
# This should fail because only global admin can remove perms from org
|
# This should fail because only global admin can remove perms from org
|
||||||
@@ -420,19 +397,14 @@ class TestAdminOrganizations:
|
|||||||
test_db: DB,
|
test_db: DB,
|
||||||
):
|
):
|
||||||
"""Admin should be able to delete another organization."""
|
"""Admin should be able to delete another organization."""
|
||||||
import uuid7
|
|
||||||
|
|
||||||
# Create org to delete
|
# Create org to delete
|
||||||
org_to_delete = Org(
|
org_to_delete = Org.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="Org To Delete",
|
display_name="Org To Delete",
|
||||||
permissions=[],
|
|
||||||
)
|
)
|
||||||
create_organization(org_to_delete)
|
create_org(org_to_delete)
|
||||||
|
|
||||||
# Create some org-specific permissions to test cleanup
|
# Create some org-specific permissions to test cleanup
|
||||||
org_perm = Permission(
|
org_perm = Permission.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
scope=f"test:org:{org_to_delete.uuid}:feature",
|
scope=f"test:org:{org_to_delete.uuid}:feature",
|
||||||
display_name="Org Feature",
|
display_name="Org Feature",
|
||||||
)
|
)
|
||||||
@@ -459,15 +431,12 @@ class TestAdminOrgPermissions:
|
|||||||
):
|
):
|
||||||
"""Admin should be able to add a permission to an org."""
|
"""Admin should be able to add a permission to an org."""
|
||||||
# First create a permission
|
# First create a permission
|
||||||
await client.post(
|
perm = Permission.create(scope="test:org:addable", display_name="Addable")
|
||||||
"/auth/api/admin/permissions",
|
create_permission(perm)
|
||||||
json={"scope": "test:org:addable", "display_name": "Addable"},
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
|
||||||
)
|
|
||||||
|
|
||||||
# Add it to the org
|
# Add it to the org
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=test:org:addable",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -482,8 +451,11 @@ class TestAdminOrgPermissions:
|
|||||||
test_org,
|
test_org,
|
||||||
):
|
):
|
||||||
"""Org admin cannot add permissions to org (requires global admin)."""
|
"""Org admin cannot add permissions to org (requires global admin)."""
|
||||||
|
admin_perm = next(
|
||||||
|
p for p in db.data().permissions.values() if p.scope == "auth:admin"
|
||||||
|
)
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=auth:admin",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={admin_perm.uuid}",
|
||||||
headers={**auth_headers(org_admin_session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(org_admin_session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 403
|
assert response.status_code == 403
|
||||||
@@ -494,19 +466,16 @@ class TestAdminOrgPermissions:
|
|||||||
):
|
):
|
||||||
"""Admin should be able to remove a permission from an org."""
|
"""Admin should be able to remove a permission from an org."""
|
||||||
# First create and add a permission
|
# First create and add a permission
|
||||||
|
perm = Permission.create(scope="test:org:removable", display_name="Removable")
|
||||||
|
create_permission(perm)
|
||||||
await client.post(
|
await client.post(
|
||||||
"/auth/api/admin/permissions",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={perm.uuid}",
|
||||||
json={"scope": "test:org:removable", "display_name": "Removable"},
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
|
||||||
)
|
|
||||||
await client.post(
|
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=test:org:removable",
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Remove it
|
# Remove it
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=test:org:removable",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -521,8 +490,11 @@ class TestAdminOrgPermissions:
|
|||||||
test_org,
|
test_org,
|
||||||
):
|
):
|
||||||
"""Org admin cannot remove permissions from org (requires global admin)."""
|
"""Org admin cannot remove permissions from org (requires global admin)."""
|
||||||
|
admin_perm = next(
|
||||||
|
p for p in db.data().permissions.values() if p.scope == "auth:admin"
|
||||||
|
)
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=auth:admin",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={admin_perm.uuid}",
|
||||||
headers={**auth_headers(org_admin_session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(org_admin_session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 403
|
assert response.status_code == 403
|
||||||
@@ -590,7 +562,7 @@ class TestAdminRoles:
|
|||||||
f"/auth/api/admin/orgs/{test_org.uuid}/roles",
|
f"/auth/api/admin/orgs/{test_org.uuid}/roles",
|
||||||
json={
|
json={
|
||||||
"display_name": "Role With Perms",
|
"display_name": "Role With Perms",
|
||||||
"permissions": [grantable_permission.scope],
|
"permissions": [str(grantable_permission.uuid)],
|
||||||
},
|
},
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
@@ -608,10 +580,7 @@ class TestAdminRoles:
|
|||||||
):
|
):
|
||||||
"""Creating role with non-grantable permission should fail."""
|
"""Creating role with non-grantable permission should fail."""
|
||||||
# Create permission but don't add to org
|
# Create permission but don't add to org
|
||||||
import uuid7
|
perm = Permission.create(
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(),
|
|
||||||
scope="test:not:grantable",
|
scope="test:not:grantable",
|
||||||
display_name="Not Grantable",
|
display_name="Not Grantable",
|
||||||
)
|
)
|
||||||
@@ -621,7 +590,7 @@ class TestAdminRoles:
|
|||||||
f"/auth/api/admin/orgs/{test_org.uuid}/roles",
|
f"/auth/api/admin/orgs/{test_org.uuid}/roles",
|
||||||
json={
|
json={
|
||||||
"display_name": "Bad Role",
|
"display_name": "Bad Role",
|
||||||
"permissions": ["test:not:grantable"],
|
"permissions": [str(perm.uuid)],
|
||||||
},
|
},
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
@@ -683,10 +652,7 @@ class TestAdminRoles:
|
|||||||
test_db: DB,
|
test_db: DB,
|
||||||
):
|
):
|
||||||
"""Adding non-grantable permission to role should fail."""
|
"""Adding non-grantable permission to role should fail."""
|
||||||
import uuid7
|
perm = Permission.create(
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(),
|
|
||||||
scope="test:not:grantable:update",
|
scope="test:not:grantable:update",
|
||||||
display_name="Not Grantable",
|
display_name="Not Grantable",
|
||||||
)
|
)
|
||||||
@@ -1110,12 +1076,9 @@ class TestAdminUsersInOrg:
|
|||||||
):
|
):
|
||||||
"""Creating link for user without credentials should return registration link."""
|
"""Creating link for user without credentials should return registration link."""
|
||||||
# Create user without credentials
|
# Create user without credentials
|
||||||
user_no_cred = User(
|
user_no_cred = User.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
display_name="User Without Creds",
|
display_name="User Without Creds",
|
||||||
role_uuid=user_role.uuid,
|
role=user_role.uuid,
|
||||||
created_at=datetime.now(timezone.utc),
|
|
||||||
visits=0,
|
|
||||||
)
|
)
|
||||||
create_user(user_no_cred)
|
create_user(user_no_cred)
|
||||||
|
|
||||||
@@ -1202,11 +1165,9 @@ class TestAdminSessions:
|
|||||||
):
|
):
|
||||||
"""Admin should be able to delete a user's session."""
|
"""Admin should be able to delete a user's session."""
|
||||||
# Create an additional session to delete
|
# Create an additional session to delete
|
||||||
extra_token = _create_token()
|
extra_token = create_session(
|
||||||
create_session(
|
|
||||||
user_uuid=test_user.uuid,
|
user_uuid=test_user.uuid,
|
||||||
credential_uuid=test_credential.uuid,
|
credential_uuid=test_credential.uuid,
|
||||||
key=extra_token,
|
|
||||||
host="other.host:4401",
|
host="other.host:4401",
|
||||||
ip="192.168.1.1",
|
ip="192.168.1.1",
|
||||||
user_agent="other-agent",
|
user_agent="other-agent",
|
||||||
@@ -1287,7 +1248,7 @@ class TestAdminSessions:
|
|||||||
):
|
):
|
||||||
"""Deleting non-existent session should fail."""
|
"""Deleting non-existent session should fail."""
|
||||||
# Use a valid format but non-existent key
|
# Use a valid format but non-existent key
|
||||||
fake_token = _create_token()
|
fake_token = secrets.token_urlsafe(12)
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/users/{test_user.uuid}/sessions/{fake_token}",
|
f"/auth/api/admin/orgs/{test_org.uuid}/users/{test_user.uuid}/sessions/{fake_token}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
@@ -1391,15 +1352,11 @@ class TestAdminPermissions:
|
|||||||
):
|
):
|
||||||
"""Admin should be able to update a permission."""
|
"""Admin should be able to update a permission."""
|
||||||
# Create permission first
|
# Create permission first
|
||||||
import uuid7
|
perm = Permission.create(scope="test:updateable", display_name="Updateable")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="test:updateable", display_name="Updateable"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
create_permission(perm)
|
||||||
|
|
||||||
response = await client.patch(
|
response = await client.patch(
|
||||||
"/auth/api/admin/permission?permission_id=test:updateable&display_name=Updated%20Name",
|
f"/auth/api/admin/permission?permission_uuid={perm.uuid}&display_name=Updated%20Name",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -1412,15 +1369,11 @@ class TestAdminPermissions:
|
|||||||
):
|
):
|
||||||
"""Updating permission with empty name should fail."""
|
"""Updating permission with empty name should fail."""
|
||||||
# Create permission first
|
# Create permission first
|
||||||
import uuid7
|
perm = Permission.create(scope="test:perm", display_name="Test Perm")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="test:perm", display_name="Test Perm"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
create_permission(perm)
|
||||||
|
|
||||||
response = await client.patch(
|
response = await client.patch(
|
||||||
"/auth/api/admin/permission?permission_id=test:perm&display_name=",
|
f"/auth/api/admin/permission?permission_uuid={perm.uuid}&display_name=",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 400
|
||||||
@@ -1428,47 +1381,32 @@ class TestAdminPermissions:
|
|||||||
assert "display_name is required" in data["detail"]
|
assert "display_name is required" in data["detail"]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_rename_permission(
|
async def test_update_permission_scope(
|
||||||
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
||||||
):
|
):
|
||||||
"""Admin should be able to rename a permission."""
|
"""Admin should be able to update a permission's scope via PATCH."""
|
||||||
# Create permission first
|
# Create permission first
|
||||||
import uuid7
|
perm = Permission.create(scope="test:renameable2", display_name="Renameable")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="test:renameable2", display_name="Renameable"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
create_permission(perm)
|
||||||
|
|
||||||
response = await client.post(
|
response = await client.patch(
|
||||||
"/auth/api/admin/permission/rename",
|
f"/auth/api/admin/permission?permission_uuid={perm.uuid}&scope=test:renamed2",
|
||||||
json={"old_scope": "test:renameable2", "new_scope": "test:renamed2"},
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_rename_permission_missing_ids(
|
async def test_update_permission_auth_admin_scope_fails(
|
||||||
self, client: httpx.AsyncClient, session_token: str
|
self, client: httpx.AsyncClient, session_token: str
|
||||||
):
|
):
|
||||||
"""Renaming permission without IDs should fail."""
|
"""Cannot change the auth:admin permission scope."""
|
||||||
response = await client.post(
|
# Get the auth:admin permission
|
||||||
"/auth/api/admin/permission/rename",
|
|
||||||
json={},
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
|
||||||
)
|
|
||||||
assert response.status_code == 400
|
|
||||||
data = response.json()
|
|
||||||
assert "required" in data["detail"]
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
perms = list(db.data().permissions.values())
|
||||||
async def test_rename_permission_auth_admin_fails(
|
admin_perm = next(p for p in perms if p.scope == "auth:admin")
|
||||||
self, client: httpx.AsyncClient, session_token: str
|
|
||||||
):
|
response = await client.patch(
|
||||||
"""Cannot rename the auth:admin permission."""
|
f"/auth/api/admin/permission?permission_uuid={admin_perm.uuid}&scope=auth:superadmin",
|
||||||
response = await client.post(
|
|
||||||
"/auth/api/admin/permission/rename",
|
|
||||||
json={"old_id": "auth:admin", "new_id": "auth:superadmin"},
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 400
|
||||||
@@ -1476,24 +1414,15 @@ class TestAdminPermissions:
|
|||||||
assert "Cannot rename the master admin" in data["detail"]
|
assert "Cannot rename the master admin" in data["detail"]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_rename_permission_with_display_name(
|
async def test_update_permission_scope_and_display_name(
|
||||||
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
||||||
):
|
):
|
||||||
"""Renaming permission can also update display name."""
|
"""Updating permission can change scope and display name together."""
|
||||||
import uuid7
|
perm = Permission.create(scope="test:rename:withname", display_name="Old Name")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="test:rename:withname", display_name="Old Name"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
create_permission(perm)
|
||||||
|
|
||||||
response = await client.post(
|
response = await client.patch(
|
||||||
"/auth/api/admin/permission/rename",
|
f"/auth/api/admin/permission?permission_uuid={perm.uuid}&scope=test:renamed:withname&display_name=New%20Display%20Name",
|
||||||
json={
|
|
||||||
"old_scope": "test:rename:withname",
|
|
||||||
"new_scope": "test:renamed:withname",
|
|
||||||
"display_name": "New Display Name",
|
|
||||||
},
|
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -1504,15 +1433,11 @@ class TestAdminPermissions:
|
|||||||
):
|
):
|
||||||
"""Admin should be able to delete a permission."""
|
"""Admin should be able to delete a permission."""
|
||||||
# Create permission first
|
# Create permission first
|
||||||
import uuid7
|
perm = Permission.create(scope="test:deleteable", display_name="Deleteable")
|
||||||
|
|
||||||
perm = Permission(
|
|
||||||
uuid=uuid7.create(), scope="test:deleteable", display_name="Deleteable"
|
|
||||||
)
|
|
||||||
create_permission(perm)
|
create_permission(perm)
|
||||||
|
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
"/auth/api/admin/permission?permission_id=test:deleteable",
|
f"/auth/api/admin/permission?permission_uuid={perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -1524,8 +1449,13 @@ class TestAdminPermissions:
|
|||||||
self, client: httpx.AsyncClient, session_token: str
|
self, client: httpx.AsyncClient, session_token: str
|
||||||
):
|
):
|
||||||
"""Cannot delete the only auth:admin permission (would lock out admin)."""
|
"""Cannot delete the only auth:admin permission (would lock out admin)."""
|
||||||
|
# Get the auth:admin permission
|
||||||
|
|
||||||
|
perms = list(db.data().permissions.values())
|
||||||
|
admin_perm = next(p for p in perms if p.scope == "auth:admin")
|
||||||
|
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
"/auth/api/admin/permission?permission_id=auth:admin",
|
f"/auth/api/admin/permission?permission_uuid={admin_perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 400
|
||||||
@@ -1537,19 +1467,21 @@ class TestAdminPermissions:
|
|||||||
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
||||||
):
|
):
|
||||||
"""Can delete an auth:admin permission if another accessible one exists."""
|
"""Can delete an auth:admin permission if another accessible one exists."""
|
||||||
import uuid7
|
|
||||||
|
|
||||||
from paskia.db import Permission
|
|
||||||
|
|
||||||
# Create a second auth:admin permission (no domain restriction)
|
# Create a second auth:admin permission (no domain restriction)
|
||||||
perm2 = Permission(
|
perm2 = Permission.create(scope="auth:admin", display_name="Secondary Admin")
|
||||||
uuid=uuid7.create(), scope="auth:admin", display_name="Secondary Admin"
|
|
||||||
)
|
|
||||||
create_permission(perm2)
|
create_permission(perm2)
|
||||||
|
|
||||||
|
# Get the original auth:admin permission (the one created in setup)
|
||||||
|
|
||||||
|
perms = list(db.data().permissions.values())
|
||||||
|
admin_perms = [p for p in perms if p.scope == "auth:admin"]
|
||||||
|
# Delete the first one (not the one we just created)
|
||||||
|
original_admin_perm = next(p for p in admin_perms if p.uuid != perm2.uuid)
|
||||||
|
|
||||||
# Now we can delete the original one
|
# Now we can delete the original one
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
"/auth/api/admin/permission?permission_id=auth:admin",
|
f"/auth/api/admin/permission?permission_uuid={original_admin_perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
@@ -1561,13 +1493,9 @@ class TestAdminPermissions:
|
|||||||
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
self, client: httpx.AsyncClient, session_token: str, test_db: DB
|
||||||
):
|
):
|
||||||
"""Cannot delete auth:admin if remaining one has mismatched domain."""
|
"""Cannot delete auth:admin if remaining one has mismatched domain."""
|
||||||
import uuid7
|
|
||||||
|
|
||||||
from paskia.db import Permission
|
|
||||||
|
|
||||||
# Create a second auth:admin permission with a different domain
|
# Create a second auth:admin permission with a different domain
|
||||||
perm2 = Permission(
|
perm2 = Permission.create(
|
||||||
uuid=uuid7.create(),
|
|
||||||
scope="auth:admin",
|
scope="auth:admin",
|
||||||
display_name="Other Domain Admin",
|
display_name="Other Domain Admin",
|
||||||
domain="other.example.com",
|
domain="other.example.com",
|
||||||
@@ -1575,8 +1503,14 @@ class TestAdminPermissions:
|
|||||||
create_permission(perm2)
|
create_permission(perm2)
|
||||||
|
|
||||||
# Cannot delete the original one because the remaining one is not accessible
|
# Cannot delete the original one because the remaining one is not accessible
|
||||||
|
# Get the original auth:admin permission
|
||||||
|
|
||||||
|
perms = list(db.data().permissions.values())
|
||||||
|
admin_perms = [p for p in perms if p.scope == "auth:admin" and p.domain is None]
|
||||||
|
original_admin_perm = admin_perms[0] # The one without domain
|
||||||
|
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
"/auth/api/admin/permission?permission_id=auth:admin",
|
f"/auth/api/admin/permission?permission_uuid={original_admin_perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 400
|
||||||
@@ -1588,8 +1522,11 @@ class TestAdminPermissions:
|
|||||||
self, client: httpx.AsyncClient, session_token: str, test_org
|
self, client: httpx.AsyncClient, session_token: str, test_org
|
||||||
):
|
):
|
||||||
"""Cannot remove auth:admin permission from your own organization."""
|
"""Cannot remove auth:admin permission from your own organization."""
|
||||||
|
admin_perm = next(
|
||||||
|
p for p in db.data().permissions.values() if p.scope == "auth:admin"
|
||||||
|
)
|
||||||
response = await client.delete(
|
response = await client.delete(
|
||||||
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_id=auth:admin",
|
f"/auth/api/admin/orgs/{test_org.uuid}/permission?permission_uuid={admin_perm.uuid}",
|
||||||
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
headers={**auth_headers(session_token), "Host": "localhost:4401"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 400
|
||||||
|
|||||||
+43
-48
@@ -10,12 +10,15 @@ These tests cover:
|
|||||||
- /auth/api/set-session - Set session from bearer token
|
- /auth/api/set-session - Set session from bearer token
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from datetime import datetime, timezone
|
import secrets
|
||||||
|
from datetime import UTC, datetime, timedelta
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
from paskia.authsession import EXPIRES
|
||||||
from paskia.db import create_session, delete_session
|
from paskia.db import create_session, delete_session
|
||||||
|
from paskia.util.passphrase import generate
|
||||||
from tests.conftest import auth_headers
|
from tests.conftest import auth_headers
|
||||||
|
|
||||||
|
|
||||||
@@ -76,7 +79,9 @@ class TestValidateEndpoint:
|
|||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert data["valid"] is True
|
assert data["valid"] is True
|
||||||
assert "user_uuid" in data
|
assert "ctx" in data
|
||||||
|
assert "user" in data["ctx"]
|
||||||
|
assert "uuid" in data["ctx"]["user"]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_validate_with_permission_check(
|
async def test_validate_with_permission_check(
|
||||||
@@ -243,9 +248,9 @@ class TestUserInfoEndpoint:
|
|||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert "user" in data
|
assert "ctx" in data
|
||||||
assert data["user"]["user_uuid"] == str(test_user.uuid)
|
assert data["ctx"]["user"]["uuid"] == str(test_user.uuid)
|
||||||
assert data["user"]["user_name"] == test_user.display_name
|
assert data["ctx"]["user"]["display_name"] == test_user.display_name
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_user_info_includes_credentials(
|
async def test_user_info_includes_credentials(
|
||||||
@@ -286,7 +291,8 @@ class TestUserInfoEndpoint:
|
|||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert "permissions" in data
|
assert "ctx" in data
|
||||||
|
assert "permissions" in data["ctx"]
|
||||||
|
|
||||||
|
|
||||||
class TestSetSessionEndpoint:
|
class TestSetSessionEndpoint:
|
||||||
@@ -314,7 +320,7 @@ class TestSetSessionEndpoint:
|
|||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert "user_uuid" in data
|
assert "user" in data
|
||||||
# Check that Set-Cookie header is present
|
# Check that Set-Cookie header is present
|
||||||
assert "set-cookie" in response.headers
|
assert "set-cookie" in response.headers
|
||||||
|
|
||||||
@@ -392,47 +398,43 @@ class TestForwardAuthHtmlResponse:
|
|||||||
assert data["auth"]["mode"] == "login"
|
assert data["auth"]["mode"] == "login"
|
||||||
|
|
||||||
|
|
||||||
class TestUserInfoWithResetToken:
|
class TestTokenInfoEndpoint:
|
||||||
"""Tests for user-info endpoint with reset tokens"""
|
"""Tests for token-info endpoint with reset tokens"""
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_user_info_with_invalid_reset_token(self, client: httpx.AsyncClient):
|
async def test_token_info_with_invalid_token(self, client: httpx.AsyncClient):
|
||||||
"""User info with invalid reset token format should return 401."""
|
"""Token info with invalid token format should return 400."""
|
||||||
# Invalid format - not a well-formed passphrase (wrong separator)
|
response = await client.get(
|
||||||
response = await client.post(
|
"/auth/api/token-info",
|
||||||
"/auth/api/user-info?reset=invalid-token-format",
|
headers={"Authorization": "Bearer invalid-token-format"},
|
||||||
)
|
)
|
||||||
# Invalid format raises ValueError which gets converted to 401 HTTPException
|
assert response.status_code == 400
|
||||||
assert response.status_code == 401
|
|
||||||
data = response.json()
|
|
||||||
assert "Invalid reset token" in data["detail"]
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_user_info_with_nonexistent_reset_token(
|
async def test_token_info_with_nonexistent_token(self, client: httpx.AsyncClient):
|
||||||
self, client: httpx.AsyncClient
|
"""Token info with well-formed but non-existent token should return 401."""
|
||||||
):
|
|
||||||
"""User info with well-formed but non-existent reset token should return 401."""
|
|
||||||
# We need a well-formed passphrase that doesn't exist in DB
|
|
||||||
from paskia.util.passphrase import generate
|
|
||||||
|
|
||||||
fake_token = generate() # Generates a well-formed token
|
fake_token = generate()
|
||||||
response = await client.post(
|
response = await client.get(
|
||||||
f"/auth/api/user-info?reset={fake_token}",
|
"/auth/api/token-info",
|
||||||
|
headers={"Authorization": f"Bearer {fake_token}"},
|
||||||
)
|
)
|
||||||
# Should return 401 for non-existent token
|
|
||||||
assert response.status_code == 401
|
assert response.status_code == 401
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_user_info_with_valid_reset_token(
|
async def test_token_info_with_valid_token(
|
||||||
self, client: httpx.AsyncClient, reset_token: str, test_user
|
self, client: httpx.AsyncClient, reset_token: str, test_user
|
||||||
):
|
):
|
||||||
"""User info with valid reset token should return minimal user info."""
|
"""Token info with valid reset token should return token type and display name."""
|
||||||
response = await client.post(
|
response = await client.get(
|
||||||
f"/auth/api/user-info?reset={reset_token}",
|
"/auth/api/token-info",
|
||||||
|
headers={"Authorization": f"Bearer {reset_token}"},
|
||||||
)
|
)
|
||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert "user" in data
|
assert "token_type" in data
|
||||||
|
assert "display_name" in data
|
||||||
|
assert data["display_name"] == test_user.display_name
|
||||||
|
|
||||||
|
|
||||||
class TestSetSessionErrors:
|
class TestSetSessionErrors:
|
||||||
@@ -442,7 +444,7 @@ class TestSetSessionErrors:
|
|||||||
async def test_set_session_with_invalid_bearer_token(
|
async def test_set_session_with_invalid_bearer_token(
|
||||||
self, client: httpx.AsyncClient
|
self, client: httpx.AsyncClient
|
||||||
):
|
):
|
||||||
"""Set session with invalid (malformed) bearer token should return 400."""
|
"""Set session with invalid (malformed) bearer token should return 401."""
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
"/auth/api/set-session",
|
"/auth/api/set-session",
|
||||||
headers={
|
headers={
|
||||||
@@ -450,8 +452,8 @@ class TestSetSessionErrors:
|
|||||||
"Host": "localhost:4401",
|
"Host": "localhost:4401",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
# Invalid token format returns 400
|
# Invalid token returns 401 (session not found)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 401
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_set_session_with_nonexistent_token(self, client: httpx.AsyncClient):
|
async def test_set_session_with_nonexistent_token(self, client: httpx.AsyncClient):
|
||||||
@@ -465,8 +467,8 @@ class TestSetSessionErrors:
|
|||||||
"Host": "localhost:4401",
|
"Host": "localhost:4401",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
# Non-existent session returns 400 (ValueError -> 400)
|
# Non-existent session returns 401 (session expired)
|
||||||
assert response.status_code == 400
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
class TestValidateSessionRefresh:
|
class TestValidateSessionRefresh:
|
||||||
@@ -499,10 +501,9 @@ class TestValidateSessionRefresh:
|
|||||||
self, client: httpx.AsyncClient, test_db
|
self, client: httpx.AsyncClient, test_db
|
||||||
):
|
):
|
||||||
"""Validate should handle session expiry during refresh attempt."""
|
"""Validate should handle session expiry during refresh attempt."""
|
||||||
from paskia.db.operations import _create_token
|
|
||||||
|
|
||||||
# Create a token but don't create a session for it
|
# Create a token but don't create a session for it
|
||||||
token = _create_token()
|
token = secrets.token_urlsafe(12)
|
||||||
response = await client.post(
|
response = await client.post(
|
||||||
"/auth/api/validate",
|
"/auth/api/validate",
|
||||||
headers={**auth_headers(token), "Host": "localhost:4401"},
|
headers={**auth_headers(token), "Host": "localhost:4401"},
|
||||||
@@ -519,18 +520,12 @@ class TestValidateSessionRefresh:
|
|||||||
test_credential,
|
test_credential,
|
||||||
):
|
):
|
||||||
"""Validate should return 401 if session disappears during refresh."""
|
"""Validate should return 401 if session disappears during refresh."""
|
||||||
from datetime import timedelta
|
|
||||||
|
|
||||||
from paskia.authsession import EXPIRES
|
|
||||||
from paskia.db.operations import _create_token
|
|
||||||
|
|
||||||
# Create a session with an old expiry time to trigger refresh
|
# Create a session with an old expiry time to trigger refresh
|
||||||
token = _create_token()
|
old_expiry = datetime.now(UTC) + EXPIRES - timedelta(minutes=10)
|
||||||
old_expiry = datetime.now(timezone.utc) + EXPIRES - timedelta(minutes=10)
|
token = create_session(
|
||||||
create_session(
|
|
||||||
user_uuid=test_user.uuid,
|
user_uuid=test_user.uuid,
|
||||||
credential_uuid=test_credential.uuid,
|
credential_uuid=test_credential.uuid,
|
||||||
key=token,
|
|
||||||
host="localhost:4401",
|
host="localhost:4401",
|
||||||
ip="127.0.0.1",
|
ip="127.0.0.1",
|
||||||
user_agent="pytest",
|
user_agent="pytest",
|
||||||
|
|||||||
Reference in New Issue
Block a user