Compare commits

...
604 Commits
Author SHA1 Message Date
Gani Georgiev 970b00011f updated go deps and bumped app version 2024-06-03 08:36:28 +03:00
Gani Georgiev 1682db5c72 updated ui/dist and go deps 2024-05-11 09:08:29 +03:00
Gani Georgiev 5acaa0a55c [#4865] fixed Firefox calendar picker grid layout 2024-05-04 13:48:46 +03:00
Gani Georgiev 8a75a9ab04 Merge branch 'master' into develop 2024-05-04 11:25:46 +03:00
Gani Georgiev c410f34089 updated go deps 2024-05-03 20:07:13 +03:00
Gani Georgiev 52a0a87cb2 Merge branch 'master' into develop 2024-05-03 19:59:09 +03:00
Gani Georgiev 2aeb37dcd0 [#4857] load the full record in the relation picker edit panel 2024-05-03 19:58:52 +03:00
Gani Georgiev 3a18121939 Merge branch 'master' into develop 2024-04-25 11:09:08 +03:00
Gani Georgiev 7a9dae7bdd updated changelog 2024-04-25 11:06:48 +03:00
Gani Georgiev dd7b06c00f Merge branch 'master' into develop 2024-04-25 10:52:54 +03:00
Gani Georgiev 950f796cbc added temp collections cache 2024-04-25 10:14:59 +03:00
Gani Georgiev 7675d2e07b Merge branch 'master' into develop 2024-04-24 23:34:44 +03:00
Gani Georgiev 2b82c36bdd updated test cases 2024-04-24 23:23:24 +03:00
Gani Georgiev 3df868f72a added extra extension length normalization 2024-04-24 23:20:47 +03:00
Gani Georgiev 4902b72247 Merge branch 'master' into develop 2024-04-24 22:12:31 +03:00
Gani Georgiev e7ebbd1343 updated go deps 2024-04-24 22:12:09 +03:00
Gani Georgiev ece62ebdf5 [#4824] updated the uploaded filename normalization to take double extensions in consideration 2024-04-24 22:00:18 +03:00
Gani Georgiev 0ea26d91e1 rollback to goreleaser-action v3 2024-04-15 09:44:02 +03:00
Gani Georgiev 7628dc5634 bumped goreleaser action version 2024-04-15 09:37:56 +03:00
Gani Georgiev 4cfabc61e6 updated changelog and ui/dist 2024-04-15 09:20:31 +03:00
Gani Georgiev 34a25f640b updated go deps 2024-04-13 16:19:33 +03:00
Gani Georgiev 4ac0954546 fixed zeroValue isArray check and bumped app version 2024-04-13 16:18:52 +03:00
Gani Georgiev 4745bb4286 fixed zeroValue isArray check 2024-04-13 16:16:21 +03:00
Gani Georgiev 7734d63e51 [#4737] fixed OAuth2 clear btn action 2024-04-13 15:44:20 +03:00
Gani Georgiev 4286812a09 updated go deps 2024-04-06 11:52:44 +03:00
Gani Georgiev 8264ead4de updated changelog 2024-04-05 23:15:41 +03:00
Gani Georgiev 4dc8a10af5 added aria-expanded to the dropdown triggers 2024-04-05 23:15:11 +03:00
Gani Georgiev a9d468a863 updated ui/dist 2024-04-05 20:47:28 +03:00
Gani Georgiev ebc1ed6598 replaced btn mail template outline with border for compatability 2024-04-05 20:46:15 +03:00
Gani Georgiev c951d4bc94 updated ui/dist 2024-04-05 20:35:02 +03:00
Gani Georgiev bee8e0826a [#4707] added constrasting border to the default email template btn style 2024-04-05 20:33:56 +03:00
Gani Georgiev d5fc74d973 updated go deps 2024-04-05 20:20:03 +03:00
Gani Georgiev 63bcffb223 [#4704] fixed '~' autowildcard wrapping when the string has escaped % character 2024-04-05 20:14:28 +03:00
Gani Georgiev ac76166cb2 updated changelog 2024-03-29 21:30:42 +02:00
Gani Georgiev 37dd9c8645 vendored and trimmed the s3blob driver and updated dependencies 2024-03-29 21:19:26 +02:00
Gani Georgiev 9eb3ff5833 updated ui/dist 2024-03-28 13:33:02 +02:00
Marcel van Remmerden 9090979a8d [#4650] updated GitLab logo 2024-03-28 13:29:07 +02:00
Gani Georgiev 7ce2545b8b updated go deps 2024-03-23 13:15:57 +02:00
Gani Georgiev 31f2ba89e8 added aria-hidden attr and bumped app version 2024-03-23 13:12:46 +02:00
Gani Georgiev 0122d4f527 [#4607] fixed the keyboard-accebility of the Admin UI dropdowns 2024-03-22 20:07:01 +02:00
Gani Georgiev b596bbdc3e updated go deps 2024-03-21 10:22:32 +02:00
Gani Georgiev 04927178e5 updated backup restore message 2024-03-21 10:18:00 +02:00
Gani Georgiev 98ba003921 added done channel for the cron ticker 2024-03-20 23:55:32 +02:00
Gani Georgiev 309c4fe6fe call TestApp.ResetBootstrap as finalizer of the test OnTerminate hook 2024-03-20 23:52:26 +02:00
Gani Georgiev 03cec9a5ac [#4600] autorun migrations for the test app and call the OnTerminate hook on TestApp.Cleanup 2024-03-20 22:47:16 +02:00
Gani Georgiev 48153d4542 updated restore backup warning message and changed archive.Extract to ignore irregular files 2024-03-17 15:43:27 +02:00
Gani Georgiev be40803d31 updated security.Encrypt and security.Decrypt docs 2024-03-17 15:42:40 +02:00
Gani Georgiev a5eff395b4 [#4566] fixed JSVM routerUse() example 2024-03-15 11:45:46 +02:00
Gani Georgiev 20653ef786 bumped app version 2024-03-12 23:58:18 +02:00
Gani Georgiev 0f1b73a4f5 [#4544] implemented JSVM FormData and added support for $http.send multipart/form-data requests 2024-03-12 21:35:29 +02:00
Gani Georgiev adab0da179 [#4510] fixed godoc typos 2024-03-07 11:53:54 +02:00
Gani Georgiev e5e2519f88 [#4505] removed redundant CodeBlock component styles 2024-03-06 17:47:41 +02:00
Gani Georgiev 0afc380a11 bumped app version 2024-03-06 16:47:47 +02:00
Gani Georgiev 3551dea44a [#4500] added the field name as part of the @request.data.* relations join 2024-03-06 15:45:25 +02:00
Gani Georgiev eff09852a4 updated GitHub release action min Go version 2024-03-06 11:30:02 +02:00
Gani Georgiev 90c313cf09 updated go deps 2024-03-06 11:20:32 +02:00
Gani Georgiev 5574fe39ce updated go deps 2024-03-06 11:17:56 +02:00
Gani Georgiev 6695aba758 [#4498] fixed OnAfterApiError nil error reference 2024-03-06 11:06:39 +02:00
Gani Georgiev 1eeacf0204 [#4492] fixed admin dropdown z-index on Safari 2024-03-05 19:44:20 +02:00
Gani Georgiev 4a1736a785 restored nullifyMissingField checks 2024-03-03 00:13:48 +02:00
Gani Georgiev 186d2ed328 bumped app version 2024-03-02 18:17:03 +02:00
Gani Georgiev 35d7b5f056 updated changelog 2024-03-01 17:09:39 +02:00
Gani Georgiev bb410e7e0d [#4462] fixed Admin UI record and collection panels not reinitializing properly on browser back/forward navigation 2024-03-01 17:00:26 +02:00
Gani Georgiev 9babca5f77 [#4448] added error checks to the autogenerated Go migrations 2024-02-29 04:17:59 +02:00
Gani Georgiev b845d3dbea [#4437] initialize RecordAuthWithOAuth2Event.IsNewRecord for the OnRecordBeforeAuthWithOAuth2Request hook 2024-02-27 12:14:02 +02:00
Gani Georgiev 39d24ba897 Merge branch 'master' into develop 2024-02-26 20:03:45 +02:00
Gani Georgiev 631957fa32 =fixed changelog typo and added PR link 2024-02-26 19:57:53 +02:00
Gani Georgiev f414e70ffa Merge branch 'develop' 2024-02-26 19:10:55 +02:00
Gani Georgiev d084800c45 updated go deps and bumped ui version 2024-02-26 19:05:20 +02:00
Gani Georgiev f1a6c19309 fixed logs printer dev tests 2024-02-26 16:39:35 +02:00
Gani Georgiev 53ee5212bc [#4431] always refresh the app settings before loading the backup cron job 2024-02-26 15:01:49 +02:00
Gani Georgiev 548fce20b5 added back-relation expand limit 2024-02-25 21:06:43 +02:00
Gani Georgiev 1014c92d86 sort exported collections by type and name 2024-02-25 21:06:14 +02:00
Gani Georgiev 88c56cd539 added :each support for file and relation fields 2024-02-25 12:19:19 +02:00
Gani Georgiev a8b363ed76 normalized collections export sidebar padding and reduced the waiting time for the cron test 2024-02-24 21:35:03 +02:00
Gani Georgiev 6132fb4a03 updated collections export styles 2024-02-24 17:18:06 +02:00
Gani Georgiev 20fba0f686 moved filter autocomplete to worker 2024-02-24 13:46:16 +02:00
Gani Georgievandalxjsn 4f46222de9 [#4393] added Planning Center OAuth2 provider
Co-authored-by: alxjsn <alxjsn@sameorigin.org>
2024-02-24 08:46:22 +02:00
Gani Georgiev 4fba93e834 regenerated jsvm types and added locks for the startTimer 2024-02-21 22:42:01 +02:00
Gani Georgiev 5a715cc60a [#4394] reschedule the first cron tick to start at 00 second 2024-02-21 19:49:52 +02:00
Gani Georgiev f2ed186540 added autocomplete for the back relation keys 2024-02-19 23:13:04 +02:00
Gani Georgiev 4937acb3e2 added back relation filter reference support 2024-02-19 16:55:34 +02:00
Gani Georgiev 4743c1ce72 updated jsvm types and changelog 2024-02-17 17:14:46 +02:00
Gani Georgiev a11abef84b added @request.context field 2024-02-17 15:01:09 +02:00
Gani Georgiev 6aaf98215d hide the merge collections import btn if no schema is specified 2024-02-12 12:33:31 +02:00
Gani Georgiev 2662d875b9 removed unnecessary concat 2024-02-12 11:40:04 +02:00
Gani Georgiev 959c6b6d6c [#3403] added option to import/export a subset of collections 2024-02-12 11:38:22 +02:00
Gani Georgiev d4a2f05075 added presentable file field fallback 2024-02-11 22:31:10 +02:00
Gani Georgiev 4c14c6cccf synced with master 2024-02-11 11:08:47 +02:00
Gani Georgiev aaa6e971a3 fixed changelog typo 2024-02-11 09:37:27 +02:00
Gani Georgiev 27b6d0c505 updated jstypes and ui/dist 2024-02-10 23:28:23 +02:00
Gani Georgiev a46815ed69 merged with master 2024-02-10 15:41:39 +02:00
Gani Georgiev a5fdbeae79 manually clear all TinyMCE events on editor removal 2024-02-10 15:06:49 +02:00
Gani Georgiev 8599754e45 sync with master 2024-02-10 11:16:23 +02:00
Gani Georgiev 71141dde69 aligned healthCheckResponse struct fields 2024-02-10 11:04:59 +02:00
Gani Georgiev 388f61aed6 [#4310] allow HEAD requests to the health endpoint 2024-02-10 10:59:39 +02:00
Gani Georgiev c32f272123 [#4322] disable the JS required validations for disabled OIDC providers 2024-02-09 22:17:26 +02:00
Gani Georgiev 81ef6f1127 Merge branch 'master' into develop 2024-02-08 00:29:41 +02:00
Gani Georgiev 8f8a7c3268 fixed readme typo 2024-02-08 00:29:16 +02:00
Gani Georgiev 1b89aabf14 updated github actions 2024-02-07 21:45:55 +02:00
Gani Georgiev 5f1b2fda74 updated github action node version 2024-02-07 21:22:16 +02:00
Gani Georgiev 9f1c1c2e33 updated go deps and the min github action go version 2024-02-07 21:19:45 +02:00
Gani Georgiev 7ef118581b use the email field const 2024-02-07 21:17:39 +02:00
Gani Georgiev b7447f3e27 synced with master 2024-02-07 21:13:35 +02:00
Gani Georgiev 368af1f0fc updated the readme 2024-02-07 21:04:20 +02:00
Gani Georgiev 41aa9b189c updated changelog 2024-02-07 20:15:17 +02:00
Gani Georgiev ed9cc2f33c updated changelog 2024-02-07 20:04:20 +02:00
Gani Georgiev bada2338f7 [#2173] fixed request.auth.* initialization which caused the current authenticated user email to not being returned in the authRefresh() calls 2024-02-07 19:51:09 +02:00
Gani Georgiev 722a74994f fixed the error reporting of admin update/delete commands 2024-02-06 13:55:08 +02:00
Gani Georgiev 442b286b1d updated changelog 2024-02-05 23:10:01 +02:00
Gani Georgiev 5105612a45 renamed gcp middleware file and updated go deps 2024-02-05 17:59:31 +02:00
Gani Georgiev b9029010d9 upgraded to aws-sdk-go-v2 and added a special middleware for GCP 2024-02-05 17:26:39 +02:00
Gani Georgiev 03a3f9876e sync Admin UI collection changes across browser tabs 2024-02-03 15:39:09 +02:00
Gani Georgiev 49adba6947 added jsvm.Config.OnInit optional field 2024-02-03 13:07:37 +02:00
Gani Georgiev ef965aafbb removed unused tinymce assets 2024-02-03 11:35:43 +02:00
Gani Georgiev fa8e3d83b5 updated npm deps 2024-02-02 17:32:42 +02:00
Gani Georgiev 8402938191 fixed Admin UI vertical image preview scroll 2024-02-02 15:10:33 +02:00
Gani Georgiev 8a0eed22fa update ghupdate to use the config executable name when excluding the update note from the release notes 2024-02-02 12:48:37 +02:00
Gani Georgiev 3b6fcf265a fixed RecordUpsert.RemoveFiles godoc example 2024-02-02 09:21:41 +02:00
Gani Georgiev fb78a39161 updated readme and the thumbGenSem limit 2024-01-31 11:08:40 +02:00
Gani Georgiev 9436efb7fd fixed hideControls store reactivity check 2024-01-24 17:38:22 +02:00
Gani Georgiev 05556e7cbc updated changelog and go deps 2024-01-24 11:18:12 +02:00
Gani Georgiev 2862119c1f updated serve command error reporting 2024-01-24 11:06:49 +02:00
Gani Georgiev eaf121ead7 updated ui/dist 2024-01-23 21:28:13 +02:00
Gani Georgiev aabe820e35 fixed typos and some linter suggestions 2024-01-23 20:56:14 +02:00
Gani Georgiev 80d65a198b optimized multiple records cascade delete query 2024-01-23 20:22:51 +02:00
Gani Georgiev 3013e0299a added helper admin cmd error message in case the migrations are not initialized yet 2024-01-23 20:13:30 +02:00
Gani Georgiev 6fd2e7ab0f updated min Go and Node.js verion in CONTRIBUTING.md 2024-01-22 16:54:12 +02:00
Gani Georgiev a44a73a17c fixed unverified typos 2024-01-22 08:02:48 +02:00
Gani Georgiev bf30af393e include 0 in the auto numeric suffix field name 2024-01-21 20:57:48 +02:00
Gani Georgiev f0410a7625 [#4033] added option to duplicate fields 2024-01-21 20:22:56 +02:00
Gani Georgiev ba56623245 exported .gzip() and .bodyLimit(bytes) JSVM middlewares 2024-01-21 17:13:22 +02:00
Gani Georgiev 702b4aa1c2 Merge branch 'master' into develop 2024-01-21 15:21:22 +02:00
Gani Georgiev 3f7db19fdd remove funding.yaml 2024-01-21 12:29:54 +02:00
Gani Georgiev 9855397a22 replaced the default binder with rest.MultiBinder 2024-01-20 15:03:45 +02:00
Gani Georgiev d9b219d64f [#4192] take collection minPasswordLength in consideration for the user password generator btn 2024-01-20 13:18:00 +02:00
Gani Georgiev c642a860ca rename local const redirect path vars for consistency 2024-01-20 13:16:06 +02:00
Gani Georgiev b2b792b763 [#4177] added graceful OAuth2 redirect error handling 2024-01-19 19:15:01 +02:00
Willow (GHOST) fc18e69183 [#4175] update patreon logo 2024-01-18 18:13:46 +02:00
Gani Georgiev fa65038fc1 synced with master 2024-01-16 13:14:07 +02:00
Gani Georgiev 9419d1928a [#4160] fixed the Admin UI auto indexes update when renaming fields with a common prefix 2024-01-16 12:50:44 +02:00
Gani Georgiev c9bc2f07aa added EmailTemplate.Hidden field 2024-01-16 11:38:09 +02:00
Gani Georgiev 28fc186f5c added support for loading a serialized json payload as part of multipart/form-data request 2024-01-14 22:20:46 +02:00
Gani Georgiev cdb539dcc8 updated changelog 2024-01-13 18:02:49 +02:00
Gani Georgiev af7c6d8d9b [#4066] mark user as verified on confirm password reset 2024-01-13 17:52:41 +02:00
Gani Georgiev cd2fc536ca updated Prism.js bundle 2024-01-13 16:26:59 +02:00
Gani Georgiev 2a28f6ff33 [#4106] added custom Prism.js bundle and registered new TinyMCE codesample languages 2024-01-13 16:13:32 +02:00
Gani Georgiev 036c0da05f reduce slightly the min row table height 2024-01-13 14:55:59 +02:00
Gani Georgiev d795a6671b updated go.sum 2024-01-13 13:23:21 +02:00
Gani Georgiev 8cff94f27c synced with master 2024-01-13 13:23:08 +02:00
Gani Georgiev 931f6bc0cb updated go deps 2024-01-13 11:36:22 +02:00
Gani Georgiev 2e3ae1b60a [#4145] fixed JSVM types generation for functions with omitted arg types 2024-01-13 11:28:15 +02:00
Gani Georgiev 6155e6426a ghupdate messages update 2024-01-13 11:21:55 +02:00
Gani Georgiev 0f95a11fc1 Merge branch 'master' into develop 2024-01-05 20:52:27 +02:00
Gani Georgiev eb695cc6d3 updated go crypto and other go deps 2024-01-05 20:48:15 +02:00
Gani Georgiev 28d15e86eb fixed optional migration condition
note: practically even the previous version should work ok because the json field didn't have previous options anyway and if it was nil the migration will fail
2024-01-05 20:35:09 +02:00
Gani Georgiev b033109654 synced with master 2024-01-04 21:26:55 +02:00
Gani Georgiev d0352aa3f9 [#4079] fixed popup searchbar css styles to prevent hiding the additional controls 2024-01-04 16:20:47 +02:00
Gani Georgiev 3b7d0e84f6 synced with master 2024-01-03 17:00:24 +02:00
Gani Georgiev 592b7e85c6 specify the exact license and changelog files to include in the release archive 2024-01-03 15:04:34 +02:00
Gani Georgiev 3792d44c35 updated changelog 2024-01-03 14:25:51 +02:00
Gani Georgiev a021fcaa75 [#4072] added non-json value dummy object wrap normalization 2024-01-03 14:16:12 +02:00
Gani Georgiev d123e19e61 synced with master 2024-01-03 12:46:49 +02:00
Gani Georgiev 982f876a93 updated jsvm types 2024-01-03 11:08:30 +02:00
Gani Georgiev 1fcc2d8683 updated CHANGELOG and added t.Parallel to some of the tests 2024-01-03 10:58:25 +02:00
Gani Georgiev 4f2492290e [#4068] fixed the json field query comparisons to work correctly with plain JSON values 2024-01-03 10:43:46 +02:00
Gani Georgiev 8f625daa2f updated some of the tests to use t.Parallel 2024-01-03 04:30:20 +02:00
Gani Georgiev 0599955676 sort cascadeDelete refs for deterministic tests output 2024-01-03 04:29:30 +02:00
Gani Georgiev 97a8409a65 fixed sleep example typo and synced with master 2023-12-30 12:06:07 +02:00
Gani Georgiev 422eb30797 synced with master 2023-12-29 23:31:54 +02:00
Gani Georgiev c4116e3a7d added jsvm sleep binding 2023-12-29 23:29:00 +02:00
Gani Georgiev 64cee264f0 bumped app version 2023-12-29 22:00:47 +02:00
Gani Georgiev 0ae9f24a81 updated fields query param examples for the auth actions 2023-12-29 21:47:04 +02:00
Gani Georgiev 6d942c7d30 docs fixes commits from develop 2023-12-29 21:25:32 +02:00
Gani Georgiev 705d7f48e7 synced with master 2023-12-29 10:00:06 +02:00
Gani Georgiev 9f67c5d563 regenerated jsvm types 2023-12-29 09:58:21 +02:00
mookrs 1ac7330e0b [#4043] fixed typos in godoc comments 2023-12-29 09:56:36 +02:00
Gani Georgiev e73b3a32d2 synced with master 2023-12-27 10:50:48 +02:00
Gani Georgiev 461886f64e fixed the monospace font loading in the Admin UI 2023-12-27 10:47:18 +02:00
Gani Georgiev 370862fa2e updated AppleClientSecretCreate struct comment 2023-12-26 20:50:35 +02:00
Gani Georgiev 6b3780c630 [#4035] replaced JWT token with just JWT 2023-12-26 19:57:38 +02:00
Gani Georgiev c807f66c59 synced with master 2023-12-24 11:23:50 +02:00
Gani Georgiev 8d97eb0769 [#4022] fixed multi-line text paste in the Admin UI search bar 2023-12-24 11:13:17 +02:00
Gani Georgiev 5f5f9ca426 reorder loading=lazy before src per the svelte docs 2023-12-18 07:47:17 +02:00
Gani Georgievandaabajyan 4e91be6d74 [#3948] added Bitbucket OAuth2 provider
Co-authored-by: aabajyan <arsen.abajyan@pm.me>
2023-12-17 15:47:17 +02:00
Gani Georgiev 1208edec92 regenerated jsvm types 2023-12-17 00:20:04 +02:00
Gani Georgiev 5555e63116 updated dev debug log text message color to be slightly more visible 2023-12-17 00:17:29 +02:00
Gani Georgiev d6569b445c added timestamp to the generated JSVM types file to prevent creating it every time on app startup 2023-12-16 23:20:38 +02:00
Gani Georgiev 0b4f3b2adf combine the logs listing label span 2023-12-16 23:16:07 +02:00
Gani Georgiev 8bd968ed06 split changelog in chunks 2023-12-16 18:22:18 +02:00
Gani Georgiev 5c961f8537 [#3918] added --dev flag, dev log printer and some minor log UI enhacements 2023-12-16 18:15:36 +02:00
Gani Georgiev bf5eba0384 added bool to the view query sql syntax highlighter and autocompletion 2023-12-13 09:07:18 +02:00
Gani Georgiev c213d9313e fixed changelog typo 2023-12-12 19:48:28 +02:00
Gani Georgiev b31cf984a5 [#3930] replaced the default 100ms api tests timeout in favor of new ApiScenario.Timeout field 2023-12-12 19:46:58 +02:00
Gani Georgiev 8671debc35 removed the blank current time entry from the logs chart 2023-12-11 09:49:06 +02:00
Gani Georgiev b0f027d27a updated changelog formatting and temp moved the admin only rule checks to the record_helpers 2023-12-10 21:06:02 +02:00
Gani Georgiev 98c8c98603 updated jsvm types 2023-12-10 12:50:56 +02:00
Gani Georgiev 97345f0317 skip log writes if max retention setting is zero 2023-12-10 12:40:33 +02:00
Gani Georgiev b29e404f22 updated ui/dist, go deps, docs and fixed some typos 2023-12-10 12:23:31 +02:00
Gani Georgiev d8ec36fa4c updated jsvm types 2023-12-09 22:40:45 +02:00
Gani Georgiev fb2eafe860 [#3790] added MaxSize json field option 2023-12-09 22:30:37 +02:00
Gani Georgiev b9f391cf85 revert ResetBootstrapState removal on app termination since closing the db explicitly enforces checkout and clearing the side-car wal file 2023-12-09 19:43:29 +02:00
Gani Georgiev 646f90ef43 updated logs chart 2023-12-09 16:31:17 +02:00
Gani Georgiev 5b6b4599b7 updated logs listing 2023-12-09 15:12:00 +02:00
Gani Georgiev 35fc6d0734 define Server.BaseContext to cancel globally the SSE connections on server shutdown 2023-12-08 23:14:14 +02:00
Gani Georgiev 506b759560 fixed graceful shutdown handling 2023-12-08 21:16:48 +02:00
Gani Georgiev d86e20b7f2 remove the unnecessary App.ResetBootstrapState calls as sqlite connections will be closed anyway with the process termination 2023-12-08 19:24:14 +02:00
Gani Georgiev 4c473385b2 trigger OnTerminate() hook on app.Restart() call 2023-12-08 15:46:33 +02:00
Gani Georgiev afbbc1d97c removed unnecessary logs index and updated logs ui 2023-12-08 14:26:06 +02:00
Gani Georgiev 4d3ba270c0 fix nullable non-equal comparisions 2023-12-08 13:50:12 +02:00
Gani Georgiev 1bf7f148b0 minor types.DateTime optimizations to minimize time.Time value copies 2023-12-08 10:36:12 +02:00
Gani Georgiev 6e6c873cc6 [#3896] added $apis.requireGuestOnly() middleware JSVM binding 2023-12-07 18:49:56 +02:00
Gani Georgiev 16da7d9e1a removed unused options struct 2023-12-06 20:44:47 +02:00
Gani Georgiev f7df737c45 added filesystem.NewFileFromUrl(ctx, url) 2023-12-06 20:42:30 +02:00
Gani Georgiev 64eefb44e8 added onlyVerified field to the authMethods response 2023-12-06 13:30:47 +02:00
Gani Georgiev 31317df21c added onlyVerified auth collection option 2023-12-06 11:57:04 +02:00
Gani Georgiev 865865fdeb updated jsvm $security.parse* token helpers to return the payload as plain object 2023-12-04 20:46:33 +02:00
Gani Georgiev 5b2575b754 [#3877] fixed test messages typo 2023-12-04 18:09:29 +02:00
Gani Georgiev 6327ac20da updated changelog 2023-12-04 17:18:57 +02:00
Gani GeorgievandTobias Muehlberger 8cd1c8709c [#3794] limit concurrent thumbs generation
Co-authored-by: Tobias Muehlberger <tobias@muehlberger.dev>
2023-12-04 16:52:10 +02:00
Gani Georgiev 14a2fd6215 skip wrapping sql.ErrNoRows 2023-12-04 16:23:56 +02:00
Gani Georgiev cdfc1f7b70 removed unnecessary Close call and formatted map hints 2023-12-04 16:22:49 +02:00
Gani Georgiev 41dcd9b4d4 use error.Is to handle wrapped errors 2023-12-04 16:21:57 +02:00
Gani Georgiev 0fb859c321 updated logs list min-width 2023-12-03 20:58:12 +02:00
Gani Georgiev f57d38f529 use linear thumb resample filter 2023-12-03 20:56:28 +02:00
Gani Georgiev 04024cb6b7 removed incorrect base error message 2023-12-03 20:55:15 +02:00
Gani Georgiev 58a2d3cd09 added the failed dao query to the error message 2023-12-03 20:54:48 +02:00
Gani Georgiev 4d27278c60 always show list errors if there is no filter 2023-12-03 20:54:04 +02:00
Gani Georgiev 70f1647a4c updated logs list styles 2023-12-03 14:12:44 +02:00
Gani Georgiev 5b94aced3a use a red colored stderr writer for the cobra cmd errors 2023-12-03 13:44:30 +02:00
Gani Georgiev 070a1cd6d9 removed eagerly resetting the bootstrap state to prevent concurrent access errors 2023-12-03 12:36:51 +02:00
Gani Georgiev 716f508d66 removed activity logger for the realtime connect action and added helper debug log when subscriptions are changed 2023-12-03 12:12:30 +02:00
Gani Georgiev 7013174315 removed empty local() font-face declarations 2023-12-03 12:00:11 +02:00
Gani Georgiev 559aad36a3 added the log id in the query params 2023-12-03 11:39:40 +02:00
Gani Georgiev 6416328c3b added support for specifying @collection.* aliases 2023-12-03 10:57:58 +02:00
Gani Georgiev d3713a9d7c added support for comments in the API rules and filter expressions 2023-12-02 16:37:04 +02:00
Gani Georgiev aaab643629 [#3700] allow a single OAuth2 user to be used for authentication in multiple auth collection 2023-12-02 12:43:22 +02:00
Gani Georgiev b283ee2263 added OAuth2 displayName and pkce options 2023-11-29 20:19:54 +02:00
Gani Georgiev 995733000f added filesystem.Copy(src, dest) 2023-11-28 21:09:53 +02:00
Gani Georgiev 99bdb4e701 [#3617] added expiry field to the OAuth2 user 2023-11-27 20:32:28 +02:00
Gani Georgiev 3b79535dc7 sort the auth providers by their Name field 2023-11-27 20:05:06 +02:00
Gani Georgiev 05cc3f9e6c updated confirm password reset docs example 2023-11-26 15:00:48 +02:00
Gani Georgiev 3f2e38ca82 updated API preview examples 2023-11-26 14:59:14 +02:00
Gani Georgiev 531a7abec9 updated links formatting in the autogenerated html->text mail body 2023-11-26 14:47:26 +02:00
Gani Georgiev 821aae4a62 logs refactoring 2023-11-26 13:33:17 +02:00
Gani Georgiev ff5535f4de synced with master 2023-11-11 12:51:26 +02:00
Gani Georgiev 985ab1e5b7 updated changelog 2023-11-11 12:50:39 +02:00
Gani Georgiev 69a805d0d1 synced with master 2023-11-11 12:50:20 +02:00
Gani Georgiev d240649497 updated ui/dist 2023-11-11 12:48:11 +02:00
Gani Georgiev 9957919d9a updated tygoja and the generated jsvm types 2023-11-11 12:46:46 +02:00
Gani Georgiev 890a0904cf [#3697] allowed hyphens in usernames 2023-11-11 12:19:33 +02:00
Gani Georgiev cdd32512d5 synced with master 2023-11-10 15:18:14 +02:00
Gani Georgiev 5835193900 [#3735] fixed text field min/max validators to properly count multi-byte characters 2023-11-10 14:58:00 +02:00
Gani Georgiev 4abe199acc [#3715] fixed TinyMCE source code viewer textarea styles 2023-11-08 21:19:16 +02:00
Gani Georgiev a170923637 synced with master 2023-11-06 11:42:59 +02:00
Gani Georgiev f4f3724b7a updated ui/dist 2023-11-06 11:35:59 +02:00
Gani Georgievandsergeypdev ba7cf8bf8e [#3689] relaxed the OAuth2 redirect url validation to allow any string value
Co-authored-by: sergeypdev <sergeypoznyak@protonmail.com>
2023-11-06 11:33:10 +02:00
Gani Georgiev 500615c1ee added missing documention for the JSVM $mails.* bindings 2023-11-06 11:26:38 +02:00
Gani Georgiev 8961232a44 [#3685] added the release notes to the success ghupdate output 2023-11-06 11:19:12 +02:00
Gani Georgiev 907167e696 synced with master 2023-11-03 09:56:50 +02:00
Gani Georgiev 4e51e393a2 updated ui/dist 2023-11-03 05:50:49 +02:00
Gani Georgiev 5ea784609f Merge branch 'master' into develop 2023-10-28 18:46:43 +03:00
Gani Georgiev ea5ca009de [#3627] updated tygoja to fallback to []number for the generated TS []byte union type when used in 'M extends T' declarations 2023-10-28 16:38:18 +03:00
Gani Georgiev d13802133a fixed changelog typos 2023-10-28 00:27:59 +03:00
Gani Georgiev f3a40001a4 updated codemirror deps and regenerated ui/dist 2023-10-27 22:40:18 +03:00
Gani Georgiev 1ae570921b added negative string number normalizations for the json field type 2023-10-27 22:37:11 +03:00
Gani Georgiev f889a3fcb3 synced with master 2023-10-27 22:28:15 +03:00
Gani Georgiev 34fed679fd removed old comment 2023-10-27 17:38:53 +03:00
Gani Georgiev b7a49efa88 fixed excerpt modifier to properly add spaces after block tags 2023-10-27 17:36:26 +03:00
Gani Georgiev 01e8c0f9f7 [#3616] fixed tokenizer whitespace characters trimming 2023-10-27 15:19:06 +03:00
Gani Georgiev 1d67a35acf added changelog rc note 2023-10-27 07:29:26 +03:00
Gani Georgiev e2d8028d0a [#3602] use the auth collection name in the OAuth2 examples 2023-10-25 22:18:18 +03:00
Gani Georgiev d8a1875f84 fix the node version as latest seems to cause some issue with sass 2023-10-24 15:12:19 +03:00
Gani Georgiev 79617e6d99 =added experimental expand, filter, fields, custom query and headers parameters support for the realtime subscriptions 2023-10-24 14:46:03 +03:00
Gani Georgiev e6f1b3dfe4 updated relation field validation message 2023-10-21 15:52:19 +03:00
Gani Georgiev 94253f0dd5 updated the supported non-cgo build targets list 2023-10-16 20:27:37 +03:00
Kunal Singh 6cfaf343ac [#3531] updated README.md and CONTRIBUTING.md formatting 2023-10-16 20:18:05 +03:00
Gani Georgiev 9c562294ff set a default id column width and updated ui dist 2023-10-15 14:40:16 +03:00
Gani Georgiev 3c5409d607 updated changelog 2023-10-15 14:17:09 +03:00
Gani Georgiev 8868fa9ae6 use a custom tinymce svelte component and other minor optimizations 2023-10-15 14:04:44 +03:00
Gani Georgiev c0fa53a2ab check the mime type of the collections file field and updated field styles to minimize the layout shifts 2023-10-15 06:49:32 +03:00
Gani Georgiev 007b6a04ff updated dependencies and regenerated jsvm types 2023-10-14 23:28:32 +03:00
Gani Georgiev 731383a915 added .cmd() as alias for .exec() 2023-10-14 20:08:21 +03:00
Gani Georgiev 866d38caf9 updated jsvm types and removed unused helper 2023-10-14 19:14:27 +03:00
Gani Georgiev 1f6ab24b34 updated replaceQueryParams to use the last ? 2023-10-14 15:11:32 +03:00
Gani Georgievandthisni1s 01e33c07fe [#3364] added mailcow OAuth2 provider
Co-authored-by: thisni1s <nils@jn2p.de>
2023-10-14 14:52:35 +03:00
Gani Georgiev 69983bff5e removed legacy fonts 2023-10-12 23:34:34 +03:00
Gani Georgiev 2567659696 dragline z-index fix 2023-10-10 21:38:15 +03:00
Gani Georgiev 3e487f7e9d updated api preview docs 2023-10-09 21:03:02 +03:00
Gani Georgiev ca1a395628 minor styles adjustments 2023-10-09 19:55:53 +03:00
Gani Georgiev 1a47c70ccf Added support to manually resize the collections sidebar 2023-10-09 16:11:49 +03:00
Gani Georgiev 1f4bdfb867 [#3112] added options to pin collections 2023-10-09 14:26:56 +03:00
Gani Georgiev eae16cc42c synced with master 2023-10-09 12:01:21 +03:00
Gani Georgiev 1527b5ea4f updated CHANGELOG 2023-10-08 23:43:58 +03:00
Gani Georgiev ba6e17b3be updated jsvm types 2023-10-08 23:26:23 +03:00
Gani Georgiev b8219af941 [#3476] added raw template function 2023-10-08 23:17:38 +03:00
Gani Georgiev 8865cc1431 renamed record upsert local requestInfo to requestData to distinguish better from models.RequestInfo 2023-10-08 22:52:14 +03:00
Gani Georgiev 20b6ce4b84 excluded expand from the record draft and applied some lint fields alignment suggestions 2023-10-08 15:22:03 +03:00
Gani Georgiev e2f806d8bb added jsvm subscriptions.Message binding 2023-10-07 16:11:38 +03:00
Gani Georgiev 49e3f4ad93 [#3447] added jsvm http.Cookie binding 2023-10-07 15:35:20 +03:00
Gani Georgiev 6d672348e7 rearanged the DefaultClient struct fields to reduce its size from ~72 to ~32 bytes 2023-10-07 13:17:32 +03:00
Gani Georgiev 80d774a8ef [#3461] removed content-type charset and deprecated keep-alive header field 2023-10-07 12:57:07 +03:00
Gani Georgiev 5a5125383a Merge branch 'master' into develop 2023-10-05 09:33:46 +03:00
Gani Georgiev 7fa1ff53c9 trim view query semicolon chars and allow single quotes for column aliases 2023-10-05 09:31:24 +03:00
Gani Georgiev 0f4e27a11f updated nonempty label styles 2023-10-04 10:19:00 +03:00
Gani Georgiev 9997223923 fixed comment 2023-10-04 01:27:50 +03:00
Gani Georgiev 632ade795f updated file picker thumbs size 2023-10-03 16:18:32 +03:00
Gani Georgiev 91bd739b71 load only records with non-empty file fields and fupdated files list styles 2023-10-03 15:00:45 +03:00
Gani Georgiev 957064d70b extract the thumb sizes only from the selected file field 2023-10-03 12:51:55 +03:00
Gani Georgiev 609792a355 added records file picker support for the editor field 2023-10-03 10:36:46 +03:00
Gani Georgiev 2f5cfcfe87 replaced interface{} with any 2023-10-01 18:45:27 +03:00
Gani Georgiev 5732bc38e3 reload the records counter and remove drafts failures from LocalStorage 2023-10-01 15:57:20 +03:00
Gani Georgiev d69181cfef added helper class to disable the tabs animation to avoid the flickering 2023-10-01 15:56:29 +03:00
Gani Georgiev 8908d03b8c added support for linking to the record preview/update form and some other minor improvements 2023-10-01 12:55:30 +03:00
Gani Georgiev ebf73f5602 updated ui/dist 2023-09-30 14:43:12 +03:00
Gani Georgiev 8416f03bcf show local date on hover 2023-09-30 12:13:00 +03:00
Gani Georgiev 5d87385170 synced with master 2023-09-30 10:20:57 +03:00
Gani Georgiev 837134559f updated changelog 2023-09-30 09:37:26 +03:00
Gani Georgiev fadd12cd22 update tygoja and the generated jsvm typings 2023-09-28 23:04:34 +03:00
Gani Georgiev 469769d270 updated go deps 2023-09-25 23:30:10 +03:00
Gani Georgiev e4b7303a5d synced with master 2023-09-25 23:26:07 +03:00
Gani Georgiev e1fb5d26a5 [#3382] replaced filepath with path when extracting the filekey parent prefix 2023-09-25 22:47:48 +03:00
Gani Georgiev 4f396ca439 synced with master 2023-09-24 11:57:45 +03:00
Gani Georgiev ff08fc0fa4 remove the created and updated fields from the view API Preview and listings if the query doesn't have them 2023-09-24 11:27:10 +03:00
Gani Georgiev 4b511475ff updated go deps 2023-09-24 11:07:35 +03:00
Gani Georgiev 2550a9de54 [#3344, #2505] optimized records listing 2023-09-24 11:05:12 +03:00
Gani Georgiev 0f5dad7ede synced with master 2023-09-22 21:24:46 +03:00
Gani Georgiev d0b1c9d998 updated the invalid rel ids reactivity handling 2023-09-22 18:40:00 +03:00
Gani Georgiev fd9e120434 updated go deps and regenerated jsvm types 2023-09-22 18:24:34 +03:00
Gani Georgiev 92731ddd50 [#3372] fixed Admin UI listing error on invalid record relation 2023-09-22 18:19:05 +03:00
Gani Georgiev 4b4aaf2112 use goccy/go-json to speedup serialization 2023-09-18 22:52:36 +03:00
Gani Georgiev 6013d14bc6 added support for :excerpt(max, withEllipsis?) fields modifier 2023-09-18 15:20:10 +03:00
Gani Georgiev f3bcd7d3df added tokenizer.IgnoreParenthesis() to allow ignoring the parenthesis characters boundary checks 2023-09-17 12:14:57 +03:00
Gani GeorgievandGHOST 71f9be3cb0 [#3323] added Patreon OAuth2 provider
Co-authored-by: GHOST <ghostdevbusiness@gmail.com>
2023-09-16 08:20:49 +03:00
Gani Georgiev f605521208 updated js types docs 2023-09-16 07:07:35 +03:00
Gani Georgiev 4927583790 updated changelog and go deps 2023-09-16 06:54:19 +03:00
Gani Georgiev 6e80cb8136 added more descriptive internal password reset error message 2023-09-15 20:45:28 +03:00
Gani Georgiev bb0a2dd698 [#3310] added headers and cookies fields to the .send result 2023-09-14 14:47:47 +03:00
Gani Georgiev 2608efb56c added array fallback in case of missing joinNonEmpty items 2023-09-12 19:57:42 +03:00
Gani Georgiev eb2aa1cfc6 [#2197] added escape character support for the select field options 2023-09-12 10:29:54 +03:00
Gani Georgiev e1528aedac updated migration comment 2023-09-10 18:33:27 +03:00
Gani Georgiev 22b0a2b586 updated changelog 2023-09-10 10:57:51 +03:00
Gani Georgiev 0ca86a0c87 [#3273] added readerToString() JSVM helper 2023-09-10 10:46:19 +03:00
Gani Georgiev b2c8f394af fixed changelog typo 2023-09-09 12:29:09 +03:00
Gani Georgiev 56b2641469 added hmac jsvm primitives and updated docs 2023-09-09 12:03:34 +03:00
Gani Georgiev f266621a0f updated go deps 2023-09-06 14:13:02 +03:00
Gani Georgiev ca136c5dc1 [#3265] silent the localStorage quota error to prevent breaking the record form panel 2023-09-06 14:11:58 +03:00
Gani Georgiev abfe18bcce [#3261] exclude the local temp dir from the backups 2023-09-06 07:09:28 +03:00
Gani Georgiev f3fc7f78d7 updated changelog 2023-09-05 11:54:09 +03:00
Gani Georgiev 26fd069d11 updated npm deps 2023-09-05 11:49:01 +03:00
Gani Georgiev 62bde9e1f3 updated jsvm types 2023-09-05 11:25:14 +03:00
Gani Georgiev b945cd2fdf updated changelog and ui/dist 2023-09-05 11:23:41 +03:00
Gani Georgiev 89a0520f7d [#3257] normalize pasted text in the editor field 2023-09-05 10:36:42 +03:00
Gani Georgiev a88b3c5db3 updated api docs and enabled paste_as_text editor option 2023-09-05 08:06:21 +03:00
Gani Georgiev a37d9cfb84 updated ui/dist 2023-09-04 11:35:39 +03:00
Gani Georgiev 5b084bbbfd fixed grammar 2023-09-01 14:28:49 +03:00
Gani Georgiev 322508f6d1 registered a custom Deflate compressor to speedup the backups generation 2023-09-01 14:27:23 +03:00
Gani Georgiev 78e70bd52b there is no need to nil the app.settings on ResetBootstrapState 2023-09-01 13:46:33 +03:00
Gani Georgiev 58401459bf updated ui/dist 2023-09-01 12:44:43 +03:00
Gani Georgiev f172754775 removed cgo build archives artifacts 2023-09-01 12:41:49 +03:00
Gani Georgiev 31670ab3e1 log cron job errors 2023-09-01 11:17:09 +03:00
Gani Georgiev baacf4913b updated api preview tabs style 2023-09-01 09:30:49 +03:00
Gani Georgiev 8a94ccea42 updated to Svelte 4 2023-09-01 09:22:49 +03:00
Gani Georgiev b2bab9573a removed forgotten svg icon font declaration 2023-08-30 20:39:39 +03:00
Gani Georgiev e5b5c1f76f added option to auto generate admin and auth record passwords from the Admin UI 2023-08-30 14:59:00 +03:00
Gani Georgiev ccb1c42220 updated jsvm types and removed unnecessary comment 2023-08-29 22:31:33 +03:00
Gani Georgiev a394777264 [#3191] added client-side validation and syntax highlight for the json field 2023-08-29 22:10:57 +03:00
Gani Georgiev 7d10b3c502 updated go deps 2023-08-29 19:23:01 +03:00
Gani Georgiev 64ffb308bb use singular NoDecimal option name 2023-08-29 18:41:20 +03:00
Gani Georgiev 916c74c218 [#3113] added NoDecimal number field option 2023-08-29 18:35:57 +03:00
Gani Georgiev 17974d534e added ellipsis for long backup titles 2023-08-28 21:05:35 +03:00
Gani Georgiev bde7a86b30 [#3066] added the application name as part of the autogenerated backup name for easier identification 2023-08-28 20:36:55 +03:00
Gani Georgiev f7f8f09336 [#2599] added option to upload a backup file from the Admin UI 2023-08-28 20:06:48 +03:00
Gani Georgiev 2a6b891a9b merged with master 2023-08-26 14:46:30 +03:00
Gani Georgiev b4cb35483b updated jsvm types 2023-08-26 14:43:55 +03:00
Gani Georgiev 1606cfd6e2 fixed cronAdd example 2023-08-26 14:34:47 +03:00
Gani Georgiev 08f97ef0bb Merge branch 'master' into develop 2023-08-26 10:38:13 +03:00
Gani Georgiev 0dc263a40c updated go deps and use the new fileblob NoTempDir option 2023-08-26 10:37:12 +03:00
Gani Georgiev 824031e1a4 updated changelog 2023-08-26 10:30:26 +03:00
Gani Georgiev 311bc74b7e [#3025] updated tests.ApiScenario fields 2023-08-25 22:14:04 +03:00
Gani Georgiev 4f3d1682de synced with master 2023-08-25 18:11:09 +03:00
Gani Georgiev 18728732b9 updated ui/dist and go deps 2023-08-25 16:49:59 +03:00
Gani Georgiev bb4f27cfb5 updated automigrate template test 2023-08-25 16:47:41 +03:00
impact-merlinmarekandMerlin Marek d423acad3b [#3192] fixed autogenerated down migration not preserving old rules state
Co-authored-by: Merlin Marek <merlin.marek@posteo.net>
2023-08-25 16:21:04 +03:00
Gani Georgiev ef73052546 added httpAddr default when domain name is missing 2023-08-25 12:06:24 +03:00
Gani Georgiev c89c68a4dc poc of serve domain args 2023-08-25 11:16:31 +03:00
Gani Georgiev 02495554cf [#3175] added jsvm crypto primitives 2023-08-24 11:25:00 +03:00
Gani Georgiev cdbe6d78d3 added basic fields wildcard support 2023-08-23 20:56:38 +03:00
Gani Georgiev ff6904f1f8 removed unnecessary test cases prefix 2023-08-23 16:50:43 +03:00
Gani Georgiev bc0222dcb4 [#3176] skip fields query param transformations for non 20x responses 2023-08-23 16:49:09 +03:00
Gani Georgiev 04826ba588 reduced the default prewarmed goja vms to 25 2023-08-22 22:09:33 +03:00
Gani Georgiev 6ca1f5c431 use crypto if available 2023-08-22 22:01:05 +03:00
Gani Georgiev 2863763a27 added option to control the default TinyMCE urls convert behavior 2023-08-22 14:39:21 +03:00
Gani Georgiev 9c0d952543 fixed isNew checks 2023-08-22 13:01:08 +03:00
Gani Georgiev ed4f7c7358 updated presentable fields sorting 2023-08-22 12:29:54 +03:00
Gani Georgiev 49f1c869c0 synced with master 2023-08-22 10:54:22 +03:00
Gani Georgiev 5e6949062f updated changelog 2023-08-21 18:15:34 +03:00
Gani Georgiev dc063b20fa bumped app version 2023-08-21 18:14:23 +03:00
Gani Georgiev f42bbfd927 updated go deps 2023-08-21 18:10:07 +03:00
Gani Georgiev f0af24d78f use the presentable prop when displaying relations 2023-08-21 18:06:35 +03:00
Gani Georgiev 26fd3d48df added migration to copy existing DisplayFields to the new Presentable field 2023-08-21 12:58:51 +03:00
Gani Georgiev 864bbe7e12 added SchemaField.Presentable field 2023-08-21 12:58:18 +03:00
Gani Georgiev 1e995552c8 updated apis.Serve godoc 2023-08-20 18:31:56 +03:00
Gani Georgiev 6baae97b5d added json marshal fallback for complex structs as placeholder param 2023-08-19 16:33:45 +03:00
Gani Georgiev bcfbbc53f8 [#3147] don't silence connectivity errors 2023-08-18 19:14:13 +03:00
Gani Georgiev 8a916cd636 added datetime macros 2023-08-18 08:48:33 +03:00
Gani Georgiev 75f58a28ac added placeholder params support for Dao.FindRecordsByFilter and Dao.FindFirstRecordByFilter 2023-08-18 06:31:14 +03:00
Gani Georgiev e87ef431c5 added jsvm .* binds 2023-08-17 20:50:00 +03:00
Gani Georgiev b2ac538580 [#3097] added SmtpConfig.LocalName option 2023-08-17 19:07:56 +03:00
Gani Georgiev 53b20ec104 updated LastVerificationSentAt and LastResetSentAt fill sequence 2023-08-17 14:03:11 +03:00
Gani Georgiev c8ef3c4050 updated ui/dist 2023-08-16 22:38:56 +03:00
Gani Georgiev 9113d30103 fixed typo 2023-08-16 17:38:08 +03:00
Gani Georgiev fef6e584b7 updated jsvm types and godoc list formatting 2023-08-15 12:35:52 +03:00
Gani Georgiev 67fa47b1bb [#3132] updated godoc 2023-08-15 12:25:24 +03:00
Gani Georgiev 5f21c4119f [#3132] added common cron expression macros 2023-08-15 12:21:33 +03:00
Gani Georgiev 734f35c504 synced with master 2023-08-15 12:10:21 +03:00
Gani Georgiev 038ae8f803 updated changelog 2023-08-15 01:14:10 +03:00
Gani Georgiev 2236288e57 quoted the wrapped view query columns 2023-08-15 01:09:53 +03:00
Gani Georgiev 8f10c66160 updated ui/dist 2023-08-15 00:58:34 +03:00
Gani Georgiev 5960dc5f2d removed js sdk dto helpers 2023-08-14 21:20:49 +03:00
Gani Georgiev cbf1215bb1 updated jsvm types 2023-08-11 14:40:25 +03:00
Gani Georgiev 1b633720be updated views migrations to use SaveCollection 2023-08-11 14:37:53 +03:00
Gani Georgiev adb5d6e998 [#3110] normalized view queries with numeric or expression ids 2023-08-11 14:29:18 +03:00
Gani Georgiev 3841946b61 downgraded temp gocloud until the os.NoTempDir is released 2023-08-10 21:38:13 +03:00
Gani Georgiev 369e2703c2 updated ui/dist 2023-08-10 08:56:26 +03:00
Gani Georgiev 4a45ad91fa [#3106] always refresh the Admins UI initial admins counter cache when there are none 2023-08-10 08:50:48 +03:00
Nikita Zhenev 265dac45ce [#3103] fixed jsvm registerMigrations error message typo 2023-08-09 18:40:34 +03:00
Gani Georgiev f152787578 updated changelog 2023-08-09 13:17:24 +03:00
Gani Georgiev 640623f4fd [#3098] fixed incorrect cascade delete tooltip message 2023-08-09 12:58:20 +03:00
Gani Georgiev 1aff89f377 use the logs maxDays before firing the goroutine 2023-08-09 12:23:49 +03:00
Gani Georgiev 7d6b12a4ef updated ui/dist 2023-08-08 14:33:08 +03:00
Gani Georgiev 7a3223e415 [#3089] use a temp dir inside pb_data to prevent backups cross-device link error 2023-08-08 14:15:29 +03:00
Gani Georgiev bd18688f35 [#3090] fixed relation to view error message 2023-08-08 12:41:51 +03:00
Gani Georgiev f90da96820 enabled lazy loading for the Admin UI thumb images 2023-08-06 21:51:55 +03:00
Gani Georgiev 6c8f2d2cd6 use scrollbar-gutter to minimize the table records listing layout shifts 2023-08-06 21:44:26 +03:00
Gani Georgiev 5e84305922 fixed changelog grammar 2023-08-05 10:12:22 +03:00
Gani Georgiev b3421861e6 updated jsvm types 2023-08-05 09:51:00 +03:00
Gani Georgiev 872492ad22 updated changelog 2023-08-05 07:26:53 +03:00
Gung JodiandGung Jodi 5c14c7cf5e [#3068] fixed RequestData log deprecation note
Co-authored-by: Gung Jodi <agung.pratama@dana.id>
2023-08-05 07:24:20 +03:00
Gani Georgiev b59f0f418e updated ui/dist 2023-08-03 12:42:00 +03:00
Gani Georgiev 06d3e27e03 [#3054] added core.RealtimeConnectEvent.IdleTimeout field 2023-08-03 12:38:02 +03:00
Gani Georgiev b1093baef7 [#3058] soft-deprecated 'data' prop in favour of 'body' to allow raw strings 2023-08-03 12:32:04 +03:00
Gani Georgiev b3f09ff045 updated changelog and go deps 2023-07-31 22:47:35 +03:00
Gani Georgiev b33ad36f64 renamed variable name 2023-07-31 17:48:26 +03:00
Gani Georgiev 9254ce46eb trigger the jsvm cron ticker only on app serve 2023-07-31 14:18:59 +03:00
Gani Georgiev cc8c855306 updated ui/dist 2023-07-31 13:09:39 +03:00
Gani Georgiev b74994b906 [#3026] use relative path for the oauth2 provider page link 2023-07-31 13:07:30 +03:00
Gani Georgiev f652dc71bb [#3025] manually trigger the OnBeforeServe hook for tests.ApiScenario 2023-07-31 12:27:22 +03:00
Gani Georgiev 6d2677a5e3 fixed cronRemove docs declaration 2023-07-30 18:02:16 +03:00
Gani Georgiev 0c5305c174 fixed incomplete sentence in the changelog 2023-07-30 16:20:58 +03:00
Gani Georgiev 5398576f4f updated changelog formatting 2023-07-30 15:22:10 +03:00
Gani Georgiev 3d1f570b38 updated jsvm types 2023-07-30 14:12:35 +03:00
Gani Georgiev fa057502f1 updated ui/dist deps 2023-07-30 14:10:42 +03:00
Gani Georgiev bb4a5cfe83 updated ui/dist and some lint warnings 2023-07-30 13:40:22 +03:00
Gani Georgiev ac1fd74942 updated tests 2023-07-30 10:02:44 +03:00
Gani Georgiev db660ac780 revert the default max perPage limit to 500 for now 2023-07-29 21:44:31 +03:00
Gani Georgiev cdeb9a94ed added action arg to the before Dao hook to allow skipping the default persist behavior 2023-07-29 19:52:36 +03:00
Gani Georgiev 6da94aef8d updated jsvm panic handling when HooksWatch is set 2023-07-29 16:01:53 +03:00
Gani Georgiev 0a4fdc17a5 enabled tokens binds and removed primitive constructors overwrites 2023-07-29 13:56:31 +03:00
Gani Georgiev 1bbba7a0ae added cron expression UTC timezone note 2023-07-28 22:24:21 +03:00
Gani Georgiev fcc4e305e0 removed unnecessary large timeout (for reordering the queue 0 should be enough) 2023-07-27 16:10:59 +03:00
Gani Georgiev 854796a8dd [#3000] disallowed relations to views from non-view collections 2023-07-27 15:57:20 +03:00
Gani Georgiev e6a41773ca [#2588] added warning message in case the update command is run in a Docker container or NixOS 2023-07-26 13:18:59 +03:00
Gani Georgiev f4a6d8af49 excluded unnecessary types to reduce the size of the generated declarations file 2023-07-26 10:45:06 +03:00
Gani Georgiev 1563855251 [#2992] added migration to reset already inserted null values 2023-07-26 00:40:48 +03:00
Gani Georgiev 1330e2e1e7 [#2992] fixed zero-default value not being used if the field is not explicitly set when manually creating records 2023-07-25 20:37:19 +03:00
Gani Georgiev 34fe55686d wrapped tests.ApiScenario execution in a subtest 2023-07-25 13:37:43 +03:00
Gani Georgiev b0aa387235 removed extra param unescaping as it was fixed in echo 2023-07-25 13:36:57 +03:00
Gani Georgiev c3f7aeb856 register LoadAuthContext as Pre so that the auth context is aavailable other Pre middlewares 2023-07-25 12:45:41 +03:00
Gani Georgiev 54a6ae6710 updated record.PublicExport comment 2023-07-25 06:11:50 +03:00
Gani Georgiev 8dfc90985b added native echo.HandlerFunc support and .staticDirectoryHandler bind 2023-07-24 21:11:55 +03:00
Gani Georgiev 99ea916c14 renamed expand fetchFunc args to optFetchFunc and updated jsvm types 2023-07-24 16:59:13 +03:00
Gani Georgiev 70151a3c19 added bindings 2023-07-24 16:39:11 +03:00
Gani Georgiev 543fb350ec added jsvm .* helpers 2023-07-24 13:59:13 +03:00
Gani Georgiev ea4e3128ca updated jsvm types 2023-07-24 12:45:23 +03:00
Gani Georgiev ae8cbc8f45 added template.Registry.LoadFS method 2023-07-24 12:33:46 +03:00
Gani Georgiev cb156e1903 increased the default sqlite cache_size to 16mb 2023-07-24 10:35:42 +03:00
Gani Georgiev edcb6950e5 watch pb_hooks subdirectories 2023-07-23 23:45:41 +03:00
Gani Georgiev 085fb1601e added jsvm binding 2023-07-23 16:43:38 +03:00
Gani Georgiev 132a8c0aab added template.Registry.LoadString test 2023-07-23 15:48:01 +03:00
Gani Georgiev 4f3ca6fe2b added helper html template rendering utils 2023-07-23 15:37:30 +03:00
Gani Georgiev 13c0572fe1 updated jsvm types 2023-07-22 19:01:20 +03:00
Gani Georgiev fda4b67dbc updated ui/dist 2023-07-22 18:59:45 +03:00
Gani Georgiev aefbccbfea replaced os.IsNotExists 2023-07-22 18:59:33 +03:00
Gani Georgiev d1336da339 make use of skipTotal 2023-07-22 18:50:40 +03:00
Gani Georgiev f453cefc0b updated go.mod and jsvm types 2023-07-21 23:36:37 +03:00
Gani Georgiev b6bc09fee1 updated jsvm types 2023-07-21 23:29:01 +03:00
Gani Georgiev 437843084b added search skipTotal support 2023-07-21 23:24:36 +03:00
Gani Georgiev 1e4c665b53 [#2957] added support for wrapped api errors 2023-07-20 22:01:58 +03:00
Gani Georgiev ac52befb5b changed subscription.Message.Data to []byte and added client.Send(m) helper 2023-07-20 21:25:13 +03:00
Gani Georgiev 50d7df45eb added ?download file serve query param support to force file download 2023-07-20 15:04:26 +03:00
Gani Georgiev 7e0a4e61b4 updated ui/dist 2023-07-20 14:33:45 +03:00
Gani Georgiev 689ad644c1 updated npm deps 2023-07-20 13:16:16 +03:00
Gani Georgiev f660707712 added e.action to the realtime docs preview 2023-07-20 13:10:23 +03:00
Gani Georgiev 06016722d1 removed legacy 404 check and preserved collections sort order within each type group 2023-07-20 13:06:49 +03:00
Gani Georgiev 832d7f360c updated tinymce 2023-07-20 12:00:23 +03:00
Gani Georgiev 939653ecc0 added after hooks error response tests 2023-07-20 11:42:57 +03:00
Gani Georgiev 610a948dcc added Response.Committed checks 2023-07-20 10:40:03 +03:00
Gani Georgiev b2284b5f4b updated OnModel hooks comment for consistency with the site docs 2023-07-19 18:23:39 +03:00
Gani Georgiev d9e1a759a1 make use of the after hook finalizer 2023-07-18 15:31:36 +03:00
Gani Georgiev 624b443f98 removed unnecessary collection queries 2023-07-18 13:41:14 +03:00
Gani Georgiev 71a70bac9d updated jsvm errors handling 2023-07-18 12:36:04 +03:00
Gani Georgiev 0110869c89 soft deprecated apis.RequestData(c) in favor of apis.RequestInfo(c) and updated jsvm bindings 2023-07-17 23:13:39 +03:00
Gani Georgiev 7d4017225c synced with master 2023-07-17 13:17:17 +03:00
Gani Georgiev 94a1cc07d5 [#2930] added extra normalizations to ensure that newly created multiple fields has the correct zero-default for already inserted records 2023-07-17 11:38:19 +03:00
Gani Georgiev 1720c82570 updated comment 2023-07-17 00:08:06 +03:00
Gani Georgiev 81bd1a1732 reset the requestData Admin and AuthRecord fields 2023-07-17 00:05:15 +03:00
Gani Georgiev f421da4b9b use Dao.CanAccessRecord when checking for protected file access 2023-07-17 00:03:09 +03:00
Gani Georgiev 3eaa3ca1b5 fixed .send binding tests 2023-07-16 23:38:49 +03:00
Gani Georgiev 2d1ad16b4f updated cron jsvm bindings and generated types 2023-07-16 23:24:10 +03:00
Gani Georgiev 6179864828 return the http.Server instance to allow manual shutdowns 2023-07-16 23:13:15 +03:00
Gani Georgiev 64d7ab22f3 treat returned false bool from a jsvm hook as hook.stopPropagation 2023-07-14 16:50:35 +03:00
Gani Georgiev 4962dc618b added record.ExpandedOne(rel) and record.ExpandedAll(rel) helpers 2023-07-14 15:21:59 +03:00
Gani Georgiev 8e2246113a synced with master 2023-07-14 12:44:26 +03:00
Gani Georgiev b9993aaa73 updated changelog 2023-07-14 12:16:56 +03:00
Gani Georgiev f77fb0cc1c updated tests with some clarification code comments 2023-07-14 12:13:44 +03:00
Gani Georgiev 460cc35bb6 updated ui/dist 2023-07-14 12:01:48 +03:00
Gani Georgiev f0bcffec8b [#2914] register the eagerRequestDataCache middleware only for the api grroup to avoid conflicts with custom routes 2023-07-14 11:55:29 +03:00
Gani Georgiev 2b465b0646 load a default fetchFunc for dao.ExpandRecord(s) 2023-07-14 08:36:01 +03:00
Gani Georgiev fdccdcebad added option to call Dao.RecordQuery() with the collection id or name 2023-07-13 22:38:55 +03:00
Gani Georgiev a38bd5bedc tests types.d.ts in gitignore 2023-07-12 17:31:26 +03:00
Gani Georgiev 6fe04bd280 returned OnAfterBootstrap error and added more jsvm tests 2023-07-12 17:12:45 +03:00
Gani Georgiev d0a68da7e7 fixed jsvm docs path 2023-07-11 18:19:33 +03:00
Gani Georgiev ede67dbc20 added jsvm bindings and updateing the workflow to generate the jsvm types 2023-07-11 18:09:55 +03:00
Gani Georgiev 3d3fe5c614 updated Dao.CanAccessRecord to return the invalid filter or db error 2023-07-11 11:50:10 +03:00
Gani Georgiev 7bb33d4453 updated Application URL input label for consistency 2023-07-09 16:44:29 +03:00
Gani Georgiev 0ad4dbc65a synced with master 2023-07-08 21:29:21 +03:00
Gani Georgiev c3844250e8 updated go deps 2023-07-08 20:06:34 +03:00
Gani Georgiev c293994d2b added hooksPool flag and updated doc comments 2023-07-08 20:02:03 +03:00
Gani Georgiev a557aa35f5 updated readme and Find*ByFilter godoc comment 2023-07-08 14:01:02 +03:00
Gani Georgiev 736e9673ad updated generated types 2023-07-08 13:59:24 +03:00
Gani Georgiev 13d96e793b (no tests) updated jsvm bindings 2023-07-08 13:51:16 +03:00
Gani Georgiev 5e37c90dde added cron.Total method 2023-07-08 13:51:00 +03:00
Gani Georgiev 7bcd00a87e synced with master 2023-07-06 23:38:37 +03:00
Gani Georgiev ebfbb55f91 allow no space between the index table name and columns list 2023-07-06 23:20:31 +03:00
Gani Georgiev d77479131a [#2868] fixed unique validator detailed error message not being returned when camelCase field name is used 2023-07-06 23:14:18 +03:00
Gani Georgiev 8ef00efe84 allow no space between the index table name and columns list 2023-07-06 15:47:16 +03:00
Gani Georgiev a4101f7670 synced with master 2023-07-03 20:53:09 +03:00
Gani Georgiev 08b4fc20a9 fixed changelog typo 2023-07-03 19:46:40 +03:00
Gani Georgiev 320f990f84 [#2818] fixed text field regex pattern example 2023-06-30 19:46:00 +03:00
Gani Georgiev 7297f55ca4 [#2817] allowed 0 as RelationOptions.MinSelect value 2023-06-30 18:13:56 +03:00
Gani Georgiev 2cb642bbf7 aliased and soft-deprecated NewToken with NewJWT, added encrypt/decrypt goja bindings and other minor doc changes 2023-06-28 22:56:03 +03:00
Gani Georgiev ecdf9c26cd added comments and typedoc group tags to the generated docs 2023-06-28 21:39:57 +03:00
Gani Georgiev a672ab959f merged jsvm migrations and hooks and updated the ambient TS types location 2023-06-27 14:45:04 +03:00
Gani Georgiev 1571ebe4eb use microseconds when inserting the auto generated migration 2023-06-27 00:35:17 +03:00
Gani Georgiev b8bb5e8d72 fixed migrate down not returning the correct migrations order when the stored applied time is in seconds 2023-06-27 00:33:31 +03:00
Gani Georgiev af77554250 enabled baseBinds for the goja migrations 2023-06-26 23:17:37 +03:00
Gani Georgiev 3b68782cfb synced with master 2023-06-26 18:21:49 +03:00
Gani Georgiev 8388e36f28 updated readme note 2023-06-26 11:03:54 +03:00
Gani Georgiev 051b3702b0 replaced DynamicList with a more generic (model) helper to allow creating pointer slice of any type 2023-06-25 20:19:12 +03:00
Gani Georgiev 39accdba58 updated go deps 2023-06-23 22:26:05 +03:00
Gani Georgiev 32de0aa40a use direct string comparison in the ApiError message test 2023-06-23 22:23:03 +03:00
Gani Georgiev 9bfcdd086a replaced .* errors with constructors and added apisBinds tests 2023-06-23 22:20:13 +03:00
Gani Georgiev 1d20124467 updated changelog 2023-06-23 14:15:53 +03:00
Gani GeorgievandValentine 435eca6f35 [#2762] added Yandex OAuth2 provider
Co-authored-by: Valentine <xb2w1z@gmail.com>
2023-06-23 14:13:43 +03:00
Gani Georgiev 0a61db6efd add @todo note to the RequestData struct 2023-06-23 13:38:13 +03:00
Gani Georgiev 7ab8405946 removed goja middlewares that don't make much sense in the goja context 2023-06-23 12:54:08 +03:00
Gani Georgiev 6fa3e99be2 use inflector.UcFirst instead of strings.Title 2023-06-22 21:54:39 +03:00
Gani Georgiev 3160fb2d99 added DynamicModel form tag and removed unused helper 2023-06-22 16:32:21 +03:00
Gani Georgiev 1cbf16b3bf ucfirst the DynamicModel field name so that we can use later the same FieldMapper resolver rules 2023-06-22 16:29:58 +03:00
Gani Georgiev dad289b90d bind hooksWatch flag 2023-06-21 21:46:13 +03:00
Gani Georgiev c795ecd21e updated jsvm generated types 2023-06-21 20:40:43 +03:00
Gani Georgiev 21607f0f66 updated cobra.Command constructor and update structConstructor to use goja.Object.Set 2023-06-21 20:36:57 +03:00
Gani Georgiev fc311a8d28 removed the temp len binding as the issue was fixed in goja#521 2023-06-21 13:40:59 +03:00
Gani Georgiev 93606c6647 added DynamicModel and DynamicList goja bindings 2023-06-21 11:21:54 +03:00
Gani Georgiev 1adcfcc03b adde json map Get and Set helpers 2023-06-20 22:57:51 +03:00
Gani Georgiev ed4304dc30 added jsvm typings and docs generation 2023-06-20 08:54:02 +03:00
Gani Georgiev c0a6a21b9e updated code comments and added some notes 2023-06-19 21:45:45 +03:00
Gani Georgiev a7bb599cd0 Merge branch 'master' into develop 2023-06-16 14:48:52 +03:00
Gani Georgiev 28ba4655c1 added note about the disk space when creating backups 2023-06-16 12:24:28 +03:00
Gani Georgiev bd95a5b74c [#2693] removed the implicit autosnapshot migration creation as it is not clear to the users when it happens 2023-06-14 13:21:37 +03:00
Gani Georgiev 745b230097 updated jsvm mapper and updated godoc formatting 2023-06-14 13:14:30 +03:00
Gani Georgiev ec303a60ed [#2271] added dao.CanAccessRecord() helper 2023-06-14 13:13:21 +03:00
Gani Georgiev e99b0627d6 synced with master 2023-06-10 23:43:09 +03:00
Gani Georgiev b77b6d1a18 updated ui/dist 2023-06-09 13:45:42 +03:00
Gani Georgiev 779b23d919 synced with master 2023-06-09 13:33:38 +03:00
Gani Georgiev a5b27cce5c fixed apis.NewUnauthorizedError test 2023-06-08 18:16:00 +03:00
Gani Georgiev ebd6891471 updated broken tests 2023-06-08 18:14:01 +03:00
Gani Georgiev 3cf3e04866 restructered some of the internals and added basic js app hooks support 2023-06-08 17:59:08 +03:00
Gani Georgiev ff5508cb79 synced with master 2023-06-02 19:53:53 +03:00
Gani Georgiev f07f7a1e35 synced with master 2023-06-02 19:38:36 +03:00
Gani Georgiev 881b625177 updated changelog 2023-06-02 19:23:43 +03:00
Gani Georgiev 4c2dcac61a added dao.WithoutHooks() helper 2023-06-01 15:42:38 +03:00
Gani Georgiev dcb00a3917 updated changelog 2023-05-31 21:51:01 +03:00
Gani Georgiev ddca49ba16 [#2309] added query by filter record helpers 2023-05-31 11:49:16 +03:00
Gani Georgiev 0fb92720f8 Merge branch 'master' into develop 2023-05-30 21:22:39 +03:00
Gani Georgiev 7de346b532 fixed realtime delete event to be called after the record was deleted from the db 2023-05-29 22:28:07 +03:00
Gani Georgiev 729f9f142e check after hook errors 2023-05-29 21:50:07 +03:00
Gani Georgiev 45b73e3dfb Merge branch 'master' into develop 2023-05-29 17:01:16 +03:00
Gani Georgiev d3711b0503 added new core.ServeEvent fields 2023-05-29 16:57:50 +03:00
Gani Georgiev 9d8df8d05d added option to remove single registered hook handler 2023-05-29 14:51:03 +03:00
Gani Georgiev 97f29e4305 synced with master 2023-05-28 23:57:19 +03:00
Gani Georgiev fcfcaa0628 refresh the cached logged admin and auth record 2023-05-28 17:36:56 +03:00
Gani Georgiev d5314b028b synced with master 2023-05-27 14:12:34 +03:00
Gani Georgiev 3be5875ea9 Merge branch 'master' into develop 2023-05-25 21:00:45 +03:00
Gani Georgiev 94680c41f7 synced with master 2023-05-24 23:31:31 +03:00
Gani GeorgievandValentine af71b63f23 [#2533] added VK OAuth2 provider
Co-authored-by: Valentine <xb2w1z@gmail.com>
2023-05-24 15:41:58 +03:00
Gani Georgiev e40cf46b33 synced with master 2023-05-24 11:07:29 +03:00
Gani Georgiev 5b330ab5b4 updated compareVersions tests 2023-05-23 23:15:56 +03:00
Gani Georgiev 7dcfa65146 updated ui/dist 2023-05-23 22:47:04 +03:00
Gani GeorgievandPedro Costa a6bb1bf096 [#2534] added Instagram OAuth2 provider
Co-authored-by: Pedro Costa <550684+pnmcosta@users.noreply.github.com>
2023-05-23 22:37:44 +03:00
Gani Georgiev 728427cecf Merge branch 'master' into develop 2023-05-23 21:56:23 +03:00
Gani Georgiev 651b439096 synced with master 2023-05-23 11:36:55 +03:00
Gani Georgiev 7a41ff0127 added version check todo 2023-05-23 10:46:54 +03:00
626 changed files with 52071 additions and 50957 deletions
-13
View File
@@ -1,13 +0,0 @@
# These are supported funding model platforms
github: # Replace with up to 4 GitHub Sponsors-enabled usernames e.g., [user1, user2]
patreon: # Replace with a single Patreon username
open_collective: pocketbase
ko_fi: # Replace with a single Ko-fi username
tidelift: # Replace with a single Tidelift platform-name/package-name e.g., npm/babel
community_bridge: # Replace with a single Community Bridge project-name e.g., cloud-foundry
liberapay: # Replace with a single Liberapay username
issuehunt: # Replace with a single IssueHunt username
otechie: # Replace with a single Otechie username
lfx_crowdfunding: # Replace with a single LFX Crowdfunding project-name e.g., cloud-foundry
custom: ['https://www.paypal.com/donate/?hosted_button_id=4DVXNL4B8WT98']
+12 -5
View File
@@ -9,19 +9,19 @@ jobs:
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- name: Checkout - name: Checkout
uses: actions/checkout@v3 uses: actions/checkout@v4
with: with:
fetch-depth: 0 fetch-depth: 0
- name: Set up Node.js - name: Set up Node.js
uses: actions/setup-node@v3 uses: actions/setup-node@v4
with: with:
node-version: latest node-version: 20.11.0
- name: Set up Go - name: Set up Go
uses: actions/setup-go@v3 uses: actions/setup-go@v5
with: with:
go-version: '>=1.20.3' go-version: '>=1.22.3'
# This step usually is not needed because the /ui/dist is pregenerated locally # This step usually is not needed because the /ui/dist is pregenerated locally
# but its here to ensure that each release embeds the latest admin ui artifacts. # but its here to ensure that each release embeds the latest admin ui artifacts.
@@ -29,6 +29,13 @@ jobs:
- name: Build Admin dashboard UI - name: Build Admin dashboard UI
run: npm --prefix=./ui ci && npm --prefix=./ui run build run: npm --prefix=./ui ci && npm --prefix=./ui run build
# Temporary disable as the types can have random generated identifiers making it non-deterministic.
#
# # Similar to the above, the jsvm types are pregenerated locally
# # but its here to ensure that it wasn't forgotten to be executed.
# - name: Generate jsvm types
# run: go run ./plugins/jsvm/internal/types/types.go
# The prebuilt golangci-lint doesn't support go 1.18+ yet # The prebuilt golangci-lint doesn't support go 1.18+ yet
# https://github.com/golangci/golangci-lint/issues/2649 # https://github.com/golangci/golangci-lint/issues/2649
# - name: Run linter # - name: Run linter
+3 -10
View File
@@ -7,6 +7,7 @@ before:
- go mod tidy - go mod tidy
builds: builds:
# used only for tests
- id: build_cgo - id: build_cgo
main: ./examples/base main: ./examples/base
binary: pocketbase binary: pocketbase
@@ -46,20 +47,12 @@ release:
draft: true draft: true
archives: archives:
- id: archive_cgo
builds: [build_cgo]
name_template: '{{ .ProjectName }}_{{ .Version }}_{{ .Os }}_{{ .Arch }}_cgo'
format: zip
files:
- LICENSE*
- CHANGELOG*
- id: archive_noncgo - id: archive_noncgo
builds: [build_noncgo] builds: [build_noncgo]
format: zip format: zip
files: files:
- LICENSE* - LICENSE.md
- CHANGELOG* - CHANGELOG.md
checksum: checksum:
name_template: 'checksums.txt' name_template: 'checksums.txt'
+901 -1384
View File
File diff suppressed because it is too large Load Diff
+1384
View File
File diff suppressed because it is too large Load Diff
+22 -20
View File
@@ -1,5 +1,4 @@
Contributing to PocketBase # Contributing to PocketBase
======================================================================
Thanks for taking the time to improve PocketBase! Thanks for taking the time to improve PocketBase!
@@ -9,26 +8,26 @@ This document describes how to prepare a PR for a change in the main repository.
- [Making changes in the Go code](#making-changes-in-the-go-code) - [Making changes in the Go code](#making-changes-in-the-go-code)
- [Making changes in the Admin UI](#making-changes-in-the-admin-ui) - [Making changes in the Admin UI](#making-changes-in-the-admin-ui)
## Prerequisites ## Prerequisites
- Go 1.18+ (for making changes in the Go code) - Go 1.21+ (for making changes in the Go code)
- Node 16+ (for making changes in the Admin UI) - Node 18+ (for making changes in the Admin UI)
If you haven't already, you can fork the main repository and clone your fork so that you can work locally: If you haven't already, you can fork the main repository and clone your fork so that you can work locally:
``` ```
git clone https://github.com/your_username/pocketbase.git git clone https://github.com/your_username/pocketbase.git
``` ```
> [!IMPORTANT]
> It is recommended to create a new branch from master for each of your bugfixes and features. > It is recommended to create a new branch from master for each of your bugfixes and features.
> This is required if you are planning to submit multiple PRs in order to keep the changes separate for review until they eventually get merged. > This is required if you are planning to submit multiple PRs in order to keep the changes separate for review until they eventually get merged.
## Making changes in the Go code ## Making changes in the Go code
PocketBase is distributed as a Go package, which means that in order to run the project you'll have to create a Go `main` program that imports the package. PocketBase is distributed as a Go package, which means that in order to run the project you'll have to create a Go `main` program that imports the package.
The repository already includes such program, located in `/examples/base`, that is also used for the prebuilt executables. The repository already includes such program, located in `examples/base`, that is also used for the prebuilt executables.
So, let's assume that you already done some changes in the PocketBase Go code and you want now to run them: So, let's assume that you already done some changes in the PocketBase Go code and you want now to run them:
@@ -41,20 +40,22 @@ This will start a web server on `http://localhost:8090` with the embedded prebui
- Add unit/integration tests for your changes (we are using the standard `testing` go package). - Add unit/integration tests for your changes (we are using the standard `testing` go package).
To run the tests, you could execute (while in the root project directory): To run the tests, you could execute (while in the root project directory):
```sh
go test ./...
# or using the Makefile ```sh
make test go test ./...
```
# or using the Makefile
make test
```
- Run the linter - **golangci** ([see how to install](https://golangci-lint.run/usage/install/#local-installation)): - Run the linter - **golangci** ([see how to install](https://golangci-lint.run/usage/install/#local-installation)):
```sh
golangci-lint run -c ./golangci.yml ./...
# or using the Makefile ```sh
make lint golangci-lint run -c ./golangci.yml ./...
```
# or using the Makefile
make lint
```
## Making changes in the Admin UI ## Making changes in the Admin UI
@@ -65,14 +66,15 @@ To start the Admin UI:
1. Navigate to the `ui` project directory 1. Navigate to the `ui` project directory
2. Run `npm install` to install the node dependencies 2. Run `npm install` to install the node dependencies
3. Start vite's dev server 3. Start vite's dev server
```sh ```sh
npm run dev npm run dev
``` ```
You could open the browser and access the running Admin UI at `http://localhost:3000`. You could open the browser and access the running Admin UI at `http://localhost:3000`.
Since the Admin UI is just a client-side application, you need to have the PocketBase backend server also running in the background (either manually running the `examples/base/main.go` or download a prebuilt executable). Since the Admin UI is just a client-side application, you need to have the PocketBase backend server also running in the background (either manually running the `examples/base/main.go` or download a prebuilt executable).
> [!NOTE]
> By default, the Admin UI is expecting the backend server to be started at `http://localhost:8090`, but you could change that by creating a new `ui/.env.development.local` file with `PB_BACKEND_URL = YOUR_ADDRESS` variable inside it. > By default, the Admin UI is expecting the backend server to be started at `http://localhost:8090`, but you could change that by creating a new `ui/.env.development.local` file with `PB_BACKEND_URL = YOUR_ADDRESS` variable inside it.
Every change you make in the Admin UI should be automatically reflected in the browser at `http://localhost:3000` without reloading the page. Every change you make in the Admin UI should be automatically reflected in the browser at `http://localhost:3000` without reloading the page.
+3
View File
@@ -4,6 +4,9 @@ lint:
test: test:
go test ./... -v --cover go test ./... -v --cover
jstypes:
go run ./plugins/jsvm/internal/types/types.go
test-report: test-report:
go test ./... -v --cover -coverprofile=coverage.out go test ./... -v --cover -coverprofile=coverage.out
go tool cover -html=coverage.out go tool cover -html=coverage.out
+67 -60
View File
@@ -7,7 +7,7 @@
<p align="center"> <p align="center">
<a href="https://github.com/pocketbase/pocketbase/actions/workflows/release.yaml" target="_blank" rel="noopener"><img src="https://github.com/pocketbase/pocketbase/actions/workflows/release.yaml/badge.svg" alt="build" /></a> <a href="https://github.com/pocketbase/pocketbase/actions/workflows/release.yaml" target="_blank" rel="noopener"><img src="https://github.com/pocketbase/pocketbase/actions/workflows/release.yaml/badge.svg" alt="build" /></a>
<a href="https://github.com/pocketbase/pocketbase/releases" target="_blank" rel="noopener"><img src="https://img.shields.io/github/release/pocketbase/pocketbase.svg" alt="Latest releases" /></a> <a href="https://github.com/pocketbase/pocketbase/releases" target="_blank" rel="noopener"><img src="https://img.shields.io/github/release/pocketbase/pocketbase.svg" alt="Latest releases" /></a>
<a href="https://pkg.go.dev/github.com/pocketbase/pocketbase" target="_blank" rel="noopener"><img src="https://godoc.org/github.com/ganigeorgiev/fexpr?status.svg" alt="Go package documentation" /></a> <a href="https://pkg.go.dev/github.com/pocketbase/pocketbase" target="_blank" rel="noopener"><img src="https://godoc.org/github.com/pocketbase/pocketbase?status.svg" alt="Go package documentation" /></a>
</p> </p>
[PocketBase](https://pocketbase.io) is an open source Go backend, consisting of: [PocketBase](https://pocketbase.io) is an open source Go backend, consisting of:
@@ -19,10 +19,10 @@
**For documentation and examples, please visit https://pocketbase.io/docs.** **For documentation and examples, please visit https://pocketbase.io/docs.**
> ⚠️ Please keep in mind that PocketBase is still under active development > [!WARNING]
> Please keep in mind that PocketBase is still under active development
> and therefore full backward compatibility is not guaranteed before reaching v1.0.0. > and therefore full backward compatibility is not guaranteed before reaching v1.0.0.
## API SDK clients ## API SDK clients
The easiest way to interact with the API is to use one of the official SDK clients: The easiest way to interact with the API is to use one of the official SDK clients:
@@ -30,79 +30,90 @@ The easiest way to interact with the API is to use one of the official SDK clien
- **JavaScript - [pocketbase/js-sdk](https://github.com/pocketbase/js-sdk)** (_browser and node_) - **JavaScript - [pocketbase/js-sdk](https://github.com/pocketbase/js-sdk)** (_browser and node_)
- **Dart - [pocketbase/dart-sdk](https://github.com/pocketbase/dart-sdk)** (_web, mobile, desktop_) - **Dart - [pocketbase/dart-sdk](https://github.com/pocketbase/dart-sdk)** (_web, mobile, desktop_)
## Overview ## Overview
PocketBase could be [downloaded directly as a standalone app](https://github.com/pocketbase/pocketbase/releases) or it could be used as a Go framework/toolkit which allows you to build ### Use as standalone app
You could download the prebuilt executable for your platform from the [Releases page](https://github.com/pocketbase/pocketbase/releases).
Once downloaded, extract the archive and run `./pocketbase serve` in the extracted directory.
The prebuilt executables are based on the [`examples/base/main.go` file](https://github.com/pocketbase/pocketbase/blob/master/examples/base/main.go) and comes with the JS VM plugin enabled by default which allows to extend PocketBase with JavaScript (_for more details please refer to [Extend with JavaScript](https://pocketbase.io/docs/js-overview/)_).
### Use as a Go framework/toolkit
PocketBase is distributed as a regular Go library package which allows you to build
your own custom app specific business logic and still have a single portable executable at the end. your own custom app specific business logic and still have a single portable executable at the end.
### Installation Here is a minimal example:
```sh 0. [Install Go 1.21+](https://go.dev/doc/install) (_if you haven't already_)
# go 1.18+
go get github.com/pocketbase/pocketbase
```
> For Windows, you may have to use go 1.19+ due to an incorrect js mime type in the Windows Registry (see [issue#6](https://github.com/pocketbase/pocketbase/issues/6)).
### Example 1. Create a new project directory with the following `main.go` file inside it:
```go
package main
```go import (
package main "log"
"net/http"
import ( "github.com/labstack/echo/v5"
"log" "github.com/pocketbase/pocketbase"
"net/http" "github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/core"
)
"github.com/labstack/echo/v5" func main() {
"github.com/pocketbase/pocketbase" app := pocketbase.New()
"github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/core"
)
func main() { app.OnBeforeServe().Add(func(e *core.ServeEvent) error {
app := pocketbase.New() // add new "GET /hello" route to the app router (echo)
e.Router.AddRoute(echo.Route{
Method: http.MethodGet,
Path: "/hello",
Handler: func(c echo.Context) error {
return c.String(200, "Hello world!")
},
Middlewares: []echo.MiddlewareFunc{
apis.ActivityLogger(app),
},
})
app.OnBeforeServe().Add(func(e *core.ServeEvent) error { return nil
// add new "GET /hello" route to the app router (echo)
e.Router.AddRoute(echo.Route{
Method: http.MethodGet,
Path: "/hello",
Handler: func(c echo.Context) error {
return c.String(200, "Hello world!")
},
Middlewares: []echo.MiddlewareFunc{
apis.ActivityLogger(app),
},
}) })
return nil if err := app.Start(); err != nil {
}) log.Fatal(err)
}
if err := app.Start(); err != nil {
log.Fatal(err)
} }
} ```
```
### Running and building 2. To init the dependencies, run `go mod init myapp && go mod tidy`.
Running/building the application is the same as for any other Go program, aka. just `go run` and `go build`. 3. To start the application, run `go run main.go serve`.
**PocketBase embeds SQLite, but doesn't require CGO.** 4. To build a statically linked executable, you can run `CGO_ENABLED=0 go build` and then start the created executable with `./myapp serve`.
If CGO is enabled (aka. `CGO_ENABLED=1`), it will use [mattn/go-sqlite3](https://pkg.go.dev/github.com/mattn/go-sqlite3) driver, otherwise - [modernc.org/sqlite](https://pkg.go.dev/modernc.org/sqlite). > [!NOTE]
Enable CGO only if you really need to squeeze the read/write query performance at the expense of complicating cross compilation. > PocketBase embeds SQLite, but doesn't require CGO.
>
> If CGO is enabled (aka. `CGO_ENABLED=1`), it will use [mattn/go-sqlite3](https://pkg.go.dev/github.com/mattn/go-sqlite3) driver, otherwise - [modernc.org/sqlite](https://pkg.go.dev/modernc.org/sqlite).
> Enable CGO only if you really need to squeeze the read/write query performance at the expense of complicating cross compilation.
_For more details please refer to [Extend with Go](https://pocketbase.io/docs/go-overview/)._
### Building and running the repo main.go example
To build the minimal standalone executable, like the prebuilt ones in the releases page, you can simply run `go build` inside the `examples/base` directory: To build the minimal standalone executable, like the prebuilt ones in the releases page, you can simply run `go build` inside the `examples/base` directory:
0. [Install Go 1.18+](https://go.dev/doc/install) (_if you haven't already_) 0. [Install Go 1.21+](https://go.dev/doc/install) (_if you haven't already_)
1. Clone/download the repo 1. Clone/download the repo
2. Navigate to `examples/base` 2. Navigate to `examples/base`
3. Run `GOOS=linux GOARCH=amd64 CGO_ENABLED=0 go build` 3. Run `GOOS=linux GOARCH=amd64 CGO_ENABLED=0 go build`
(_https://go.dev/doc/install/source#environment_) (_https://go.dev/doc/install/source#environment_)
4. Start the created executable by running `./base serve`. 4. Start the created executable by running `./base serve`.
The supported build targets by the non-cgo driver at the moment are: Note that the supported build targets by the pure Go SQLite driver at the moment are:
``` ```
darwin amd64 darwin amd64
darwin arm64 darwin arm64
@@ -114,6 +125,7 @@ linux arm
linux arm64 linux arm64
linux ppc64le linux ppc64le
linux riscv64 linux riscv64
linux s390x
windows amd64 windows amd64
windows arm64 windows arm64
``` ```
@@ -121,7 +133,8 @@ windows arm64
### Testing ### Testing
PocketBase comes with mixed bag of unit and integration tests. PocketBase comes with mixed bag of unit and integration tests.
To run them, use the default `go test` command: To run them, use the standard `go test` command:
```sh ```sh
go test ./... go test ./...
``` ```
@@ -134,7 +147,6 @@ If you discover a security vulnerability within PocketBase, please send an e-mai
All reports will be promptly addressed, and you'll be credited accordingly. All reports will be promptly addressed, and you'll be credited accordingly.
## Contributing ## Contributing
PocketBase is free and open source project licensed under the [MIT License](LICENSE.md). PocketBase is free and open source project licensed under the [MIT License](LICENSE.md).
@@ -144,16 +156,11 @@ You could help continuing its development by:
- [Contribute to the source code](CONTRIBUTING.md) - [Contribute to the source code](CONTRIBUTING.md)
- [Suggest new features and report issues](https://github.com/pocketbase/pocketbase/issues) - [Suggest new features and report issues](https://github.com/pocketbase/pocketbase/issues)
- [Donate a small amount](https://pocketbase.io/support-us)
PRs for _small features_ (eg. adding new OAuth2 providers), bug and documentation fixes, etc. are more than welcome. PRs for new OAuth2 providers, bug fixes, code optimizations and documentation improvements are more than welcome.
But please refrain creating PRs for _big features_ without previously discussing the implementation details. Reviewing big PRs often requires a lot of time and tedious back-and-forth communication. But please refrain creating PRs for _new features_ without previously discussing the implementation details.
PocketBase has a [roadmap](https://github.com/orgs/pocketbase/projects/2) PocketBase has a [roadmap](https://github.com/orgs/pocketbase/projects/2) and I try to work on issues in specific order and such PRs often come in out of nowhere and skew all initial planning with tedious back-and-forth communication.
and I try to work on issues in a specific order and such PRs often come in out of nowhere and skew all initial planning.
Don't get upset if I close your PR, even if it is well executed and tested. This doesn't mean that it will never be merged. Don't get upset if I close your PR, even if it is well executed and tested. This doesn't mean that it will never be merged.
Later we can always refer to it and/or take pieces of your implementation when the time comes to work on the issue (don't worry you'll be credited in the release notes). Later we can always refer to it and/or take pieces of your implementation when the time comes to work on the issue (don't worry you'll be credited in the release notes).
_Please also note that PocketBase was initially created to serve as a new backend for my other open source project - [Presentator](https://presentator.io) (see [#183](https://github.com/presentator/presentator/issues/183)),
so all feature requests will be first aligned with what we need for Presentator v3._
+67 -63
View File
@@ -1,7 +1,6 @@
package apis package apis
import ( import (
"log"
"net/http" "net/http"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
@@ -33,18 +32,28 @@ type adminApi struct {
app core.App app core.App
} }
func (api *adminApi) authResponse(c echo.Context, admin *models.Admin) error { func (api *adminApi) authResponse(c echo.Context, admin *models.Admin, finalizers ...func(token string) error) error {
token, tokenErr := tokens.NewAdminAuthToken(api.app, admin) token, tokenErr := tokens.NewAdminAuthToken(api.app, admin)
if tokenErr != nil { if tokenErr != nil {
return NewBadRequestError("Failed to create auth token.", tokenErr) return NewBadRequestError("Failed to create auth token.", tokenErr)
} }
for _, f := range finalizers {
if err := f(token); err != nil {
return err
}
}
event := new(core.AdminAuthEvent) event := new(core.AdminAuthEvent)
event.HttpContext = c event.HttpContext = c
event.Admin = admin event.Admin = admin
event.Token = token event.Token = token
return api.app.OnAdminAuthRequest().Trigger(event, func(e *core.AdminAuthEvent) error { return api.app.OnAdminAuthRequest().Trigger(event, func(e *core.AdminAuthEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(200, map[string]any{ return e.HttpContext.JSON(200, map[string]any{
"token": e.Token, "token": e.Token,
"admin": e.Admin, "admin": e.Admin,
@@ -62,17 +71,11 @@ func (api *adminApi) authRefresh(c echo.Context) error {
event.HttpContext = c event.HttpContext = c
event.Admin = admin event.Admin = admin
handlerErr := api.app.OnAdminBeforeAuthRefreshRequest().Trigger(event, func(e *core.AdminAuthRefreshEvent) error { return api.app.OnAdminBeforeAuthRefreshRequest().Trigger(event, func(e *core.AdminAuthRefreshEvent) error {
return api.authResponse(e.HttpContext, e.Admin) return api.app.OnAdminAfterAuthRefreshRequest().Trigger(event, func(e *core.AdminAuthRefreshEvent) error {
return api.authResponse(e.HttpContext, e.Admin)
})
}) })
if handlerErr == nil {
if err := api.app.OnAdminAfterAuthRefreshRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return handlerErr
} }
func (api *adminApi) authWithPassword(c echo.Context) error { func (api *adminApi) authWithPassword(c echo.Context) error {
@@ -95,17 +98,13 @@ func (api *adminApi) authWithPassword(c echo.Context) error {
return NewBadRequestError("Failed to authenticate.", err) return NewBadRequestError("Failed to authenticate.", err)
} }
return api.authResponse(e.HttpContext, e.Admin) return api.app.OnAdminAfterAuthWithPasswordRequest().Trigger(event, func(e *core.AdminAuthWithPasswordEvent) error {
return api.authResponse(e.HttpContext, e.Admin)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnAdminAfterAuthWithPasswordRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -129,30 +128,29 @@ func (api *adminApi) requestPasswordReset(c echo.Context) error {
return api.app.OnAdminBeforeRequestPasswordResetRequest().Trigger(event, func(e *core.AdminRequestPasswordResetEvent) error { return api.app.OnAdminBeforeRequestPasswordResetRequest().Trigger(event, func(e *core.AdminRequestPasswordResetEvent) error {
// run in background because we don't need to show the result to the client // run in background because we don't need to show the result to the client
routine.FireAndForget(func() { routine.FireAndForget(func() {
if err := next(e.Admin); err != nil && api.app.IsDebug() { if err := next(e.Admin); err != nil {
log.Println(err) api.app.Logger().Error("Failed to send admin password reset request.", "error", err)
} }
}) })
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnAdminAfterRequestPasswordResetRequest().Trigger(event, func(e *core.AdminRequestPasswordResetEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil { // eagerly write 204 response and skip submit errors
if err := api.app.OnAdminAfterRequestPasswordResetRequest().Trigger(event); err != nil && api.app.IsDebug() { // as a measure against admins enumeration
log.Println(err)
}
} else if api.app.IsDebug() {
log.Println(submitErr)
}
// don't return the response error to prevent emails enumeration
if !c.Response().Committed { if !c.Response().Committed {
c.NoContent(http.StatusNoContent) c.NoContent(http.StatusNoContent)
} }
return nil return submitErr
} }
func (api *adminApi) confirmPasswordReset(c echo.Context) error { func (api *adminApi) confirmPasswordReset(c echo.Context) error {
@@ -173,17 +171,17 @@ func (api *adminApi) confirmPasswordReset(c echo.Context) error {
return NewBadRequestError("Failed to set new password.", err) return NewBadRequestError("Failed to set new password.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnAdminAfterConfirmPasswordResetRequest().Trigger(event, func(e *core.AdminConfirmPasswordResetEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnAdminAfterConfirmPasswordResetRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -208,6 +206,10 @@ func (api *adminApi) list(c echo.Context) error {
event.Result = result event.Result = result
return api.app.OnAdminsListRequest().Trigger(event, func(e *core.AdminsListEvent) error { return api.app.OnAdminsListRequest().Trigger(event, func(e *core.AdminsListEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Result) return e.HttpContext.JSON(http.StatusOK, e.Result)
}) })
} }
@@ -228,6 +230,10 @@ func (api *adminApi) view(c echo.Context) error {
event.Admin = admin event.Admin = admin
return api.app.OnAdminViewRequest().Trigger(event, func(e *core.AdminViewEvent) error { return api.app.OnAdminViewRequest().Trigger(event, func(e *core.AdminViewEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Admin) return e.HttpContext.JSON(http.StatusOK, e.Admin)
}) })
} }
@@ -256,17 +262,17 @@ func (api *adminApi) create(c echo.Context) error {
return NewBadRequestError("Failed to create admin.", err) return NewBadRequestError("Failed to create admin.", err)
} }
return e.HttpContext.JSON(http.StatusOK, e.Admin) return api.app.OnAdminAfterCreateRequest().Trigger(event, func(e *core.AdminCreateEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Admin)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnAdminAfterCreateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -302,17 +308,17 @@ func (api *adminApi) update(c echo.Context) error {
return NewBadRequestError("Failed to update admin.", err) return NewBadRequestError("Failed to update admin.", err)
} }
return e.HttpContext.JSON(http.StatusOK, e.Admin) return api.app.OnAdminAfterUpdateRequest().Trigger(event, func(e *core.AdminUpdateEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Admin)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnAdminAfterUpdateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -331,19 +337,17 @@ func (api *adminApi) delete(c echo.Context) error {
event.HttpContext = c event.HttpContext = c
event.Admin = admin event.Admin = admin
handlerErr := api.app.OnAdminBeforeDeleteRequest().Trigger(event, func(e *core.AdminDeleteEvent) error { return api.app.OnAdminBeforeDeleteRequest().Trigger(event, func(e *core.AdminDeleteEvent) error {
if err := api.app.Dao().DeleteAdmin(e.Admin); err != nil { if err := api.app.Dao().DeleteAdmin(e.Admin); err != nil {
return NewBadRequestError("Failed to delete admin.", err) return NewBadRequestError("Failed to delete admin.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnAdminAfterDeleteRequest().Trigger(event, func(e *core.AdminDeleteEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
if handlerErr == nil {
if err := api.app.OnAdminAfterDeleteRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return handlerErr
} }
+158
View File
@@ -1,6 +1,7 @@
package apis_test package apis_test
import ( import (
"errors"
"net/http" "net/http"
"strings" "strings"
"testing" "testing"
@@ -8,6 +9,7 @@ import (
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
@@ -15,6 +17,8 @@ import (
) )
func TestAdminAuthWithPassword(t *testing.T) { func TestAdminAuthWithPassword(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "empty data", Name: "empty data",
@@ -89,6 +93,26 @@ func TestAdminAuthWithPassword(t *testing.T) {
"OnAdminAuthRequest": 1, "OnAdminAuthRequest": 1,
}, },
}, },
{
Name: "OnAdminAfterAuthWithPasswordRequest error response",
Method: http.MethodPost,
Url: "/api/admins/auth-with-password",
Body: strings.NewReader(`{"identity":"test@example.com","password":"1234567890"}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4MTYwMH0.han3_sG65zLddpcX2ic78qgy7FKecuPfOpFa8Dvi5Bg",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnAdminAfterAuthWithPasswordRequest().Add(func(e *core.AdminAuthWithPasswordEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnAdminBeforeAuthWithPasswordRequest": 1,
"OnAdminAfterAuthWithPasswordRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -97,6 +121,8 @@ func TestAdminAuthWithPassword(t *testing.T) {
} }
func TestAdminRequestPasswordReset(t *testing.T) { func TestAdminRequestPasswordReset(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "empty data", Name: "empty data",
@@ -166,6 +192,8 @@ func TestAdminRequestPasswordReset(t *testing.T) {
} }
func TestAdminConfirmPasswordReset(t *testing.T) { func TestAdminConfirmPasswordReset(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "empty data", Name: "empty data",
@@ -224,6 +252,29 @@ func TestAdminConfirmPasswordReset(t *testing.T) {
"OnAdminAfterConfirmPasswordResetRequest": 1, "OnAdminAfterConfirmPasswordResetRequest": 1,
}, },
}, },
{
Name: "OnAdminAfterConfirmPasswordResetRequest error response",
Method: http.MethodPost,
Url: "/api/admins/confirm-password-reset",
Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImV4cCI6MjIwODk4MTYwMH0.kwFEler6KSMKJNstuaSDvE1QnNdCta5qSnjaIQ0hhhc",
"password":"1234567891",
"passwordConfirm":"1234567891"
}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnAdminAfterConfirmPasswordResetRequest().Add(func(e *core.AdminConfirmPasswordResetEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeUpdate": 1,
"OnModelAfterUpdate": 1,
"OnAdminBeforeConfirmPasswordResetRequest": 1,
"OnAdminAfterConfirmPasswordResetRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -232,6 +283,8 @@ func TestAdminConfirmPasswordReset(t *testing.T) {
} }
func TestAdminRefresh(t *testing.T) { func TestAdminRefresh(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -278,6 +331,25 @@ func TestAdminRefresh(t *testing.T) {
"OnAdminAfterAuthRefreshRequest": 1, "OnAdminAfterAuthRefreshRequest": 1,
}, },
}, },
{
Name: "OnAdminAfterAuthRefreshRequest error response",
Method: http.MethodPost,
Url: "/api/admins/auth-refresh",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnAdminAfterAuthRefreshRequest().Add(func(e *core.AdminAuthRefreshEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnAdminBeforeAuthRefreshRequest": 1,
"OnAdminAfterAuthRefreshRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -286,6 +358,8 @@ func TestAdminRefresh(t *testing.T) {
} }
func TestAdminsList(t *testing.T) { func TestAdminsList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -389,6 +463,8 @@ func TestAdminsList(t *testing.T) {
} }
func TestAdminView(t *testing.T) { func TestAdminView(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -444,6 +520,8 @@ func TestAdminView(t *testing.T) {
} }
func TestAdminDelete(t *testing.T) { func TestAdminDelete(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -510,6 +588,27 @@ func TestAdminDelete(t *testing.T) {
"OnAdminBeforeDeleteRequest": 1, "OnAdminBeforeDeleteRequest": 1,
}, },
}, },
{
Name: "OnAdminAfterDeleteRequest error response",
Method: http.MethodDelete,
Url: "/api/admins/sbmbsdb40jyxf7h",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnAdminAfterDeleteRequest().Add(func(e *core.AdminDeleteEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeDelete": 1,
"OnModelAfterDelete": 1,
"OnAdminBeforeDeleteRequest": 1,
"OnAdminAfterDeleteRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -518,6 +617,8 @@ func TestAdminDelete(t *testing.T) {
} }
func TestAdminCreate(t *testing.T) { func TestAdminCreate(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized (while having at least 1 existing admin)", Name: "unauthorized (while having at least 1 existing admin)",
@@ -637,6 +738,33 @@ func TestAdminCreate(t *testing.T) {
"OnAdminAfterCreateRequest": 1, "OnAdminAfterCreateRequest": 1,
}, },
}, },
{
Name: "OnAdminAfterCreateRequest error response",
Method: http.MethodPost,
Url: "/api/admins",
Body: strings.NewReader(`{
"email":"testnew@example.com",
"password":"1234567890",
"passwordConfirm":"1234567890",
"avatar":3
}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnAdminAfterCreateRequest().Add(func(e *core.AdminCreateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeCreate": 1,
"OnModelAfterCreate": 1,
"OnAdminBeforeCreateRequest": 1,
"OnAdminAfterCreateRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -645,6 +773,8 @@ func TestAdminCreate(t *testing.T) {
} }
func TestAdminUpdate(t *testing.T) { func TestAdminUpdate(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -729,6 +859,7 @@ func TestAdminUpdate(t *testing.T) {
}, },
}, },
{ {
Name: "authorized as admin + valid data",
Method: http.MethodPatch, Method: http.MethodPatch,
Url: "/api/admins/sbmbsdb40jyxf7h", Url: "/api/admins/sbmbsdb40jyxf7h",
Body: strings.NewReader(`{ Body: strings.NewReader(`{
@@ -759,6 +890,33 @@ func TestAdminUpdate(t *testing.T) {
"OnAdminAfterUpdateRequest": 1, "OnAdminAfterUpdateRequest": 1,
}, },
}, },
{
Name: "OnAdminAfterUpdateRequest error response",
Method: http.MethodPatch,
Url: "/api/admins/sbmbsdb40jyxf7h",
Body: strings.NewReader(`{
"email":"testnew@example.com",
"password":"1234567891",
"passwordConfirm":"1234567891",
"avatar":5
}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnAdminAfterUpdateRequest().Add(func(e *core.AdminUpdateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeUpdate": 1,
"OnModelAfterUpdate": 1,
"OnAdminBeforeUpdateRequest": 1,
"OnAdminAfterUpdateRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
+50 -26
View File
@@ -66,43 +66,67 @@ func NewUnauthorizedError(message string, data any) *ApiError {
// NewApiError creates and returns new normalized `ApiError` instance. // NewApiError creates and returns new normalized `ApiError` instance.
func NewApiError(status int, message string, data any) *ApiError { func NewApiError(status int, message string, data any) *ApiError {
message = inflector.Sentenize(message)
formattedData := map[string]any{}
if v, ok := data.(validation.Errors); ok {
formattedData = resolveValidationErrors(v)
}
return &ApiError{ return &ApiError{
rawData: data, rawData: data,
Data: formattedData, Data: safeErrorsData(data),
Code: status, Code: status,
Message: strings.TrimSpace(message), Message: strings.TrimSpace(inflector.Sentenize(message)),
} }
} }
func resolveValidationErrors(validationErrors validation.Errors) map[string]any { func safeErrorsData(data any) map[string]any {
switch v := data.(type) {
case validation.Errors:
return resolveSafeErrorsData[error](v)
case map[string]validation.Error:
return resolveSafeErrorsData[validation.Error](v)
case map[string]error:
return resolveSafeErrorsData[error](v)
case map[string]any:
return resolveSafeErrorsData[any](v)
default:
return map[string]any{} // not nil to ensure that is json serialized as object
}
}
func resolveSafeErrorsData[T any](data map[string]T) map[string]any {
result := map[string]any{} result := map[string]any{}
// extract from each validation error its error code and message. for name, err := range data {
for name, err := range validationErrors { if isNestedError(err) {
// check for nested errors result[name] = safeErrorsData(err)
if nestedErrs, ok := err.(validation.Errors); ok {
result[name] = resolveValidationErrors(nestedErrs)
continue continue
} }
result[name] = resolveSafeErrorItem(err)
errCode := "validation_invalid_value" // default
if errObj, ok := err.(validation.ErrorObject); ok {
errCode = errObj.Code()
}
result[name] = map[string]string{
"code": errCode,
"message": inflector.Sentenize(err.Error()),
}
} }
return result return result
} }
func isNestedError(err any) bool {
switch err.(type) {
case validation.Errors, map[string]validation.Error, map[string]error, map[string]any:
return true
}
return false
}
// resolveSafeErrorItem extracts from each validation error its
// public safe error code and message.
func resolveSafeErrorItem(err any) map[string]string {
// default public safe error values
code := "validation_invalid_value"
msg := "Invalid value."
// only validation errors are public safe
if obj, ok := err.(validation.Error); ok {
code = obj.Code()
msg = inflector.Sentenize(obj.Error())
}
return map[string]string{
"code": code,
"message": msg,
}
}
+24 -12
View File
@@ -10,6 +10,8 @@ import (
) )
func TestNewApiErrorWithRawData(t *testing.T) { func TestNewApiErrorWithRawData(t *testing.T) {
t.Parallel()
e := apis.NewApiError( e := apis.NewApiError(
300, 300,
"message_test", "message_test",
@@ -33,14 +35,16 @@ func TestNewApiErrorWithRawData(t *testing.T) {
} }
func TestNewApiErrorWithValidationData(t *testing.T) { func TestNewApiErrorWithValidationData(t *testing.T) {
t.Parallel()
e := apis.NewApiError( e := apis.NewApiError(
300, 300,
"message_test", "message_test",
validation.Errors{ validation.Errors{
"err1": errors.New("test error"), "err1": errors.New("test error"), // should be normalized
"err2": validation.ErrRequired, "err2": validation.ErrRequired,
"err3": validation.Errors{ "err3": validation.Errors{
"sub1": errors.New("test error"), "sub1": errors.New("test error"), // should be normalized
"sub2": validation.ErrRequired, "sub2": validation.ErrRequired,
"sub3": validation.Errors{ "sub3": validation.Errors{
"sub11": validation.ErrRequired, "sub11": validation.ErrRequired,
@@ -50,10 +54,10 @@ func TestNewApiErrorWithValidationData(t *testing.T) {
) )
result, _ := json.Marshal(e) result, _ := json.Marshal(e)
expected := `{"code":300,"message":"Message_test.","data":{"err1":{"code":"validation_invalid_value","message":"Test error."},"err2":{"code":"validation_required","message":"Cannot be blank."},"err3":{"sub1":{"code":"validation_invalid_value","message":"Test error."},"sub2":{"code":"validation_required","message":"Cannot be blank."},"sub3":{"sub11":{"code":"validation_required","message":"Cannot be blank."}}}}}` expected := `{"code":300,"message":"Message_test.","data":{"err1":{"code":"validation_invalid_value","message":"Invalid value."},"err2":{"code":"validation_required","message":"Cannot be blank."},"err3":{"sub1":{"code":"validation_invalid_value","message":"Invalid value."},"sub2":{"code":"validation_required","message":"Cannot be blank."},"sub3":{"sub11":{"code":"validation_required","message":"Cannot be blank."}}}}}`
if string(result) != expected { if string(result) != expected {
t.Errorf("Expected %v, got %v", expected, string(result)) t.Errorf("Expected \n%v, \ngot \n%v", expected, string(result))
} }
if e.Error() != "Message_test." { if e.Error() != "Message_test." {
@@ -66,6 +70,8 @@ func TestNewApiErrorWithValidationData(t *testing.T) {
} }
func TestNewNotFoundError(t *testing.T) { func TestNewNotFoundError(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
message string message string
data any data any
@@ -73,7 +79,7 @@ func TestNewNotFoundError(t *testing.T) {
}{ }{
{"", nil, `{"code":404,"message":"The requested resource wasn't found.","data":{}}`}, {"", nil, `{"code":404,"message":"The requested resource wasn't found.","data":{}}`},
{"demo", "rawData_test", `{"code":404,"message":"Demo.","data":{}}`}, {"demo", "rawData_test", `{"code":404,"message":"Demo.","data":{}}`},
{"demo", validation.Errors{"err1": errors.New("test error")}, `{"code":404,"message":"Demo.","data":{"err1":{"code":"validation_invalid_value","message":"Test error."}}}`}, {"demo", validation.Errors{"err1": validation.NewError("test_code", "test_message")}, `{"code":404,"message":"Demo.","data":{"err1":{"code":"test_code","message":"Test_message."}}}`},
} }
for i, scenario := range scenarios { for i, scenario := range scenarios {
@@ -81,12 +87,14 @@ func TestNewNotFoundError(t *testing.T) {
result, _ := json.Marshal(e) result, _ := json.Marshal(e)
if string(result) != scenario.expected { if string(result) != scenario.expected {
t.Errorf("(%d) Expected %v, got %v", i, scenario.expected, string(result)) t.Errorf("(%d) Expected \n%v, \ngot \n%v", i, scenario.expected, string(result))
} }
} }
} }
func TestNewBadRequestError(t *testing.T) { func TestNewBadRequestError(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
message string message string
data any data any
@@ -94,7 +102,7 @@ func TestNewBadRequestError(t *testing.T) {
}{ }{
{"", nil, `{"code":400,"message":"Something went wrong while processing your request.","data":{}}`}, {"", nil, `{"code":400,"message":"Something went wrong while processing your request.","data":{}}`},
{"demo", "rawData_test", `{"code":400,"message":"Demo.","data":{}}`}, {"demo", "rawData_test", `{"code":400,"message":"Demo.","data":{}}`},
{"demo", validation.Errors{"err1": errors.New("test error")}, `{"code":400,"message":"Demo.","data":{"err1":{"code":"validation_invalid_value","message":"Test error."}}}`}, {"demo", validation.Errors{"err1": validation.NewError("test_code", "test_message")}, `{"code":400,"message":"Demo.","data":{"err1":{"code":"test_code","message":"Test_message."}}}`},
} }
for i, scenario := range scenarios { for i, scenario := range scenarios {
@@ -102,12 +110,14 @@ func TestNewBadRequestError(t *testing.T) {
result, _ := json.Marshal(e) result, _ := json.Marshal(e)
if string(result) != scenario.expected { if string(result) != scenario.expected {
t.Errorf("(%d) Expected %v, got %v", i, scenario.expected, string(result)) t.Errorf("(%d) Expected \n%v, \ngot \n%v", i, scenario.expected, string(result))
} }
} }
} }
func TestNewForbiddenError(t *testing.T) { func TestNewForbiddenError(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
message string message string
data any data any
@@ -115,7 +125,7 @@ func TestNewForbiddenError(t *testing.T) {
}{ }{
{"", nil, `{"code":403,"message":"You are not allowed to perform this request.","data":{}}`}, {"", nil, `{"code":403,"message":"You are not allowed to perform this request.","data":{}}`},
{"demo", "rawData_test", `{"code":403,"message":"Demo.","data":{}}`}, {"demo", "rawData_test", `{"code":403,"message":"Demo.","data":{}}`},
{"demo", validation.Errors{"err1": errors.New("test error")}, `{"code":403,"message":"Demo.","data":{"err1":{"code":"validation_invalid_value","message":"Test error."}}}`}, {"demo", validation.Errors{"err1": validation.NewError("test_code", "test_message")}, `{"code":403,"message":"Demo.","data":{"err1":{"code":"test_code","message":"Test_message."}}}`},
} }
for i, scenario := range scenarios { for i, scenario := range scenarios {
@@ -123,12 +133,14 @@ func TestNewForbiddenError(t *testing.T) {
result, _ := json.Marshal(e) result, _ := json.Marshal(e)
if string(result) != scenario.expected { if string(result) != scenario.expected {
t.Errorf("(%d) Expected %v, got %v", i, scenario.expected, string(result)) t.Errorf("(%d) Expected \n%v, \ngot \n%v", i, scenario.expected, string(result))
} }
} }
} }
func TestNewUnauthorizedError(t *testing.T) { func TestNewUnauthorizedError(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
message string message string
data any data any
@@ -136,7 +148,7 @@ func TestNewUnauthorizedError(t *testing.T) {
}{ }{
{"", nil, `{"code":401,"message":"Missing or invalid authentication token.","data":{}}`}, {"", nil, `{"code":401,"message":"Missing or invalid authentication token.","data":{}}`},
{"demo", "rawData_test", `{"code":401,"message":"Demo.","data":{}}`}, {"demo", "rawData_test", `{"code":401,"message":"Demo.","data":{}}`},
{"demo", validation.Errors{"err1": errors.New("test error")}, `{"code":401,"message":"Demo.","data":{"err1":{"code":"validation_invalid_value","message":"Test error."}}}`}, {"demo", validation.Errors{"err1": validation.NewError("test_code", "test_message")}, `{"code":401,"message":"Demo.","data":{"err1":{"code":"test_code","message":"Test_message."}}}`},
} }
for i, scenario := range scenarios { for i, scenario := range scenarios {
@@ -144,7 +156,7 @@ func TestNewUnauthorizedError(t *testing.T) {
result, _ := json.Marshal(e) result, _ := json.Marshal(e)
if string(result) != scenario.expected { if string(result) != scenario.expected {
t.Errorf("(%d) Expected %v, got %v", i, scenario.expected, string(result)) t.Errorf("(%d) Expected \n%v, \ngot \n%v", i, scenario.expected, string(result))
} }
} }
} }
+31 -9
View File
@@ -2,9 +2,7 @@ package apis
import ( import (
"context" "context"
"log"
"net/http" "net/http"
"net/url"
"path/filepath" "path/filepath"
"time" "time"
@@ -12,6 +10,8 @@ import (
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/forms" "github.com/pocketbase/pocketbase/forms"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tools/filesystem"
"github.com/pocketbase/pocketbase/tools/rest"
"github.com/pocketbase/pocketbase/tools/types" "github.com/pocketbase/pocketbase/tools/types"
"github.com/spf13/cast" "github.com/spf13/cast"
) )
@@ -25,6 +25,7 @@ func bindBackupApi(app core.App, rg *echo.Group) {
subGroup := rg.Group("/backups", ActivityLogger(app)) subGroup := rg.Group("/backups", ActivityLogger(app))
subGroup.GET("", api.list, RequireAdminAuth()) subGroup.GET("", api.list, RequireAdminAuth())
subGroup.POST("", api.create, RequireAdminAuth()) subGroup.POST("", api.create, RequireAdminAuth())
subGroup.POST("/upload", api.upload, RequireAdminAuth())
subGroup.GET("/:key", api.download) subGroup.GET("/:key", api.download)
subGroup.DELETE("/:key", api.delete, RequireAdminAuth()) subGroup.DELETE("/:key", api.delete, RequireAdminAuth())
subGroup.POST("/:key/restore", api.restore, RequireAdminAuth()) subGroup.POST("/:key/restore", api.restore, RequireAdminAuth())
@@ -67,7 +68,7 @@ func (api *backupApi) list(c echo.Context) error {
} }
func (api *backupApi) create(c echo.Context) error { func (api *backupApi) create(c echo.Context) error {
if api.app.Cache().Has(core.CacheKeyActiveBackup) { if api.app.Store().Has(core.StoreKeyActiveBackup) {
return NewBadRequestError("Try again later - another backup/restore process has already been started", nil) return NewBadRequestError("Try again later - another backup/restore process has already been started", nil)
} }
@@ -89,6 +90,28 @@ func (api *backupApi) create(c echo.Context) error {
}) })
} }
func (api *backupApi) upload(c echo.Context) error {
files, err := rest.FindUploadedFiles(c.Request(), "file")
if err != nil {
return NewBadRequestError("Missing or invalid uploaded file.", err)
}
form := forms.NewBackupUpload(api.app)
form.File = files[0]
return form.Submit(func(next forms.InterceptorNextFunc[*filesystem.File]) forms.InterceptorNextFunc[*filesystem.File] {
return func(file *filesystem.File) error {
if err := next(file); err != nil {
return NewBadRequestError("Failed to upload backup.", err)
}
// we don't retrieve the generated backup file because it may not be
// available yet due to the eventually consistent nature of some S3 providers
return c.NoContent(http.StatusNoContent)
}
})
}
func (api *backupApi) download(c echo.Context) error { func (api *backupApi) download(c echo.Context) error {
fileToken := c.QueryParam("token") fileToken := c.QueryParam("token")
@@ -128,12 +151,11 @@ func (api *backupApi) download(c echo.Context) error {
} }
func (api *backupApi) restore(c echo.Context) error { func (api *backupApi) restore(c echo.Context) error {
if api.app.Cache().Has(core.CacheKeyActiveBackup) { if api.app.Store().Has(core.StoreKeyActiveBackup) {
return NewBadRequestError("Try again later - another backup/restore process has already been started.", nil) return NewBadRequestError("Try again later - another backup/restore process has already been started.", nil)
} }
// @todo remove the extra unescape after https://github.com/labstack/echo/issues/2447 key := c.PathParam("key")
key, _ := url.PathUnescape(c.PathParam("key"))
existsCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second) existsCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel() defer cancel()
@@ -158,8 +180,8 @@ func (api *backupApi) restore(c echo.Context) error {
// give some optimistic time to write the response // give some optimistic time to write the response
time.Sleep(1 * time.Second) time.Sleep(1 * time.Second)
if err := api.app.RestoreBackup(ctx, key); err != nil && api.app.IsDebug() { if err := api.app.RestoreBackup(ctx, key); err != nil {
log.Println(err) api.app.Logger().Error("Failed to restore backup", "key", key, "error", err.Error())
} }
}() }()
@@ -180,7 +202,7 @@ func (api *backupApi) delete(c echo.Context) error {
key := c.PathParam("key") key := c.PathParam("key")
if key != "" && cast.ToString(api.app.Cache().Get(core.CacheKeyActiveBackup)) == key { if key != "" && cast.ToString(api.app.Store().Get(core.StoreKeyActiveBackup)) == key {
return NewBadRequestError("The backup is currently being used and cannot be deleted.", nil) return NewBadRequestError("The backup is currently being used and cannot be deleted.", nil)
} }
+219 -20
View File
@@ -1,7 +1,11 @@
package apis_test package apis_test
import ( import (
"archive/zip"
"bytes"
"context" "context"
"io"
"mime/multipart"
"net/http" "net/http"
"strings" "strings"
"testing" "testing"
@@ -13,6 +17,8 @@ import (
) )
func TestBackupsList(t *testing.T) { func TestBackupsList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -80,12 +86,14 @@ func TestBackupsList(t *testing.T) {
} }
func TestBackupsCreate(t *testing.T) { func TestBackupsCreate(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
Method: http.MethodPost, Method: http.MethodPost,
Url: "/api/backups", Url: "/api/backups",
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app) ensureNoBackups(t, app)
}, },
ExpectedStatus: 401, ExpectedStatus: 401,
@@ -98,7 +106,7 @@ func TestBackupsCreate(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc", "Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app) ensureNoBackups(t, app)
}, },
ExpectedStatus: 401, ExpectedStatus: 401,
@@ -112,9 +120,9 @@ func TestBackupsCreate(t *testing.T) {
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.Cache().Set(core.CacheKeyActiveBackup, "") app.Store().Set(core.StoreKeyActiveBackup, "")
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app) ensureNoBackups(t, app)
}, },
ExpectedStatus: 400, ExpectedStatus: 400,
@@ -127,7 +135,7 @@ func TestBackupsCreate(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
files, err := getBackupFiles(app) files, err := getBackupFiles(app)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -152,7 +160,7 @@ func TestBackupsCreate(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app) ensureNoBackups(t, app)
}, },
ExpectedStatus: 400, ExpectedStatus: 400,
@@ -169,7 +177,7 @@ func TestBackupsCreate(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
files, err := getBackupFiles(app) files, err := getBackupFiles(app)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -193,7 +201,143 @@ func TestBackupsCreate(t *testing.T) {
} }
} }
func TestBackupsUpload(t *testing.T) {
t.Parallel()
// create dummy form data bodies
type body struct {
buffer io.Reader
contentType string
}
bodies := make([]body, 10)
for i := 0; i < 10; i++ {
func() {
zb := new(bytes.Buffer)
zw := zip.NewWriter(zb)
if err := zw.Close(); err != nil {
t.Fatal(err)
}
b := new(bytes.Buffer)
mw := multipart.NewWriter(b)
mfw, err := mw.CreateFormFile("file", "test")
if err != nil {
t.Fatal(err)
}
if _, err := io.Copy(mfw, zb); err != nil {
t.Fatal(err)
}
mw.Close()
bodies[i] = body{
buffer: b,
contentType: mw.FormDataContentType(),
}
}()
}
// ---
scenarios := []tests.ApiScenario{
{
Name: "unauthorized",
Method: http.MethodPost,
Url: "/api/backups/upload",
Body: bodies[0].buffer,
RequestHeaders: map[string]string{
"Content-Type": bodies[0].contentType,
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app)
},
ExpectedStatus: 401,
ExpectedContent: []string{`"data":{}`},
},
{
Name: "authorized as auth record",
Method: http.MethodPost,
Url: "/api/backups/upload",
Body: bodies[1].buffer,
RequestHeaders: map[string]string{
"Content-Type": bodies[1].contentType,
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc",
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app)
},
ExpectedStatus: 401,
ExpectedContent: []string{`"data":{}`},
},
{
Name: "authorized as admin (missing file)",
Method: http.MethodPost,
Url: "/api/backups/upload",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureNoBackups(t, app)
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{`},
},
{
Name: "authorized as admin (existing backup name)",
Method: http.MethodPost,
Url: "/api/backups/upload",
Body: bodies[3].buffer,
RequestHeaders: map[string]string{
"Content-Type": bodies[3].contentType,
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
fsys, err := app.NewBackupsFilesystem()
if err != nil {
t.Fatal(err)
}
defer fsys.Close()
// create a dummy backup file to simulate existing backups
if err := fsys.Upload([]byte("123"), "test"); err != nil {
t.Fatal(err)
}
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
files, _ := getBackupFiles(app)
if total := len(files); total != 1 {
t.Fatalf("Expected %d backup file, got %d", 1, total)
}
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{"file":{`},
},
{
Name: "authorized as admin (valid file)",
Method: http.MethodPost,
Url: "/api/backups/upload",
Body: bodies[4].buffer,
RequestHeaders: map[string]string{
"Content-Type": bodies[4].contentType,
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
files, _ := getBackupFiles(app)
if total := len(files); total != 1 {
t.Fatalf("Expected %d backup file, got %d", 1, total)
}
},
ExpectedStatus: 204,
},
}
for _, scenario := range scenarios {
scenario.Test(t)
}
}
func TestBackupsDownload(t *testing.T) { func TestBackupsDownload(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -300,7 +444,7 @@ func TestBackupsDownload(t *testing.T) {
{ {
Name: "with valid admin file token but missing backup name", Name: "with valid admin file token but missing backup name",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/backups/mizzing?token=eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsImV4cCI6MTg5MzQ1MjQ2MSwidHlwZSI6ImFkbWluIn0.LyAMpSfaHVsuUqIlqqEbhDQSdFzoPz_EIDcb2VJMBsU", Url: "/api/backups/missing?token=eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsImV4cCI6MTg5MzQ1MjQ2MSwidHlwZSI6ImFkbWluIn0.LyAMpSfaHVsuUqIlqqEbhDQSdFzoPz_EIDcb2VJMBsU",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := createTestBackups(app); err != nil { if err := createTestBackups(app); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -325,6 +469,22 @@ func TestBackupsDownload(t *testing.T) {
`logs.db`, `logs.db`,
}, },
}, },
{
Name: "with valid admin file token and backup name with escaped char",
Method: http.MethodGet,
Url: "/api/backups/%40test4.zip?token=eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsImV4cCI6MTg5MzQ1MjQ2MSwidHlwZSI6ImFkbWluIn0.LyAMpSfaHVsuUqIlqqEbhDQSdFzoPz_EIDcb2VJMBsU",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := createTestBackups(app); err != nil {
t.Fatal(err)
}
},
ExpectedStatus: 200,
ExpectedContent: []string{
`storage/`,
`data.db`,
`logs.db`,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -333,13 +493,15 @@ func TestBackupsDownload(t *testing.T) {
} }
func TestBackupsDelete(t *testing.T) { func TestBackupsDelete(t *testing.T) {
t.Parallel()
noTestBackupFilesChanges := func(t *testing.T, app *tests.TestApp) { noTestBackupFilesChanges := func(t *testing.T, app *tests.TestApp) {
files, err := getBackupFiles(app) files, err := getBackupFiles(app)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
expected := 3 expected := 4
if total := len(files); total != expected { if total := len(files); total != expected {
t.Fatalf("Expected %d backup(s), got %d", expected, total) t.Fatalf("Expected %d backup(s), got %d", expected, total)
} }
@@ -355,7 +517,7 @@ func TestBackupsDelete(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
noTestBackupFilesChanges(t, app) noTestBackupFilesChanges(t, app)
}, },
ExpectedStatus: 401, ExpectedStatus: 401,
@@ -373,7 +535,7 @@ func TestBackupsDelete(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
noTestBackupFilesChanges(t, app) noTestBackupFilesChanges(t, app)
}, },
ExpectedStatus: 401, ExpectedStatus: 401,
@@ -391,7 +553,7 @@ func TestBackupsDelete(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
noTestBackupFilesChanges(t, app) noTestBackupFilesChanges(t, app)
}, },
ExpectedStatus: 400, ExpectedStatus: 400,
@@ -410,10 +572,9 @@ func TestBackupsDelete(t *testing.T) {
} }
// mock active backup with the same name to delete // mock active backup with the same name to delete
app.Cache().Set(core.CacheKeyActiveBackup, "test1.zip") app.Store().Set(core.StoreKeyActiveBackup, "test1.zip")
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
noTestBackupFilesChanges(t, app) noTestBackupFilesChanges(t, app)
}, },
ExpectedStatus: 400, ExpectedStatus: 400,
@@ -432,16 +593,16 @@ func TestBackupsDelete(t *testing.T) {
} }
// mock active backup with different name // mock active backup with different name
app.Cache().Set(core.CacheKeyActiveBackup, "new.zip") app.Store().Set(core.StoreKeyActiveBackup, "new.zip")
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
files, err := getBackupFiles(app) files, err := getBackupFiles(app)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
if total := len(files); total != 2 { if total := len(files); total != 3 {
t.Fatalf("Expected 2 backup files, got %d", total) t.Fatalf("Expected %d backup files, got %d", 3, total)
} }
deletedFile := "test1.zip" deletedFile := "test1.zip"
@@ -454,6 +615,38 @@ func TestBackupsDelete(t *testing.T) {
}, },
ExpectedStatus: 204, ExpectedStatus: 204,
}, },
{
Name: "authorized as admin (backup with escaped character)",
Method: http.MethodDelete,
Url: "/api/backups/%40test4.zip",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := createTestBackups(app); err != nil {
t.Fatal(err)
}
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
files, err := getBackupFiles(app)
if err != nil {
t.Fatal(err)
}
if total := len(files); total != 3 {
t.Fatalf("Expected %d backup files, got %d", 3, total)
}
deletedFile := "@test4.zip"
for _, f := range files {
if f.Key == deletedFile {
t.Fatalf("Expected backup %q to be deleted", deletedFile)
}
}
},
ExpectedStatus: 204,
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -462,6 +655,8 @@ func TestBackupsDelete(t *testing.T) {
} }
func TestBackupsRestore(t *testing.T) { func TestBackupsRestore(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -517,7 +712,7 @@ func TestBackupsRestore(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
app.Cache().Set(core.CacheKeyActiveBackup, "") app.Store().Set(core.StoreKeyActiveBackup, "")
}, },
ExpectedStatus: 400, ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`}, ExpectedContent: []string{`"data":{}`},
@@ -546,6 +741,10 @@ func createTestBackups(app core.App) error {
return err return err
} }
if err := app.CreateBackup(ctx, "@test4.zip"); err != nil {
return err
}
return nil return nil
} }
+67 -66
View File
@@ -2,14 +2,16 @@
package apis package apis
import ( import (
"database/sql"
"errors" "errors"
"fmt" "fmt"
"io/fs" "io/fs"
"log" "log/slog"
"net/http" "net/http"
"net/url" "net/url"
"path/filepath" "path/filepath"
"strings" "strings"
"time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/labstack/echo/v5/middleware" "github.com/labstack/echo/v5/middleware"
@@ -25,15 +27,17 @@ const trailedAdminPath = "/_/"
// system and app specific routes and middlewares. // system and app specific routes and middlewares.
func InitApi(app core.App) (*echo.Echo, error) { func InitApi(app core.App) (*echo.Echo, error) {
e := echo.New() e := echo.New()
e.Debug = app.IsDebug() e.Debug = false
e.Binder = &rest.MultiBinder{}
e.JSONSerializer = &rest.Serializer{ e.JSONSerializer = &rest.Serializer{
FieldsParam: "fields", FieldsParam: fieldsQueryParam,
} }
// configure a custom router // configure a custom router
e.ResetRouterCreator(func(ec *echo.Echo) echo.Router { e.ResetRouterCreator(func(ec *echo.Echo) echo.Router {
return echo.NewRouter(echo.RouterConfig{ return echo.NewRouter(echo.RouterConfig{
UnescapePathParamValues: true, UnescapePathParamValues: true,
AllowOverwritingRoute: true,
}) })
}) })
@@ -44,38 +48,42 @@ func InitApi(app core.App) (*echo.Echo, error) {
return !strings.HasPrefix(c.Request().URL.Path, "/api/") return !strings.HasPrefix(c.Request().URL.Path, "/api/")
}, },
})) }))
e.Pre(LoadAuthContext(app))
e.Use(middleware.Recover()) e.Use(middleware.Recover())
e.Use(middleware.Secure()) e.Use(middleware.Secure())
e.Use(LoadAuthContext(app)) e.Use(func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
c.Set(ContextExecStartKey, time.Now())
return next(c)
}
})
// custom error handler // custom error handler
e.HTTPErrorHandler = func(c echo.Context, err error) { e.HTTPErrorHandler = func(c echo.Context, err error) {
if c.Response().Committed { if err == nil {
if app.IsDebug() { return // no error
log.Println("HTTPErrorHandler response was already committed:", err)
}
return
} }
var apiErr *ApiError var apiErr *ApiError
switch v := err.(type) { if errors.As(err, &apiErr) {
case *echo.HTTPError: // already an api error...
if v.Internal != nil && app.IsDebug() { } else if v := new(echo.HTTPError); errors.As(err, &v) {
log.Println(v.Internal)
}
msg := fmt.Sprintf("%v", v.Message) msg := fmt.Sprintf("%v", v.Message)
apiErr = NewApiError(v.Code, msg, v) apiErr = NewApiError(v.Code, msg, v)
case *ApiError: } else {
if app.IsDebug() && v.RawData() != nil { if errors.Is(err, sql.ErrNoRows) {
log.Println(v.RawData()) apiErr = NewNotFoundError("", err)
} else {
apiErr = NewBadRequestError("", err)
} }
apiErr = v }
default:
if err != nil && app.IsDebug() { logRequest(app, c, apiErr)
log.Println(err)
} if c.Response().Committed {
apiErr = NewBadRequestError("", err) return // already committed
} }
event := new(core.ApiErrorEvent) event := new(core.ApiErrorEvent)
@@ -84,6 +92,10 @@ func InitApi(app core.App) (*echo.Echo, error) {
// send error response // send error response
hookErr := app.OnBeforeApiError().Trigger(event, func(e *core.ApiErrorEvent) error { hookErr := app.OnBeforeApiError().Trigger(event, func(e *core.ApiErrorEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
// @see https://github.com/labstack/echo/issues/608 // @see https://github.com/labstack/echo/issues/608
if e.HttpContext.Request().Method == http.MethodHead { if e.HttpContext.Request().Method == http.MethodHead {
return e.HttpContext.NoContent(apiErr.Code) return e.HttpContext.NoContent(apiErr.Code)
@@ -92,19 +104,20 @@ func InitApi(app core.App) (*echo.Echo, error) {
return e.HttpContext.JSON(apiErr.Code, apiErr) return e.HttpContext.JSON(apiErr.Code, apiErr)
}) })
// truly rare case; eg. client already disconnected if hookErr == nil {
if hookErr != nil && app.IsDebug() { if err := app.OnAfterApiError().Trigger(event); err != nil {
log.Println(hookErr) app.Logger().Debug("OnAfterApiError failure", slog.String("error", err.Error()))
}
} else {
app.Logger().Debug("OnBeforeApiError error (truly rare case, eg. client already disconnected)", slog.String("error", hookErr.Error()))
} }
app.OnAfterApiError().Trigger(event)
} }
// admin ui routes // admin ui routes
bindStaticAdminUI(app, e) bindStaticAdminUI(app, e)
// default routes // default routes
api := e.Group("/api") api := e.Group("/api", eagerRequestInfoCache(app))
bindSettingsApi(app, api) bindSettingsApi(app, api)
bindAdminApi(app, api) bindAdminApi(app, api)
bindCollectionApi(app, api) bindCollectionApi(app, api)
@@ -116,20 +129,6 @@ func InitApi(app core.App) (*echo.Echo, error) {
bindHealthApi(app, api) bindHealthApi(app, api)
bindBackupApi(app, api) bindBackupApi(app, api)
// trigger the custom BeforeServe hook for the created api router
// allowing users to further adjust its options or register new routes
serveEvent := &core.ServeEvent{
App: app,
Router: e,
}
if err := app.OnBeforeServe().Trigger(serveEvent); err != nil {
return nil, err
}
// note: it is after the OnBeforeServe hook to ensure that the implicit
// cache is after any user custom defined middlewares
e.Use(eagerRequestDataCache(app))
// catch all any route // catch all any route
api.Any("/*", func(c echo.Context) error { api.Any("/*", func(c echo.Context) error {
return echo.ErrNotFound return echo.ErrNotFound
@@ -192,19 +191,6 @@ func bindStaticAdminUI(app core.App, e *echo.Echo) error {
return nil return nil
} }
const totalAdminsCacheKey = "@totalAdmins"
func updateTotalAdminsCache(app core.App) error {
total, err := app.Dao().TotalAdmins()
if err != nil {
return err
}
app.Cache().Set(totalAdminsCacheKey, total)
return nil
}
func uiCacheControl() echo.MiddlewareFunc { func uiCacheControl() echo.MiddlewareFunc {
return func(next echo.HandlerFunc) echo.HandlerFunc { return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error { return func(c echo.Context) error {
@@ -219,15 +205,29 @@ func uiCacheControl() echo.MiddlewareFunc {
} }
} }
const hasAdminsCacheKey = "@hasAdmins"
func updateHasAdminsCache(app core.App) error {
total, err := app.Dao().TotalAdmins()
if err != nil {
return err
}
app.Store().Set(hasAdminsCacheKey, total > 0)
return nil
}
// installerRedirect redirects the user to the installer admin UI page // installerRedirect redirects the user to the installer admin UI page
// when the application needs some preliminary configurations to be done. // when the application needs some preliminary configurations to be done.
func installerRedirect(app core.App) echo.MiddlewareFunc { func installerRedirect(app core.App) echo.MiddlewareFunc {
// keep totalAdminsCacheKey value up-to-date // keep hasAdminsCacheKey value up-to-date
app.OnAdminAfterCreateRequest().Add(func(data *core.AdminCreateEvent) error { app.OnAdminAfterCreateRequest().Add(func(data *core.AdminCreateEvent) error {
return updateTotalAdminsCache(app) return updateHasAdminsCache(app)
}) })
app.OnAdminAfterDeleteRequest().Add(func(data *core.AdminDeleteEvent) error { app.OnAdminAfterDeleteRequest().Add(func(data *core.AdminDeleteEvent) error {
return updateTotalAdminsCache(app) return updateHasAdminsCache(app)
}) })
return func(next echo.HandlerFunc) echo.HandlerFunc { return func(next echo.HandlerFunc) echo.HandlerFunc {
@@ -238,23 +238,24 @@ func installerRedirect(app core.App) echo.MiddlewareFunc {
return next(c) return next(c)
} }
// load into cache (if not already) hasAdmins := cast.ToBool(app.Store().Get(hasAdminsCacheKey))
if !app.Cache().Has(totalAdminsCacheKey) {
if err := updateTotalAdminsCache(app); err != nil { if !hasAdmins {
// update the cache to make sure that the admin wasn't created by another process
if err := updateHasAdminsCache(app); err != nil {
return err return err
} }
hasAdmins = cast.ToBool(app.Store().Get(hasAdminsCacheKey))
} }
totalAdmins := cast.ToInt(app.Cache().Get(totalAdminsCacheKey))
_, hasInstallerParam := c.Request().URL.Query()["installer"] _, hasInstallerParam := c.Request().URL.Query()["installer"]
if totalAdmins == 0 && !hasInstallerParam { if !hasAdmins && !hasInstallerParam {
// redirect to the installer page // redirect to the installer page
return c.Redirect(http.StatusTemporaryRedirect, "?installer#") return c.Redirect(http.StatusTemporaryRedirect, "?installer#")
} }
if totalAdmins != 0 && hasInstallerParam { if hasAdmins && hasInstallerParam {
// clear the installer param // clear the installer param
return c.Redirect(http.StatusTemporaryRedirect, "?") return c.Redirect(http.StatusTemporaryRedirect, "?")
} }
+176 -48
View File
@@ -1,6 +1,7 @@
package apis_test package apis_test
import ( import (
"database/sql"
"errors" "errors"
"fmt" "fmt"
"net/http" "net/http"
@@ -10,10 +11,13 @@ import (
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/apis" "github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/rest"
"github.com/spf13/cast" "github.com/spf13/cast"
) )
func Test404(t *testing.T) { func Test404(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Method: http.MethodGet, Method: http.MethodGet,
@@ -52,6 +56,8 @@ func Test404(t *testing.T) {
} }
func TestCustomRoutesAndErrorsHandling(t *testing.T) { func TestCustomRoutesAndErrorsHandling(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "custom route", Name: "custom route",
@@ -141,6 +147,8 @@ func TestCustomRoutesAndErrorsHandling(t *testing.T) {
} }
func TestRemoveTrailingSlashMiddleware(t *testing.T) { func TestRemoveTrailingSlashMiddleware(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "non /api/* route (exact match)", Name: "non /api/* route (exact match)",
@@ -213,16 +221,27 @@ func TestRemoveTrailingSlashMiddleware(t *testing.T) {
} }
} }
func TestEagerRequestDataCache(t *testing.T) { func TestMultiBinder(t *testing.T) {
t.Parallel()
rawJson := `{"name":"test123"}`
formData, mp, err := tests.MockMultipartData(map[string]string{
rest.MultipartJsonKey: rawJson,
})
if err != nil {
t.Fatal(err)
}
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "[UNKNOWN] unsupported eager cached request method", Name: "non-api group route",
Method: "UNKNOWN", Method: "POST",
Url: "/custom", Url: "/custom",
Body: strings.NewReader(`{"name":"test123"}`), Body: strings.NewReader(rawJson),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.AddRoute(echo.Route{ e.AddRoute(echo.Route{
Method: "UNKNOWN", Method: "POST",
Path: "/custom", Path: "/custom",
Handler: func(c echo.Context) error { Handler: func(c echo.Context) error {
data := &struct { data := &struct {
@@ -233,59 +252,168 @@ func TestEagerRequestDataCache(t *testing.T) {
return err return err
} }
// since the unknown method is not eager cache support // try to read the body again
// it should fail reading the json body twice r := apis.RequestInfo(c)
r := apis.RequestData(c) if v := cast.ToString(r.Data["name"]); v != "test123" {
if v := cast.ToString(r.Data["name"]); v != "" { t.Fatalf("Expected request data with name %q, got, %q", "test123", v)
t.Fatalf("Expected empty request data body, got, %v", r.Data)
} }
return c.String(200, data.Name) return c.NoContent(200)
}, },
}) })
}, },
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedContent: []string{"test123"},
}, },
} {
Name: "api group route",
Method: "GET",
Url: "/api/admins",
Body: strings.NewReader(rawJson),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.Use(func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
// it is not important whether the route handler return an error since
// we just need to ensure that the eagerRequestInfoCache was registered
next(c)
// supported eager cache request methods // ensure that the body was read at least once
supportedMethods := []string{"POST", "PUT", "PATCH", "DELETE"} data := &struct {
for _, m := range supportedMethods { Name string `json:"name"`
scenarios = append( }{}
scenarios, c.Bind(data)
tests.ApiScenario{
Name: fmt.Sprintf("[%s] valid cached json body request", m),
Method: http.MethodPost,
Url: "/custom",
Body: strings.NewReader(`{"name":"test123"}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.AddRoute(echo.Route{
Method: http.MethodPost,
Path: "/custom",
Handler: func(c echo.Context) error {
data := &struct {
Name string `json:"name"`
}{}
if err := c.Bind(data); err != nil { // try to read the body again
return err r := apis.RequestInfo(c)
} if v := cast.ToString(r.Data["name"]); v != "test123" {
t.Fatalf("Expected request data with name %q, got, %q", "test123", v)
}
// try to read the body again return nil
r := apis.RequestData(c) }
if v := cast.ToString(r.Data["name"]); v != "test123" { })
t.Fatalf("Expected request data with name %q, got, %q", "test123", v)
}
return c.String(200, data.Name)
},
})
},
ExpectedStatus: 200,
ExpectedContent: []string{"test123"},
}, },
) ExpectedStatus: 200,
},
{
Name: "custom route with @jsonPayload as multipart body",
Method: "POST",
Url: "/custom",
Body: formData,
RequestHeaders: map[string]string{
"Content-Type": mp.FormDataContentType(),
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.AddRoute(echo.Route{
Method: "POST",
Path: "/custom",
Handler: func(c echo.Context) error {
data := &struct {
Name string `json:"name"`
}{}
if err := c.Bind(data); err != nil {
return err
}
// try to read the body again
r := apis.RequestInfo(c)
if v := cast.ToString(r.Data["name"]); v != "test123" {
t.Fatalf("Expected request data with name %q, got, %q", "test123", v)
}
return c.NoContent(200)
},
})
},
ExpectedStatus: 200,
},
}
for _, scenario := range scenarios {
scenario.Test(t)
}
}
func TestErrorHandler(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{
{
Name: "apis.ApiError",
Method: http.MethodGet,
Url: "/test",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.GET("/test", func(c echo.Context) error {
return apis.NewApiError(418, "test", nil)
})
},
ExpectedStatus: 418,
ExpectedContent: []string{`"message":"Test."`},
},
{
Name: "wrapped apis.ApiError",
Method: http.MethodGet,
Url: "/test",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.GET("/test", func(c echo.Context) error {
return fmt.Errorf("example 123: %w", apis.NewApiError(418, "test", nil))
})
},
ExpectedStatus: 418,
ExpectedContent: []string{`"message":"Test."`},
NotExpectedContent: []string{"example", "123"},
},
{
Name: "echo.HTTPError",
Method: http.MethodGet,
Url: "/test",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.GET("/test", func(c echo.Context) error {
return echo.NewHTTPError(418, "test")
})
},
ExpectedStatus: 418,
ExpectedContent: []string{`"message":"Test."`},
},
{
Name: "wrapped echo.HTTPError",
Method: http.MethodGet,
Url: "/test",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.GET("/test", func(c echo.Context) error {
return fmt.Errorf("example 123: %w", echo.NewHTTPError(418, "test"))
})
},
ExpectedStatus: 418,
ExpectedContent: []string{`"message":"Test."`},
NotExpectedContent: []string{"example", "123"},
},
{
Name: "wrapped sql.ErrNoRows",
Method: http.MethodGet,
Url: "/test",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.GET("/test", func(c echo.Context) error {
return fmt.Errorf("example 123: %w", sql.ErrNoRows)
})
},
ExpectedStatus: 404,
ExpectedContent: []string{`"data":{}`},
NotExpectedContent: []string{"example", "123"},
},
{
Name: "custom error",
Method: http.MethodGet,
Url: "/test",
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
e.GET("/test", func(c echo.Context) error {
return fmt.Errorf("example 123")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
NotExpectedContent: []string{"example", "123"},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
+40 -41
View File
@@ -1,7 +1,6 @@
package apis package apis
import ( import (
"log"
"net/http" "net/http"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
@@ -49,6 +48,10 @@ func (api *collectionApi) list(c echo.Context) error {
event.Result = result event.Result = result
return api.app.OnCollectionsListRequest().Trigger(event, func(e *core.CollectionsListEvent) error { return api.app.OnCollectionsListRequest().Trigger(event, func(e *core.CollectionsListEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Result) return e.HttpContext.JSON(http.StatusOK, e.Result)
}) })
} }
@@ -64,6 +67,10 @@ func (api *collectionApi) view(c echo.Context) error {
event.Collection = collection event.Collection = collection
return api.app.OnCollectionViewRequest().Trigger(event, func(e *core.CollectionViewEvent) error { return api.app.OnCollectionViewRequest().Trigger(event, func(e *core.CollectionViewEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Collection) return e.HttpContext.JSON(http.StatusOK, e.Collection)
}) })
} }
@@ -83,7 +90,7 @@ func (api *collectionApi) create(c echo.Context) error {
event.Collection = collection event.Collection = collection
// create the collection // create the collection
submitErr := form.Submit(func(next forms.InterceptorNextFunc[*models.Collection]) forms.InterceptorNextFunc[*models.Collection] { return form.Submit(func(next forms.InterceptorNextFunc[*models.Collection]) forms.InterceptorNextFunc[*models.Collection] {
return func(m *models.Collection) error { return func(m *models.Collection) error {
event.Collection = m event.Collection = m
@@ -92,18 +99,16 @@ func (api *collectionApi) create(c echo.Context) error {
return NewBadRequestError("Failed to create the collection.", err) return NewBadRequestError("Failed to create the collection.", err)
} }
return e.HttpContext.JSON(http.StatusOK, e.Collection) return api.app.OnCollectionAfterCreateRequest().Trigger(event, func(e *core.CollectionCreateEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Collection)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnCollectionAfterCreateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr
} }
func (api *collectionApi) update(c echo.Context) error { func (api *collectionApi) update(c echo.Context) error {
@@ -124,7 +129,7 @@ func (api *collectionApi) update(c echo.Context) error {
event.Collection = collection event.Collection = collection
// update the collection // update the collection
submitErr := form.Submit(func(next forms.InterceptorNextFunc[*models.Collection]) forms.InterceptorNextFunc[*models.Collection] { return form.Submit(func(next forms.InterceptorNextFunc[*models.Collection]) forms.InterceptorNextFunc[*models.Collection] {
return func(m *models.Collection) error { return func(m *models.Collection) error {
event.Collection = m event.Collection = m
@@ -133,18 +138,16 @@ func (api *collectionApi) update(c echo.Context) error {
return NewBadRequestError("Failed to update the collection.", err) return NewBadRequestError("Failed to update the collection.", err)
} }
return e.HttpContext.JSON(http.StatusOK, e.Collection) return api.app.OnCollectionAfterUpdateRequest().Trigger(event, func(e *core.CollectionUpdateEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Collection)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnCollectionAfterUpdateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr
} }
func (api *collectionApi) delete(c echo.Context) error { func (api *collectionApi) delete(c echo.Context) error {
@@ -157,21 +160,19 @@ func (api *collectionApi) delete(c echo.Context) error {
event.HttpContext = c event.HttpContext = c
event.Collection = collection event.Collection = collection
handlerErr := api.app.OnCollectionBeforeDeleteRequest().Trigger(event, func(e *core.CollectionDeleteEvent) error { return api.app.OnCollectionBeforeDeleteRequest().Trigger(event, func(e *core.CollectionDeleteEvent) error {
if err := api.app.Dao().DeleteCollection(e.Collection); err != nil { if err := api.app.Dao().DeleteCollection(e.Collection); err != nil {
return NewBadRequestError("Failed to delete collection due to existing dependency.", err) return NewBadRequestError("Failed to delete collection due to existing dependency.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnCollectionAfterDeleteRequest().Trigger(event, func(e *core.CollectionDeleteEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
if handlerErr == nil {
if err := api.app.OnCollectionAfterDeleteRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return handlerErr
} }
func (api *collectionApi) bulkImport(c echo.Context) error { func (api *collectionApi) bulkImport(c echo.Context) error {
@@ -187,7 +188,7 @@ func (api *collectionApi) bulkImport(c echo.Context) error {
event.Collections = form.Collections event.Collections = form.Collections
// import collections // import collections
submitErr := form.Submit(func(next forms.InterceptorNextFunc[[]*models.Collection]) forms.InterceptorNextFunc[[]*models.Collection] { return form.Submit(func(next forms.InterceptorNextFunc[[]*models.Collection]) forms.InterceptorNextFunc[[]*models.Collection] {
return func(imports []*models.Collection) error { return func(imports []*models.Collection) error {
event.Collections = imports event.Collections = imports
@@ -196,16 +197,14 @@ func (api *collectionApi) bulkImport(c echo.Context) error {
return NewBadRequestError("Failed to import the submitted collections.", err) return NewBadRequestError("Failed to import the submitted collections.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnCollectionsAfterImportRequest().Trigger(event, func(e *core.CollectionsImportEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnCollectionsAfterImportRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr
} }
+154 -105
View File
@@ -1,6 +1,7 @@
package apis_test package apis_test
import ( import (
"errors"
"net/http" "net/http"
"os" "os"
"path/filepath" "path/filepath"
@@ -9,13 +10,15 @@ import (
"time" "time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/models/schema"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/list" "github.com/pocketbase/pocketbase/tools/list"
) )
func TestCollectionsList(t *testing.T) { func TestCollectionsList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -45,7 +48,7 @@ func TestCollectionsList(t *testing.T) {
ExpectedContent: []string{ ExpectedContent: []string{
`"page":1`, `"page":1`,
`"perPage":30`, `"perPage":30`,
`"totalItems":10`, `"totalItems":11`,
`"items":[{`, `"items":[{`,
`"id":"_pb_users_auth_"`, `"id":"_pb_users_auth_"`,
`"id":"v851q4r790rhknl"`, `"id":"v851q4r790rhknl"`,
@@ -55,6 +58,7 @@ func TestCollectionsList(t *testing.T) {
`"id":"wzlqyes4orhoygb"`, `"id":"wzlqyes4orhoygb"`,
`"id":"4d1blo5cuycfaca"`, `"id":"4d1blo5cuycfaca"`,
`"id":"9n89pl5vkct6330"`, `"id":"9n89pl5vkct6330"`,
`"id":"ib3m2700k5hlsjz"`,
`"type":"auth"`, `"type":"auth"`,
`"type":"base"`, `"type":"base"`,
}, },
@@ -73,9 +77,9 @@ func TestCollectionsList(t *testing.T) {
ExpectedContent: []string{ ExpectedContent: []string{
`"page":2`, `"page":2`,
`"perPage":2`, `"perPage":2`,
`"totalItems":10`, `"totalItems":11`,
`"items":[{`, `"items":[{`,
`"id":"kpv709sk2lqbqk8"`, `"id":"v9gwnfh02gjq1q0"`,
`"id":"9n89pl5vkct6330"`, `"id":"9n89pl5vkct6330"`,
}, },
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
@@ -123,6 +127,8 @@ func TestCollectionsList(t *testing.T) {
} }
func TestCollectionView(t *testing.T) { func TestCollectionView(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -191,6 +197,8 @@ func TestCollectionView(t *testing.T) {
} }
func TestCollectionDelete(t *testing.T) { func TestCollectionDelete(t *testing.T) {
t.Parallel()
ensureDeletedFiles := func(app *tests.TestApp, collectionId string) { ensureDeletedFiles := func(app *tests.TestApp, collectionId string) {
storageDir := filepath.Join(app.DataDir(), "storage", collectionId) storageDir := filepath.Join(app.DataDir(), "storage", collectionId)
@@ -243,7 +251,7 @@ func TestCollectionDelete(t *testing.T) {
"OnCollectionBeforeDeleteRequest": 1, "OnCollectionBeforeDeleteRequest": 1,
"OnCollectionAfterDeleteRequest": 1, "OnCollectionAfterDeleteRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureDeletedFiles(app, "9n89pl5vkct6330") ensureDeletedFiles(app, "9n89pl5vkct6330")
}, },
}, },
@@ -262,7 +270,7 @@ func TestCollectionDelete(t *testing.T) {
"OnCollectionBeforeDeleteRequest": 1, "OnCollectionBeforeDeleteRequest": 1,
"OnCollectionAfterDeleteRequest": 1, "OnCollectionAfterDeleteRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureDeletedFiles(app, "9n89pl5vkct6330") ensureDeletedFiles(app, "9n89pl5vkct6330")
}, },
}, },
@@ -299,7 +307,6 @@ func TestCollectionDelete(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
Delay: 100 * time.Millisecond,
ExpectedStatus: 204, ExpectedStatus: 204,
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnModelBeforeDelete": 1, "OnModelBeforeDelete": 1,
@@ -308,6 +315,27 @@ func TestCollectionDelete(t *testing.T) {
"OnCollectionAfterDeleteRequest": 1, "OnCollectionAfterDeleteRequest": 1,
}, },
}, },
{
Name: "OnCollectionAfterDeleteRequest error response",
Method: http.MethodDelete,
Url: "/api/collections/view2",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnCollectionAfterDeleteRequest().Add(func(e *core.CollectionDeleteEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeDelete": 1,
"OnModelAfterDelete": 1,
"OnCollectionBeforeDeleteRequest": 1,
"OnCollectionAfterDeleteRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -316,6 +344,8 @@ func TestCollectionDelete(t *testing.T) {
} }
func TestCollectionCreate(t *testing.T) { func TestCollectionCreate(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -378,7 +408,7 @@ func TestCollectionCreate(t *testing.T) {
`"name":"new"`, `"name":"new"`,
`"type":"base"`, `"type":"base"`,
`"system":false`, `"system":false`,
`"schema":[{"system":false,"id":"12345789","name":"test","type":"text","required":false,"unique":false,"options":{"min":null,"max":null,"pattern":""}}]`, `"schema":[{"system":false,"id":"12345789","name":"test","type":"text","required":false,"presentable":false,"unique":false,"options":{"min":null,"max":null,"pattern":""}}]`,
`"options":{}`, `"options":{}`,
}, },
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
@@ -402,8 +432,8 @@ func TestCollectionCreate(t *testing.T) {
`"name":"new"`, `"name":"new"`,
`"type":"auth"`, `"type":"auth"`,
`"system":false`, `"system":false`,
`"schema":[{"system":false,"id":"12345789","name":"test","type":"text","required":false,"unique":false,"options":{"min":null,"max":null,"pattern":""}}]`, `"schema":[{"system":false,"id":"12345789","name":"test","type":"text","required":false,"presentable":false,"unique":false,"options":{"min":null,"max":null,"pattern":""}}]`,
`"options":{"allowEmailAuth":false,"allowOAuth2Auth":false,"allowUsernameAuth":false,"exceptEmailDomains":null,"manageRule":null,"minPasswordLength":0,"onlyEmailDomains":null,"requireEmail":false}`, `"options":{"allowEmailAuth":false,"allowOAuth2Auth":false,"allowUsernameAuth":false,"exceptEmailDomains":null,"manageRule":null,"minPasswordLength":0,"onlyEmailDomains":null,"onlyVerified":false,"requireEmail":false}`,
}, },
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnModelBeforeCreate": 1, "OnModelBeforeCreate": 1,
@@ -536,6 +566,28 @@ func TestCollectionCreate(t *testing.T) {
`"options":{"minPasswordLength":{"code":"validation_required"`, `"options":{"minPasswordLength":{"code":"validation_required"`,
}, },
}, },
{
Name: "OnCollectionAfterCreateRequest error response",
Method: http.MethodPost,
Url: "/api/collections",
Body: strings.NewReader(`{"name":"new","type":"base","schema":[{"type":"text","id":"12345789","name":"test"}]}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnCollectionAfterCreateRequest().Add(func(e *core.CollectionCreateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeCreate": 1,
"OnModelAfterCreate": 1,
"OnCollectionBeforeCreateRequest": 1,
"OnCollectionAfterCreateRequest": 1,
},
},
// view // view
// ----------------------------------------------------------- // -----------------------------------------------------------
@@ -649,7 +701,7 @@ func TestCollectionCreate(t *testing.T) {
"OnCollectionBeforeCreateRequest": 1, "OnCollectionBeforeCreateRequest": 1,
"OnCollectionAfterCreateRequest": 1, "OnCollectionAfterCreateRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
indexes, err := app.Dao().TableIndexes("new") indexes, err := app.Dao().TableIndexes("new")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -671,6 +723,8 @@ func TestCollectionCreate(t *testing.T) {
} }
func TestCollectionUpdate(t *testing.T) { func TestCollectionUpdate(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -720,6 +774,28 @@ func TestCollectionUpdate(t *testing.T) {
"OnModelBeforeUpdate": 1, "OnModelBeforeUpdate": 1,
}, },
}, },
{
Name: "OnCollectionAfterUpdateRequest error response",
Method: http.MethodPatch,
Url: "/api/collections/demo1",
Body: strings.NewReader(`{}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnCollectionAfterUpdateRequest().Add(func(e *core.CollectionUpdateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnCollectionAfterUpdateRequest": 1,
"OnCollectionBeforeUpdateRequest": 1,
"OnModelAfterUpdate": 1,
"OnModelBeforeUpdate": 1,
},
},
{ {
Name: "authorized as admin + invalid data (eg. existing name)", Name: "authorized as admin + invalid data (eg. existing name)",
Method: http.MethodPatch, Method: http.MethodPatch,
@@ -757,7 +833,7 @@ func TestCollectionUpdate(t *testing.T) {
"OnCollectionBeforeUpdateRequest": 1, "OnCollectionBeforeUpdateRequest": 1,
"OnCollectionAfterUpdateRequest": 1, "OnCollectionAfterUpdateRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
// check if the record table was renamed // check if the record table was renamed
if !app.Dao().HasTable("new") { if !app.Dao().HasTable("new") {
t.Fatal("Couldn't find record table 'new'.") t.Fatal("Couldn't find record table 'new'.")
@@ -817,7 +893,8 @@ func TestCollectionUpdate(t *testing.T) {
{"type":"text","name":"password"}, {"type":"text","name":"password"},
{"type":"text","name":"passwordConfirm"}, {"type":"text","name":"passwordConfirm"},
{"type":"text","name":"oldPassword"} {"type":"text","name":"oldPassword"}
] ],
"indexes": []
}`), }`),
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
@@ -893,88 +970,6 @@ func TestCollectionUpdate(t *testing.T) {
}, },
}, },
// rel field change displayFields propagation
// -----------------------------------------------------------
{
Name: "renaming a display field should also update the referenced displayFields value",
Method: http.MethodPatch,
Url: "/api/collections/demo3",
Body: strings.NewReader(`{
"schema":[
{
"id": "w5z2x0nq",
"type": "text",
"name": "title_change"
}
]
}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
ExpectedStatus: 200,
ExpectedContent: []string{
`"name":"title_change"`,
},
ExpectedEvents: map[string]int{
"OnModelBeforeUpdate": 2,
"OnModelAfterUpdate": 2,
"OnCollectionBeforeUpdateRequest": 1,
"OnCollectionAfterUpdateRequest": 1,
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
collection, err := app.Dao().FindCollectionByNameOrId("demo4")
if err != nil {
t.Fatal(err)
}
relField := collection.Schema.GetFieldByName("rel_many_no_cascade_required")
options := relField.Options.(*schema.RelationOptions)
expectedDisplayFields := []string{"title_change", "id"}
if len(list.SubtractSlice(options.DisplayFields, expectedDisplayFields)) != 0 {
t.Fatalf("Expected displayFields %v, got %v", expectedDisplayFields, options.DisplayFields)
}
},
},
{
Name: "deleting a display field should also update the referenced displayFields value",
Method: http.MethodPatch,
Url: "/api/collections/demo3",
Body: strings.NewReader(`{
"schema":[
{
"type": "text",
"name": "new_field"
}
]
}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
ExpectedStatus: 200,
ExpectedContent: []string{
`"name":"new_field"`,
},
ExpectedEvents: map[string]int{
"OnModelBeforeUpdate": 2,
"OnModelAfterUpdate": 2,
"OnCollectionBeforeUpdateRequest": 1,
"OnCollectionAfterUpdateRequest": 1,
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
collection, err := app.Dao().FindCollectionByNameOrId("demo4")
if err != nil {
t.Fatal(err)
}
relField := collection.Schema.GetFieldByName("rel_many_no_cascade_required")
options := relField.Options.(*schema.RelationOptions)
expectedDisplayFields := []string{"id"}
if len(list.SubtractSlice(options.DisplayFields, expectedDisplayFields)) != 0 {
t.Fatalf("Expected displayFields %v, got %v", expectedDisplayFields, options.DisplayFields)
}
},
},
// view // view
// ----------------------------------------------------------- // -----------------------------------------------------------
{ {
@@ -1076,7 +1071,7 @@ func TestCollectionUpdate(t *testing.T) {
"OnCollectionBeforeUpdateRequest": 1, "OnCollectionBeforeUpdateRequest": 1,
"OnCollectionAfterUpdateRequest": 1, "OnCollectionAfterUpdateRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
indexes, err := app.Dao().TableIndexes("new") indexes, err := app.Dao().TableIndexes("new")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1098,7 +1093,9 @@ func TestCollectionUpdate(t *testing.T) {
} }
func TestCollectionsImport(t *testing.T) { func TestCollectionsImport(t *testing.T) {
totalCollections := 10 t.Parallel()
totalCollections := 11
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
@@ -1131,7 +1128,7 @@ func TestCollectionsImport(t *testing.T) {
`"data":{`, `"data":{`,
`"collections":{"code":"validation_required"`, `"collections":{"code":"validation_required"`,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
collections := []*models.Collection{} collections := []*models.Collection{}
if err := app.Dao().CollectionQuery().All(&collections); err != nil { if err := app.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1157,9 +1154,9 @@ func TestCollectionsImport(t *testing.T) {
}, },
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnCollectionsBeforeImportRequest": 1, "OnCollectionsBeforeImportRequest": 1,
"OnModelBeforeDelete": 4, "OnModelBeforeDelete": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
collections := []*models.Collection{} collections := []*models.Collection{}
if err := app.Dao().CollectionQuery().All(&collections); err != nil { if err := app.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1201,7 +1198,7 @@ func TestCollectionsImport(t *testing.T) {
"OnCollectionsBeforeImportRequest": 1, "OnCollectionsBeforeImportRequest": 1,
"OnModelBeforeCreate": 2, "OnModelBeforeCreate": 2,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
collections := []*models.Collection{} collections := []*models.Collection{}
if err := app.Dao().CollectionQuery().All(&collections); err != nil { if err := app.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1257,7 +1254,7 @@ func TestCollectionsImport(t *testing.T) {
"OnModelBeforeCreate": 3, "OnModelBeforeCreate": 3,
"OnModelAfterCreate": 3, "OnModelAfterCreate": 3,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
collections := []*models.Collection{} collections := []*models.Collection{}
if err := app.Dao().CollectionQuery().All(&collections); err != nil { if err := app.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1355,14 +1352,14 @@ func TestCollectionsImport(t *testing.T) {
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnCollectionsAfterImportRequest": 1, "OnCollectionsAfterImportRequest": 1,
"OnCollectionsBeforeImportRequest": 1, "OnCollectionsBeforeImportRequest": 1,
"OnModelBeforeDelete": 8, "OnModelBeforeDelete": 9,
"OnModelAfterDelete": 8, "OnModelAfterDelete": 9,
"OnModelBeforeUpdate": 2, "OnModelBeforeUpdate": 2,
"OnModelAfterUpdate": 2, "OnModelAfterUpdate": 2,
"OnModelBeforeCreate": 1, "OnModelBeforeCreate": 1,
"OnModelAfterCreate": 1, "OnModelAfterCreate": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
collections := []*models.Collection{} collections := []*models.Collection{}
if err := app.Dao().CollectionQuery().All(&collections); err != nil { if err := app.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1373,6 +1370,58 @@ func TestCollectionsImport(t *testing.T) {
} }
}, },
}, },
{
Name: "authorized as admin + successful collections save",
Method: http.MethodPut,
Url: "/api/collections/import",
Body: strings.NewReader(`{
"collections":[
{
"name": "import1",
"schema": [
{
"id": "koih1lqx",
"name": "test",
"type": "text"
}
]
},
{
"name": "import2",
"schema": [
{
"id": "koih1lqx",
"name": "test",
"type": "text"
}
],
"indexes": [
"create index idx_test on import2 (test)"
]
},
{
"name": "auth_without_schema",
"type": "auth"
}
]
}`),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnCollectionsAfterImportRequest().Add(func(e *core.CollectionsImportEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnCollectionsBeforeImportRequest": 1,
"OnCollectionsAfterImportRequest": 1,
"OnModelBeforeCreate": 3,
"OnModelAfterCreate": 3,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
+87 -62
View File
@@ -1,23 +1,26 @@
package apis package apis
import ( import (
"context"
"errors" "errors"
"fmt" "fmt"
"log" "log/slog"
"net/http" "net/http"
"runtime"
"strings" "strings"
"time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/models/schema" "github.com/pocketbase/pocketbase/models/schema"
"github.com/pocketbase/pocketbase/resolvers"
"github.com/pocketbase/pocketbase/tokens" "github.com/pocketbase/pocketbase/tokens"
"github.com/pocketbase/pocketbase/tools/filesystem"
"github.com/pocketbase/pocketbase/tools/list" "github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/search"
"github.com/pocketbase/pocketbase/tools/security" "github.com/pocketbase/pocketbase/tools/security"
"github.com/spf13/cast" "github.com/spf13/cast"
"golang.org/x/sync/semaphore"
"golang.org/x/sync/singleflight"
) )
var imageContentTypes = []string{"image/png", "image/jpg", "image/jpeg", "image/gif"} var imageContentTypes = []string{"image/png", "image/jpg", "image/jpeg", "image/gif"}
@@ -25,7 +28,12 @@ var defaultThumbSizes = []string{"100x100"}
// bindFileApi registers the file api endpoints and the corresponding handlers. // bindFileApi registers the file api endpoints and the corresponding handlers.
func bindFileApi(app core.App, rg *echo.Group) { func bindFileApi(app core.App, rg *echo.Group) {
api := fileApi{app: app} api := fileApi{
app: app,
thumbGenSem: semaphore.NewWeighted(int64(runtime.NumCPU() + 2)), // the value is arbitrary chosen and may change in the future
thumbGenPending: new(singleflight.Group),
thumbGenMaxWait: 60 * time.Second,
}
subGroup := rg.Group("/files", ActivityLogger(app)) subGroup := rg.Group("/files", ActivityLogger(app))
subGroup.POST("/token", api.fileToken) subGroup.POST("/token", api.fileToken)
@@ -35,6 +43,18 @@ func bindFileApi(app core.App, rg *echo.Group) {
type fileApi struct { type fileApi struct {
app core.App app core.App
// thumbGenSem is a semaphore to prevent too much concurrent
// requests generating new thumbs at the same time.
thumbGenSem *semaphore.Weighted
// thumbGenPending represents a group of currently pending
// thumb generation processes.
thumbGenPending *singleflight.Group
// thumbGenMaxWait is the maximum waiting time for starting a new
// thumb generation process.
thumbGenMaxWait time.Duration
} }
func (api *fileApi) fileToken(c echo.Context) error { func (api *fileApi) fileToken(c echo.Context) error {
@@ -49,23 +69,21 @@ func (api *fileApi) fileToken(c echo.Context) error {
event.Token, _ = tokens.NewRecordFileToken(api.app, record) event.Token, _ = tokens.NewRecordFileToken(api.app, record)
} }
handlerErr := api.app.OnFileBeforeTokenRequest().Trigger(event, func(e *core.FileTokenEvent) error { return api.app.OnFileBeforeTokenRequest().Trigger(event, func(e *core.FileTokenEvent) error {
if e.Model == nil || e.Token == "" { if e.Model == nil || e.Token == "" {
return NewBadRequestError("Failed to generate file token.", nil) return NewBadRequestError("Failed to generate file token.", nil)
} }
return e.HttpContext.JSON(http.StatusOK, map[string]string{ return api.app.OnFileAfterTokenRequest().Trigger(event, func(e *core.FileTokenEvent) error {
"token": e.Token, if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, map[string]string{
"token": e.Token,
})
}) })
}) })
if handlerErr == nil {
if err := api.app.OnFileAfterTokenRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return handlerErr
} }
func (api *fileApi) download(c echo.Context) error { func (api *fileApi) download(c echo.Context) error {
@@ -102,7 +120,20 @@ func (api *fileApi) download(c echo.Context) error {
adminOrAuthRecord, _ := api.findAdminOrAuthRecordByFileToken(token) adminOrAuthRecord, _ := api.findAdminOrAuthRecordByFileToken(token)
if !api.canAccessRecord(adminOrAuthRecord, record, record.Collection().ViewRule) { // create a copy of the cached request data and adjust it for the current auth model
requestInfo := *RequestInfo(c)
requestInfo.Context = models.RequestInfoContextProtectedFile
requestInfo.Admin = nil
requestInfo.AuthRecord = nil
if adminOrAuthRecord != nil {
if admin, _ := adminOrAuthRecord.(*models.Admin); admin != nil {
requestInfo.Admin = admin
} else if record, _ := adminOrAuthRecord.(*models.Record); record != nil {
requestInfo.AuthRecord = record
}
}
if ok, _ := api.app.Dao().CanAccessRecord(record, &requestInfo, record.Collection().ViewRule); !ok {
return NewForbiddenError("Insufficient permissions to access the file resource.", nil) return NewForbiddenError("Insufficient permissions to access the file resource.", nil)
} }
} }
@@ -118,11 +149,11 @@ func (api *fileApi) download(c echo.Context) error {
baseFilesPath = fileRecord.BaseFilesPath() baseFilesPath = fileRecord.BaseFilesPath()
} }
fs, err := api.app.NewFilesystem() fsys, err := api.app.NewFilesystem()
if err != nil { if err != nil {
return NewBadRequestError("Filesystem initialization failure.", err) return NewBadRequestError("Filesystem initialization failure.", err)
} }
defer fs.Close() defer fsys.Close()
originalPath := baseFilesPath + "/" + filename originalPath := baseFilesPath + "/" + filename
servedPath := originalPath servedPath := originalPath
@@ -132,7 +163,7 @@ func (api *fileApi) download(c echo.Context) error {
thumbSize := c.QueryParam("thumb") thumbSize := c.QueryParam("thumb")
if thumbSize != "" && (list.ExistInSlice(thumbSize, defaultThumbSizes) || list.ExistInSlice(thumbSize, options.Thumbs)) { if thumbSize != "" && (list.ExistInSlice(thumbSize, defaultThumbSizes) || list.ExistInSlice(thumbSize, options.Thumbs)) {
// extract the original file meta attributes and check it existence // extract the original file meta attributes and check it existence
oAttrs, oAttrsErr := fs.Attributes(originalPath) oAttrs, oAttrsErr := fsys.Attributes(originalPath)
if oAttrsErr != nil { if oAttrsErr != nil {
return NewNotFoundError("", err) return NewNotFoundError("", err)
} }
@@ -143,10 +174,19 @@ func (api *fileApi) download(c echo.Context) error {
servedName = thumbSize + "_" + filename servedName = thumbSize + "_" + filename
servedPath = baseFilesPath + "/thumbs_" + filename + "/" + servedName servedPath = baseFilesPath + "/thumbs_" + filename + "/" + servedName
// create a new thumb if it doesn exists // create a new thumb if it doesn't exist
if exists, _ := fs.Exists(servedPath); !exists { if exists, _ := fsys.Exists(servedPath); !exists {
if err := fs.CreateThumb(originalPath, servedPath, thumbSize); err != nil { if err := api.createThumb(c, fsys, originalPath, servedPath, thumbSize); err != nil {
servedPath = originalPath // fallback to the original api.app.Logger().Warn(
"Fallback to original - failed to create thumb "+servedName,
slog.Any("error", err),
slog.String("original", originalPath),
slog.String("thumb", servedPath),
)
// fallback to the original
servedName = filename
servedPath = originalPath
} }
} }
} }
@@ -166,9 +206,11 @@ func (api *fileApi) download(c echo.Context) error {
c.Response().Header().Del("X-Frame-Options") c.Response().Header().Del("X-Frame-Options")
return api.app.OnFileDownloadRequest().Trigger(event, func(e *core.FileDownloadEvent) error { return api.app.OnFileDownloadRequest().Trigger(event, func(e *core.FileDownloadEvent) error {
res := e.HttpContext.Response() if e.HttpContext.Response().Committed {
req := e.HttpContext.Request() return nil
if err := fs.Serve(res, req, e.ServedPath, e.ServedName); err != nil { }
if err := fsys.Serve(e.HttpContext.Response(), e.HttpContext.Request(), e.ServedPath, e.ServedName); err != nil {
return NewNotFoundError("", err) return NewNotFoundError("", err)
} }
@@ -207,45 +249,28 @@ func (api *fileApi) findAdminOrAuthRecordByFileToken(fileToken string) (models.M
return nil, errors.New("missing or invalid file token") return nil, errors.New("missing or invalid file token")
} }
// @todo move to a helper and maybe combine with the realtime checks when refactoring the realtime service func (api *fileApi) createThumb(
func (api *fileApi) canAccessRecord(adminOrAuthRecord models.Model, record *models.Record, accessRule *string) bool { c echo.Context,
admin, _ := adminOrAuthRecord.(*models.Admin) fsys *filesystem.System,
if admin != nil { originalPath string,
// admins can access everything thumbPath string,
return true thumbSize string,
} ) error {
ch := api.thumbGenPending.DoChan(thumbPath, func() (any, error) {
ctx, cancel := context.WithTimeout(c.Request().Context(), api.thumbGenMaxWait)
defer cancel()
if accessRule == nil { if err := api.thumbGenSem.Acquire(ctx, 1); err != nil {
// only admins can access this record return nil, err
return false
}
ruleFunc := func(q *dbx.SelectQuery) error {
if *accessRule == "" {
return nil // empty public rule
} }
defer api.thumbGenSem.Release(1)
// mock request data return nil, fsys.CreateThumb(originalPath, thumbPath, thumbSize)
requestData := &models.RequestData{ })
Method: "GET",
}
requestData.AuthRecord, _ = adminOrAuthRecord.(*models.Record)
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), record.Collection(), requestData, true) res := <-ch
expr, err := search.FilterData(*accessRule).BuildExpr(resolver)
if err != nil {
return err
}
resolver.UpdateQuery(q)
q.AndWhere(expr)
return nil api.thumbGenPending.Forget(thumbPath)
}
foundRecord, err := api.app.Dao().FindRecordById(record.Collection().Id, record.Id, ruleFunc) return res.Err
if err == nil && foundRecord != nil {
return true
}
return false
} }
+87
View File
@@ -2,20 +2,26 @@ package apis_test
import ( import (
"net/http" "net/http"
"net/http/httptest"
"os" "os"
"path" "path"
"path/filepath" "path/filepath"
"runtime" "runtime"
"sync"
"testing" "testing"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/apis"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/models/schema"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/types" "github.com/pocketbase/pocketbase/tools/types"
) )
func TestFileToken(t *testing.T) { func TestFileToken(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -88,6 +94,8 @@ func TestFileToken(t *testing.T) {
} }
func TestFileDownload(t *testing.T) { func TestFileDownload(t *testing.T) {
t.Parallel()
_, currentFile, _, _ := runtime.Caller(0) _, currentFile, _, _ := runtime.Caller(0)
dataDirRelPath := "../tests/data/" dataDirRelPath := "../tests/data/"
@@ -385,3 +393,82 @@ func TestFileDownload(t *testing.T) {
scenario.Test(t) scenario.Test(t)
} }
} }
func TestConcurrentThumbsGeneration(t *testing.T) {
t.Parallel()
app, err := tests.NewTestApp()
if err != nil {
t.Fatal(err)
}
defer app.Cleanup()
fsys, err := app.NewFilesystem()
if err != nil {
t.Fatal(err)
}
defer fsys.Close()
// create a dummy file field collection
demo1, err := app.Dao().FindCollectionByNameOrId("demo1")
if err != nil {
t.Fatal(err)
}
fileField := demo1.Schema.GetFieldByName("file_one")
fileField.Options = &schema.FileOptions{
Protected: false,
MaxSelect: 1,
MaxSize: 999999,
// new thumbs
Thumbs: []string{"111x111", "111x222", "111x333"},
}
demo1.Schema.AddField(fileField)
if err := app.Dao().SaveCollection(demo1); err != nil {
t.Fatal(err)
}
fileKey := "wsmn24bux7wo113/al1h9ijdeojtsjy/300_Jsjq7RdBgA.png"
e, err := apis.InitApi(app)
if err != nil {
t.Fatal(err)
}
urls := []string{
"/api/files/" + fileKey + "?thumb=111x111",
"/api/files/" + fileKey + "?thumb=111x111", // should still result in single thumb
"/api/files/" + fileKey + "?thumb=111x222",
"/api/files/" + fileKey + "?thumb=111x333",
}
var wg sync.WaitGroup
wg.Add(len(urls))
for _, url := range urls {
url := url
go func() {
defer wg.Done()
recorder := httptest.NewRecorder()
req := httptest.NewRequest("GET", url, nil)
e.ServeHTTP(recorder, req)
}()
}
wg.Wait()
// ensure that all new requested thumbs were created
thumbKeys := []string{
"wsmn24bux7wo113/al1h9ijdeojtsjy/thumbs_300_Jsjq7RdBgA.png/111x111_" + filepath.Base(fileKey),
"wsmn24bux7wo113/al1h9ijdeojtsjy/thumbs_300_Jsjq7RdBgA.png/111x222_" + filepath.Base(fileKey),
"wsmn24bux7wo113/al1h9ijdeojtsjy/thumbs_300_Jsjq7RdBgA.png/111x333_" + filepath.Base(fileKey),
}
for _, k := range thumbKeys {
if exists, _ := fsys.Exists(k); !exists {
t.Fatalf("Missing thumb %q: %v", k, err)
}
}
}
+7 -2
View File
@@ -12,6 +12,7 @@ func bindHealthApi(app core.App, rg *echo.Group) {
api := healthApi{app: app} api := healthApi{app: app}
subGroup := rg.Group("/health") subGroup := rg.Group("/health")
subGroup.HEAD("", api.healthCheck)
subGroup.GET("", api.healthCheck) subGroup.GET("", api.healthCheck)
} }
@@ -20,8 +21,8 @@ type healthApi struct {
} }
type healthCheckResponse struct { type healthCheckResponse struct {
Code int `json:"code"`
Message string `json:"message"` Message string `json:"message"`
Code int `json:"code"`
Data struct { Data struct {
CanBackup bool `json:"canBackup"` CanBackup bool `json:"canBackup"`
} `json:"data"` } `json:"data"`
@@ -29,10 +30,14 @@ type healthCheckResponse struct {
// healthCheck returns a 200 OK response if the server is healthy. // healthCheck returns a 200 OK response if the server is healthy.
func (api *healthApi) healthCheck(c echo.Context) error { func (api *healthApi) healthCheck(c echo.Context) error {
if c.Request().Method == http.MethodHead {
return c.NoContent(http.StatusOK)
}
resp := new(healthCheckResponse) resp := new(healthCheckResponse)
resp.Code = http.StatusOK resp.Code = http.StatusOK
resp.Message = "API is healthy." resp.Message = "API is healthy."
resp.Data.CanBackup = !api.app.Cache().Has(core.CacheKeyActiveBackup) resp.Data.CanBackup = !api.app.Store().Has(core.StoreKeyActiveBackup)
return c.JSON(http.StatusOK, resp) return c.JSON(http.StatusOK, resp)
} }
+9 -1
View File
@@ -8,9 +8,17 @@ import (
) )
func TestHealthAPI(t *testing.T) { func TestHealthAPI(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "health status returns 200", Name: "HEAD health status",
Method: http.MethodHead,
Url: "/api/health",
ExpectedStatus: 200,
},
{
Name: "GET health status",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/health", Url: "/api/health",
ExpectedStatus: 200, ExpectedStatus: 200,
+18 -18
View File
@@ -15,27 +15,27 @@ func bindLogsApi(app core.App, rg *echo.Group) {
api := logsApi{app: app} api := logsApi{app: app}
subGroup := rg.Group("/logs", RequireAdminAuth()) subGroup := rg.Group("/logs", RequireAdminAuth())
subGroup.GET("/requests", api.requestsList) subGroup.GET("", api.list)
subGroup.GET("/requests/stats", api.requestsStats) subGroup.GET("/stats", api.stats)
subGroup.GET("/requests/:id", api.requestView) subGroup.GET("/:id", api.view)
} }
type logsApi struct { type logsApi struct {
app core.App app core.App
} }
var requestFilterFields = []string{ var logFilterFields = []string{
"rowid", "id", "created", "updated", "rowid", "id", "created", "updated",
"url", "method", "status", "auth", "level", "message", "data",
"remoteIp", "userIp", "referer", "userAgent", `^data\.[\w\.\:]*\w+$`,
} }
func (api *logsApi) requestsList(c echo.Context) error { func (api *logsApi) list(c echo.Context) error {
fieldResolver := search.NewSimpleFieldResolver(requestFilterFields...) fieldResolver := search.NewSimpleFieldResolver(logFilterFields...)
result, err := search.NewProvider(fieldResolver). result, err := search.NewProvider(fieldResolver).
Query(api.app.LogsDao().RequestQuery()). Query(api.app.LogsDao().LogQuery()).
ParseAndExec(c.QueryParams().Encode(), &[]*models.Request{}) ParseAndExec(c.QueryParams().Encode(), &[]*models.Log{})
if err != nil { if err != nil {
return NewBadRequestError("", err) return NewBadRequestError("", err)
@@ -44,8 +44,8 @@ func (api *logsApi) requestsList(c echo.Context) error {
return c.JSON(http.StatusOK, result) return c.JSON(http.StatusOK, result)
} }
func (api *logsApi) requestsStats(c echo.Context) error { func (api *logsApi) stats(c echo.Context) error {
fieldResolver := search.NewSimpleFieldResolver(requestFilterFields...) fieldResolver := search.NewSimpleFieldResolver(logFilterFields...)
filter := c.QueryParam(search.FilterQueryParam) filter := c.QueryParam(search.FilterQueryParam)
@@ -58,24 +58,24 @@ func (api *logsApi) requestsStats(c echo.Context) error {
} }
} }
stats, err := api.app.LogsDao().RequestsStats(expr) stats, err := api.app.LogsDao().LogsStats(expr)
if err != nil { if err != nil {
return NewBadRequestError("Failed to generate requests stats.", err) return NewBadRequestError("Failed to generate logs stats.", err)
} }
return c.JSON(http.StatusOK, stats) return c.JSON(http.StatusOK, stats)
} }
func (api *logsApi) requestView(c echo.Context) error { func (api *logsApi) view(c echo.Context) error {
id := c.PathParam("id") id := c.PathParam("id")
if id == "" { if id == "" {
return NewNotFoundError("", nil) return NewNotFoundError("", nil)
} }
request, err := api.app.LogsDao().FindRequestById(id) log, err := api.app.LogsDao().FindLogById(id)
if err != nil || request == nil { if err != nil || log == nil {
return NewNotFoundError("", err) return NewNotFoundError("", err)
} }
return c.JSON(http.StatusOK, request) return c.JSON(http.StatusOK, log)
} }
+27 -21
View File
@@ -8,19 +8,21 @@ import (
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
) )
func TestRequestsList(t *testing.T) { func TestLogsList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests", Url: "/api/logs",
ExpectedStatus: 401, ExpectedStatus: 401,
ExpectedContent: []string{`"data":{}`}, ExpectedContent: []string{`"data":{}`},
}, },
{ {
Name: "authorized as auth record", Name: "authorized as auth record",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests", Url: "/api/logs",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc", "Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc",
}, },
@@ -30,12 +32,12 @@ func TestRequestsList(t *testing.T) {
{ {
Name: "authorized as admin", Name: "authorized as admin",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests", Url: "/api/logs",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := tests.MockRequestLogsData(app); err != nil { if err := tests.MockLogsData(app); err != nil {
t.Fatal(err) t.Fatal(err)
} }
}, },
@@ -52,12 +54,12 @@ func TestRequestsList(t *testing.T) {
{ {
Name: "authorized as admin + filter", Name: "authorized as admin + filter",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests?filter=status>200", Url: "/api/logs?filter=data.status>200",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := tests.MockRequestLogsData(app); err != nil { if err := tests.MockLogsData(app); err != nil {
t.Fatal(err) t.Fatal(err)
} }
}, },
@@ -77,19 +79,21 @@ func TestRequestsList(t *testing.T) {
} }
} }
func TestRequestView(t *testing.T) { func TestLogView(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/873f2133-9f38-44fb-bf82-c8f53b310d91", Url: "/api/logs/873f2133-9f38-44fb-bf82-c8f53b310d91",
ExpectedStatus: 401, ExpectedStatus: 401,
ExpectedContent: []string{`"data":{}`}, ExpectedContent: []string{`"data":{}`},
}, },
{ {
Name: "authorized as auth record", Name: "authorized as auth record",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/873f2133-9f38-44fb-bf82-c8f53b310d91", Url: "/api/logs/873f2133-9f38-44fb-bf82-c8f53b310d91",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc", "Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc",
}, },
@@ -99,12 +103,12 @@ func TestRequestView(t *testing.T) {
{ {
Name: "authorized as admin (nonexisting request log)", Name: "authorized as admin (nonexisting request log)",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/missing1-9f38-44fb-bf82-c8f53b310d91", Url: "/api/logs/missing1-9f38-44fb-bf82-c8f53b310d91",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := tests.MockRequestLogsData(app); err != nil { if err := tests.MockLogsData(app); err != nil {
t.Fatal(err) t.Fatal(err)
} }
}, },
@@ -114,12 +118,12 @@ func TestRequestView(t *testing.T) {
{ {
Name: "authorized as admin (existing request log)", Name: "authorized as admin (existing request log)",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/873f2133-9f38-44fb-bf82-c8f53b310d91", Url: "/api/logs/873f2133-9f38-44fb-bf82-c8f53b310d91",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := tests.MockRequestLogsData(app); err != nil { if err := tests.MockLogsData(app); err != nil {
t.Fatal(err) t.Fatal(err)
} }
}, },
@@ -135,19 +139,21 @@ func TestRequestView(t *testing.T) {
} }
} }
func TestRequestsStats(t *testing.T) { func TestLogsStats(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/stats", Url: "/api/logs/stats",
ExpectedStatus: 401, ExpectedStatus: 401,
ExpectedContent: []string{`"data":{}`}, ExpectedContent: []string{`"data":{}`},
}, },
{ {
Name: "authorized as auth record", Name: "authorized as auth record",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/stats", Url: "/api/logs/stats",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc", "Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc",
}, },
@@ -157,12 +163,12 @@ func TestRequestsStats(t *testing.T) {
{ {
Name: "authorized as admin", Name: "authorized as admin",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/stats", Url: "/api/logs/stats",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := tests.MockRequestLogsData(app); err != nil { if err := tests.MockLogsData(app); err != nil {
t.Fatal(err) t.Fatal(err)
} }
}, },
@@ -174,12 +180,12 @@ func TestRequestsStats(t *testing.T) {
{ {
Name: "authorized as admin + filter", Name: "authorized as admin + filter",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/logs/requests/stats?filter=status>200", Url: "/api/logs/stats?filter=data.status>200",
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
if err := tests.MockRequestLogsData(app); err != nil { if err := tests.MockLogsData(app); err != nil {
t.Fatal(err) t.Fatal(err)
} }
}, },
+87 -76
View File
@@ -2,9 +2,10 @@ package apis
import ( import (
"fmt" "fmt"
"log" "log/slog"
"net" "net"
"net/http" "net/http"
"net/url"
"strings" "strings"
"time" "time"
@@ -15,7 +16,6 @@ import (
"github.com/pocketbase/pocketbase/tools/list" "github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/routine" "github.com/pocketbase/pocketbase/tools/routine"
"github.com/pocketbase/pocketbase/tools/security" "github.com/pocketbase/pocketbase/tools/security"
"github.com/pocketbase/pocketbase/tools/types"
"github.com/spf13/cast" "github.com/spf13/cast"
) )
@@ -24,6 +24,7 @@ const (
ContextAdminKey string = "admin" ContextAdminKey string = "admin"
ContextAuthRecordKey string = "authRecord" ContextAuthRecordKey string = "authRecord"
ContextCollectionKey string = "collection" ContextCollectionKey string = "collection"
ContextExecStartKey string = "execStart"
) )
// RequireGuestOnly middleware requires a request to NOT have a valid // RequireGuestOnly middleware requires a request to NOT have a valid
@@ -260,7 +261,7 @@ func LoadCollectionContext(app core.App, optCollectionTypes ...string) echo.Midd
return func(next echo.HandlerFunc) echo.HandlerFunc { return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error { return func(c echo.Context) error {
if param := c.PathParam("collection"); param != "" { if param := c.PathParam("collection"); param != "" {
collection, err := app.Dao().FindCollectionByNameOrId(param) collection, err := core.FindCachedCollectionByNameOrId(app, param)
if err != nil || collection == nil { if err != nil || collection == nil {
return NewNotFoundError("", err) return NewNotFoundError("", err)
} }
@@ -285,84 +286,92 @@ func LoadCollectionContext(app core.App, optCollectionTypes ...string) echo.Midd
func ActivityLogger(app core.App) echo.MiddlewareFunc { func ActivityLogger(app core.App) echo.MiddlewareFunc {
return func(next echo.HandlerFunc) echo.HandlerFunc { return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error { return func(c echo.Context) error {
err := next(c) if err := next(c); err != nil {
// no logs retention
if app.Settings().Logs.MaxDays == 0 {
return err return err
} }
httpRequest := c.Request() logRequest(app, c, nil)
httpResponse := c.Response()
status := httpResponse.Status
meta := types.JsonMap{}
if err != nil { return nil
switch v := err.(type) {
case *echo.HTTPError:
status = v.Code
meta["errorMessage"] = v.Message
meta["errorDetails"] = fmt.Sprint(v.Internal)
case *ApiError:
status = v.Code
meta["errorMessage"] = v.Message
meta["errorDetails"] = fmt.Sprint(v.RawData())
default:
status = http.StatusBadRequest
meta["errorMessage"] = v.Error()
}
}
requestAuth := models.RequestAuthGuest
if c.Get(ContextAuthRecordKey) != nil {
requestAuth = models.RequestAuthRecord
} else if c.Get(ContextAdminKey) != nil {
requestAuth = models.RequestAuthAdmin
}
ip, _, _ := net.SplitHostPort(httpRequest.RemoteAddr)
model := &models.Request{
Url: httpRequest.URL.RequestURI(),
Method: strings.ToUpper(httpRequest.Method),
Status: status,
Auth: requestAuth,
UserIp: realUserIp(httpRequest, ip),
RemoteIp: ip,
Referer: httpRequest.Referer(),
UserAgent: httpRequest.UserAgent(),
Meta: meta,
}
// set timestamp fields before firing a new go routine
model.RefreshCreated()
model.RefreshUpdated()
routine.FireAndForget(func() {
if err := app.LogsDao().SaveRequest(model); err != nil && app.IsDebug() {
log.Println("Log save failed:", err)
}
// Delete old request logs
// ---
now := time.Now()
lastLogsDeletedAt := cast.ToTime(app.Cache().Get("lastLogsDeletedAt"))
daysDiff := now.Sub(lastLogsDeletedAt).Hours() * 24
if daysDiff > float64(app.Settings().Logs.MaxDays) {
deleteErr := app.LogsDao().DeleteOldRequests(now.AddDate(0, 0, -1*app.Settings().Logs.MaxDays))
if deleteErr == nil {
app.Cache().Set("lastLogsDeletedAt", now)
} else if app.IsDebug() {
log.Println("Logs delete failed:", deleteErr)
}
}
})
return err
} }
} }
} }
func logRequest(app core.App, c echo.Context, err *ApiError) {
// no logs retention
if app.Settings().Logs.MaxDays == 0 {
return
}
attrs := make([]any, 0, 15)
attrs = append(attrs, slog.String("type", "request"))
started := cast.ToTime(c.Get(ContextExecStartKey))
if !started.IsZero() {
attrs = append(attrs, slog.Float64("execTime", float64(time.Since(started))/float64(time.Millisecond)))
}
httpRequest := c.Request()
httpResponse := c.Response()
method := strings.ToUpper(httpRequest.Method)
status := httpResponse.Status
requestUri := httpRequest.URL.RequestURI()
// parse the request error
if err != nil {
status = err.Code
attrs = append(
attrs,
slog.String("error", err.Message),
slog.Any("details", err.RawData()),
)
}
requestAuth := models.RequestAuthGuest
if c.Get(ContextAuthRecordKey) != nil {
requestAuth = models.RequestAuthRecord
} else if c.Get(ContextAdminKey) != nil {
requestAuth = models.RequestAuthAdmin
}
attrs = append(
attrs,
slog.String("url", requestUri),
slog.String("method", method),
slog.Int("status", status),
slog.String("auth", requestAuth),
slog.String("referer", httpRequest.Referer()),
slog.String("userAgent", httpRequest.UserAgent()),
)
if app.Settings().Logs.LogIp {
ip, _, _ := net.SplitHostPort(httpRequest.RemoteAddr)
attrs = append(
attrs,
slog.String("userIp", realUserIp(httpRequest, ip)),
slog.String("remoteIp", ip),
)
}
// don't block on logs write
routine.FireAndForget(func() {
message := method + " "
if escaped, err := url.PathUnescape(requestUri); err == nil {
message += escaped
} else {
message += requestUri
}
if err != nil {
app.Logger().Error(message, attrs...)
} else {
app.Logger().Info(message, attrs...)
}
})
}
// Returns the "real" user IP from common proxy headers (or fallbackIp if none is found). // Returns the "real" user IP from common proxy headers (or fallbackIp if none is found).
// //
// The returned IP value shouldn't be trusted if not behind a trusted reverse proxy! // The returned IP value shouldn't be trusted if not behind a trusted reverse proxy!
@@ -393,15 +402,17 @@ func realUserIp(r *http.Request, fallbackIp string) string {
return fallbackIp return fallbackIp
} }
// eagerRequestDataCache ensures that the request data is cached in the request // @todo consider removing as this may no longer be needed due to the custom rest.MultiBinder.
//
// eagerRequestInfoCache ensures that the request data is cached in the request
// context to allow reading for example the json request body data more than once. // context to allow reading for example the json request body data more than once.
func eagerRequestDataCache(app core.App) echo.MiddlewareFunc { func eagerRequestInfoCache(app core.App) echo.MiddlewareFunc {
return func(next echo.HandlerFunc) echo.HandlerFunc { return func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error { return func(c echo.Context) error {
switch c.Request().Method { switch c.Request().Method {
// currently we are eagerly caching only the requests with body // currently we are eagerly caching only the requests with body
case "POST", "PUT", "PATCH", "DELETE": case "POST", "PUT", "PATCH", "DELETE":
RequestData(c) RequestInfo(c)
} }
return next(c) return next(c)
+16
View File
@@ -10,6 +10,8 @@ import (
) )
func TestRequireGuestOnly(t *testing.T) { func TestRequireGuestOnly(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "valid record token", Name: "valid record token",
@@ -104,6 +106,8 @@ func TestRequireGuestOnly(t *testing.T) {
} }
func TestRequireRecordAuth(t *testing.T) { func TestRequireRecordAuth(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "guest", Name: "guest",
@@ -242,6 +246,8 @@ func TestRequireRecordAuth(t *testing.T) {
} }
func TestRequireSameContextRecordAuth(t *testing.T) { func TestRequireSameContextRecordAuth(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "guest", Name: "guest",
@@ -358,6 +364,8 @@ func TestRequireSameContextRecordAuth(t *testing.T) {
} }
func TestRequireAdminAuth(t *testing.T) { func TestRequireAdminAuth(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "guest", Name: "guest",
@@ -452,6 +460,8 @@ func TestRequireAdminAuth(t *testing.T) {
} }
func TestRequireAdminAuthOnlyIfAny(t *testing.T) { func TestRequireAdminAuthOnlyIfAny(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "guest (while having at least 1 existing admin)", Name: "guest (while having at least 1 existing admin)",
@@ -571,6 +581,8 @@ func TestRequireAdminAuthOnlyIfAny(t *testing.T) {
} }
func TestRequireAdminOrRecordAuth(t *testing.T) { func TestRequireAdminOrRecordAuth(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "guest", Name: "guest",
@@ -731,6 +743,8 @@ func TestRequireAdminOrRecordAuth(t *testing.T) {
} }
func TestRequireAdminOrOwnerAuth(t *testing.T) { func TestRequireAdminOrOwnerAuth(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "guest", Name: "guest",
@@ -869,6 +883,8 @@ func TestRequireAdminOrOwnerAuth(t *testing.T) {
} }
func TestLoadCollectionContext(t *testing.T) { func TestLoadCollectionContext(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "missing collection", Name: "missing collection",
+276 -153
View File
@@ -4,8 +4,7 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "log/slog"
"log"
"net/http" "net/http"
"strings" "strings"
"time" "time"
@@ -16,18 +15,20 @@ import (
"github.com/pocketbase/pocketbase/forms" "github.com/pocketbase/pocketbase/forms"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/resolvers" "github.com/pocketbase/pocketbase/resolvers"
"github.com/pocketbase/pocketbase/tools/rest"
"github.com/pocketbase/pocketbase/tools/routine" "github.com/pocketbase/pocketbase/tools/routine"
"github.com/pocketbase/pocketbase/tools/search" "github.com/pocketbase/pocketbase/tools/search"
"github.com/pocketbase/pocketbase/tools/subscriptions" "github.com/pocketbase/pocketbase/tools/subscriptions"
"github.com/spf13/cast"
) )
// bindRealtimeApi registers the realtime api endpoints. // bindRealtimeApi registers the realtime api endpoints.
func bindRealtimeApi(app core.App, rg *echo.Group) { func bindRealtimeApi(app core.App, rg *echo.Group) {
api := realtimeApi{app: app} api := realtimeApi{app: app}
subGroup := rg.Group("/realtime", ActivityLogger(app)) subGroup := rg.Group("/realtime")
subGroup.GET("", api.connect) subGroup.GET("", api.connect)
subGroup.POST("", api.setSubscriptions) subGroup.POST("", api.setSubscriptions, ActivityLogger(app))
api.bindEvents() api.bindEvents()
} }
@@ -50,16 +51,19 @@ func (api *realtimeApi) connect(c echo.Context) error {
Client: client, Client: client,
} }
if err := api.app.OnRealtimeDisconnectRequest().Trigger(disconnectEvent); err != nil && api.app.IsDebug() { if err := api.app.OnRealtimeDisconnectRequest().Trigger(disconnectEvent); err != nil {
log.Println(err) api.app.Logger().Debug(
"OnRealtimeDisconnectRequest error",
slog.String("clientId", client.Id()),
slog.String("error", err.Error()),
)
} }
api.app.SubscriptionsBroker().Unregister(client.Id()) api.app.SubscriptionsBroker().Unregister(client.Id())
}() }()
c.Response().Header().Set("Content-Type", "text/event-stream; charset=UTF-8") c.Response().Header().Set("Content-Type", "text/event-stream")
c.Response().Header().Set("Cache-Control", "no-store") c.Response().Header().Set("Cache-Control", "no-store")
c.Response().Header().Set("Connection", "keep-alive")
// https://github.com/pocketbase/pocketbase/discussions/480#discussioncomment-3657640 // https://github.com/pocketbase/pocketbase/discussions/480#discussioncomment-3657640
// https://nginx.org/en/docs/http/ngx_http_proxy_module.html#proxy_buffering // https://nginx.org/en/docs/http/ngx_http_proxy_module.html#proxy_buffering
c.Response().Header().Set("X-Accel-Buffering", "no") c.Response().Header().Set("X-Accel-Buffering", "no")
@@ -67,15 +71,14 @@ func (api *realtimeApi) connect(c echo.Context) error {
connectEvent := &core.RealtimeConnectEvent{ connectEvent := &core.RealtimeConnectEvent{
HttpContext: c, HttpContext: c,
Client: client, Client: client,
IdleTimeout: 5 * time.Minute,
} }
if err := api.app.OnRealtimeConnectRequest().Trigger(connectEvent); err != nil { if err := api.app.OnRealtimeConnectRequest().Trigger(connectEvent); err != nil {
return err return err
} }
if api.app.IsDebug() { api.app.Logger().Debug("Realtime connection established.", slog.String("clientId", client.Id()))
log.Printf("Realtime connection established: %s\n", client.Id())
}
// signalize established connection (aka. fire "connect" message) // signalize established connection (aka. fire "connect" message)
connectMsgEvent := &core.RealtimeMessageEvent{ connectMsgEvent := &core.RealtimeMessageEvent{
@@ -83,30 +86,31 @@ func (api *realtimeApi) connect(c echo.Context) error {
Client: client, Client: client,
Message: &subscriptions.Message{ Message: &subscriptions.Message{
Name: "PB_CONNECT", Name: "PB_CONNECT",
Data: `{"clientId":"` + client.Id() + `"}`, Data: []byte(`{"clientId":"` + client.Id() + `"}`),
}, },
} }
connectMsgErr := api.app.OnRealtimeBeforeMessageSend().Trigger(connectMsgEvent, func(e *core.RealtimeMessageEvent) error { connectMsgErr := api.app.OnRealtimeBeforeMessageSend().Trigger(connectMsgEvent, func(e *core.RealtimeMessageEvent) error {
w := e.HttpContext.Response() w := e.HttpContext.Response()
fmt.Fprint(w, "id:"+client.Id()+"\n") w.Write([]byte("id:" + client.Id() + "\n"))
fmt.Fprint(w, "event:"+e.Message.Name+"\n") w.Write([]byte("event:" + e.Message.Name + "\n"))
fmt.Fprint(w, "data:"+e.Message.Data+"\n\n") w.Write([]byte("data:"))
w.Write(e.Message.Data)
w.Write([]byte("\n\n"))
w.Flush() w.Flush()
return nil return api.app.OnRealtimeAfterMessageSend().Trigger(e)
}) })
if connectMsgErr != nil { if connectMsgErr != nil {
if api.app.IsDebug() { api.app.Logger().Debug(
log.Println("Realtime connection closed (failed to deliver PB_CONNECT):", client.Id(), connectMsgErr) "Realtime connection closed (failed to deliver PB_CONNECT)",
} slog.String("clientId", client.Id()),
slog.String("error", connectMsgErr.Error()),
)
return nil return nil
} }
if err := api.app.OnRealtimeAfterMessageSend().Trigger(connectMsgEvent); err != nil && api.app.IsDebug() {
log.Println("OnRealtimeAfterMessageSend PB_CONNECT error:", err)
}
// start an idle timer to keep track of inactive/forgotten connections // start an idle timer to keep track of inactive/forgotten connections
idleDuration := 5 * time.Minute idleTimeout := connectEvent.IdleTimeout
idleTimer := time.NewTimer(idleDuration) idleTimer := time.NewTimer(idleTimeout)
defer idleTimer.Stop() defer idleTimer.Stop()
for { for {
@@ -116,9 +120,10 @@ func (api *realtimeApi) connect(c echo.Context) error {
case msg, ok := <-client.Channel(): case msg, ok := <-client.Channel():
if !ok { if !ok {
// channel is closed // channel is closed
if api.app.IsDebug() { api.app.Logger().Debug(
log.Println("Realtime connection closed (closed channel):", client.Id()) "Realtime connection closed (closed channel)",
} slog.String("clientId", client.Id()),
)
return nil return nil
} }
@@ -129,30 +134,31 @@ func (api *realtimeApi) connect(c echo.Context) error {
} }
msgErr := api.app.OnRealtimeBeforeMessageSend().Trigger(msgEvent, func(e *core.RealtimeMessageEvent) error { msgErr := api.app.OnRealtimeBeforeMessageSend().Trigger(msgEvent, func(e *core.RealtimeMessageEvent) error {
w := e.HttpContext.Response() w := e.HttpContext.Response()
fmt.Fprint(w, "id:"+e.Client.Id()+"\n") w.Write([]byte("id:" + e.Client.Id() + "\n"))
fmt.Fprint(w, "event:"+e.Message.Name+"\n") w.Write([]byte("event:" + e.Message.Name + "\n"))
fmt.Fprint(w, "data:"+e.Message.Data+"\n\n") w.Write([]byte("data:"))
w.Write(e.Message.Data)
w.Write([]byte("\n\n"))
w.Flush() w.Flush()
return nil return api.app.OnRealtimeAfterMessageSend().Trigger(msgEvent)
}) })
if msgErr != nil { if msgErr != nil {
if api.app.IsDebug() { api.app.Logger().Debug(
log.Println("Realtime connection closed (failed to deliver message):", client.Id(), msgErr) "Realtime connection closed (failed to deliver message)",
} slog.String("clientId", client.Id()),
slog.String("error", msgErr.Error()),
)
return nil return nil
} }
if err := api.app.OnRealtimeAfterMessageSend().Trigger(msgEvent); err != nil && api.app.IsDebug() {
log.Println("OnRealtimeAfterMessageSend error:", err)
}
idleTimer.Stop() idleTimer.Stop()
idleTimer.Reset(idleDuration) idleTimer.Reset(idleTimeout)
case <-c.Request().Context().Done(): case <-c.Request().Context().Done():
// connection is closed // connection is closed
if api.app.IsDebug() { api.app.Logger().Debug(
log.Println("Realtime connection closed (cancelled request):", client.Id()) "Realtime connection closed (cancelled request)",
} slog.String("clientId", client.Id()),
)
return nil return nil
} }
} }
@@ -191,7 +197,7 @@ func (api *realtimeApi) setSubscriptions(c echo.Context) error {
Subscriptions: form.Subscriptions, Subscriptions: form.Subscriptions,
} }
handlerErr := api.app.OnRealtimeBeforeSubscribeRequest().Trigger(event, func(e *core.RealtimeSubscribeEvent) error { return api.app.OnRealtimeBeforeSubscribeRequest().Trigger(event, func(e *core.RealtimeSubscribeEvent) error {
// update auth state // update auth state
e.Client.Set(ContextAdminKey, e.HttpContext.Get(ContextAdminKey)) e.Client.Set(ContextAdminKey, e.HttpContext.Get(ContextAdminKey))
e.Client.Set(ContextAuthRecordKey, e.HttpContext.Get(ContextAuthRecordKey)) e.Client.Set(ContextAuthRecordKey, e.HttpContext.Get(ContextAuthRecordKey))
@@ -202,14 +208,20 @@ func (api *realtimeApi) setSubscriptions(c echo.Context) error {
// subscribe to the new subscriptions // subscribe to the new subscriptions
e.Client.Subscribe(e.Subscriptions...) e.Client.Subscribe(e.Subscriptions...)
return e.HttpContext.NoContent(http.StatusNoContent) api.app.Logger().Debug(
"Realtime subscriptions updated.",
slog.String("clientId", e.Client.Id()),
slog.Any("subscriptions", e.Subscriptions),
)
return api.app.OnRealtimeAfterSubscribeRequest().Trigger(event, func(e *core.RealtimeSubscribeEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
if handlerErr == nil {
api.app.OnRealtimeAfterSubscribeRequest().Trigger(event)
}
return handlerErr
} }
// updateClientsAuthModel updates the existing clients auth model with the new one (matched by ID). // updateClientsAuthModel updates the existing clients auth model with the new one (matched by ID).
@@ -269,8 +281,13 @@ func (api *realtimeApi) bindEvents() {
api.app.OnModelAfterCreate().PreAdd(func(e *core.ModelEvent) error { api.app.OnModelAfterCreate().PreAdd(func(e *core.ModelEvent) error {
if record := api.resolveRecord(e.Model); record != nil { if record := api.resolveRecord(e.Model); record != nil {
if err := api.broadcastRecord("create", record); err != nil && api.app.IsDebug() { if err := api.broadcastRecord("create", record, false); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to broadcast record create",
slog.String("id", record.Id),
slog.String("collectionName", record.Collection().Name),
slog.String("error", err.Error()),
)
} }
} }
return nil return nil
@@ -278,8 +295,13 @@ func (api *realtimeApi) bindEvents() {
api.app.OnModelAfterUpdate().PreAdd(func(e *core.ModelEvent) error { api.app.OnModelAfterUpdate().PreAdd(func(e *core.ModelEvent) error {
if record := api.resolveRecord(e.Model); record != nil { if record := api.resolveRecord(e.Model); record != nil {
if err := api.broadcastRecord("update", record); err != nil && api.app.IsDebug() { if err := api.broadcastRecord("update", record, false); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to broadcast record update",
slog.String("id", record.Id),
slog.String("collectionName", record.Collection().Name),
slog.String("error", err.Error()),
)
} }
} }
return nil return nil
@@ -287,8 +309,27 @@ func (api *realtimeApi) bindEvents() {
api.app.OnModelBeforeDelete().Add(func(e *core.ModelEvent) error { api.app.OnModelBeforeDelete().Add(func(e *core.ModelEvent) error {
if record := api.resolveRecord(e.Model); record != nil { if record := api.resolveRecord(e.Model); record != nil {
if err := api.broadcastRecord("delete", record); err != nil && api.app.IsDebug() { if err := api.broadcastRecord("delete", record, true); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to dry cache record delete",
slog.String("id", record.Id),
slog.String("collectionName", record.Collection().Name),
slog.String("error", err.Error()),
)
}
}
return nil
})
api.app.OnModelAfterDelete().Add(func(e *core.ModelEvent) error {
if record := api.resolveRecord(e.Model); record != nil {
if err := api.broadcastDryCachedRecord("delete", record); err != nil {
api.app.Logger().Debug(
"Failed to broadcast record delete",
slog.String("id", record.Id),
slog.String("collectionName", record.Collection().Name),
slog.String("error", err.Error()),
)
} }
} }
return nil return nil
@@ -321,58 +362,16 @@ func (api *realtimeApi) resolveRecordCollection(model models.Model) (collection
return collection return collection
} }
// canAccessRecord checks if the subscription client has access to the specified record model. // recordData represents the broadcasted record subscrition message data.
func (api *realtimeApi) canAccessRecord(client subscriptions.Client, record *models.Record, accessRule *string) bool {
admin, _ := client.Get(ContextAdminKey).(*models.Admin)
if admin != nil {
// admins can access everything
return true
}
if accessRule == nil {
// only admins can access this record
return false
}
ruleFunc := func(q *dbx.SelectQuery) error {
if *accessRule == "" {
return nil // empty public rule
}
// mock request data
requestData := &models.RequestData{
Method: "GET",
}
requestData.AuthRecord, _ = client.Get(ContextAuthRecordKey).(*models.Record)
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), record.Collection(), requestData, true)
expr, err := search.FilterData(*accessRule).BuildExpr(resolver)
if err != nil {
return err
}
resolver.UpdateQuery(q)
q.AndWhere(expr)
return nil
}
foundRecord, err := api.app.Dao().FindRecordById(record.Collection().Id, record.Id, ruleFunc)
if err == nil && foundRecord != nil {
return true
}
return false
}
type recordData struct { type recordData struct {
Action string `json:"action"` Record any `json:"record"` /* map or models.Record */
Record *models.Record `json:"record"` Action string `json:"action"`
} }
func (api *realtimeApi) broadcastRecord(action string, record *models.Record) error { func (api *realtimeApi) broadcastRecord(action string, record *models.Record, dryCache bool) error {
collection := record.Collection() collection := record.Collection()
if collection == nil { if collection == nil {
return errors.New("Record collection not set.") return errors.New("[broadcastRecord] Record collection not set.")
} }
clients := api.app.SubscriptionsBroker().Clients() clients := api.app.SubscriptionsBroker().Clients()
@@ -380,77 +379,159 @@ func (api *realtimeApi) broadcastRecord(action string, record *models.Record) er
return nil // no subscribers return nil // no subscribers
} }
// create a clean record copy without expand and unknown fields
// because we don't know if the clients have permissions to view them
cleanRecord := record.CleanCopy()
subscriptionRuleMap := map[string]*string{ subscriptionRuleMap := map[string]*string{
(collection.Name + "/" + cleanRecord.Id): collection.ViewRule, (collection.Name + "/" + record.Id + "?"): collection.ViewRule,
(collection.Id + "/" + cleanRecord.Id): collection.ViewRule, (collection.Id + "/" + record.Id + "?"): collection.ViewRule,
(collection.Name + "/*"): collection.ListRule, (collection.Name + "/*?"): collection.ListRule,
(collection.Id + "/*"): collection.ListRule, (collection.Id + "/*?"): collection.ListRule,
// @deprecated: the same as the wildcard topic but kept for backward compatibility // @deprecated: the same as the wildcard topic but kept for backward compatibility
collection.Name: collection.ListRule, (collection.Name + "?"): collection.ListRule,
collection.Id: collection.ListRule, (collection.Id + "?"): collection.ListRule,
} }
data := &recordData{ dryCacheKey := action + "/" + record.Id
Action: action,
Record: cleanRecord,
}
dataBytes, err := json.Marshal(data)
if err != nil {
if api.app.IsDebug() {
log.Println(err)
}
return err
}
encodedData := string(dataBytes)
for _, client := range clients { for _, client := range clients {
client := client client := client
for subscription, rule := range subscriptionRuleMap { // note: not executed concurrently to avoid races and to ensure
if !client.HasSubscription(subscription) { // that the access checks are applied for the current record db state
for prefix, rule := range subscriptionRuleMap {
subs := client.Subscriptions(prefix)
if len(subs) == 0 {
continue continue
} }
if !api.canAccessRecord(client, data.Record, rule) { for sub, options := range subs {
continue // create a clean record copy without expand and unknown fields
} // because we don't know yet which exact fields the client subscription has permissions to access
cleanRecord := record.CleanCopy()
msg := subscriptions.Message{ // mock request data
Name: subscription, requestInfo := &models.RequestInfo{
Data: encodedData, Context: models.RequestInfoContextRealtime,
} Method: "GET",
Query: options.Query,
Headers: options.Headers,
}
requestInfo.Admin, _ = client.Get(ContextAdminKey).(*models.Admin)
requestInfo.AuthRecord, _ = client.Get(ContextAuthRecordKey).(*models.Record)
// ignore the auth record email visibility checks for if !api.canAccessRecord(cleanRecord, requestInfo, rule) {
// auth owner, admin or manager continue
if collection.IsAuth() { }
authId := extractAuthIdFromGetter(client)
if authId == data.Record.Id || rawExpand := cast.ToString(options.Query[expandQueryParam])
api.canAccessRecord(client, data.Record, collection.AuthOptions().ManageRule) { if rawExpand != "" {
data.Record.IgnoreEmailVisibility(true) // ignore expandErrs := api.app.Dao().ExpandRecord(cleanRecord, strings.Split(rawExpand, ","), expandFetch(api.app.Dao(), requestInfo))
if newData, err := json.Marshal(data); err == nil { if len(expandErrs) > 0 {
msg.Data = string(newData) api.app.Logger().Debug(
"[broadcastRecord] expand errors",
slog.String("id", cleanRecord.Id),
slog.String("collectionName", cleanRecord.Collection().Name),
slog.String("sub", sub),
slog.String("expand", rawExpand),
slog.Any("errors", expandErrs),
)
} }
data.Record.IgnoreEmailVisibility(false) // restore }
// ignore the auth record email visibility checks
// for auth owner, admin or manager
if collection.IsAuth() {
authId := extractAuthIdFromGetter(client)
if authId == cleanRecord.Id {
if api.canAccessRecord(cleanRecord, requestInfo, collection.AuthOptions().ManageRule) {
cleanRecord.IgnoreEmailVisibility(true)
}
}
}
data := &recordData{
Action: action,
Record: cleanRecord,
}
// check fields
rawFields := cast.ToString(options.Query[fieldsQueryParam])
if rawFields != "" {
decoded, err := rest.PickFields(cleanRecord, rawFields)
if err == nil {
data.Record = decoded
} else {
api.app.Logger().Debug(
"[broadcastRecord] pick fields error",
slog.String("id", cleanRecord.Id),
slog.String("collectionName", cleanRecord.Collection().Name),
slog.String("sub", sub),
slog.String("fields", rawFields),
slog.String("error", err.Error()),
)
}
}
dataBytes, err := json.Marshal(data)
if err != nil {
api.app.Logger().Debug(
"[broadcastRecord] data marshal error",
slog.String("id", cleanRecord.Id),
slog.String("collectionName", cleanRecord.Collection().Name),
slog.String("error", err.Error()),
)
continue
}
msg := subscriptions.Message{
Name: sub,
Data: dataBytes,
}
if dryCache {
messages, ok := client.Get(dryCacheKey).([]subscriptions.Message)
if !ok {
messages = []subscriptions.Message{msg}
} else {
messages = append(messages, msg)
}
client.Set(dryCacheKey, messages)
} else {
routine.FireAndForget(func() {
client.Send(msg)
})
} }
} }
routine.FireAndForget(func() {
if !client.IsDiscarded() {
client.Channel() <- msg
}
})
} }
} }
return nil return nil
} }
// broadcastDryCachedRecord broadcasts all cached record related messages.
func (api *realtimeApi) broadcastDryCachedRecord(action string, record *models.Record) error {
key := action + "/" + record.Id
clients := api.app.SubscriptionsBroker().Clients()
for _, client := range clients {
messages, ok := client.Get(key).([]subscriptions.Message)
if !ok {
continue
}
client.Unset(key)
client := client
routine.FireAndForget(func() {
for _, msg := range messages {
client.Send(msg)
}
})
}
return nil
}
type getter interface { type getter interface {
Get(string) any Get(string) any
} }
@@ -468,3 +549,45 @@ func extractAuthIdFromGetter(val getter) string {
return "" return ""
} }
// canAccessRecord checks if the subscription client has access to the specified record model.
func (api *realtimeApi) canAccessRecord(
record *models.Record,
requestInfo *models.RequestInfo,
accessRule *string,
) bool {
// check the access rule
// ---
if ok, _ := api.app.Dao().CanAccessRecord(record, requestInfo, accessRule); !ok {
return false
}
// check the subscription client-side filter (if any)
// ---
filter := cast.ToString(requestInfo.Query[search.FilterQueryParam])
if filter == "" {
return true // no further checks needed
}
if err := checkForAdminOnlyRuleFields(requestInfo); err != nil {
return false
}
ruleFunc := func(q *dbx.SelectQuery) error {
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), record.Collection(), requestInfo, false)
expr, err := search.FilterData(filter).BuildExpr(resolver)
if err != nil {
return err
}
q.AndWhere(expr)
resolver.UpdateQuery(q)
return nil
}
_, err := api.app.Dao().FindRecordById(record.Collection().Id, record.Id, ruleFunc)
return err == nil
}
+12 -9
View File
@@ -5,6 +5,7 @@ import (
"net/http" "net/http"
"strings" "strings"
"testing" "testing"
"time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
@@ -22,6 +23,7 @@ func TestRealtimeConnect(t *testing.T) {
{ {
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/realtime", Url: "/api/realtime",
Timeout: 100 * time.Millisecond,
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedContent: []string{ ExpectedContent: []string{
`id:`, `id:`,
@@ -34,7 +36,7 @@ func TestRealtimeConnect(t *testing.T) {
"OnRealtimeAfterMessageSend": 1, "OnRealtimeAfterMessageSend": 1,
"OnRealtimeDisconnectRequest": 1, "OnRealtimeDisconnectRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if len(app.SubscriptionsBroker().Clients()) != 0 { if len(app.SubscriptionsBroker().Clients()) != 0 {
t.Errorf("Expected the subscribers to be removed after connection close, found %d", len(app.SubscriptionsBroker().Clients())) t.Errorf("Expected the subscribers to be removed after connection close, found %d", len(app.SubscriptionsBroker().Clients()))
} }
@@ -44,6 +46,7 @@ func TestRealtimeConnect(t *testing.T) {
Name: "PB_CONNECT interrupt", Name: "PB_CONNECT interrupt",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/realtime", Url: "/api/realtime",
Timeout: 100 * time.Millisecond,
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnRealtimeConnectRequest": 1, "OnRealtimeConnectRequest": 1,
@@ -58,7 +61,7 @@ func TestRealtimeConnect(t *testing.T) {
return nil return nil
}) })
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if len(app.SubscriptionsBroker().Clients()) != 0 { if len(app.SubscriptionsBroker().Clients()) != 0 {
t.Errorf("Expected the subscribers to be removed after connection close, found %d", len(app.SubscriptionsBroker().Clients())) t.Errorf("Expected the subscribers to be removed after connection close, found %d", len(app.SubscriptionsBroker().Clients()))
} }
@@ -68,11 +71,11 @@ func TestRealtimeConnect(t *testing.T) {
Name: "Skipping/ignoring messages", Name: "Skipping/ignoring messages",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/realtime", Url: "/api/realtime",
Timeout: 100 * time.Millisecond,
ExpectedStatus: 200, ExpectedStatus: 200,
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnRealtimeConnectRequest": 1, "OnRealtimeConnectRequest": 1,
"OnRealtimeBeforeMessageSend": 1, "OnRealtimeBeforeMessageSend": 1,
"OnRealtimeAfterMessageSend": 1,
"OnRealtimeDisconnectRequest": 1, "OnRealtimeDisconnectRequest": 1,
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
@@ -80,7 +83,7 @@ func TestRealtimeConnect(t *testing.T) {
return hook.StopPropagation return hook.StopPropagation
}) })
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if len(app.SubscriptionsBroker().Clients()) != 0 { if len(app.SubscriptionsBroker().Clients()) != 0 {
t.Errorf("Expected the subscribers to be removed after connection close, found %d", len(app.SubscriptionsBroker().Clients())) t.Errorf("Expected the subscribers to be removed after connection close, found %d", len(app.SubscriptionsBroker().Clients()))
} }
@@ -125,7 +128,7 @@ func TestRealtimeSubscribe(t *testing.T) {
client.Subscribe("test0") client.Subscribe("test0")
app.SubscriptionsBroker().Register(client) app.SubscriptionsBroker().Register(client)
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if len(client.Subscriptions()) != 0 { if len(client.Subscriptions()) != 0 {
t.Errorf("Expected no subscriptions, got %v", client.Subscriptions()) t.Errorf("Expected no subscriptions, got %v", client.Subscriptions())
} }
@@ -146,7 +149,7 @@ func TestRealtimeSubscribe(t *testing.T) {
client.Subscribe("test0") client.Subscribe("test0")
app.SubscriptionsBroker().Register(client) app.SubscriptionsBroker().Register(client)
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
expectedSubs := []string{"test1", "test2"} expectedSubs := []string{"test1", "test2"}
if len(expectedSubs) != len(client.Subscriptions()) { if len(expectedSubs) != len(client.Subscriptions()) {
t.Errorf("Expected subscriptions %v, got %v", expectedSubs, client.Subscriptions()) t.Errorf("Expected subscriptions %v, got %v", expectedSubs, client.Subscriptions())
@@ -176,7 +179,7 @@ func TestRealtimeSubscribe(t *testing.T) {
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.SubscriptionsBroker().Register(client) app.SubscriptionsBroker().Register(client)
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
admin, _ := client.Get(apis.ContextAdminKey).(*models.Admin) admin, _ := client.Get(apis.ContextAdminKey).(*models.Admin)
if admin == nil { if admin == nil {
t.Errorf("Expected admin auth model, got nil") t.Errorf("Expected admin auth model, got nil")
@@ -200,7 +203,7 @@ func TestRealtimeSubscribe(t *testing.T) {
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.SubscriptionsBroker().Register(client) app.SubscriptionsBroker().Register(client)
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
authRecord, _ := client.Get(apis.ContextAuthRecordKey).(*models.Record) authRecord, _ := client.Get(apis.ContextAuthRecordKey).(*models.Record)
if authRecord == nil { if authRecord == nil {
t.Errorf("Expected auth record model, got nil") t.Errorf("Expected auth record model, got nil")
@@ -225,7 +228,7 @@ func TestRealtimeSubscribe(t *testing.T) {
app.SubscriptionsBroker().Register(client) app.SubscriptionsBroker().Register(client)
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
authRecord, _ := client.Get(apis.ContextAuthRecordKey).(*models.Record) authRecord, _ := client.Get(apis.ContextAuthRecordKey).(*models.Record)
if authRecord == nil { if authRecord == nil {
t.Errorf("Expected auth record model, got nil") t.Errorf("Expected auth record model, got nil")
+165 -134
View File
@@ -4,8 +4,9 @@ import (
"encoding/json" "encoding/json"
"errors" "errors"
"fmt" "fmt"
"log" "log/slog"
"net/http" "net/http"
"sort"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
@@ -65,26 +66,23 @@ func (api *recordAuthApi) authRefresh(c echo.Context) error {
event.Collection = record.Collection() event.Collection = record.Collection()
event.Record = record event.Record = record
handlerErr := api.app.OnRecordBeforeAuthRefreshRequest().Trigger(event, func(e *core.RecordAuthRefreshEvent) error { return api.app.OnRecordBeforeAuthRefreshRequest().Trigger(event, func(e *core.RecordAuthRefreshEvent) error {
return RecordAuthResponse(api.app, e.HttpContext, e.Record, nil) return api.app.OnRecordAfterAuthRefreshRequest().Trigger(event, func(e *core.RecordAuthRefreshEvent) error {
return RecordAuthResponse(api.app, e.HttpContext, e.Record, nil)
})
}) })
if handlerErr == nil {
if err := api.app.OnRecordAfterAuthRefreshRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return handlerErr
} }
type providerInfo struct { type providerInfo struct {
Name string `json:"name"` Name string `json:"name"`
State string `json:"state"` DisplayName string `json:"displayName"`
State string `json:"state"`
AuthUrl string `json:"authUrl"`
// technically could be omitted if the provider doesn't support PKCE,
// but to avoid breaking existing typed clients we'll return them as empty string
CodeVerifier string `json:"codeVerifier"` CodeVerifier string `json:"codeVerifier"`
CodeChallenge string `json:"codeChallenge"` CodeChallenge string `json:"codeChallenge"`
CodeChallengeMethod string `json:"codeChallengeMethod"` CodeChallengeMethod string `json:"codeChallengeMethod"`
AuthUrl string `json:"authUrl"`
} }
func (api *recordAuthApi) authMethods(c echo.Context) error { func (api *recordAuthApi) authMethods(c echo.Context) error {
@@ -96,12 +94,14 @@ func (api *recordAuthApi) authMethods(c echo.Context) error {
authOptions := collection.AuthOptions() authOptions := collection.AuthOptions()
result := struct { result := struct {
AuthProviders []providerInfo `json:"authProviders"`
UsernamePassword bool `json:"usernamePassword"` UsernamePassword bool `json:"usernamePassword"`
EmailPassword bool `json:"emailPassword"` EmailPassword bool `json:"emailPassword"`
AuthProviders []providerInfo `json:"authProviders"` OnlyVerified bool `json:"onlyVerified"`
}{ }{
UsernamePassword: authOptions.AllowUsernameAuth, UsernamePassword: authOptions.AllowUsernameAuth,
EmailPassword: authOptions.AllowEmailAuth, EmailPassword: authOptions.AllowEmailAuth,
OnlyVerified: authOptions.OnlyVerified,
AuthProviders: []providerInfo{}, AuthProviders: []providerInfo{},
} }
@@ -117,46 +117,61 @@ func (api *recordAuthApi) authMethods(c echo.Context) error {
provider, err := auth.NewProviderByName(name) provider, err := auth.NewProviderByName(name)
if err != nil { if err != nil {
if api.app.IsDebug() { api.app.Logger().Debug("Missing or invalid provider name", slog.String("name", name))
log.Println(err)
}
continue // skip provider continue // skip provider
} }
if err := config.SetupProvider(provider); err != nil { if err := config.SetupProvider(provider); err != nil {
if api.app.IsDebug() { api.app.Logger().Debug(
log.Println(err) "Failed to setup provider",
} slog.String("name", name),
slog.String("error", err.Error()),
)
continue // skip provider continue // skip provider
} }
state := security.RandomString(30) info := providerInfo{
codeVerifier := security.RandomString(43) Name: name,
codeChallenge := security.S256Challenge(codeVerifier) DisplayName: provider.DisplayName(),
codeChallengeMethod := "S256" State: security.RandomString(30),
urlOpts := []oauth2.AuthCodeOption{
oauth2.SetAuthURLParam("code_challenge", codeChallenge),
oauth2.SetAuthURLParam("code_challenge_method", codeChallengeMethod),
} }
if name == auth.NameApple { if info.DisplayName == "" {
info.DisplayName = name
}
urlOpts := []oauth2.AuthCodeOption{}
// custom providers url options
switch name {
case auth.NameApple:
// see https://developer.apple.com/documentation/sign_in_with_apple/sign_in_with_apple_js/incorporating_sign_in_with_apple_into_other_platforms#3332113 // see https://developer.apple.com/documentation/sign_in_with_apple/sign_in_with_apple_js/incorporating_sign_in_with_apple_into_other_platforms#3332113
urlOpts = append(urlOpts, oauth2.SetAuthURLParam("response_mode", "query")) urlOpts = append(urlOpts, oauth2.SetAuthURLParam("response_mode", "query"))
} }
result.AuthProviders = append(result.AuthProviders, providerInfo{ if provider.PKCE() {
Name: name, info.CodeVerifier = security.RandomString(43)
State: state, info.CodeChallenge = security.S256Challenge(info.CodeVerifier)
CodeVerifier: codeVerifier, info.CodeChallengeMethod = "S256"
CodeChallenge: codeChallenge, urlOpts = append(urlOpts,
CodeChallengeMethod: codeChallengeMethod, oauth2.SetAuthURLParam("code_challenge", info.CodeChallenge),
AuthUrl: provider.BuildAuthUrl( oauth2.SetAuthURLParam("code_challenge_method", info.CodeChallengeMethod),
state, )
urlOpts..., }
) + "&redirect_uri=", // empty redirect_uri so that users can append their redirect url
}) info.AuthUrl = provider.BuildAuthUrl(
info.State,
urlOpts...,
) + "&redirect_uri=" // empty redirect_uri so that users can append their redirect url
result.AuthProviders = append(result.AuthProviders, info)
} }
// sort providers
sort.SliceStable(result.AuthProviders, func(i, j int) bool {
return result.AuthProviders[i].Name < result.AuthProviders[j].Name
})
return c.JSON(http.StatusOK, result) return c.JSON(http.StatusOK, result)
} }
@@ -186,14 +201,15 @@ func (api *recordAuthApi) authWithOAuth2(c echo.Context) error {
event.HttpContext = c event.HttpContext = c
event.Collection = collection event.Collection = collection
event.ProviderName = form.Provider event.ProviderName = form.Provider
event.IsNewRecord = false
form.SetBeforeNewRecordCreateFunc(func(createForm *forms.RecordUpsert, authRecord *models.Record, authUser *auth.AuthUser) error { form.SetBeforeNewRecordCreateFunc(func(createForm *forms.RecordUpsert, authRecord *models.Record, authUser *auth.AuthUser) error {
return createForm.DrySubmit(func(txDao *daos.Dao) error { return createForm.DrySubmit(func(txDao *daos.Dao) error {
event.IsNewRecord = true event.IsNewRecord = true
// clone the current request data and assign the form create data as its body data // clone the current request data and assign the form create data as its body data
requestData := *RequestData(c) requestInfo := *RequestInfo(c)
requestData.Data = form.CreateData requestInfo.Context = models.RequestInfoContextOAuth2
requestInfo.Data = form.CreateData
createRuleFunc := func(q *dbx.SelectQuery) error { createRuleFunc := func(q *dbx.SelectQuery) error {
admin, _ := c.Get(ContextAdminKey).(*models.Admin) admin, _ := c.Get(ContextAdminKey).(*models.Admin)
@@ -206,7 +222,7 @@ func (api *recordAuthApi) authWithOAuth2(c echo.Context) error {
} }
if *collection.CreateRule != "" { if *collection.CreateRule != "" {
resolver := resolvers.NewRecordFieldResolver(txDao, collection, &requestData, true) resolver := resolvers.NewRecordFieldResolver(txDao, collection, &requestInfo, true)
expr, err := search.FilterData(*collection.CreateRule).BuildExpr(resolver) expr, err := search.FilterData(*collection.CreateRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -231,6 +247,7 @@ func (api *recordAuthApi) authWithOAuth2(c echo.Context) error {
event.Record = data.Record event.Record = data.Record
event.OAuth2User = data.OAuth2User event.OAuth2User = data.OAuth2User
event.ProviderClient = data.ProviderClient event.ProviderClient = data.ProviderClient
event.IsNewRecord = data.Record == nil
return api.app.OnRecordBeforeAuthWithOAuth2Request().Trigger(event, func(e *core.RecordAuthWithOAuth2Event) error { return api.app.OnRecordBeforeAuthWithOAuth2Request().Trigger(event, func(e *core.RecordAuthWithOAuth2Event) error {
data.Record = e.Record data.Record = e.Record
@@ -251,17 +268,13 @@ func (api *recordAuthApi) authWithOAuth2(c echo.Context) error {
IsNew: event.IsNewRecord, IsNew: event.IsNewRecord,
} }
return RecordAuthResponse(api.app, e.HttpContext, e.Record, meta) return api.app.OnRecordAfterAuthWithOAuth2Request().Trigger(event, func(e *core.RecordAuthWithOAuth2Event) error {
return RecordAuthResponse(api.app, e.HttpContext, e.Record, meta)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterAuthWithOAuth2Request().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -291,17 +304,13 @@ func (api *recordAuthApi) authWithPassword(c echo.Context) error {
return NewBadRequestError("Failed to authenticate.", err) return NewBadRequestError("Failed to authenticate.", err)
} }
return RecordAuthResponse(api.app, e.HttpContext, e.Record, nil) return api.app.OnRecordAfterAuthWithPasswordRequest().Trigger(event, func(e *core.RecordAuthWithPasswordEvent) error {
return RecordAuthResponse(api.app, e.HttpContext, e.Record, nil)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterAuthWithPasswordRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -336,30 +345,32 @@ func (api *recordAuthApi) requestPasswordReset(c echo.Context) error {
return api.app.OnRecordBeforeRequestPasswordResetRequest().Trigger(event, func(e *core.RecordRequestPasswordResetEvent) error { return api.app.OnRecordBeforeRequestPasswordResetRequest().Trigger(event, func(e *core.RecordRequestPasswordResetEvent) error {
// run in background because we don't need to show the result to the client // run in background because we don't need to show the result to the client
routine.FireAndForget(func() { routine.FireAndForget(func() {
if err := next(e.Record); err != nil && api.app.IsDebug() { if err := next(e.Record); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to send password reset email",
slog.String("error", err.Error()),
)
} }
}) })
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterRequestPasswordResetRequest().Trigger(event, func(e *core.RecordRequestPasswordResetEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil { // eagerly write 204 response and skip submit errors
if err := api.app.OnRecordAfterRequestPasswordResetRequest().Trigger(event); err != nil && api.app.IsDebug() { // as a measure against emails enumeration
log.Println(err)
}
} else if api.app.IsDebug() {
log.Println(submitErr)
}
// don't return the response error to prevent emails enumeration
if !c.Response().Committed { if !c.Response().Committed {
c.NoContent(http.StatusNoContent) c.NoContent(http.StatusNoContent)
} }
return nil return submitErr
} }
func (api *recordAuthApi) confirmPasswordReset(c echo.Context) error { func (api *recordAuthApi) confirmPasswordReset(c echo.Context) error {
@@ -386,17 +397,17 @@ func (api *recordAuthApi) confirmPasswordReset(c echo.Context) error {
return NewBadRequestError("Failed to set new password.", err) return NewBadRequestError("Failed to set new password.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterConfirmPasswordResetRequest().Trigger(event, func(e *core.RecordConfirmPasswordResetEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterConfirmPasswordResetRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -426,30 +437,32 @@ func (api *recordAuthApi) requestVerification(c echo.Context) error {
return api.app.OnRecordBeforeRequestVerificationRequest().Trigger(event, func(e *core.RecordRequestVerificationEvent) error { return api.app.OnRecordBeforeRequestVerificationRequest().Trigger(event, func(e *core.RecordRequestVerificationEvent) error {
// run in background because we don't need to show the result to the client // run in background because we don't need to show the result to the client
routine.FireAndForget(func() { routine.FireAndForget(func() {
if err := next(e.Record); err != nil && api.app.IsDebug() { if err := next(e.Record); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to send verification email",
slog.String("error", err.Error()),
)
} }
}) })
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterRequestVerificationRequest().Trigger(event, func(e *core.RecordRequestVerificationEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil { // eagerly write 204 response and skip submit errors
if err := api.app.OnRecordAfterRequestVerificationRequest().Trigger(event); err != nil && api.app.IsDebug() { // as a measure against users enumeration
log.Println(err)
}
} else if api.app.IsDebug() {
log.Println(submitErr)
}
// don't return the response error to prevent emails enumeration
if !c.Response().Committed { if !c.Response().Committed {
c.NoContent(http.StatusNoContent) c.NoContent(http.StatusNoContent)
} }
return nil return submitErr
} }
func (api *recordAuthApi) confirmVerification(c echo.Context) error { func (api *recordAuthApi) confirmVerification(c echo.Context) error {
@@ -476,17 +489,17 @@ func (api *recordAuthApi) confirmVerification(c echo.Context) error {
return NewBadRequestError("An error occurred while submitting the form.", err) return NewBadRequestError("An error occurred while submitting the form.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterConfirmVerificationRequest().Trigger(event, func(e *core.RecordConfirmVerificationEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterConfirmVerificationRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -511,23 +524,23 @@ func (api *recordAuthApi) requestEmailChange(c echo.Context) error {
event.Collection = collection event.Collection = collection
event.Record = record event.Record = record
submitErr := form.Submit(func(next forms.InterceptorNextFunc[*models.Record]) forms.InterceptorNextFunc[*models.Record] { return form.Submit(func(next forms.InterceptorNextFunc[*models.Record]) forms.InterceptorNextFunc[*models.Record] {
return func(record *models.Record) error { return func(record *models.Record) error {
return api.app.OnRecordBeforeRequestEmailChangeRequest().Trigger(event, func(e *core.RecordRequestEmailChangeEvent) error { return api.app.OnRecordBeforeRequestEmailChangeRequest().Trigger(event, func(e *core.RecordRequestEmailChangeEvent) error {
if err := next(e.Record); err != nil { if err := next(e.Record); err != nil {
return NewBadRequestError("Failed to request email change.", err) return NewBadRequestError("Failed to request email change.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterRequestEmailChangeRequest().Trigger(event, func(e *core.RecordRequestEmailChangeEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil {
api.app.OnRecordAfterRequestEmailChangeRequest().Trigger(event)
}
return submitErr
} }
func (api *recordAuthApi) confirmEmailChange(c echo.Context) error { func (api *recordAuthApi) confirmEmailChange(c echo.Context) error {
@@ -554,17 +567,17 @@ func (api *recordAuthApi) confirmEmailChange(c echo.Context) error {
return NewBadRequestError("Failed to confirm email change.", err) return NewBadRequestError("Failed to confirm email change.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterConfirmEmailChangeRequest().Trigger(event, func(e *core.RecordConfirmEmailChangeEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterConfirmEmailChangeRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr return submitErr
} }
@@ -628,54 +641,72 @@ func (api *recordAuthApi) unlinkExternalAuth(c echo.Context) error {
event.Record = record event.Record = record
event.ExternalAuth = externalAuth event.ExternalAuth = externalAuth
handlerErr := api.app.OnRecordBeforeUnlinkExternalAuthRequest().Trigger(event, func(e *core.RecordUnlinkExternalAuthEvent) error { return api.app.OnRecordBeforeUnlinkExternalAuthRequest().Trigger(event, func(e *core.RecordUnlinkExternalAuthEvent) error {
if err := api.app.Dao().DeleteExternalAuth(externalAuth); err != nil { if err := api.app.Dao().DeleteExternalAuth(externalAuth); err != nil {
return NewBadRequestError("Cannot unlink the external auth provider.", err) return NewBadRequestError("Cannot unlink the external auth provider.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterUnlinkExternalAuthRequest().Trigger(event, func(e *core.RecordUnlinkExternalAuthEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
if handlerErr == nil {
api.app.OnRecordAfterUnlinkExternalAuthRequest().Trigger(event)
}
return handlerErr
} }
// ------------------------------------------------------------------- // -------------------------------------------------------------------
const oauth2SubscriptionTopic = "@oauth2" const (
oauth2SubscriptionTopic string = "@oauth2"
oauth2RedirectFailurePath string = "../_/#/auth/oauth2-redirect-failure"
oauth2RedirectSuccessPath string = "../_/#/auth/oauth2-redirect-success"
)
type oauth2EventMessage struct {
State string `json:"state"`
Code string `json:"code"`
Error string `json:"error,omitempty"`
}
func (api *recordAuthApi) oauth2SubscriptionRedirect(c echo.Context) error { func (api *recordAuthApi) oauth2SubscriptionRedirect(c echo.Context) error {
state := c.QueryParam("state") state := c.QueryParam("state")
code := c.QueryParam("code") if state == "" {
api.app.Logger().Debug("Missing OAuth2 state parameter")
if code == "" || state == "" { return c.Redirect(http.StatusTemporaryRedirect, oauth2RedirectFailurePath)
return NewBadRequestError("Invalid OAuth2 redirect parameters.", nil)
} }
client, err := api.app.SubscriptionsBroker().ClientById(state) client, err := api.app.SubscriptionsBroker().ClientById(state)
if err != nil || client.IsDiscarded() || !client.HasSubscription(oauth2SubscriptionTopic) { if err != nil || client.IsDiscarded() || !client.HasSubscription(oauth2SubscriptionTopic) {
return NewNotFoundError("Missing or invalid OAuth2 subscription client.", err) api.app.Logger().Debug("Missing or invalid OAuth2 subscription client", "error", err, "clientId", state)
return c.Redirect(http.StatusTemporaryRedirect, oauth2RedirectFailurePath)
} }
defer client.Unsubscribe(oauth2SubscriptionTopic)
data := map[string]string{ data := oauth2EventMessage{
"state": state, State: state,
"code": code, Code: c.QueryParam("code"),
Error: c.QueryParam("error"),
} }
encodedData, err := json.Marshal(data) encodedData, err := json.Marshal(data)
if err != nil { if err != nil {
return NewBadRequestError("Failed to marshalize OAuth2 redirect data.", err) api.app.Logger().Debug("Failed to marshalize OAuth2 redirect data", "error", err)
return c.Redirect(http.StatusTemporaryRedirect, oauth2RedirectFailurePath)
} }
msg := subscriptions.Message{ msg := subscriptions.Message{
Name: oauth2SubscriptionTopic, Name: oauth2SubscriptionTopic,
Data: string(encodedData), Data: encodedData,
} }
client.Channel() <- msg client.Send(msg)
return c.Redirect(http.StatusTemporaryRedirect, "../_/#/auth/oauth2-redirect") if data.Error != "" || data.Code == "" {
api.app.Logger().Debug("Failed OAuth2 redirect due to an error or missing code parameter", "error", data.Error, "clientId", data.State)
return c.Redirect(http.StatusTemporaryRedirect, oauth2RedirectFailurePath)
}
return c.Redirect(http.StatusTemporaryRedirect, oauth2RedirectSuccessPath)
} }
+545 -99
View File
@@ -2,12 +2,14 @@ package apis_test
import ( import (
"context" "context"
"errors"
"net/http" "net/http"
"strings" "strings"
"testing" "testing"
"time" "time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/subscriptions" "github.com/pocketbase/pocketbase/tools/subscriptions"
@@ -15,6 +17,8 @@ import (
) )
func TestRecordAuthMethodsList(t *testing.T) { func TestRecordAuthMethodsList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "missing collection", Name: "missing collection",
@@ -38,6 +42,7 @@ func TestRecordAuthMethodsList(t *testing.T) {
ExpectedContent: []string{ ExpectedContent: []string{
`"usernamePassword":true`, `"usernamePassword":true`,
`"emailPassword":true`, `"emailPassword":true`,
`"onlyVerified":false`,
`"authProviders":[{`, `"authProviders":[{`,
`"name":"gitlab"`, `"name":"gitlab"`,
`"state":`, `"state":`,
@@ -56,6 +61,7 @@ func TestRecordAuthMethodsList(t *testing.T) {
ExpectedContent: []string{ ExpectedContent: []string{
`"usernamePassword":false`, `"usernamePassword":false`,
`"emailPassword":true`, `"emailPassword":true`,
`"onlyVerified":true`,
`"authProviders":[]`, `"authProviders":[]`,
}, },
}, },
@@ -67,6 +73,8 @@ func TestRecordAuthMethodsList(t *testing.T) {
} }
func TestRecordAuthWithPassword(t *testing.T) { func TestRecordAuthWithPassword(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "invalid body format", Name: "invalid body format",
@@ -210,7 +218,7 @@ func TestRecordAuthWithPassword(t *testing.T) {
}, },
}, },
{ {
Name: "valid email and valid password in allowed collection", Name: "valid email (unverified) and valid password in allowed collection",
Method: http.MethodPost, Method: http.MethodPost,
Url: "/api/collections/users/auth-with-password", Url: "/api/collections/users/auth-with-password",
Body: strings.NewReader(`{ Body: strings.NewReader(`{
@@ -223,6 +231,48 @@ func TestRecordAuthWithPassword(t *testing.T) {
`"token":"`, `"token":"`,
`"id":"4q1xlclmfloku33"`, `"id":"4q1xlclmfloku33"`,
`"email":"test@example.com"`, `"email":"test@example.com"`,
`"verified":false`,
},
ExpectedEvents: map[string]int{
"OnRecordBeforeAuthWithPasswordRequest": 1,
"OnRecordAfterAuthWithPasswordRequest": 1,
"OnRecordAuthRequest": 1,
},
},
// onlyVerified collection check
{
Name: "unverified user in onlyVerified collection",
Method: http.MethodPost,
Url: "/api/collections/clients/auth-with-password",
Body: strings.NewReader(`{
"identity":"test2@example.com",
"password":"1234567890"
}`),
ExpectedStatus: 403,
ExpectedContent: []string{
`"data":{}`,
},
ExpectedEvents: map[string]int{
"OnRecordBeforeAuthWithPasswordRequest": 1,
"OnRecordAfterAuthWithPasswordRequest": 1,
},
},
{
Name: "verified user in onlyVerified collection",
Method: http.MethodPost,
Url: "/api/collections/clients/auth-with-password",
Body: strings.NewReader(`{
"identity":"test@example.com",
"password":"1234567890"
}`),
ExpectedStatus: 200,
ExpectedContent: []string{
`"record":{`,
`"token":"`,
`"id":"gk390qegs4y47wn"`,
`"email":"test@example.com"`,
`"verified":true`,
}, },
ExpectedEvents: map[string]int{ ExpectedEvents: map[string]int{
"OnRecordBeforeAuthWithPasswordRequest": 1, "OnRecordBeforeAuthWithPasswordRequest": 1,
@@ -280,6 +330,28 @@ func TestRecordAuthWithPassword(t *testing.T) {
"OnRecordAuthRequest": 1, "OnRecordAuthRequest": 1,
}, },
}, },
// after hooks error checks
{
Name: "OnRecordAfterAuthWithPasswordRequest error response",
Method: http.MethodPost,
Url: "/api/collections/users/auth-with-password",
Body: strings.NewReader(`{
"identity":"test2_username",
"password":"1234567890"
}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterAuthWithPasswordRequest().Add(func(e *core.RecordAuthWithPasswordEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnRecordBeforeAuthWithPasswordRequest": 1,
"OnRecordAfterAuthWithPasswordRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -288,6 +360,8 @@ func TestRecordAuthWithPassword(t *testing.T) {
} }
func TestRecordAuthRefresh(t *testing.T) { func TestRecordAuthRefresh(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -353,6 +427,60 @@ func TestRecordAuthRefresh(t *testing.T) {
"OnRecordAfterAuthRefreshRequest": 1, "OnRecordAfterAuthRefreshRequest": 1,
}, },
}, },
{
Name: "unverified auth record in onlyVerified collection",
Method: http.MethodPost,
Url: "/api/collections/clients/auth-refresh",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6Im8xeTBkZDBzcGQ3ODZtZCIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoidjg1MXE0cjc5MHJoa25sIiwiZXhwIjoyMjA4OTg1MjYxfQ.-JYlrz5DcGzvb0nYx-xqnSFMu9dupyKY7Vg_FUm0OaM",
},
ExpectedStatus: 403,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnRecordBeforeAuthRefreshRequest": 1,
"OnRecordAfterAuthRefreshRequest": 1,
},
},
{
Name: "verified auth record in onlyVerified collection",
Method: http.MethodPost,
Url: "/api/collections/clients/auth-refresh",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6ImdrMzkwcWVnczR5NDd3biIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoidjg1MXE0cjc5MHJoa25sIiwiZXhwIjoyMjA4OTg1MjYxfQ.q34IWXrRWsjLvbbVNRfAs_J4SoTHloNBfdGEiLmy-D8",
},
ExpectedStatus: 200,
ExpectedContent: []string{
`"token":`,
`"record":`,
`"id":"gk390qegs4y47wn"`,
`"verified":true`,
`"email":"test@example.com"`,
},
ExpectedEvents: map[string]int{
"OnRecordBeforeAuthRefreshRequest": 1,
"OnRecordAuthRequest": 1,
"OnRecordAfterAuthRefreshRequest": 1,
},
},
{
Name: "OnRecordAfterAuthRefreshRequest error response",
Method: http.MethodPost,
Url: "/api/collections/users/auth-refresh?expand=rel,missing",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiY29sbGVjdGlvbklkIjoiX3BiX3VzZXJzX2F1dGhfIiwiZXhwIjoyMjA4OTg1MjYxfQ.UwD8JvkbQtXpymT09d7J6fdA0aP9g4FJ1GPh_ggEkzc",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterAuthRefreshRequest().Add(func(e *core.RecordAuthRefreshEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnRecordBeforeAuthRefreshRequest": 1,
"OnRecordAfterAuthRefreshRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -361,6 +489,8 @@ func TestRecordAuthRefresh(t *testing.T) {
} }
func TestRecordAuthRequestPasswordReset(t *testing.T) { func TestRecordAuthRequestPasswordReset(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "not an auth collection", Name: "not an auth collection",
@@ -446,6 +576,8 @@ func TestRecordAuthRequestPasswordReset(t *testing.T) {
} }
func TestRecordAuthConfirmPasswordReset(t *testing.T) { func TestRecordAuthConfirmPasswordReset(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "empty data", Name: "empty data",
@@ -494,10 +626,8 @@ func TestRecordAuthConfirmPasswordReset(t *testing.T) {
"password":"12345678", "password":"12345678",
"passwordConfirm":"12345678" "passwordConfirm":"12345678"
}`), }`),
ExpectedStatus: 400, ExpectedStatus: 400,
ExpectedContent: []string{ ExpectedContent: []string{`"data":{}`},
`"data":{}`,
},
}, },
{ {
Name: "different auth collection", Name: "different auth collection",
@@ -514,7 +644,7 @@ func TestRecordAuthConfirmPasswordReset(t *testing.T) {
}, },
}, },
{ {
Name: "valid token and data", Name: "valid token and data (unverified user)",
Method: http.MethodPost, Method: http.MethodPost,
Url: "/api/collections/users/confirm-password-reset", Url: "/api/collections/users/confirm-password-reset",
Body: strings.NewReader(`{ Body: strings.NewReader(`{
@@ -529,6 +659,155 @@ func TestRecordAuthConfirmPasswordReset(t *testing.T) {
"OnRecordBeforeConfirmPasswordResetRequest": 1, "OnRecordBeforeConfirmPasswordResetRequest": 1,
"OnRecordAfterConfirmPasswordResetRequest": 1, "OnRecordAfterConfirmPasswordResetRequest": 1,
}, },
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
user, err := app.Dao().FindAuthRecordByEmail("users", "test@example.com")
if err != nil {
t.Fatalf("Failed to fetch confirm password user: %v", err)
}
if user.Verified() {
t.Fatalf("Expected the user to be unverified")
}
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
user, err := app.Dao().FindAuthRecordByToken(
"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg",
app.Settings().RecordPasswordResetToken.Secret,
)
if err == nil {
t.Fatalf("Expected the password reset token to be invalidated")
}
user, err = app.Dao().FindAuthRecordByEmail("users", "test@example.com")
if err != nil {
t.Fatalf("Failed to fetch confirm password user: %v", err)
}
if !user.Verified() {
t.Fatalf("Expected the user to be marked as verified")
}
},
},
{
Name: "valid token and data (unverified user with different email from the one in the token)",
Method: http.MethodPost,
Url: "/api/collections/users/confirm-password-reset",
Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg",
"password":"12345678",
"passwordConfirm":"12345678"
}`),
ExpectedStatus: 204,
ExpectedEvents: map[string]int{
"OnModelAfterUpdate": 1,
"OnModelBeforeUpdate": 1,
"OnRecordBeforeConfirmPasswordResetRequest": 1,
"OnRecordAfterConfirmPasswordResetRequest": 1,
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
user, err := app.Dao().FindAuthRecordByEmail("users", "test@example.com")
if err != nil {
t.Fatalf("Failed to fetch confirm password user: %v", err)
}
if user.Verified() {
t.Fatalf("Expected the user to be unverified")
}
// manually change the email to check whether the verified state will be updated
user.SetEmail("test_update@example.com")
if err := app.Dao().WithoutHooks().SaveRecord(user); err != nil {
t.Fatalf("Failed to update user test email")
}
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
user, err := app.Dao().FindAuthRecordByToken(
"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg",
app.Settings().RecordPasswordResetToken.Secret,
)
if err == nil {
t.Fatalf("Expected the password reset token to be invalidated")
}
user, err = app.Dao().FindAuthRecordByEmail("users", "test_update@example.com")
if err != nil {
t.Fatalf("Failed to fetch confirm password user: %v", err)
}
if user.Verified() {
t.Fatalf("Expected the user to remain unverified")
}
},
},
{
Name: "valid token and data (verified user)",
Method: http.MethodPost,
Url: "/api/collections/users/confirm-password-reset",
Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg",
"password":"12345678",
"passwordConfirm":"12345678"
}`),
ExpectedStatus: 204,
ExpectedEvents: map[string]int{
"OnModelAfterUpdate": 1,
"OnModelBeforeUpdate": 1,
"OnRecordBeforeConfirmPasswordResetRequest": 1,
"OnRecordAfterConfirmPasswordResetRequest": 1,
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
user, err := app.Dao().FindAuthRecordByEmail("users", "test@example.com")
if err != nil {
t.Fatalf("Failed to fetch confirm password user: %v", err)
}
// ensure that the user is already verified
user.SetVerified(true)
if err := app.Dao().WithoutHooks().SaveRecord(user); err != nil {
t.Fatalf("Failed to update user verified state")
}
},
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
user, err := app.Dao().FindAuthRecordByToken(
"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg",
app.Settings().RecordPasswordResetToken.Secret,
)
if err == nil {
t.Fatalf("Expected the password reset token to be invalidated")
}
user, err = app.Dao().FindAuthRecordByEmail("users", "test@example.com")
if err != nil {
t.Fatalf("Failed to fetch confirm password user: %v", err)
}
if !user.Verified() {
t.Fatalf("Expected the user to remain verified")
}
},
},
{
Name: "OnRecordAfterConfirmPasswordResetRequest error response",
Method: http.MethodPost,
Url: "/api/collections/users/confirm-password-reset",
Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg",
"password":"12345678",
"passwordConfirm":"12345678"
}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterConfirmPasswordResetRequest().Add(func(e *core.RecordConfirmPasswordResetEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelAfterUpdate": 1,
"OnModelBeforeUpdate": 1,
"OnRecordBeforeConfirmPasswordResetRequest": 1,
"OnRecordAfterConfirmPasswordResetRequest": 1,
},
}, },
} }
@@ -538,6 +817,8 @@ func TestRecordAuthConfirmPasswordReset(t *testing.T) {
} }
func TestRecordAuthRequestVerification(t *testing.T) { func TestRecordAuthRequestVerification(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "not an auth collection", Name: "not an auth collection",
@@ -631,6 +912,8 @@ func TestRecordAuthRequestVerification(t *testing.T) {
} }
func TestRecordAuthConfirmVerification(t *testing.T) { func TestRecordAuthConfirmVerification(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "empty data", Name: "empty data",
@@ -671,10 +954,8 @@ func TestRecordAuthConfirmVerification(t *testing.T) {
Body: strings.NewReader(`{ Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg" "token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.R_4FOSUHIuJQ5Crl3PpIPCXMsoHzuTaNlccpXg_3FOg"
}`), }`),
ExpectedStatus: 400, ExpectedStatus: 400,
ExpectedContent: []string{ ExpectedContent: []string{`"data":{}`},
`"data":{}`,
},
}, },
{ {
Name: "different auth collection", Name: "different auth collection",
@@ -732,6 +1013,27 @@ func TestRecordAuthConfirmVerification(t *testing.T) {
"OnRecordAfterConfirmVerificationRequest": 1, "OnRecordAfterConfirmVerificationRequest": 1,
}, },
}, },
{
Name: "OnRecordAfterConfirmVerificationRequest error response",
Method: http.MethodPost,
Url: "/api/collections/users/confirm-verification",
Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImVtYWlsIjoidGVzdEBleGFtcGxlLmNvbSIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZXhwIjoyMjA4OTg1MjYxfQ.hL16TVmStHFdHLc4a860bRqJ3sFfzjv0_NRNzwsvsrc"
}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterConfirmVerificationRequest().Add(func(e *core.RecordConfirmVerificationEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelAfterUpdate": 1,
"OnModelBeforeUpdate": 1,
"OnRecordBeforeConfirmVerificationRequest": 1,
"OnRecordAfterConfirmVerificationRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -740,6 +1042,8 @@ func TestRecordAuthConfirmVerification(t *testing.T) {
} }
func TestRecordAuthRequestEmailChange(t *testing.T) { func TestRecordAuthRequestEmailChange(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -815,7 +1119,7 @@ func TestRecordAuthRequestEmailChange(t *testing.T) {
ExpectedStatus: 400, ExpectedStatus: 400,
ExpectedContent: []string{ ExpectedContent: []string{
`"data":`, `"data":`,
`"newEmail":{"code":"validation_record_email_exists"`, `"newEmail":{"code":"validation_record_email_invalid"`,
}, },
}, },
{ {
@@ -842,6 +1146,8 @@ func TestRecordAuthRequestEmailChange(t *testing.T) {
} }
func TestRecordAuthConfirmEmailChange(t *testing.T) { func TestRecordAuthConfirmEmailChange(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "not an auth collection", Name: "not an auth collection",
@@ -932,6 +1238,28 @@ func TestRecordAuthConfirmEmailChange(t *testing.T) {
`"token":{"code":"validation_token_collection_mismatch"`, `"token":{"code":"validation_token_collection_mismatch"`,
}, },
}, },
{
Name: "OnRecordAfterConfirmEmailChangeRequest error response",
Method: http.MethodPost,
Url: "/api/collections/users/confirm-email-change",
Body: strings.NewReader(`{
"token":"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6IjRxMXhsY2xtZmxva3UzMyIsImNvbGxlY3Rpb25JZCI6Il9wYl91c2Vyc19hdXRoXyIsInR5cGUiOiJhdXRoUmVjb3JkIiwiZW1haWwiOiJ0ZXN0QGV4YW1wbGUuY29tIiwibmV3RW1haWwiOiJjaGFuZ2VAZXhhbXBsZS5jb20iLCJleHAiOjIyMDg5ODUyNjF9.1sG6cL708pRXXjiHRZhG-in0X5fnttSf5nNcadKoYRs",
"password":"1234567890"
}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterConfirmEmailChangeRequest().Add(func(e *core.RecordConfirmEmailChangeEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelAfterUpdate": 1,
"OnModelBeforeUpdate": 1,
"OnRecordBeforeConfirmEmailChangeRequest": 1,
"OnRecordAfterConfirmEmailChangeRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -940,6 +1268,8 @@ func TestRecordAuthConfirmEmailChange(t *testing.T) {
} }
func TestRecordAuthListExternalsAuths(t *testing.T) { func TestRecordAuthListExternalsAuths(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -1040,6 +1370,8 @@ func TestRecordAuthListExternalsAuths(t *testing.T) {
} }
func TestRecordAuthUnlinkExternalsAuth(t *testing.T) { func TestRecordAuthUnlinkExternalsAuth(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -1083,7 +1415,7 @@ func TestRecordAuthUnlinkExternalsAuth(t *testing.T) {
"OnRecordAfterUnlinkExternalAuthRequest": 1, "OnRecordAfterUnlinkExternalAuthRequest": 1,
"OnRecordBeforeUnlinkExternalAuthRequest": 1, "OnRecordBeforeUnlinkExternalAuthRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
record, err := app.Dao().FindRecordById("users", "4q1xlclmfloku33") record, err := app.Dao().FindRecordById("users", "4q1xlclmfloku33")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1129,7 +1461,7 @@ func TestRecordAuthUnlinkExternalsAuth(t *testing.T) {
"OnRecordAfterUnlinkExternalAuthRequest": 1, "OnRecordAfterUnlinkExternalAuthRequest": 1,
"OnRecordBeforeUnlinkExternalAuthRequest": 1, "OnRecordBeforeUnlinkExternalAuthRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
record, err := app.Dao().FindRecordById("users", "4q1xlclmfloku33") record, err := app.Dao().FindRecordById("users", "4q1xlclmfloku33")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -1140,6 +1472,27 @@ func TestRecordAuthUnlinkExternalsAuth(t *testing.T) {
} }
}, },
}, },
{
Name: "OnRecordBeforeUnlinkExternalAuthRequest error response",
Method: http.MethodDelete,
Url: "/api/collections/users/records/4q1xlclmfloku33/external-auths/google",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterUnlinkExternalAuthRequest().Add(func(e *core.RecordUnlinkExternalAuthEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelAfterDelete": 1,
"OnModelBeforeDelete": 1,
"OnRecordAfterUnlinkExternalAuthRequest": 1,
"OnRecordBeforeUnlinkExternalAuthRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -1148,114 +1501,207 @@ func TestRecordAuthUnlinkExternalsAuth(t *testing.T) {
} }
func TestRecordAuthOAuth2Redirect(t *testing.T) { func TestRecordAuthOAuth2Redirect(t *testing.T) {
c1 := subscriptions.NewDefaultClient() t.Parallel()
c2 := subscriptions.NewDefaultClient() clientStubs := make([]map[string]subscriptions.Client, 0, 10)
c2.Subscribe("@oauth2")
c3 := subscriptions.NewDefaultClient() for i := 0; i < 10; i++ {
c3.Subscribe("test1", "@oauth2") c1 := subscriptions.NewDefaultClient()
c4 := subscriptions.NewDefaultClient() c2 := subscriptions.NewDefaultClient()
c4.Subscribe("test1", "test2") c2.Subscribe("@oauth2")
c5 := subscriptions.NewDefaultClient() c3 := subscriptions.NewDefaultClient()
c5.Subscribe("@oauth2") c3.Subscribe("test1", "@oauth2")
c5.Discard()
beforeTestFunc := func(t *testing.T, app *tests.TestApp, e *echo.Echo) { c4 := subscriptions.NewDefaultClient()
app.SubscriptionsBroker().Register(c1) c4.Subscribe("test1", "test2")
app.SubscriptionsBroker().Register(c2)
app.SubscriptionsBroker().Register(c3) c5 := subscriptions.NewDefaultClient()
app.SubscriptionsBroker().Register(c4) c5.Subscribe("@oauth2")
app.SubscriptionsBroker().Register(c5) c5.Discard()
clientStubs = append(clientStubs, map[string]subscriptions.Client{
"c1": c1,
"c2": c2,
"c3": c3,
"c4": c4,
"c5": c5,
})
}
checkFailureRedirect := func(t *testing.T, app *tests.TestApp, res *http.Response) {
loc := res.Header.Get("Location")
if !strings.Contains(loc, "/oauth2-redirect-failure") {
t.Fatalf("Expected failure redirect, got %q", loc)
}
}
checkSuccessRedirect := func(t *testing.T, app *tests.TestApp, res *http.Response) {
loc := res.Header.Get("Location")
if !strings.Contains(loc, "/oauth2-redirect-success") {
t.Fatalf("Expected success redirect, got %q", loc)
}
}
checkClientMessages := func(t *testing.T, clientId string, msg subscriptions.Message, expectedMessages map[string][]string) {
if len(expectedMessages[clientId]) == 0 {
t.Fatalf("Unexpected client %q message, got %s:\n%s", clientId, msg.Name, msg.Data)
}
if msg.Name != "@oauth2" {
t.Fatalf("Expected @oauth2 msg.Name, got %q", msg.Name)
}
for _, txt := range expectedMessages[clientId] {
if !strings.Contains(string(msg.Data), txt) {
t.Fatalf("Failed to find %q in \n%s", txt, msg.Data)
}
}
}
beforeTestFunc := func(
clients map[string]subscriptions.Client,
expectedMessages map[string][]string,
) func(*testing.T, *tests.TestApp, *echo.Echo) {
return func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
for _, client := range clients {
app.SubscriptionsBroker().Register(client)
}
ctx, cancelFunc := context.WithTimeout(context.Background(), 100*time.Millisecond)
// add to the app store so that it can be cancelled manually after test completion
app.Store().Set("cancelFunc", cancelFunc)
go func() {
defer cancelFunc()
for {
select {
case msg := <-clients["c1"].Channel():
checkClientMessages(t, "c1", msg, expectedMessages)
case msg := <-clients["c2"].Channel():
checkClientMessages(t, "c2", msg, expectedMessages)
case msg := <-clients["c3"].Channel():
checkClientMessages(t, "c3", msg, expectedMessages)
case msg := <-clients["c4"].Channel():
checkClientMessages(t, "c4", msg, expectedMessages)
case msg := <-clients["c5"].Channel():
checkClientMessages(t, "c5", msg, expectedMessages)
case <-ctx.Done():
for _, c := range clients {
close(c.Channel())
}
return
}
}
}()
}
} }
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "no state query param", Name: "no state query param",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/oauth2-redirect?code=123", Url: "/api/oauth2-redirect?code=123",
ExpectedStatus: 400, BeforeTestFunc: beforeTestFunc(clientStubs[0], nil),
ExpectedContent: []string{`"data":{}`}, ExpectedStatus: http.StatusTemporaryRedirect,
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkFailureRedirect(t, app, res)
},
}, },
{ {
Name: "no code query param", Name: "invalid or missing client",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/oauth2-redirect?state=" + c3.Id(), Url: "/api/oauth2-redirect?code=123&state=missing",
ExpectedStatus: 400, BeforeTestFunc: beforeTestFunc(clientStubs[1], nil),
ExpectedContent: []string{`"data":{}`}, ExpectedStatus: http.StatusTemporaryRedirect,
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkFailureRedirect(t, app, res)
},
}, },
{ {
Name: "missing client", Name: "no code query param",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/oauth2-redirect?code=123&state=missing", Url: "/api/oauth2-redirect?state=" + clientStubs[2]["c3"].Id(),
ExpectedStatus: 404, BeforeTestFunc: beforeTestFunc(clientStubs[2], map[string][]string{
ExpectedContent: []string{`"data":{}`}, "c3": {`"state":"` + clientStubs[2]["c3"].Id(), `"code":""`},
}),
ExpectedStatus: http.StatusTemporaryRedirect,
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkFailureRedirect(t, app, res)
if clientStubs[2]["c3"].HasSubscription("@oauth2") {
t.Fatalf("Expected oauth2 subscription to be removed")
}
},
}, },
{ {
Name: "discarded client with @oauth2 subscription", Name: "error query param",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/oauth2-redirect?code=123&state=" + c5.Id(), Url: "/api/oauth2-redirect?error=example&code=123&state=" + clientStubs[3]["c3"].Id(),
BeforeTestFunc: beforeTestFunc, BeforeTestFunc: beforeTestFunc(clientStubs[3], map[string][]string{
ExpectedStatus: 404, "c3": {`"state":"` + clientStubs[3]["c3"].Id(), `"code":"123"`, `"error":"example"`},
ExpectedContent: []string{`"data":{}`}, }),
ExpectedStatus: http.StatusTemporaryRedirect,
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkFailureRedirect(t, app, res)
if clientStubs[3]["c3"].HasSubscription("@oauth2") {
t.Fatalf("Expected oauth2 subscription to be removed")
}
},
}, },
{ {
Name: "client without @oauth2 subscription", Name: "discarded client with @oauth2 subscription",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/oauth2-redirect?code=123&state=" + c4.Id(), Url: "/api/oauth2-redirect?code=123&state=" + clientStubs[4]["c5"].Id(),
BeforeTestFunc: beforeTestFunc, BeforeTestFunc: beforeTestFunc(clientStubs[4], nil),
ExpectedStatus: 404, ExpectedStatus: http.StatusTemporaryRedirect,
ExpectedContent: []string{`"data":{}`}, AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkFailureRedirect(t, app, res)
},
},
{
Name: "client without @oauth2 subscription",
Method: http.MethodGet,
Url: "/api/oauth2-redirect?code=123&state=" + clientStubs[4]["c4"].Id(),
BeforeTestFunc: beforeTestFunc(clientStubs[5], nil),
ExpectedStatus: http.StatusTemporaryRedirect,
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkFailureRedirect(t, app, res)
},
}, },
{ {
Name: "client with @oauth2 subscription", Name: "client with @oauth2 subscription",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/oauth2-redirect?code=123&state=" + c3.Id(), Url: "/api/oauth2-redirect?code=123&state=" + clientStubs[6]["c3"].Id(),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { BeforeTestFunc: beforeTestFunc(clientStubs[6], map[string][]string{
beforeTestFunc(t, app, e) "c3": {`"state":"` + clientStubs[6]["c3"].Id(), `"code":"123"`},
}),
ctx, cancelFunc := context.WithTimeout(context.Background(), 1*time.Second)
go func() {
defer cancelFunc()
L:
for {
select {
case <-c1.Channel():
t.Error("Unexpected c1 message")
break L
case <-c2.Channel():
t.Error("Unexpected c2 message")
break L
case msg := <-c3.Channel():
if msg.Name != "@oauth2" {
t.Errorf("Expected @oauth2 msg.Name, got %q", msg.Name)
}
expectedParams := []string{`"state"`, `"code"`}
for _, p := range expectedParams {
if !strings.Contains(msg.Data, p) {
t.Errorf("Couldn't find %s in \n%v", p, msg.Data)
}
}
break L
case <-c4.Channel():
t.Error("Unexpected c4 message")
break L
case <-c5.Channel():
t.Error("Unexpected c5 message")
break L
case <-ctx.Done():
t.Error("Context timeout reached")
break L
}
}
}()
},
ExpectedStatus: http.StatusTemporaryRedirect, ExpectedStatus: http.StatusTemporaryRedirect,
AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
app.Store().Get("cancelFunc").(context.CancelFunc)()
checkSuccessRedirect(t, app, res)
if clientStubs[6]["c3"].HasSubscription("@oauth2") {
t.Fatalf("Expected oauth2 subscription to be removed")
}
},
}, },
} }
+89 -98
View File
@@ -2,9 +2,8 @@ package apis
import ( import (
"fmt" "fmt"
"log" "log/slog"
"net/http" "net/http"
"strings"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
@@ -16,8 +15,6 @@ import (
"github.com/pocketbase/pocketbase/tools/search" "github.com/pocketbase/pocketbase/tools/search"
) )
const expandQueryParam = "expand"
// bindRecordCrudApi registers the record crud api endpoints and // bindRecordCrudApi registers the record crud api endpoints and
// the corresponding handlers. // the corresponding handlers.
func bindRecordCrudApi(app core.App, rg *echo.Group) { func bindRecordCrudApi(app core.App, rg *echo.Group) {
@@ -45,14 +42,14 @@ func (api *recordApi) list(c echo.Context) error {
return NewNotFoundError("", "Missing collection context.") return NewNotFoundError("", "Missing collection context.")
} }
requestInfo := RequestInfo(c)
// forbid users and guests to query special filter/sort fields // forbid users and guests to query special filter/sort fields
if err := api.checkForForbiddenQueryFields(c); err != nil { if err := checkForAdminOnlyRuleFields(requestInfo); err != nil {
return err return err
} }
requestData := RequestData(c) if requestInfo.Admin == nil && collection.ListRule == nil {
if requestData.Admin == nil && collection.ListRule == nil {
// only admins can access if the rule is nil // only admins can access if the rule is nil
return NewForbiddenError("Only admins can perform this action.", nil) return NewForbiddenError("Only admins can perform this action.", nil)
} }
@@ -60,20 +57,15 @@ func (api *recordApi) list(c echo.Context) error {
fieldsResolver := resolvers.NewRecordFieldResolver( fieldsResolver := resolvers.NewRecordFieldResolver(
api.app.Dao(), api.app.Dao(),
collection, collection,
requestData, requestInfo,
// hidden fields are searchable only by admins // hidden fields are searchable only by admins
requestData.Admin != nil, requestInfo.Admin != nil,
) )
searchProvider := search.NewProvider(fieldsResolver). searchProvider := search.NewProvider(fieldsResolver).
Query(api.app.Dao().RecordQuery(collection)) Query(api.app.Dao().RecordQuery(collection))
// views don't have "rowid" so we fallback to "id" if requestInfo.Admin == nil && collection.ListRule != nil {
if collection.IsView() {
searchProvider.CountCol("id")
}
if requestData.Admin == nil && collection.ListRule != nil {
searchProvider.AddFilter(search.FilterData(*collection.ListRule)) searchProvider.AddFilter(search.FilterData(*collection.ListRule))
} }
@@ -81,7 +73,7 @@ func (api *recordApi) list(c echo.Context) error {
result, err := searchProvider.ParseAndExec(c.QueryParams().Encode(), &records) result, err := searchProvider.ParseAndExec(c.QueryParams().Encode(), &records)
if err != nil { if err != nil {
return NewBadRequestError("Invalid filter parameters.", err) return NewBadRequestError("", err)
} }
event := new(core.RecordsListEvent) event := new(core.RecordsListEvent)
@@ -91,8 +83,12 @@ func (api *recordApi) list(c echo.Context) error {
event.Result = result event.Result = result
return api.app.OnRecordsListRequest().Trigger(event, func(e *core.RecordsListEvent) error { return api.app.OnRecordsListRequest().Trigger(event, func(e *core.RecordsListEvent) error {
if err := EnrichRecords(e.HttpContext, api.app.Dao(), e.Records); err != nil && api.app.IsDebug() { if e.HttpContext.Response().Committed {
log.Println(err) return nil
}
if err := EnrichRecords(e.HttpContext, api.app.Dao(), e.Records); err != nil {
api.app.Logger().Debug("Failed to enrich list records", slog.String("error", err.Error()))
} }
return e.HttpContext.JSON(http.StatusOK, e.Result) return e.HttpContext.JSON(http.StatusOK, e.Result)
@@ -110,16 +106,16 @@ func (api *recordApi) view(c echo.Context) error {
return NewNotFoundError("", nil) return NewNotFoundError("", nil)
} }
requestData := RequestData(c) requestInfo := RequestInfo(c)
if requestData.Admin == nil && collection.ViewRule == nil { if requestInfo.Admin == nil && collection.ViewRule == nil {
// only admins can access if the rule is nil // only admins can access if the rule is nil
return NewForbiddenError("Only admins can perform this action.", nil) return NewForbiddenError("Only admins can perform this action.", nil)
} }
ruleFunc := func(q *dbx.SelectQuery) error { ruleFunc := func(q *dbx.SelectQuery) error {
if requestData.Admin == nil && collection.ViewRule != nil && *collection.ViewRule != "" { if requestInfo.Admin == nil && collection.ViewRule != nil && *collection.ViewRule != "" {
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestData, true) resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestInfo, true)
expr, err := search.FilterData(*collection.ViewRule).BuildExpr(resolver) expr, err := search.FilterData(*collection.ViewRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -141,8 +137,17 @@ func (api *recordApi) view(c echo.Context) error {
event.Record = record event.Record = record
return api.app.OnRecordViewRequest().Trigger(event, func(e *core.RecordViewEvent) error { return api.app.OnRecordViewRequest().Trigger(event, func(e *core.RecordViewEvent) error {
if err := EnrichRecord(e.HttpContext, api.app.Dao(), e.Record); err != nil && api.app.IsDebug() { if e.HttpContext.Response().Committed {
log.Println(err) return nil
}
if err := EnrichRecord(e.HttpContext, api.app.Dao(), e.Record); err != nil {
api.app.Logger().Debug(
"Failed to enrich view record",
slog.String("id", e.Record.Id),
slog.String("collectionName", e.Record.Collection().Name),
slog.String("error", err.Error()),
)
} }
return e.HttpContext.JSON(http.StatusOK, e.Record) return e.HttpContext.JSON(http.StatusOK, e.Record)
@@ -155,23 +160,23 @@ func (api *recordApi) create(c echo.Context) error {
return NewNotFoundError("", "Missing collection context.") return NewNotFoundError("", "Missing collection context.")
} }
requestData := RequestData(c) requestInfo := RequestInfo(c)
if requestData.Admin == nil && collection.CreateRule == nil { if requestInfo.Admin == nil && collection.CreateRule == nil {
// only admins can access if the rule is nil // only admins can access if the rule is nil
return NewForbiddenError("Only admins can perform this action.", nil) return NewForbiddenError("Only admins can perform this action.", nil)
} }
hasFullManageAccess := requestData.Admin != nil hasFullManageAccess := requestInfo.Admin != nil
// temporary save the record and check it against the create rule // temporary save the record and check it against the create rule
if requestData.Admin == nil && collection.CreateRule != nil { if requestInfo.Admin == nil && collection.CreateRule != nil {
testRecord := models.NewRecord(collection) testRecord := models.NewRecord(collection)
// replace modifiers fields so that the resolved value is always // replace modifiers fields so that the resolved value is always
// available when accessing requestData.Data using just the field name // available when accessing requestInfo.Data using just the field name
if requestData.HasModifierDataKeys() { if requestInfo.HasModifierDataKeys() {
requestData.Data = testRecord.ReplaceModifers(requestData.Data) requestInfo.Data = testRecord.ReplaceModifers(requestInfo.Data)
} }
testForm := forms.NewRecordUpsert(api.app, testRecord) testForm := forms.NewRecordUpsert(api.app, testRecord)
@@ -185,7 +190,7 @@ func (api *recordApi) create(c echo.Context) error {
return nil // no create rule to resolve return nil // no create rule to resolve
} }
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestData, true) resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestInfo, true)
expr, err := search.FilterData(*collection.CreateRule).BuildExpr(resolver) expr, err := search.FilterData(*collection.CreateRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -200,7 +205,7 @@ func (api *recordApi) create(c echo.Context) error {
if err != nil { if err != nil {
return fmt.Errorf("DrySubmit create rule failure: %w", err) return fmt.Errorf("DrySubmit create rule failure: %w", err)
} }
hasFullManageAccess = hasAuthManageAccess(txDao, foundRecord, requestData) hasFullManageAccess = hasAuthManageAccess(txDao, foundRecord, requestInfo)
return nil return nil
}) })
@@ -225,7 +230,7 @@ func (api *recordApi) create(c echo.Context) error {
event.UploadedFiles = form.FilesToUpload() event.UploadedFiles = form.FilesToUpload()
// create the record // create the record
submitErr := form.Submit(func(next forms.InterceptorNextFunc[*models.Record]) forms.InterceptorNextFunc[*models.Record] { return form.Submit(func(next forms.InterceptorNextFunc[*models.Record]) forms.InterceptorNextFunc[*models.Record] {
return func(m *models.Record) error { return func(m *models.Record) error {
event.Record = m event.Record = m
@@ -234,22 +239,25 @@ func (api *recordApi) create(c echo.Context) error {
return NewBadRequestError("Failed to create record.", err) return NewBadRequestError("Failed to create record.", err)
} }
if err := EnrichRecord(e.HttpContext, api.app.Dao(), e.Record); err != nil && api.app.IsDebug() { if err := EnrichRecord(e.HttpContext, api.app.Dao(), e.Record); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to enrich create record",
slog.String("id", e.Record.Id),
slog.String("collectionName", e.Record.Collection().Name),
slog.String("error", err.Error()),
)
} }
return e.HttpContext.JSON(http.StatusOK, e.Record) return api.app.OnRecordAfterCreateRequest().Trigger(event, func(e *core.RecordCreateEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Record)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterCreateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr
} }
func (api *recordApi) update(c echo.Context) error { func (api *recordApi) update(c echo.Context) error {
@@ -263,26 +271,26 @@ func (api *recordApi) update(c echo.Context) error {
return NewNotFoundError("", nil) return NewNotFoundError("", nil)
} }
requestData := RequestData(c) requestInfo := RequestInfo(c)
if requestData.Admin == nil && collection.UpdateRule == nil { if requestInfo.Admin == nil && collection.UpdateRule == nil {
// only admins can access if the rule is nil // only admins can access if the rule is nil
return NewForbiddenError("Only admins can perform this action.", nil) return NewForbiddenError("Only admins can perform this action.", nil)
} }
// eager fetch the record so that the modifier field values are replaced // eager fetch the record so that the modifier field values are replaced
// and available when accessing requestData.Data using just the field name // and available when accessing requestInfo.Data using just the field name
if requestData.HasModifierDataKeys() { if requestInfo.HasModifierDataKeys() {
record, err := api.app.Dao().FindRecordById(collection.Id, recordId) record, err := api.app.Dao().FindRecordById(collection.Id, recordId)
if err != nil || record == nil { if err != nil || record == nil {
return NewNotFoundError("", err) return NewNotFoundError("", err)
} }
requestData.Data = record.ReplaceModifers(requestData.Data) requestInfo.Data = record.ReplaceModifers(requestInfo.Data)
} }
ruleFunc := func(q *dbx.SelectQuery) error { ruleFunc := func(q *dbx.SelectQuery) error {
if requestData.Admin == nil && collection.UpdateRule != nil && *collection.UpdateRule != "" { if requestInfo.Admin == nil && collection.UpdateRule != nil && *collection.UpdateRule != "" {
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestData, true) resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestInfo, true)
expr, err := search.FilterData(*collection.UpdateRule).BuildExpr(resolver) expr, err := search.FilterData(*collection.UpdateRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -300,7 +308,7 @@ func (api *recordApi) update(c echo.Context) error {
} }
form := forms.NewRecordUpsert(api.app, record) form := forms.NewRecordUpsert(api.app, record)
form.SetFullManageAccess(requestData.Admin != nil || hasAuthManageAccess(api.app.Dao(), record, requestData)) form.SetFullManageAccess(requestInfo.Admin != nil || hasAuthManageAccess(api.app.Dao(), record, requestInfo))
// load request // load request
if err := form.LoadRequest(c.Request(), ""); err != nil { if err := form.LoadRequest(c.Request(), ""); err != nil {
@@ -314,7 +322,7 @@ func (api *recordApi) update(c echo.Context) error {
event.UploadedFiles = form.FilesToUpload() event.UploadedFiles = form.FilesToUpload()
// update the record // update the record
submitErr := form.Submit(func(next forms.InterceptorNextFunc[*models.Record]) forms.InterceptorNextFunc[*models.Record] { return form.Submit(func(next forms.InterceptorNextFunc[*models.Record]) forms.InterceptorNextFunc[*models.Record] {
return func(m *models.Record) error { return func(m *models.Record) error {
event.Record = m event.Record = m
@@ -323,22 +331,25 @@ func (api *recordApi) update(c echo.Context) error {
return NewBadRequestError("Failed to update record.", err) return NewBadRequestError("Failed to update record.", err)
} }
if err := EnrichRecord(e.HttpContext, api.app.Dao(), e.Record); err != nil && api.app.IsDebug() { if err := EnrichRecord(e.HttpContext, api.app.Dao(), e.Record); err != nil {
log.Println(err) api.app.Logger().Debug(
"Failed to enrich update record",
slog.String("id", e.Record.Id),
slog.String("collectionName", e.Record.Collection().Name),
slog.String("error", err.Error()),
)
} }
return e.HttpContext.JSON(http.StatusOK, e.Record) return api.app.OnRecordAfterUpdateRequest().Trigger(event, func(e *core.RecordUpdateEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.Record)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnRecordAfterUpdateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr
} }
func (api *recordApi) delete(c echo.Context) error { func (api *recordApi) delete(c echo.Context) error {
@@ -352,16 +363,16 @@ func (api *recordApi) delete(c echo.Context) error {
return NewNotFoundError("", nil) return NewNotFoundError("", nil)
} }
requestData := RequestData(c) requestInfo := RequestInfo(c)
if requestData.Admin == nil && collection.DeleteRule == nil { if requestInfo.Admin == nil && collection.DeleteRule == nil {
// only admins can access if the rule is nil // only admins can access if the rule is nil
return NewForbiddenError("Only admins can perform this action.", nil) return NewForbiddenError("Only admins can perform this action.", nil)
} }
ruleFunc := func(q *dbx.SelectQuery) error { ruleFunc := func(q *dbx.SelectQuery) error {
if requestData.Admin == nil && collection.DeleteRule != nil && *collection.DeleteRule != "" { if requestInfo.Admin == nil && collection.DeleteRule != nil && *collection.DeleteRule != "" {
resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestData, true) resolver := resolvers.NewRecordFieldResolver(api.app.Dao(), collection, requestInfo, true)
expr, err := search.FilterData(*collection.DeleteRule).BuildExpr(resolver) expr, err := search.FilterData(*collection.DeleteRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -382,38 +393,18 @@ func (api *recordApi) delete(c echo.Context) error {
event.Collection = collection event.Collection = collection
event.Record = record event.Record = record
handlerErr := api.app.OnRecordBeforeDeleteRequest().Trigger(event, func(e *core.RecordDeleteEvent) error { return api.app.OnRecordBeforeDeleteRequest().Trigger(event, func(e *core.RecordDeleteEvent) error {
// delete the record // delete the record
if err := api.app.Dao().DeleteRecord(e.Record); err != nil { if err := api.app.Dao().DeleteRecord(e.Record); err != nil {
return NewBadRequestError("Failed to delete record. Make sure that the record is not part of a required relation reference.", err) return NewBadRequestError("Failed to delete record. Make sure that the record is not part of a required relation reference.", err)
} }
return e.HttpContext.NoContent(http.StatusNoContent) return api.app.OnRecordAfterDeleteRequest().Trigger(event, func(e *core.RecordDeleteEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.NoContent(http.StatusNoContent)
})
}) })
if handlerErr == nil {
if err := api.app.OnRecordAfterDeleteRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return handlerErr
}
func (api *recordApi) checkForForbiddenQueryFields(c echo.Context) error {
admin, _ := c.Get(ContextAdminKey).(*models.Admin)
if admin != nil {
return nil // admins are allowed to query everything
}
decodedQuery := c.QueryParam(search.FilterQueryParam) + c.QueryParam(search.SortQueryParam)
forbiddenFields := []string{"@collection.", "@request."}
for _, field := range forbiddenFields {
if strings.Contains(decodedQuery, field) {
return NewForbiddenError("Only admins can filter by @collection and @request query params", nil)
}
}
return nil
} }
+249 -7
View File
@@ -1,6 +1,7 @@
package apis_test package apis_test
import ( import (
"errors"
"net/http" "net/http"
"net/url" "net/url"
"os" "os"
@@ -10,11 +11,16 @@ import (
"time" "time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/rest"
"github.com/pocketbase/pocketbase/tools/types"
) )
func TestRecordCrudList(t *testing.T) { func TestRecordCrudList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "missing collection", Name: "missing collection",
@@ -41,9 +47,16 @@ func TestRecordCrudList(t *testing.T) {
ExpectedContent: []string{`"data":{}`}, ExpectedContent: []string{`"data":{}`},
}, },
{ {
Name: "public collection but with admin only filter/sort (aka. @collection)", Name: "public collection but with admin only filter param (aka. @collection, @request, etc.)",
Method: http.MethodGet, Method: http.MethodGet,
Url: "/api/collections/demo2/records?filter=@collection.demo2.title='test1'", Url: "/api/collections/demo2/records?filter=%40collection.demo2.title='test1'",
ExpectedStatus: 403,
ExpectedContent: []string{`"data":{}`},
},
{
Name: "public collection but with admin only sort param (aka. @collection, @request, etc.)",
Method: http.MethodGet,
Url: "/api/collections/demo2/records?sort=@request.auth.title",
ExpectedStatus: 403, ExpectedStatus: 403,
ExpectedContent: []string{`"data":{}`}, ExpectedContent: []string{`"data":{}`},
}, },
@@ -460,6 +473,22 @@ func TestRecordCrudList(t *testing.T) {
}, },
ExpectedEvents: map[string]int{"OnRecordsListRequest": 1}, ExpectedEvents: map[string]int{"OnRecordsListRequest": 1},
}, },
{
Name: "view collection with numeric ids",
Method: http.MethodGet,
Url: "/api/collections/numeric_id_view/records",
ExpectedStatus: 200,
ExpectedContent: []string{
`"page":1`,
`"perPage":30`,
`"totalPages":1`,
`"totalItems":2`,
`"items":[{`,
`"id":"1"`,
`"id":"2"`,
},
ExpectedEvents: map[string]int{"OnRecordsListRequest": 1},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -468,6 +497,8 @@ func TestRecordCrudList(t *testing.T) {
} }
func TestRecordCrudView(t *testing.T) { func TestRecordCrudView(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "missing collection", Name: "missing collection",
@@ -729,6 +760,16 @@ func TestRecordCrudView(t *testing.T) {
}, },
ExpectedEvents: map[string]int{"OnRecordViewRequest": 1}, ExpectedEvents: map[string]int{"OnRecordViewRequest": 1},
}, },
{
Name: "view record with numeric id",
Method: http.MethodGet,
Url: "/api/collections/numeric_id_view/records/1",
ExpectedStatus: 200,
ExpectedContent: []string{
`"id":"1"`,
},
ExpectedEvents: map[string]int{"OnRecordViewRequest": 1},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -737,6 +778,8 @@ func TestRecordCrudView(t *testing.T) {
} }
func TestRecordCrudDelete(t *testing.T) { func TestRecordCrudDelete(t *testing.T) {
t.Parallel()
ensureDeletedFiles := func(app *tests.TestApp, collectionId string, recordId string) { ensureDeletedFiles := func(app *tests.TestApp, collectionId string, recordId string) {
storageDir := filepath.Join(app.DataDir(), "storage", collectionId, recordId) storageDir := filepath.Join(app.DataDir(), "storage", collectionId, recordId)
@@ -836,6 +879,27 @@ func TestRecordCrudDelete(t *testing.T) {
"OnRecordBeforeDeleteRequest": 1, "OnRecordBeforeDeleteRequest": 1,
}, },
}, },
{
Name: "OnRecordAfterDeleteRequest error response",
Method: http.MethodDelete,
Url: "/api/collections/clients/records/o1y0dd0spd786md",
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterDeleteRequest().Add(func(e *core.RecordDeleteEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelAfterDelete": 1,
"OnModelBeforeDelete": 1,
"OnRecordAfterDeleteRequest": 1,
"OnRecordBeforeDeleteRequest": 1,
},
},
{ {
Name: "authenticated record that match the collection delete rule", Name: "authenticated record that match the collection delete rule",
Method: http.MethodDelete, Method: http.MethodDelete,
@@ -854,7 +918,7 @@ func TestRecordCrudDelete(t *testing.T) {
"OnRecordAfterDeleteRequest": 1, "OnRecordAfterDeleteRequest": 1,
"OnRecordBeforeDeleteRequest": 1, "OnRecordBeforeDeleteRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
ensureDeletedFiles(app, "_pb_users_auth_", "4q1xlclmfloku33") ensureDeletedFiles(app, "_pb_users_auth_", "4q1xlclmfloku33")
// check if all the external auths records were deleted // check if all the external auths records were deleted
@@ -941,7 +1005,7 @@ func TestRecordCrudDelete(t *testing.T) {
"OnRecordBeforeDeleteRequest": 1, "OnRecordBeforeDeleteRequest": 1,
"OnRecordAfterDeleteRequest": 1, "OnRecordAfterDeleteRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
recId := "84nmscqy84lsi1t" recId := "84nmscqy84lsi1t"
rec, _ := app.Dao().FindRecordById("demo1", recId, nil) rec, _ := app.Dao().FindRecordById("demo1", recId, nil)
if rec != nil { if rec != nil {
@@ -959,6 +1023,8 @@ func TestRecordCrudDelete(t *testing.T) {
} }
func TestRecordCrudCreate(t *testing.T) { func TestRecordCrudCreate(t *testing.T) {
t.Parallel()
formData, mp, err := tests.MockMultipartData(map[string]string{ formData, mp, err := tests.MockMultipartData(map[string]string{
"title": "title_test", "title": "title_test",
}, "files") }, "files")
@@ -966,6 +1032,20 @@ func TestRecordCrudCreate(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
formData2, mp2, err2 := tests.MockMultipartData(map[string]string{
rest.MultipartJsonKey: `{"title": "title_test2", "testPayload": 123}`,
}, "files")
if err2 != nil {
t.Fatal(err2)
}
formData3, mp3, err3 := tests.MockMultipartData(map[string]string{
rest.MultipartJsonKey: `{"title": "title_test3", "testPayload": 123}`,
}, "files")
if err3 != nil {
t.Fatal(err3)
}
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "missing collection", Name: "missing collection",
@@ -1173,6 +1253,60 @@ func TestRecordCrudCreate(t *testing.T) {
"OnModelAfterCreate": 1, "OnModelAfterCreate": 1,
}, },
}, },
{
Name: "submit via multipart form data with @jsonPayload key and unsatisfied @request.data rule",
Method: http.MethodPost,
Url: "/api/collections/demo3/records",
Body: formData2,
RequestHeaders: map[string]string{
"Content-Type": mp2.FormDataContentType(),
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
collection, err := app.Dao().FindCollectionByNameOrId("demo3")
if err != nil {
t.Fatalf("failed to find demo3 collection: %v", err)
}
collection.CreateRule = types.Pointer("@request.data.testPayload != 123")
if err := app.Dao().WithoutHooks().SaveCollection(collection); err != nil {
t.Fatalf("failed to update demo3 collection create rule: %v", err)
}
core.ReloadCachedCollections(app)
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
},
{
Name: "submit via multipart form data with @jsonPayload key and satisfied @request.data rule",
Method: http.MethodPost,
Url: "/api/collections/demo3/records",
Body: formData3,
RequestHeaders: map[string]string{
"Content-Type": mp3.FormDataContentType(),
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
collection, err := app.Dao().FindCollectionByNameOrId("demo3")
if err != nil {
t.Fatalf("failed to find demo3 collection: %v", err)
}
collection.CreateRule = types.Pointer("@request.data.testPayload = 123")
if err := app.Dao().WithoutHooks().SaveCollection(collection); err != nil {
t.Fatalf("failed to update demo3 collection create rule: %v", err)
}
core.ReloadCachedCollections(app)
},
ExpectedStatus: 200,
ExpectedContent: []string{
`"id":"`,
`"title":"title_test3"`,
`"files":["`,
},
ExpectedEvents: map[string]int{
"OnRecordBeforeCreateRequest": 1,
"OnRecordAfterCreateRequest": 1,
"OnModelBeforeCreate": 1,
"OnModelAfterCreate": 1,
},
},
{ {
Name: "unique field error check", Name: "unique field error check",
Method: http.MethodPost, Method: http.MethodPost,
@@ -1187,6 +1321,25 @@ func TestRecordCrudCreate(t *testing.T) {
`"code":"validation_not_unique"`, `"code":"validation_not_unique"`,
}, },
}, },
{
Name: "OnRecordAfterCreateRequest error response",
Method: http.MethodPost,
Url: "/api/collections/demo2/records",
Body: strings.NewReader(`{"title":"new"}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterCreateRequest().Add(func(e *core.RecordCreateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnRecordBeforeCreateRequest": 1,
"OnRecordAfterCreateRequest": 1,
"OnModelBeforeCreate": 1,
"OnModelAfterCreate": 1,
},
},
// ID checks // ID checks
// ----------------------------------------------------------- // -----------------------------------------------------------
@@ -1516,6 +1669,8 @@ func TestRecordCrudCreate(t *testing.T) {
} }
func TestRecordCrudUpdate(t *testing.T) { func TestRecordCrudUpdate(t *testing.T) {
t.Parallel()
formData, mp, err := tests.MockMultipartData(map[string]string{ formData, mp, err := tests.MockMultipartData(map[string]string{
"title": "title_test", "title": "title_test",
}, "files") }, "files")
@@ -1523,6 +1678,20 @@ func TestRecordCrudUpdate(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
formData2, mp2, err2 := tests.MockMultipartData(map[string]string{
rest.MultipartJsonKey: `{"title": "title_test2", "testPayload": 123}`,
}, "files")
if err2 != nil {
t.Fatal(err2)
}
formData3, mp3, err3 := tests.MockMultipartData(map[string]string{
rest.MultipartJsonKey: `{"title": "title_test3", "testPayload": 123}`,
}, "files")
if err3 != nil {
t.Fatal(err3)
}
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "missing collection", Name: "missing collection",
@@ -1745,6 +1914,79 @@ func TestRecordCrudUpdate(t *testing.T) {
"OnModelAfterUpdate": 1, "OnModelAfterUpdate": 1,
}, },
}, },
{
Name: "submit via multipart form data with @jsonPayload key and unsatisfied @request.data rule",
Method: http.MethodPatch,
Url: "/api/collections/demo3/records/mk5fmymtx4wsprk",
Body: formData2,
RequestHeaders: map[string]string{
"Content-Type": mp2.FormDataContentType(),
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
collection, err := app.Dao().FindCollectionByNameOrId("demo3")
if err != nil {
t.Fatalf("failed to find demo3 collection: %v", err)
}
collection.UpdateRule = types.Pointer("@request.data.testPayload != 123")
if err := app.Dao().WithoutHooks().SaveCollection(collection); err != nil {
t.Fatalf("failed to update demo3 collection update rule: %v", err)
}
core.ReloadCachedCollections(app)
},
ExpectedStatus: 404,
ExpectedContent: []string{`"data":{}`},
},
{
Name: "submit via multipart form data with @jsonPayload key and satisfied @request.data rule",
Method: http.MethodPatch,
Url: "/api/collections/demo3/records/mk5fmymtx4wsprk",
Body: formData3,
RequestHeaders: map[string]string{
"Content-Type": mp3.FormDataContentType(),
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
collection, err := app.Dao().FindCollectionByNameOrId("demo3")
if err != nil {
t.Fatalf("failed to find demo3 collection: %v", err)
}
collection.UpdateRule = types.Pointer("@request.data.testPayload = 123")
if err := app.Dao().WithoutHooks().SaveCollection(collection); err != nil {
t.Fatalf("failed to update demo3 collection update rule: %v", err)
}
core.ReloadCachedCollections(app)
},
ExpectedStatus: 200,
ExpectedContent: []string{
`"id":"mk5fmymtx4wsprk"`,
`"title":"title_test3"`,
`"files":["`,
},
ExpectedEvents: map[string]int{
"OnRecordBeforeUpdateRequest": 1,
"OnRecordAfterUpdateRequest": 1,
"OnModelBeforeUpdate": 1,
"OnModelAfterUpdate": 1,
},
},
{
Name: "OnRecordAfterUpdateRequest error response",
Method: http.MethodPatch,
Url: "/api/collections/demo2/records/0yxhwia2amd8gec",
Body: strings.NewReader(`{"title":"new"}`),
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnRecordAfterUpdateRequest().Add(func(e *core.RecordUpdateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnRecordBeforeUpdateRequest": 1,
"OnRecordAfterUpdateRequest": 1,
"OnModelBeforeUpdate": 1,
"OnModelAfterUpdate": 1,
},
},
{ {
Name: "try to change the id of an existing record", Name: "try to change the id of an existing record",
Method: http.MethodPatch, Method: http.MethodPatch,
@@ -1945,7 +2187,7 @@ func TestRecordCrudUpdate(t *testing.T) {
"OnRecordAfterUpdateRequest": 1, "OnRecordAfterUpdateRequest": 1,
"OnRecordBeforeUpdateRequest": 1, "OnRecordBeforeUpdateRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
record, _ := app.Dao().FindRecordById("nologin", "phhq3wr65cap535") record, _ := app.Dao().FindRecordById("nologin", "phhq3wr65cap535")
if !record.ValidatePassword("12345678") { if !record.ValidatePassword("12345678") {
t.Fatal("Password update failed.") t.Fatal("Password update failed.")
@@ -1988,7 +2230,7 @@ func TestRecordCrudUpdate(t *testing.T) {
"OnRecordAfterUpdateRequest": 1, "OnRecordAfterUpdateRequest": 1,
"OnRecordBeforeUpdateRequest": 1, "OnRecordBeforeUpdateRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
record, _ := app.Dao().FindRecordById("users", "oap640cot4yru2s") record, _ := app.Dao().FindRecordById("users", "oap640cot4yru2s")
if !record.ValidatePassword("12345678") { if !record.ValidatePassword("12345678") {
t.Fatal("Password update failed.") t.Fatal("Password update failed.")
@@ -2050,7 +2292,7 @@ func TestRecordCrudUpdate(t *testing.T) {
"OnRecordAfterUpdateRequest": 1, "OnRecordAfterUpdateRequest": 1,
"OnRecordBeforeUpdateRequest": 1, "OnRecordBeforeUpdateRequest": 1,
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
record, _ := app.Dao().FindRecordById("nologin", "dc49k6jgejn40h3") record, _ := app.Dao().FindRecordById("nologin", "dc49k6jgejn40h3")
if !record.ValidatePassword("123456789") { if !record.ValidatePassword("123456789") {
t.Fatal("Password update failed.") t.Fatal("Password update failed.")
+97 -32
View File
@@ -3,6 +3,7 @@ package apis
import ( import (
"fmt" "fmt"
"log" "log"
"log/slog"
"net/http" "net/http"
"strings" "strings"
@@ -13,23 +14,37 @@ import (
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/resolvers" "github.com/pocketbase/pocketbase/resolvers"
"github.com/pocketbase/pocketbase/tokens" "github.com/pocketbase/pocketbase/tokens"
"github.com/pocketbase/pocketbase/tools/inflector"
"github.com/pocketbase/pocketbase/tools/rest" "github.com/pocketbase/pocketbase/tools/rest"
"github.com/pocketbase/pocketbase/tools/search" "github.com/pocketbase/pocketbase/tools/search"
) )
const ContextRequestDataKey = "requestData" const ContextRequestInfoKey = "requestInfo"
// RequestData exports cached common request data fields const expandQueryParam = "expand"
const fieldsQueryParam = "fields"
// Deprecated: Use RequestInfo instead.
func RequestData(c echo.Context) *models.RequestInfo {
log.Println("RequestData(c) is deprecated and will be removed in the future! You can replace it with RequestInfo(c).")
return RequestInfo(c)
}
// RequestInfo exports cached common request data fields
// (query, body, logged auth state, etc.) from the provided context. // (query, body, logged auth state, etc.) from the provided context.
func RequestData(c echo.Context) *models.RequestData { func RequestInfo(c echo.Context) *models.RequestInfo {
// return cached to avoid copying the body multiple times // return cached to avoid copying the body multiple times
if v := c.Get(ContextRequestDataKey); v != nil { if v := c.Get(ContextRequestInfoKey); v != nil {
if data, ok := v.(*models.RequestData); ok { if data, ok := v.(*models.RequestInfo); ok {
// refresh auth state
data.AuthRecord, _ = c.Get(ContextAuthRecordKey).(*models.Record)
data.Admin, _ = c.Get(ContextAdminKey).(*models.Admin)
return data return data
} }
} }
result := &models.RequestData{ result := &models.RequestInfo{
Context: models.RequestInfoContextDefault,
Method: c.Request().Method, Method: c.Request().Method,
Query: map[string]any{}, Query: map[string]any{},
Data: map[string]any{}, Data: map[string]any{},
@@ -40,7 +55,7 @@ func RequestData(c echo.Context) *models.RequestData {
// ("X-Token" is converted to "x_token") // ("X-Token" is converted to "x_token")
for k, v := range c.Request().Header { for k, v := range c.Request().Header {
if len(v) > 0 { if len(v) > 0 {
result.Headers[strings.ToLower(strings.ReplaceAll(k, "-", "_"))] = v[0] result.Headers[inflector.Snakecase(k)] = v[0]
} }
} }
@@ -49,12 +64,24 @@ func RequestData(c echo.Context) *models.RequestData {
echo.BindQueryParams(c, &result.Query) echo.BindQueryParams(c, &result.Query)
rest.BindBody(c, &result.Data) rest.BindBody(c, &result.Data)
c.Set(ContextRequestDataKey, result) c.Set(ContextRequestInfoKey, result)
return result return result
} }
func RecordAuthResponse(app core.App, c echo.Context, authRecord *models.Record, meta any) error { // RecordAuthResponse writes standardised json record auth response
// into the specified request context.
func RecordAuthResponse(
app core.App,
c echo.Context,
authRecord *models.Record,
meta any,
finalizers ...func(token string) error,
) error {
if !authRecord.Verified() && authRecord.Collection().AuthOptions().OnlyVerified {
return NewForbiddenError("Please verify your email first.", nil)
}
token, tokenErr := tokens.NewRecordAuthToken(app, authRecord) token, tokenErr := tokens.NewRecordAuthToken(app, authRecord)
if tokenErr != nil { if tokenErr != nil {
return NewBadRequestError("Failed to create auth token.", tokenErr) return NewBadRequestError("Failed to create auth token.", tokenErr)
@@ -68,6 +95,10 @@ func RecordAuthResponse(app core.App, c echo.Context, authRecord *models.Record,
event.Meta = meta event.Meta = meta
return app.OnRecordAuthRequest().Trigger(event, func(e *core.RecordAuthEvent) error { return app.OnRecordAuthRequest().Trigger(event, func(e *core.RecordAuthEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
// allow always returning the email address of the authenticated account // allow always returning the email address of the authenticated account
e.Record.IgnoreEmailVisibility(true) e.Record.IgnoreEmailVisibility(true)
@@ -75,16 +106,16 @@ func RecordAuthResponse(app core.App, c echo.Context, authRecord *models.Record,
expands := strings.Split(c.QueryParam(expandQueryParam), ",") expands := strings.Split(c.QueryParam(expandQueryParam), ",")
if len(expands) > 0 { if len(expands) > 0 {
// create a copy of the cached request data and adjust it to the current auth record // create a copy of the cached request data and adjust it to the current auth record
requestData := *RequestData(e.HttpContext) requestInfo := *RequestInfo(e.HttpContext)
requestData.Admin = nil requestInfo.Admin = nil
requestData.AuthRecord = e.Record requestInfo.AuthRecord = e.Record
failed := app.Dao().ExpandRecord( failed := app.Dao().ExpandRecord(
e.Record, e.Record,
expands, expands,
expandFetch(app.Dao(), &requestData), expandFetch(app.Dao(), &requestInfo),
) )
if len(failed) > 0 && app.IsDebug() { if len(failed) > 0 {
log.Println("Failed to expand relations: ", failed) app.Logger().Debug("[RecordAuthResponse] Failed to expand relations", slog.Any("errors", failed))
} }
} }
@@ -97,6 +128,12 @@ func RecordAuthResponse(app core.App, c echo.Context, authRecord *models.Record,
result["meta"] = e.Meta result["meta"] = e.Meta
} }
for _, f := range finalizers {
if err := f(e.Token); err != nil {
return err
}
}
return e.HttpContext.JSON(http.StatusOK, result) return e.HttpContext.JSON(http.StatusOK, result)
}) })
} }
@@ -104,7 +141,7 @@ func RecordAuthResponse(app core.App, c echo.Context, authRecord *models.Record,
// EnrichRecord parses the request context and enrich the provided record: // EnrichRecord parses the request context and enrich the provided record:
// - expands relations (if defaultExpands and/or ?expand query param is set) // - expands relations (if defaultExpands and/or ?expand query param is set)
// - ensures that the emails of the auth record and its expanded auth relations // - ensures that the emails of the auth record and its expanded auth relations
// are visibe only for the current logged admin, record owner or record with manage access // are visible only for the current logged admin, record owner or record with manage access
func EnrichRecord(c echo.Context, dao *daos.Dao, record *models.Record, defaultExpands ...string) error { func EnrichRecord(c echo.Context, dao *daos.Dao, record *models.Record, defaultExpands ...string) error {
return EnrichRecords(c, dao, []*models.Record{record}, defaultExpands...) return EnrichRecords(c, dao, []*models.Record{record}, defaultExpands...)
} }
@@ -112,11 +149,11 @@ func EnrichRecord(c echo.Context, dao *daos.Dao, record *models.Record, defaultE
// EnrichRecords parses the request context and enriches the provided records: // EnrichRecords parses the request context and enriches the provided records:
// - expands relations (if defaultExpands and/or ?expand query param is set) // - expands relations (if defaultExpands and/or ?expand query param is set)
// - ensures that the emails of the auth records and their expanded auth relations // - ensures that the emails of the auth records and their expanded auth relations
// are visibe only for the current logged admin, record owner or record with manage access // are visible only for the current logged admin, record owner or record with manage access
func EnrichRecords(c echo.Context, dao *daos.Dao, records []*models.Record, defaultExpands ...string) error { func EnrichRecords(c echo.Context, dao *daos.Dao, records []*models.Record, defaultExpands ...string) error {
requestData := RequestData(c) requestInfo := RequestInfo(c)
if err := autoIgnoreAuthRecordsEmailVisibility(dao, records, requestData); err != nil { if err := autoIgnoreAuthRecordsEmailVisibility(dao, records, requestInfo); err != nil {
return fmt.Errorf("Failed to resolve email visibility: %w", err) return fmt.Errorf("Failed to resolve email visibility: %w", err)
} }
@@ -128,7 +165,7 @@ func EnrichRecords(c echo.Context, dao *daos.Dao, records []*models.Record, defa
return nil // nothing to expand return nil // nothing to expand
} }
errs := dao.ExpandRecords(records, expands, expandFetch(dao, requestData)) errs := dao.ExpandRecords(records, expands, expandFetch(dao, requestInfo))
if len(errs) > 0 { if len(errs) > 0 {
return fmt.Errorf("Failed to expand: %v", errs) return fmt.Errorf("Failed to expand: %v", errs)
} }
@@ -139,11 +176,11 @@ func EnrichRecords(c echo.Context, dao *daos.Dao, records []*models.Record, defa
// expandFetch is the records fetch function that is used to expand related records. // expandFetch is the records fetch function that is used to expand related records.
func expandFetch( func expandFetch(
dao *daos.Dao, dao *daos.Dao,
requestData *models.RequestData, requestInfo *models.RequestInfo,
) daos.ExpandFetchFunc { ) daos.ExpandFetchFunc {
return func(relCollection *models.Collection, relIds []string) ([]*models.Record, error) { return func(relCollection *models.Collection, relIds []string) ([]*models.Record, error) {
records, err := dao.FindRecordsByIds(relCollection.Id, relIds, func(q *dbx.SelectQuery) error { records, err := dao.FindRecordsByIds(relCollection.Id, relIds, func(q *dbx.SelectQuery) error {
if requestData.Admin != nil { if requestInfo.Admin != nil {
return nil // admins can access everything return nil // admins can access everything
} }
@@ -152,7 +189,7 @@ func expandFetch(
} }
if *relCollection.ViewRule != "" { if *relCollection.ViewRule != "" {
resolver := resolvers.NewRecordFieldResolver(dao, relCollection, requestData, true) resolver := resolvers.NewRecordFieldResolver(dao, relCollection, requestInfo, true)
expr, err := search.FilterData(*(relCollection.ViewRule)).BuildExpr(resolver) expr, err := search.FilterData(*(relCollection.ViewRule)).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -165,7 +202,7 @@ func expandFetch(
}) })
if err == nil && len(records) > 0 { if err == nil && len(records) > 0 {
autoIgnoreAuthRecordsEmailVisibility(dao, records, requestData) autoIgnoreAuthRecordsEmailVisibility(dao, records, requestInfo)
} }
return records, err return records, err
@@ -179,13 +216,13 @@ func expandFetch(
func autoIgnoreAuthRecordsEmailVisibility( func autoIgnoreAuthRecordsEmailVisibility(
dao *daos.Dao, dao *daos.Dao,
records []*models.Record, records []*models.Record,
requestData *models.RequestData, requestInfo *models.RequestInfo,
) error { ) error {
if len(records) == 0 || !records[0].Collection().IsAuth() { if len(records) == 0 || !records[0].Collection().IsAuth() {
return nil // nothing to check return nil // nothing to check
} }
if requestData.Admin != nil { if requestInfo.Admin != nil {
for _, rec := range records { for _, rec := range records {
rec.IgnoreEmailVisibility(true) rec.IgnoreEmailVisibility(true)
} }
@@ -201,8 +238,8 @@ func autoIgnoreAuthRecordsEmailVisibility(
recordIds[i] = rec.Id recordIds[i] = rec.Id
} }
if requestData != nil && requestData.AuthRecord != nil && mappedRecords[requestData.AuthRecord.Id] != nil { if requestInfo != nil && requestInfo.AuthRecord != nil && mappedRecords[requestInfo.AuthRecord.Id] != nil {
mappedRecords[requestData.AuthRecord.Id].IgnoreEmailVisibility(true) mappedRecords[requestInfo.AuthRecord.Id].IgnoreEmailVisibility(true)
} }
authOptions := collection.AuthOptions() authOptions := collection.AuthOptions()
@@ -218,7 +255,7 @@ func autoIgnoreAuthRecordsEmailVisibility(
Select(dao.DB().QuoteSimpleColumnName(collection.Name) + ".id"). Select(dao.DB().QuoteSimpleColumnName(collection.Name) + ".id").
AndWhere(dbx.In(dao.DB().QuoteSimpleColumnName(collection.Name)+".id", recordIds...)) AndWhere(dbx.In(dao.DB().QuoteSimpleColumnName(collection.Name)+".id", recordIds...))
resolver := resolvers.NewRecordFieldResolver(dao, collection, requestData, true) resolver := resolvers.NewRecordFieldResolver(dao, collection, requestInfo, true)
expr, err := search.FilterData(*authOptions.ManageRule).BuildExpr(resolver) expr, err := search.FilterData(*authOptions.ManageRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -247,7 +284,7 @@ func autoIgnoreAuthRecordsEmailVisibility(
func hasAuthManageAccess( func hasAuthManageAccess(
dao *daos.Dao, dao *daos.Dao,
record *models.Record, record *models.Record,
requestData *models.RequestData, requestInfo *models.RequestInfo,
) bool { ) bool {
if !record.Collection().IsAuth() { if !record.Collection().IsAuth() {
return false return false
@@ -259,12 +296,12 @@ func hasAuthManageAccess(
return false // only for admins (manageRule can't be empty) return false // only for admins (manageRule can't be empty)
} }
if requestData == nil || requestData.AuthRecord == nil { if requestInfo == nil || requestInfo.AuthRecord == nil {
return false // no auth record return false // no auth record
} }
ruleFunc := func(q *dbx.SelectQuery) error { ruleFunc := func(q *dbx.SelectQuery) error {
resolver := resolvers.NewRecordFieldResolver(dao, record.Collection(), requestData, true) resolver := resolvers.NewRecordFieldResolver(dao, record.Collection(), requestInfo, true)
expr, err := search.FilterData(*manageRule).BuildExpr(resolver) expr, err := search.FilterData(*manageRule).BuildExpr(resolver)
if err != nil { if err != nil {
return err return err
@@ -278,3 +315,31 @@ func hasAuthManageAccess(
return findErr == nil return findErr == nil
} }
var ruleQueryParams = []string{search.FilterQueryParam, search.SortQueryParam}
var adminOnlyRuleFields = []string{"@collection.", "@request."}
// @todo consider moving the rules check to the RecordFieldResolver.
//
// checkForAdminOnlyRuleFields loosely checks and returns an error if
// the provided RequestInfo contains rule fields that only the admin can use.
func checkForAdminOnlyRuleFields(requestInfo *models.RequestInfo) error {
if requestInfo.Admin != nil || len(requestInfo.Query) == 0 {
return nil // admin or nothing to check
}
for _, param := range ruleQueryParams {
v, _ := requestInfo.Query[param].(string)
if v == "" {
continue
}
for _, field := range adminOnlyRuleFields {
if strings.Contains(v, field) {
return NewForbiddenError("Only admins can filter by "+field, nil)
}
}
}
return nil
}
+19 -3
View File
@@ -13,7 +13,9 @@ import (
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
) )
func TestRequestData(t *testing.T) { func TestRequestInfo(t *testing.T) {
t.Parallel()
e := echo.New() e := echo.New()
req := httptest.NewRequest(http.MethodPost, "/?test=123", strings.NewReader(`{"test":456}`)) req := httptest.NewRequest(http.MethodPost, "/?test=123", strings.NewReader(`{"test":456}`))
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON) req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
@@ -29,10 +31,10 @@ func TestRequestData(t *testing.T) {
dummyAdmin.Id = "id2" dummyAdmin.Id = "id2"
c.Set(apis.ContextAdminKey, dummyAdmin) c.Set(apis.ContextAdminKey, dummyAdmin)
result := apis.RequestData(c) result := apis.RequestInfo(c)
if result == nil { if result == nil {
t.Fatal("Expected *models.RequestData instance, got nil") t.Fatal("Expected *models.RequestInfo instance, got nil")
} }
if result.Method != http.MethodPost { if result.Method != http.MethodPost {
@@ -67,6 +69,8 @@ func TestRequestData(t *testing.T) {
} }
func TestRecordAuthResponse(t *testing.T) { func TestRecordAuthResponse(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -83,6 +87,11 @@ func TestRecordAuthResponse(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
unverifiedAuthRecord, err := app.Dao().FindRecordById("clients", "o1y0dd0spd786md")
if err != nil {
t.Fatal(err)
}
scenarios := []struct { scenarios := []struct {
name string name string
record *models.Record record *models.Record
@@ -97,6 +106,11 @@ func TestRecordAuthResponse(t *testing.T) {
record: nonAuthRecord, record: nonAuthRecord,
expectError: true, expectError: true,
}, },
{
name: "valid auth record but with unverified email in onlyVerified collection",
record: unverifiedAuthRecord,
expectError: true,
},
{ {
name: "valid auth record - without meta", name: "valid auth record - without meta",
record: authRecord, record: authRecord,
@@ -179,6 +193,8 @@ func TestRecordAuthResponse(t *testing.T) {
} }
func TestEnrichRecords(t *testing.T) { func TestEnrichRecords(t *testing.T) {
t.Parallel()
e := echo.New() e := echo.New()
req := httptest.NewRequest(http.MethodGet, "/?expand=rel_many", nil) req := httptest.NewRequest(http.MethodGet, "/?expand=rel_many", nil)
req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON) req.Header.Set(echo.HeaderContentType, echo.MIMEApplicationJSON)
+141 -38
View File
@@ -8,41 +8,64 @@ import (
"net/http" "net/http"
"path/filepath" "path/filepath"
"strings" "strings"
"sync"
"time" "time"
"github.com/fatih/color" "github.com/fatih/color"
"github.com/labstack/echo/v5"
"github.com/labstack/echo/v5/middleware" "github.com/labstack/echo/v5/middleware"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/migrations" "github.com/pocketbase/pocketbase/migrations"
"github.com/pocketbase/pocketbase/migrations/logs" "github.com/pocketbase/pocketbase/migrations/logs"
"github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/migrate" "github.com/pocketbase/pocketbase/tools/migrate"
"golang.org/x/crypto/acme" "golang.org/x/crypto/acme"
"golang.org/x/crypto/acme/autocert" "golang.org/x/crypto/acme/autocert"
) )
// ServeOptions defines an optional struct for apis.Serve(). // ServeConfig defines a configuration struct for apis.Serve().
type ServeOptions struct { type ServeConfig struct {
// ShowStartBanner indicates whether to show or hide the server start console message.
ShowStartBanner bool ShowStartBanner bool
HttpAddr string
HttpsAddr string // HttpAddr is the TCP address to listen for the HTTP server (eg. `127.0.0.1:80`).
AllowedOrigins []string // optional list of CORS origins (default to "*") HttpAddr string
BeforeServeFunc func(server *http.Server) error
// HttpsAddr is the TCP address to listen for the HTTPS server (eg. `127.0.0.1:443`).
HttpsAddr string
// Optional domains list to use when issuing the TLS certificate.
//
// If not set, the host from the bound server address will be used.
//
// For convenience, for each "non-www" domain a "www" entry and
// redirect will be automatically added.
CertificateDomains []string
// AllowedOrigins is an optional list of CORS origins (default to "*").
AllowedOrigins []string
} }
// Serve starts a new app web server. // Serve starts a new app web server.
func Serve(app core.App, options *ServeOptions) error { //
if options == nil { // NB! The app should be bootstrapped before starting the web server.
options = &ServeOptions{} //
} // Example:
//
if len(options.AllowedOrigins) == 0 { // app.Bootstrap()
options.AllowedOrigins = []string{"*"} // apis.Serve(app, apis.ServeConfig{
// HttpAddr: "127.0.0.1:8080",
// ShowStartBanner: false,
// })
func Serve(app core.App, config ServeConfig) (*http.Server, error) {
if len(config.AllowedOrigins) == 0 {
config.AllowedOrigins = []string{"*"}
} }
// ensure that the latest migrations are applied before starting the server // ensure that the latest migrations are applied before starting the server
if err := runMigrations(app); err != nil { if err := runMigrations(app); err != nil {
return err return nil, err
} }
// reload app settings in case a new default value was set with a migration // reload app settings in case a new default value was set with a migration
@@ -56,33 +79,75 @@ func Serve(app core.App, options *ServeOptions) error {
router, err := InitApi(app) router, err := InitApi(app)
if err != nil { if err != nil {
return err return nil, err
} }
// configure cors // configure cors
router.Use(middleware.CORSWithConfig(middleware.CORSConfig{ router.Use(middleware.CORSWithConfig(middleware.CORSConfig{
Skipper: middleware.DefaultSkipper, Skipper: middleware.DefaultSkipper,
AllowOrigins: options.AllowedOrigins, AllowOrigins: config.AllowedOrigins,
AllowMethods: []string{http.MethodGet, http.MethodHead, http.MethodPut, http.MethodPatch, http.MethodPost, http.MethodDelete}, AllowMethods: []string{http.MethodGet, http.MethodHead, http.MethodPut, http.MethodPatch, http.MethodPost, http.MethodDelete},
})) }))
// start http server // start http server
// --- // ---
mainAddr := options.HttpAddr mainAddr := config.HttpAddr
if options.HttpsAddr != "" { if config.HttpsAddr != "" {
mainAddr = options.HttpsAddr mainAddr = config.HttpsAddr
} }
mainHost, _, _ := net.SplitHostPort(mainAddr) var wwwRedirects []string
certManager := autocert.Manager{ // extract the host names for the certificate host policy
hostNames := config.CertificateDomains
if len(hostNames) == 0 {
host, _, _ := net.SplitHostPort(mainAddr)
hostNames = append(hostNames, host)
}
for _, host := range hostNames {
if strings.HasPrefix(host, "www.") {
continue // explicitly set www host
}
wwwHost := "www." + host
if !list.ExistInSlice(wwwHost, hostNames) {
hostNames = append(hostNames, wwwHost)
wwwRedirects = append(wwwRedirects, wwwHost)
}
}
// implicit www->non-www redirect(s)
if len(wwwRedirects) > 0 {
router.Pre(func(next echo.HandlerFunc) echo.HandlerFunc {
return func(c echo.Context) error {
host := c.Request().Host
if strings.HasPrefix(host, "www.") && list.ExistInSlice(host, wwwRedirects) {
return c.Redirect(
http.StatusTemporaryRedirect,
(c.Scheme() + "://" + host[4:] + c.Request().RequestURI),
)
}
return next(c)
}
})
}
certManager := &autocert.Manager{
Prompt: autocert.AcceptTOS, Prompt: autocert.AcceptTOS,
Cache: autocert.DirCache(filepath.Join(app.DataDir(), ".autocert_cache")), Cache: autocert.DirCache(filepath.Join(app.DataDir(), ".autocert_cache")),
HostPolicy: autocert.HostWhitelist(mainHost, "www."+mainHost), HostPolicy: autocert.HostWhitelist(hostNames...),
} }
serverConfig := &http.Server{ // base request context used for cancelling long running requests
// like the SSE connections
baseCtx, cancelBaseCtx := context.WithCancel(context.Background())
defer cancelBaseCtx()
server := &http.Server{
TLSConfig: &tls.Config{ TLSConfig: &tls.Config{
MinVersion: tls.VersionTLS12,
GetCertificate: certManager.GetCertificate, GetCertificate: certManager.GetCertificate,
NextProtos: []string{acme.ALPNProto}, NextProtos: []string{acme.ALPNProto},
}, },
@@ -91,18 +156,31 @@ func Serve(app core.App, options *ServeOptions) error {
// WriteTimeout: 60 * time.Second, // breaks sse! // WriteTimeout: 60 * time.Second, // breaks sse!
Handler: router, Handler: router,
Addr: mainAddr, Addr: mainAddr,
BaseContext: func(l net.Listener) context.Context {
return baseCtx
},
} }
if options.BeforeServeFunc != nil { serveEvent := &core.ServeEvent{
if err := options.BeforeServeFunc(serverConfig); err != nil { App: app,
return err Router: router,
} Server: server,
CertManager: certManager,
}
if err := app.OnBeforeServe().Trigger(serveEvent); err != nil {
return nil, err
} }
if options.ShowStartBanner { if config.ShowStartBanner {
schema := "http" schema := "http"
if options.HttpsAddr != "" { addr := server.Addr
if config.HttpsAddr != "" {
schema = "https" schema = "https"
if len(config.CertificateDomains) > 0 {
addr = config.CertificateDomains[0]
}
} }
date := new(strings.Builder) date := new(strings.Builder)
@@ -112,34 +190,59 @@ func Serve(app core.App, options *ServeOptions) error {
bold.Printf( bold.Printf(
"%s Server started at %s\n", "%s Server started at %s\n",
strings.TrimSpace(date.String()), strings.TrimSpace(date.String()),
color.CyanString("%s://%s", schema, serverConfig.Addr), color.CyanString("%s://%s", schema, addr),
) )
regular := color.New() regular := color.New()
regular.Printf(" ➜ REST API: %s\n", color.CyanString("%s://%s/api/", schema, serverConfig.Addr)) regular.Printf("├─ REST API: %s\n", color.CyanString("%s://%s/api/", schema, addr))
regular.Printf(" ➜ Admin UI: %s\n", color.CyanString("%s://%s/_/", schema, serverConfig.Addr)) regular.Printf("└─ Admin UI: %s\n", color.CyanString("%s://%s/_/", schema, addr))
} }
// WaitGroup to block until server.ShutDown() returns because Serve and similar methods exit immediately.
// Note that the WaitGroup would not do anything if the app.OnTerminate() hook isn't triggered.
var wg sync.WaitGroup
// try to gracefully shutdown the server on app termination // try to gracefully shutdown the server on app termination
app.OnTerminate().Add(func(e *core.TerminateEvent) error { app.OnTerminate().Add(func(e *core.TerminateEvent) error {
cancelBaseCtx()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel() defer cancel()
serverConfig.Shutdown(ctx)
wg.Add(1)
server.Shutdown(ctx)
if e.IsRestart {
// wait for execve and other handlers up to 5 seconds before exit
time.AfterFunc(5*time.Second, func() {
wg.Done()
})
} else {
wg.Done()
}
return nil return nil
}) })
// wait for the graceful shutdown to complete before exit
defer wg.Wait()
// ---
// @todo consider removing the server return value because it is
// not really useful when combined with the blocking serve calls
// ---
// start HTTPS server // start HTTPS server
if options.HttpsAddr != "" { if config.HttpsAddr != "" {
// if httpAddr is set, start an HTTP server to redirect the traffic to the HTTPS version // if httpAddr is set, start an HTTP server to redirect the traffic to the HTTPS version
if options.HttpAddr != "" { if config.HttpAddr != "" {
go http.ListenAndServe(options.HttpAddr, certManager.HTTPHandler(nil)) go http.ListenAndServe(config.HttpAddr, certManager.HTTPHandler(nil))
} }
return serverConfig.ListenAndServeTLS("", "") return server, server.ListenAndServeTLS("", "")
} }
// OR start HTTP server // OR start HTTP server
return serverConfig.ListenAndServe() return server, server.ListenAndServe()
} }
type migrationsConnection struct { type migrationsConnection struct {
+16 -15
View File
@@ -1,7 +1,6 @@
package apis package apis
import ( import (
"log"
"net/http" "net/http"
validation "github.com/go-ozzo/ozzo-validation/v4" validation "github.com/go-ozzo/ozzo-validation/v4"
@@ -38,6 +37,10 @@ func (api *settingsApi) list(c echo.Context) error {
event.RedactedSettings = settings event.RedactedSettings = settings
return api.app.OnSettingsListRequest().Trigger(event, func(e *core.SettingsListEvent) error { return api.app.OnSettingsListRequest().Trigger(event, func(e *core.SettingsListEvent) error {
if e.HttpContext.Response().Committed {
return nil
}
return e.HttpContext.JSON(http.StatusOK, e.RedactedSettings) return e.HttpContext.JSON(http.StatusOK, e.RedactedSettings)
}) })
} }
@@ -55,7 +58,7 @@ func (api *settingsApi) set(c echo.Context) error {
event.OldSettings = api.app.Settings() event.OldSettings = api.app.Settings()
// update the settings // update the settings
submitErr := form.Submit(func(next forms.InterceptorNextFunc[*settings.Settings]) forms.InterceptorNextFunc[*settings.Settings] { return form.Submit(func(next forms.InterceptorNextFunc[*settings.Settings]) forms.InterceptorNextFunc[*settings.Settings] {
return func(s *settings.Settings) error { return func(s *settings.Settings) error {
event.NewSettings = s event.NewSettings = s
@@ -64,23 +67,21 @@ func (api *settingsApi) set(c echo.Context) error {
return NewBadRequestError("An error occurred while submitting the form.", err) return NewBadRequestError("An error occurred while submitting the form.", err)
} }
redactedSettings, err := api.app.Settings().RedactClone() return api.app.OnSettingsAfterUpdateRequest().Trigger(event, func(e *core.SettingsUpdateEvent) error {
if err != nil { if e.HttpContext.Response().Committed {
return NewBadRequestError("", err) return nil
} }
return e.HttpContext.JSON(http.StatusOK, redactedSettings) redactedSettings, err := api.app.Settings().RedactClone()
if err != nil {
return NewBadRequestError("", err)
}
return e.HttpContext.JSON(http.StatusOK, redactedSettings)
})
}) })
} }
}) })
if submitErr == nil {
if err := api.app.OnSettingsAfterUpdateRequest().Trigger(event); err != nil && api.app.IsDebug() {
log.Println(err)
}
}
return submitErr
} }
func (api *settingsApi) testS3(c echo.Context) error { func (api *settingsApi) testS3(c echo.Context) error {
+58 -3
View File
@@ -6,16 +6,20 @@ import (
"crypto/rand" "crypto/rand"
"crypto/x509" "crypto/x509"
"encoding/pem" "encoding/pem"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"strings" "strings"
"testing" "testing"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
) )
func TestSettingsList(t *testing.T) { func TestSettingsList(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -75,6 +79,13 @@ func TestSettingsList(t *testing.T) {
`"oidc2Auth":{`, `"oidc2Auth":{`,
`"oidc3Auth":{`, `"oidc3Auth":{`,
`"appleAuth":{`, `"appleAuth":{`,
`"instagramAuth":{`,
`"vkAuth":{`,
`"yandexAuth":{`,
`"patreonAuth":{`,
`"mailcowAuth":{`,
`"bitbucketAuth":{`,
`"planningcenterAuth":{`,
`"secret":"******"`, `"secret":"******"`,
`"clientSecret":"******"`, `"clientSecret":"******"`,
}, },
@@ -90,6 +101,8 @@ func TestSettingsList(t *testing.T) {
} }
func TestSettingsSet(t *testing.T) { func TestSettingsSet(t *testing.T) {
t.Parallel()
validData := `{"meta":{"appName":"update_test"}}` validData := `{"meta":{"appName":"update_test"}}`
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
@@ -153,6 +166,13 @@ func TestSettingsSet(t *testing.T) {
`"oidc2Auth":{`, `"oidc2Auth":{`,
`"oidc3Auth":{`, `"oidc3Auth":{`,
`"appleAuth":{`, `"appleAuth":{`,
`"instagramAuth":{`,
`"vkAuth":{`,
`"yandexAuth":{`,
`"patreonAuth":{`,
`"mailcowAuth":{`,
`"bitbucketAuth":{`,
`"planningcenterAuth":{`,
`"secret":"******"`, `"secret":"******"`,
`"clientSecret":"******"`, `"clientSecret":"******"`,
`"appName":"acme_test"`, `"appName":"acme_test"`,
@@ -220,6 +240,13 @@ func TestSettingsSet(t *testing.T) {
`"oidc2Auth":{`, `"oidc2Auth":{`,
`"oidc3Auth":{`, `"oidc3Auth":{`,
`"appleAuth":{`, `"appleAuth":{`,
`"instagramAuth":{`,
`"vkAuth":{`,
`"yandexAuth":{`,
`"patreonAuth":{`,
`"mailcowAuth":{`,
`"bitbucketAuth":{`,
`"planningcenterAuth":{`,
`"secret":"******"`, `"secret":"******"`,
`"clientSecret":"******"`, `"clientSecret":"******"`,
`"appName":"update_test"`, `"appName":"update_test"`,
@@ -231,6 +258,28 @@ func TestSettingsSet(t *testing.T) {
"OnSettingsAfterUpdateRequest": 1, "OnSettingsAfterUpdateRequest": 1,
}, },
}, },
{
Name: "OnSettingsAfterUpdateRequest error response",
Method: http.MethodPatch,
Url: "/api/settings",
Body: strings.NewReader(validData),
RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
},
BeforeTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) {
app.OnSettingsAfterUpdateRequest().Add(func(e *core.SettingsUpdateEvent) error {
return errors.New("error")
})
},
ExpectedStatus: 400,
ExpectedContent: []string{`"data":{}`},
ExpectedEvents: map[string]int{
"OnModelBeforeUpdate": 1,
"OnModelAfterUpdate": 1,
"OnSettingsBeforeUpdateRequest": 1,
"OnSettingsAfterUpdateRequest": 1,
},
},
} }
for _, scenario := range scenarios { for _, scenario := range scenarios {
@@ -239,6 +288,8 @@ func TestSettingsSet(t *testing.T) {
} }
func TestSettingsTestS3(t *testing.T) { func TestSettingsTestS3(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -305,6 +356,8 @@ func TestSettingsTestS3(t *testing.T) {
} }
func TestSettingsTestEmail(t *testing.T) { func TestSettingsTestEmail(t *testing.T) {
t.Parallel()
scenarios := []tests.ApiScenario{ scenarios := []tests.ApiScenario{
{ {
Name: "unauthorized", Name: "unauthorized",
@@ -367,7 +420,7 @@ func TestSettingsTestEmail(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if app.TestMailer.TotalSend != 1 { if app.TestMailer.TotalSend != 1 {
t.Fatalf("[verification] Expected 1 sent email, got %d", app.TestMailer.TotalSend) t.Fatalf("[verification] Expected 1 sent email, got %d", app.TestMailer.TotalSend)
} }
@@ -402,7 +455,7 @@ func TestSettingsTestEmail(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if app.TestMailer.TotalSend != 1 { if app.TestMailer.TotalSend != 1 {
t.Fatalf("[password-reset] Expected 1 sent email, got %d", app.TestMailer.TotalSend) t.Fatalf("[password-reset] Expected 1 sent email, got %d", app.TestMailer.TotalSend)
} }
@@ -437,7 +490,7 @@ func TestSettingsTestEmail(t *testing.T) {
RequestHeaders: map[string]string{ RequestHeaders: map[string]string{
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8", "Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJpZCI6InN5d2JoZWNuaDQ2cmhtMCIsInR5cGUiOiJhZG1pbiIsImV4cCI6MjIwODk4NTI2MX0.M1m--VOqGyv0d23eeUc0r9xE8ZzHaYVmVFw1VZW6gT8",
}, },
AfterTestFunc: func(t *testing.T, app *tests.TestApp, e *echo.Echo) { AfterTestFunc: func(t *testing.T, app *tests.TestApp, res *http.Response) {
if app.TestMailer.TotalSend != 1 { if app.TestMailer.TotalSend != 1 {
t.Fatalf("[email-change] Expected 1 sent email, got %d", app.TestMailer.TotalSend) t.Fatalf("[email-change] Expected 1 sent email, got %d", app.TestMailer.TotalSend)
} }
@@ -469,6 +522,8 @@ func TestSettingsTestEmail(t *testing.T) {
} }
func TestGenerateAppleClientSecret(t *testing.T) { func TestGenerateAppleClientSecret(t *testing.T) {
t.Parallel()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
+24 -18
View File
@@ -28,12 +28,10 @@ func NewAdminCommand(app core.App) *cobra.Command {
func adminCreateCommand(app core.App) *cobra.Command { func adminCreateCommand(app core.App) *cobra.Command {
command := &cobra.Command{ command := &cobra.Command{
Use: "create", Use: "create",
Example: "admin create test@example.com 1234567890", Example: "admin create test@example.com 1234567890",
Short: "Creates a new admin account", Short: "Creates a new admin account",
// prevents printing the error log twice SilenceUsage: true,
SilenceErrors: true,
SilenceUsage: true,
RunE: func(command *cobra.Command, args []string) error { RunE: func(command *cobra.Command, args []string) error {
if len(args) != 2 { if len(args) != 2 {
return errors.New("Missing email and password arguments.") return errors.New("Missing email and password arguments.")
@@ -51,6 +49,10 @@ func adminCreateCommand(app core.App) *cobra.Command {
admin.Email = args[0] admin.Email = args[0]
admin.SetPassword(args[1]) admin.SetPassword(args[1])
if !app.Dao().HasTable(admin.TableName()) {
return errors.New("Migration are not initialized yet. Please run 'migrate up' and try again.")
}
if err := app.Dao().SaveAdmin(admin); err != nil { if err := app.Dao().SaveAdmin(admin); err != nil {
return fmt.Errorf("Failed to create new admin account: %v", err) return fmt.Errorf("Failed to create new admin account: %v", err)
} }
@@ -65,12 +67,10 @@ func adminCreateCommand(app core.App) *cobra.Command {
func adminUpdateCommand(app core.App) *cobra.Command { func adminUpdateCommand(app core.App) *cobra.Command {
command := &cobra.Command{ command := &cobra.Command{
Use: "update", Use: "update",
Example: "admin update test@example.com 1234567890", Example: "admin update test@example.com 1234567890",
Short: "Changes the password of a single admin account", Short: "Changes the password of a single admin account",
// prevents printing the error log twice SilenceUsage: true,
SilenceErrors: true,
SilenceUsage: true,
RunE: func(command *cobra.Command, args []string) error { RunE: func(command *cobra.Command, args []string) error {
if len(args) != 2 { if len(args) != 2 {
return errors.New("Missing email and password arguments.") return errors.New("Missing email and password arguments.")
@@ -84,6 +84,10 @@ func adminUpdateCommand(app core.App) *cobra.Command {
return errors.New("The new password must be at least 8 chars long.") return errors.New("The new password must be at least 8 chars long.")
} }
if !app.Dao().HasTable((&models.Admin{}).TableName()) {
return errors.New("Migration are not initialized yet. Please run 'migrate up' and try again.")
}
admin, err := app.Dao().FindAdminByEmail(args[0]) admin, err := app.Dao().FindAdminByEmail(args[0])
if err != nil { if err != nil {
return fmt.Errorf("Admin with email %s doesn't exist.", args[0]) return fmt.Errorf("Admin with email %s doesn't exist.", args[0])
@@ -105,17 +109,19 @@ func adminUpdateCommand(app core.App) *cobra.Command {
func adminDeleteCommand(app core.App) *cobra.Command { func adminDeleteCommand(app core.App) *cobra.Command {
command := &cobra.Command{ command := &cobra.Command{
Use: "delete", Use: "delete",
Example: "admin delete test@example.com", Example: "admin delete test@example.com",
Short: "Deletes an existing admin account", Short: "Deletes an existing admin account",
// prevents printing the error log twice SilenceUsage: true,
SilenceErrors: true,
SilenceUsage: true,
RunE: func(command *cobra.Command, args []string) error { RunE: func(command *cobra.Command, args []string) error {
if len(args) == 0 || args[0] == "" || is.EmailFormat.Validate(args[0]) != nil { if len(args) == 0 || args[0] == "" || is.EmailFormat.Validate(args[0]) != nil {
return errors.New("Invalid or missing email address.") return errors.New("Invalid or missing email address.")
} }
if !app.Dao().HasTable((&models.Admin{}).TableName()) {
return errors.New("Migration are not initialized yet. Please run 'migrate up' and try again.")
}
admin, err := app.Dao().FindAdminByEmail(args[0]) admin, err := app.Dao().FindAdminByEmail(args[0])
if err != nil { if err != nil {
color.Yellow("Admin %s is already deleted.", args[0]) color.Yellow("Admin %s is already deleted.", args[0])
+2
View File
@@ -8,6 +8,8 @@ import (
) )
func TestAdminCreateCommand(t *testing.T) { func TestAdminCreateCommand(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+33 -14
View File
@@ -1,7 +1,7 @@
package cmd package cmd
import ( import (
"log" "errors"
"net/http" "net/http"
"github.com/pocketbase/pocketbase/apis" "github.com/pocketbase/pocketbase/apis"
@@ -17,19 +17,38 @@ func NewServeCommand(app core.App, showStartBanner bool) *cobra.Command {
var httpsAddr string var httpsAddr string
command := &cobra.Command{ command := &cobra.Command{
Use: "serve", Use: "serve [domain(s)]",
Short: "Starts the web server (default to 127.0.0.1:8090)", Args: cobra.ArbitraryArgs,
Run: func(command *cobra.Command, args []string) { Short: "Starts the web server (default to 127.0.0.1:8090 if no domain is specified)",
err := apis.Serve(app, &apis.ServeOptions{ SilenceUsage: true,
HttpAddr: httpAddr, RunE: func(command *cobra.Command, args []string) error {
HttpsAddr: httpsAddr, // set default listener addresses if at least one domain is specified
ShowStartBanner: showStartBanner, if len(args) > 0 {
AllowedOrigins: allowedOrigins, if httpAddr == "" {
httpAddr = "0.0.0.0:80"
}
if httpsAddr == "" {
httpsAddr = "0.0.0.0:443"
}
} else {
if httpAddr == "" {
httpAddr = "127.0.0.1:8090"
}
}
_, err := apis.Serve(app, apis.ServeConfig{
HttpAddr: httpAddr,
HttpsAddr: httpsAddr,
ShowStartBanner: showStartBanner,
AllowedOrigins: allowedOrigins,
CertificateDomains: args,
}) })
if err != http.ErrServerClosed { if errors.Is(err, http.ErrServerClosed) {
log.Fatalln(err) return nil
} }
return err
}, },
} }
@@ -43,15 +62,15 @@ func NewServeCommand(app core.App, showStartBanner bool) *cobra.Command {
command.PersistentFlags().StringVar( command.PersistentFlags().StringVar(
&httpAddr, &httpAddr,
"http", "http",
"127.0.0.1:8090", "",
"api HTTP server address", "TCP address to listen for the HTTP server\n(if domain args are specified - default to 0.0.0.0:80, otherwise - default to 127.0.0.1:8090)",
) )
command.PersistentFlags().StringVar( command.PersistentFlags().StringVar(
&httpsAddr, &httpsAddr,
"https", "https",
"", "",
"api HTTPS server address (auto TLS via Let's Encrypt)\nthe incoming --http address traffic also will be redirected to this address", "TCP address to listen for the HTTPS server\n(if domain args are specified - default to 0.0.0.0:443, otherwise - default to empty string, aka. no TLS)\nThe incoming HTTP traffic also will be auto redirected to the HTTPS version",
) )
return command return command
+88 -91
View File
@@ -5,6 +5,7 @@ package core
import ( import (
"context" "context"
"log/slog"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
@@ -48,6 +49,9 @@ type App interface {
// the users table from LogsDao will result in error. // the users table from LogsDao will result in error.
LogsDao() *daos.Dao LogsDao() *daos.Dao
// Logger returns the active app logger.
Logger() *slog.Logger
// DataDir returns the app data directory path. // DataDir returns the app data directory path.
DataDir() string DataDir() string
@@ -55,16 +59,18 @@ type App interface {
// (used for settings encryption). // (used for settings encryption).
EncryptionEnv() string EncryptionEnv() string
// IsDebug returns whether the app is in debug mode // IsDev returns whether the app is in dev mode.
// (showing more detailed error logs, executed sql statements, etc.). IsDev() bool
IsDebug() bool
// Settings returns the loaded app settings. // Settings returns the loaded app settings.
Settings() *settings.Settings Settings() *settings.Settings
// Cache returns the app internal cache store. // Deprecated: Use app.Store() instead.
Cache() *store.Store[any] Cache() *store.Store[any]
// Store returns the app runtime store.
Store() *store.Store[any]
// SubscriptionsBroker returns the app realtime subscriptions broker instance. // SubscriptionsBroker returns the app realtime subscriptions broker instance.
SubscriptionsBroker() *subscriptions.Broker SubscriptionsBroker() *subscriptions.Broker
@@ -74,14 +80,14 @@ type App interface {
// NewFilesystem creates and returns a configured filesystem.System instance // NewFilesystem creates and returns a configured filesystem.System instance
// for managing regular app files (eg. collection uploads). // for managing regular app files (eg. collection uploads).
// //
// NB! Make sure to call `Close()` on the returned result // NB! Make sure to call Close() on the returned result
// after you are done working with it. // after you are done working with it.
NewFilesystem() (*filesystem.System, error) NewFilesystem() (*filesystem.System, error)
// NewBackupsFilesystem creates and returns a configured filesystem.System instance // NewBackupsFilesystem creates and returns a configured filesystem.System instance
// for managing app backups. // for managing app backups.
// //
// NB! Make sure to call `Close()` on the returned result // NB! Make sure to call Close() on the returned result
// after you are done working with it. // after you are done working with it.
NewBackupsFilesystem() (*filesystem.System, error) NewBackupsFilesystem() (*filesystem.System, error)
@@ -131,21 +137,21 @@ type App interface {
// App event hooks // App event hooks
// --------------------------------------------------------------- // ---------------------------------------------------------------
// OnBeforeBootstrap hook is triggered before initializing the base // OnBeforeBootstrap hook is triggered before initializing the main
// application resources (eg. before db open and initial settings load). // application resources (eg. before db open and initial settings load).
OnBeforeBootstrap() *hook.Hook[*BootstrapEvent] OnBeforeBootstrap() *hook.Hook[*BootstrapEvent]
// OnAfterBootstrap hook is triggered after initializing the base // OnAfterBootstrap hook is triggered after initializing the main
// application resources (eg. after db open and initial settings load). // application resources (eg. after db open and initial settings load).
OnAfterBootstrap() *hook.Hook[*BootstrapEvent] OnAfterBootstrap() *hook.Hook[*BootstrapEvent]
// OnBeforeServe hook is triggered before serving the internal router (echo), // OnBeforeServe hook is triggered before serving the internal router (echo),
// allowing you to adjust its options and attach new routes. // allowing you to adjust its options and attach new routes or middlewares.
OnBeforeServe() *hook.Hook[*ServeEvent] OnBeforeServe() *hook.Hook[*ServeEvent]
// OnBeforeApiError hook is triggered right before sending an error API // OnBeforeApiError hook is triggered right before sending an error API
// response to the client, allowing you to further modify the error data // response to the client, allowing you to further modify the error data
// or to return a completely different API response (using [hook.StopPropagation]). // or to return a completely different API response.
OnBeforeApiError() *hook.Hook[*ApiErrorEvent] OnBeforeApiError() *hook.Hook[*ApiErrorEvent]
// OnAfterApiError hook is triggered right after sending an error API // OnAfterApiError hook is triggered right after sending an error API
@@ -162,7 +168,7 @@ type App interface {
// --------------------------------------------------------------- // ---------------------------------------------------------------
// OnModelBeforeCreate hook is triggered before inserting a new // OnModelBeforeCreate hook is triggered before inserting a new
// entry in the DB, allowing you to modify or validate the stored data. // model in the DB, allowing you to modify or validate the stored data.
// //
// If the optional "tags" list (table names and/or the Collection id for Record models) // If the optional "tags" list (table names and/or the Collection id for Record models)
// is specified, then all event handlers registered via the created hook // is specified, then all event handlers registered via the created hook
@@ -170,7 +176,7 @@ type App interface {
OnModelBeforeCreate(tags ...string) *hook.TaggedHook[*ModelEvent] OnModelBeforeCreate(tags ...string) *hook.TaggedHook[*ModelEvent]
// OnModelAfterCreate hook is triggered after successfully // OnModelAfterCreate hook is triggered after successfully
// inserting a new entry in the DB. // inserting a new model in the DB.
// //
// If the optional "tags" list (table names and/or the Collection id for Record models) // If the optional "tags" list (table names and/or the Collection id for Record models)
// is specified, then all event handlers registered via the created hook // is specified, then all event handlers registered via the created hook
@@ -178,7 +184,7 @@ type App interface {
OnModelAfterCreate(tags ...string) *hook.TaggedHook[*ModelEvent] OnModelAfterCreate(tags ...string) *hook.TaggedHook[*ModelEvent]
// OnModelBeforeUpdate hook is triggered before updating existing // OnModelBeforeUpdate hook is triggered before updating existing
// entry in the DB, allowing you to modify or validate the stored data. // model in the DB, allowing you to modify or validate the stored data.
// //
// If the optional "tags" list (table names and/or the Collection id for Record models) // If the optional "tags" list (table names and/or the Collection id for Record models)
// is specified, then all event handlers registered via the created hook // is specified, then all event handlers registered via the created hook
@@ -186,7 +192,7 @@ type App interface {
OnModelBeforeUpdate(tags ...string) *hook.TaggedHook[*ModelEvent] OnModelBeforeUpdate(tags ...string) *hook.TaggedHook[*ModelEvent]
// OnModelAfterUpdate hook is triggered after successfully updating // OnModelAfterUpdate hook is triggered after successfully updating
// existing entry in the DB. // existing model in the DB.
// //
// If the optional "tags" list (table names and/or the Collection id for Record models) // If the optional "tags" list (table names and/or the Collection id for Record models)
// is specified, then all event handlers registered via the created hook // is specified, then all event handlers registered via the created hook
@@ -194,15 +200,15 @@ type App interface {
OnModelAfterUpdate(tags ...string) *hook.TaggedHook[*ModelEvent] OnModelAfterUpdate(tags ...string) *hook.TaggedHook[*ModelEvent]
// OnModelBeforeDelete hook is triggered before deleting an // OnModelBeforeDelete hook is triggered before deleting an
// existing entry from the DB. // existing model from the DB.
// //
// If the optional "tags" list (table names and/or the Collection id for Record models) // If the optional "tags" list (table names and/or the Collection id for Record models)
// is specified, then all event handlers registered via the created hook // is specified, then all event handlers registered via the created hook
// will be triggered and called only if their event data origin matches the tags. // will be triggered and called only if their event data origin matches the tags.
OnModelBeforeDelete(tags ...string) *hook.TaggedHook[*ModelEvent] OnModelBeforeDelete(tags ...string) *hook.TaggedHook[*ModelEvent]
// OnModelAfterDelete is triggered after successfully deleting an // OnModelAfterDelete hook is triggered after successfully deleting an
// existing entry from the DB. // existing model from the DB.
// //
// If the optional "tags" list (table names and/or the Collection id for Record models) // If the optional "tags" list (table names and/or the Collection id for Record models)
// is specified, then all event handlers registered via the created hook // is specified, then all event handlers registered via the created hook
@@ -213,22 +219,18 @@ type App interface {
// Mailer event hooks // Mailer event hooks
// --------------------------------------------------------------- // ---------------------------------------------------------------
// OnMailerBeforeAdminResetPasswordSend hook is triggered right before // OnMailerBeforeAdminResetPasswordSend hook is triggered right
// sending a password reset email to an admin. // before sending a password reset email to an admin, allowing you
// // to inspect and customize the email message that is being sent.
// Could be used to send your own custom email template if
// [hook.StopPropagation] is returned in one of its listeners.
OnMailerBeforeAdminResetPasswordSend() *hook.Hook[*MailerAdminEvent] OnMailerBeforeAdminResetPasswordSend() *hook.Hook[*MailerAdminEvent]
// OnMailerAfterAdminResetPasswordSend hook is triggered after // OnMailerAfterAdminResetPasswordSend hook is triggered after
// admin password reset email was successfully sent. // admin password reset email was successfully sent.
OnMailerAfterAdminResetPasswordSend() *hook.Hook[*MailerAdminEvent] OnMailerAfterAdminResetPasswordSend() *hook.Hook[*MailerAdminEvent]
// OnMailerBeforeRecordResetPasswordSend hook is triggered right before // OnMailerBeforeRecordResetPasswordSend hook is triggered right
// sending a password reset email to an auth record. // before sending a password reset email to an auth record, allowing
// // you to inspect and customize the email message that is being sent.
// Could be used to send your own custom email template if
// [hook.StopPropagation] is returned in one of its listeners.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -243,11 +245,9 @@ type App interface {
// triggered and called only if their event data origin matches the tags. // triggered and called only if their event data origin matches the tags.
OnMailerAfterRecordResetPasswordSend(tags ...string) *hook.TaggedHook[*MailerRecordEvent] OnMailerAfterRecordResetPasswordSend(tags ...string) *hook.TaggedHook[*MailerRecordEvent]
// OnMailerBeforeRecordVerificationSend hook is triggered right before // OnMailerBeforeRecordVerificationSend hook is triggered right
// sending a verification email to an auth record. // before sending a verification email to an auth record, allowing
// // you to inspect and customize the email message that is being sent.
// Could be used to send your own custom email template if
// [hook.StopPropagation] is returned in one of its listeners.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -263,10 +263,8 @@ type App interface {
OnMailerAfterRecordVerificationSend(tags ...string) *hook.TaggedHook[*MailerRecordEvent] OnMailerAfterRecordVerificationSend(tags ...string) *hook.TaggedHook[*MailerRecordEvent]
// OnMailerBeforeRecordChangeEmailSend hook is triggered right before // OnMailerBeforeRecordChangeEmailSend hook is triggered right before
// sending a confirmation new address email to an auth record. // sending a confirmation new address email to an auth record, allowing
// // you to inspect and customize the email message that is being sent.
// Could be used to send your own custom email template if
// [hook.StopPropagation] is returned in one of its listeners.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -293,14 +291,14 @@ type App interface {
// SSE client connection. // SSE client connection.
OnRealtimeDisconnectRequest() *hook.Hook[*RealtimeDisconnectEvent] OnRealtimeDisconnectRequest() *hook.Hook[*RealtimeDisconnectEvent]
// OnRealtimeBeforeMessage hook is triggered right before sending // OnRealtimeBeforeMessageSend hook is triggered right before sending
// an SSE message to a client. // an SSE message to a client.
// //
// Returning [hook.StopPropagation] will prevent sending the message. // Returning [hook.StopPropagation] will prevent sending the message.
// Returning any other non-nil error will close the realtime connection. // Returning any other non-nil error will close the realtime connection.
OnRealtimeBeforeMessageSend() *hook.Hook[*RealtimeMessageEvent] OnRealtimeBeforeMessageSend() *hook.Hook[*RealtimeMessageEvent]
// OnRealtimeBeforeMessage hook is triggered right after sending // OnRealtimeAfterMessageSend hook is triggered right after sending
// an SSE message to a client. // an SSE message to a client.
OnRealtimeAfterMessageSend() *hook.Hook[*RealtimeMessageEvent] OnRealtimeAfterMessageSend() *hook.Hook[*RealtimeMessageEvent]
@@ -328,8 +326,7 @@ type App interface {
// Settings update request (after request data load and before settings persistence). // Settings update request (after request data load and before settings persistence).
// //
// Could be used to additionally validate the request data or // Could be used to additionally validate the request data or
// implement completely different persistence behavior // implement completely different persistence behavior.
// (returning [hook.StopPropagation]).
OnSettingsBeforeUpdateRequest() *hook.Hook[*SettingsUpdateEvent] OnSettingsBeforeUpdateRequest() *hook.Hook[*SettingsUpdateEvent]
// OnSettingsAfterUpdateRequest hook is triggered after each // OnSettingsAfterUpdateRequest hook is triggered after each
@@ -383,7 +380,7 @@ type App interface {
// Admin create request (after request data load and before model persistence). // Admin create request (after request data load and before model persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
OnAdminBeforeCreateRequest() *hook.Hook[*AdminCreateEvent] OnAdminBeforeCreateRequest() *hook.Hook[*AdminCreateEvent]
// OnAdminAfterCreateRequest hook is triggered after each // OnAdminAfterCreateRequest hook is triggered after each
@@ -394,7 +391,7 @@ type App interface {
// Admin update request (after request data load and before model persistence). // Admin update request (after request data load and before model persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
OnAdminBeforeUpdateRequest() *hook.Hook[*AdminUpdateEvent] OnAdminBeforeUpdateRequest() *hook.Hook[*AdminUpdateEvent]
// OnAdminAfterUpdateRequest hook is triggered after each // OnAdminAfterUpdateRequest hook is triggered after each
@@ -405,7 +402,7 @@ type App interface {
// Admin delete request (after model load and before actual deletion). // Admin delete request (after model load and before actual deletion).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different delete behavior (returning [hook.StopPropagation]). // completely different delete behavior.
OnAdminBeforeDeleteRequest() *hook.Hook[*AdminDeleteEvent] OnAdminBeforeDeleteRequest() *hook.Hook[*AdminDeleteEvent]
// OnAdminAfterDeleteRequest hook is triggered after each // OnAdminAfterDeleteRequest hook is triggered after each
@@ -434,7 +431,7 @@ type App interface {
// auth refresh API request (right before generating a new auth token). // auth refresh API request (right before generating a new auth token).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different auth refresh behavior (returning [hook.StopPropagation]). // completely different auth refresh behavior.
OnAdminBeforeAuthRefreshRequest() *hook.Hook[*AdminAuthRefreshEvent] OnAdminBeforeAuthRefreshRequest() *hook.Hook[*AdminAuthRefreshEvent]
// OnAdminAfterAuthRefreshRequest hook is triggered after each // OnAdminAfterAuthRefreshRequest hook is triggered after each
@@ -445,7 +442,7 @@ type App interface {
// request password reset API request (after request data load and before sending the reset email). // request password reset API request (after request data load and before sending the reset email).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different password reset behavior (returning [hook.StopPropagation]). // completely different password reset behavior.
OnAdminBeforeRequestPasswordResetRequest() *hook.Hook[*AdminRequestPasswordResetEvent] OnAdminBeforeRequestPasswordResetRequest() *hook.Hook[*AdminRequestPasswordResetEvent]
// OnAdminAfterRequestPasswordResetRequest hook is triggered after each // OnAdminAfterRequestPasswordResetRequest hook is triggered after each
@@ -456,7 +453,7 @@ type App interface {
// confirm password reset API request (after request data load and before persistence). // confirm password reset API request (after request data load and before persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
OnAdminBeforeConfirmPasswordResetRequest() *hook.Hook[*AdminConfirmPasswordResetEvent] OnAdminBeforeConfirmPasswordResetRequest() *hook.Hook[*AdminConfirmPasswordResetEvent]
// OnAdminAfterConfirmPasswordResetRequest hook is triggered after each // OnAdminAfterConfirmPasswordResetRequest hook is triggered after each
@@ -482,7 +479,7 @@ type App interface {
// auth with password API request (after request data load and before password validation). // auth with password API request (after request data load and before password validation).
// //
// Could be used to implement for example a custom password validation // Could be used to implement for example a custom password validation
// or to locate a different Record identity (by assigning [RecordAuthWithPasswordEvent.Record]). // or to locate a different Record model (by reassigning [RecordAuthWithPasswordEvent.Record]).
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -500,11 +497,11 @@ type App interface {
// OnRecordBeforeAuthWithOAuth2Request hook is triggered before each Record // OnRecordBeforeAuthWithOAuth2Request hook is triggered before each Record
// OAuth2 sign-in/sign-up API request (after token exchange and before external provider linking). // OAuth2 sign-in/sign-up API request (after token exchange and before external provider linking).
// //
// If the [RecordAuthWithOAuth2Event.Record] is nil, then the OAuth2 // If the [RecordAuthWithOAuth2Event.Record] is not set, then the OAuth2
// request will try to create a new auth Record. // request will try to create a new auth Record.
// //
// To assign or link a different existing record model you can // To assign or link a different existing record model you can
// overwrite/modify the [RecordAuthWithOAuth2Event.Record] field. // change the [RecordAuthWithOAuth2Event.Record] field.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -523,7 +520,7 @@ type App interface {
// auth refresh API request (right before generating a new auth token). // auth refresh API request (right before generating a new auth token).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different auth refresh behavior (returning [hook.StopPropagation]). // completely different auth refresh behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -538,11 +535,39 @@ type App interface {
// triggered and called only if their event data origin matches the tags. // triggered and called only if their event data origin matches the tags.
OnRecordAfterAuthRefreshRequest(tags ...string) *hook.TaggedHook[*RecordAuthRefreshEvent] OnRecordAfterAuthRefreshRequest(tags ...string) *hook.TaggedHook[*RecordAuthRefreshEvent]
// OnRecordListExternalAuthsRequest hook is triggered on each API record external auths list request.
//
// Could be used to validate or modify the response before returning it to the client.
//
// If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be
// triggered and called only if their event data origin matches the tags.
OnRecordListExternalAuthsRequest(tags ...string) *hook.TaggedHook[*RecordListExternalAuthsEvent]
// OnRecordBeforeUnlinkExternalAuthRequest hook is triggered before each API record
// external auth unlink request (after models load and before the actual relation deletion).
//
// Could be used to additionally validate the request data or implement
// completely different delete behavior.
//
// If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be
// triggered and called only if their event data origin matches the tags.
OnRecordBeforeUnlinkExternalAuthRequest(tags ...string) *hook.TaggedHook[*RecordUnlinkExternalAuthEvent]
// OnRecordAfterUnlinkExternalAuthRequest hook is triggered after each
// successful API record external auth unlink request.
//
// If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be
// triggered and called only if their event data origin matches the tags.
OnRecordAfterUnlinkExternalAuthRequest(tags ...string) *hook.TaggedHook[*RecordUnlinkExternalAuthEvent]
// OnRecordBeforeRequestPasswordResetRequest hook is triggered before each Record // OnRecordBeforeRequestPasswordResetRequest hook is triggered before each Record
// request password reset API request (after request data load and before sending the reset email). // request password reset API request (after request data load and before sending the reset email).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different password reset behavior (returning [hook.StopPropagation]). // completely different password reset behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -561,7 +586,7 @@ type App interface {
// confirm password reset API request (after request data load and before persistence). // confirm password reset API request (after request data load and before persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -580,7 +605,7 @@ type App interface {
// request verification API request (after request data load and before sending the verification email). // request verification API request (after request data load and before sending the verification email).
// //
// Could be used to additionally validate the loaded request data or implement // Could be used to additionally validate the loaded request data or implement
// completely different verification behavior (returning [hook.StopPropagation]). // completely different verification behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -599,7 +624,7 @@ type App interface {
// confirm verification API request (after request data load and before persistence). // confirm verification API request (after request data load and before persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -618,7 +643,7 @@ type App interface {
// (after request data load and before sending the email link to confirm the change). // (after request data load and before sending the email link to confirm the change).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different request email change behavior (returning [hook.StopPropagation]). // completely different request email change behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -637,7 +662,7 @@ type App interface {
// confirm email change API request (after request data load and before persistence). // confirm email change API request (after request data load and before persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -652,34 +677,6 @@ type App interface {
// triggered and called only if their event data origin matches the tags. // triggered and called only if their event data origin matches the tags.
OnRecordAfterConfirmEmailChangeRequest(tags ...string) *hook.TaggedHook[*RecordConfirmEmailChangeEvent] OnRecordAfterConfirmEmailChangeRequest(tags ...string) *hook.TaggedHook[*RecordConfirmEmailChangeEvent]
// OnRecordListExternalAuthsRequest hook is triggered on each API record external auths list request.
//
// Could be used to validate or modify the response before returning it to the client.
//
// If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be
// triggered and called only if their event data origin matches the tags.
OnRecordListExternalAuthsRequest(tags ...string) *hook.TaggedHook[*RecordListExternalAuthsEvent]
// OnRecordBeforeUnlinkExternalAuthRequest hook is triggered before each API record
// external auth unlink request (after models load and before the actual relation deletion).
//
// Could be used to additionally validate the request data or implement
// completely different delete behavior (returning [hook.StopPropagation]).
//
// If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be
// triggered and called only if their event data origin matches the tags.
OnRecordBeforeUnlinkExternalAuthRequest(tags ...string) *hook.TaggedHook[*RecordUnlinkExternalAuthEvent]
// OnRecordAfterUnlinkExternalAuthRequest hook is triggered after each
// successful API record external auth unlink request.
//
// If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be
// triggered and called only if their event data origin matches the tags.
OnRecordAfterUnlinkExternalAuthRequest(tags ...string) *hook.TaggedHook[*RecordUnlinkExternalAuthEvent]
// --------------------------------------------------------------- // ---------------------------------------------------------------
// Record CRUD API event hooks // Record CRUD API event hooks
// --------------------------------------------------------------- // ---------------------------------------------------------------
@@ -706,7 +703,7 @@ type App interface {
// create request (after request data load and before model persistence). // create request (after request data load and before model persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -725,7 +722,7 @@ type App interface {
// update request (after request data load and before model persistence). // update request (after request data load and before model persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -744,7 +741,7 @@ type App interface {
// delete request (after model load and before actual deletion). // delete request (after model load and before actual deletion).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different delete behavior (returning [hook.StopPropagation]). // completely different delete behavior.
// //
// If the optional "tags" list (Collection ids or names) is specified, // If the optional "tags" list (Collection ids or names) is specified,
// then all event handlers registered via the created hook will be // then all event handlers registered via the created hook will be
@@ -777,7 +774,7 @@ type App interface {
// create request (after request data load and before model persistence). // create request (after request data load and before model persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
OnCollectionBeforeCreateRequest() *hook.Hook[*CollectionCreateEvent] OnCollectionBeforeCreateRequest() *hook.Hook[*CollectionCreateEvent]
// OnCollectionAfterCreateRequest hook is triggered after each // OnCollectionAfterCreateRequest hook is triggered after each
@@ -788,7 +785,7 @@ type App interface {
// update request (after request data load and before model persistence). // update request (after request data load and before model persistence).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different persistence behavior (returning [hook.StopPropagation]). // completely different persistence behavior.
OnCollectionBeforeUpdateRequest() *hook.Hook[*CollectionUpdateEvent] OnCollectionBeforeUpdateRequest() *hook.Hook[*CollectionUpdateEvent]
// OnCollectionAfterUpdateRequest hook is triggered after each // OnCollectionAfterUpdateRequest hook is triggered after each
@@ -799,7 +796,7 @@ type App interface {
// Collection delete request (after model load and before actual deletion). // Collection delete request (after model load and before actual deletion).
// //
// Could be used to additionally validate the request data or implement // Could be used to additionally validate the request data or implement
// completely different delete behavior (returning [hook.StopPropagation]). // completely different delete behavior.
OnCollectionBeforeDeleteRequest() *hook.Hook[*CollectionDeleteEvent] OnCollectionBeforeDeleteRequest() *hook.Hook[*CollectionDeleteEvent]
// OnCollectionAfterDeleteRequest hook is triggered after each // OnCollectionAfterDeleteRequest hook is triggered after each
@@ -810,7 +807,7 @@ type App interface {
// collections import request (after request data load and before the actual import). // collections import request (after request data load and before the actual import).
// //
// Could be used to additionally validate the imported collections or // Could be used to additionally validate the imported collections or
// to implement completely different import behavior (returning [hook.StopPropagation]). // to implement completely different import behavior.
OnCollectionsBeforeImportRequest() *hook.Hook[*CollectionsImportEvent] OnCollectionsBeforeImportRequest() *hook.Hook[*CollectionsImportEvent]
// OnCollectionsAfterImportRequest hook is triggered after each // OnCollectionsAfterImportRequest hook is triggered after each
+213 -67
View File
@@ -5,6 +5,7 @@ import (
"database/sql" "database/sql"
"errors" "errors"
"log" "log"
"log/slog"
"os" "os"
"path/filepath" "path/filepath"
"runtime" "runtime"
@@ -18,10 +19,14 @@ import (
"github.com/pocketbase/pocketbase/models/settings" "github.com/pocketbase/pocketbase/models/settings"
"github.com/pocketbase/pocketbase/tools/filesystem" "github.com/pocketbase/pocketbase/tools/filesystem"
"github.com/pocketbase/pocketbase/tools/hook" "github.com/pocketbase/pocketbase/tools/hook"
"github.com/pocketbase/pocketbase/tools/logger"
"github.com/pocketbase/pocketbase/tools/mailer" "github.com/pocketbase/pocketbase/tools/mailer"
"github.com/pocketbase/pocketbase/tools/routine" "github.com/pocketbase/pocketbase/tools/routine"
"github.com/pocketbase/pocketbase/tools/security"
"github.com/pocketbase/pocketbase/tools/store" "github.com/pocketbase/pocketbase/tools/store"
"github.com/pocketbase/pocketbase/tools/subscriptions" "github.com/pocketbase/pocketbase/tools/subscriptions"
"github.com/pocketbase/pocketbase/tools/types"
"github.com/spf13/cast"
) )
const ( const (
@@ -39,8 +44,10 @@ var _ App = (*BaseApp)(nil)
// BaseApp implements core.App and defines the base PocketBase app structure. // BaseApp implements core.App and defines the base PocketBase app structure.
type BaseApp struct { type BaseApp struct {
// @todo consider introducing a mutex to allow safe concurrent config changes during runtime
// configurable parameters // configurable parameters
isDebug bool isDev bool
dataDir string dataDir string
encryptionEnv string encryptionEnv string
dataMaxOpenConns int dataMaxOpenConns int
@@ -49,11 +56,12 @@ type BaseApp struct {
logsMaxIdleConns int logsMaxIdleConns int
// internals // internals
cache *store.Store[any] store *store.Store[any]
settings *settings.Settings settings *settings.Settings
dao *daos.Dao dao *daos.Dao
logsDao *daos.Dao logsDao *daos.Dao
subscriptionsBroker *subscriptions.Broker subscriptionsBroker *subscriptions.Broker
logger *slog.Logger
// app event hooks // app event hooks
onBeforeBootstrap *hook.Hook[*BootstrapEvent] onBeforeBootstrap *hook.Hook[*BootstrapEvent]
@@ -167,9 +175,9 @@ type BaseApp struct {
// BaseAppConfig defines a BaseApp configuration option // BaseAppConfig defines a BaseApp configuration option
type BaseAppConfig struct { type BaseAppConfig struct {
IsDev bool
DataDir string DataDir string
EncryptionEnv string EncryptionEnv string
IsDebug bool
DataMaxOpenConns int // default to 500 DataMaxOpenConns int // default to 500
DataMaxIdleConns int // default 20 DataMaxIdleConns int // default 20
LogsMaxOpenConns int // default to 100 LogsMaxOpenConns int // default to 100
@@ -180,16 +188,16 @@ type BaseAppConfig struct {
// configured with the provided arguments. // configured with the provided arguments.
// //
// To initialize the app, you need to call `app.Bootstrap()`. // To initialize the app, you need to call `app.Bootstrap()`.
func NewBaseApp(config *BaseAppConfig) *BaseApp { func NewBaseApp(config BaseAppConfig) *BaseApp {
app := &BaseApp{ app := &BaseApp{
isDev: config.IsDev,
dataDir: config.DataDir, dataDir: config.DataDir,
isDebug: config.IsDebug,
encryptionEnv: config.EncryptionEnv, encryptionEnv: config.EncryptionEnv,
dataMaxOpenConns: config.DataMaxOpenConns, dataMaxOpenConns: config.DataMaxOpenConns,
dataMaxIdleConns: config.DataMaxIdleConns, dataMaxIdleConns: config.DataMaxIdleConns,
logsMaxOpenConns: config.LogsMaxOpenConns, logsMaxOpenConns: config.LogsMaxOpenConns,
logsMaxIdleConns: config.LogsMaxIdleConns, logsMaxIdleConns: config.LogsMaxIdleConns,
cache: store.New[any](nil), store: store.New[any](nil),
settings: settings.New(), settings: settings.New(),
subscriptionsBroker: subscriptions.NewBroker(), subscriptionsBroker: subscriptions.NewBroker(),
@@ -314,6 +322,17 @@ func (app *BaseApp) IsBootstrapped() bool {
return app.dao != nil && app.logsDao != nil && app.settings != nil return app.dao != nil && app.logsDao != nil && app.settings != nil
} }
// Logger returns the default app logger.
//
// If the application is not bootstrapped yet, fallbacks to slog.Default().
func (app *BaseApp) Logger() *slog.Logger {
if app.logger == nil {
return slog.Default()
}
return app.logger
}
// Bootstrap initializes the application // Bootstrap initializes the application
// (aka. create data dir, open db connections, load settings, etc.). // (aka. create data dir, open db connections, load settings, etc.).
// //
@@ -343,17 +362,17 @@ func (app *BaseApp) Bootstrap() error {
return err return err
} }
if err := app.initLogger(); err != nil {
return err
}
// we don't check for an error because the db migrations may have not been executed yet // we don't check for an error because the db migrations may have not been executed yet
app.RefreshSettings() app.RefreshSettings()
// cleanup the pb_data temp directory (if any) // cleanup the pb_data temp directory (if any)
os.RemoveAll(filepath.Join(app.DataDir(), LocalTempDirName)) os.RemoveAll(filepath.Join(app.DataDir(), LocalTempDirName))
if err := app.OnAfterBootstrap().Trigger(event); err != nil && app.IsDebug() { return app.OnAfterBootstrap().Trigger(event)
log.Println(err)
}
return nil
} }
// ResetBootstrapState takes care for releasing initialized app resources // ResetBootstrapState takes care for releasing initialized app resources
@@ -379,7 +398,6 @@ func (app *BaseApp) ResetBootstrapState() error {
app.dao = nil app.dao = nil
app.logsDao = nil app.logsDao = nil
app.settings = nil
return nil return nil
} }
@@ -443,10 +461,11 @@ func (app *BaseApp) EncryptionEnv() string {
return app.encryptionEnv return app.encryptionEnv
} }
// IsDebug returns whether the app is in debug mode // IsDev returns whether the app is in dev mode.
// (showing more detailed error logs, executed sql statements, etc.). //
func (app *BaseApp) IsDebug() bool { // When enabled logs, executed sql statements, etc. are printed to the stderr.
return app.isDebug func (app *BaseApp) IsDev() bool {
return app.isDev
} }
// Settings returns the loaded app settings. // Settings returns the loaded app settings.
@@ -454,9 +473,15 @@ func (app *BaseApp) Settings() *settings.Settings {
return app.settings return app.settings
} }
// Cache returns the app internal cache store. // Deprecated: Use app.Store() instead.
func (app *BaseApp) Cache() *store.Store[any] { func (app *BaseApp) Cache() *store.Store[any] {
return app.cache color.Yellow("app.Store() is soft-deprecated. Please replace it with app.Store().")
return app.Store()
}
// Store returns the app internal runtime store.
func (app *BaseApp) Store() *store.Store[any] {
return app.store
} }
// SubscriptionsBroker returns the app realtime subscriptions broker instance. // SubscriptionsBroker returns the app realtime subscriptions broker instance.
@@ -475,6 +500,7 @@ func (app *BaseApp) NewMailClient() mailer.Mailer {
Password: app.Settings().Smtp.Password, Password: app.Settings().Smtp.Password,
Tls: app.Settings().Smtp.Tls, Tls: app.Settings().Smtp.Tls,
AuthMethod: app.Settings().Smtp.AuthMethod, AuthMethod: app.Settings().Smtp.AuthMethod,
LocalName: app.Settings().Smtp.LocalName,
} }
} }
@@ -485,7 +511,7 @@ func (app *BaseApp) NewMailClient() mailer.Mailer {
// for managing regular app files (eg. collection uploads) // for managing regular app files (eg. collection uploads)
// based on the current app settings. // based on the current app settings.
// //
// NB! Make sure to call `Close()` on the returned result // NB! Make sure to call Close() on the returned result
// after you are done working with it. // after you are done working with it.
func (app *BaseApp) NewFilesystem() (*filesystem.System, error) { func (app *BaseApp) NewFilesystem() (*filesystem.System, error) {
if app.settings != nil && app.settings.S3.Enabled { if app.settings != nil && app.settings.S3.Enabled {
@@ -506,7 +532,7 @@ func (app *BaseApp) NewFilesystem() (*filesystem.System, error) {
// NewFilesystem creates a new local or S3 filesystem instance // NewFilesystem creates a new local or S3 filesystem instance
// for managing app backups based on the current app settings. // for managing app backups based on the current app settings.
// //
// NB! Make sure to call `Close()` on the returned result // NB! Make sure to call Close() on the returned result
// after you are done working with it. // after you are done working with it.
func (app *BaseApp) NewBackupsFilesystem() (*filesystem.System, error) { func (app *BaseApp) NewBackupsFilesystem() (*filesystem.System, error) {
if app.settings != nil && app.settings.Backups.S3.Enabled { if app.settings != nil && app.settings.Backups.S3.Enabled {
@@ -537,17 +563,17 @@ func (app *BaseApp) Restart() error {
return err return err
} }
// optimistically reset the app bootstrap state return app.OnTerminate().Trigger(&TerminateEvent{
app.ResetBootstrapState() App: app,
IsRestart: true,
}, func(e *TerminateEvent) error {
e.App.ResetBootstrapState()
if err := syscall.Exec(execPath, os.Args, os.Environ()); err != nil { // attempt to restart the bootstrap process in case execve returns an error for some reason
// restart the app bootstrap state defer e.App.Bootstrap()
app.Bootstrap()
return err return syscall.Exec(execPath, os.Args, os.Environ())
} })
return nil
} }
// RefreshSettings reinitializes and reloads the stored application settings. // RefreshSettings reinitializes and reloads the stored application settings.
@@ -559,7 +585,7 @@ func (app *BaseApp) RefreshSettings() error {
encryptionKey := os.Getenv(app.EncryptionEnv()) encryptionKey := os.Getenv(app.EncryptionEnv())
storedSettings, err := app.Dao().FindSettings(encryptionKey) storedSettings, err := app.Dao().FindSettings(encryptionKey)
if err != nil && err != sql.ErrNoRows { if err != nil && !errors.Is(err, sql.ErrNoRows) {
return err return err
} }
@@ -573,6 +599,13 @@ func (app *BaseApp) RefreshSettings() error {
return err return err
} }
// reload handler level (if initialized)
if app.Logger() != nil {
if h, ok := app.Logger().Handler().(*logger.BatchHandler); ok {
h.SetLevel(app.getLoggerMinLevel())
}
}
return nil return nil
} }
@@ -992,7 +1025,7 @@ func (app *BaseApp) initLogsDB() error {
} }
concurrentDB.DB().SetMaxOpenConns(maxOpenConns) concurrentDB.DB().SetMaxOpenConns(maxOpenConns)
concurrentDB.DB().SetMaxIdleConns(maxIdleConns) concurrentDB.DB().SetMaxIdleConns(maxIdleConns)
concurrentDB.DB().SetConnMaxIdleTime(5 * time.Minute) concurrentDB.DB().SetConnMaxIdleTime(3 * time.Minute)
nonconcurrentDB, err := connectDB(filepath.Join(app.DataDir(), "logs.db")) nonconcurrentDB, err := connectDB(filepath.Join(app.DataDir(), "logs.db"))
if err != nil { if err != nil {
@@ -1000,7 +1033,7 @@ func (app *BaseApp) initLogsDB() error {
} }
nonconcurrentDB.DB().SetMaxOpenConns(1) nonconcurrentDB.DB().SetMaxOpenConns(1)
nonconcurrentDB.DB().SetMaxIdleConns(1) nonconcurrentDB.DB().SetMaxIdleConns(1)
nonconcurrentDB.DB().SetConnMaxIdleTime(5 * time.Minute) nonconcurrentDB.DB().SetConnMaxIdleTime(3 * time.Minute)
app.logsDao = daos.NewMultiDB(concurrentDB, nonconcurrentDB) app.logsDao = daos.NewMultiDB(concurrentDB, nonconcurrentDB)
@@ -1023,7 +1056,7 @@ func (app *BaseApp) initDataDB() error {
} }
concurrentDB.DB().SetMaxOpenConns(maxOpenConns) concurrentDB.DB().SetMaxOpenConns(maxOpenConns)
concurrentDB.DB().SetMaxIdleConns(maxIdleConns) concurrentDB.DB().SetMaxIdleConns(maxIdleConns)
concurrentDB.DB().SetConnMaxIdleTime(5 * time.Minute) concurrentDB.DB().SetConnMaxIdleTime(3 * time.Minute)
nonconcurrentDB, err := connectDB(filepath.Join(app.DataDir(), "data.db")) nonconcurrentDB, err := connectDB(filepath.Join(app.DataDir(), "data.db"))
if err != nil { if err != nil {
@@ -1031,17 +1064,16 @@ func (app *BaseApp) initDataDB() error {
} }
nonconcurrentDB.DB().SetMaxOpenConns(1) nonconcurrentDB.DB().SetMaxOpenConns(1)
nonconcurrentDB.DB().SetMaxIdleConns(1) nonconcurrentDB.DB().SetMaxIdleConns(1)
nonconcurrentDB.DB().SetConnMaxIdleTime(5 * time.Minute) nonconcurrentDB.DB().SetConnMaxIdleTime(3 * time.Minute)
if app.IsDebug() { if app.IsDev() {
nonconcurrentDB.QueryLogFunc = func(ctx context.Context, t time.Duration, sql string, rows *sql.Rows, err error) { nonconcurrentDB.QueryLogFunc = func(ctx context.Context, t time.Duration, sql string, rows *sql.Rows, err error) {
color.HiBlack("[%.2fms] %v\n", float64(t.Milliseconds()), sql) color.HiBlack("[%.2fms] %v\n", float64(t.Milliseconds()), sql)
} }
concurrentDB.QueryLogFunc = nonconcurrentDB.QueryLogFunc
nonconcurrentDB.ExecLogFunc = func(ctx context.Context, t time.Duration, sql string, result sql.Result, err error) { nonconcurrentDB.ExecLogFunc = func(ctx context.Context, t time.Duration, sql string, result sql.Result, err error) {
color.HiBlack("[%.2fms] %v\n", float64(t.Milliseconds()), sql) color.HiBlack("[%.2fms] %v\n", float64(t.Milliseconds()), sql)
} }
concurrentDB.QueryLogFunc = nonconcurrentDB.QueryLogFunc
concurrentDB.ExecLogFunc = nonconcurrentDB.ExecLogFunc concurrentDB.ExecLogFunc = nonconcurrentDB.ExecLogFunc
} }
@@ -1053,58 +1085,58 @@ func (app *BaseApp) initDataDB() error {
func (app *BaseApp) createDaoWithHooks(concurrentDB, nonconcurrentDB dbx.Builder) *daos.Dao { func (app *BaseApp) createDaoWithHooks(concurrentDB, nonconcurrentDB dbx.Builder) *daos.Dao {
dao := daos.NewMultiDB(concurrentDB, nonconcurrentDB) dao := daos.NewMultiDB(concurrentDB, nonconcurrentDB)
dao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { dao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
e := new(ModelEvent) e := new(ModelEvent)
e.Dao = eventDao e.Dao = eventDao
e.Model = m e.Model = m
return app.OnModelBeforeCreate().Trigger(e) return app.OnModelBeforeCreate().Trigger(e, func(e *ModelEvent) error {
return action()
})
} }
dao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) { dao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
e := new(ModelEvent) e := new(ModelEvent)
e.Dao = eventDao e.Dao = eventDao
e.Model = m e.Model = m
if err := app.OnModelAfterCreate().Trigger(e); err != nil && app.isDebug { return app.OnModelAfterCreate().Trigger(e)
log.Println(err)
}
} }
dao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { dao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
e := new(ModelEvent) e := new(ModelEvent)
e.Dao = eventDao e.Dao = eventDao
e.Model = m e.Model = m
return app.OnModelBeforeUpdate().Trigger(e) return app.OnModelBeforeUpdate().Trigger(e, func(e *ModelEvent) error {
return action()
})
} }
dao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) { dao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) error {
e := new(ModelEvent) e := new(ModelEvent)
e.Dao = eventDao e.Dao = eventDao
e.Model = m e.Model = m
if err := app.OnModelAfterUpdate().Trigger(e); err != nil && app.isDebug { return app.OnModelAfterUpdate().Trigger(e)
log.Println(err)
}
} }
dao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { dao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
e := new(ModelEvent) e := new(ModelEvent)
e.Dao = eventDao e.Dao = eventDao
e.Model = m e.Model = m
return app.OnModelBeforeDelete().Trigger(e) return app.OnModelBeforeDelete().Trigger(e, func(e *ModelEvent) error {
return action()
})
} }
dao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) { dao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) error {
e := new(ModelEvent) e := new(ModelEvent)
e.Dao = eventDao e.Dao = eventDao
e.Model = m e.Model = m
if err := app.OnModelAfterDelete().Trigger(e); err != nil && app.isDebug { return app.OnModelAfterDelete().Trigger(e)
log.Println(err)
}
} }
return dao return dao
@@ -1133,14 +1165,13 @@ func (app *BaseApp) registerDefaultHooks() {
// run in the background for "optimistic" delete to avoid // run in the background for "optimistic" delete to avoid
// blocking the delete transaction // blocking the delete transaction
//
// @todo consider creating a bg process queue so that the
// call could be "retried" in case of a failure.
routine.FireAndForget(func() { routine.FireAndForget(func() {
if err := deletePrefix(prefix); err != nil && app.IsDebug() { if err := deletePrefix(prefix); err != nil {
// non critical error - only log for debug app.Logger().Error(
// (usually could happen because of S3 api limits) "Failed to delete storage prefix (non critical error; usually could happen because of S3 api limits)",
log.Println(err) slog.String("prefix", prefix),
slog.String("error", err.Error()),
)
} }
}) })
} }
@@ -1148,12 +1179,127 @@ func (app *BaseApp) registerDefaultHooks() {
return nil return nil
}) })
app.OnTerminate().Add(func(e *TerminateEvent) error { if err := app.initAutobackupHooks(); err != nil {
app.ResetBootstrapState() app.Logger().Error("Failed to init auto backup hooks", slog.String("error", err.Error()))
}
registerCachedCollectionsAppHooks(app)
}
// getLoggerMinLevel returns the logger min level based on the
// app configurations (dev mode, settings, etc.).
//
// If not in dev mode - returns the level from the app settings.
//
// If the app is in dev mode it returns -9999 level allowing to print
// practically all logs to the terminal.
// In this case DB logs are still filtered but the checks for the min level are done
// in the BatchOptions.BeforeAddFunc instead of the slog.Handler.Enabled() method.
func (app *BaseApp) getLoggerMinLevel() slog.Level {
var minLevel slog.Level
if app.IsDev() {
minLevel = -9999
} else if app.Settings() != nil {
minLevel = slog.Level(app.Settings().Logs.MinLevel)
}
return minLevel
}
func (app *BaseApp) initLogger() error {
duration := 3 * time.Second
ticker := time.NewTicker(duration)
done := make(chan bool)
handler := logger.NewBatchHandler(logger.BatchOptions{
Level: app.getLoggerMinLevel(),
BatchSize: 200,
BeforeAddFunc: func(ctx context.Context, log *logger.Log) bool {
if app.IsDev() {
printLog(log)
// manually check the log level and skip if necessary
if log.Level < slog.Level(app.Settings().Logs.MinLevel) {
return false
}
}
ticker.Reset(duration)
return app.Settings().Logs.MaxDays > 0
},
WriteFunc: func(ctx context.Context, logs []*logger.Log) error {
if !app.IsBootstrapped() || app.Settings().Logs.MaxDays == 0 {
return nil
}
// write the accumulated logs
// (note: based on several local tests there is no significant performance difference between small number of separate write queries vs 1 big INSERT)
app.LogsDao().RunInTransaction(func(txDao *daos.Dao) error {
model := &models.Log{}
for _, l := range logs {
model.MarkAsNew()
// note: using pseudorandom for a slightly better performance
model.Id = security.PseudorandomStringWithAlphabet(models.DefaultIdLength, models.DefaultIdAlphabet)
model.Level = int(l.Level)
model.Message = l.Message
model.Data = l.Data
model.Created, _ = types.ParseDateTime(l.Time)
model.Updated = model.Created
if err := txDao.SaveLog(model); err != nil {
log.Println("Failed to write log", model, err)
}
}
return nil
})
// delete old logs
// ---
logsMaxDays := app.Settings().Logs.MaxDays
now := time.Now()
lastLogsDeletedAt := cast.ToTime(app.Store().Get("lastLogsDeletedAt"))
daysDiff := now.Sub(lastLogsDeletedAt).Hours() * 24
if daysDiff > float64(logsMaxDays) {
deleteErr := app.LogsDao().DeleteOldLogs(now.AddDate(0, 0, -1*logsMaxDays))
if deleteErr == nil {
app.Store().Set("lastLogsDeletedAt", now)
} else {
log.Println("Logs delete failed", deleteErr)
}
}
return nil
},
})
go func() {
ctx := context.Background()
for {
select {
case <-done:
return
case <-ticker.C:
handler.WriteAll(ctx)
}
}
}()
app.logger = slog.New(handler)
app.OnTerminate().PreAdd(func(e *TerminateEvent) error {
// write all remaining logs before ticker.Stop to avoid races with ResetBootstrap user calls
handler.WriteAll(context.Background())
ticker.Stop()
done <- true
return nil return nil
}) })
if err := app.initAutobackupHooks(); err != nil && app.IsDebug() { return nil
log.Println(err)
}
} }
+97 -51
View File
@@ -5,7 +5,7 @@ import (
"errors" "errors"
"fmt" "fmt"
"io" "io"
"log" "log/slog"
"os" "os"
"path/filepath" "path/filepath"
"runtime" "runtime"
@@ -17,12 +17,16 @@ import (
"github.com/pocketbase/pocketbase/tools/archive" "github.com/pocketbase/pocketbase/tools/archive"
"github.com/pocketbase/pocketbase/tools/cron" "github.com/pocketbase/pocketbase/tools/cron"
"github.com/pocketbase/pocketbase/tools/filesystem" "github.com/pocketbase/pocketbase/tools/filesystem"
"github.com/pocketbase/pocketbase/tools/inflector"
"github.com/pocketbase/pocketbase/tools/osutils" "github.com/pocketbase/pocketbase/tools/osutils"
"github.com/pocketbase/pocketbase/tools/security" "github.com/pocketbase/pocketbase/tools/security"
) )
// Deprecated: Replaced with StoreKeyActiveBackup.
const CacheKeyActiveBackup string = "@activeBackup" const CacheKeyActiveBackup string = "@activeBackup"
const StoreKeyActiveBackup string = "@activeBackup"
// CreateBackup creates a new backup of the current app pb_data directory. // CreateBackup creates a new backup of the current app pb_data directory.
// //
// If name is empty, it will be autogenerated. // If name is empty, it will be autogenerated.
@@ -31,6 +35,9 @@ const CacheKeyActiveBackup string = "@activeBackup"
// The backup is executed within a transaction, meaning that new writes // The backup is executed within a transaction, meaning that new writes
// will be temporary "blocked" until the backup file is generated. // will be temporary "blocked" until the backup file is generated.
// //
// To safely perform the backup, it is recommended to have free disk space
// for at least 2x the size of the pb_data directory.
//
// By default backups are stored in pb_data/backups // By default backups are stored in pb_data/backups
// (the backups directory itself is excluded from the generated backup). // (the backups directory itself is excluded from the generated backup).
// //
@@ -39,31 +46,35 @@ const CacheKeyActiveBackup string = "@activeBackup"
// //
// Backups can be stored on S3 if it is configured in app.Settings().Backups. // Backups can be stored on S3 if it is configured in app.Settings().Backups.
func (app *BaseApp) CreateBackup(ctx context.Context, name string) error { func (app *BaseApp) CreateBackup(ctx context.Context, name string) error {
if app.Cache().Has(CacheKeyActiveBackup) { if app.Store().Has(StoreKeyActiveBackup) {
return errors.New("try again later - another backup/restore operation has already been started") return errors.New("try again later - another backup/restore operation has already been started")
} }
// auto generate backup name
if name == "" { if name == "" {
name = fmt.Sprintf( name = app.generateBackupName("pb_backup_")
"pb_backup_%s.zip",
time.Now().UTC().Format("20060102150405"),
)
} }
app.Cache().Set(CacheKeyActiveBackup, name) app.Store().Set(StoreKeyActiveBackup, name)
defer app.Cache().Remove(CacheKeyActiveBackup) defer app.Store().Remove(StoreKeyActiveBackup)
// Archive pb_data in a temp directory, exluding the "backups" dir itself (if exist). // root dir entries to exclude from the backup generation
exclude := []string{LocalBackupsDirName, LocalTempDirName}
// make sure that the special temp directory exists
// note: it needs to be inside the current pb_data to avoid "cross-device link" errors
localTempDir := filepath.Join(app.DataDir(), LocalTempDirName)
if err := os.MkdirAll(localTempDir, os.ModePerm); err != nil {
return fmt.Errorf("failed to create a temp dir: %w", err)
}
// Archive pb_data in a temp directory, exluding the "backups" and the temp dirs.
// //
// Run in transaction to temporary block other writes (transactions uses the NonconcurrentDB connection). // Run in transaction to temporary block other writes (transactions uses the NonconcurrentDB connection).
// --- // ---
tempPath := filepath.Join(os.TempDir(), "pb_backup_"+security.PseudorandomString(4)) tempPath := filepath.Join(localTempDir, "pb_backup_"+security.PseudorandomString(4))
createErr := app.Dao().RunInTransaction(func(txDao *daos.Dao) error { createErr := app.Dao().RunInTransaction(func(txDao *daos.Dao) error {
if err := archive.Create(app.DataDir(), tempPath, LocalBackupsDirName); err != nil { // @todo consider experimenting with temp switching the readonly pragma after the db interface change
return err return archive.Create(app.DataDir(), tempPath, exclude...)
}
return nil
}) })
if createErr != nil { if createErr != nil {
return createErr return createErr
@@ -118,7 +129,7 @@ func (app *BaseApp) CreateBackup(ctx context.Context, name string) error {
// //
// 4. Move the extracted dir content to the app "pb_data". // 4. Move the extracted dir content to the app "pb_data".
// //
// 5. Restart the app (on successfull app bootstap it will also remove the old pb_data). // 5. Restart the app (on successful app bootstap it will also remove the old pb_data).
// //
// If a failure occure during the restore process the dir changes are reverted. // If a failure occure during the restore process the dir changes are reverted.
// If for whatever reason the revert is not possible, it panics. // If for whatever reason the revert is not possible, it panics.
@@ -127,12 +138,12 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
return errors.New("restore is not supported on windows") return errors.New("restore is not supported on windows")
} }
if app.Cache().Has(CacheKeyActiveBackup) { if app.Store().Has(StoreKeyActiveBackup) {
return errors.New("try again later - another backup/restore operation has already been started") return errors.New("try again later - another backup/restore operation has already been started")
} }
app.Cache().Set(CacheKeyActiveBackup, name) app.Store().Set(StoreKeyActiveBackup, name)
defer app.Cache().Remove(CacheKeyActiveBackup) defer app.Store().Remove(StoreKeyActiveBackup)
fsys, err := app.NewBackupsFilesystem() fsys, err := app.NewBackupsFilesystem()
if err != nil { if err != nil {
@@ -149,7 +160,15 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
} }
defer br.Close() defer br.Close()
tempZip, err := os.CreateTemp(os.TempDir(), "pb_restore") // make sure that the special temp directory exists
// note: it needs to be inside the current pb_data to avoid "cross-device link" errors
localTempDir := filepath.Join(app.DataDir(), LocalTempDirName)
if err := os.MkdirAll(localTempDir, os.ModePerm); err != nil {
return fmt.Errorf("failed to create a temp dir: %w", err)
}
// create a temp zip file from the blob.Reader and try to extract it
tempZip, err := os.CreateTemp(localTempDir, "pb_restore_zip")
if err != nil { if err != nil {
return err return err
} }
@@ -159,13 +178,7 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
return err return err
} }
// make sure that the special temp directory extractedDataDir := filepath.Join(localTempDir, "pb_restore_"+security.PseudorandomString(4))
if err := os.MkdirAll(filepath.Join(app.DataDir(), LocalTempDirName), os.ModePerm); err != nil {
return fmt.Errorf("failed to create a temp dir: %w", err)
}
// note: it needs to be inside the current pb_data to avoid "cross-device link" errors
extractedDataDir := filepath.Join(app.DataDir(), LocalTempDirName, "pb_restore_"+security.PseudorandomString(4))
defer os.RemoveAll(extractedDataDir) defer os.RemoveAll(extractedDataDir)
if err := archive.Extract(tempZip.Name(), extractedDataDir); err != nil { if err := archive.Extract(tempZip.Name(), extractedDataDir); err != nil {
return err return err
@@ -179,8 +192,12 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
// remove the extracted zip file since we no longer need it // remove the extracted zip file since we no longer need it
// (this is in case the app restarts and the defer calls are not called) // (this is in case the app restarts and the defer calls are not called)
if err := os.Remove(tempZip.Name()); err != nil && app.IsDebug() { if err := os.Remove(tempZip.Name()); err != nil {
log.Println(err) app.Logger().Debug(
"[RestoreBackup] Failed to remove the temp zip backup file",
slog.String("file", tempZip.Name()),
slog.String("error", err.Error()),
)
} }
// root dir entries to exclude from the backup restore // root dir entries to exclude from the backup restore
@@ -189,7 +206,7 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
// move the current pb_data content to a special temp location // move the current pb_data content to a special temp location
// that will hold the old data between dirs replace // that will hold the old data between dirs replace
// (the temp dir will be automatically removed on the next app start) // (the temp dir will be automatically removed on the next app start)
oldTempDataDir := filepath.Join(app.DataDir(), LocalTempDirName, "old_pb_data_"+security.PseudorandomString(4)) oldTempDataDir := filepath.Join(localTempDir, "old_pb_data_"+security.PseudorandomString(4))
if err := osutils.MoveDirContent(app.DataDir(), oldTempDataDir, exclude...); err != nil { if err := osutils.MoveDirContent(app.DataDir(), oldTempDataDir, exclude...); err != nil {
return fmt.Errorf("failed to move the current pb_data content to a temp location: %w", err) return fmt.Errorf("failed to move the current pb_data content to a temp location: %w", err)
} }
@@ -213,8 +230,8 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
// restart the app // restart the app
if err := app.Restart(); err != nil { if err := app.Restart(); err != nil {
if err := revertDataDirChanges(); err != nil { if revertErr := revertDataDirChanges(); revertErr != nil {
panic(err) panic(revertErr)
} }
return fmt.Errorf("failed to restart the app process: %w", err) return fmt.Errorf("failed to restart the app process: %w", err)
@@ -224,7 +241,6 @@ func (app *BaseApp) RestoreBackup(ctx context.Context, name string) error {
} }
// initAutobackupHooks registers the autobackup app serve hooks. // initAutobackupHooks registers the autobackup app serve hooks.
// @todo add tests
func (app *BaseApp) initAutobackupHooks() error { func (app *BaseApp) initAutobackupHooks() error {
c := cron.New() c := cron.New()
isServe := false isServe := false
@@ -232,23 +248,32 @@ func (app *BaseApp) initAutobackupHooks() error {
loadJob := func() { loadJob := func() {
c.Stop() c.Stop()
// make sure that app.Settings() is always up to date
//
// @todo remove with the refactoring as core.App and daos.Dao will be one.
if err := app.RefreshSettings(); err != nil {
app.Logger().Debug(
"[Backup cron] Failed to get the latest app settings",
slog.String("error", err.Error()),
)
}
rawSchedule := app.Settings().Backups.Cron rawSchedule := app.Settings().Backups.Cron
if rawSchedule == "" || !isServe || !app.IsBootstrapped() { if rawSchedule == "" || !isServe || !app.IsBootstrapped() {
return return
} }
c.Add("@autobackup", rawSchedule, func() { c.Add("@autobackup", rawSchedule, func() {
autoPrefix := "@auto_pb_backup_" const autoPrefix = "@auto_pb_backup_"
name := fmt.Sprintf( name := app.generateBackupName(autoPrefix)
"%s%s.zip",
autoPrefix,
time.Now().UTC().Format("20060102150405"),
)
if err := app.CreateBackup(context.Background(), name); err != nil && app.IsDebug() { if err := app.CreateBackup(context.Background(), name); err != nil {
// @todo replace after logs generalization app.Logger().Debug(
log.Println(err) "[Backup cron] Failed to create backup",
slog.String("name", name),
slog.String("error", err.Error()),
)
} }
maxKeep := app.Settings().Backups.CronMaxKeep maxKeep := app.Settings().Backups.CronMaxKeep
@@ -258,17 +283,21 @@ func (app *BaseApp) initAutobackupHooks() error {
} }
fsys, err := app.NewBackupsFilesystem() fsys, err := app.NewBackupsFilesystem()
if err != nil && app.IsDebug() { if err != nil {
// @todo replace after logs generalization app.Logger().Debug(
log.Println(err) "[Backup cron] Failed to initialize the backup filesystem",
slog.String("error", err.Error()),
)
return return
} }
defer fsys.Close() defer fsys.Close()
files, err := fsys.List(autoPrefix) files, err := fsys.List(autoPrefix)
if err != nil && app.IsDebug() { if err != nil {
// @todo replace after logs generalization app.Logger().Debug(
log.Println(err) "[Backup cron] Failed to list autogenerated backups",
slog.String("error", err.Error()),
)
return return
} }
@@ -285,9 +314,12 @@ func (app *BaseApp) initAutobackupHooks() error {
toRemove := files[maxKeep:] toRemove := files[maxKeep:]
for _, f := range toRemove { for _, f := range toRemove {
if err := fsys.Delete(f.Key); err != nil && app.IsDebug() { if err := fsys.Delete(f.Key); err != nil {
// @todo replace after logs generalization app.Logger().Debug(
log.Println(err) "[Backup cron] Failed to remove old autogenerated backup",
slog.String("key", f.Key),
slog.String("error", err.Error()),
)
} }
} }
}) })
@@ -323,3 +355,17 @@ func (app *BaseApp) initAutobackupHooks() error {
return nil return nil
} }
func (app *BaseApp) generateBackupName(prefix string) string {
appName := inflector.Snakecase(app.Settings().Meta.AppName)
if len(appName) > 50 {
appName = appName[:50]
}
return fmt.Sprintf(
"%s%s_%s.zip",
prefix,
appName,
time.Now().UTC().Format("20060102150405"),
)
}
+11 -6
View File
@@ -19,12 +19,17 @@ func TestCreateBackup(t *testing.T) {
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
// set some long app name with spaces and special characters
app.Settings().Meta.AppName = "test @! " + strings.Repeat("a", 100)
expectedAppNamePrefix := "test_" + strings.Repeat("a", 45)
// test pending error // test pending error
app.Cache().Set(core.CacheKeyActiveBackup, "") app.Store().Set(core.StoreKeyActiveBackup, "")
if err := app.CreateBackup(context.Background(), "test.zip"); err == nil { if err := app.CreateBackup(context.Background(), "test.zip"); err == nil {
t.Fatal("Expected pending error, got nil") t.Fatal("Expected pending error, got nil")
} }
app.Cache().Remove(core.CacheKeyActiveBackup) app.Store().Remove(core.StoreKeyActiveBackup)
// create with auto generated name // create with auto generated name
if err := app.CreateBackup(context.Background(), ""); err != nil { if err := app.CreateBackup(context.Background(), ""); err != nil {
@@ -49,8 +54,8 @@ func TestCreateBackup(t *testing.T) {
} }
expectedFiles := []string{ expectedFiles := []string{
`^pb_backup_\w+\.zip$`, `^pb_backup_` + expectedAppNamePrefix + `_\w+\.zip$`,
`^pb_backup_\w+\.zip.attrs$`, `^pb_backup_` + expectedAppNamePrefix + `_\w+\.zip.attrs$`,
"custom", "custom",
"custom.attrs", "custom.attrs",
} }
@@ -93,11 +98,11 @@ func TestRestoreBackup(t *testing.T) {
} }
// test pending error // test pending error
app.Cache().Set(core.CacheKeyActiveBackup, "") app.Store().Set(core.StoreKeyActiveBackup, "")
if err := app.RestoreBackup(context.Background(), "test"); err == nil { if err := app.RestoreBackup(context.Background(), "test"); err == nil {
t.Fatal("Expected pending error, got nil") t.Fatal("Expected pending error, got nil")
} }
app.Cache().Remove(core.CacheKeyActiveBackup) app.Store().Remove(core.StoreKeyActiveBackup)
// missing backup // missing backup
if err := app.RestoreBackup(context.Background(), "missing"); err == nil { if err := app.RestoreBackup(context.Background(), "missing"); err == nil {
+1 -1
View File
@@ -20,7 +20,7 @@ func TestBaseAppRefreshSettings(t *testing.T) {
// check if the new settings are saved in the db // check if the new settings are saved in the db
app.ResetEventCalls() app.ResetEventCalls()
if err := app.RefreshSettings(); err != nil { if err := app.RefreshSettings(); err != nil {
t.Fatal("Failed to refresh the settings after delete") t.Fatalf("Failed to refresh the settings after delete: %v", err)
} }
testEventCalls(t, app, map[string]int{ testEventCalls(t, app, map[string]int{
"OnModelBeforeCreate": 1, "OnModelBeforeCreate": 1,
+327 -43
View File
@@ -1,20 +1,30 @@
package core package core
import ( import (
"fmt"
"log/slog"
"os" "os"
"testing" "testing"
"time"
"github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/migrations"
"github.com/pocketbase/pocketbase/migrations/logs"
"github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/logger"
"github.com/pocketbase/pocketbase/tools/mailer" "github.com/pocketbase/pocketbase/tools/mailer"
"github.com/pocketbase/pocketbase/tools/migrate"
) )
func TestNewBaseApp(t *testing.T) { func TestNewBaseApp(t *testing.T) {
const testDataDir = "./pb_base_app_test_data_dir/" const testDataDir = "./pb_base_app_test_data_dir/"
defer os.RemoveAll(testDataDir) defer os.RemoveAll(testDataDir)
app := NewBaseApp(&BaseAppConfig{ app := NewBaseApp(BaseAppConfig{
DataDir: testDataDir, DataDir: testDataDir,
EncryptionEnv: "test_env", EncryptionEnv: "test_env",
IsDebug: true, IsDev: true,
}) })
if app.dataDir != testDataDir { if app.dataDir != testDataDir {
@@ -25,12 +35,12 @@ func TestNewBaseApp(t *testing.T) {
t.Fatalf("expected encryptionEnv test_env, got %q", app.dataDir) t.Fatalf("expected encryptionEnv test_env, got %q", app.dataDir)
} }
if !app.isDebug { if !app.isDev {
t.Fatalf("expected isDebug true, got %v", app.isDebug) t.Fatalf("expected isDev true, got %v", app.isDev)
} }
if app.cache == nil { if app.store == nil {
t.Fatal("expected cache to be set, got nil") t.Fatal("expected store to be set, got nil")
} }
if app.settings == nil { if app.settings == nil {
@@ -46,10 +56,9 @@ func TestBaseAppBootstrap(t *testing.T) {
const testDataDir = "./pb_base_app_test_data_dir/" const testDataDir = "./pb_base_app_test_data_dir/"
defer os.RemoveAll(testDataDir) defer os.RemoveAll(testDataDir)
app := NewBaseApp(&BaseAppConfig{ app := NewBaseApp(BaseAppConfig{
DataDir: testDataDir, DataDir: testDataDir,
EncryptionEnv: "pb_test_env", EncryptionEnv: "pb_test_env",
IsDebug: false,
}) })
defer app.ResetBootstrapState() defer app.ResetBootstrapState()
@@ -57,7 +66,6 @@ func TestBaseAppBootstrap(t *testing.T) {
t.Fatal("Didn't expect the application to be bootstrapped.") t.Fatal("Didn't expect the application to be bootstrapped.")
} }
// bootstrap
if err := app.Bootstrap(); err != nil { if err := app.Bootstrap(); err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -106,6 +114,14 @@ func TestBaseAppBootstrap(t *testing.T) {
t.Fatal("Expected app.settings to be initialized, got nil.") t.Fatal("Expected app.settings to be initialized, got nil.")
} }
if app.logger == nil {
t.Fatal("Expected app.logger to be initialized, got nil.")
}
if _, ok := app.logger.Handler().(*logger.BatchHandler); !ok {
t.Fatal("Expected app.logger handler to be initialized.")
}
// reset // reset
if err := app.ResetBootstrapState(); err != nil { if err := app.ResetBootstrapState(); err != nil {
t.Fatal(err) t.Fatal(err)
@@ -118,20 +134,16 @@ func TestBaseAppBootstrap(t *testing.T) {
if app.logsDao != nil { if app.logsDao != nil {
t.Fatalf("Expected app.logsDao to be nil, got %v.", app.logsDao) t.Fatalf("Expected app.logsDao to be nil, got %v.", app.logsDao)
} }
if app.settings != nil {
t.Fatalf("Expected app.settings to be nil, got %v.", app.settings)
}
} }
func TestBaseAppGetters(t *testing.T) { func TestBaseAppGetters(t *testing.T) {
const testDataDir = "./pb_base_app_test_data_dir/" const testDataDir = "./pb_base_app_test_data_dir/"
defer os.RemoveAll(testDataDir) defer os.RemoveAll(testDataDir)
app := NewBaseApp(&BaseAppConfig{ app := NewBaseApp(BaseAppConfig{
DataDir: testDataDir, DataDir: testDataDir,
EncryptionEnv: "pb_test_env", EncryptionEnv: "pb_test_env",
IsDebug: false, IsDev: true,
}) })
defer app.ResetBootstrapState() defer app.ResetBootstrapState()
@@ -163,16 +175,20 @@ func TestBaseAppGetters(t *testing.T) {
t.Fatalf("Expected app.EncryptionEnv %v, got %v", app.EncryptionEnv(), app.encryptionEnv) t.Fatalf("Expected app.EncryptionEnv %v, got %v", app.EncryptionEnv(), app.encryptionEnv)
} }
if app.isDebug != app.IsDebug() { if app.isDev != app.IsDev() {
t.Fatalf("Expected app.IsDebug %v, got %v", app.IsDebug(), app.isDebug) t.Fatalf("Expected app.IsDev %v, got %v", app.IsDev(), app.isDev)
} }
if app.settings != app.Settings() { if app.settings != app.Settings() {
t.Fatalf("Expected app.Settings %v, got %v", app.Settings(), app.settings) t.Fatalf("Expected app.Settings %v, got %v", app.Settings(), app.settings)
} }
if app.cache != app.Cache() { if app.store != app.Store() {
t.Fatalf("Expected app.Cache %v, got %v", app.Cache(), app.cache) t.Fatalf("Expected app.Store %v, got %v", app.Store(), app.store)
}
if app.logger != app.Logger() {
t.Fatalf("Expected app.Logger %v, got %v", app.Logger(), app.logger)
} }
if app.subscriptionsBroker != app.SubscriptionsBroker() { if app.subscriptionsBroker != app.SubscriptionsBroker() {
@@ -185,14 +201,11 @@ func TestBaseAppGetters(t *testing.T) {
} }
func TestBaseAppNewMailClient(t *testing.T) { func TestBaseAppNewMailClient(t *testing.T) {
const testDataDir = "./pb_base_app_test_data_dir/" app, cleanup, err := initTestBaseApp()
defer os.RemoveAll(testDataDir) if err != nil {
t.Fatal(err)
app := NewBaseApp(&BaseAppConfig{ }
DataDir: testDataDir, defer cleanup()
EncryptionEnv: "pb_test_env",
IsDebug: false,
})
client1 := app.NewMailClient() client1 := app.NewMailClient()
if val, ok := client1.(*mailer.Sendmail); !ok { if val, ok := client1.(*mailer.Sendmail); !ok {
@@ -208,14 +221,11 @@ func TestBaseAppNewMailClient(t *testing.T) {
} }
func TestBaseAppNewFilesystem(t *testing.T) { func TestBaseAppNewFilesystem(t *testing.T) {
const testDataDir = "./pb_base_app_test_data_dir/" app, cleanup, err := initTestBaseApp()
defer os.RemoveAll(testDataDir) if err != nil {
t.Fatal(err)
app := NewBaseApp(&BaseAppConfig{ }
DataDir: testDataDir, defer cleanup()
EncryptionEnv: "pb_test_env",
IsDebug: false,
})
// local // local
local, localErr := app.NewFilesystem() local, localErr := app.NewFilesystem()
@@ -238,14 +248,11 @@ func TestBaseAppNewFilesystem(t *testing.T) {
} }
func TestBaseAppNewBackupsFilesystem(t *testing.T) { func TestBaseAppNewBackupsFilesystem(t *testing.T) {
const testDataDir = "./pb_base_app_test_data_dir/" app, cleanup, err := initTestBaseApp()
defer os.RemoveAll(testDataDir) if err != nil {
t.Fatal(err)
app := NewBaseApp(&BaseAppConfig{ }
DataDir: testDataDir, defer cleanup()
EncryptionEnv: "pb_test_env",
IsDebug: false,
})
// local // local
local, localErr := app.NewBackupsFilesystem() local, localErr := app.NewBackupsFilesystem()
@@ -266,3 +273,280 @@ func TestBaseAppNewBackupsFilesystem(t *testing.T) {
t.Fatalf("Expected nil s3 backups filesystem, got %v", s3) t.Fatalf("Expected nil s3 backups filesystem, got %v", s3)
} }
} }
func TestBaseAppLoggerWrites(t *testing.T) {
app, cleanup, err := initTestBaseApp()
if err != nil {
t.Fatal(err)
}
defer cleanup()
threshold := 200
totalLogs := func(app App, t *testing.T) int {
var total int
err := app.LogsDao().LogQuery().Select("count(*)").Row(&total)
if err != nil {
t.Fatalf("Failed to fetch total logs: %v", err)
}
return total
}
// disabled logs retention
{
app.Settings().Logs.MaxDays = 0
for i := 0; i < threshold+1; i++ {
app.Logger().Error("test")
}
if total := totalLogs(app, t); total != 0 {
t.Fatalf("Expected no logs, got %d", total)
}
}
// test batch logs writes
{
app.Settings().Logs.MaxDays = 1
for i := 0; i < threshold-1; i++ {
app.Logger().Error("test")
}
if total := totalLogs(app, t); total != 0 {
t.Fatalf("Expected no logs, got %d", total)
}
// should trigger batch write
app.Logger().Error("test")
// should be added for the next batch write
app.Logger().Error("test")
if total := totalLogs(app, t); total != threshold {
t.Fatalf("Expected %d logs, got %d", threshold, total)
}
// wait for ~3 secs to check the timer trigger
time.Sleep(3200 * time.Millisecond)
if total := totalLogs(app, t); total != threshold+1 {
t.Fatalf("Expected %d logs, got %d", threshold+1, total)
}
}
}
func TestBaseAppRefreshSettingsLoggerMinLevelEnabled(t *testing.T) {
app, cleanup, err := initTestBaseApp()
if err != nil {
t.Fatal(err)
}
defer cleanup()
handler, ok := app.Logger().Handler().(*logger.BatchHandler)
if !ok {
t.Fatalf("Expected BatchHandler, got %v", app.Logger().Handler())
}
scenarios := []struct {
name string
isDev bool
level int
// level->enabled map
expectations map[int]bool
}{
{
"dev mode",
true,
4,
map[int]bool{
3: true,
4: true,
5: true,
},
},
{
"nondev mode",
false,
4,
map[int]bool{
3: false,
4: true,
5: true,
},
},
}
for _, s := range scenarios {
t.Run(s.name, func(t *testing.T) {
app.isDev = s.isDev
app.Settings().Logs.MinLevel = s.level
if err := app.Dao().SaveSettings(app.Settings()); err != nil {
t.Fatalf("Failed to save settings: %v", err)
}
if err := app.RefreshSettings(); err != nil {
t.Fatalf("Failed to refresh app settings: %v", err)
}
for level, enabled := range s.expectations {
if v := handler.Enabled(nil, slog.Level(level)); v != enabled {
t.Fatalf("Expected level %d Enabled() to be %v, got %v", level, enabled, v)
}
}
})
}
}
func TestBaseAppLoggerLevelDevPrint(t *testing.T) {
app, cleanup, err := initTestBaseApp()
if err != nil {
t.Fatal(err)
}
defer cleanup()
testLogLevel := 4
app.Settings().Logs.MinLevel = testLogLevel
if err := app.Dao().SaveSettings(app.Settings()); err != nil {
t.Fatal(err)
}
scenarios := []struct {
name string
isDev bool
levels []int
printedLevels []int
persistedLevels []int
}{
{
"dev mode",
true,
[]int{testLogLevel - 1, testLogLevel, testLogLevel + 1},
[]int{testLogLevel - 1, testLogLevel, testLogLevel + 1},
[]int{testLogLevel, testLogLevel + 1},
},
{
"nondev mode",
false,
[]int{testLogLevel - 1, testLogLevel, testLogLevel + 1},
[]int{},
[]int{testLogLevel, testLogLevel + 1},
},
}
for _, s := range scenarios {
t.Run(s.name, func(t *testing.T) {
var printedLevels []int
var persistedLevels []int
app.isDev = s.isDev
// trigger slog handler min level refresh
if err := app.RefreshSettings(); err != nil {
t.Fatal(err)
}
// track printed logs
originalPrintLog := printLog
defer func() {
printLog = originalPrintLog
}()
printLog = func(log *logger.Log) {
printedLevels = append(printedLevels, int(log.Level))
}
// track persisted logs
app.LogsDao().AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
l, ok := m.(*models.Log)
if ok {
persistedLevels = append(persistedLevels, l.Level)
}
return nil
}
// write and persist logs
for _, l := range s.levels {
app.Logger().Log(nil, slog.Level(l), "test")
}
handler, ok := app.Logger().Handler().(*logger.BatchHandler)
if !ok {
t.Fatalf("Expected BatchHandler, got %v", app.Logger().Handler())
}
if err := handler.WriteAll(nil); err != nil {
t.Fatalf("Failed to write all logs: %v", err)
}
// check persisted log levels
if len(s.persistedLevels) != len(persistedLevels) {
t.Fatalf("Expected persisted levels \n%v\ngot\n%v", s.persistedLevels, persistedLevels)
}
for _, l := range persistedLevels {
if !list.ExistInSlice(l, s.persistedLevels) {
t.Fatalf("Missing expected persisted level %v in %v", l, persistedLevels)
}
}
// check printed log levels
if len(s.printedLevels) != len(printedLevels) {
t.Fatalf("Expected printed levels \n%v\ngot\n%v", s.printedLevels, printedLevels)
}
for _, l := range printedLevels {
if !list.ExistInSlice(l, s.printedLevels) {
t.Fatalf("Missing expected printed level %v in %v", l, printedLevels)
}
}
})
}
}
// -------------------------------------------------------------------
// note: make sure to call `defer cleanup()` when the app is no longer needed.
func initTestBaseApp() (app *BaseApp, cleanup func(), err error) {
testDataDir, err := os.MkdirTemp("", "test_base_app")
if err != nil {
return nil, nil, err
}
cleanup = func() {
os.RemoveAll(testDataDir)
}
app = NewBaseApp(BaseAppConfig{
DataDir: testDataDir,
})
initErr := func() error {
if err := app.Bootstrap(); err != nil {
return fmt.Errorf("bootstrap error: %w", err)
}
logsRunner, err := migrate.NewRunner(app.LogsDB(), logs.LogsMigrations)
if err != nil {
return fmt.Errorf("logsRunner error: %w", err)
}
if _, err := logsRunner.Up(); err != nil {
return fmt.Errorf("logsRunner migrations execution error: %w", err)
}
dataRunner, err := migrate.NewRunner(app.DB(), migrations.AppMigrations)
if err != nil {
return fmt.Errorf("logsRunner error: %w", err)
}
if _, err := dataRunner.Up(); err != nil {
return fmt.Errorf("dataRunner migrations execution error: %w", err)
}
return nil
}()
if initErr != nil {
cleanup()
return nil, nil, initErr
}
return app, cleanup, nil
}
+72
View File
@@ -0,0 +1,72 @@
package core
// -------------------------------------------------------------------
// This is a small optimization ported from the [ongoing refactoring branch](https://github.com/pocketbase/pocketbase/discussions/4355).
//
// @todo remove after the refactoring is finalized.
// -------------------------------------------------------------------
import (
"strings"
"github.com/pocketbase/pocketbase/models"
)
const storeCachedCollectionsKey = "@cachedCollectionsContext"
func registerCachedCollectionsAppHooks(app App) {
collectionsChangeFunc := func(e *ModelEvent) error {
if _, ok := e.Model.(*models.Collection); !ok {
return nil
}
_ = ReloadCachedCollections(app)
return nil
}
app.OnModelAfterCreate().Add(collectionsChangeFunc)
app.OnModelAfterUpdate().Add(collectionsChangeFunc)
app.OnModelAfterDelete().Add(collectionsChangeFunc)
app.OnBeforeServe().Add(func(e *ServeEvent) error {
_ = ReloadCachedCollections(e.App)
return nil
})
}
func ReloadCachedCollections(app App) error {
collections := []*models.Collection{}
err := app.Dao().CollectionQuery().All(&collections)
if err != nil {
return err
}
app.Store().Set(storeCachedCollectionsKey, collections)
return nil
}
func FindCachedCollectionByNameOrId(app App, nameOrId string) (*models.Collection, error) {
// retrieve from the app cache
// ---
collections, _ := app.Store().Get(storeCachedCollectionsKey).([]*models.Collection)
for _, c := range collections {
if strings.EqualFold(c.Name, nameOrId) || c.Id == nameOrId {
return c, nil
}
}
// retrieve from the database
// ---
found, err := app.Dao().FindCollectionByNameOrId(nameOrId)
if err != nil {
return nil, err
}
err = ReloadCachedCollections(app)
if err != nil {
app.Logger().Warn("Failed to reload collections cache", "error", err)
}
return found, nil
}
+2
View File
@@ -28,6 +28,8 @@ func init() {
PRAGMA journal_size_limit = 200000000; PRAGMA journal_size_limit = 200000000;
PRAGMA synchronous = NORMAL; PRAGMA synchronous = NORMAL;
PRAGMA foreign_keys = ON; PRAGMA foreign_keys = ON;
PRAGMA temp_store = MEMORY;
PRAGMA cache_size = -16000;
`, nil) `, nil)
return err return err
+1 -1
View File
@@ -11,7 +11,7 @@ func connectDB(dbPath string) (*dbx.DB, error) {
// Note: the busy_timeout pragma must be first because // Note: the busy_timeout pragma must be first because
// the connection needs to be set to block on busy before WAL mode // the connection needs to be set to block on busy before WAL mode
// is set in case it hasn't been already set by another connection. // is set in case it hasn't been already set by another connection.
pragmas := "?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)&_pragma=journal_size_limit(200000000)&_pragma=synchronous(NORMAL)&_pragma=foreign_keys(ON)" pragmas := "?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)&_pragma=journal_size_limit(200000000)&_pragma=synchronous(NORMAL)&_pragma=foreign_keys(ON)&_pragma=temp_store(MEMORY)&_pragma=cache_size(-16000)"
db, err := dbx.Open("sqlite", dbPath+pragmas) db, err := dbx.Open("sqlite", dbPath+pragmas)
if err != nil { if err != nil {
+11 -3
View File
@@ -1,6 +1,9 @@
package core package core
import ( import (
"net/http"
"time"
"github.com/labstack/echo/v5" "github.com/labstack/echo/v5"
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
@@ -12,6 +15,7 @@ import (
"github.com/pocketbase/pocketbase/tools/mailer" "github.com/pocketbase/pocketbase/tools/mailer"
"github.com/pocketbase/pocketbase/tools/search" "github.com/pocketbase/pocketbase/tools/search"
"github.com/pocketbase/pocketbase/tools/subscriptions" "github.com/pocketbase/pocketbase/tools/subscriptions"
"golang.org/x/crypto/acme/autocert"
) )
var ( var (
@@ -66,12 +70,15 @@ type BootstrapEvent struct {
} }
type TerminateEvent struct { type TerminateEvent struct {
App App App App
IsRestart bool
} }
type ServeEvent struct { type ServeEvent struct {
App App App App
Router *echo.Echo Router *echo.Echo
Server *http.Server
CertManager *autocert.Manager
} }
type ApiErrorEvent struct { type ApiErrorEvent struct {
@@ -116,6 +123,7 @@ type MailerAdminEvent struct {
type RealtimeConnectEvent struct { type RealtimeConnectEvent struct {
HttpContext echo.Context HttpContext echo.Context
Client subscriptions.Client Client subscriptions.Client
IdleTimeout time.Duration
} }
type RealtimeDisconnectEvent struct { type RealtimeDisconnectEvent struct {
+67
View File
@@ -0,0 +1,67 @@
package core
import (
"fmt"
"log/slog"
"strings"
"github.com/fatih/color"
"github.com/pocketbase/pocketbase/tools/logger"
"github.com/pocketbase/pocketbase/tools/store"
"github.com/spf13/cast"
)
var cachedColors = store.New[*color.Color](nil)
// getColor returns [color.Color] object and cache it (if not already).
func getColor(attrs ...color.Attribute) (c *color.Color) {
cacheKey := fmt.Sprint(attrs)
if c = cachedColors.Get(cacheKey); c == nil {
c = color.New(attrs...)
cachedColors.Set(cacheKey, c)
}
return
}
// printLog prints the provided log to the stderr.
// (note: defined as variable to overwriting in the tests)
var printLog = func(log *logger.Log) {
var str strings.Builder
switch log.Level {
case slog.LevelDebug:
str.WriteString(getColor(color.Bold, color.FgHiBlack).Sprint("DEBUG "))
str.WriteString(getColor(color.FgWhite).Sprint(log.Message))
case slog.LevelInfo:
str.WriteString(getColor(color.Bold, color.FgWhite).Sprint("INFO "))
str.WriteString(getColor(color.FgWhite).Sprint(log.Message))
case slog.LevelWarn:
str.WriteString(getColor(color.Bold, color.FgYellow).Sprint("WARN "))
str.WriteString(getColor(color.FgYellow).Sprint(log.Message))
case slog.LevelError:
str.WriteString(getColor(color.Bold, color.FgRed).Sprint("ERROR "))
str.WriteString(getColor(color.FgRed).Sprint(log.Message))
default:
str.WriteString(getColor(color.Bold, color.FgCyan).Sprintf("[%d] ", log.Level))
str.WriteString(getColor(color.FgCyan).Sprint(log.Message))
}
str.WriteString("\n")
if v, ok := log.Data["type"]; ok && cast.ToString(v) == "request" {
padding := 0
keys := []string{"error", "details"}
for _, k := range keys {
if v := log.Data[k]; v != nil {
str.WriteString(getColor(color.FgHiRed).Sprintf("%s└─ %v", strings.Repeat(" ", padding), v))
str.WriteString("\n")
padding += 3
}
}
} else if len(log.Data) > 0 {
str.WriteString(getColor(color.FgHiBlack).Sprintf("└─ %v", log.Data))
str.WriteString("\n")
}
fmt.Print(str.String())
}
+4 -4
View File
@@ -46,9 +46,9 @@ func (dao *Dao) FindAdminByEmail(email string) (*models.Admin, error) {
return model, nil return model, nil
} }
// FindAdminByToken finds the admin associated with the provided JWT token. // FindAdminByToken finds the admin associated with the provided JWT.
// //
// Returns an error if the JWT token is invalid or expired. // Returns an error if the JWT is invalid or expired.
func (dao *Dao) FindAdminByToken(token string, baseTokenKey string) (*models.Admin, error) { func (dao *Dao) FindAdminByToken(token string, baseTokenKey string) (*models.Admin, error) {
// @todo consider caching the unverified claims // @todo consider caching the unverified claims
unverifiedClaims, err := security.ParseUnverifiedJWT(token) unverifiedClaims, err := security.ParseUnverifiedJWT(token)
@@ -59,7 +59,7 @@ func (dao *Dao) FindAdminByToken(token string, baseTokenKey string) (*models.Adm
// check required claims // check required claims
id, _ := unverifiedClaims["id"].(string) id, _ := unverifiedClaims["id"].(string)
if id == "" { if id == "" {
return nil, errors.New("Missing or invalid token claims.") return nil, errors.New("missing or invalid token claims")
} }
admin, err := dao.FindAdminById(id) admin, err := dao.FindAdminById(id)
@@ -116,7 +116,7 @@ func (dao *Dao) DeleteAdmin(admin *models.Admin) error {
} }
if total == 1 { if total == 1 {
return errors.New("You cannot delete the only existing admin.") return errors.New("you cannot delete the only existing admin")
} }
return dao.Delete(admin) return dao.Delete(admin)
+16
View File
@@ -8,6 +8,8 @@ import (
) )
func TestAdminQuery(t *testing.T) { func TestAdminQuery(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -20,6 +22,8 @@ func TestAdminQuery(t *testing.T) {
} }
func TestFindAdminById(t *testing.T) { func TestFindAdminById(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -47,6 +51,8 @@ func TestFindAdminById(t *testing.T) {
} }
func TestFindAdminByEmail(t *testing.T) { func TestFindAdminByEmail(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -76,6 +82,8 @@ func TestFindAdminByEmail(t *testing.T) {
} }
func TestFindAdminByToken(t *testing.T) { func TestFindAdminByToken(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -131,6 +139,8 @@ func TestFindAdminByToken(t *testing.T) {
} }
func TestTotalAdmins(t *testing.T) { func TestTotalAdmins(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -155,6 +165,8 @@ func TestTotalAdmins(t *testing.T) {
} }
func TestIsAdminEmailUnique(t *testing.T) { func TestIsAdminEmailUnique(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -180,6 +192,8 @@ func TestIsAdminEmailUnique(t *testing.T) {
} }
func TestDeleteAdmin(t *testing.T) { func TestDeleteAdmin(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -225,6 +239,8 @@ func TestDeleteAdmin(t *testing.T) {
} }
func TestSaveAdmin(t *testing.T) { func TestSaveAdmin(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+132 -99
View File
@@ -5,6 +5,7 @@ package daos
import ( import (
"errors" "errors"
"fmt"
"time" "time"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
@@ -17,7 +18,7 @@ func New(db dbx.Builder) *Dao {
return NewMultiDB(db, db) return NewMultiDB(db, db)
} }
// New creates a new Dao instance with the provided dedicated // NewMultiDB creates a new Dao instance with the provided dedicated
// async and sync db builders. // async and sync db builders.
func NewMultiDB(concurrentDB, nonconcurrentDB dbx.Builder) *Dao { func NewMultiDB(concurrentDB, nonconcurrentDB dbx.Builder) *Dao {
return &Dao{ return &Dao{
@@ -29,7 +30,8 @@ func NewMultiDB(concurrentDB, nonconcurrentDB dbx.Builder) *Dao {
} }
// Dao handles various db operations. // Dao handles various db operations.
// Think of Dao as a repository and service layer in one. //
// You can think of Dao as a repository and service layer in one.
type Dao struct { type Dao struct {
// in a transaction both refer to the same *dbx.TX instance // in a transaction both refer to the same *dbx.TX instance
concurrentDB dbx.Builder concurrentDB dbx.Builder
@@ -43,12 +45,13 @@ type Dao struct {
// This field has no effect if an explicit query context is already specified. // This field has no effect if an explicit query context is already specified.
ModelQueryTimeout time.Duration ModelQueryTimeout time.Duration
BeforeCreateFunc func(eventDao *Dao, m models.Model) error // write hooks
AfterCreateFunc func(eventDao *Dao, m models.Model) BeforeCreateFunc func(eventDao *Dao, m models.Model, action func() error) error
BeforeUpdateFunc func(eventDao *Dao, m models.Model) error AfterCreateFunc func(eventDao *Dao, m models.Model) error
AfterUpdateFunc func(eventDao *Dao, m models.Model) BeforeUpdateFunc func(eventDao *Dao, m models.Model, action func() error) error
BeforeDeleteFunc func(eventDao *Dao, m models.Model) error AfterUpdateFunc func(eventDao *Dao, m models.Model) error
AfterDeleteFunc func(eventDao *Dao, m models.Model) BeforeDeleteFunc func(eventDao *Dao, m models.Model, action func() error) error
AfterDeleteFunc func(eventDao *Dao, m models.Model) error
} }
// DB returns the default dao db builder (*dbx.DB or *dbx.TX). // DB returns the default dao db builder (*dbx.DB or *dbx.TX).
@@ -81,6 +84,21 @@ func (dao *Dao) Clone() *Dao {
return &clone return &clone
} }
// WithoutHooks returns a new Dao with the same configuration options
// as the current one, but without create/update/delete hooks.
func (dao *Dao) WithoutHooks() *Dao {
clone := dao.Clone()
clone.BeforeCreateFunc = nil
clone.AfterCreateFunc = nil
clone.BeforeUpdateFunc = nil
clone.AfterUpdateFunc = nil
clone.BeforeDeleteFunc = nil
clone.AfterDeleteFunc = nil
return clone
}
// ModelQuery creates a new preconfigured select query with preset // ModelQuery creates a new preconfigured select query with preset
// SELECT, FROM and other common fields based on the provided model. // SELECT, FROM and other common fields based on the provided model.
func (dao *Dao) ModelQuery(m models.Model) *dbx.SelectQuery { func (dao *Dao) ModelQuery(m models.Model) *dbx.SelectQuery {
@@ -101,9 +119,9 @@ func (dao *Dao) FindById(m models.Model, id string) error {
} }
type afterCallGroup struct { type afterCallGroup struct {
Action string
EventDao *Dao
Model models.Model Model models.Model
EventDao *Dao
Action string
} }
// RunInTransaction wraps fn into a transaction. // RunInTransaction wraps fn into a transaction.
@@ -134,56 +152,69 @@ func (dao *Dao) RunInTransaction(fn func(txDao *Dao) error) error {
txDao := New(tx) txDao := New(tx)
if dao.BeforeCreateFunc != nil { if dao.BeforeCreateFunc != nil {
txDao.BeforeCreateFunc = func(eventDao *Dao, m models.Model) error { txDao.BeforeCreateFunc = func(eventDao *Dao, m models.Model, action func() error) error {
return dao.BeforeCreateFunc(eventDao, m) return dao.BeforeCreateFunc(eventDao, m, action)
} }
} }
if dao.BeforeUpdateFunc != nil { if dao.BeforeUpdateFunc != nil {
txDao.BeforeUpdateFunc = func(eventDao *Dao, m models.Model) error { txDao.BeforeUpdateFunc = func(eventDao *Dao, m models.Model, action func() error) error {
return dao.BeforeUpdateFunc(eventDao, m) return dao.BeforeUpdateFunc(eventDao, m, action)
} }
} }
if dao.BeforeDeleteFunc != nil { if dao.BeforeDeleteFunc != nil {
txDao.BeforeDeleteFunc = func(eventDao *Dao, m models.Model) error { txDao.BeforeDeleteFunc = func(eventDao *Dao, m models.Model, action func() error) error {
return dao.BeforeDeleteFunc(eventDao, m) return dao.BeforeDeleteFunc(eventDao, m, action)
} }
} }
if dao.AfterCreateFunc != nil { if dao.AfterCreateFunc != nil {
txDao.AfterCreateFunc = func(eventDao *Dao, m models.Model) { txDao.AfterCreateFunc = func(eventDao *Dao, m models.Model) error {
afterCalls = append(afterCalls, afterCallGroup{"create", eventDao, m}) afterCalls = append(afterCalls, afterCallGroup{m, eventDao, "create"})
return nil
} }
} }
if dao.AfterUpdateFunc != nil { if dao.AfterUpdateFunc != nil {
txDao.AfterUpdateFunc = func(eventDao *Dao, m models.Model) { txDao.AfterUpdateFunc = func(eventDao *Dao, m models.Model) error {
afterCalls = append(afterCalls, afterCallGroup{"update", eventDao, m}) afterCalls = append(afterCalls, afterCallGroup{m, eventDao, "update"})
return nil
} }
} }
if dao.AfterDeleteFunc != nil { if dao.AfterDeleteFunc != nil {
txDao.AfterDeleteFunc = func(eventDao *Dao, m models.Model) { txDao.AfterDeleteFunc = func(eventDao *Dao, m models.Model) error {
afterCalls = append(afterCalls, afterCallGroup{"delete", eventDao, m}) afterCalls = append(afterCalls, afterCallGroup{m, eventDao, "delete"})
return nil
} }
} }
return fn(txDao) return fn(txDao)
}) })
if txError != nil {
if txError == nil { return txError
// execute after event calls on successful transaction
// (note: using the non-transaction dao to allow following queries in the after hooks)
for _, call := range afterCalls {
switch call.Action {
case "create":
dao.AfterCreateFunc(dao, call.Model)
case "update":
dao.AfterUpdateFunc(dao, call.Model)
case "delete":
dao.AfterDeleteFunc(dao, call.Model)
}
}
} }
return txError // execute after event calls on successful transaction
// (note: using the non-transaction dao to allow following queries in the after hooks)
var errs []error
for _, call := range afterCalls {
var err error
switch call.Action {
case "create":
err = dao.AfterCreateFunc(dao, call.Model)
case "update":
err = dao.AfterUpdateFunc(dao, call.Model)
case "delete":
err = dao.AfterDeleteFunc(dao, call.Model)
}
if err != nil {
errs = append(errs, err)
}
}
if len(errs) > 0 {
return fmt.Errorf("after transaction errors: %w", errors.Join(errs...))
}
return nil
} }
return errors.New("failed to start transaction (unknown dao.NonconcurrentDB() instance)") return errors.New("failed to start transaction (unknown dao.NonconcurrentDB() instance)")
@@ -196,21 +227,23 @@ func (dao *Dao) Delete(m models.Model) error {
} }
return dao.lockRetry(func(retryDao *Dao) error { return dao.lockRetry(func(retryDao *Dao) error {
if retryDao.BeforeDeleteFunc != nil { action := func() error {
if err := retryDao.BeforeDeleteFunc(retryDao, m); err != nil { if err := retryDao.NonconcurrentDB().Model(m).Delete(); err != nil {
return err return err
} }
if retryDao.AfterDeleteFunc != nil {
retryDao.AfterDeleteFunc(retryDao, m)
}
return nil
} }
if err := retryDao.NonconcurrentDB().Model(m).Delete(); err != nil { if retryDao.BeforeDeleteFunc != nil {
return err return retryDao.BeforeDeleteFunc(retryDao, m, action)
} }
if retryDao.AfterDeleteFunc != nil { return action()
retryDao.AfterDeleteFunc(retryDao, m)
}
return nil
}) })
} }
@@ -241,35 +274,35 @@ func (dao *Dao) update(m models.Model) error {
m.RefreshUpdated() m.RefreshUpdated()
action := func() error {
if v, ok := any(m).(models.ColumnValueMapper); ok {
dataMap := v.ColumnValueMap()
_, err := dao.NonconcurrentDB().Update(
m.TableName(),
dataMap,
dbx.HashExp{"id": m.GetId()},
).Execute()
if err != nil {
return err
}
} else if err := dao.NonconcurrentDB().Model(m).Update(); err != nil {
return err
}
if dao.AfterUpdateFunc != nil {
return dao.AfterUpdateFunc(dao, m)
}
return nil
}
if dao.BeforeUpdateFunc != nil { if dao.BeforeUpdateFunc != nil {
if err := dao.BeforeUpdateFunc(dao, m); err != nil { return dao.BeforeUpdateFunc(dao, m, action)
return err
}
} }
if v, ok := any(m).(models.ColumnValueMapper); ok { return action()
dataMap := v.ColumnValueMap()
_, err := dao.NonconcurrentDB().Update(
m.TableName(),
dataMap,
dbx.HashExp{"id": m.GetId()},
).Execute()
if err != nil {
return err
}
} else {
if err := dao.NonconcurrentDB().Model(m).Update(); err != nil {
return err
}
}
if dao.AfterUpdateFunc != nil {
dao.AfterUpdateFunc(dao, m)
}
return nil
} }
func (dao *Dao) create(m models.Model) error { func (dao *Dao) create(m models.Model) error {
@@ -289,36 +322,36 @@ func (dao *Dao) create(m models.Model) error {
m.RefreshUpdated() m.RefreshUpdated()
} }
action := func() error {
if v, ok := any(m).(models.ColumnValueMapper); ok {
dataMap := v.ColumnValueMap()
if _, ok := dataMap["id"]; !ok {
dataMap["id"] = m.GetId()
}
_, err := dao.NonconcurrentDB().Insert(m.TableName(), dataMap).Execute()
if err != nil {
return err
}
} else if err := dao.NonconcurrentDB().Model(m).Insert(); err != nil {
return err
}
// clears the "new" model flag
m.MarkAsNotNew()
if dao.AfterCreateFunc != nil {
return dao.AfterCreateFunc(dao, m)
}
return nil
}
if dao.BeforeCreateFunc != nil { if dao.BeforeCreateFunc != nil {
if err := dao.BeforeCreateFunc(dao, m); err != nil { return dao.BeforeCreateFunc(dao, m, action)
return err
}
} }
if v, ok := any(m).(models.ColumnValueMapper); ok { return action()
dataMap := v.ColumnValueMap()
if _, ok := dataMap["id"]; !ok {
dataMap["id"] = m.GetId()
}
_, err := dao.NonconcurrentDB().Insert(m.TableName(), dataMap).Execute()
if err != nil {
return err
}
} else {
if err := dao.NonconcurrentDB().Model(m).Insert(); err != nil {
return err
}
}
// clears the "new" model flag
m.MarkAsNotNew()
if dao.AfterCreateFunc != nil {
dao.AfterCreateFunc(dao, m)
}
return nil
} }
func (dao *Dao) lockRetry(op func(retryDao *Dao) error) error { func (dao *Dao) lockRetry(op func(retryDao *Dao) error) error {
+9 -1
View File
@@ -2,6 +2,9 @@ package daos
import ( import (
"context" "context"
"database/sql"
"errors"
"fmt"
"strings" "strings"
"time" "time"
@@ -23,9 +26,14 @@ func execLockRetry(timeout time.Duration, maxRetries int) dbx.ExecHookFunc {
q.WithContext(cancelCtx) q.WithContext(cancelCtx)
} }
return baseLockRetry(func(attempt int) error { execErr := baseLockRetry(func(attempt int) error {
return op() return op()
}, maxRetries) }, maxRetries)
if execErr != nil && !errors.Is(execErr, sql.ErrNoRows) {
execErr = fmt.Errorf("%w; failed query: %s", execErr, q.SQL())
}
return execErr
} }
} }
+4
View File
@@ -6,6 +6,8 @@ import (
) )
func TestGetDefaultRetryInterval(t *testing.T) { func TestGetDefaultRetryInterval(t *testing.T) {
t.Parallel()
if i := getDefaultRetryInterval(-1); i.Milliseconds() != 1000 { if i := getDefaultRetryInterval(-1); i.Milliseconds() != 1000 {
t.Fatalf("Expected 1000ms, got %v", i) t.Fatalf("Expected 1000ms, got %v", i)
} }
@@ -20,6 +22,8 @@ func TestGetDefaultRetryInterval(t *testing.T) {
} }
func TestBaseLockRetry(t *testing.T) { func TestBaseLockRetry(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
err error err error
failUntilAttempt int failUntilAttempt int
+131 -47
View File
@@ -49,33 +49,37 @@ func TestDaoClone(t *testing.T) {
dao := daos.NewMultiDB(testApp.Dao().ConcurrentDB(), testApp.Dao().NonconcurrentDB()) dao := daos.NewMultiDB(testApp.Dao().ConcurrentDB(), testApp.Dao().NonconcurrentDB())
dao.MaxLockRetries = 1 dao.MaxLockRetries = 1
dao.ModelQueryTimeout = 2 dao.ModelQueryTimeout = 2
dao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { dao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
hookCalls["BeforeDeleteFunc"]++ hookCalls["BeforeDeleteFunc"]++
return nil return action()
} }
dao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { dao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
hookCalls["BeforeUpdateFunc"]++ hookCalls["BeforeUpdateFunc"]++
return nil return action()
} }
dao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { dao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
hookCalls["BeforeCreateFunc"]++ hookCalls["BeforeCreateFunc"]++
return action()
}
dao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["AfterDeleteFunc"]++
return nil return nil
} }
dao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) { dao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["AfterDeleteFunc"]++
}
dao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) {
hookCalls["AfterUpdateFunc"]++ hookCalls["AfterUpdateFunc"]++
return nil
} }
dao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) { dao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["AfterCreateFunc"]++ hookCalls["AfterCreateFunc"]++
return nil
} }
clone := dao.Clone() clone := dao.Clone()
clone.MaxLockRetries = 3 clone.MaxLockRetries = 3
clone.ModelQueryTimeout = 4 clone.ModelQueryTimeout = 4
clone.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) { clone.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["NewAfterCreateFunc"]++ hookCalls["NewAfterCreateFunc"]++
return nil
} }
if dao.MaxLockRetries == clone.MaxLockRetries { if dao.MaxLockRetries == clone.MaxLockRetries {
@@ -86,16 +90,18 @@ func TestDaoClone(t *testing.T) {
t.Fatal("Expected different ModelQueryTimeout") t.Fatal("Expected different ModelQueryTimeout")
} }
emptyAction := func() error { return nil }
// trigger hooks // trigger hooks
dao.BeforeDeleteFunc(nil, nil) dao.BeforeDeleteFunc(nil, nil, emptyAction)
dao.BeforeUpdateFunc(nil, nil) dao.BeforeUpdateFunc(nil, nil, emptyAction)
dao.BeforeCreateFunc(nil, nil) dao.BeforeCreateFunc(nil, nil, emptyAction)
dao.AfterDeleteFunc(nil, nil) dao.AfterDeleteFunc(nil, nil)
dao.AfterUpdateFunc(nil, nil) dao.AfterUpdateFunc(nil, nil)
dao.AfterCreateFunc(nil, nil) dao.AfterCreateFunc(nil, nil)
clone.BeforeDeleteFunc(nil, nil) clone.BeforeDeleteFunc(nil, nil, emptyAction)
clone.BeforeUpdateFunc(nil, nil) clone.BeforeUpdateFunc(nil, nil, emptyAction)
clone.BeforeCreateFunc(nil, nil) clone.BeforeCreateFunc(nil, nil, emptyAction)
clone.AfterDeleteFunc(nil, nil) clone.AfterDeleteFunc(nil, nil)
clone.AfterUpdateFunc(nil, nil) clone.AfterUpdateFunc(nil, nil)
clone.AfterCreateFunc(nil, nil) clone.AfterCreateFunc(nil, nil)
@@ -120,6 +126,75 @@ func TestDaoClone(t *testing.T) {
} }
} }
func TestDaoWithoutHooks(t *testing.T) {
testApp, _ := tests.NewTestApp()
defer testApp.Cleanup()
hookCalls := map[string]int{}
dao := daos.NewMultiDB(testApp.Dao().ConcurrentDB(), testApp.Dao().NonconcurrentDB())
dao.MaxLockRetries = 1
dao.ModelQueryTimeout = 2
dao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
hookCalls["BeforeDeleteFunc"]++
return action()
}
dao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
hookCalls["BeforeUpdateFunc"]++
return action()
}
dao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
hookCalls["BeforeCreateFunc"]++
return action()
}
dao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["AfterDeleteFunc"]++
return nil
}
dao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["AfterUpdateFunc"]++
return nil
}
dao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
hookCalls["AfterCreateFunc"]++
return nil
}
new := dao.WithoutHooks()
if new.MaxLockRetries != dao.MaxLockRetries {
t.Fatalf("Expected MaxLockRetries %d, got %d", new.Clone().MaxLockRetries, dao.MaxLockRetries)
}
if new.ModelQueryTimeout != dao.ModelQueryTimeout {
t.Fatalf("Expected ModelQueryTimeout %d, got %d", new.Clone().ModelQueryTimeout, dao.ModelQueryTimeout)
}
if new.BeforeDeleteFunc != nil {
t.Fatal("Expected BeforeDeleteFunc to be nil")
}
if new.BeforeUpdateFunc != nil {
t.Fatal("Expected BeforeUpdateFunc to be nil")
}
if new.BeforeCreateFunc != nil {
t.Fatal("Expected BeforeCreateFunc to be nil")
}
if new.AfterDeleteFunc != nil {
t.Fatal("Expected AfterDeleteFunc to be nil")
}
if new.AfterUpdateFunc != nil {
t.Fatal("Expected AfterUpdateFunc to be nil")
}
if new.AfterCreateFunc != nil {
t.Fatal("Expected AfterCreateFunc to be nil")
}
}
func TestDaoModelQuery(t *testing.T) { func TestDaoModelQuery(t *testing.T) {
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -415,12 +490,13 @@ func TestDaoRetryCreate(t *testing.T) {
retryBeforeCreateHookCalls := 0 retryBeforeCreateHookCalls := 0
retryAfterCreateHookCalls := 0 retryAfterCreateHookCalls := 0
retryDao := daos.New(testApp.DB()) retryDao := daos.New(testApp.DB())
retryDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { retryDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
retryBeforeCreateHookCalls++ retryBeforeCreateHookCalls++
return errors.New("database is locked") return errors.New("database is locked")
} }
retryDao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) { retryDao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
retryAfterCreateHookCalls++ retryAfterCreateHookCalls++
return nil
} }
model := &models.Admin{Email: "new@example.com"} model := &models.Admin{Email: "new@example.com"}
@@ -441,7 +517,7 @@ func TestDaoRetryCreate(t *testing.T) {
// with non-locking error // with non-locking error
retryBeforeCreateHookCalls = 0 retryBeforeCreateHookCalls = 0
retryAfterCreateHookCalls = 0 retryAfterCreateHookCalls = 0
retryDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { retryDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
retryBeforeCreateHookCalls++ retryBeforeCreateHookCalls++
return errors.New("non-locking error") return errors.New("non-locking error")
} }
@@ -473,12 +549,13 @@ func TestDaoRetryUpdate(t *testing.T) {
retryBeforeUpdateHookCalls := 0 retryBeforeUpdateHookCalls := 0
retryAfterUpdateHookCalls := 0 retryAfterUpdateHookCalls := 0
retryDao := daos.New(testApp.DB()) retryDao := daos.New(testApp.DB())
retryDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { retryDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
retryBeforeUpdateHookCalls++ retryBeforeUpdateHookCalls++
return errors.New("database is locked") return errors.New("database is locked")
} }
retryDao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) { retryDao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) error {
retryAfterUpdateHookCalls++ retryAfterUpdateHookCalls++
return nil
} }
if err := retryDao.Save(model); err != nil { if err := retryDao.Save(model); err != nil {
@@ -498,7 +575,7 @@ func TestDaoRetryUpdate(t *testing.T) {
// with non-locking error // with non-locking error
retryBeforeUpdateHookCalls = 0 retryBeforeUpdateHookCalls = 0
retryAfterUpdateHookCalls = 0 retryAfterUpdateHookCalls = 0
retryDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { retryDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
retryBeforeUpdateHookCalls++ retryBeforeUpdateHookCalls++
return errors.New("non-locking error") return errors.New("non-locking error")
} }
@@ -524,12 +601,13 @@ func TestDaoRetryDelete(t *testing.T) {
retryBeforeDeleteHookCalls := 0 retryBeforeDeleteHookCalls := 0
retryAfterDeleteHookCalls := 0 retryAfterDeleteHookCalls := 0
retryDao := daos.New(testApp.DB()) retryDao := daos.New(testApp.DB())
retryDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { retryDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
retryBeforeDeleteHookCalls++ retryBeforeDeleteHookCalls++
return errors.New("database is locked") return errors.New("database is locked")
} }
retryDao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) { retryDao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) error {
retryAfterDeleteHookCalls++ retryAfterDeleteHookCalls++
return nil
} }
model, _ := retryDao.FindAdminByEmail("test@example.com") model, _ := retryDao.FindAdminByEmail("test@example.com")
@@ -550,7 +628,7 @@ func TestDaoRetryDelete(t *testing.T) {
// with non-locking error // with non-locking error
retryBeforeDeleteHookCalls = 0 retryBeforeDeleteHookCalls = 0
retryAfterDeleteHookCalls = 0 retryAfterDeleteHookCalls = 0
retryDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { retryDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
retryBeforeDeleteHookCalls++ retryBeforeDeleteHookCalls++
return errors.New("non-locking error") return errors.New("non-locking error")
} }
@@ -577,13 +655,13 @@ func TestDaoBeforeHooksError(t *testing.T) {
baseDao := testApp.Dao() baseDao := testApp.Dao()
baseDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
return errors.New("before_create") return errors.New("before_create")
} }
baseDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
return errors.New("before_update") return errors.New("before_update")
} }
baseDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
return errors.New("before_delete") return errors.New("before_delete")
} }
@@ -622,27 +700,30 @@ func TestDaoTransactionHooksCallsOnFailure(t *testing.T) {
baseDao := testApp.Dao() baseDao := testApp.Dao()
baseDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
beforeCreateFuncCalls++ beforeCreateFuncCalls++
return nil return action()
} }
baseDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
beforeUpdateFuncCalls++ beforeUpdateFuncCalls++
return nil return action()
} }
baseDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
beforeDeleteFuncCalls++ beforeDeleteFuncCalls++
return nil return action()
} }
baseDao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) { baseDao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
afterCreateFuncCalls++ afterCreateFuncCalls++
return nil
} }
baseDao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) { baseDao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) error {
afterUpdateFuncCalls++ afterUpdateFuncCalls++
return nil
} }
baseDao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) { baseDao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) error {
afterDeleteFuncCalls++ afterDeleteFuncCalls++
return nil
} }
existingModel, _ := testApp.Dao().FindAdminByEmail("test@example.com") existingModel, _ := testApp.Dao().FindAdminByEmail("test@example.com")
@@ -710,27 +791,30 @@ func TestDaoTransactionHooksCallsOnSuccess(t *testing.T) {
baseDao := testApp.Dao() baseDao := testApp.Dao()
baseDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeCreateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
beforeCreateFuncCalls++ beforeCreateFuncCalls++
return nil return action()
} }
baseDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeUpdateFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
beforeUpdateFuncCalls++ beforeUpdateFuncCalls++
return nil return action()
} }
baseDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model) error { baseDao.BeforeDeleteFunc = func(eventDao *daos.Dao, m models.Model, action func() error) error {
beforeDeleteFuncCalls++ beforeDeleteFuncCalls++
return nil return action()
} }
baseDao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) { baseDao.AfterCreateFunc = func(eventDao *daos.Dao, m models.Model) error {
afterCreateFuncCalls++ afterCreateFuncCalls++
return nil
} }
baseDao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) { baseDao.AfterUpdateFunc = func(eventDao *daos.Dao, m models.Model) error {
afterUpdateFuncCalls++ afterUpdateFuncCalls++
return nil
} }
baseDao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) { baseDao.AfterDeleteFunc = func(eventDao *daos.Dao, m models.Model) error {
afterDeleteFuncCalls++ afterDeleteFuncCalls++
return nil
} }
existingModel, _ := testApp.Dao().FindAdminByEmail("test@example.com") existingModel, _ := testApp.Dao().FindAdminByEmail("test@example.com")
+64 -10
View File
@@ -121,7 +121,7 @@ func (dao *Dao) FindCollectionReferences(collection *models.Collection, excludeI
// - is referenced as part of a relation field in another collection // - is referenced as part of a relation field in another collection
func (dao *Dao) DeleteCollection(collection *models.Collection) error { func (dao *Dao) DeleteCollection(collection *models.Collection) error {
if collection.System { if collection.System {
return fmt.Errorf("System collection %q cannot be deleted.", collection.Name) return fmt.Errorf("system collection %q cannot be deleted", collection.Name)
} }
// ensure that there aren't any existing references. // ensure that there aren't any existing references.
@@ -135,7 +135,7 @@ func (dao *Dao) DeleteCollection(collection *models.Collection) error {
for ref := range result { for ref := range result {
names = append(names, ref.Name) names = append(names, ref.Name)
} }
return fmt.Errorf("The collection %q has external relation field references (%s).", collection.Name, strings.Join(names, ", ")) return fmt.Errorf("the collection %q has external relation field references (%s)", collection.Name, strings.Join(names, ", "))
} }
return dao.RunInTransaction(func(txDao *Dao) error { return dao.RunInTransaction(func(txDao *Dao) error {
@@ -152,7 +152,7 @@ func (dao *Dao) DeleteCollection(collection *models.Collection) error {
// trigger views resave to check for dependencies // trigger views resave to check for dependencies
if err := txDao.resaveViewsWithChangedSchema(collection.Id); err != nil { if err := txDao.resaveViewsWithChangedSchema(collection.Id); err != nil {
return fmt.Errorf("The collection has a view dependency - %w", err) return fmt.Errorf("the collection has a view dependency - %w", err)
} }
return txDao.Delete(collection) return txDao.Delete(collection)
@@ -162,8 +162,8 @@ func (dao *Dao) DeleteCollection(collection *models.Collection) error {
// SaveCollection persists the provided Collection model and updates // SaveCollection persists the provided Collection model and updates
// its related records table schema. // its related records table schema.
// //
// If collecction.IsNew() is true, the method will perform a create, otherwise an update. // If collection.IsNew() is true, the method will perform a create, otherwise an update.
// To explicitly mark a collection for update you can use collecction.MarkAsNotNew(). // To explicitly mark a collection for update you can use collection.MarkAsNotNew().
func (dao *Dao) SaveCollection(collection *models.Collection) error { func (dao *Dao) SaveCollection(collection *models.Collection) error {
var oldCollection *models.Collection var oldCollection *models.Collection
@@ -227,7 +227,7 @@ func (dao *Dao) ImportCollections(
afterSync func(txDao *Dao, mappedImported, mappedExisting map[string]*models.Collection) error, afterSync func(txDao *Dao, mappedImported, mappedExisting map[string]*models.Collection) error,
) error { ) error {
if len(importedCollections) == 0 { if len(importedCollections) == 0 {
return errors.New("No collections to import") return errors.New("no collections to import")
} }
return dao.RunInTransaction(func(txDao *Dao) error { return dao.RunInTransaction(func(txDao *Dao) error {
@@ -263,11 +263,11 @@ func (dao *Dao) ImportCollections(
// extend existing schema // extend existing schema
if !deleteMissing { if !deleteMissing {
schema, _ := existing.Schema.Clone() schemaClone, _ := existing.Schema.Clone()
for _, f := range imported.Schema.Fields() { for _, f := range imported.Schema.Fields() {
schema.AddField(f) // add or replace schemaClone.AddField(f) // add or replace
} }
imported.Schema = *schema imported.Schema = *schemaClone
} }
} else { } else {
imported.MarkAsNew() imported.MarkAsNew()
@@ -285,7 +285,7 @@ func (dao *Dao) ImportCollections(
} }
if existing.System { if existing.System {
return fmt.Errorf("System collection %q cannot be deleted.", existing.Name) return fmt.Errorf("system collection %q cannot be deleted", existing.Name)
} }
// delete the related records table or view // delete the related records table or view
@@ -378,6 +378,12 @@ func (dao *Dao) saveViewCollection(newCollection, oldCollection *models.Collecti
} }
} }
// wrap view query if necessary
query, err = txDao.normalizeViewQueryId(query)
if err != nil {
return fmt.Errorf("failed to normalize view query id: %w", err)
}
// (re)create the view // (re)create the view
if err := txDao.SaveView(newCollection.Name, query); err != nil { if err := txDao.SaveView(newCollection.Name, query); err != nil {
return err return err
@@ -389,6 +395,54 @@ func (dao *Dao) saveViewCollection(newCollection, oldCollection *models.Collecti
}) })
} }
// @todo consider removing once custom id types are supported
//
// normalizeViewQueryId wraps (if necessary) the provided view query
// with a subselect to ensure that the id column is a text since
// currently we don't support non-string model ids
// (see https://github.com/pocketbase/pocketbase/issues/3110).
func (dao *Dao) normalizeViewQueryId(query string) (string, error) {
query = strings.Trim(strings.TrimSpace(query), ";")
parsed, err := dao.parseQueryToFields(query)
if err != nil {
return "", err
}
needWrapping := true
idField := parsed[schema.FieldNameId]
if idField != nil && idField.field != nil &&
idField.field.Type != schema.FieldTypeJson &&
idField.field.Type != schema.FieldTypeNumber &&
idField.field.Type != schema.FieldTypeBool {
needWrapping = false
}
if !needWrapping {
return query, nil // no changes needed
}
// raw parse to preserve the columns order
rawParsed := new(identifiersParser)
if err := rawParsed.parse(query); err != nil {
return "", err
}
columns := make([]string, 0, len(rawParsed.columns))
for _, col := range rawParsed.columns {
if col.alias == schema.FieldNameId {
columns = append(columns, fmt.Sprintf("cast([[%s]] as text) [[%s]]", col.alias, col.alias))
} else {
columns = append(columns, "[["+col.alias+"]]")
}
}
query = fmt.Sprintf("SELECT %s FROM (%s)", strings.Join(columns, ","), query)
return query, nil
}
// resaveViewsWithChangedSchema updates all view collections with changed schemas. // resaveViewsWithChangedSchema updates all view collections with changed schemas.
func (dao *Dao) resaveViewsWithChangedSchema(excludeIds ...string) error { func (dao *Dao) resaveViewsWithChangedSchema(excludeIds ...string) error {
collections, err := dao.FindCollectionsByType(models.CollectionTypeView) collections, err := dao.FindCollectionsByType(models.CollectionTypeView)
+165 -30
View File
@@ -16,6 +16,8 @@ import (
) )
func TestCollectionQuery(t *testing.T) { func TestCollectionQuery(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -28,6 +30,8 @@ func TestCollectionQuery(t *testing.T) {
} }
func TestFindCollectionsByType(t *testing.T) { func TestFindCollectionsByType(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -63,6 +67,8 @@ func TestFindCollectionsByType(t *testing.T) {
} }
func TestFindCollectionByNameOrId(t *testing.T) { func TestFindCollectionByNameOrId(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -92,6 +98,8 @@ func TestFindCollectionByNameOrId(t *testing.T) {
} }
func TestIsCollectionNameUnique(t *testing.T) { func TestIsCollectionNameUnique(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -116,6 +124,8 @@ func TestIsCollectionNameUnique(t *testing.T) {
} }
func TestFindCollectionReferences(t *testing.T) { func TestFindCollectionReferences(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -143,9 +153,11 @@ func TestFindCollectionReferences(t *testing.T) {
"rel_one_no_cascade", "rel_one_no_cascade",
"rel_one_no_cascade_required", "rel_one_no_cascade_required",
"rel_one_cascade", "rel_one_cascade",
"rel_one_unique",
"rel_many_no_cascade", "rel_many_no_cascade",
"rel_many_no_cascade_required", "rel_many_no_cascade_required",
"rel_many_cascade", "rel_many_cascade",
"rel_many_unique",
} }
for col, fields := range result { for col, fields := range result {
@@ -164,6 +176,8 @@ func TestFindCollectionReferences(t *testing.T) {
} }
func TestDeleteCollection(t *testing.T) { func TestDeleteCollection(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -210,11 +224,12 @@ func TestDeleteCollection(t *testing.T) {
{colUnsaved, true}, {colUnsaved, true},
{colReferenced, true}, {colReferenced, true},
{colSystem, true}, {colSystem, true},
{colBase, true}, // depend on view1, view2 and view2
{colView1, true}, // view2 depend on it {colView1, true}, // view2 depend on it
{colView2, false}, {colView2, false},
{colView1, false}, // no longer has dependent collections {colView1, false}, // no longer has dependent collections
{colBase, false}, {colBase, false}, // no longer has dependent views
{colAuth, false}, // should delete also its related external auths {colAuth, false}, // should delete also its related external auths
} }
for i, s := range scenarios { for i, s := range scenarios {
@@ -250,6 +265,8 @@ func TestDeleteCollection(t *testing.T) {
} }
func TestSaveCollectionCreate(t *testing.T) { func TestSaveCollectionCreate(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -296,6 +313,8 @@ func TestSaveCollectionCreate(t *testing.T) {
} }
func TestSaveCollectionUpdate(t *testing.T) { func TestSaveCollectionUpdate(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -335,6 +354,8 @@ func TestSaveCollectionUpdate(t *testing.T) {
// indirect update of a field used in view should cause view(s) update // indirect update of a field used in view should cause view(s) update
func TestSaveCollectionIndirectViewsUpdate(t *testing.T) { func TestSaveCollectionIndirectViewsUpdate(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -393,8 +414,121 @@ func TestSaveCollectionIndirectViewsUpdate(t *testing.T) {
} }
} }
func TestSaveCollectionViewWrapping(t *testing.T) {
t.Parallel()
viewName := "test_wrapping"
scenarios := []struct {
name string
query string
expected string
}{
{
"no wrapping - text field",
"select text as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select text as id, bool from demo1)",
},
{
"no wrapping - id field",
"select text as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select text as id, bool from demo1)",
},
{
"no wrapping - relation field",
"select rel_one as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select rel_one as id, bool from demo1)",
},
{
"no wrapping - select field",
"select select_many as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select select_many as id, bool from demo1)",
},
{
"no wrapping - email field",
"select email as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select email as id, bool from demo1)",
},
{
"no wrapping - datetime field",
"select datetime as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select datetime as id, bool from demo1)",
},
{
"no wrapping - url field",
"select url as id, bool from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select url as id, bool from demo1)",
},
{
"wrapping - bool field",
"select bool as id, text as txt, url from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (SELECT cast(`id` as text) `id`,`txt`,`url` FROM (select bool as id, text as txt, url from demo1))",
},
{
"wrapping - bool field (different order)",
"select text as txt, url, bool as id from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (SELECT `txt`,`url`,cast(`id` as text) `id` FROM (select text as txt, url, bool as id from demo1))",
},
{
"wrapping - json field",
"select json as id, text, url from demo1",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (SELECT cast(`id` as text) `id`,`text`,`url` FROM (select json as id, text, url from demo1))",
},
{
"wrapping - numeric id",
"select 1 as id",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (SELECT cast(`id` as text) `id` FROM (select 1 as id))",
},
{
"wrapping - expresion",
"select ('test') as id",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (SELECT cast(`id` as text) `id` FROM (select ('test') as id))",
},
{
"no wrapping - cast as text",
"select cast('test' as text) as id",
"CREATE VIEW `test_wrapping` AS SELECT * FROM (select cast('test' as text) as id)",
},
}
for _, s := range scenarios {
t.Run(s.name, func(t *testing.T) {
app, _ := tests.NewTestApp()
defer app.Cleanup()
collection := &models.Collection{
Name: viewName,
Type: models.CollectionTypeView,
Options: types.JsonMap{
"query": s.query,
},
}
err := app.Dao().SaveCollection(collection)
if err != nil {
t.Fatal(err)
}
var sql string
rowErr := app.Dao().DB().NewQuery("SELECT sql FROM sqlite_master WHERE type='view' AND name={:name}").
Bind(dbx.Params{"name": viewName}).
Row(&sql)
if rowErr != nil {
t.Fatalf("Failed to retrieve view sql: %v", rowErr)
}
if sql != s.expected {
t.Fatalf("Expected query \n%v, \ngot \n%v", s.expected, sql)
}
})
}
}
func TestImportCollections(t *testing.T) { func TestImportCollections(t *testing.T) {
totalCollections := 10 t.Parallel()
totalCollections := 11
scenarios := []struct { scenarios := []struct {
name string name string
@@ -624,7 +758,7 @@ func TestImportCollections(t *testing.T) {
"demo1": 15, "demo1": 15,
"demo2": 2, "demo2": 2,
"demo3": 2, "demo3": 2,
"demo4": 11, "demo4": 13,
"demo5": 6, "demo5": 6,
"new_import": 1, "new_import": 1,
} }
@@ -642,37 +776,38 @@ func TestImportCollections(t *testing.T) {
}, },
} }
for _, scenario := range scenarios { for _, s := range scenarios {
testApp, _ := tests.NewTestApp() t.Run(s.name, func(t *testing.T) {
defer testApp.Cleanup() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup()
importedCollections := []*models.Collection{} importedCollections := []*models.Collection{}
// load data // load data
loadErr := json.Unmarshal([]byte(scenario.jsonData), &importedCollections) loadErr := json.Unmarshal([]byte(s.jsonData), &importedCollections)
if loadErr != nil { if loadErr != nil {
t.Fatalf("[%s] Failed to load data: %v", scenario.name, loadErr) t.Fatalf("Failed to load data: %v", loadErr)
continue }
}
err := testApp.Dao().ImportCollections(importedCollections, scenario.deleteMissing, scenario.beforeRecordsSync) err := testApp.Dao().ImportCollections(importedCollections, s.deleteMissing, s.beforeRecordsSync)
hasErr := err != nil hasErr := err != nil
if hasErr != scenario.expectError { if hasErr != s.expectError {
t.Errorf("[%s] Expected hasErr to be %v, got %v (%v)", scenario.name, scenario.expectError, hasErr, err) t.Fatalf("Expected hasErr to be %v, got %v (%v)", s.expectError, hasErr, err)
} }
// check collections count // check collections count
collections := []*models.Collection{} collections := []*models.Collection{}
if err := testApp.Dao().CollectionQuery().All(&collections); err != nil { if err := testApp.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err) t.Fatal(err)
} }
if len(collections) != scenario.expectCollectionsCount { if len(collections) != s.expectCollectionsCount {
t.Errorf("[%s] Expected %d collections, got %d", scenario.name, scenario.expectCollectionsCount, len(collections)) t.Fatalf("Expected %d collections, got %d", s.expectCollectionsCount, len(collections))
} }
if scenario.afterTestFunc != nil { if s.afterTestFunc != nil {
scenario.afterTestFunc(testApp, collections) s.afterTestFunc(testApp, collections)
} }
})
} }
} }
+18 -21
View File
@@ -32,27 +32,6 @@ func (dao *Dao) FindAllExternalAuthsByRecord(authRecord *models.Record) ([]*mode
return auths, nil return auths, nil
} }
// FindExternalAuthByProvider returns the first available
// ExternalAuth model for the specified provider and providerId.
func (dao *Dao) FindExternalAuthByProvider(provider, providerId string) (*models.ExternalAuth, error) {
model := &models.ExternalAuth{}
err := dao.ExternalAuthQuery().
AndWhere(dbx.Not(dbx.HashExp{"providerId": ""})). // exclude empty providerIds
AndWhere(dbx.HashExp{
"provider": provider,
"providerId": providerId,
}).
Limit(1).
One(model)
if err != nil {
return nil, err
}
return model, nil
}
// FindExternalAuthByRecordAndProvider returns the first available // FindExternalAuthByRecordAndProvider returns the first available
// ExternalAuth model for the specified record data and provider. // ExternalAuth model for the specified record data and provider.
func (dao *Dao) FindExternalAuthByRecordAndProvider(authRecord *models.Record, provider string) (*models.ExternalAuth, error) { func (dao *Dao) FindExternalAuthByRecordAndProvider(authRecord *models.Record, provider string) (*models.ExternalAuth, error) {
@@ -74,6 +53,24 @@ func (dao *Dao) FindExternalAuthByRecordAndProvider(authRecord *models.Record, p
return model, nil return model, nil
} }
// FindFirstExternalAuthByExpr returns the first available
// ExternalAuth model that satisfies the non-nil expression.
func (dao *Dao) FindFirstExternalAuthByExpr(expr dbx.Expression) (*models.ExternalAuth, error) {
model := &models.ExternalAuth{}
err := dao.ExternalAuthQuery().
AndWhere(dbx.Not(dbx.HashExp{"providerId": ""})). // exclude empty providerIds
AndWhere(expr).
Limit(1).
One(model)
if err != nil {
return nil, err
}
return model, nil
}
// SaveExternalAuth upserts the provided ExternalAuth model. // SaveExternalAuth upserts the provided ExternalAuth model.
func (dao *Dao) SaveExternalAuth(model *models.ExternalAuth) error { func (dao *Dao) SaveExternalAuth(model *models.ExternalAuth) error {
// extra check the model data in case the provider's API response // extra check the model data in case the provider's API response
+26 -11
View File
@@ -3,11 +3,14 @@ package daos_test
import ( import (
"testing" "testing"
"github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tests" "github.com/pocketbase/pocketbase/tests"
) )
func TestExternalAuthQuery(t *testing.T) { func TestExternalAuthQuery(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -20,6 +23,8 @@ func TestExternalAuthQuery(t *testing.T) {
} }
func TestFindAllExternalAuthsByRecord(t *testing.T) { func TestFindAllExternalAuthsByRecord(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -56,25 +61,25 @@ func TestFindAllExternalAuthsByRecord(t *testing.T) {
} }
} }
func TestFindExternalAuthByProvider(t *testing.T) { func TestFindFirstExternalAuthByExpr(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
scenarios := []struct { scenarios := []struct {
provider string expr dbx.Expression
providerId string
expectedId string expectedId string
}{ }{
{"", "", ""}, {dbx.HashExp{"provider": "github", "providerId": ""}, ""},
{"github", "", ""}, {dbx.HashExp{"provider": "github", "providerId": "id1"}, ""},
{"github", "id1", ""}, {dbx.HashExp{"provider": "github", "providerId": "id2"}, ""},
{"github", "id2", ""}, {dbx.HashExp{"provider": "google", "providerId": "test123"}, "clmflokuq1xl341"},
{"google", "test123", "clmflokuq1xl341"}, {dbx.HashExp{"provider": "gitlab", "providerId": "test123"}, "dlmflokuq1xl342"},
{"gitlab", "test123", "dlmflokuq1xl342"},
} }
for i, s := range scenarios { for i, s := range scenarios {
auth, err := app.Dao().FindExternalAuthByProvider(s.provider, s.providerId) auth, err := app.Dao().FindFirstExternalAuthByExpr(s.expr)
hasErr := err != nil hasErr := err != nil
expectErr := s.expectedId == "" expectErr := s.expectedId == ""
@@ -90,6 +95,8 @@ func TestFindExternalAuthByProvider(t *testing.T) {
} }
func TestFindExternalAuthByRecordAndProvider(t *testing.T) { func TestFindExternalAuthByRecordAndProvider(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -126,6 +133,8 @@ func TestFindExternalAuthByRecordAndProvider(t *testing.T) {
} }
func TestSaveExternalAuth(t *testing.T) { func TestSaveExternalAuth(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -147,7 +156,11 @@ func TestSaveExternalAuth(t *testing.T) {
} }
// check if it was really saved // check if it was really saved
foundAuth, err := app.Dao().FindExternalAuthByProvider("test", "test_id") foundAuth, err := app.Dao().FindFirstExternalAuthByExpr(dbx.HashExp{
"collectionId": "v851q4r790rhknl",
"provider": "test",
"providerId": "test_id",
})
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -158,6 +171,8 @@ func TestSaveExternalAuth(t *testing.T) {
} }
func TestDeleteExternalAuth(t *testing.T) { func TestDeleteExternalAuth(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+67
View File
@@ -0,0 +1,67 @@
package daos
import (
"time"
"github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tools/types"
)
// LogQuery returns a new Log select query.
func (dao *Dao) LogQuery() *dbx.SelectQuery {
return dao.ModelQuery(&models.Log{})
}
// FindLogById finds a single Log entry by its id.
func (dao *Dao) FindLogById(id string) (*models.Log, error) {
model := &models.Log{}
err := dao.LogQuery().
AndWhere(dbx.HashExp{"id": id}).
Limit(1).
One(model)
if err != nil {
return nil, err
}
return model, nil
}
type LogsStatsItem struct {
Total int `db:"total" json:"total"`
Date types.DateTime `db:"date" json:"date"`
}
// LogsStats returns hourly grouped requests logs statistics.
func (dao *Dao) LogsStats(expr dbx.Expression) ([]*LogsStatsItem, error) {
result := []*LogsStatsItem{}
query := dao.LogQuery().
Select("count(id) as total", "strftime('%Y-%m-%d %H:00:00', created) as date").
GroupBy("date")
if expr != nil {
query.AndWhere(expr)
}
err := query.All(&result)
return result, err
}
// DeleteOldLogs delete all requests that are created before createdBefore.
func (dao *Dao) DeleteOldLogs(createdBefore time.Time) error {
formattedDate := createdBefore.UTC().Format(types.DefaultDateLayout)
expr := dbx.NewExp("[[created]] <= {:date}", dbx.Params{"date": formattedDate})
_, err := dao.NonconcurrentDB().Delete((&models.Log{}).TableName(), expr).Execute()
return err
}
// SaveLog upserts the provided Log model.
func (dao *Dao) SaveLog(log *models.Log) error {
return dao.Save(log)
}
+42 -32
View File
@@ -11,23 +11,27 @@ import (
"github.com/pocketbase/pocketbase/tools/types" "github.com/pocketbase/pocketbase/tools/types"
) )
func TestRequestQuery(t *testing.T) { func TestLogQuery(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
expected := "SELECT {{_requests}}.* FROM `_requests`" expected := "SELECT {{_logs}}.* FROM `_logs`"
sql := app.Dao().RequestQuery().Build().SQL() sql := app.Dao().LogQuery().Build().SQL()
if sql != expected { if sql != expected {
t.Errorf("Expected sql %s, got %s", expected, sql) t.Errorf("Expected sql %s, got %s", expected, sql)
} }
} }
func TestFindRequestById(t *testing.T) { func TestFindLogById(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
tests.MockRequestLogsData(app) tests.MockLogsData(app)
scenarios := []struct { scenarios := []struct {
id string id string
@@ -40,7 +44,7 @@ func TestFindRequestById(t *testing.T) {
} }
for i, scenario := range scenarios { for i, scenario := range scenarios {
admin, err := app.LogsDao().FindRequestById(scenario.id) admin, err := app.LogsDao().FindLogById(scenario.id)
hasErr := err != nil hasErr := err != nil
if hasErr != scenario.expectError { if hasErr != scenario.expectError {
@@ -53,17 +57,19 @@ func TestFindRequestById(t *testing.T) {
} }
} }
func TestRequestsStats(t *testing.T) { func TestLogsStats(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
tests.MockRequestLogsData(app) tests.MockLogsData(app)
expected := `[{"total":1,"date":"2022-05-01 10:00:00.000Z"},{"total":1,"date":"2022-05-02 10:00:00.000Z"}]` expected := `[{"total":1,"date":"2022-05-01 10:00:00.000Z"},{"total":1,"date":"2022-05-02 10:00:00.000Z"}]`
now := time.Now().UTC().Format(types.DefaultDateLayout) now := time.Now().UTC().Format(types.DefaultDateLayout)
exp := dbx.NewExp("[[created]] <= {:date}", dbx.Params{"date": now}) exp := dbx.NewExp("[[created]] <= {:date}", dbx.Params{"date": now})
result, err := app.LogsDao().RequestsStats(exp) result, err := app.LogsDao().LogsStats(exp)
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
} }
@@ -74,20 +80,22 @@ func TestRequestsStats(t *testing.T) {
} }
} }
func TestDeleteOldRequests(t *testing.T) { func TestDeleteOldLogs(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
tests.MockRequestLogsData(app) tests.MockLogsData(app)
scenarios := []struct { scenarios := []struct {
date string date string
expectedTotal int expectedTotal int
}{ }{
{"2022-01-01 10:00:00.000Z", 2}, // no requests to delete before that time {"2022-01-01 10:00:00.000Z", 2}, // no logs to delete before that time
{"2022-05-01 11:00:00.000Z", 1}, // only 1 request should have left {"2022-05-01 11:00:00.000Z", 1}, // only 1 log should have left
{"2022-05-03 11:00:00.000Z", 0}, // no more requests should have left {"2022-05-03 11:00:00.000Z", 0}, // no more logs should have left
{"2022-05-04 11:00:00.000Z", 0}, // no more requests should have left {"2022-05-04 11:00:00.000Z", 0}, // no more logs should have left
} }
for i, scenario := range scenarios { for i, scenario := range scenarios {
@@ -96,53 +104,55 @@ func TestDeleteOldRequests(t *testing.T) {
t.Errorf("(%d) Date error %v", i, dateErr) t.Errorf("(%d) Date error %v", i, dateErr)
} }
deleteErr := app.LogsDao().DeleteOldRequests(date) deleteErr := app.LogsDao().DeleteOldLogs(date)
if deleteErr != nil { if deleteErr != nil {
t.Errorf("(%d) Delete error %v", i, deleteErr) t.Errorf("(%d) Delete error %v", i, deleteErr)
} }
// check total remaining requests // check total remaining logs
var total int var total int
countErr := app.LogsDao().RequestQuery().Select("count(*)").Row(&total) countErr := app.LogsDao().LogQuery().Select("count(*)").Row(&total)
if countErr != nil { if countErr != nil {
t.Errorf("(%d) Count error %v", i, countErr) t.Errorf("(%d) Count error %v", i, countErr)
} }
if total != scenario.expectedTotal { if total != scenario.expectedTotal {
t.Errorf("(%d) Expected %d remaining requests, got %d", i, scenario.expectedTotal, total) t.Errorf("(%d) Expected %d remaining logs, got %d", i, scenario.expectedTotal, total)
} }
} }
} }
func TestSaveRequest(t *testing.T) { func TestSaveLog(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
tests.MockRequestLogsData(app) tests.MockLogsData(app)
// create new request // create new log
newRequest := &models.Request{} newLog := &models.Log{}
newRequest.Method = "get" newLog.Level = -4
newRequest.Meta = types.JsonMap{} newLog.Data = types.JsonMap{}
createErr := app.LogsDao().SaveRequest(newRequest) createErr := app.LogsDao().SaveLog(newLog)
if createErr != nil { if createErr != nil {
t.Fatal(createErr) t.Fatal(createErr)
} }
// check if it was really created // check if it was really created
existingRequest, fetchErr := app.LogsDao().FindRequestById(newRequest.Id) existingLog, fetchErr := app.LogsDao().FindLogById(newLog.Id)
if fetchErr != nil { if fetchErr != nil {
t.Fatal(fetchErr) t.Fatal(fetchErr)
} }
existingRequest.Method = "post" existingLog.Level = 4
updateErr := app.LogsDao().SaveRequest(existingRequest) updateErr := app.LogsDao().SaveLog(existingLog)
if updateErr != nil { if updateErr != nil {
t.Fatal(updateErr) t.Fatal(updateErr)
} }
// refresh instance to check if it was really updated // refresh instance to check if it was really updated
existingRequest, _ = app.LogsDao().FindRequestById(existingRequest.Id) existingLog, _ = app.LogsDao().FindLogById(existingLog.Id)
if existingRequest.Method != "post" { if existingLog.Level != 4 {
t.Fatalf("Expected request method to be %s, got %s", "post", existingRequest.Method) t.Fatalf("Expected log level to be %d, got %d", 4, existingLog.Level)
} }
} }
+10
View File
@@ -11,6 +11,8 @@ import (
) )
func TestParamQuery(t *testing.T) { func TestParamQuery(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -23,6 +25,8 @@ func TestParamQuery(t *testing.T) {
} }
func TestFindParamByKey(t *testing.T) { func TestFindParamByKey(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -50,6 +54,8 @@ func TestFindParamByKey(t *testing.T) {
} }
func TestSaveParam(t *testing.T) { func TestSaveParam(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -92,6 +98,8 @@ func TestSaveParam(t *testing.T) {
} }
func TestSaveParamEncrypted(t *testing.T) { func TestSaveParamEncrypted(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -126,6 +134,8 @@ func TestSaveParamEncrypted(t *testing.T) {
} }
func TestDeleteParam(t *testing.T) { func TestDeleteParam(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+292 -85
View File
@@ -1,93 +1,133 @@
package daos package daos
import ( import (
"context"
"database/sql"
"errors" "errors"
"fmt" "fmt"
"sort"
"strings" "strings"
"github.com/pocketbase/dbx" "github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/models/schema" "github.com/pocketbase/pocketbase/models/schema"
"github.com/pocketbase/pocketbase/resolvers"
"github.com/pocketbase/pocketbase/tools/inflector" "github.com/pocketbase/pocketbase/tools/inflector"
"github.com/pocketbase/pocketbase/tools/list" "github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/search"
"github.com/pocketbase/pocketbase/tools/security" "github.com/pocketbase/pocketbase/tools/security"
"github.com/pocketbase/pocketbase/tools/types" "github.com/pocketbase/pocketbase/tools/types"
"github.com/spf13/cast" "github.com/spf13/cast"
) )
// RecordQuery returns a new Record select query. // RecordQuery returns a new Record select query from a collection model, id or name.
func (dao *Dao) RecordQuery(collection *models.Collection) *dbx.SelectQuery { //
tableName := collection.Name // In case a collection id or name is provided and that collection doesn't
// actually exists, the generated query will be created with a cancelled context
// and will fail once an executor (Row(), One(), All(), etc.) is called.
func (dao *Dao) RecordQuery(collectionModelOrIdentifier any) *dbx.SelectQuery {
var tableName string
var collection *models.Collection
var collectionErr error
switch c := collectionModelOrIdentifier.(type) {
case *models.Collection:
collection = c
tableName = collection.Name
case models.Collection:
collection = &c
tableName = collection.Name
case string:
collection, collectionErr = dao.FindCollectionByNameOrId(c)
if collection != nil {
tableName = collection.Name
}
default:
collectionErr = errors.New("unsupported collection identifier, must be collection model, id or name")
}
// update with some fake table name for easier debugging
if tableName == "" {
tableName = "@@__invalidCollectionModelOrIdentifier"
}
selectCols := fmt.Sprintf("%s.*", dao.DB().QuoteSimpleColumnName(tableName)) selectCols := fmt.Sprintf("%s.*", dao.DB().QuoteSimpleColumnName(tableName))
return dao.DB(). query := dao.DB().Select(selectCols).From(tableName)
Select(selectCols).
From(tableName).
WithBuildHook(func(query *dbx.Query) {
query.WithExecHook(execLockRetry(dao.ModelQueryTimeout, dao.MaxLockRetries)).
WithOneHook(func(q *dbx.Query, a any, op func(b any) error) error {
switch v := a.(type) {
case *models.Record:
if v == nil {
return op(a)
}
row := dbx.NullStringMap{} // in case of an error attach a new context and cancel it immediately with the error
if err := op(&row); err != nil { if collectionErr != nil {
return err // @todo consider changing to WithCancelCause when upgrading
} // the min Go requirement to 1.20, so that we can pass the error
ctx, cancelFunc := context.WithCancel(context.Background())
query.WithContext(ctx)
cancelFunc()
}
record := models.NewRecordFromNullStringMap(collection, row) return query.WithBuildHook(func(q *dbx.Query) {
q.WithExecHook(execLockRetry(dao.ModelQueryTimeout, dao.MaxLockRetries)).
*v = *record WithOneHook(func(q *dbx.Query, a any, op func(b any) error) error {
switch v := a.(type) {
return nil case *models.Record:
default: if v == nil {
return op(a) return op(a)
} }
}).
WithAllHook(func(q *dbx.Query, sliceA any, op func(sliceB any) error) error {
switch v := sliceA.(type) {
case *[]*models.Record:
if v == nil {
return op(sliceA)
}
rows := []dbx.NullStringMap{} row := dbx.NullStringMap{}
if err := op(&rows); err != nil { if err := op(&row); err != nil {
return err return err
} }
records := models.NewRecordsFromNullStringMaps(collection, rows) record := models.NewRecordFromNullStringMap(collection, row)
*v = records *v = *record
return nil return nil
case *[]models.Record: default:
if v == nil { return op(a)
return op(sliceA) }
} }).
WithAllHook(func(q *dbx.Query, sliceA any, op func(sliceB any) error) error {
rows := []dbx.NullStringMap{} switch v := sliceA.(type) {
if err := op(&rows); err != nil { case *[]*models.Record:
return err if v == nil {
}
records := models.NewRecordsFromNullStringMaps(collection, rows)
nonPointers := make([]models.Record, len(records))
for i, r := range records {
nonPointers[i] = *r
}
*v = nonPointers
return nil
default:
return op(sliceA) return op(sliceA)
} }
})
}) rows := []dbx.NullStringMap{}
if err := op(&rows); err != nil {
return err
}
records := models.NewRecordsFromNullStringMaps(collection, rows)
*v = records
return nil
case *[]models.Record:
if v == nil {
return op(sliceA)
}
rows := []dbx.NullStringMap{}
if err := op(&rows); err != nil {
return err
}
records := models.NewRecordsFromNullStringMaps(collection, rows)
nonPointers := make([]models.Record, len(records))
for i, r := range records {
nonPointers[i] = *r
}
*v = nonPointers
return nil
default:
return op(sliceA)
}
})
})
} }
// FindRecordById finds the Record model by its id. // FindRecordById finds the Record model by its id.
@@ -170,12 +210,7 @@ func (dao *Dao) FindRecordsByIds(
// expr2 := dbx.NewExp("LOWER(username) = {:username}", dbx.Params{"username": "test"}) // expr2 := dbx.NewExp("LOWER(username) = {:username}", dbx.Params{"username": "test"})
// dao.FindRecordsByExpr("example", expr1, expr2) // dao.FindRecordsByExpr("example", expr1, expr2)
func (dao *Dao) FindRecordsByExpr(collectionNameOrId string, exprs ...dbx.Expression) ([]*models.Record, error) { func (dao *Dao) FindRecordsByExpr(collectionNameOrId string, exprs ...dbx.Expression) ([]*models.Record, error) {
collection, err := dao.FindCollectionByNameOrId(collectionNameOrId) query := dao.RecordQuery(collectionNameOrId)
if err != nil {
return nil, err
}
query := dao.RecordQuery(collection)
// add only the non-nil expressions // add only the non-nil expressions
for _, expr := range exprs { for _, expr := range exprs {
@@ -200,14 +235,9 @@ func (dao *Dao) FindFirstRecordByData(
key string, key string,
value any, value any,
) (*models.Record, error) { ) (*models.Record, error) {
collection, err := dao.FindCollectionByNameOrId(collectionNameOrId)
if err != nil {
return nil, err
}
record := &models.Record{} record := &models.Record{}
err = dao.RecordQuery(collection). err := dao.RecordQuery(collectionNameOrId).
AndWhere(dbx.HashExp{inflector.Columnify(key): value}). AndWhere(dbx.HashExp{inflector.Columnify(key): value}).
Limit(1). Limit(1).
One(record) One(record)
@@ -218,6 +248,113 @@ func (dao *Dao) FindFirstRecordByData(
return record, nil return record, nil
} }
// FindRecordsByFilter returns limit number of records matching the
// provided string filter.
//
// NB! Use the last "params" argument to bind untrusted user variables!
//
// The sort argument is optional and can be empty string OR the same format
// used in the web APIs, eg. "-created,title".
//
// If the limit argument is <= 0, no limit is applied to the query and
// all matching records are returned.
//
// Example:
//
// dao.FindRecordsByFilter(
// "posts",
// "title ~ {:title} && visible = {:visible}",
// "-created",
// 10,
// 0,
// dbx.Params{"title": "lorem ipsum", "visible": true}
// )
func (dao *Dao) FindRecordsByFilter(
collectionNameOrId string,
filter string,
sort string,
limit int,
offset int,
params ...dbx.Params,
) ([]*models.Record, error) {
collection, err := dao.FindCollectionByNameOrId(collectionNameOrId)
if err != nil {
return nil, err
}
q := dao.RecordQuery(collection)
// build a fields resolver and attach the generated conditions to the query
// ---
resolver := resolvers.NewRecordFieldResolver(
dao,
collection, // the base collection
nil, // no request data
true, // allow searching hidden/protected fields like "email"
)
expr, err := search.FilterData(filter).BuildExpr(resolver, params...)
if err != nil || expr == nil {
return nil, errors.New("invalid or empty filter expression")
}
q.AndWhere(expr)
if sort != "" {
for _, sortField := range search.ParseSortFromString(sort) {
expr, err := sortField.BuildExpr(resolver)
if err != nil {
return nil, err
}
if expr != "" {
q.AndOrderBy(expr)
}
}
}
resolver.UpdateQuery(q) // attaches any adhoc joins and aliases
// ---
if offset > 0 {
q.Offset(int64(offset))
}
if limit > 0 {
q.Limit(int64(limit))
}
records := []*models.Record{}
if err := q.All(&records); err != nil {
return nil, err
}
return records, nil
}
// FindFirstRecordByFilter returns the first available record matching the provided filter.
//
// NB! Use the last params argument to bind untrusted user variables!
//
// Example:
//
// dao.FindFirstRecordByFilter("posts", "slug={:slug} && status='public'", dbx.Params{"slug": "test"})
func (dao *Dao) FindFirstRecordByFilter(
collectionNameOrId string,
filter string,
params ...dbx.Params,
) (*models.Record, error) {
result, err := dao.FindRecordsByFilter(collectionNameOrId, filter, "", 1, 0, params...)
if err != nil {
return nil, err
}
if len(result) == 0 {
return nil, sql.ErrNoRows
}
return result[0], nil
}
// IsRecordValueUnique checks if the provided key-value pair is a unique Record value. // IsRecordValueUnique checks if the provided key-value pair is a unique Record value.
// //
// For correctness, if the collection is "auth" and the key is "username", // For correctness, if the collection is "auth" and the key is "username",
@@ -271,9 +408,9 @@ func (dao *Dao) IsRecordValueUnique(
return query.Row(&exists) == nil && !exists return query.Row(&exists) == nil && !exists
} }
// FindAuthRecordByToken finds the auth record associated with the provided JWT token. // FindAuthRecordByToken finds the auth record associated with the provided JWT.
// //
// Returns an error if the JWT token is invalid, expired or not associated to an auth collection record. // Returns an error if the JWT is invalid, expired or not associated to an auth collection record.
func (dao *Dao) FindAuthRecordByToken(token string, baseTokenKey string) (*models.Record, error) { func (dao *Dao) FindAuthRecordByToken(token string, baseTokenKey string) (*models.Record, error) {
unverifiedClaims, err := security.ParseUnverifiedJWT(token) unverifiedClaims, err := security.ParseUnverifiedJWT(token)
if err != nil { if err != nil {
@@ -293,7 +430,7 @@ func (dao *Dao) FindAuthRecordByToken(token string, baseTokenKey string) (*model
} }
if !record.Collection().IsAuth() { if !record.Collection().IsAuth() {
return nil, errors.New("The token is not associated to an auth collection record.") return nil, errors.New("the token is not associated to an auth collection record")
} }
verificationKey := record.TokenKey() + baseTokenKey verificationKey := record.TokenKey() + baseTokenKey
@@ -386,6 +523,62 @@ func (dao *Dao) SuggestUniqueAuthRecordUsername(
return username return username
} }
// CanAccessRecord checks if a record is allowed to be accessed by the
// specified requestInfo and accessRule.
//
// Rule and db checks are ignored in case requestInfo.Admin is set.
//
// The returned error indicate that something unexpected happened during
// the check (eg. invalid rule or db error).
//
// The method always return false on invalid access rule or db error.
//
// Example:
//
// requestInfo := apis.RequestInfo(c /* echo.Context */)
// record, _ := dao.FindRecordById("example", "RECORD_ID")
// rule := types.Pointer("@request.auth.id != '' || status = 'public'")
// // ... or use one of the record collection's rule, eg. record.Collection().ViewRule
//
// if ok, _ := dao.CanAccessRecord(record, requestInfo, rule); ok { ... }
func (dao *Dao) CanAccessRecord(record *models.Record, requestInfo *models.RequestInfo, accessRule *string) (bool, error) {
if requestInfo.Admin != nil {
// admins can access everything
return true, nil
}
if accessRule == nil {
// only admins can access this record
return false, nil
}
if *accessRule == "" {
// empty public rule, aka. everyone can access
return true, nil
}
var exists bool
query := dao.RecordQuery(record.Collection()).
Select("(1)").
AndWhere(dbx.HashExp{record.Collection().Name + ".id": record.Id})
// parse and apply the access rule filter
resolver := resolvers.NewRecordFieldResolver(dao, record.Collection(), requestInfo, true)
expr, err := search.FilterData(*accessRule).BuildExpr(resolver)
if err != nil {
return false, err
}
resolver.UpdateQuery(query)
query.AndWhere(expr)
if err := query.Limit(1).Row(&exists); err != nil && !errors.Is(err, sql.ErrNoRows) {
return false, err
}
return exists, nil
}
// SaveRecord persists the provided Record model in the database. // SaveRecord persists the provided Record model in the database.
// //
// If record.IsNew() is true, the method will perform a create, otherwise an update. // If record.IsNew() is true, the method will perform a create, otherwise an update.
@@ -464,26 +657,40 @@ func (dao *Dao) DeleteRecord(record *models.Record) error {
// //
// NB! This method is expected to be called inside a transaction. // NB! This method is expected to be called inside a transaction.
func (dao *Dao) cascadeRecordDelete(mainRecord *models.Record, refs map[*models.Collection][]*schema.SchemaField) error { func (dao *Dao) cascadeRecordDelete(mainRecord *models.Record, refs map[*models.Collection][]*schema.SchemaField) error {
uniqueJsonEachAlias := "__je__" + security.PseudorandomString(4) // @todo consider changing refs to a slice
//
// Sort the refs keys to ensure that the cascade events firing order is always the same.
// This is not necessary for the operation to function correctly but it helps having deterministic output during testing.
sortedRefKeys := make([]*models.Collection, 0, len(refs))
for k := range refs {
sortedRefKeys = append(sortedRefKeys, k)
}
sort.Slice(sortedRefKeys, func(i, j int) bool {
return sortedRefKeys[i].Name < sortedRefKeys[j].Name
})
for refCollection, fields := range refs { for _, refCollection := range sortedRefKeys {
if refCollection.IsView() { fields, ok := refs[refCollection]
continue // skip view collections
if refCollection.IsView() || !ok {
continue // skip missing or view collections
} }
for _, field := range fields { for _, field := range fields {
recordTableName := inflector.Columnify(refCollection.Name) recordTableName := inflector.Columnify(refCollection.Name)
prefixedFieldName := recordTableName + "." + inflector.Columnify(field.Name) prefixedFieldName := recordTableName + "." + inflector.Columnify(field.Name)
query := dao.RecordQuery(refCollection).Distinct(true) query := dao.RecordQuery(refCollection)
if opt, ok := field.Options.(schema.MultiValuer); !ok || !opt.IsMultiple() { if opt, ok := field.Options.(schema.MultiValuer); !ok || !opt.IsMultiple() {
query.AndWhere(dbx.HashExp{prefixedFieldName: mainRecord.Id}) query.AndWhere(dbx.HashExp{prefixedFieldName: mainRecord.Id})
} else { } else {
query.InnerJoin(fmt.Sprintf( query.AndWhere(dbx.Exists(dbx.NewExp(fmt.Sprintf(
`json_each(CASE WHEN json_valid([[%s]]) THEN [[%s]] ELSE json_array([[%s]]) END) as {{%s}}`, `SELECT 1 FROM json_each(CASE WHEN json_valid([[%s]]) THEN [[%s]] ELSE json_array([[%s]]) END) {{__je__}} WHERE [[__je__.value]]={:jevalue}`,
prefixedFieldName, prefixedFieldName, prefixedFieldName, uniqueJsonEachAlias, prefixedFieldName, prefixedFieldName, prefixedFieldName,
), dbx.HashExp{uniqueJsonEachAlias + ".value": mainRecord.Id}) ), dbx.Params{
"jevalue": mainRecord.Id,
})))
} }
if refCollection.Id == mainRecord.Collection().Id { if refCollection.Id == mainRecord.Collection().Id {
+75 -45
View File
@@ -3,6 +3,7 @@ package daos
import ( import (
"errors" "errors"
"fmt" "fmt"
"log"
"regexp" "regexp"
"strings" "strings"
@@ -10,13 +11,14 @@ import (
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/models/schema" "github.com/pocketbase/pocketbase/models/schema"
"github.com/pocketbase/pocketbase/tools/dbutils" "github.com/pocketbase/pocketbase/tools/dbutils"
"github.com/pocketbase/pocketbase/tools/inflector"
"github.com/pocketbase/pocketbase/tools/list" "github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/security" "github.com/pocketbase/pocketbase/tools/security"
"github.com/pocketbase/pocketbase/tools/types" "github.com/pocketbase/pocketbase/tools/types"
) )
// MaxExpandDepth specifies the max allowed nested expand depth path. // MaxExpandDepth specifies the max allowed nested expand depth path.
//
// @todo Consider eventually reusing resolvers.maxNestedRels
const MaxExpandDepth = 6 const MaxExpandDepth = 6
// ExpandFetchFunc defines the function that is used to fetch the expanded relation records. // ExpandFetchFunc defines the function that is used to fetch the expanded relation records.
@@ -24,21 +26,27 @@ type ExpandFetchFunc func(relCollection *models.Collection, relIds []string) ([]
// ExpandRecord expands the relations of a single Record model. // ExpandRecord expands the relations of a single Record model.
// //
// If optFetchFunc is not set, then a default function will be used
// that returns all relation records.
//
// Returns a map with the failed expand parameters and their errors. // Returns a map with the failed expand parameters and their errors.
func (dao *Dao) ExpandRecord(record *models.Record, expands []string, fetchFunc ExpandFetchFunc) map[string]error { func (dao *Dao) ExpandRecord(record *models.Record, expands []string, optFetchFunc ExpandFetchFunc) map[string]error {
return dao.ExpandRecords([]*models.Record{record}, expands, fetchFunc) return dao.ExpandRecords([]*models.Record{record}, expands, optFetchFunc)
} }
// ExpandRecords expands the relations of the provided Record models list. // ExpandRecords expands the relations of the provided Record models list.
// //
// If optFetchFunc is not set, then a default function will be used
// that returns all relation records.
//
// Returns a map with the failed expand parameters and their errors. // Returns a map with the failed expand parameters and their errors.
func (dao *Dao) ExpandRecords(records []*models.Record, expands []string, fetchFunc ExpandFetchFunc) map[string]error { func (dao *Dao) ExpandRecords(records []*models.Record, expands []string, optFetchFunc ExpandFetchFunc) map[string]error {
normalized := normalizeExpands(expands) normalized := normalizeExpands(expands)
failed := map[string]error{} failed := map[string]error{}
for _, expand := range normalized { for _, expand := range normalized {
if err := dao.expandRecords(records, expand, fetchFunc, 1); err != nil { if err := dao.expandRecords(records, expand, optFetchFunc, 1); err != nil {
failed[expand] = err failed[expand] = err
} }
} }
@@ -46,16 +54,21 @@ func (dao *Dao) ExpandRecords(records []*models.Record, expands []string, fetchF
return failed return failed
} }
var indirectExpandRegex = regexp.MustCompile(`^(\w+)\((\w+)\)$`) // Deprecated
var indirectExpandRegexOld = regexp.MustCompile(`^(\w+)\((\w+)\)$`)
var indirectExpandRegex = regexp.MustCompile(`^(\w+)_via_(\w+)$`)
// notes: // notes:
// - fetchFunc must be non-nil func // - if fetchFunc is nil, dao.FindRecordsByIds will be used
// - all records are expected to be from the same collection // - all records are expected to be from the same collection
// - if MaxExpandDepth is reached, the function returns nil ignoring the remaining expand path // - if MaxExpandDepth is reached, the function returns nil ignoring the remaining expand path
// - indirect expands are supported only with single relation fields
func (dao *Dao) expandRecords(records []*models.Record, expandPath string, fetchFunc ExpandFetchFunc, recursionLevel int) error { func (dao *Dao) expandRecords(records []*models.Record, expandPath string, fetchFunc ExpandFetchFunc, recursionLevel int) error {
if fetchFunc == nil { if fetchFunc == nil {
return errors.New("Relation records fetchFunc is not set.") // load a default fetchFunc
fetchFunc = func(relCollection *models.Collection, relIds []string) ([]*models.Record, error) {
return dao.FindRecordsByIds(relCollection.Id, relIds)
}
} }
if expandPath == "" || recursionLevel > MaxExpandDepth || len(records) == 0 { if expandPath == "" || recursionLevel > MaxExpandDepth || len(records) == 0 {
@@ -69,70 +82,87 @@ func (dao *Dao) expandRecords(records []*models.Record, expandPath string, fetch
var relCollection *models.Collection var relCollection *models.Collection
parts := strings.SplitN(expandPath, ".", 2) parts := strings.SplitN(expandPath, ".", 2)
matches := indirectExpandRegex.FindStringSubmatch(parts[0]) var matches []string
// @todo remove the old syntax support
if strings.Contains(parts[0], "(") {
matches = indirectExpandRegexOld.FindStringSubmatch(parts[0])
if len(matches) == 3 {
log.Printf(
"%s expand format is deprecated and will be removed in the future. Consider replacing it with %s_via_%s.\n",
matches[0],
matches[1],
matches[2],
)
}
} else {
matches = indirectExpandRegex.FindStringSubmatch(parts[0])
}
if len(matches) == 3 { if len(matches) == 3 {
indirectRel, _ := dao.FindCollectionByNameOrId(matches[1]) indirectRel, _ := dao.FindCollectionByNameOrId(matches[1])
if indirectRel == nil { if indirectRel == nil {
return fmt.Errorf("Couldn't find indirect related collection %q.", matches[1]) return fmt.Errorf("couldn't find back-related collection %q", matches[1])
} }
indirectRelField := indirectRel.Schema.GetFieldByName(matches[2]) indirectRelField := indirectRel.Schema.GetFieldByName(matches[2])
if indirectRelField == nil || indirectRelField.Type != schema.FieldTypeRelation { if indirectRelField == nil || indirectRelField.Type != schema.FieldTypeRelation {
return fmt.Errorf("Couldn't find indirect relation field %q in collection %q.", matches[2], mainCollection.Name) return fmt.Errorf("couldn't find back-relation field %q in collection %q", matches[2], indirectRel.Name)
} }
indirectRelField.InitOptions() indirectRelField.InitOptions()
indirectRelFieldOptions, _ := indirectRelField.Options.(*schema.RelationOptions) indirectRelFieldOptions, _ := indirectRelField.Options.(*schema.RelationOptions)
if indirectRelFieldOptions == nil || indirectRelFieldOptions.CollectionId != mainCollection.Id { if indirectRelFieldOptions == nil || indirectRelFieldOptions.CollectionId != mainCollection.Id {
return fmt.Errorf("Invalid indirect relation field path %q.", parts[0]) return fmt.Errorf("invalid back-relation field path %q", parts[0])
}
if indirectRelFieldOptions.IsMultiple() {
// for now don't allow multi-relation indirect fields expand
// due to eventual poor query performance with large data sets.
return fmt.Errorf("Multi-relation fields cannot be indirectly expanded in %q.", parts[0])
} }
recordIds := make([]any, len(records)) // add the related id(s) as a dynamic relation field value to
for i, record := range records { // allow further expand checks at later stage in a more unified manner
recordIds[i] = record.Id prepErr := func() error {
} q := dao.DB().Select("id").
From(indirectRel.Name).
Limit(1000) // the limit is arbitrary chosen and may change in the future
// @todo after the index optimizations consider allowing if indirectRelFieldOptions.IsMultiple() {
// indirect expand for multi-relation fields q.AndWhere(dbx.Exists(dbx.NewExp(fmt.Sprintf(
indirectRecords, err := dao.FindRecordsByExpr( "SELECT 1 FROM %s je WHERE je.value = {:id}",
indirectRel.Id, dbutils.JsonEach(indirectRelField.Name),
dbx.In(inflector.Columnify(matches[2]), recordIds...), ))))
) } else {
if err != nil { q.AndWhere(dbx.NewExp("[[" + indirectRelField.Name + "]] = {:id}"))
return err
}
mappedIndirectRecordIds := make(map[string][]string, len(indirectRecords))
for _, indirectRecord := range indirectRecords {
recId := indirectRecord.GetString(matches[2])
if recId != "" {
mappedIndirectRecordIds[recId] = append(mappedIndirectRecordIds[recId], indirectRecord.Id)
} }
}
// add the indirect relation ids as a new relation field value pq := q.Build().Prepare()
for _, record := range records {
relIds, ok := mappedIndirectRecordIds[record.Id] for _, record := range records {
if ok && len(relIds) > 0 { var relIds []string
record.Set(parts[0], relIds)
err := pq.Bind(dbx.Params{"id": record.Id}).Column(&relIds)
if err != nil {
return errors.Join(err, pq.Close())
}
if len(relIds) > 0 {
record.Set(parts[0], relIds)
}
} }
return pq.Close()
}()
if prepErr != nil {
return prepErr
} }
relFieldOptions = &schema.RelationOptions{ relFieldOptions = &schema.RelationOptions{
MaxSelect: nil, MaxSelect: nil,
CollectionId: indirectRel.Id, CollectionId: indirectRel.Id,
} }
if isRelFieldUnique(indirectRel, indirectRelField.Name) { if dbutils.HasSingleColumnUniqueIndex(indirectRelField.Name, indirectRel.Indexes) {
relFieldOptions.MaxSelect = types.Pointer(1) relFieldOptions.MaxSelect = types.Pointer(1)
} }
// indirect relation // indirect/back relation
relField = &schema.SchemaField{ relField = &schema.SchemaField{
Id: "indirect_" + security.PseudorandomString(5), Id: "_" + parts[0] + security.PseudorandomString(3),
Type: schema.FieldTypeRelation, Type: schema.FieldTypeRelation,
Name: parts[0], Name: parts[0],
Options: relFieldOptions, Options: relFieldOptions,
+93 -31
View File
@@ -14,6 +14,8 @@ import (
) )
func TestExpandRecords(t *testing.T) { func TestExpandRecords(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -48,15 +50,6 @@ func TestExpandRecords(t *testing.T) {
0, 0,
0, 0,
}, },
{
"empty fetchFunc",
"demo4",
[]string{"i9naidtvr6qsgb4", "qzaqccwrmva4o1n"},
[]string{"self_rel_one", "self_rel_many.self_rel_one"},
nil,
0,
2,
},
{ {
"fetchFunc with error", "fetchFunc with error",
"demo4", "demo4",
@@ -101,6 +94,19 @@ func TestExpandRecords(t *testing.T) {
0, 0,
1, 1,
}, },
{
"with nil fetchfunc",
"users",
[]string{
"bgs820n361vj1qd",
"4q1xlclmfloku33",
"oap640cot4yru2s", // no rels
},
[]string{"rel"},
nil,
2,
0,
},
{ {
"expand normalizations", "expand normalizations",
"demo4", "demo4",
@@ -132,6 +138,19 @@ func TestExpandRecords(t *testing.T) {
2, 2,
0, 0,
}, },
{
"with nil fetchfunc",
"users",
[]string{
"bgs820n361vj1qd",
"4q1xlclmfloku33",
"oap640cot4yru2s", // no rels
},
[]string{"rel"},
nil,
2,
0,
},
{ {
"maxExpandDepth reached", "maxExpandDepth reached",
"demo4", "demo4",
@@ -144,7 +163,7 @@ func TestExpandRecords(t *testing.T) {
0, 0,
}, },
{ {
"simple indirect expand", "simple back single relation field expand (deprecated syntax)",
"demo3", "demo3",
[]string{"lcl9d87w22ml6jy"}, []string{"lcl9d87w22ml6jy"},
[]string{"demo4(rel_one_no_cascade_required)"}, []string{"demo4(rel_one_no_cascade_required)"},
@@ -155,11 +174,22 @@ func TestExpandRecords(t *testing.T) {
0, 0,
}, },
{ {
"nested indirect expand", "simple back expand via single relation field",
"demo3",
[]string{"lcl9d87w22ml6jy"},
[]string{"demo4_via_rel_one_no_cascade_required"},
func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil)
},
1,
0,
},
{
"nested back expand via single relation field",
"demo3", "demo3",
[]string{"lcl9d87w22ml6jy"}, []string{"lcl9d87w22ml6jy"},
[]string{ []string{
"demo4(rel_one_no_cascade_required).self_rel_many.self_rel_many.self_rel_one", "demo4_via_rel_one_no_cascade_required.self_rel_many.self_rel_many.self_rel_one",
}, },
func(c *models.Collection, ids []string) ([]*models.Record, error) { func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil) return app.Dao().FindRecordsByIds(c.Id, ids, nil)
@@ -167,6 +197,19 @@ func TestExpandRecords(t *testing.T) {
5, 5,
0, 0,
}, },
{
"nested back expand via multiple relation field",
"demo3",
[]string{"lcl9d87w22ml6jy"},
[]string{
"demo4_via_rel_many_no_cascade_required.self_rel_many.rel_many_no_cascade_required.demo4_via_rel_many_no_cascade_required",
},
func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil)
},
7,
0,
},
{ {
"expand multiple relations sharing a common path", "expand multiple relations sharing a common path",
"demo4", "demo4",
@@ -205,6 +248,8 @@ func TestExpandRecords(t *testing.T) {
} }
func TestExpandRecord(t *testing.T) { func TestExpandRecord(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -228,15 +273,6 @@ func TestExpandRecord(t *testing.T) {
0, 0,
0, 0,
}, },
{
"empty fetchFunc",
"demo4",
"i9naidtvr6qsgb4",
[]string{"self_rel_one", "self_rel_many.self_rel_one"},
nil,
0,
2,
},
{ {
"fetchFunc with error", "fetchFunc with error",
"demo4", "demo4",
@@ -320,7 +356,7 @@ func TestExpandRecord(t *testing.T) {
0, 0,
}, },
{ {
"simple indirect expand", "simple indirect expand via single relation field (deprecated syntax)",
"demo3", "demo3",
"lcl9d87w22ml6jy", "lcl9d87w22ml6jy",
[]string{"demo4(rel_one_no_cascade_required)"}, []string{"demo4(rel_one_no_cascade_required)"},
@@ -331,7 +367,18 @@ func TestExpandRecord(t *testing.T) {
0, 0,
}, },
{ {
"nested indirect expand", "simple indirect expand via single relation field",
"demo3",
"lcl9d87w22ml6jy",
[]string{"demo4_via_rel_one_no_cascade_required"},
func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil)
},
1,
0,
},
{
"nested indirect expand via single relation field",
"demo3", "demo3",
"lcl9d87w22ml6jy", "lcl9d87w22ml6jy",
[]string{ []string{
@@ -343,6 +390,19 @@ func TestExpandRecord(t *testing.T) {
5, 5,
0, 0,
}, },
{
"nested indirect expand via single relation field",
"demo3",
"lcl9d87w22ml6jy",
[]string{
"demo4_via_rel_many_no_cascade_required.self_rel_many.rel_many_no_cascade_required.demo4_via_rel_many_no_cascade_required",
},
func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil)
},
7,
0,
},
} }
for _, s := range scenarios { for _, s := range scenarios {
@@ -364,6 +424,8 @@ func TestExpandRecord(t *testing.T) {
} }
func TestIndirectExpandSingeVsArrayResult(t *testing.T) { func TestIndirectExpandSingeVsArrayResult(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -374,21 +436,23 @@ func TestIndirectExpandSingeVsArrayResult(t *testing.T) {
// non-unique indirect expand // non-unique indirect expand
{ {
errs := app.Dao().ExpandRecord(record, []string{"demo4(rel_one_cascade)"}, func(c *models.Collection, ids []string) ([]*models.Record, error) { errs := app.Dao().ExpandRecord(record, []string{"demo4_via_rel_one_cascade"}, func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil) return app.Dao().FindRecordsByIds(c.Id, ids, nil)
}) })
if len(errs) > 0 { if len(errs) > 0 {
t.Fatal(errs) t.Fatal(errs)
} }
result, ok := record.Expand()["demo4(rel_one_cascade)"].([]*models.Record) result, ok := record.Expand()["demo4_via_rel_one_cascade"].([]*models.Record)
if !ok { if !ok {
t.Fatalf("Expected the expanded result to be a slice, got %v", result) t.Fatalf("Expected the expanded result to be a slice, got %v", result)
} }
} }
// mock a unique constraint for the rel_one_cascade field // unique indirect expand
{ {
// mock a unique constraint for the rel_one_cascade field
// ---
demo4, err := app.Dao().FindCollectionByNameOrId("demo4") demo4, err := app.Dao().FindCollectionByNameOrId("demo4")
if err != nil { if err != nil {
t.Fatal(err) t.Fatal(err)
@@ -399,18 +463,16 @@ func TestIndirectExpandSingeVsArrayResult(t *testing.T) {
if err := app.Dao().SaveCollection(demo4); err != nil { if err := app.Dao().SaveCollection(demo4); err != nil {
t.Fatalf("Failed to mock unique constraint: %v", err) t.Fatalf("Failed to mock unique constraint: %v", err)
} }
} // ---
// non-unique indirect expand errs := app.Dao().ExpandRecord(record, []string{"demo4_via_rel_one_cascade"}, func(c *models.Collection, ids []string) ([]*models.Record, error) {
{
errs := app.Dao().ExpandRecord(record, []string{"demo4(rel_one_cascade)"}, func(c *models.Collection, ids []string) ([]*models.Record, error) {
return app.Dao().FindRecordsByIds(c.Id, ids, nil) return app.Dao().FindRecordsByIds(c.Id, ids, nil)
}) })
if len(errs) > 0 { if len(errs) > 0 {
t.Fatal(errs) t.Fatal(errs)
} }
result, ok := record.Expand()["demo4(rel_one_cascade)"].(*models.Record) result, ok := record.Expand()["demo4_via_rel_one_cascade"].(*models.Record)
if !ok { if !ok {
t.Fatalf("Expected the expanded result to be a single model, got %v", result) t.Fatalf("Expected the expanded result to be a single model, got %v", result)
} }
+64 -76
View File
@@ -10,7 +10,6 @@ import (
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/models/schema" "github.com/pocketbase/pocketbase/models/schema"
"github.com/pocketbase/pocketbase/tools/dbutils" "github.com/pocketbase/pocketbase/tools/dbutils"
"github.com/pocketbase/pocketbase/tools/list"
"github.com/pocketbase/pocketbase/tools/security" "github.com/pocketbase/pocketbase/tools/security"
) )
@@ -159,10 +158,6 @@ func (dao *Dao) SyncRecordTableSchema(newCollection *models.Collection, oldColle
return err return err
} }
if err := txDao.syncRelationDisplayFieldsChanges(newCollection, renamedFieldNames, deletedFieldNames); err != nil {
return err
}
return txDao.createCollectionIndexes(newCollection) return txDao.createCollectionIndexes(newCollection)
}) })
} }
@@ -173,10 +168,23 @@ func (dao *Dao) normalizeSingleVsMultipleFieldChanges(newCollection, oldCollecti
} }
return dao.RunInTransaction(func(txDao *Dao) error { return dao.RunInTransaction(func(txDao *Dao) error {
// temporary disable the schema error checks to prevent view and trigger errors
// when "altering" (aka. deleting and recreating) the non-normalized columns
if _, err := txDao.DB().NewQuery("PRAGMA writable_schema = ON").Execute(); err != nil {
return err
}
// executed with defer to make sure that the pragma is always reverted
// in case of an error and when nested transactions are used
defer txDao.DB().NewQuery("PRAGMA writable_schema = RESET").Execute()
for _, newField := range newCollection.Schema.Fields() { for _, newField := range newCollection.Schema.Fields() {
oldField := oldCollection.Schema.GetFieldById(newField.Id) // allow to continue even if there is no old field for the cases
if oldField == nil { // when a new field is added and there are already inserted data
continue var isOldMultiple bool
if oldField := oldCollection.Schema.GetFieldById(newField.Id); oldField != nil {
if opt, ok := oldField.Options.(schema.MultiValuer); ok {
isOldMultiple = opt.IsMultiple()
}
} }
var isNewMultiple bool var isNewMultiple bool
@@ -184,20 +192,30 @@ func (dao *Dao) normalizeSingleVsMultipleFieldChanges(newCollection, oldCollecti
isNewMultiple = opt.IsMultiple() isNewMultiple = opt.IsMultiple()
} }
var isOldMultiple bool
if opt, ok := oldField.Options.(schema.MultiValuer); ok {
isOldMultiple = opt.IsMultiple()
}
if isOldMultiple == isNewMultiple { if isOldMultiple == isNewMultiple {
continue // no change continue // no change
} }
var updateQuery *dbx.Query // update the column definition by:
// 1. inserting a new column with the new definition
// 2. copy normalized values from the original column to the new one
// 3. drop the original column
// 4. rename the new column to the original column
// -------------------------------------------------------
originalName := newField.Name
tempName := "_" + newField.Name + security.PseudorandomString(5)
_, err := txDao.DB().AddColumn(newCollection.Name, tempName, newField.ColDefinition()).Execute()
if err != nil {
return err
}
var copyQuery *dbx.Query
if !isOldMultiple && isNewMultiple { if !isOldMultiple && isNewMultiple {
// single -> multiple (convert to array) // single -> multiple (convert to array)
updateQuery = txDao.DB().NewQuery(fmt.Sprintf( copyQuery = txDao.DB().NewQuery(fmt.Sprintf(
`UPDATE {{%s}} set [[%s]] = ( `UPDATE {{%s}} set [[%s]] = (
CASE CASE
WHEN COALESCE([[%s]], '') = '' WHEN COALESCE([[%s]], '') = ''
@@ -212,19 +230,19 @@ func (dao *Dao) normalizeSingleVsMultipleFieldChanges(newCollection, oldCollecti
END END
)`, )`,
newCollection.Name, newCollection.Name,
newField.Name, tempName,
newField.Name, originalName,
newField.Name, originalName,
newField.Name, originalName,
newField.Name, originalName,
newField.Name, originalName,
)) ))
} else { } else {
// multiple -> single (keep only the last element) // multiple -> single (keep only the last element)
// //
// note: for file fields the actual files are not deleted // note: for file fields the actual file objects are not
// allowing additional custom handling via migration. // deleted allowing additional custom handling via migration
updateQuery = txDao.DB().NewQuery(fmt.Sprintf( copyQuery = txDao.DB().NewQuery(fmt.Sprintf(
`UPDATE {{%s}} set [[%s]] = ( `UPDATE {{%s}} set [[%s]] = (
CASE CASE
WHEN COALESCE([[%s]], '[]') = '[]' WHEN COALESCE([[%s]], '[]') = '[]'
@@ -239,68 +257,38 @@ func (dao *Dao) normalizeSingleVsMultipleFieldChanges(newCollection, oldCollecti
END END
)`, )`,
newCollection.Name, newCollection.Name,
newField.Name, tempName,
newField.Name, originalName,
newField.Name, originalName,
newField.Name, originalName,
newField.Name, originalName,
newField.Name, originalName,
)) ))
} }
if _, err := updateQuery.Execute(); err != nil { // copy the normalized values
if _, err := copyQuery.Execute(); err != nil {
return err
}
// drop the original column
if _, err := txDao.DB().DropColumn(newCollection.Name, originalName).Execute(); err != nil {
return err
}
// rename the new column back to the original
if _, err := txDao.DB().RenameColumn(newCollection.Name, tempName, originalName).Execute(); err != nil {
return err return err
} }
} }
return nil // revert the pragma and reload the schema
_, revertErr := txDao.DB().NewQuery("PRAGMA writable_schema = RESET").Execute()
return revertErr
}) })
} }
func (dao *Dao) syncRelationDisplayFieldsChanges(collection *models.Collection, renamedFieldNames map[string]string, deletedFieldNames []string) error {
if len(renamedFieldNames) == 0 && len(deletedFieldNames) == 0 {
return nil // nothing to sync
}
refs, err := dao.FindCollectionReferences(collection)
if err != nil {
return err
}
for refCollection, refFields := range refs {
for _, refField := range refFields {
options, _ := refField.Options.(*schema.RelationOptions)
if options == nil {
continue
}
// remove deleted (if any)
newDisplayFields := list.SubtractSlice(options.DisplayFields, deletedFieldNames)
for old, new := range renamedFieldNames {
for i, name := range newDisplayFields {
if name == old {
newDisplayFields[i] = new
}
}
}
// has changes
if len(list.SubtractSlice(options.DisplayFields, newDisplayFields)) > 0 {
options.DisplayFields = newDisplayFields
// direct collection save to prevent self-referencing
// recursion and unnecessary records table sync checks
if err := dao.Save(refCollection); err != nil {
return err
}
}
}
}
return nil
}
func (dao *Dao) dropCollectionIndex(collection *models.Collection) error { func (dao *Dao) dropCollectionIndex(collection *models.Collection) error {
if collection.IsView() { if collection.IsView() {
return nil // views don't have indexes return nil // views don't have indexes
+121 -57
View File
@@ -14,6 +14,8 @@ import (
) )
func TestSyncRecordTableSchema(t *testing.T) { func TestSyncRecordTableSchema(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -132,6 +134,8 @@ func TestSyncRecordTableSchema(t *testing.T) {
} }
func TestSingleVsMultipleValuesNormalization(t *testing.T) { func TestSingleVsMultipleValuesNormalization(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -171,89 +175,149 @@ func TestSingleVsMultipleValuesNormalization(t *testing.T) {
opt := relManyField.Options.(*schema.RelationOptions) opt := relManyField.Options.(*schema.RelationOptions)
opt.MaxSelect = types.Pointer(1) opt.MaxSelect = types.Pointer(1)
} }
{
// new multivaluer field to check whether the array normalization
// will be applied for already inserted data
collection.Schema.AddField(&schema.SchemaField{
Name: "new_multiple",
Type: schema.FieldTypeSelect,
Options: &schema.SelectOptions{
Values: []string{"a", "b", "c"},
MaxSelect: 3,
},
})
}
if err := app.Dao().SaveCollection(collection); err != nil { if err := app.Dao().SaveCollection(collection); err != nil {
t.Fatal(err) t.Fatal(err)
} }
type expectation struct { // ensures that the writable schema was reverted to its expected default
SelectOne string `db:"select_one"` var writableSchema bool
SelectMany string `db:"select_many"` app.Dao().DB().NewQuery("PRAGMA writable_schema").Row(&writableSchema)
FileOne string `db:"file_one"` if writableSchema == true {
FileMany string `db:"file_many"` t.Fatalf("Expected writable_schema to be OFF, got %v", writableSchema)
RelOne string `db:"rel_one"`
RelMany string `db:"rel_many"`
} }
scenarios := []struct { // check whether the columns DEFAULT definition was updated
// ---------------------------------------------------------------
tableInfo, err := app.Dao().TableInfo(collection.Name)
if err != nil {
t.Fatal(err)
}
tableInfoExpectations := map[string]string{
"select_one": `'[]'`,
"select_many": `''`,
"file_one": `'[]'`,
"file_many": `''`,
"rel_one": `'[]'`,
"rel_many": `''`,
"new_multiple": `'[]'`,
}
for col, dflt := range tableInfoExpectations {
t.Run("check default for "+col, func(t *testing.T) {
var row *models.TableInfoRow
for _, r := range tableInfo {
if r.Name == col {
row = r
break
}
}
if row == nil {
t.Fatalf("Missing info for column %q", col)
}
if v := row.DefaultValue.String(); v != dflt {
t.Fatalf("Expected default value %q, got %q", dflt, v)
}
})
}
// check whether the values were normalized
// ---------------------------------------------------------------
type fieldsExpectation struct {
SelectOne string `db:"select_one"`
SelectMany string `db:"select_many"`
FileOne string `db:"file_one"`
FileMany string `db:"file_many"`
RelOne string `db:"rel_one"`
RelMany string `db:"rel_many"`
NewMultiple string `db:"new_multiple"`
}
fieldsScenarios := []struct {
recordId string recordId string
expected expectation expected fieldsExpectation
}{ }{
{ {
"imy661ixudk5izi", "imy661ixudk5izi",
expectation{ fieldsExpectation{
SelectOne: `[]`, SelectOne: `[]`,
SelectMany: ``, SelectMany: ``,
FileOne: `[]`, FileOne: `[]`,
FileMany: ``, FileMany: ``,
RelOne: `[]`, RelOne: `[]`,
RelMany: ``, RelMany: ``,
NewMultiple: `[]`,
}, },
}, },
{ {
"al1h9ijdeojtsjy", "al1h9ijdeojtsjy",
expectation{ fieldsExpectation{
SelectOne: `["optionB"]`, SelectOne: `["optionB"]`,
SelectMany: `optionB`, SelectMany: `optionB`,
FileOne: `["300_Jsjq7RdBgA.png"]`, FileOne: `["300_Jsjq7RdBgA.png"]`,
FileMany: ``, FileMany: ``,
RelOne: `["84nmscqy84lsi1t"]`, RelOne: `["84nmscqy84lsi1t"]`,
RelMany: `oap640cot4yru2s`, RelMany: `oap640cot4yru2s`,
NewMultiple: `[]`,
}, },
}, },
{ {
"84nmscqy84lsi1t", "84nmscqy84lsi1t",
expectation{ fieldsExpectation{
SelectOne: `["optionB"]`, SelectOne: `["optionB"]`,
SelectMany: `optionC`, SelectMany: `optionC`,
FileOne: `["test_d61b33QdDU.txt"]`, FileOne: `["test_d61b33QdDU.txt"]`,
FileMany: `test_tC1Yc87DfC.txt`, FileMany: `test_tC1Yc87DfC.txt`,
RelOne: `[]`, RelOne: `[]`,
RelMany: `oap640cot4yru2s`, RelMany: `oap640cot4yru2s`,
NewMultiple: `[]`,
}, },
}, },
} }
for _, s := range scenarios { for _, s := range fieldsScenarios {
result := new(expectation) t.Run("check fields for record "+s.recordId, func(t *testing.T) {
result := new(fieldsExpectation)
err := app.Dao().DB().Select( err := app.Dao().DB().Select(
"select_one", "select_one",
"select_many", "select_many",
"file_one", "file_one",
"file_many", "file_many",
"rel_one", "rel_one",
"rel_many", "rel_many",
).From(collection.Name).Where(dbx.HashExp{"id": s.recordId}).One(result) "new_multiple",
if err != nil { ).From(collection.Name).Where(dbx.HashExp{"id": s.recordId}).One(result)
t.Errorf("[%s] Failed to load record: %v", s.recordId, err) if err != nil {
continue t.Fatalf("Failed to load record: %v", err)
} }
encodedResult, err := json.Marshal(result) encodedResult, err := json.Marshal(result)
if err != nil { if err != nil {
t.Errorf("[%s] Failed to encode result: %v", s.recordId, err) t.Fatalf("Failed to encode result: %v", err)
continue }
}
encodedExpectation, err := json.Marshal(s.expected) encodedExpectation, err := json.Marshal(s.expected)
if err != nil { if err != nil {
t.Errorf("[%s] Failed to encode expectation: %v", s.recordId, err) t.Fatalf("Failed to encode expectation: %v", err)
continue }
}
if !bytes.EqualFold(encodedExpectation, encodedResult) { if !bytes.EqualFold(encodedExpectation, encodedResult) {
t.Errorf("[%s] Expected \n%s, \ngot \n%s", s.recordId, encodedExpectation, encodedResult) t.Fatalf("Expected \n%s, \ngot \n%s", encodedExpectation, encodedResult)
} }
})
} }
} }
+467 -8
View File
@@ -4,7 +4,6 @@ import (
"context" "context"
"database/sql" "database/sql"
"errors" "errors"
"fmt"
"regexp" "regexp"
"strings" "strings"
"testing" "testing"
@@ -19,7 +18,9 @@ import (
"github.com/pocketbase/pocketbase/tools/types" "github.com/pocketbase/pocketbase/tools/types"
) )
func TestRecordQuery(t *testing.T) { func TestRecordQueryWithDifferentCollectionValues(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -28,15 +29,39 @@ func TestRecordQuery(t *testing.T) {
t.Fatal(err) t.Fatal(err)
} }
expected := fmt.Sprintf("SELECT `%s`.* FROM `%s`", collection.Name, collection.Name) scenarios := []struct {
name any
collection any
expectedTotal int
expectError bool
}{
{"with nil value", nil, 0, true},
{"with invalid or missing collection id/name", "missing", 0, true},
{"with pointer model", collection, 3, false},
{"with value model", *collection, 3, false},
{"with name", "demo1", 3, false},
{"with id", "wsmn24bux7wo113", 3, false},
}
sql := app.Dao().RecordQuery(collection).Build().SQL() for _, s := range scenarios {
if sql != expected { var records []*models.Record
t.Errorf("Expected sql %s, got %s", expected, sql) err := app.Dao().RecordQuery(s.collection).All(&records)
hasErr := err != nil
if hasErr != s.expectError {
t.Errorf("[%s] Expected hasError %v, got %v", s.name, s.expectError, hasErr)
continue
}
if total := len(records); total != s.expectedTotal {
t.Errorf("[%s] Expected %d records, got %d", s.name, s.expectedTotal, total)
}
} }
} }
func TestRecordQueryOneWithRecord(t *testing.T) { func TestRecordQueryOneWithRecord(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -61,6 +86,8 @@ func TestRecordQueryOneWithRecord(t *testing.T) {
} }
func TestRecordQueryAllWithRecordsSlices(t *testing.T) { func TestRecordQueryAllWithRecordsSlices(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -122,6 +149,8 @@ func TestRecordQueryAllWithRecordsSlices(t *testing.T) {
} }
func TestFindRecordById(t *testing.T) { func TestFindRecordById(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -182,6 +211,8 @@ func TestFindRecordById(t *testing.T) {
} }
func TestFindRecordsByIds(t *testing.T) { func TestFindRecordsByIds(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -273,6 +304,8 @@ func TestFindRecordsByIds(t *testing.T) {
} }
func TestFindRecordsByExpr(t *testing.T) { func TestFindRecordsByExpr(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -343,6 +376,8 @@ func TestFindRecordsByExpr(t *testing.T) {
} }
func TestFindFirstRecordByData(t *testing.T) { func TestFindFirstRecordByData(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -405,7 +440,415 @@ func TestFindFirstRecordByData(t *testing.T) {
} }
} }
func TestFindRecordsByFilter(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp()
defer app.Cleanup()
scenarios := []struct {
name string
collectionIdOrName string
filter string
sort string
limit int
offset int
params []dbx.Params
expectError bool
expectRecordIds []string
}{
{
"missing collection",
"missing",
"id != ''",
"",
0,
0,
nil,
true,
nil,
},
{
"missing filter",
"demo2",
"",
"",
0,
0,
nil,
true,
nil,
},
{
"invalid filter",
"demo2",
"someMissingField > 1",
"",
0,
0,
nil,
true,
nil,
},
{
"simple filter",
"demo2",
"id != ''",
"",
0,
0,
nil,
false,
[]string{
"llvuca81nly1qls",
"achvryl401bhse3",
"0yxhwia2amd8gec",
},
},
{
"multi-condition filter with sort",
"demo2",
"id != '' && active=true",
"-created,title",
-1, // should behave the same as 0
0,
nil,
false,
[]string{
"0yxhwia2amd8gec",
"achvryl401bhse3",
},
},
{
"with limit and offset",
"demo2",
"id != ''",
"title",
2,
1,
nil,
false,
[]string{
"achvryl401bhse3",
"0yxhwia2amd8gec",
},
},
{
"with placeholder params",
"demo2",
"active = {:active}",
"",
10,
0,
[]dbx.Params{{"active": false}},
false,
[]string{
"llvuca81nly1qls",
},
},
{
"with json filter and sort",
"demo4",
"json_object != null && json_object.a.b = 'test'",
"-json_object.a",
10,
0,
[]dbx.Params{{"active": false}},
false,
[]string{
"i9naidtvr6qsgb4",
},
},
}
for _, s := range scenarios {
t.Run(s.name, func(t *testing.T) {
records, err := app.Dao().FindRecordsByFilter(
s.collectionIdOrName,
s.filter,
s.sort,
s.limit,
s.offset,
s.params...,
)
hasErr := err != nil
if hasErr != s.expectError {
t.Fatalf("[%s] Expected hasErr to be %v, got %v (%v)", s.name, s.expectError, hasErr, err)
}
if hasErr {
return
}
if len(records) != len(s.expectRecordIds) {
t.Fatalf("[%s] Expected %d records, got %d", s.name, len(s.expectRecordIds), len(records))
}
for i, id := range s.expectRecordIds {
if id != records[i].Id {
t.Fatalf("[%s] Expected record with id %q, got %q at index %d", s.name, id, records[i].Id, i)
}
}
})
}
}
func TestFindFirstRecordByFilter(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp()
defer app.Cleanup()
scenarios := []struct {
name string
collectionIdOrName string
filter string
params []dbx.Params
expectError bool
expectRecordId string
}{
{
"missing collection",
"missing",
"id != ''",
nil,
true,
"",
},
{
"missing filter",
"demo2",
"",
nil,
true,
"",
},
{
"invalid filter",
"demo2",
"someMissingField > 1",
nil,
true,
"",
},
{
"valid filter but no matches",
"demo2",
"id = 'test'",
nil,
true,
"",
},
{
"valid filter and multiple matches",
"demo2",
"id != ''",
nil,
false,
"llvuca81nly1qls",
},
{
"with placeholder params",
"demo2",
"active = {:active}",
[]dbx.Params{{"active": false}},
false,
"llvuca81nly1qls",
},
}
for _, s := range scenarios {
record, err := app.Dao().FindFirstRecordByFilter(s.collectionIdOrName, s.filter, s.params...)
hasErr := err != nil
if hasErr != s.expectError {
t.Errorf("[%s] Expected hasErr to be %v, got %v (%v)", s.name, s.expectError, hasErr, err)
continue
}
if hasErr {
continue
}
if record.Id != s.expectRecordId {
t.Errorf("[%s] Expected record with id %q, got %q", s.name, s.expectRecordId, record.Id)
}
}
}
func TestCanAccessRecord(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp()
defer app.Cleanup()
admin, err := app.Dao().FindAdminByEmail("test@example.com")
if err != nil {
t.Fatal(err)
}
authRecord, err := app.Dao().FindAuthRecordByEmail("users", "test@example.com")
if err != nil {
t.Fatal(err)
}
record, err := app.Dao().FindRecordById("demo1", "imy661ixudk5izi")
if err != nil {
t.Fatal(err)
}
scenarios := []struct {
name string
record *models.Record
requestInfo *models.RequestInfo
rule *string
expected bool
expectError bool
}{
{
"as admin with nil rule",
record,
&models.RequestInfo{
Admin: admin,
},
nil,
true,
false,
},
{
"as admin with non-empty rule",
record,
&models.RequestInfo{
Admin: admin,
},
types.Pointer("id = ''"), // the filter rule should be ignored
true,
false,
},
{
"as admin with invalid rule",
record,
&models.RequestInfo{
Admin: admin,
},
types.Pointer("id ?!@ 1"), // the filter rule should be ignored
true,
false,
},
{
"as guest with nil rule",
record,
&models.RequestInfo{},
nil,
false,
false,
},
{
"as guest with empty rule",
record,
&models.RequestInfo{},
types.Pointer(""),
true,
false,
},
{
"as guest with invalid rule",
record,
&models.RequestInfo{},
types.Pointer("id ?!@ 1"),
false,
true,
},
{
"as guest with mismatched rule",
record,
&models.RequestInfo{},
types.Pointer("@request.auth.id != ''"),
false,
false,
},
{
"as guest with matched rule",
record,
&models.RequestInfo{
Data: map[string]any{"test": 1},
},
types.Pointer("@request.auth.id != '' || @request.data.test = 1"),
true,
false,
},
{
"as auth record with nil rule",
record,
&models.RequestInfo{
AuthRecord: authRecord,
},
nil,
false,
false,
},
{
"as auth record with empty rule",
record,
&models.RequestInfo{
AuthRecord: authRecord,
},
types.Pointer(""),
true,
false,
},
{
"as auth record with invalid rule",
record,
&models.RequestInfo{
AuthRecord: authRecord,
},
types.Pointer("id ?!@ 1"),
false,
true,
},
{
"as auth record with mismatched rule",
record,
&models.RequestInfo{
AuthRecord: authRecord,
Data: map[string]any{"test": 1},
},
types.Pointer("@request.auth.id != '' && @request.data.test > 1"),
false,
false,
},
{
"as auth record with matched rule",
record,
&models.RequestInfo{
AuthRecord: authRecord,
Data: map[string]any{"test": 2},
},
types.Pointer("@request.auth.id != '' && @request.data.test > 1"),
true,
false,
},
}
for _, s := range scenarios {
result, err := app.Dao().CanAccessRecord(s.record, s.requestInfo, s.rule)
if result != s.expected {
t.Errorf("[%s] Expected %v, got %v", s.name, s.expected, result)
}
hasErr := err != nil
if hasErr != s.expectError {
t.Errorf("[%s] Expected hasErr %v, got %v (%v)", s.name, s.expectError, hasErr, err)
}
}
}
func TestIsRecordValueUnique(t *testing.T) { func TestIsRecordValueUnique(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -455,6 +898,8 @@ func TestIsRecordValueUnique(t *testing.T) {
} }
func TestFindAuthRecordByToken(t *testing.T) { func TestFindAuthRecordByToken(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -517,6 +962,8 @@ func TestFindAuthRecordByToken(t *testing.T) {
} }
func TestFindAuthRecordByEmail(t *testing.T) { func TestFindAuthRecordByEmail(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -548,6 +995,8 @@ func TestFindAuthRecordByEmail(t *testing.T) {
} }
func TestFindAuthRecordByUsername(t *testing.T) { func TestFindAuthRecordByUsername(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -580,6 +1029,8 @@ func TestFindAuthRecordByUsername(t *testing.T) {
} }
func TestSuggestUniqueAuthRecordUsername(t *testing.T) { func TestSuggestUniqueAuthRecordUsername(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -617,6 +1068,8 @@ func TestSuggestUniqueAuthRecordUsername(t *testing.T) {
} }
func TestSaveRecord(t *testing.T) { func TestSaveRecord(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -650,6 +1103,8 @@ func TestSaveRecord(t *testing.T) {
} }
func TestSaveRecordWithIdFromOtherCollection(t *testing.T) { func TestSaveRecordWithIdFromOtherCollection(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -682,6 +1137,8 @@ func TestSaveRecordWithIdFromOtherCollection(t *testing.T) {
} }
func TestDeleteRecord(t *testing.T) { func TestDeleteRecord(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -748,17 +1205,19 @@ func TestDeleteRecord(t *testing.T) {
} }
// ensure that the json rel fields were prefixed // ensure that the json rel fields were prefixed
joinedQueries := strings.Join(calledQueries, " ") joinedQueries := strings.Join(calledQueries, " ")
expectedRelManyPart := "`demo1` INNER JOIN json_each(CASE WHEN json_valid([[demo1.rel_many]]) THEN [[demo1.rel_many]] ELSE json_array([[demo1.rel_many]]) END)" expectedRelManyPart := "SELECT `demo1`.* FROM `demo1` WHERE EXISTS (SELECT 1 FROM json_each(CASE WHEN json_valid([[demo1.rel_many]]) THEN [[demo1.rel_many]] ELSE json_array([[demo1.rel_many]]) END) {{__je__}} WHERE [[__je__.value]]='"
if !strings.Contains(joinedQueries, expectedRelManyPart) { if !strings.Contains(joinedQueries, expectedRelManyPart) {
t.Fatalf("(rec3) Expected the cascade delete to call the query \n%v, got \n%v", expectedRelManyPart, calledQueries) t.Fatalf("(rec3) Expected the cascade delete to call the query \n%v, got \n%v", expectedRelManyPart, calledQueries)
} }
expectedRelOnePart := "SELECT DISTINCT `demo1`.* FROM `demo1` WHERE (`demo1`.`rel_one`=" expectedRelOnePart := "SELECT `demo1`.* FROM `demo1` WHERE (`demo1`.`rel_one`='"
if !strings.Contains(joinedQueries, expectedRelOnePart) { if !strings.Contains(joinedQueries, expectedRelOnePart) {
t.Fatalf("(rec3) Expected the cascade delete to call the query \n%v, got \n%v", expectedRelOnePart, calledQueries) t.Fatalf("(rec3) Expected the cascade delete to call the query \n%v, got \n%v", expectedRelOnePart, calledQueries)
} }
} }
func TestDeleteRecordBatchProcessing(t *testing.T) { func TestDeleteRecordBatchProcessing(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
-70
View File
@@ -1,70 +0,0 @@
package daos
import (
"time"
"github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tools/types"
)
// RequestQuery returns a new Request logs select query.
func (dao *Dao) RequestQuery() *dbx.SelectQuery {
return dao.ModelQuery(&models.Request{})
}
// FindRequestById finds a single Request log by its id.
func (dao *Dao) FindRequestById(id string) (*models.Request, error) {
model := &models.Request{}
err := dao.RequestQuery().
AndWhere(dbx.HashExp{"id": id}).
Limit(1).
One(model)
if err != nil {
return nil, err
}
return model, nil
}
type RequestsStatsItem struct {
Total int `db:"total" json:"total"`
Date types.DateTime `db:"date" json:"date"`
}
// RequestsStats returns hourly grouped requests logs statistics.
func (dao *Dao) RequestsStats(expr dbx.Expression) ([]*RequestsStatsItem, error) {
result := []*RequestsStatsItem{}
query := dao.RequestQuery().
Select("count(id) as total", "strftime('%Y-%m-%d %H:00:00', created) as date").
GroupBy("date")
if expr != nil {
query.AndWhere(expr)
}
err := query.All(&result)
return result, err
}
// DeleteOldRequests delete all requests that are created before createdBefore.
func (dao *Dao) DeleteOldRequests(createdBefore time.Time) error {
m := models.Request{}
tableName := m.TableName()
formattedDate := createdBefore.UTC().Format(types.DefaultDateLayout)
expr := dbx.NewExp("[[created]] <= {:date}", dbx.Params{"date": formattedDate})
_, err := dao.NonconcurrentDB().Delete(tableName, expr).Execute()
return err
}
// SaveRequest upserts the provided Request model.
func (dao *Dao) SaveRequest(request *models.Request) error {
return dao.Save(request)
}
+2
View File
@@ -8,6 +8,8 @@ import (
) )
func TestSaveAndFindSettings(t *testing.T) { func TestSaveAndFindSettings(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+12
View File
@@ -12,6 +12,8 @@ import (
) )
func TestHasTable(t *testing.T) { func TestHasTable(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -36,6 +38,8 @@ func TestHasTable(t *testing.T) {
} }
func TestTableColumns(t *testing.T) { func TestTableColumns(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -64,6 +68,8 @@ func TestTableColumns(t *testing.T) {
} }
func TestTableInfo(t *testing.T) { func TestTableInfo(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -92,6 +98,8 @@ func TestTableInfo(t *testing.T) {
} }
func TestDeleteTable(t *testing.T) { func TestDeleteTable(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -116,6 +124,8 @@ func TestDeleteTable(t *testing.T) {
} }
func TestVacuum(t *testing.T) { func TestVacuum(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -141,6 +151,8 @@ func TestVacuum(t *testing.T) {
} }
func TestTableIndexes(t *testing.T) { func TestTableIndexes(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+53 -16
View File
@@ -43,10 +43,10 @@ func (dao *Dao) SaveView(name string, selectQuery string) error {
return err return err
} }
trimmed := strings.Trim(selectQuery, ";") selectQuery = strings.Trim(strings.TrimSpace(selectQuery), ";")
// try to eagerly detect multiple inline statements // try to eagerly detect multiple inline statements
tk := tokenizer.NewFromString(trimmed) tk := tokenizer.NewFromString(selectQuery)
tk.Separators(';') tk.Separators(';')
if queryParts, _ := tk.ScanAll(); len(queryParts) > 1 { if queryParts, _ := tk.ScanAll(); len(queryParts) > 1 {
return errors.New("multiple statements are not supported") return errors.New("multiple statements are not supported")
@@ -56,7 +56,7 @@ func (dao *Dao) SaveView(name string, selectQuery string) error {
// //
// note: the query is wrapped in a secondary SELECT as a rudimentary // note: the query is wrapped in a secondary SELECT as a rudimentary
// measure to discourage multiple inline sql statements execution. // measure to discourage multiple inline sql statements execution.
viewQuery := fmt.Sprintf("CREATE VIEW {{%s}} AS SELECT * FROM (%s)", name, trimmed) viewQuery := fmt.Sprintf("CREATE VIEW {{%s}} AS SELECT * FROM (%s)", name, selectQuery)
if _, err := txDao.DB().NewQuery(viewQuery).Execute(); err != nil { if _, err := txDao.DB().NewQuery(viewQuery).Execute(); err != nil {
return err return err
} }
@@ -232,9 +232,14 @@ func defaultViewField(name string) *schema.SchemaField {
return &schema.SchemaField{ return &schema.SchemaField{
Name: name, Name: name,
Type: schema.FieldTypeJson, Type: schema.FieldTypeJson,
Options: &schema.JsonOptions{
MaxSize: 1, // the size doesn't matter in this case
},
} }
} }
var castRegex = regexp.MustCompile(`(?i)^cast\s*\(.*\s+as\s+(\w+)\s*\)$`)
func (dao *Dao) parseQueryToFields(selectQuery string) (map[string]*queryField, error) { func (dao *Dao) parseQueryToFields(selectQuery string) (map[string]*queryField, error) {
p := new(identifiersParser) p := new(identifiersParser)
if err := p.parse(selectQuery); err != nil { if err := p.parse(selectQuery); err != nil {
@@ -257,15 +262,8 @@ func (dao *Dao) parseQueryToFields(selectQuery string) (map[string]*queryField,
for _, col := range p.columns { for _, col := range p.columns {
colLower := strings.ToLower(col.original) colLower := strings.ToLower(col.original)
// numeric expression cast // numeric aggregations
if strings.Contains(colLower, "(") && if strings.HasPrefix(colLower, "count(") || strings.HasPrefix(colLower, "total(") {
(strings.HasPrefix(colLower, "count(") ||
strings.HasPrefix(colLower, "total(") ||
strings.Contains(colLower, " as numeric") ||
strings.Contains(colLower, " as real") ||
strings.Contains(colLower, " as int") ||
strings.Contains(colLower, " as integer") ||
strings.Contains(colLower, " as decimal")) {
result[col.alias] = &queryField{ result[col.alias] = &queryField{
field: &schema.SchemaField{ field: &schema.SchemaField{
Name: col.alias, Name: col.alias,
@@ -275,6 +273,38 @@ func (dao *Dao) parseQueryToFields(selectQuery string) (map[string]*queryField,
continue continue
} }
castMatch := castRegex.FindStringSubmatch(colLower)
// numeric casts
if len(castMatch) == 2 {
switch castMatch[1] {
case "real", "integer", "int", "decimal", "numeric":
result[col.alias] = &queryField{
field: &schema.SchemaField{
Name: col.alias,
Type: schema.FieldTypeNumber,
},
}
continue
case "text":
result[col.alias] = &queryField{
field: &schema.SchemaField{
Name: col.alias,
Type: schema.FieldTypeText,
},
}
continue
case "boolean", "bool":
result[col.alias] = &queryField{
field: &schema.SchemaField{
Name: col.alias,
Type: schema.FieldTypeBool,
},
}
continue
}
}
parts := strings.Split(col.original, ".") parts := strings.Split(col.original, ".")
var fieldName string var fieldName string
@@ -431,7 +461,7 @@ type identifiersParser struct {
} }
func (p *identifiersParser) parse(selectQuery string) error { func (p *identifiersParser) parse(selectQuery string) error {
str := strings.Trim(selectQuery, ";") str := strings.Trim(strings.TrimSpace(selectQuery), ";")
str = joinReplaceRegex.ReplaceAllString(str, " _join_ ") str = joinReplaceRegex.ReplaceAllString(str, " _join_ ")
str = discardReplaceRegex.ReplaceAllString(str, " _discard_ ") str = discardReplaceRegex.ReplaceAllString(str, " _discard_ ")
str = commentsReplaceRegex.ReplaceAllString(str, "") str = commentsReplaceRegex.ReplaceAllString(str, "")
@@ -572,13 +602,20 @@ func identifierFromParts(parts []string) (identifier, error) {
} }
result.original = trimRawIdentifier(result.original) result.original = trimRawIdentifier(result.original)
result.alias = trimRawIdentifier(result.alias)
// we trim the single quote even though it is not a valid column quote character
// because SQLite allows it if the context expects an identifier and not string literal
// (https://www.sqlite.org/lang_keywords.html)
result.alias = trimRawIdentifier(result.alias, "'")
return result, nil return result, nil
} }
func trimRawIdentifier(rawIdentifier string) string { func trimRawIdentifier(rawIdentifier string, extraTrimChars ...string) string {
const trimChars = "`\"[];" trimChars := "`\"[];"
if len(extraTrimChars) > 0 {
trimChars += strings.Join(extraTrimChars, "")
}
parts := strings.Split(rawIdentifier, ".") parts := strings.Split(rawIdentifier, ".")
+58 -41
View File
@@ -33,6 +33,8 @@ func ensureNoTempViews(app core.App, t *testing.T) {
} }
func TestDeleteView(t *testing.T) { func TestDeleteView(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -60,6 +62,8 @@ func TestDeleteView(t *testing.T) {
} }
func TestSaveView(t *testing.T) { func TestSaveView(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -147,40 +151,41 @@ func TestSaveView(t *testing.T) {
} }
for _, s := range scenarios { for _, s := range scenarios {
err := app.Dao().SaveView(s.viewName, s.query) t.Run(s.scenarioName, func(t *testing.T) {
err := app.Dao().SaveView(s.viewName, s.query)
hasErr := err != nil hasErr := err != nil
if hasErr != s.expectError { if hasErr != s.expectError {
t.Errorf("[%s] Expected hasErr %v, got %v (%v)", s.scenarioName, s.expectError, hasErr, err) t.Fatalf("Expected hasErr %v, got %v (%v)", s.expectError, hasErr, err)
continue
}
if hasErr {
continue
}
infoRows, err := app.Dao().TableInfo(s.viewName)
if err != nil {
t.Errorf("[%s] Failed to fetch table info for %s: %v", s.scenarioName, s.viewName, err)
continue
}
if len(s.expectColumns) != len(infoRows) {
t.Errorf("[%s] Expected %d columns, got %d", s.scenarioName, len(s.expectColumns), len(infoRows))
continue
}
for _, row := range infoRows {
if !list.ExistInSlice(row.Name, s.expectColumns) {
t.Errorf("[%s] Missing %q column in %v", s.scenarioName, row.Name, s.expectColumns)
} }
}
if hasErr {
return
}
infoRows, err := app.Dao().TableInfo(s.viewName)
if err != nil {
t.Fatalf("Failed to fetch table info for %s: %v", s.viewName, err)
}
if len(s.expectColumns) != len(infoRows) {
t.Fatalf("Expected %d columns, got %d", len(s.expectColumns), len(infoRows))
}
for _, row := range infoRows {
if !list.ExistInSlice(row.Name, s.expectColumns) {
t.Fatalf("Missing %q column in %v", row.Name, s.expectColumns)
}
}
})
} }
ensureNoTempViews(app, t) ensureNoTempViews(app, t)
} }
func TestCreateViewSchemaWithDiscardedNestedTransaction(t *testing.T) { func TestCreateViewSchemaWithDiscardedNestedTransaction(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -197,6 +202,8 @@ func TestCreateViewSchemaWithDiscardedNestedTransaction(t *testing.T) {
} }
func TestCreateViewSchema(t *testing.T) { func TestCreateViewSchema(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -272,24 +279,26 @@ func TestCreateViewSchema(t *testing.T) {
"datetime", "datetime",
"json", "json",
"rel_one", "rel_one",
"rel_many" "rel_many",
'single_quoted_custom_literal' as 'single_quoted_column'
from demo1 from demo1
`, `,
false, false,
map[string]string{ map[string]string{
"text": schema.FieldTypeText, "text": schema.FieldTypeText,
"bool": schema.FieldTypeBool, "bool": schema.FieldTypeBool,
"url": schema.FieldTypeUrl, "url": schema.FieldTypeUrl,
"select_one": schema.FieldTypeSelect, "select_one": schema.FieldTypeSelect,
"select_many": schema.FieldTypeSelect, "select_many": schema.FieldTypeSelect,
"file_one": schema.FieldTypeFile, "file_one": schema.FieldTypeFile,
"file_many": schema.FieldTypeFile, "file_many": schema.FieldTypeFile,
"number_alias": schema.FieldTypeNumber, "number_alias": schema.FieldTypeNumber,
"email": schema.FieldTypeEmail, "email": schema.FieldTypeEmail,
"datetime": schema.FieldTypeDate, "datetime": schema.FieldTypeDate,
"json": schema.FieldTypeJson, "json": schema.FieldTypeJson,
"rel_one": schema.FieldTypeRelation, "rel_one": schema.FieldTypeRelation,
"rel_many": schema.FieldTypeRelation, "rel_many": schema.FieldTypeRelation,
"single_quoted_column": schema.FieldTypeJson,
}, },
}, },
{ {
@@ -330,7 +339,7 @@ func TestCreateViewSchema(t *testing.T) {
}, },
}, },
{ {
"query with numeric casts", "query with casts",
`select `select
a.id, a.id,
count(a.id) count, count(a.id) count,
@@ -339,6 +348,9 @@ func TestCreateViewSchema(t *testing.T) {
cast(a.id as real) cast_real, cast(a.id as real) cast_real,
cast(a.id as decimal) cast_decimal, cast(a.id as decimal) cast_decimal,
cast(a.id as numeric) cast_numeric, cast(a.id as numeric) cast_numeric,
cast(a.id as text) cast_text,
cast(a.id as bool) cast_bool,
cast(a.id as boolean) cast_boolean,
avg(a.id) avg, avg(a.id) avg,
sum(a.id) sum, sum(a.id) sum,
total(a.id) total, total(a.id) total,
@@ -354,6 +366,9 @@ func TestCreateViewSchema(t *testing.T) {
"cast_real": schema.FieldTypeNumber, "cast_real": schema.FieldTypeNumber,
"cast_decimal": schema.FieldTypeNumber, "cast_decimal": schema.FieldTypeNumber,
"cast_numeric": schema.FieldTypeNumber, "cast_numeric": schema.FieldTypeNumber,
"cast_text": schema.FieldTypeText,
"cast_bool": schema.FieldTypeBool,
"cast_boolean": schema.FieldTypeBool,
// json because they are nullable // json because they are nullable
"sum": schema.FieldTypeJson, "sum": schema.FieldTypeJson,
"avg": schema.FieldTypeJson, "avg": schema.FieldTypeJson,
@@ -479,6 +494,8 @@ func TestCreateViewSchema(t *testing.T) {
} }
func TestFindRecordByViewFile(t *testing.T) { func TestFindRecordByViewFile(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+34 -6
View File
@@ -22,6 +22,30 @@ func main() {
// Optional plugin flags: // Optional plugin flags:
// --------------------------------------------------------------- // ---------------------------------------------------------------
var hooksDir string
app.RootCmd.PersistentFlags().StringVar(
&hooksDir,
"hooksDir",
"",
"the directory with the JS app hooks",
)
var hooksWatch bool
app.RootCmd.PersistentFlags().BoolVar(
&hooksWatch,
"hooksWatch",
true,
"auto restart the app on pb_hooks file change",
)
var hooksPool int
app.RootCmd.PersistentFlags().IntVar(
&hooksPool,
"hooksPool",
25,
"the total prewarm goja.Runtime instances for the JS app hooks execution",
)
var migrationsDir string var migrationsDir string
app.RootCmd.PersistentFlags().StringVar( app.RootCmd.PersistentFlags().StringVar(
&migrationsDir, &migrationsDir,
@@ -68,22 +92,25 @@ func main() {
// Plugins and hooks: // Plugins and hooks:
// --------------------------------------------------------------- // ---------------------------------------------------------------
// load js pb_migrations // load jsvm (hooks and migrations)
jsvm.MustRegisterMigrations(app, &jsvm.MigrationsOptions{ jsvm.MustRegister(app, jsvm.Config{
Dir: migrationsDir, MigrationsDir: migrationsDir,
HooksDir: hooksDir,
HooksWatch: hooksWatch,
HooksPoolSize: hooksPool,
}) })
// migrate command (with js templates) // migrate command (with js templates)
migratecmd.MustRegister(app, app.RootCmd, &migratecmd.Options{ migratecmd.MustRegister(app, app.RootCmd, migratecmd.Config{
TemplateLang: migratecmd.TemplateLangJS, TemplateLang: migratecmd.TemplateLangJS,
Automigrate: automigrate, Automigrate: automigrate,
Dir: migrationsDir, Dir: migrationsDir,
}) })
// GitHub selfupdate // GitHub selfupdate
ghupdate.MustRegister(app, app.RootCmd, nil) ghupdate.MustRegister(app, app.RootCmd, ghupdate.Config{})
app.OnAfterBootstrap().Add(func(e *core.BootstrapEvent) error { app.OnAfterBootstrap().PreAdd(func(e *core.BootstrapEvent) error {
app.Dao().ModelQueryTimeout = time.Duration(queryTimeout) * time.Second app.Dao().ModelQueryTimeout = time.Duration(queryTimeout) * time.Second
return nil return nil
}) })
@@ -105,5 +132,6 @@ func defaultPublicDir() string {
// most likely ran with go run // most likely ran with go run
return "./pb_public" return "./pb_public"
} }
return filepath.Join(os.Args[0], "../pb_public") return filepath.Join(os.Args[0], "../pb_public")
} }
+4
View File
@@ -10,6 +10,8 @@ import (
) )
func TestAdminLoginValidateAndSubmit(t *testing.T) { func TestAdminLoginValidateAndSubmit(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -51,6 +53,8 @@ func TestAdminLoginValidateAndSubmit(t *testing.T) {
} }
func TestAdminLoginInterceptors(t *testing.T) { func TestAdminLoginInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -11,6 +11,8 @@ import (
) )
func TestAdminPasswordResetConfirmValidateAndSubmit(t *testing.T) { func TestAdminPasswordResetConfirmValidateAndSubmit(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -99,6 +101,8 @@ func TestAdminPasswordResetConfirmValidateAndSubmit(t *testing.T) {
} }
func TestAdminPasswordResetConfirmInterceptors(t *testing.T) { func TestAdminPasswordResetConfirmInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
+5 -4
View File
@@ -2,6 +2,7 @@ package forms
import ( import (
"errors" "errors"
"fmt"
"time" "time"
validation "github.com/go-ozzo/ozzo-validation/v4" validation "github.com/go-ozzo/ozzo-validation/v4"
@@ -66,7 +67,7 @@ func (form *AdminPasswordResetRequest) Submit(interceptors ...InterceptorFunc[*m
admin, err := form.dao.FindAdminByEmail(form.Email) admin, err := form.dao.FindAdminByEmail(form.Email)
if err != nil { if err != nil {
return err return fmt.Errorf("Failed to fetch admin with email %s: %w", form.Email, err)
} }
now := time.Now().UTC() now := time.Now().UTC()
@@ -75,14 +76,14 @@ func (form *AdminPasswordResetRequest) Submit(interceptors ...InterceptorFunc[*m
return errors.New("You have already requested a password reset.") return errors.New("You have already requested a password reset.")
} }
// update last sent timestamp
admin.LastResetSentAt = types.NowDateTime()
return runInterceptors(admin, func(m *models.Admin) error { return runInterceptors(admin, func(m *models.Admin) error {
if err := mails.SendAdminPasswordReset(form.app, m); err != nil { if err := mails.SendAdminPasswordReset(form.app, m); err != nil {
return err return err
} }
// update last sent timestamp
m.LastResetSentAt = types.NowDateTime()
return form.dao.SaveAdmin(m) return form.dao.SaveAdmin(m)
}, interceptors...) }, interceptors...)
} }
+6 -2
View File
@@ -10,6 +10,8 @@ import (
) )
func TestAdminPasswordResetRequestValidateAndSubmit(t *testing.T) { func TestAdminPasswordResetRequestValidateAndSubmit(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -74,6 +76,8 @@ func TestAdminPasswordResetRequestValidateAndSubmit(t *testing.T) {
} }
func TestAdminPasswordResetRequestInterceptors(t *testing.T) { func TestAdminPasswordResetRequestInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -117,7 +121,7 @@ func TestAdminPasswordResetRequestInterceptors(t *testing.T) {
t.Fatalf("Expected interceptor2 to be called") t.Fatalf("Expected interceptor2 to be called")
} }
if interceptorLastResetSentAt.String() == admin.LastResetSentAt.String() { if interceptorLastResetSentAt.String() != admin.LastResetSentAt.String() {
t.Fatalf("Expected the form model to be filled before calling the interceptors") t.Fatalf("Expected the form model to NOT be filled before calling the interceptors")
} }
} }
+8
View File
@@ -12,6 +12,8 @@ import (
) )
func TestNewAdminUpsert(t *testing.T) { func TestNewAdminUpsert(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -31,6 +33,8 @@ func TestNewAdminUpsert(t *testing.T) {
} }
func TestAdminUpsertValidateAndSubmit(t *testing.T) { func TestAdminUpsertValidateAndSubmit(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -183,6 +187,8 @@ func TestAdminUpsertValidateAndSubmit(t *testing.T) {
} }
func TestAdminUpsertSubmitInterceptors(t *testing.T) { func TestAdminUpsertSubmitInterceptors(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -231,6 +237,8 @@ func TestAdminUpsertSubmitInterceptors(t *testing.T) {
} }
func TestAdminUpsertWithCustomId(t *testing.T) { func TestAdminUpsertWithCustomId(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+2 -2
View File
@@ -12,7 +12,7 @@ import (
var privateKeyRegex = regexp.MustCompile(`(?m)-----BEGIN PRIVATE KEY----[\s\S]+-----END PRIVATE KEY-----`) var privateKeyRegex = regexp.MustCompile(`(?m)-----BEGIN PRIVATE KEY----[\s\S]+-----END PRIVATE KEY-----`)
// AppleClientSecretCreate is a [models.Admin] upsert (create/update) form. // AppleClientSecretCreate is a form struct to generate a new Apple Client Secret.
// //
// Reference: https://developer.apple.com/documentation/sign_in_with_apple/generate_and_validate_tokens // Reference: https://developer.apple.com/documentation/sign_in_with_apple/generate_and_validate_tokens
type AppleClientSecretCreate struct { type AppleClientSecretCreate struct {
@@ -33,7 +33,7 @@ type AppleClientSecretCreate struct {
// Usually wrapped within -----BEGIN PRIVATE KEY----- X -----END PRIVATE KEY-----. // Usually wrapped within -----BEGIN PRIVATE KEY----- X -----END PRIVATE KEY-----.
PrivateKey string `form:"privateKey" json:"privateKey"` PrivateKey string `form:"privateKey" json:"privateKey"`
// Duration specifies how long the generated JWT token should be considered valid. // Duration specifies how long the generated JWT should be considered valid.
// The specified value must be in seconds and max 15777000 (~6months). // The specified value must be in seconds and max 15777000 (~6months).
Duration int `form:"duration" json:"duration"` Duration int `form:"duration" json:"duration"`
} }
+2
View File
@@ -15,6 +15,8 @@ import (
) )
func TestAppleClientSecretCreateValidateAndSubmit(t *testing.T) { func TestAppleClientSecretCreateValidateAndSubmit(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+12 -10
View File
@@ -10,6 +10,8 @@ import (
) )
func TestBackupCreateValidateAndSubmit(t *testing.T) { func TestBackupCreateValidateAndSubmit(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
name string name string
backupName string backupName string
@@ -38,7 +40,7 @@ func TestBackupCreateValidateAndSubmit(t *testing.T) {
} }
for _, s := range scenarios { for _, s := range scenarios {
func() { t.Run(s.name, func(t *testing.T) {
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -56,47 +58,47 @@ func TestBackupCreateValidateAndSubmit(t *testing.T) {
// parse errors // parse errors
errs, ok := result.(validation.Errors) errs, ok := result.(validation.Errors)
if !ok && result != nil { if !ok && result != nil {
t.Errorf("[%s] Failed to parse errors %v", s.name, result) t.Fatalf("Failed to parse errors %v", result)
return return
} }
// check errors // check errors
if len(errs) > len(s.expectedErrors) { if len(errs) > len(s.expectedErrors) {
t.Errorf("[%s] Expected error keys %v, got %v", s.name, s.expectedErrors, errs) t.Fatalf("Expected error keys %v, got %v", s.expectedErrors, errs)
} }
for _, k := range s.expectedErrors { for _, k := range s.expectedErrors {
if _, ok := errs[k]; !ok { if _, ok := errs[k]; !ok {
t.Errorf("[%s] Missing expected error key %q in %v", s.name, k, errs) t.Fatalf("Missing expected error key %q in %v", k, errs)
} }
} }
// retrieve all created backup files // retrieve all created backup files
files, err := fsys.List("") files, err := fsys.List("")
if err != nil { if err != nil {
t.Errorf("[%s] Failed to retrieve backup files", s.name) t.Fatal("Failed to retrieve backup files")
return return
} }
if result != nil { if result != nil {
if total := len(files); total != 0 { if total := len(files); total != 0 {
t.Errorf("[%s] Didn't expected backup files, found %d", s.name, total) t.Fatalf("Didn't expected backup files, found %d", total)
} }
return return
} }
if total := len(files); total != 1 { if total := len(files); total != 1 {
t.Errorf("[%s] Expected 1 backup file, got %d", s.name, total) t.Fatalf("Expected 1 backup file, got %d", total)
return return
} }
if s.backupName == "" { if s.backupName == "" {
prefix := "pb_backup_" prefix := "pb_backup_"
if !strings.HasPrefix(files[0].Key, prefix) { if !strings.HasPrefix(files[0].Key, prefix) {
t.Errorf("[%s] Expected the backup file, to have prefix %q: %q", s.name, prefix, files[0].Key) t.Fatalf("Expected the backup file, to have prefix %q: %q", prefix, files[0].Key)
} }
} else if s.backupName != files[0].Key { } else if s.backupName != files[0].Key {
t.Errorf("[%s] Expected backup file %q, got %q", s.name, s.backupName, files[0].Key) t.Fatalf("Expected backup file %q, got %q", s.backupName, files[0].Key)
} }
}() })
} }
} }
+85
View File
@@ -0,0 +1,85 @@
package forms
import (
"context"
validation "github.com/go-ozzo/ozzo-validation/v4"
"github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/forms/validators"
"github.com/pocketbase/pocketbase/tools/filesystem"
)
// BackupUpload is a request form for uploading a new app backup.
type BackupUpload struct {
app core.App
ctx context.Context
File *filesystem.File `json:"file"`
}
// NewBackupUpload creates new BackupUpload request form.
func NewBackupUpload(app core.App) *BackupUpload {
return &BackupUpload{
app: app,
ctx: context.Background(),
}
}
// SetContext replaces the default form upload context with the provided one.
func (form *BackupUpload) SetContext(ctx context.Context) {
form.ctx = ctx
}
// Validate makes the form validatable by implementing [validation.Validatable] interface.
func (form *BackupUpload) Validate() error {
return validation.ValidateStruct(form,
validation.Field(
&form.File,
validation.Required,
validation.By(validators.UploadedFileMimeType([]string{"application/zip"})),
validation.By(form.checkUniqueName),
),
)
}
func (form *BackupUpload) checkUniqueName(value any) error {
v, _ := value.(*filesystem.File)
if v == nil {
return nil // nothing to check
}
fsys, err := form.app.NewBackupsFilesystem()
if err != nil {
return err
}
defer fsys.Close()
fsys.SetContext(form.ctx)
if exists, err := fsys.Exists(v.OriginalName); err != nil || exists {
return validation.NewError("validation_backup_name_exists", "Backup file with the specified name already exists.")
}
return nil
}
// Submit validates the form and upload the backup file.
//
// You can optionally provide a list of InterceptorFunc to further
// modify the form behavior before uploading the backup.
func (form *BackupUpload) Submit(interceptors ...InterceptorFunc[*filesystem.File]) error {
if err := form.Validate(); err != nil {
return err
}
return runInterceptors(form.File, func(file *filesystem.File) error {
fsys, err := form.app.NewBackupsFilesystem()
if err != nil {
return err
}
fsys.SetContext(form.ctx)
return fsys.UploadFile(file, file.OriginalName)
}, interceptors...)
}
+120
View File
@@ -0,0 +1,120 @@
package forms_test
import (
"archive/zip"
"bytes"
"testing"
validation "github.com/go-ozzo/ozzo-validation/v4"
"github.com/pocketbase/pocketbase/forms"
"github.com/pocketbase/pocketbase/tests"
"github.com/pocketbase/pocketbase/tools/filesystem"
)
func TestBackupUploadValidateAndSubmit(t *testing.T) {
t.Parallel()
var zb bytes.Buffer
zw := zip.NewWriter(&zb)
if err := zw.Close(); err != nil {
t.Fatal(err)
}
f0, _ := filesystem.NewFileFromBytes([]byte("test"), "existing")
f1, _ := filesystem.NewFileFromBytes([]byte("456"), "nozip")
f2, _ := filesystem.NewFileFromBytes(zb.Bytes(), "existing")
f3, _ := filesystem.NewFileFromBytes(zb.Bytes(), "zip")
scenarios := []struct {
name string
file *filesystem.File
expectedErrors []string
}{
{
"missing file",
nil,
[]string{"file"},
},
{
"non-zip file",
f1,
[]string{"file"},
},
{
"zip file with non-unique name",
f2,
[]string{"file"},
},
{
"zip file with unique name",
f3,
[]string{},
},
}
for _, s := range scenarios {
t.Run(s.name, func(t *testing.T) {
app, _ := tests.NewTestApp()
defer app.Cleanup()
fsys, err := app.NewBackupsFilesystem()
if err != nil {
t.Fatal(err)
}
defer fsys.Close()
// create a dummy backup file to simulate existing backups
if err := fsys.UploadFile(f0, f0.OriginalName); err != nil {
t.Fatal(err)
}
form := forms.NewBackupUpload(app)
form.File = s.file
result := form.Submit()
// parse errors
errs, ok := result.(validation.Errors)
if !ok && result != nil {
t.Fatalf("Failed to parse errors %v", result)
}
// check errors
if len(errs) > len(s.expectedErrors) {
t.Fatalf("Expected error keys %v, got %v", s.expectedErrors, errs)
}
for _, k := range s.expectedErrors {
if _, ok := errs[k]; !ok {
t.Fatalf("Missing expected error key %q in %v", k, errs)
}
}
expectedFiles := []*filesystem.File{f0}
if result == nil {
expectedFiles = append(expectedFiles, s.file)
}
// retrieve all uploaded backup files
files, err := fsys.List("")
if err != nil {
t.Fatal("Failed to retrieve backup files")
}
if len(files) != len(expectedFiles) {
t.Fatalf("Expected %d files, got %d", len(expectedFiles), len(files))
}
for _, ef := range expectedFiles {
exists := false
for _, f := range files {
if f.Key == ef.OriginalName {
exists = true
break
}
}
if !exists {
t.Fatalf("Missing expected backup file %v", ef.OriginalName)
}
}
})
}
}
+28 -22
View File
@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"regexp" "regexp"
"strconv" "strconv"
"strings"
validation "github.com/go-ozzo/ozzo-validation/v4" validation "github.com/go-ozzo/ozzo-validation/v4"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
@@ -131,6 +132,7 @@ func (form *CollectionUpsert) Validate() error {
validation.Match(collectionNameRegex), validation.Match(collectionNameRegex),
validation.By(form.ensureNoSystemNameChange), validation.By(form.ensureNoSystemNameChange),
validation.By(form.checkUniqueName), validation.By(form.checkUniqueName),
validation.By(form.checkForVia),
), ),
// validates using the type's own validation rules + some collection's specifics // validates using the type's own validation rules + some collection's specifics
validation.Field( validation.Field(
@@ -163,6 +165,19 @@ func (form *CollectionUpsert) Validate() error {
) )
} }
func (form *CollectionUpsert) checkForVia(value any) error {
v, _ := value.(string)
if v == "" {
return nil
}
if strings.Contains(strings.ToLower(v), "_via_") {
return validation.NewError("validation_invalid_name", "The name of the collection cannot contain '_via_'.")
}
return nil
}
func (form *CollectionUpsert) checkUniqueName(value any) error { func (form *CollectionUpsert) checkUniqueName(value any) error {
v, _ := value.(string) v, _ := value.(string)
@@ -229,14 +244,6 @@ func (form *CollectionUpsert) ensureNoFieldsTypeChange(value any) error {
func (form *CollectionUpsert) checkRelationFields(value any) error { func (form *CollectionUpsert) checkRelationFields(value any) error {
v, _ := value.(schema.Schema) v, _ := value.(schema.Schema)
systemDisplayFields := schema.BaseModelFieldNames()
systemDisplayFields = append(systemDisplayFields,
schema.FieldNameUsername,
schema.FieldNameEmail,
schema.FieldNameEmailVisibility,
schema.FieldNameVerified,
)
for i, field := range v.Fields() { for i, field := range v.Fields() {
if field.Type != schema.FieldTypeRelation { if field.Type != schema.FieldTypeRelation {
continue continue
@@ -268,10 +275,10 @@ func (form *CollectionUpsert) checkRelationFields(value any) error {
} }
} }
collection, err := form.dao.FindCollectionByNameOrId(options.CollectionId) relCollection, _ := form.dao.FindCollectionByNameOrId(options.CollectionId)
// validate collectionId // validate collectionId
if err != nil || collection.Id != options.CollectionId { if relCollection == nil || relCollection.Id != options.CollectionId {
return validation.Errors{fmt.Sprint(i): validation.Errors{ return validation.Errors{fmt.Sprint(i): validation.Errors{
"options": validation.Errors{ "options": validation.Errors{
"collectionId": validation.NewError( "collectionId": validation.NewError(
@@ -282,17 +289,16 @@ func (form *CollectionUpsert) checkRelationFields(value any) error {
} }
} }
// validate displayFields (if any) // allow only views to have relations to other views
for _, name := range options.DisplayFields { // (see https://github.com/pocketbase/pocketbase/issues/3000)
if collection.Schema.GetFieldByName(name) == nil && !list.ExistInSlice(name, systemDisplayFields) { if form.Type != models.CollectionTypeView && relCollection.IsView() {
return validation.Errors{fmt.Sprint(i): validation.Errors{ return validation.Errors{fmt.Sprint(i): validation.Errors{
"options": validation.Errors{ "options": validation.Errors{
"displayFields": validation.NewError( "collectionId": validation.NewError(
"validation_field_invalid_relation_displayFields", "validation_field_non_view_base_relation_collection",
fmt.Sprintf("%q does not exist in the related %q collection.", name, collection.Name), "Non view collections are not allowed to have a view relation.",
), ),
}}, }},
}
} }
} }
} }
@@ -379,7 +385,7 @@ func (form *CollectionUpsert) checkRule(value any) error {
_, err := search.FilterData(*v).BuildExpr(r) _, err := search.FilterData(*v).BuildExpr(r)
if err != nil { if err != nil {
return validation.NewError("validation_invalid_rule", "Invalid filter rule.") return validation.NewError("validation_invalid_rule", "Invalid filter rule. Raw error: "+err.Error())
} }
return nil return nil
+189 -146
View File
@@ -16,6 +16,8 @@ import (
) )
func TestNewCollectionUpsert(t *testing.T) { func TestNewCollectionUpsert(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -88,6 +90,8 @@ func TestNewCollectionUpsert(t *testing.T) {
} }
func TestCollectionUpsertValidateAndSubmit(t *testing.T) { func TestCollectionUpsertValidateAndSubmit(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -101,6 +105,17 @@ func TestCollectionUpsertValidateAndSubmit(t *testing.T) {
{"empty create (auth)", "", `{"type":"auth"}`, []string{"name"}}, {"empty create (auth)", "", `{"type":"auth"}`, []string{"name"}},
{"empty create (view)", "", `{"type":"view"}`, []string{"name", "options"}}, {"empty create (view)", "", `{"type":"view"}`, []string{"name", "options"}},
{"empty update", "demo2", "{}", []string{}}, {"empty update", "demo2", "{}", []string{}},
{
"collection and field with _via_ names",
"",
`{
"name": "a_via_b",
"schema": [
{"name":"c_via_d","type":"text"}
]
}`,
[]string{"name", "schema"},
},
{ {
"create failure", "create failure",
"", "",
@@ -171,25 +186,6 @@ func TestCollectionUpsertValidateAndSubmit(t *testing.T) {
}`, }`,
[]string{"schema"}, []string{"schema"},
}, },
{
"create failure - missing relation display field",
"",
`{
"name": "test_new",
"type": "base",
"schema": [
{
"name":"test",
"type":"relation",
"options":{
"collectionId":"wsmn24bux7wo113",
"displayFields":["text", "missing"]
}
}
]
}`,
[]string{"schema"},
},
{ {
"create failure - check auth options validators", "create failure - check auth options validators",
"", "",
@@ -400,6 +396,54 @@ func TestCollectionUpsertValidateAndSubmit(t *testing.T) {
// view tests // view tests
// ----------------------------------------------------------- // -----------------------------------------------------------
{
"base->view relation",
"",
`{
"name": "test_view_relation",
"type": "base",
"schema": [
{
"name": "test",
"type": "relation",
"options":{
"collectionId": "v9gwnfh02gjq1q0"
}
}
]
}`,
[]string{"schema"}, // not allowed
},
{
"auth->view relation",
"",
`{
"name": "test_view_relation",
"type": "auth",
"schema": [
{
"name": "test",
"type": "relation",
"options": {
"collectionId": "v9gwnfh02gjq1q0"
}
}
]
}`,
[]string{"schema"}, // not allowed
},
{
"view->view relation",
"",
`{
"name": "test_view_relation",
"type": "view",
"options": {
"query": "select view1.id, view1.id as rel from view1"
}
}`,
[]string{}, // allowed
},
{ {
"view create failure", "view create failure",
"", "",
@@ -495,141 +539,138 @@ func TestCollectionUpsertValidateAndSubmit(t *testing.T) {
} }
for _, s := range scenarios { for _, s := range scenarios {
collection := &models.Collection{} t.Run(s.testName, func(t *testing.T) {
if s.existingName != "" { collection := &models.Collection{}
var err error if s.existingName != "" {
collection, err = app.Dao().FindCollectionByNameOrId(s.existingName) var err error
if err != nil { collection, err = app.Dao().FindCollectionByNameOrId(s.existingName)
t.Fatal(err) if err != nil {
} t.Fatal(err)
}
form := forms.NewCollectionUpsert(app, collection)
// load data
loadErr := json.Unmarshal([]byte(s.jsonData), form)
if loadErr != nil {
t.Errorf("[%s] Failed to load form data: %v", s.testName, loadErr)
continue
}
interceptorCalls := 0
interceptor := func(next forms.InterceptorNextFunc[*models.Collection]) forms.InterceptorNextFunc[*models.Collection] {
return func(c *models.Collection) error {
interceptorCalls++
return next(c)
}
}
// parse errors
result := form.Submit(interceptor)
errs, ok := result.(validation.Errors)
if !ok && result != nil {
t.Errorf("[%s] Failed to parse errors %v", s.testName, result)
continue
}
// check interceptor calls
expectInterceptorCalls := 1
if len(s.expectedErrors) > 0 {
expectInterceptorCalls = 0
}
if interceptorCalls != expectInterceptorCalls {
t.Errorf("[%s] Expected interceptor to be called %d, got %d", s.testName, expectInterceptorCalls, interceptorCalls)
}
// check errors
if len(errs) > len(s.expectedErrors) {
t.Errorf("[%s] Expected error keys %v, got %v", s.testName, s.expectedErrors, errs)
}
for _, k := range s.expectedErrors {
if _, ok := errs[k]; !ok {
t.Errorf("[%s] Missing expected error key %q in %v", s.testName, k, errs)
}
}
if len(s.expectedErrors) > 0 {
continue
}
collection, _ = app.Dao().FindCollectionByNameOrId(form.Name)
if collection == nil {
t.Errorf("[%s] Expected to find collection %q, got nil", s.testName, form.Name)
continue
}
if form.Name != collection.Name {
t.Errorf("[%s] Expected Name %q, got %q", s.testName, collection.Name, form.Name)
}
if form.Type != collection.Type {
t.Errorf("[%s] Expected Type %q, got %q", s.testName, collection.Type, form.Type)
}
if form.System != collection.System {
t.Errorf("[%s] Expected System %v, got %v", s.testName, collection.System, form.System)
}
if cast.ToString(form.ListRule) != cast.ToString(collection.ListRule) {
t.Errorf("[%s] Expected ListRule %v, got %v", s.testName, collection.ListRule, form.ListRule)
}
if cast.ToString(form.ViewRule) != cast.ToString(collection.ViewRule) {
t.Errorf("[%s] Expected ViewRule %v, got %v", s.testName, collection.ViewRule, form.ViewRule)
}
if cast.ToString(form.CreateRule) != cast.ToString(collection.CreateRule) {
t.Errorf("[%s] Expected CreateRule %v, got %v", s.testName, collection.CreateRule, form.CreateRule)
}
if cast.ToString(form.UpdateRule) != cast.ToString(collection.UpdateRule) {
t.Errorf("[%s] Expected UpdateRule %v, got %v", s.testName, collection.UpdateRule, form.UpdateRule)
}
if cast.ToString(form.DeleteRule) != cast.ToString(collection.DeleteRule) {
t.Errorf("[%s] Expected DeleteRule %v, got %v", s.testName, collection.DeleteRule, form.DeleteRule)
}
rawFormSchema, _ := form.Schema.MarshalJSON()
rawCollectionSchema, _ := collection.Schema.MarshalJSON()
if len(form.Schema.Fields()) != len(collection.Schema.Fields()) {
t.Errorf("[%s] Expected Schema \n%v, \ngot \n%v", s.testName, string(rawCollectionSchema), string(rawFormSchema))
continue
}
for _, f := range form.Schema.Fields() {
if collection.Schema.GetFieldByName(f.Name) == nil {
t.Errorf("[%s] Missing field %s \nin \n%v", s.testName, f.Name, string(rawFormSchema))
continue
}
}
// check indexes (if any)
allIndexes, _ := app.Dao().TableIndexes(form.Name)
for _, formIdx := range form.Indexes {
parsed := dbutils.ParseIndex(formIdx)
parsed.TableName = form.Name
normalizedIdx := parsed.Build()
var exists bool
for _, idx := range allIndexes {
if dbutils.ParseIndex(idx).Build() == normalizedIdx {
exists = true
continue
} }
} }
if !exists { form := forms.NewCollectionUpsert(app, collection)
t.Errorf(
"[%s] Missing index %s \nin \n%v", s.testName, normalizedIdx, allIndexes) // load data
continue loadErr := json.Unmarshal([]byte(s.jsonData), form)
if loadErr != nil {
t.Fatalf("Failed to load form data: %v", loadErr)
} }
}
interceptorCalls := 0
interceptor := func(next forms.InterceptorNextFunc[*models.Collection]) forms.InterceptorNextFunc[*models.Collection] {
return func(c *models.Collection) error {
interceptorCalls++
return next(c)
}
}
// parse errors
result := form.Submit(interceptor)
errs, ok := result.(validation.Errors)
if !ok && result != nil {
t.Fatalf("Failed to parse errors %v", result)
}
// check interceptor calls
expectInterceptorCalls := 1
if len(s.expectedErrors) > 0 {
expectInterceptorCalls = 0
}
if interceptorCalls != expectInterceptorCalls {
t.Fatalf("Expected interceptor to be called %d, got %d", expectInterceptorCalls, interceptorCalls)
}
// check errors
if len(errs) > len(s.expectedErrors) {
t.Fatalf("Expected error keys %v, got %v", s.expectedErrors, errs)
}
for _, k := range s.expectedErrors {
if _, ok := errs[k]; !ok {
t.Fatalf("Missing expected error key %q in %v", k, errs)
}
}
if len(s.expectedErrors) > 0 {
return
}
collection, _ = app.Dao().FindCollectionByNameOrId(form.Name)
if collection == nil {
t.Fatalf("Expected to find collection %q, got nil", form.Name)
}
if form.Name != collection.Name {
t.Fatalf("Expected Name %q, got %q", collection.Name, form.Name)
}
if form.Type != collection.Type {
t.Fatalf("Expected Type %q, got %q", collection.Type, form.Type)
}
if form.System != collection.System {
t.Fatalf("Expected System %v, got %v", collection.System, form.System)
}
if cast.ToString(form.ListRule) != cast.ToString(collection.ListRule) {
t.Fatalf("Expected ListRule %v, got %v", collection.ListRule, form.ListRule)
}
if cast.ToString(form.ViewRule) != cast.ToString(collection.ViewRule) {
t.Fatalf("Expected ViewRule %v, got %v", collection.ViewRule, form.ViewRule)
}
if cast.ToString(form.CreateRule) != cast.ToString(collection.CreateRule) {
t.Fatalf("Expected CreateRule %v, got %v", collection.CreateRule, form.CreateRule)
}
if cast.ToString(form.UpdateRule) != cast.ToString(collection.UpdateRule) {
t.Fatalf("Expected UpdateRule %v, got %v", collection.UpdateRule, form.UpdateRule)
}
if cast.ToString(form.DeleteRule) != cast.ToString(collection.DeleteRule) {
t.Fatalf("Expected DeleteRule %v, got %v", collection.DeleteRule, form.DeleteRule)
}
rawFormSchema, _ := form.Schema.MarshalJSON()
rawCollectionSchema, _ := collection.Schema.MarshalJSON()
if len(form.Schema.Fields()) != len(collection.Schema.Fields()) {
t.Fatalf("Expected Schema \n%v, \ngot \n%v", string(rawCollectionSchema), string(rawFormSchema))
}
for _, f := range form.Schema.Fields() {
if collection.Schema.GetFieldByName(f.Name) == nil {
t.Fatalf("Missing field %s \nin \n%v", f.Name, string(rawFormSchema))
}
}
// check indexes (if any)
allIndexes, _ := app.Dao().TableIndexes(form.Name)
for _, formIdx := range form.Indexes {
parsed := dbutils.ParseIndex(formIdx)
parsed.TableName = form.Name
normalizedIdx := parsed.Build()
var exists bool
for _, idx := range allIndexes {
if dbutils.ParseIndex(idx).Build() == normalizedIdx {
exists = true
continue
}
}
if !exists {
t.Fatalf("Missing index %s \nin \n%v", normalizedIdx, allIndexes)
}
}
})
} }
} }
func TestCollectionUpsertSubmitInterceptors(t *testing.T) { func TestCollectionUpsertSubmitInterceptors(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -680,6 +721,8 @@ func TestCollectionUpsertSubmitInterceptors(t *testing.T) {
} }
func TestCollectionUpsertWithCustomId(t *testing.T) { func TestCollectionUpsertWithCustomId(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+1 -5
View File
@@ -3,7 +3,6 @@ package forms
import ( import (
"encoding/json" "encoding/json"
"fmt" "fmt"
"log"
validation "github.com/go-ozzo/ozzo-validation/v4" validation "github.com/go-ozzo/ozzo-validation/v4"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
@@ -78,12 +77,9 @@ func (form *CollectionsImport) Submit(interceptors ...InterceptorFunc[[]*models.
} }
// generic/db failure // generic/db failure
if form.app.IsDebug() {
log.Println("Internal import failure:", importErr)
}
return validation.Errors{"collections": validation.NewError( return validation.Errors{"collections": validation.NewError(
"collections_import_failure", "collections_import_failure",
"Failed to import the collections configuration.", "Failed to import the collections configuration. Raw error:\n"+importErr.Error(),
)} )}
}) })
}, interceptors...) }, interceptors...)
+44 -37
View File
@@ -11,6 +11,8 @@ import (
) )
func TestCollectionsImportValidate(t *testing.T) { func TestCollectionsImportValidate(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -38,7 +40,9 @@ func TestCollectionsImportValidate(t *testing.T) {
} }
func TestCollectionsImportSubmit(t *testing.T) { func TestCollectionsImportSubmit(t *testing.T) {
totalCollections := 10 t.Parallel()
totalCollections := 11
scenarios := []struct { scenarios := []struct {
name string name string
@@ -206,7 +210,7 @@ func TestCollectionsImportSubmit(t *testing.T) {
expectError: true, expectError: true,
expectCollectionsCount: totalCollections, expectCollectionsCount: totalCollections,
expectEvents: map[string]int{ expectEvents: map[string]int{
"OnModelBeforeDelete": 4, "OnModelBeforeDelete": 1,
}, },
}, },
{ {
@@ -418,48 +422,51 @@ func TestCollectionsImportSubmit(t *testing.T) {
} }
for _, s := range scenarios { for _, s := range scenarios {
testApp, _ := tests.NewTestApp() t.Run(s.name, func(t *testing.T) {
defer testApp.Cleanup() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup()
form := forms.NewCollectionsImport(testApp) form := forms.NewCollectionsImport(testApp)
// load data // load data
loadErr := json.Unmarshal([]byte(s.jsonData), form) loadErr := json.Unmarshal([]byte(s.jsonData), form)
if loadErr != nil { if loadErr != nil {
t.Errorf("[%s] Failed to load form data: %v", s.name, loadErr) t.Fatalf("Failed to load form data: %v", loadErr)
continue
}
err := form.Submit()
hasErr := err != nil
if hasErr != s.expectError {
t.Errorf("[%s] Expected hasErr to be %v, got %v (%v)", s.name, s.expectError, hasErr, err)
}
// check collections count
collections := []*models.Collection{}
if err := testApp.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err)
}
if len(collections) != s.expectCollectionsCount {
t.Errorf("[%s] Expected %d collections, got %d", s.name, s.expectCollectionsCount, len(collections))
}
// check events
if len(testApp.EventCalls) > len(s.expectEvents) {
t.Errorf("[%s] Expected events %v, got %v", s.name, s.expectEvents, testApp.EventCalls)
}
for event, expectedCalls := range s.expectEvents {
actualCalls := testApp.EventCalls[event]
if actualCalls != expectedCalls {
t.Errorf("[%s] Expected event %s to be called %d, got %d", s.name, event, expectedCalls, actualCalls)
} }
}
err := form.Submit()
hasErr := err != nil
if hasErr != s.expectError {
t.Fatalf("Expected hasErr to be %v, got %v (%v)", s.expectError, hasErr, err)
}
// check collections count
collections := []*models.Collection{}
if err := testApp.Dao().CollectionQuery().All(&collections); err != nil {
t.Fatal(err)
}
if len(collections) != s.expectCollectionsCount {
t.Fatalf("Expected %d collections, got %d", s.expectCollectionsCount, len(collections))
}
// check events
if len(testApp.EventCalls) > len(s.expectEvents) {
t.Fatalf("Expected events %v, got %v", s.expectEvents, testApp.EventCalls)
}
for event, expectedCalls := range s.expectEvents {
actualCalls := testApp.EventCalls[event]
if actualCalls != expectedCalls {
t.Fatalf("Expected event %s to be called %d, got %d", event, expectedCalls, actualCalls)
}
}
})
} }
} }
func TestCollectionsImportSubmitInterceptors(t *testing.T) { func TestCollectionsImportSubmitInterceptors(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
+2
View File
@@ -8,6 +8,8 @@ import (
) )
func TestRealtimeSubscribeValidate(t *testing.T) { func TestRealtimeSubscribeValidate(t *testing.T) {
t.Parallel()
scenarios := []struct { scenarios := []struct {
clientId string clientId string
expectError bool expectError bool
+2
View File
@@ -128,6 +128,8 @@ func (form *RecordEmailChangeConfirm) Submit(interceptors ...InterceptorFunc[*mo
authRecord.SetEmail(newEmail) authRecord.SetEmail(newEmail)
authRecord.SetVerified(true) authRecord.SetVerified(true)
// @todo consider removing if not necessary anymore
authRecord.RefreshTokenKey() // invalidate old tokens authRecord.RefreshTokenKey() // invalidate old tokens
interceptorsErr := runInterceptors(authRecord, func(m *models.Record) error { interceptorsErr := runInterceptors(authRecord, func(m *models.Record) error {
@@ -13,6 +13,8 @@ import (
) )
func TestRecordEmailChangeConfirmValidateAndSubmit(t *testing.T) { func TestRecordEmailChangeConfirmValidateAndSubmit(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -145,6 +147,8 @@ func TestRecordEmailChangeConfirmValidateAndSubmit(t *testing.T) {
} }
func TestRecordEmailChangeConfirmInterceptors(t *testing.T) { func TestRecordEmailChangeConfirmInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
+1 -1
View File
@@ -54,7 +54,7 @@ func (form *RecordEmailChangeRequest) checkUniqueEmail(value any) error {
v, _ := value.(string) v, _ := value.(string)
if !form.dao.IsRecordValueUnique(form.record.Collection().Id, schema.FieldNameEmail, v) { if !form.dao.IsRecordValueUnique(form.record.Collection().Id, schema.FieldNameEmail, v) {
return validation.NewError("validation_record_email_exists", "User email already exists.") return validation.NewError("validation_record_email_invalid", "User email already exists or it is invalid.")
} }
return nil return nil
@@ -12,6 +12,8 @@ import (
) )
func TestRecordEmailChangeRequestValidateAndSubmit(t *testing.T) { func TestRecordEmailChangeRequestValidateAndSubmit(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -106,6 +108,8 @@ func TestRecordEmailChangeRequestValidateAndSubmit(t *testing.T) {
} }
func TestRecordEmailChangeRequestInterceptors(t *testing.T) { func TestRecordEmailChangeRequestInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
+15 -9
View File
@@ -7,7 +7,7 @@ import (
"time" "time"
validation "github.com/go-ozzo/ozzo-validation/v4" validation "github.com/go-ozzo/ozzo-validation/v4"
"github.com/go-ozzo/ozzo-validation/v4/is" "github.com/pocketbase/dbx"
"github.com/pocketbase/pocketbase/core" "github.com/pocketbase/pocketbase/core"
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
@@ -46,7 +46,7 @@ type RecordOAuth2Login struct {
// The authorization code returned from the initial request. // The authorization code returned from the initial request.
Code string `form:"code" json:"code"` Code string `form:"code" json:"code"`
// The code verifier sent with the initial request as part of the code_challenge. // The optional PKCE code verifier as part of the code_challenge sent with the initial request.
CodeVerifier string `form:"codeVerifier" json:"codeVerifier"` CodeVerifier string `form:"codeVerifier" json:"codeVerifier"`
// The redirect url sent with the initial request. // The redirect url sent with the initial request.
@@ -88,8 +88,7 @@ func (form *RecordOAuth2Login) Validate() error {
return validation.ValidateStruct(form, return validation.ValidateStruct(form,
validation.Field(&form.Provider, validation.Required, validation.By(form.checkProviderName)), validation.Field(&form.Provider, validation.Required, validation.By(form.checkProviderName)),
validation.Field(&form.Code, validation.Required), validation.Field(&form.Code, validation.Required),
validation.Field(&form.CodeVerifier, validation.Required), validation.Field(&form.RedirectUrl, validation.Required),
validation.Field(&form.RedirectUrl, validation.Required, is.URL),
) )
} }
@@ -143,11 +142,14 @@ func (form *RecordOAuth2Login) Submit(
provider.SetRedirectUrl(form.RedirectUrl) provider.SetRedirectUrl(form.RedirectUrl)
var opts []oauth2.AuthCodeOption
if provider.PKCE() {
opts = append(opts, oauth2.SetAuthURLParam("code_verifier", form.CodeVerifier))
}
// fetch token // fetch token
token, err := provider.FetchToken( token, err := provider.FetchToken(form.Code, opts...)
form.Code,
oauth2.SetAuthURLParam("code_verifier", form.CodeVerifier),
)
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
@@ -161,7 +163,11 @@ func (form *RecordOAuth2Login) Submit(
var authRecord *models.Record var authRecord *models.Record
// check for existing relation with the auth record // check for existing relation with the auth record
rel, _ := form.dao.FindExternalAuthByProvider(form.Provider, authUser.Id) rel, _ := form.dao.FindFirstExternalAuthByExpr(dbx.HashExp{
"collectionId": form.collection.Id,
"provider": form.Provider,
"providerId": authUser.Id,
})
switch { switch {
case rel != nil: case rel != nil:
authRecord, err = form.dao.FindRecordById(form.collection.Id, rel.RecordId) authRecord, err = form.dao.FindRecordById(form.collection.Id, rel.RecordId)
+10 -2
View File
@@ -10,6 +10,8 @@ import (
) )
func TestUserOauth2LoginValidate(t *testing.T) { func TestUserOauth2LoginValidate(t *testing.T) {
t.Parallel()
app, _ := tests.NewTestApp() app, _ := tests.NewTestApp()
defer app.Cleanup() defer app.Cleanup()
@@ -23,13 +25,13 @@ func TestUserOauth2LoginValidate(t *testing.T) {
"empty payload", "empty payload",
"users", "users",
"{}", "{}",
[]string{"provider", "code", "codeVerifier", "redirectUrl"}, []string{"provider", "code", "redirectUrl"},
}, },
{ {
"empty data", "empty data",
"users", "users",
`{"provider":"","code":"","codeVerifier":"","redirectUrl":""}`, `{"provider":"","code":"","codeVerifier":"","redirectUrl":""}`,
[]string{"provider", "code", "codeVerifier", "redirectUrl"}, []string{"provider", "code", "redirectUrl"},
}, },
{ {
"missing provider", "missing provider",
@@ -49,6 +51,12 @@ func TestUserOauth2LoginValidate(t *testing.T) {
`{"provider":"gitlab","code":"123","codeVerifier":"123","redirectUrl":"https://example.com"}`, `{"provider":"gitlab","code":"123","codeVerifier":"123","redirectUrl":"https://example.com"}`,
[]string{}, []string{},
}, },
{
"[#3689] any redirectUrl value",
"users",
`{"provider":"gitlab","code":"123","codeVerifier":"123","redirectUrl":"something"}`,
[]string{},
},
} }
for _, s := range scenarios { for _, s := range scenarios {
+4
View File
@@ -10,6 +10,8 @@ import (
) )
func TestRecordPasswordLoginValidateAndSubmit(t *testing.T) { func TestRecordPasswordLoginValidateAndSubmit(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -132,6 +134,8 @@ func TestRecordPasswordLoginValidateAndSubmit(t *testing.T) {
} }
func TestRecordPasswordLoginInterceptors(t *testing.T) { func TestRecordPasswordLoginInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
+15 -1
View File
@@ -6,6 +6,8 @@ import (
"github.com/pocketbase/pocketbase/daos" "github.com/pocketbase/pocketbase/daos"
"github.com/pocketbase/pocketbase/forms/validators" "github.com/pocketbase/pocketbase/forms/validators"
"github.com/pocketbase/pocketbase/models" "github.com/pocketbase/pocketbase/models"
"github.com/pocketbase/pocketbase/tools/security"
"github.com/spf13/cast"
) )
// RecordPasswordResetConfirm is an auth record password reset confirmation form. // RecordPasswordResetConfirm is an auth record password reset confirmation form.
@@ -91,9 +93,21 @@ func (form *RecordPasswordResetConfirm) Submit(interceptors ...InterceptorFunc[*
return nil, err return nil, err
} }
if !authRecord.Verified() {
payload, err := security.ParseUnverifiedJWT(form.Token)
if err != nil {
return nil, err
}
// mark as verified if the email hasn't changed
if authRecord.Email() == cast.ToString(payload["email"]) {
authRecord.SetVerified(true)
}
}
interceptorsErr := runInterceptors(authRecord, func(m *models.Record) error { interceptorsErr := runInterceptors(authRecord, func(m *models.Record) error {
authRecord = m authRecord = m
return form.dao.SaveRecord(m) return form.dao.SaveRecord(authRecord)
}, interceptors...) }, interceptors...)
if interceptorsErr != nil { if interceptorsErr != nil {
@@ -13,6 +13,8 @@ import (
) )
func TestRecordPasswordResetConfirmValidateAndSubmit(t *testing.T) { func TestRecordPasswordResetConfirmValidateAndSubmit(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()
@@ -136,6 +138,8 @@ func TestRecordPasswordResetConfirmValidateAndSubmit(t *testing.T) {
} }
func TestRecordPasswordResetConfirmInterceptors(t *testing.T) { func TestRecordPasswordResetConfirmInterceptors(t *testing.T) {
t.Parallel()
testApp, _ := tests.NewTestApp() testApp, _ := tests.NewTestApp()
defer testApp.Cleanup() defer testApp.Cleanup()

Some files were not shown because too many files have changed in this diff Show More