Compare commits

..

3 Commits

Author SHA1 Message Date
Sven van Ginkel
f0f1f7985c feat: persist view preferences and language to user settings (#1831)
Co-authored-by: ChangkeunJ <reiot92@gmail.com>
2026-09-16 20:29:36 -04:00
Luís Palma
50f6fc075d Merge commit from fork
* fix: make first-user bootstrap atomic

* add tests

---------

Co-authored-by: henrygd <hank@henrygd.me>
2026-09-16 20:13:32 -04:00
henrygd
a0bf338796 hub: raise max batch requests and lower max batch body size 2026-09-16 19:39:07 -04:00
11 changed files with 489 additions and 72 deletions

View File

@@ -6,12 +6,16 @@ import (
"fmt"
"io"
"net/http"
"net/http/httptest"
"sort"
"testing"
"time"
beszelTests "github.com/henrygd/beszel/internal/tests"
"github.com/henrygd/beszel/internal/migrations"
"github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/core"
pbTests "github.com/pocketbase/pocketbase/tests"
"github.com/stretchr/testify/require"
@@ -26,6 +30,59 @@ func jsonReader(v any) io.Reader {
return bytes.NewReader(data)
}
type gatedReader struct {
data []byte
started chan struct{}
release chan struct{}
offset int
}
func (r *gatedReader) Read(p []byte) (int, error) {
if r.offset == 0 {
close(r.started)
<-r.release
}
if r.offset >= len(r.data) {
return 0, io.EOF
}
n := copy(p, r.data[r.offset:])
r.offset += n
return n, nil
}
func firstUserTestMux(t *testing.T) (*beszelTests.TestHub, http.Handler) {
t.Helper()
hub, err := beszelTests.NewTestHub(t.TempDir())
require.NoError(t, err)
_ = hub.StartHub()
router, err := apis.NewRouter(hub.TestApp)
require.NoError(t, err)
serveEvent := &core.ServeEvent{App: hub.TestApp, Router: router}
var handler http.Handler
err = hub.TestApp.OnServe().Trigger(serveEvent, func(e *core.ServeEvent) error {
var buildErr error
handler, buildErr = e.Router.BuildMux()
return buildErr
})
require.NoError(t, err)
require.NotNil(t, handler)
return hub, handler
}
func postFirstUser(handler http.Handler, email string) *httptest.ResponseRecorder {
body, _ := json.Marshal(map[string]string{
"email": email,
"password": "password123",
})
req := httptest.NewRequest(http.MethodPost, "/api/beszel/create-user", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
recorder := httptest.NewRecorder()
handler.ServeHTTP(recorder, req)
return recorder
}
func TestApiRoutesAuthentication(t *testing.T) {
hub, user := beszelTests.GetHubWithUser(t)
defer hub.Cleanup()
@@ -789,6 +846,87 @@ func TestFirstUserCreation(t *testing.T) {
})
}
func TestFirstUserBootstrapAtomicity(t *testing.T) {
t.Run("concurrent complete requests produce exactly one winner", func(t *testing.T) {
hub, handler := firstUserTestMux(t)
defer hub.Cleanup()
start := make(chan struct{})
statuses := make(chan int, 2)
for _, email := range []string{"first@example.com", "second@example.com"} {
go func(email string) {
<-start
statuses <- postFirstUser(handler, email).Code
}(email)
}
close(start)
got := []int{<-statuses, <-statuses}
sort.Ints(got)
require.Equal(t, []int{http.StatusOK, http.StatusForbidden}, got)
users, err := hub.FindAllRecords("users")
require.NoError(t, err)
require.Len(t, users, 1)
superusers, err := hub.FindAllRecords(core.CollectionNameSuperusers)
require.NoError(t, err)
require.Len(t, superusers, 1)
require.NotEqual(t, migrations.TempAdminEmail, superusers[0].Email())
})
t.Run("partial body cannot retain stale bootstrap authorization", func(t *testing.T) {
hub, handler := firstUserTestMux(t)
defer hub.Cleanup()
body, err := json.Marshal(map[string]string{
"email": "parked@example.com",
"password": "password123",
})
require.NoError(t, err)
gated := &gatedReader{
data: body,
started: make(chan struct{}),
release: make(chan struct{}),
}
parkedRequest := httptest.NewRequest(http.MethodPost, "/api/beszel/create-user", gated)
parkedRequest.Header.Set("Content-Type", "application/json")
parkedRecorder := httptest.NewRecorder()
parkedDone := make(chan struct{})
go func() {
handler.ServeHTTP(parkedRecorder, parkedRequest)
close(parkedDone)
}()
select {
case <-gated.started:
case <-time.After(2 * time.Second):
t.Fatal("parked request did not begin reading its body")
}
operatorRecorder := postFirstUser(handler, "operator@example.com")
require.Equal(t, http.StatusOK, operatorRecorder.Code)
lateRecorder := postFirstUser(handler, "late@example.com")
require.Equal(t, http.StatusForbidden, lateRecorder.Code)
close(gated.release)
select {
case <-parkedDone:
case <-time.After(2 * time.Second):
t.Fatal("parked request did not finish")
}
require.Equal(t, http.StatusForbidden, parkedRecorder.Code)
users, err := hub.FindAllRecords("users")
require.NoError(t, err)
require.Len(t, users, 1)
require.Equal(t, "operator@example.com", users[0].Email())
superusers, err := hub.FindAllRecords(core.CollectionNameSuperusers)
require.NoError(t, err)
require.Len(t, superusers, 1)
require.Equal(t, "operator@example.com", superusers[0].Email())
})
}
func TestCreateUserEndpointAvailability(t *testing.T) {
t.Run("CreateUserEndpoint available when no users exist", func(t *testing.T) {
hub, _ := beszelTests.NewTestHub(t.TempDir())

View File

@@ -122,6 +122,8 @@ func (h *Hub) initialize(app core.App) error {
settings := app.Settings()
// batch requests (for alerts)
settings.Batch.Enabled = true
settings.Batch.MaxRequests = 100
settings.Batch.MaxBodySize = 1 << 20 // 1 MiB
// set URL if APP_URL env is set
if appURL, isSet := utils.GetEnv("APP_URL"); isSet {
h.appURL = appURL

View File

@@ -63,7 +63,7 @@ export default function SettingsProfilePage({ userSettings }: { userSettings: Us
<Label className="block" htmlFor="lang">
<Trans>Preferred Language</Trans>
</Label>
<Select value={i18n.locale} onValueChange={(lang: string) => dynamicActivate(lang)}>
<Select name="lang" value={i18n.locale} onValueChange={(lang: string) => dynamicActivate(lang)}>
<SelectTrigger id="lang">
<SelectValue />
</SelectTrigger>

View File

@@ -14,7 +14,7 @@ import { lazy, useEffect } from "react"
import { $router } from "@/components/router.tsx"
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card.tsx"
import { toast } from "@/components/ui/use-toast.ts"
import { pb } from "@/lib/api"
import { saveUserSettings } from "@/lib/api"
import { $userSettings } from "@/lib/stores.ts"
import type { UserSettings } from "@/types"
import { Separator } from "../../ui/separator"
@@ -36,24 +36,13 @@ const HeartbeatSettings = lazy(heartbeatSettingsImport)
export async function saveSettings(newSettings: Partial<UserSettings>) {
try {
// get fresh copy of settings
const req = await pb.collection("user_settings").getFirstListItem("", {
fields: "id,settings",
})
// update user settings
const updatedSettings = await pb.collection("user_settings").update(req.id, {
settings: {
...req.settings,
...newSettings,
},
})
$userSettings.set(updatedSettings.settings)
await saveUserSettings(newSettings)
toast({
title: t`Settings saved`,
description: t`Your user settings have been updated.`,
})
} catch (e) {
// console.error('update settings', e)
console.error("save settings", e)
toast({
title: t`Failed to save settings`,
description: t`Check logs for more details.`,

View File

@@ -1,9 +1,9 @@
import { useStore } from "@nanostores/react"
import { getPagePath } from "@nanostores/router"
import { subscribeKeys } from "nanostores"
import { useEffect, useMemo, useRef, useState } from "react"
import { useCallback, useEffect, useMemo, useRef, useState } from "react"
import { useContainerChartConfigs } from "@/components/charts/hooks"
import { pb } from "@/lib/api"
import { pb, queueUserSettings } from "@/lib/api"
import { SystemStatus } from "@/lib/enums"
import {
$allSystemsById,
@@ -15,7 +15,7 @@ import {
$systems,
$userSettings,
} from "@/lib/stores"
import { chartTimeData, listen, parseSemVer, useBrowserStorage } from "@/lib/utils"
import { chartTimeData, listen, parseSemVer } from "@/lib/utils"
import type {
ChartData,
ContainerStatsRecord,
@@ -35,8 +35,42 @@ export function useSystemData(id: string) {
const systems = useStore($systems)
const chartTime = useStore($chartTime)
const maxValues = useStore($maxValues)
const [grid, setGrid] = useBrowserStorage("grid", true)
const [displayMode, setDisplayMode] = useBrowserStorage<"default" | "tabs">("displayMode", "default")
const [grid, _setGrid] = useState<boolean>(
() => $userSettings.get().grid ?? JSON.parse(localStorage.getItem("besz-grid") ?? "null") ?? true
)
const [displayMode, _setDisplayMode] = useState<"default" | "tabs">(
() =>
$userSettings.get().displayMode ??
(JSON.parse(localStorage.getItem("besz-displayMode") || "null") as "default" | "tabs" | null) ??
"default"
)
const applied = useRef(new Set<string>())
useEffect(() => {
return subscribeKeys($userSettings, ["grid", "displayMode"], (vals) => {
if (!applied.current.has("grid") && vals.grid !== undefined) {
applied.current.add("grid")
_setGrid(vals.grid)
}
if (!applied.current.has("displayMode") && vals.displayMode !== undefined) {
applied.current.add("displayMode")
_setDisplayMode(vals.displayMode)
}
})
}, [])
const setGrid = useCallback((v: boolean) => {
_setGrid(v)
localStorage.setItem("besz-grid", JSON.stringify(v))
$userSettings.setKey("grid", v)
queueUserSettings({ grid: v })
}, [])
const setDisplayMode = useCallback((v: "default" | "tabs") => {
_setDisplayMode(v)
localStorage.setItem("besz-displayMode", JSON.stringify(v))
$userSettings.setKey("displayMode", v)
queueUserSettings({ displayMode: v })
}, [])
const [activeTab, setActiveTabRaw] = useState("core")
const [mountedTabs, setMountedTabs] = useState(() => new Set<string>(["core"]))
const tabsRef = useRef<string[]>(["core", "disk"])

View File

@@ -1,5 +1,6 @@
import { Trans, useLingui } from "@lingui/react/macro"
import { useStore } from "@nanostores/react"
import { subscribeKeys } from "nanostores"
import { getPagePath } from "@nanostores/router"
import {
type ColumnDef,
@@ -26,7 +27,7 @@ import {
Settings2Icon,
XIcon,
} from "lucide-react"
import { memo, useEffect, useMemo, useRef, useState } from "react"
import { memo, useCallback, useEffect, useMemo, useRef, useState } from "react"
import { Button } from "@/components/ui/button"
import {
DropdownMenu,
@@ -42,8 +43,9 @@ import {
import { Input } from "@/components/ui/input"
import { TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"
import { SystemStatus } from "@/lib/enums"
import { $downSystems, $pausedSystems, $systems, $upSystems } from "@/lib/stores"
import { cn, runOnce, useBrowserStorage } from "@/lib/utils"
import { queueUserSettings } from "@/lib/api"
import { $downSystems, $pausedSystems, $systems, $upSystems, $userSettings } from "@/lib/stores"
import { cn, runOnce } from "@/lib/utils"
import type { SystemRecord } from "@/types"
import AlertButton from "../alerts/alert-button"
import { $router, Link } from "../router"
@@ -62,14 +64,83 @@ export default function SystemsTable() {
const pausedSystems = $pausedSystems.get()
const { i18n, t } = useLingui()
const [filter, setFilter] = useState<string>("")
const [statusFilter, setStatusFilter] = useState<StatusFilter>("all")
const [sorting, setSorting] = useBrowserStorage<SortingState>(
"sortMode",
[{ id: "system", desc: false }],
sessionStorage
const [statusFilter, setStatusFilter] = useState<StatusFilter>(
() =>
$userSettings.get().statusFilter ??
(JSON.parse(localStorage.getItem("besz-statusFilter") || "null") as StatusFilter | null) ??
"all"
)
const [sorting, setSorting] = useState<SortingState>(
() =>
$userSettings.get().sortMode ??
JSON.parse(sessionStorage.getItem("besz-sortMode") || "null") ?? [{ id: "system", desc: false }]
)
const [columnFilters, setColumnFilters] = useState<ColumnFiltersState>([])
const [columnVisibility, setColumnVisibility] = useBrowserStorage<VisibilityState>("cols", {})
const [columnVisibility, setColumnVisibility] = useState<VisibilityState>(
() => $userSettings.get().cols ?? JSON.parse(localStorage.getItem("besz-cols") || "{}")
)
// Apply settings from server once they load (handles incognito / new devices)
const applied = useRef(new Set<string>())
useEffect(() => {
return subscribeKeys($userSettings, ["cols", "statusFilter", "viewMode", "sortMode"], (vals) => {
if (!applied.current.has("cols") && vals.cols !== undefined) {
applied.current.add("cols")
setColumnVisibility(vals.cols)
}
if (!applied.current.has("statusFilter") && vals.statusFilter !== undefined) {
applied.current.add("statusFilter")
setStatusFilter(vals.statusFilter)
}
if (!applied.current.has("viewMode") && vals.viewMode !== undefined) {
applied.current.add("viewMode")
setViewMode(vals.viewMode)
}
if (!applied.current.has("sortMode") && vals.sortMode !== undefined) {
applied.current.add("sortMode")
setSorting(vals.sortMode)
}
})
}, [])
const handleColumnVisibilityChange = useCallback(
(updater: VisibilityState | ((prev: VisibilityState) => VisibilityState)) => {
setColumnVisibility((prev) => {
const next = typeof updater === "function" ? updater(prev) : updater
localStorage.setItem("besz-cols", JSON.stringify(next))
$userSettings.setKey("cols", next)
queueUserSettings({ cols: next })
return next
})
},
[]
)
const handleStatusFilterChange = useCallback((value: string) => {
const next = value as StatusFilter
setStatusFilter(next)
localStorage.setItem("besz-statusFilter", JSON.stringify(next))
$userSettings.setKey("statusFilter", next)
queueUserSettings({ statusFilter: next })
}, [])
const handleViewModeChange = useCallback((view: string) => {
const next = view as ViewMode
setViewMode(next)
localStorage.setItem("besz-viewMode", JSON.stringify(next))
$userSettings.setKey("viewMode", next)
queueUserSettings({ viewMode: next })
}, [])
const handleSortingChange = useCallback((updater: SortingState | ((prev: SortingState) => SortingState)) => {
setSorting((prev) => {
const next = typeof updater === "function" ? updater(prev) : updater
sessionStorage.setItem("besz-sortMode", JSON.stringify(next))
$userSettings.setKey("sortMode", next)
queueUserSettings({ sortMode: next })
return next
})
}, [])
const locale = i18n.locale
@@ -87,10 +158,12 @@ export default function SystemsTable() {
return Object.values(pausedSystems) ?? []
}, [data, statusFilter])
const [viewMode, setViewMode] = useBrowserStorage<ViewMode>(
"viewMode",
// show grid view on mobile if there are less than 200 systems (looks better but table is more efficient)
window.innerWidth < 1024 && filteredData.length < 200 ? "grid" : "table"
const [viewMode, setViewMode] = useState<ViewMode>(
() =>
$userSettings.get().viewMode ??
(JSON.parse(localStorage.getItem("besz-viewMode") || "null") as ViewMode | null) ??
// show grid view on mobile if there are less than 200 systems (looks better but table is more efficient)
(window.innerWidth < 1024 && filteredData.length < 200 ? "grid" : "table")
)
useEffect(() => {
@@ -105,11 +178,11 @@ export default function SystemsTable() {
data: filteredData,
columns: columnDefs,
getCoreRowModel: getCoreRowModel(),
onSortingChange: setSorting,
onSortingChange: handleSortingChange,
getSortedRowModel: getSortedRowModel(),
onColumnFiltersChange: setColumnFilters,
getFilteredRowModel: getFilteredRowModel(),
onColumnVisibilityChange: setColumnVisibility,
onColumnVisibilityChange: handleColumnVisibilityChange,
state: {
sorting,
columnFilters,
@@ -181,11 +254,7 @@ export default function SystemsTable() {
<Trans>Layout</Trans>
</DropdownMenuLabel>
<DropdownMenuSeparator />
<DropdownMenuRadioGroup
className="px-1 pb-1"
value={viewMode}
onValueChange={(view) => setViewMode(view as ViewMode)}
>
<DropdownMenuRadioGroup className="px-1 pb-1" value={viewMode} onValueChange={handleViewModeChange}>
<DropdownMenuRadioItem value="table" onSelect={(e) => e.preventDefault()} className="gap-2">
<LayoutListIcon className="size-4" />
<Trans>Table</Trans>
@@ -206,7 +275,7 @@ export default function SystemsTable() {
<DropdownMenuRadioGroup
className="px-1 pb-1"
value={statusFilter}
onValueChange={(value) => setStatusFilter(value as StatusFilter)}
onValueChange={handleStatusFilterChange}
>
<DropdownMenuRadioItem value="all" onSelect={(e) => e.preventDefault()}>
<Trans>All Systems</Trans>
@@ -245,7 +314,9 @@ export default function SystemsTable() {
<DropdownMenuItem
onSelect={(e) => {
e.preventDefault()
setSorting([{ id: column.id, desc: sorting[0]?.id === column.id && !sorting[0]?.desc }])
handleSortingChange([
{ id: column.id, desc: sorting[0]?.id === column.id && !sorting[0]?.desc },
])
}}
key={column.id}
>

View File

@@ -2,6 +2,7 @@ import { t } from "@lingui/core/macro"
import PocketBase from "pocketbase"
import { basePath } from "@/components/router"
import { toast } from "@/components/ui/use-toast"
import { dynamicActivate, getLocale } from "@/lib/i18n"
import type { ChartTimes, UserSettings } from "@/types"
import { $alerts, $allSystemsById, $allSystemsByName, $userSettings } from "./stores"
import { chartTimeData, debounce } from "./utils"
@@ -52,11 +53,45 @@ export function logOut() {
pb.realtime.unsubscribe()
}
/** Save a partial update to user settings in database immediately */
export async function saveUserSettings(newSettings: Partial<UserSettings>) {
// get fresh copy of settings so concurrent changes aren't overwritten
const req = await pb.collection("user_settings").getFirstListItem("", { fields: "id,settings" })
const updatedSettings = await pb.collection("user_settings").update(req.id, {
settings: {
...req.settings,
...newSettings,
},
})
$userSettings.set(updatedSettings.settings)
}
// keys queued by queueUserSettings, flushed together in a single request so that
// two debounced saves for different keys can't race each other's read-modify-write
// and silently drop one of the changes
let queuedSettings: Partial<UserSettings> = {}
const flushQueuedSettings = debounce(() => {
const toSave = queuedSettings
queuedSettings = {}
if (Object.keys(toSave).length === 0) {
return
}
saveUserSettings(toSave).catch(console.error)
}, 1000)
/** Queue a partial user settings update, merging with any other pending keys and saving them together after a debounce window */
export function queueUserSettings(newSettings: Partial<UserSettings>) {
queuedSettings = { ...queuedSettings, ...newSettings }
flushQueuedSettings()
}
/** Fetch or create user settings in database */
export async function updateUserSettings() {
try {
const req = await pb.collection("user_settings").getFirstListItem("", { fields: "settings" })
$userSettings.set(req.settings)
dynamicActivate(req.settings.lang || getLocale())
return
} catch (e) {
console.error("get settings", e)
@@ -65,6 +100,7 @@ export async function updateUserSettings() {
try {
const createdSettings = await pb.collection("user_settings").create({ user: pb.authStore.record?.id })
$userSettings.set(createdSettings.settings)
dynamicActivate(createdSettings.settings.lang || getLocale())
} catch (e) {
console.error("create settings", e)
}

View File

@@ -121,6 +121,7 @@ const Layout = () => {
const I18nApp = () => {
useEffect(() => {
// Activate a locale so I18nProvider can mount App and load the account settings.
dynamicActivate(getLocale())
}, [])

View File

@@ -370,6 +370,13 @@ export interface UserSettings {
colorCrit?: number
hourFormat?: HourFormat
layoutWidth?: number
lang?: string
cols?: Record<string, boolean>
statusFilter?: "all" | "up" | "down" | "paused" | "pending"
viewMode?: "table" | "grid"
sortMode?: Array<{ id: string; desc: boolean }>
grid?: boolean
displayMode?: "default" | "tabs"
}
type ChartDataContainer = {

View File

@@ -2,6 +2,7 @@
package users
import (
"errors"
"log"
"net/http"
@@ -15,6 +16,8 @@ type UserManager struct {
app core.App
}
var errBootstrapUnavailable = errors.New("bootstrap unavailable")
func NewUserManager(app core.App) *UserManager {
return &UserManager{
app: app,
@@ -59,17 +62,7 @@ func (um *UserManager) InitializeUserSettings(e *core.RecordEvent) error {
// Custom API endpoint to create the first user.
// Mimics previous default behavior in PocketBase < 0.23.0 allowing user to be created through the Beszel UI.
func (um *UserManager) CreateFirstUser(e *core.RequestEvent) error {
// check that there are no users
totalUsers, err := um.app.CountRecords("users")
if err != nil || totalUsers > 0 {
return e.JSON(http.StatusForbidden, map[string]string{"err": "Forbidden"})
}
// check that there is only one superuser and the email matches the email of the superuser we set up in initial-settings.go
adminUsers, err := um.app.FindAllRecords(core.CollectionNameSuperusers)
if err != nil || len(adminUsers) != 1 || adminUsers[0].GetString("email") != migrations.TempAdminEmail {
return e.JSON(http.StatusForbidden, map[string]string{"err": "Forbidden"})
}
// create first user using supplied email and password in request body
// Consume the complete body before evaluating the one-time bootstrap state.
data := struct {
Email string `json:"email"`
Password string `json:"password"`
@@ -81,26 +74,55 @@ func (um *UserManager) CreateFirstUser(e *core.RequestEvent) error {
return e.JSON(http.StatusBadRequest, map[string]string{"err": "Bad request"})
}
collection, _ := um.app.FindCollectionByNameOrId("users")
user := core.NewRecord(collection)
user.SetEmail(data.Email)
user.SetPassword(data.Password)
user.Set("role", "admin")
user.Set("verified", true)
if err := um.app.Save(user); err != nil {
return e.JSON(http.StatusInternalServerError, map[string]string{"err": err.Error()})
}
// create superuser using the email of the first user
collection, _ = um.app.FindCollectionByNameOrId(core.CollectionNameSuperusers)
adminUser := core.NewRecord(collection)
adminUser.SetEmail(data.Email)
adminUser.SetPassword(data.Password)
if err := um.app.Save(adminUser); err != nil {
return e.JSON(http.StatusInternalServerError, map[string]string{"err": err.Error()})
}
// delete the intial superuser
if err := um.app.Delete(adminUsers[0]); err != nil {
err := um.app.RunInTransaction(func(txApp core.App) error {
totalUsers, err := txApp.CountRecords("users")
if err != nil {
return err
}
if totalUsers > 0 {
return errBootstrapUnavailable
}
adminUsers, err := txApp.FindAllRecords(core.CollectionNameSuperusers)
if err != nil {
return err
}
if len(adminUsers) != 1 || adminUsers[0].GetString("email") != migrations.TempAdminEmail {
return errBootstrapUnavailable
}
collection, err := txApp.FindCollectionByNameOrId("users")
if err != nil {
return err
}
user := core.NewRecord(collection)
user.SetEmail(data.Email)
user.SetPassword(data.Password)
user.Set("role", "admin")
user.Set("verified", true)
if err := txApp.Save(user); err != nil {
return err
}
collection, err = txApp.FindCollectionByNameOrId(core.CollectionNameSuperusers)
if err != nil {
return err
}
adminUser := core.NewRecord(collection)
adminUser.SetEmail(data.Email)
adminUser.SetPassword(data.Password)
if err := txApp.Save(adminUser); err != nil {
return err
}
return txApp.Delete(adminUsers[0])
})
if errors.Is(err, errBootstrapUnavailable) {
return e.JSON(http.StatusForbidden, map[string]string{"err": "Forbidden"})
}
if err != nil {
return e.JSON(http.StatusInternalServerError, map[string]string{"err": err.Error()})
}
return e.JSON(http.StatusOK, map[string]string{"msg": "User created"})
}

View File

@@ -0,0 +1,117 @@
//go:build testing
package users_test
import (
"errors"
"io"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/henrygd/beszel/internal/migrations"
beszelTests "github.com/henrygd/beszel/internal/tests"
"github.com/henrygd/beszel/internal/users"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/tools/router"
"github.com/stretchr/testify/require"
)
type blockedBody struct {
io.Reader
entered chan struct{}
resume chan struct{}
}
func (b *blockedBody) Read(p []byte) (int, error) {
if b.entered != nil {
close(b.entered)
b.entered = nil
<-b.resume
}
return b.Reader.Read(p)
}
func TestCreateFirstUserAtomic(t *testing.T) {
for _, scenario := range []string{"parked body", "concurrent requests", "rollback"} {
t.Run(scenario, func(t *testing.T) {
h, err := beszelTests.NewTestHub(t.TempDir())
require.NoError(t, err)
defer h.Cleanup()
h.StartHub()
um := users.NewUserManager(h.App)
invoke := func(body io.Reader) int {
req := httptest.NewRequest("POST", "/api/beszel/create-user", body)
req.Header.Set("Content-Type", "application/json")
res := httptest.NewRecorder()
if err := um.CreateFirstUser(&core.RequestEvent{App: h.App, Event: router.Event{Request: req, Response: res}}); err != nil {
t.Error(err)
}
return res.Code
}
body := func(email string) io.Reader {
return strings.NewReader(`{"email":"` + email + `","password":"password12345"}`)
}
await := func(results <-chan int) int {
select {
case status := <-results:
return status
case <-time.After(10 * time.Second):
t.Fatal("request did not finish")
return 0
}
}
switch scenario {
case "parked body":
entered, resume := make(chan struct{}), make(chan struct{})
defer func() {
select {
case <-resume:
default:
close(resume)
}
}()
result := make(chan int, 1)
go func() { result <- invoke(&blockedBody{body("attacker@example.com"), entered, resume}) }()
select {
case <-entered:
case <-time.After(10 * time.Second):
t.Fatal("request did not reach body parsing")
}
require.Equal(t, 200, invoke(body("operator@example.com")))
close(resume)
require.Equal(t, 403, await(result))
case "concurrent requests":
start := make(chan struct{})
results := make(chan int, 2)
for _, email := range []string{"one@example.com", "two@example.com"} {
go func() { <-start; results <- invoke(body(email)) }()
}
close(start)
require.ElementsMatch(t, []int{200, 403}, []int{await(results), await(results)})
case "rollback":
hook := h.OnRecordCreate(core.CollectionNameSuperusers).BindFunc(func(e *core.RecordEvent) error {
return errors.New("injected superuser creation failure")
})
require.Equal(t, 500, invoke(body("operator@example.com")))
count, err := h.CountRecords("users")
require.NoError(t, err)
require.Zero(t, count)
admins, err := h.FindAllRecords(core.CollectionNameSuperusers)
require.NoError(t, err)
require.Len(t, admins, 1)
require.Equal(t, migrations.TempAdminEmail, admins[0].Email())
h.OnRecordCreate(core.CollectionNameSuperusers).Unbind(hook)
require.Equal(t, 200, invoke(body("operator@example.com")))
}
count, err := h.CountRecords("users")
require.NoError(t, err)
require.EqualValues(t, 1, count)
admins, err := h.FindAllRecords(core.CollectionNameSuperusers)
require.NoError(t, err)
require.Len(t, admins, 1)
require.NotEqual(t, migrations.TempAdminEmail, admins[0].Email())
})
}
}