Compare commits
1038 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 10df0b21bf | |||
| 2732dafba9 | |||
| 987974b317 | |||
| 237af31f5f | |||
| f58b8943fb | |||
| 95b8c42244 | |||
| 279ae44e29 | |||
| 4b06e850d5 | |||
| dd02ff5677 | |||
| 4ee43bb9a4 | |||
| 4de574370c | |||
| 02e74185bd | |||
| fe48e2aedc | |||
| db115a3c5d | |||
| 8fe91e198e | |||
| 800d1051a7 | |||
| cb7bb22780 | |||
| 91ff221ccb | |||
| 2bb47e7631 | |||
| ad092b8c15 | |||
| e7621f1ffd | |||
| 5c96c76e89 | |||
| 5fdf5365e7 | |||
| c326bc4899 | |||
| de8922bb8e | |||
| 36aef73969 | |||
| fcf2452b62 | |||
| 87a2311234 | |||
| 178ce70a4a | |||
| f3c12a6122 | |||
| 14285567d5 | |||
| 474706da44 | |||
| 6ab9d120c4 | |||
| 6e86919336 | |||
| c03ac42bff | |||
| e173059b3f | |||
| 5c373e7224 | |||
| 3fc0d52126 | |||
| 64adce7445 | |||
| 854b0f16e3 | |||
| 672e66994a | |||
| ad3d7c0934 | |||
| 80fb2add65 | |||
| 14a18fbc3b | |||
| b4eb27b168 | |||
| 0661c46447 | |||
| 073627884c | |||
| c5b08e778d | |||
| 0232d20e42 | |||
| ebb280f685 | |||
| 3f06c1312d | |||
| e2acbdc7d3 | |||
| 8c18f0da17 | |||
| 82d90b48e2 | |||
| 4d81d275a6 | |||
| 1780adcffa | |||
| ed76bdfd9b | |||
| 2860bc9024 | |||
| e46dadbf45 | |||
| a2de72ab72 | |||
| 9d5b1e25e2 | |||
| 3631cd630c | |||
| 19b7eb2ae8 | |||
| ce72626121 | |||
| f08c0d6273 | |||
| 246ef70245 | |||
| b4b50e8340 | |||
| baf760729e | |||
| 5ae36f6890 | |||
| 4d305d166e | |||
| 940fe7c357 | |||
| 0ea652fa53 | |||
| 9c73ab8e59 | |||
| a6b81f9d42 | |||
| e46303c2f3 | |||
| e72836c027 | |||
| 3f5adadd51 | |||
| eaba8ec6a3 | |||
| 98ec537b68 | |||
| 501319c70b | |||
| ca962f8040 | |||
| 01811438dc | |||
| f983b20a8f | |||
| 182868c0d7 | |||
| 7707db6dc6 | |||
| b294f11720 | |||
| 13d451dc36 | |||
| 6ac9a55c77 | |||
| f188e81e3a | |||
| 8a1dc2aa0e | |||
| b0e1fb5305 | |||
| 69664ec6cd | |||
| 552c4ea819 | |||
| 1f6a3a5803 | |||
| 02c399239c | |||
| d2ec99e7d5 | |||
| 23d952235e | |||
| 8f00f5ddee | |||
| 67a3f96129 | |||
| d352324f58 | |||
| 9c791c4573 | |||
| b7a195d3bc | |||
| 8079db7893 | |||
| 3372a93f71 | |||
| 7cd6252bf7 | |||
| 5d8a6e26bb | |||
| d60c3cd303 | |||
| 971e578dd7 | |||
| 0c9938155b | |||
| f3eedce835 | |||
| 833af6f161 | |||
| 911a995c6a | |||
| 15f4069592 | |||
| f6e955acce | |||
| 9bff4e0d00 | |||
| 256d97136e | |||
| c336bf1eb1 | |||
| c5db5c94e2 | |||
| fe66e51996 | |||
| 2aac612f10 | |||
| defea2eadd | |||
| 24949fd05d | |||
| b1d6188b5e | |||
| 5f3d134d7e | |||
| 7633cdf4e4 | |||
| 4f38075de4 | |||
| 1aa02aea24 | |||
| 958c5950b3 | |||
| 97263cd81a | |||
| 7c7a80ee32 | |||
| c9cee854c3 | |||
| 3bf9589462 | |||
| 0b45fa6130 | |||
| 209f78d740 | |||
| 84663d53e3 | |||
| 0c070bffd8 | |||
| e1941adfc9 | |||
| 6e6f9d83a9 | |||
| 63dcf4a3f6 | |||
| 1f267892f1 | |||
| c960d05acc | |||
| 8653247e44 | |||
| 0cb2f6c362 | |||
| e3e93a0487 | |||
| ef4c513908 | |||
| c6642203c0 | |||
| 31218c8860 | |||
| 2cfacb24f0 | |||
| df19da1379 | |||
| 2a84747cd1 | |||
| 4c6596b561 | |||
| 5380becac2 | |||
| caf90f8383 | |||
| 4a2d0d8c79 | |||
| eeaf61cc41 | |||
| 492b2c643c | |||
| 6834f3ef1b | |||
| 5ff1413adf | |||
| 73c3cb36e8 | |||
| 63316249cf | |||
| 5e088995da | |||
| f148785079 | |||
| 76ce7dc834 | |||
| 90be14b249 | |||
| 50c046136d | |||
| c3d1e0161e | |||
| 8f9405b5ac | |||
| fb5d494e42 | |||
| 0b20753225 | |||
| 813820443e | |||
| 35f3d7e712 | |||
| d9b8ef7725 | |||
| eaa533760d | |||
| 08853e823c | |||
| 4eb1e0515a | |||
| 025c2382d8 | |||
| 06109fd14e | |||
| 5a025576c4 | |||
| 8717538967 | |||
| d72fcd463f | |||
| 36172160ac | |||
| 55b9da50d4 | |||
| cf6c502555 | |||
| 996deb76ff | |||
| 51cf0995f7 | |||
| d5fb0da335 | |||
| 777910b653 | |||
| ee95173de0 | |||
| e1ad8eae8c | |||
| 12bea9a861 | |||
| 8ce459fc23 | |||
| 48a93036ca | |||
| 4b60e0bd48 | |||
| faecaad306 | |||
| 662de20e89 | |||
| 8487dcd13d | |||
| f41ecc9b2c | |||
| 1416d901d0 | |||
| 01739e8c22 | |||
| d09227b61c | |||
| 3da3953a4b | |||
| cb7693c428 | |||
| a8174371f5 | |||
| 1f24b1b95b | |||
| ae8bd0ca47 | |||
| 91855b0763 | |||
| 9ded099a82 | |||
| c3e389d957 | |||
| 9e9b1fb158 | |||
| db446e9ba3 | |||
| 751b16d804 | |||
| 33eded02b1 | |||
| 96ef7d4769 | |||
| 3f8841a589 | |||
| 8656434f51 | |||
| 12bf5d9724 | |||
| 9b944d4a87 | |||
| 10edd6c6a5 | |||
| 0ef7bf23d7 | |||
| de0b674025 | |||
| b148675ba8 | |||
| 917b282614 | |||
| f96fcbe6c9 | |||
| 0b064b9d20 | |||
| 120002b394 | |||
| fdf5053ace | |||
| 0592c59174 | |||
| ac5495134b | |||
| 54585e028e | |||
| 4cf240d478 | |||
| 6c8bb4fec8 | |||
| 08fabe154a | |||
| f0f9165958 | |||
| 971d450128 | |||
| 0cb6017b50 | |||
| 3e68634681 | |||
| 57f204ece9 | |||
| 64d3f27afe | |||
| bc34ed0104 | |||
| c51ad019d2 | |||
| 74eaf6fcef | |||
| d97fda4840 | |||
| 3bfa5ad266 | |||
| 4a46f372b9 | |||
| 8e865180a5 | |||
| 82b1f89332 | |||
| 9a480e3d68 | |||
| 3557590679 | |||
| f4515aaaca | |||
| d27140f809 | |||
| 11c816c145 | |||
| 874dc0cee3 | |||
| ee82d0f5a4 | |||
| 59d851272d | |||
| 1563073bd4 | |||
| 9b32cbc91b | |||
| 6cae9e425f | |||
| 031b51b9b8 | |||
| ed75fc9da1 | |||
| a47db95129 | |||
| ce4e797bd1 | |||
| 87d85624be | |||
| 34c5d9e55e | |||
| bd35cdd268 | |||
| b2d4560039 | |||
| bb71f0aab8 | |||
| e85acd590d | |||
| 04a81ef7a3 | |||
| 46c78e6cce | |||
| f2d8f2e129 | |||
| 217217dd01 | |||
| 8e9e80f175 | |||
| a776622b1c | |||
| 424da0a343 | |||
| 3fc06ae0b8 | |||
| 7847154da4 | |||
| 3c7f0888de | |||
| 3624ceefe6 | |||
| 952463a8c9 | |||
| 008b356dd0 | |||
| 96a7231531 | |||
| f0d3e22f53 | |||
| b40b2c3453 | |||
| c9a650b863 | |||
| cd47cf2adf | |||
| 19b9cd8163 | |||
| 6cd76caebb | |||
| 76e2aaa5f1 | |||
| 603bcaad6a | |||
| fd4ea5c71b | |||
| f28d7c7068 | |||
| d631756345 | |||
| 3bc57827d9 | |||
| da77f906a5 | |||
| cfa7fdc68a | |||
| 805e67c622 | |||
| 564b6872c9 | |||
| 4fb7533ab6 | |||
| 6b92c91c02 | |||
| 213c6fc944 | |||
| 6c17e3e482 | |||
| 8408edbb93 | |||
| 62764ec9ee | |||
| 15b37950e9 | |||
| 75f4d55972 | |||
| f1afd475fc | |||
| 5046d05dea | |||
| aa5f497620 | |||
| fb3d87bd46 | |||
| 300cc2a981 | |||
| 23ee85f6b9 | |||
| 0da97eb092 | |||
| 53736a1906 | |||
| 250b008304 | |||
| e5beb48f69 | |||
| f61691b8de | |||
| db4e1f7b6d | |||
| 57910b7af5 | |||
| 7b5a0bbab9 | |||
| 38e2b77362 | |||
| 5264d6c66c | |||
| 62624bc5ff | |||
| c0fb2c4069 | |||
| 26ce680f03 | |||
| 0e3064b4c1 | |||
| 9d3da998df | |||
| ec651980a2 | |||
| 98567e0261 | |||
| eb87f4fe9b | |||
| fcbe7880a9 | |||
| cee6697c82 | |||
| 5b471a606c | |||
| f8c5270681 | |||
| bf26c05036 | |||
| c2982b8696 | |||
| 2c506da349 | |||
| a399641508 | |||
| 34e367cb15 | |||
| d865c6ac20 | |||
| 16ffa4264e | |||
| 1f8a80c125 | |||
| f78e535cc3 | |||
| 3abb66659c | |||
| d0f8e21b41 | |||
| 652c1d630b | |||
| f65f1840d5 | |||
| ab97c3a7f2 | |||
| 95af9ffe5e | |||
| 4458b18530 | |||
| ca88d7d425 | |||
| 3b0e7b27cc | |||
| 5ab7226f1a | |||
| f737f2903a | |||
| f70fa65969 | |||
| 3201733a50 | |||
| 7cb82bb6ac | |||
| 2c36b7dc95 | |||
| bb2c3ac759 | |||
| 5de7e01a86 | |||
| d3349e649e | |||
| fd23c49f20 | |||
| 0c9e413cd0 | |||
| 2aebf04ff5 | |||
| c9dd5cce11 | |||
| 77f392b29b | |||
| 4c81b1bd1f | |||
| 4345d5128f | |||
| 270e1a56b5 | |||
| 9583d6d96f | |||
| 01e27ac905 | |||
| 709964b207 | |||
| 028f592817 | |||
| b0d6f1ab60 | |||
| b483ed16fc | |||
| 13161994da | |||
| dd7f7dfd97 | |||
| 5a546dc0f1 | |||
| 75b522e36c | |||
| 5d239e904b | |||
| cdd1524def | |||
| be6b49c872 | |||
| 6de11e25bf | |||
| 8060a79b9c | |||
| bc94825877 | |||
| d7d0413bc8 | |||
| d3443b034a | |||
| 7f644a2388 | |||
| 338f4fa9bd | |||
| 44cebeab2a | |||
| 87fca21c7f | |||
| 2e087731b9 | |||
| c9f554465d | |||
| fd75961580 | |||
| b72ab3b318 | |||
| 554bb21d32 | |||
| 8898d3ccc8 | |||
| 2c53d1cd4f | |||
| d8957263ee | |||
| 50ef607293 | |||
| 14b892b6d4 | |||
| 78204c2490 | |||
| 7fb534954e | |||
| 6bc5b19da2 | |||
| 0aa1df3920 | |||
| 1e8f433419 | |||
| 2dfb96ee44 | |||
| 04cde33d01 | |||
| 3a68eeebba | |||
| 5e45b68cc1 | |||
| 10d415ae0c | |||
| 7438551080 | |||
| 1edb342e57 | |||
| 615882adf4 | |||
| 8c77982109 | |||
| 7970c1f518 | |||
| bb8f324e90 | |||
| 80dd117ebf | |||
| 3a4f7452c0 | |||
| 09c485cede | |||
| 83e83fff73 | |||
| e94844a55e | |||
| 054af636e3 | |||
| 50af0ea679 | |||
| 78a245d3a5 | |||
| 8b07f7201f | |||
| 39db5fc564 | |||
| 528f0a624d | |||
| 30f06e0b32 | |||
| 9bc23aa40b | |||
| 9c71e9aabf | |||
| 40d976ed82 | |||
| 39671ad3b9 | |||
| 1c4a261f01 | |||
| df4c91a586 | |||
| a41ded535f | |||
| 2b834d1405 | |||
| 26a6eb81a1 | |||
| acd1572b53 | |||
| 4389e64269 | |||
| 00f5e5bc1a | |||
| 6bcabfef77 | |||
| 1d5eee4281 | |||
| ab19ae5706 | |||
| 38452a51ae | |||
| c5b67fe364 | |||
| 06c728cd21 | |||
| bc1c644eec | |||
| eb22a58065 | |||
| ebf1fdb2a6 | |||
| bea9dae957 | |||
| 64577f2d92 | |||
| fdff970e59 | |||
| c306d80755 | |||
| 50a4f76f9c | |||
| 3f4562df0a | |||
| e58706e97f | |||
| 6b31ddac8e | |||
| 691ed13074 | |||
| 9f7d6c8548 | |||
| cbc04977ed | |||
| 652fad758f | |||
| a66eb4b2a2 | |||
| d6fb78715a | |||
| 40274f00d5 | |||
| 2375758a39 | |||
| 12c70b802a | |||
| 086793c05a | |||
| f7dface365 | |||
| 6e2e4a830f | |||
| 4eb04bd7f6 | |||
| d538c0879d | |||
| ba17067012 | |||
| 45713f2de9 | |||
| 70e5391649 | |||
| a50c3aadf5 | |||
| ca84a11284 | |||
| 82d67c56b3 | |||
| 6e50c61970 | |||
| 6fb1df1aad | |||
| 5d716ad78f | |||
| 9c26d412d7 | |||
| ac6e052788 | |||
| c8987d966b | |||
| 542fb46ee7 | |||
| 491c424580 | |||
| 289dcc375c | |||
| 666f739f4e | |||
| 1e5e35a26a | |||
| e95693d684 | |||
| c286e1b394 | |||
| 40d56f0396 | |||
| 653c137222 | |||
| 3b491ea84c | |||
| ff76855f29 | |||
| 35db9f88de | |||
| 73295cad33 | |||
| b4131e0d19 | |||
| 3afd406c59 | |||
| 341dad643f | |||
| 9ee3249146 | |||
| 1875986ce8 | |||
| 7685b787f2 | |||
| 95f1798908 | |||
| e709f2646c | |||
| d5ed512612 | |||
| 5fd3f0661b | |||
| 256d9cb674 | |||
| 5d5784f069 | |||
| 91549c8f5e | |||
| 4555cdb7df | |||
| 0529ae53ff | |||
| ec77c51cb5 | |||
| 2d01e71afd | |||
| 5bc1fac79c | |||
| c71f33a214 | |||
| c2d56d7475 | |||
| 80b1b01beb | |||
| 07d90758b8 | |||
| 04d51bab22 | |||
| 8fa3e9bc4a | |||
| 2132e7e8d1 | |||
| 9a342af609 | |||
| 1741ab1011 | |||
| 4acad71c8f | |||
| 8c62a6b714 | |||
| 79fc11b452 | |||
| e7d1289f1b | |||
| 61e402f785 | |||
| 5f548e29fa | |||
| aea9e227db | |||
| b71b18d184 | |||
| c374ba3ae9 | |||
| 32e73727ad | |||
| e4157e3373 | |||
| 708962b6a7 | |||
| 0b9bba7158 | |||
| 39150cb9ab | |||
| ea2e927054 | |||
| e916001e08 | |||
| b22f9a6af6 | |||
| 5fb11b978c | |||
| 5864a4ed06 | |||
| a171186f93 | |||
| 50271ff21f | |||
| 382e082198 | |||
| e5d4f8b58f | |||
| 201d7c8fdc | |||
| 84e8ea2fb2 | |||
| 14330d1a83 | |||
| 6376b6d73b | |||
| 2e241dd5f3 | |||
| 5914d9f511 | |||
| 32db11f47f | |||
| 5144f03dc8 | |||
| b701787ef4 | |||
| 0bd8b71438 | |||
| 034876afbb | |||
| 06091546b0 | |||
| a92bf8ec8f | |||
| 0f91b8bb7c | |||
| a3fd74f117 | |||
| 0b291995b9 | |||
| a2a4c6fc9a | |||
| ae075908d1 | |||
| d48e36813e | |||
| fcefd0978f | |||
| 5a77e36ac6 | |||
| 93b3959d2f | |||
| b03ea41638 | |||
| 01a9511dcf | |||
| 069433ccbc | |||
| dce75473a5 | |||
| 89bf07d2ef | |||
| caf860dd10 | |||
| af44af98b8 | |||
| 58af9e5a29 | |||
| 95fc54a16d | |||
| 8ba5d0c15b | |||
| 262551cd83 | |||
| 6d533aae33 | |||
| 8fc5a398e3 | |||
| b4c6f75786 | |||
| 5d3fe88be9 | |||
| 289eb422f3 | |||
| 7acefb5f25 | |||
| bab5967d88 | |||
| 189613b482 | |||
| 8faf164e9e | |||
| fc9caf4f9b | |||
| a7eff2c033 | |||
| 1fe6f86cc0 | |||
| 54400f8f1c | |||
| 222f3cf4e7 | |||
| fce2b4b01c | |||
| af696daf55 | |||
| 1ce9e97ad0 | |||
| 52ebbad335 | |||
| 4cbda5c228 | |||
| b181799553 | |||
| b7adcb26c7 | |||
| 2f95febf71 | |||
| a9829598c5 | |||
| fc453be71f | |||
| cfc52fd355 | |||
| 4f9f458151 | |||
| c3bb93c197 | |||
| 73119e0cef | |||
| 019807d0f5 | |||
| 8098856186 | |||
| 209e82f7e5 | |||
| ba387b8e8b | |||
| 906c1ca246 | |||
| 322aef98ae | |||
| dd3002d2c8 | |||
| d8d0aedd9e | |||
| 107ac9f4e9 | |||
| 9b4d8e2bfe | |||
| f714a9ed31 | |||
| 13d2230089 | |||
| d846704aa6 | |||
| 58733d6450 | |||
| 0edb94158d | |||
| e27bacee3b | |||
| 883f123138 | |||
| 351d6f3f4f | |||
| 2c652dcc3e | |||
| af3eed4040 | |||
| c2c7a73dee | |||
| f3627581ab | |||
| 0cda2dafdc | |||
| a7a807bb11 | |||
| 6cdca24333 | |||
| b2196a47be | |||
| 6437364800 | |||
| 81ac695c43 | |||
| 78c9df4bb0 | |||
| c4b4643f6e | |||
| 71350f2ad6 | |||
| 35abdb0e69 | |||
| 9627667d82 | |||
| ca9cf868cf | |||
| 48e3552e1a | |||
| 002cdf312b | |||
| 98ae72aef2 | |||
| bf43fc8ed2 | |||
| 223539bebb | |||
| 89378b2270 | |||
| 68c100fc7e | |||
| 8acf76fe93 | |||
| e05718e702 | |||
| dcf57b0898 | |||
| feea051c4f | |||
| 8aa8a93a6a | |||
| 17e77fa65a | |||
| 743627ecc2 | |||
| d7448c5c69 | |||
| b6a0648366 | |||
| 077672d8c1 | |||
| 211d479eb9 | |||
| 3e36f4c888 | |||
| 6175489ba0 | |||
| 94ce9b12f6 | |||
| 087f3c10e9 | |||
| 5acc55ce92 | |||
| 73727dc56a | |||
| 43bcfe5653 | |||
| 555b0a6627 | |||
| ba9f32a039 | |||
| a4e4dbc43e | |||
| c5b330cb4e | |||
| f3c85f73c3 | |||
| ec51cb0015 | |||
| e68bf13949 | |||
| 6497d35d84 | |||
| 195febfb09 | |||
| c26ecfce08 | |||
| cc6acc1f03 | |||
| 06eb6e636e | |||
| 3ac5a97467 | |||
| 19b4ea4f6b | |||
| 16b2ad3b24 | |||
| f6c6712550 | |||
| 640662e9b2 | |||
| be0a0922dc | |||
| 9e9a4f9b6a | |||
| 9a95d5cfd5 | |||
| 86f9f5f9b1 | |||
| 5a11f14681 | |||
| a1dcc1d2b5 | |||
| a580527de4 | |||
| 53f7ab37d5 | |||
| b7df2ca74b | |||
| 754ea99092 | |||
| 13e8d5098c | |||
| 8db0093562 | |||
| bcce0c0885 | |||
| e96ddc7865 | |||
| cd19ac661e | |||
| afe42270bb | |||
| 2be09efd99 | |||
| 660b18ee97 | |||
| 1cb39ca027 | |||
| 21ac4f3308 | |||
| f589fc7fa5 | |||
| 8292704703 | |||
| 8443b719a6 | |||
| 14c9432640 | |||
| bd283f19fe | |||
| 25c31a63ff | |||
| ade7c9a865 | |||
| 344d20e4e6 | |||
| ceb12c8e51 | |||
| 355a685ad8 | |||
| 335678dd41 | |||
| 4ec4b13f3f | |||
| 3a63c4dfb2 | |||
| ba85ca6005 | |||
| 7e3b82d55e | |||
| 38ec8ec274 | |||
| ebc1d1cf7f | |||
| c2d3a1b187 | |||
| 00f920aedd | |||
| 4ea6b47b91 | |||
| 1db1803e97 | |||
| b5c0d3ca88 | |||
| 2f735f562c | |||
| a0c7ff6e19 | |||
| 2d754d52ef | |||
| f03227c248 | |||
| 0f21479abc | |||
| 4109bf65d2 | |||
| f2d2aaa354 | |||
| f7c550eda7 | |||
| 3f8a95d808 | |||
| c31b72810e | |||
| 962ee09f30 | |||
| 9b8eb911c3 | |||
| d5af0903e2 | |||
| 6c07d70d93 | |||
| db40e85856 | |||
| 97616beb11 | |||
| b7bf4f352d | |||
| da47283e0a | |||
| 42dc335747 | |||
| 4de739c566 | |||
| a3218583c6 | |||
| 70e188125b | |||
| 52b3f15329 | |||
| 7e2d392791 | |||
| ff1310c23e | |||
| 58c9b5d7f0 | |||
| d36ba88de7 | |||
| 7db26f4a31 | |||
| a0f5ba1799 | |||
| 9b45ba62ba | |||
| 3aa5364e0c | |||
| 4857203d08 | |||
| d03991f7e1 | |||
| 44ea43f9cf | |||
| 33a02633b7 | |||
| 8d65f09400 | |||
| aa7d5d7c7c | |||
| 0b28e33fae | |||
| 6a14a7956f | |||
| bd9da65511 | |||
| c69dc1b8c1 | |||
| 8996715c0c | |||
| ada6c368e4 | |||
| 0854e901ac | |||
| 31b4d9387b | |||
| 5e6f44e6bf | |||
| 0955765525 | |||
| b710a90796 | |||
| 0942609d03 | |||
| 7f20dc7782 | |||
| 39c0464fe8 | |||
| 1877fc0000 | |||
| 2dd1ca0197 | |||
| b95f552ed2 | |||
| c73c9484b8 | |||
| 56f346ac20 | |||
| 9a9efae19e | |||
| f4a607df05 | |||
| db5f3d51d9 | |||
| 232691a396 | |||
| 40f0beb1bf | |||
| bdb67de5cb | |||
| de8a1e18e8 | |||
| 5f0c9b8205 | |||
| e929930bae | |||
| a18245cd1d | |||
| d256e66574 | |||
| a264ba474a | |||
| f53574dc33 | |||
| ca717aa01f | |||
| 699a93e193 | |||
| f9b9fb081a | |||
| 63905842fd | |||
| 3a1ecaf63a | |||
| e8c2854325 | |||
| 93b4dcf641 | |||
| 5a9d2e7ad6 | |||
| cfb1e2d62d | |||
| a27a0d6f01 | |||
| 1883606a5e | |||
| 79e76522e3 | |||
| 3f67b80a42 | |||
| dc1a12772a | |||
| fe8e1c41cc | |||
| 46b2f17acc | |||
| 72f96e3c02 | |||
| 064ba72d36 | |||
| fdc48c7fe9 | |||
| 97f85ce74e | |||
| 231f983429 | |||
| 9d11d741f4 | |||
| 19c1ccb11c | |||
| 6ebf66275c | |||
| aa6e2fe001 | |||
| 8a55860e86 | |||
| 4791e2fa15 | |||
| a7d428d009 | |||
| 0599cf86b3 | |||
| 68e8af9554 | |||
| ddd07c87ac | |||
| 19d8e04566 | |||
| b30bbaea99 | |||
| 758581dafb | |||
| 5b4c8cdb60 | |||
| 7efda5cf93 | |||
| 650a9925c4 | |||
| 81c905a8bb | |||
| 5b6c6e47eb | |||
| 27377d4928 | |||
| dc1ee0e2e5 | |||
| c485d4410d | |||
| 2317595829 | |||
| 99b81e9868 | |||
| d0cd4c89f0 | |||
| fa78a71492 | |||
| f5e2a7845b | |||
| ac10adac84 | |||
| 7cebb82412 | |||
| b1ce28f804 | |||
| 2b30b40a7e | |||
| 9abc46cc04 | |||
| 55d0f917ea | |||
| cd1dc74295 | |||
| 6f197e69fc | |||
| 02ac55ab1a | |||
| ba95ddb930 | |||
| f220021927 | |||
| 5a783d4e14 | |||
| b5b285ebe6 | |||
| 43f3f6bdbe | |||
| 6a967051ba | |||
| 52e3852129 | |||
| a644efe016 | |||
| e0de0fd6fb | |||
| 99df0816b0 | |||
| d318110bbe | |||
| dd207eef9b | |||
| 1203a4b75f | |||
| 48092c98e7 | |||
| 89e0c2fb50 | |||
| 2c12af95fc | |||
| 7bc8bb428e | |||
| fb43e24a3e | |||
| 10bb533f64 | |||
| 9853a27d82 | |||
| 45949f34c6 | |||
| ab532b55dc | |||
| ea7fffa3a1 | |||
| 4277e5ed5d | |||
| 59f1c46b94 | |||
| d439d9afdd | |||
| 5351d82275 | |||
| 7f521eb755 | |||
| 179f6125e8 | |||
| 9670e84862 | |||
| 70dd35dec4 | |||
| 0c3b8d0556 | |||
| fe36974b5b | |||
| 2e75b3f726 | |||
| 3392d941ee | |||
| a00e67b01f | |||
| 56f7f2351f | |||
| 4cf93dc16e | |||
| 17b0f3762c | |||
| e0352b1223 | |||
| 299ae98ebd | |||
| 891ce4bf70 | |||
| 82aee1cd0f | |||
| 681bc77803 | |||
| 158b113ac9 | |||
| e7a8582799 | |||
| 2c02d33a28 | |||
| 0f5680b86c | |||
| a4af6323ab | |||
| aab1c5fda9 | |||
| 575b4e088e | |||
| ed404b4661 | |||
| 7acac4225b | |||
| b9015764c5 | |||
| 7991ddb5bd | |||
| b958537adf | |||
| e056c54bdb | |||
| 58d0fccaa0 | |||
| 8ccb5e73e0 | |||
| fef66b16a9 | |||
| e7872607fd | |||
| 26b4d04d76 | |||
| aad62958fb | |||
| 7dae15fd29 | |||
| 174e4890dd | |||
| f4b793e9b8 | |||
| 330c3aedb6 | |||
| 507c333c15 | |||
| cb3bf809cc | |||
| f6fcaaeccc | |||
| c8bc4afc93 | |||
| 62d7ad7e9e | |||
| 0f408f67b4 | |||
| 3fd2bb3c5f | |||
| 3601e2ca5c | |||
| 5722252dcb | |||
| a8d44b189c | |||
| d7b7f1478f | |||
| 71f437e43a | |||
| af4cd966e9 | |||
| d8b7fe4148 | |||
| 98dbe76634 | |||
| 7586381026 | |||
| b6cc867c35 | |||
| e19b1af705 | |||
| 4dbdc2e03a | |||
| 3be1fa5040 | |||
| 14adbeeabb | |||
| 1d4ccb1f2f | |||
| ca07d6a876 | |||
| c1366fa970 | |||
| d6db9ad5ee | |||
| 1f0331117a | |||
| b8bd049dff | |||
| 6ff1da2496 | |||
| 00ca452f79 | |||
| 3d85c64f76 | |||
| 6fe3ecfede | |||
| 6f2eb78937 | |||
| 3080e37ddf | |||
| a2592629db | |||
| e4fbd45636 | |||
| 24e0053def | |||
| 9f1a86d554 | |||
| 08099f94e4 | |||
| fea8f5c0de | |||
| 8fee5c8d0b | |||
| 8e0bd50e88 | |||
| a79f7ba50f | |||
| d43ff23079 | |||
| 144a90fa73 | |||
| a03b3af933 | |||
| 55543be3b0 | |||
| eea99eb01d | |||
| 8e955f3c81 | |||
| 34b9d529e7 | |||
| c724b8d83a | |||
| 8d7c8e7c4e | |||
| b435957a9b | |||
| 26d5f99425 | |||
| d7d1175ba5 | |||
| 054b3491e7 | |||
| e98c2ea63e | |||
| bb21486536 | |||
| 7bb7e00927 | |||
| b87dd6083c | |||
| ebea83a750 | |||
| c0a3753e55 | |||
| 3361042186 | |||
| 716efcf73d | |||
| ffc795e5b1 | |||
| c938a7c219 | |||
| 744a433cc9 | |||
| e8e91fd29a | |||
| fe9f84ac05 | |||
| 7cb5407bcf | |||
| f7cf7e2a4c | |||
| 3d38813da5 | |||
| fc72532aea | |||
| ea0f7fb54f | |||
| 0fc8148c47 | |||
| ed8579757c | |||
| c8d456e6eb | |||
| 25532aad87 | |||
| 1e1ba1d8d7 | |||
| 698c6a222e | |||
| c7c9040321 | |||
| 1b3f35d77d | |||
| a85992ee79 | |||
| b7c78177e9 | |||
| a953b726fc | |||
| e60914127c | |||
| 9748103782 | |||
| 2842b4ac46 | |||
| bcfb91b05f | |||
| 49eb2dd8e5 | |||
| afae9e130a | |||
| ecca0ecdfa | |||
| b0cd612974 | |||
| 09e174f086 | |||
| 68ea6aabc7 | |||
| 61d7639129 | |||
| 85da309740 | |||
| d18c10996a | |||
| 32ea0c8949 | |||
| 5dcb9814e4 | |||
| 366026e909 | |||
| c626d06d8a | |||
| 1775c4e0dd | |||
| bad4743fc2 | |||
| 8abddaae08 | |||
| 2ba4c0d558 | |||
| b5cb7d7ad7 | |||
| 47783b32b1 | |||
| eb511c0a27 | |||
| a9c0b848ee | |||
| 935577edec | |||
| 3d71d2b362 | |||
| a1b1cd06c2 | |||
| 1095b41579 | |||
| bf337a1560 | |||
| 5ff7b72a80 | |||
| 52f9283155 | |||
| 404b3a7ba5 | |||
| e54333d344 | |||
| af34ce88df | |||
| af3ff0b581 | |||
| 82c524a788 |
@@ -5,6 +5,7 @@ REDIS_URL=redis://redis:6379/0
|
||||
EXTERNAL_HOSTNAME=docuelevate.example.com
|
||||
GOTENBERG_URL=http://gotenberg:3000
|
||||
ALLOW_FILE_DELETE=true # Allow deletion of file records
|
||||
COMPLIANCE_ENABLED=true # Enable compliance templates dashboard (GDPR, HIPAA, SOC 2)
|
||||
|
||||
# **UI / Appearance**
|
||||
# Default colour scheme: system (follow OS), light, or dark
|
||||
@@ -16,6 +17,12 @@ ALLOW_FILE_DELETE=true # Allow deletion of file records
|
||||
PROCESSALL_THROTTLE_THRESHOLD=20 # Number of files above which throttling is applied (default: 20)
|
||||
PROCESSALL_THROTTLE_DELAY=3 # Delay in seconds between each task submission when throttling (default: 3)
|
||||
|
||||
# **Task Retry Settings**
|
||||
# Failed tasks are automatically retried with exponential backoff and jitter.
|
||||
# TASK_RETRY_MAX_RETRIES=3 # Max retry attempts per task (default: 3)
|
||||
# TASK_RETRY_DELAYS=60,300,900 # Countdown (seconds) before each retry; 1 min, 5 min, 15 min
|
||||
# TASK_RETRY_JITTER=true # Add ±20% random jitter to prevent thundering-herd (default: true)
|
||||
|
||||
# **Client-Side Upload Throttling**
|
||||
# Controls pacing when the browser uploads files (especially large directory drops).
|
||||
# The browser auto-detects rate-limit (HTTP 429) responses and backs off accordingly.
|
||||
@@ -90,6 +97,22 @@ MAX_UPLOAD_SIZE=1073741824
|
||||
# Allowed request headers (use * to allow all)
|
||||
# CORS_ALLOWED_HEADERS=*
|
||||
|
||||
# **Audit Logging & SIEM Integration** (see docs/ConfigurationGuide.md#audit-logging)
|
||||
# Enable HTTP request audit logging middleware
|
||||
AUDIT_LOGGING_ENABLED=true
|
||||
# Include client IP in audit log entries (disable for GDPR-sensitive deployments)
|
||||
AUDIT_LOG_INCLUDE_CLIENT_IP=true
|
||||
|
||||
# Forward audit events to an external SIEM system (Syslog, Splunk, Logstash, Grafana, etc.)
|
||||
# AUDIT_SIEM_ENABLED=false
|
||||
# AUDIT_SIEM_TRANSPORT=syslog # syslog | http
|
||||
# AUDIT_SIEM_SYSLOG_HOST=localhost
|
||||
# AUDIT_SIEM_SYSLOG_PORT=514
|
||||
# AUDIT_SIEM_SYSLOG_PROTOCOL=udp # udp | tcp
|
||||
# AUDIT_SIEM_HTTP_URL= # e.g. https://splunk:8088/services/collector/event
|
||||
# AUDIT_SIEM_HTTP_TOKEN= # Bearer / HEC token
|
||||
# AUDIT_SIEM_HTTP_CUSTOM_HEADERS= # Comma-separated Key:Value pairs
|
||||
|
||||
# **Rate Limiting** (see SECURITY_AUDIT.md and docs/API.md)
|
||||
# Protects against DoS attacks and API abuse by limiting request rates per IP/user
|
||||
# Enabled by default - highly recommended for production
|
||||
@@ -123,12 +146,62 @@ ADMIN_USERNAME=admin
|
||||
ADMIN_PASSWORD=your_secure_password
|
||||
ADMIN_GROUP_NAME=admin
|
||||
|
||||
# **Multi-User Mode**
|
||||
# When enabled, each user has their own document space with isolated uploads,
|
||||
# search, and file management. Requires AUTH_ENABLED=true.
|
||||
MULTI_USER_ENABLED=false
|
||||
# Allow users to self-register with an email address and password.
|
||||
# Set to true to enable the /signup page. Requires MULTI_USER_ENABLED=true.
|
||||
# When SMTP is configured, a verification email is sent before the account is activated.
|
||||
# Without SMTP, accounts are activated immediately upon registration.
|
||||
# ALLOW_LOCAL_SIGNUP=false
|
||||
# Default upload limit per user per day (0 = unlimited)
|
||||
DEFAULT_DAILY_UPLOAD_LIMIT=0
|
||||
# Show unowned documents (owner_id=NULL) to all users (true) or only admins (false)
|
||||
UNOWNED_DOCS_VISIBLE_TO_ALL=true
|
||||
# Auto-assign this owner ID to documents ingested without a session (e.g. IMAP, API)
|
||||
# Leave empty/unset to keep them unowned until claimed.
|
||||
# DEFAULT_OWNER_ID=
|
||||
|
||||
# **Subscription / Quota Settings**
|
||||
# Soft-limit overage buffer in percent (0–200). Announced quota is multiplied by (1 + percent/100)
|
||||
# for actual enforcement. E.g. 20 means a 150-doc/month plan enforces at 180. 0 = enforce exactly.
|
||||
# Per-plan overage_percent set in the Plan Designer overrides this global default.
|
||||
# SUBSCRIPTION_OVERAGE_PERCENT=20
|
||||
|
||||
# **OpenID Connect/Authentik Settings**
|
||||
AUTHENTIK_CLIENT_ID=<yourAuthentikAppClientID>
|
||||
AUTHENTIK_CLIENT_SECRET=<yourAuthentikClientSecret>
|
||||
AUTHENTIK_CONFIG_URL=<ConfigUrlOfYourApp, e.g. https://authentik.example.com/application/o/docuelevate/.well-known/openid-configuration>
|
||||
OAUTH_PROVIDER_NAME="Authentik SSO"
|
||||
|
||||
# **Social Login Providers**
|
||||
# Enable one or more social login providers to let users sign in with existing accounts.
|
||||
# Each provider requires separate OAuth credentials. See docs/SocialLoginSetup.md for details.
|
||||
|
||||
# Google Sign-In (https://console.cloud.google.com/apis/credentials)
|
||||
# SOCIAL_AUTH_GOOGLE_ENABLED=false
|
||||
# SOCIAL_AUTH_GOOGLE_CLIENT_ID=your-google-client-id.apps.googleusercontent.com
|
||||
# SOCIAL_AUTH_GOOGLE_CLIENT_SECRET=your-google-client-secret
|
||||
|
||||
# Microsoft Sign-In / Azure AD (https://portal.azure.com/#blade/Microsoft_AAD_RegisteredApps)
|
||||
# SOCIAL_AUTH_MICROSOFT_ENABLED=false
|
||||
# SOCIAL_AUTH_MICROSOFT_CLIENT_ID=your-microsoft-application-id
|
||||
# SOCIAL_AUTH_MICROSOFT_CLIENT_SECRET=your-microsoft-client-secret
|
||||
# SOCIAL_AUTH_MICROSOFT_TENANT=common # common | organizations | consumers | <tenant-id>
|
||||
|
||||
# Apple Sign-In (https://developer.apple.com/account/resources)
|
||||
# SOCIAL_AUTH_APPLE_ENABLED=false
|
||||
# SOCIAL_AUTH_APPLE_CLIENT_ID=com.example.docuelevate
|
||||
# SOCIAL_AUTH_APPLE_TEAM_ID=ABCDE12345
|
||||
# SOCIAL_AUTH_APPLE_KEY_ID=FGHIJ67890
|
||||
# SOCIAL_AUTH_APPLE_PRIVATE_KEY="-----BEGIN PRIVATE KEY-----\n...\n-----END PRIVATE KEY-----"
|
||||
|
||||
# Dropbox Sign-In (https://www.dropbox.com/developers/apps)
|
||||
# SOCIAL_AUTH_DROPBOX_ENABLED=false
|
||||
# SOCIAL_AUTH_DROPBOX_CLIENT_ID=your-dropbox-app-key
|
||||
# SOCIAL_AUTH_DROPBOX_CLIENT_SECRET=your-dropbox-app-secret
|
||||
|
||||
# **AI/ML Services**
|
||||
# Select your AI provider: openai | azure | anthropic | gemini | ollama | openrouter | portkey | litellm
|
||||
AI_PROVIDER=openai
|
||||
@@ -168,16 +241,82 @@ OPENAI_MODEL=gpt-4o-mini
|
||||
# AI_MODEL=gpt-4o # deployment name in Azure
|
||||
|
||||
# Azure Document Intelligence (OCR – separate from AI provider above)
|
||||
# **Email Settings**
|
||||
# **Email Settings (shared SMTP – password reset, verification, and system notifications)**
|
||||
EMAIL_HOST=smtp.example.com
|
||||
EMAIL_PORT=587
|
||||
EMAIL_USERNAME=docuelevate@example.com
|
||||
EMAIL_PASSWORD=your_secure_email_password
|
||||
EMAIL_USE_TLS=True
|
||||
EMAIL_SENDER=DocuElevate System <docuelevate@example.com>
|
||||
EMAIL_DEFAULT_RECIPIENT=recipient@example.com
|
||||
# EMAIL_DEFAULT_RECIPIENT is not used for document delivery (see DEST_EMAIL_* below)
|
||||
|
||||
# **Email Destination Settings (dedicated SMTP for document delivery)**
|
||||
# These settings are intentionally separate from the shared EMAIL_* settings above.
|
||||
# Configuring EMAIL_HOST for password reset / notifications does NOT automatically
|
||||
# enable the email destination – you must set DEST_EMAIL_HOST to activate it.
|
||||
# DEST_EMAIL_ENABLED=true # Set to false to disable email delivery without removing credentials
|
||||
DEST_EMAIL_HOST=smtp.example.com
|
||||
DEST_EMAIL_PORT=587
|
||||
DEST_EMAIL_USERNAME=docuelevate@example.com
|
||||
DEST_EMAIL_PASSWORD=your_secure_email_password
|
||||
DEST_EMAIL_USE_TLS=True
|
||||
DEST_EMAIL_SENDER=DocuElevate Delivery <docuelevate@example.com>
|
||||
DEST_EMAIL_DEFAULT_RECIPIENT=recipient@example.com
|
||||
|
||||
# **Watch Folder Ingestion**
|
||||
# DocuElevate can automatically monitor directories (local, FTP, SFTP, and cloud providers) for new files.
|
||||
#
|
||||
# Local watch folders — works with any mounted path (SMB/CIFS, NFS, local disk, etc.)
|
||||
# Set WATCH_FOLDERS to a comma-separated list of absolute paths inside the container.
|
||||
WATCH_FOLDERS=
|
||||
WATCH_FOLDER_POLL_INTERVAL=1
|
||||
WATCH_FOLDER_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# FTP ingest — poll an FTP directory for new files (uses FTP connection settings above)
|
||||
FTP_INGEST_ENABLED=false
|
||||
FTP_INGEST_FOLDER=
|
||||
FTP_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# SFTP ingest — poll an SFTP directory for new files (uses SFTP connection settings above)
|
||||
SFTP_INGEST_ENABLED=false
|
||||
SFTP_INGEST_FOLDER=
|
||||
SFTP_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# Dropbox ingest — poll a Dropbox folder (uses Dropbox OAuth credentials above)
|
||||
DROPBOX_INGEST_ENABLED=false
|
||||
DROPBOX_INGEST_FOLDER=
|
||||
DROPBOX_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# Google Drive ingest — poll a Google Drive folder (uses Google Drive credentials above)
|
||||
GOOGLE_DRIVE_INGEST_ENABLED=false
|
||||
GOOGLE_DRIVE_INGEST_FOLDER_ID=
|
||||
GOOGLE_DRIVE_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# OneDrive ingest — poll a OneDrive folder (uses OneDrive MSAL credentials above)
|
||||
ONEDRIVE_INGEST_ENABLED=false
|
||||
ONEDRIVE_INGEST_FOLDER_PATH=
|
||||
ONEDRIVE_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# Nextcloud ingest — poll a Nextcloud folder (uses Nextcloud WebDAV credentials above)
|
||||
NEXTCLOUD_INGEST_ENABLED=false
|
||||
NEXTCLOUD_INGEST_FOLDER=
|
||||
NEXTCLOUD_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# Amazon S3 ingest — poll an S3 prefix (uses S3/AWS credentials above)
|
||||
S3_INGEST_ENABLED=false
|
||||
S3_INGEST_PREFIX=
|
||||
S3_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# WebDAV ingest — poll a WebDAV folder (uses WebDAV credentials above)
|
||||
WEBDAV_INGEST_ENABLED=false
|
||||
WEBDAV_INGEST_FOLDER=
|
||||
WEBDAV_INGEST_DELETE_AFTER_PROCESS=false
|
||||
|
||||
# **IMAP Settings**
|
||||
# DocuElevate polls these mailboxes for new email attachments and automatically ingests them.
|
||||
# No manual forwarding required — DocuElevate acts as an IMAP *client*.
|
||||
# For HP Scanners / Scan-to-Email: configure the scanner to send to a dedicated mailbox,
|
||||
# then point DocuElevate at that mailbox using the settings below.
|
||||
IMAP1_HOST=mail.example.com
|
||||
IMAP1_PORT=993
|
||||
IMAP1_USERNAME=<IMAP1_USERNAME>
|
||||
@@ -200,8 +339,15 @@ IMAP2_DELETE_AFTER_PROCESS=false
|
||||
# Use for pre-production instances that share a mailbox with production.
|
||||
IMAP_READONLY_MODE=false
|
||||
|
||||
# Controls which attachment types are ingested from IMAP emails.
|
||||
# 'documents_only' (default) – PDFs and office files only; images are skipped.
|
||||
# 'all' – all supported file types including images.
|
||||
# Per-user IMAP accounts can override this global default.
|
||||
IMAP_ATTACHMENT_FILTER=documents_only
|
||||
|
||||
# **Storage/Document Services**
|
||||
# Amazon S3
|
||||
# S3_ENABLED=true # Set to false to disable S3 uploads without removing credentials
|
||||
AWS_REGION=us-east-1
|
||||
AWS_ACCESS_KEY_ID=AKIAIOSFODNN7EXAMPLE
|
||||
AWS_SECRET_ACCESS_KEY=wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY
|
||||
@@ -211,12 +357,14 @@ S3_STORAGE_CLASS=STANDARD
|
||||
S3_ACL=private
|
||||
|
||||
# NextCloud
|
||||
# NEXTCLOUD_ENABLED=true # Set to false to disable NextCloud uploads without removing credentials
|
||||
NEXTCLOUD_UPLOAD_URL=https://nextcloud.example.com/remote.php/dav/files/<USERNAME>
|
||||
NEXTCLOUD_FOLDER="<NEXTCLOUD_FOLDER_PATH>"
|
||||
NEXTCLOUD_USERNAME=<NEXTCLOUD_USERNAME>
|
||||
NEXTCLOUD_PASSWORD=<NEXTCLOUD_PASSWORD>
|
||||
|
||||
# Paperless-ngx
|
||||
# PAPERLESS_ENABLED=true # Set to false to disable Paperless uploads without removing credentials
|
||||
PAPERLESS_HOST=https://paperless.example.com
|
||||
PAPERLESS_NGX_API_TOKEN=<PAPERLESS_API_TOKEN>
|
||||
# Optional: Name of the custom field in Paperless-ngx to store the "absender" (sender) value
|
||||
@@ -233,12 +381,14 @@ PAPERLESS_NGX_API_TOKEN=<PAPERLESS_API_TOKEN>
|
||||
# PAPERLESS_CUSTOM_FIELDS_MAPPING=
|
||||
|
||||
# Dropbox
|
||||
# DROPBOX_ENABLED=true # Set to false to disable Dropbox uploads without removing credentials
|
||||
DROPBOX_APP_KEY=<DROPBOX_APP_KEY>
|
||||
DROPBOX_APP_SECRET=<DROPBOX_APP_SECRET>
|
||||
DROPBOX_REFRESH_TOKEN=<DROPBOX_REFRESH_TOKEN>
|
||||
DROPBOX_FOLDER="/Documents/Uploads"
|
||||
|
||||
# Google Drive
|
||||
# GOOGLE_DRIVE_ENABLED=true # Set to false to disable Google Drive uploads without removing credentials
|
||||
# Service Account Method:
|
||||
GOOGLE_DRIVE_CREDENTIALS_JSON={"type":"service_account","project_id":"your-project","private_key_id":"key-id","private_key":"-----BEGIN PRIVATE KEY-----\nYOUR_PRIVATE_KEY\n-----END PRIVATE KEY-----\n","client_email":"service-account@project.iam.gserviceaccount.com","client_id":"client-id","auth_uri":"https://accounts.google.com/o/oauth2/auth","token_uri":"https://oauth2.googleapis.com/token","auth_provider_x509_cert_url":"https://www.googleapis.com/oauth2/v1/certs","client_x509_cert_url":"https://www.googleapis.com/robot/v1/metadata/x509/service-account%40project.iam.gserviceaccount.com"}
|
||||
GOOGLE_DRIVE_FOLDER_ID=<YOUR_FOLDER_ID>
|
||||
@@ -251,6 +401,7 @@ GOOGLE_DRIVE_CLIENT_SECRET=your-oauth-client-secret # Required for OAuth method
|
||||
GOOGLE_DRIVE_REFRESH_TOKEN=your-oauth-refresh-token # Required for OAuth method
|
||||
|
||||
# OneDrive
|
||||
# ONEDRIVE_ENABLED=true # Set to false to disable OneDrive uploads without removing credentials
|
||||
ONEDRIVE_CLIENT_ID=your-client-id
|
||||
ONEDRIVE_CLIENT_SECRET=your-client-secret
|
||||
ONEDRIVE_TENANT_ID=common
|
||||
@@ -258,6 +409,7 @@ ONEDRIVE_REFRESH_TOKEN=your-refresh-token
|
||||
ONEDRIVE_FOLDER_PATH=Documents/Uploads
|
||||
|
||||
# WebDAV
|
||||
# WEBDAV_ENABLED=true # Set to false to disable WebDAV uploads without removing credentials
|
||||
WEBDAV_URL=https://webdav.example.com/path
|
||||
WEBDAV_USERNAME=webdav_user
|
||||
WEBDAV_PASSWORD=your_secure_webdav_password
|
||||
@@ -265,6 +417,7 @@ WEBDAV_FOLDER=/Documents/Uploads
|
||||
WEBDAV_VERIFY_SSL=True
|
||||
|
||||
# FTP
|
||||
# FTP_ENABLED=true # Set to false to disable FTP uploads without removing credentials
|
||||
# Security Note: FTP_USE_TLS=True is strongly recommended for secure connections
|
||||
# Set FTP_ALLOW_PLAINTEXT=False in production to prevent unencrypted FTP
|
||||
FTP_HOST=ftp.example.com
|
||||
@@ -276,6 +429,7 @@ FTP_USE_TLS=True
|
||||
FTP_ALLOW_PLAINTEXT=True
|
||||
|
||||
# SFTP
|
||||
# SFTP_ENABLED=true # Set to false to disable SFTP uploads without removing credentials
|
||||
# Security Note: Host key verification is enabled by default (False)
|
||||
# Only set to True in development/testing environments if needed
|
||||
# When false, configure SSH known_hosts for proper host key verification
|
||||
@@ -288,6 +442,16 @@ SFTP_PASSWORD=your_secure_sftp_password
|
||||
SFTP_FOLDER=/Documents/Uploads
|
||||
SFTP_DISABLE_HOST_KEY_VERIFICATION=False # Default is False (secure); set to True only for testing
|
||||
|
||||
# iCloud Drive
|
||||
# ICLOUD_ENABLED=true # Set to false to disable iCloud uploads without removing credentials
|
||||
# Requires an Apple ID with iCloud Drive enabled.
|
||||
# For accounts with two-factor authentication (most accounts), generate an
|
||||
# app-specific password at https://appleid.apple.com/account/manage
|
||||
ICLOUD_USERNAME=your_apple_id@example.com
|
||||
ICLOUD_PASSWORD=your-app-specific-password
|
||||
ICLOUD_FOLDER=Documents/Uploads
|
||||
# ICLOUD_COOKIE_DIRECTORY=/path/to/cookie/dir # Optional: defaults to ~/.pyicloud
|
||||
|
||||
# **HTTP Request Settings**
|
||||
# Timeout for HTTP requests - set higher to handle large PDF files (up to 1GB)
|
||||
HTTP_REQUEST_TIMEOUT=120 # Timeout in seconds (default: 120 for large file operations)
|
||||
@@ -313,9 +477,27 @@ NOTIFY_ON_STARTUP=True
|
||||
NOTIFY_ON_SHUTDOWN=False
|
||||
NOTIFY_ON_FILE_PROCESSED=True
|
||||
|
||||
# Webhooks – Notify external systems via HTTP POST on document events.
|
||||
# Individual webhooks (URL, events, secret) are managed via /api/webhooks/.
|
||||
WEBHOOK_ENABLED=True
|
||||
|
||||
# Uptime Kuma
|
||||
UPTIME_KUMA_URL=https://status.example.com/api/push/abcdef123456?status=up
|
||||
UPTIME_KUMA_PING_INTERVAL=5
|
||||
UPTIME_KUMA_PING_INTERVAL=5
|
||||
|
||||
# Backup & Restore
|
||||
# Enable automatic scheduled backups (hourly, daily, weekly)
|
||||
BACKUP_ENABLED=True
|
||||
# Directory for local backup archives (defaults to <WORKDIR>/backups)
|
||||
# BACKUP_DIR=/data/backups
|
||||
# Optional remote destination: s3, dropbox, google_drive, onedrive, nextcloud, webdav, ftp, sftp, email
|
||||
# BACKUP_REMOTE_DESTINATION=s3
|
||||
# Sub-folder used when uploading backup archives to the remote destination
|
||||
BACKUP_REMOTE_FOLDER=backups
|
||||
# Retention: number of snapshots to keep per tier
|
||||
BACKUP_RETAIN_HOURLY=96 # 4 days of hourly snapshots
|
||||
BACKUP_RETAIN_DAILY=21 # 3 weeks of daily snapshots
|
||||
BACKUP_RETAIN_WEEKLY=13 # ~3 months of weekly snapshots
|
||||
|
||||
# **Full-Text Search (Meilisearch)**
|
||||
# URL for the Meilisearch instance.
|
||||
@@ -326,4 +508,84 @@ MEILISEARCH_URL=http://meilisearch:7700
|
||||
# Optional master/API key for secured Meilisearch instances
|
||||
# MEILISEARCH_API_KEY=your_master_key_here
|
||||
MEILISEARCH_INDEX_NAME=documents
|
||||
ENABLE_SEARCH=True
|
||||
ENABLE_SEARCH=True
|
||||
|
||||
# **Duplicate Detection**
|
||||
# Exact duplicate detection (SHA-256) is always on during document processing.
|
||||
# The settings below control near-duplicate detection (same scanned content,
|
||||
# different hash) and the visibility of deduplication steps.
|
||||
ENABLE_DEDUPLICATION=True
|
||||
SHOW_DEDUPLICATION_STEP=True
|
||||
# Minimum cosine similarity score (0–1) for two documents to be flagged as
|
||||
# near-duplicates. 0.85 means 85 % semantic overlap. Lower = more matches.
|
||||
NEAR_DUPLICATE_THRESHOLD=0.85
|
||||
|
||||
# **PDF/A Archival Conversion**
|
||||
# When enabled, PDF/A copies of both the original ingested file and the processed
|
||||
# file are created and saved alongside the standard copies. This may double or
|
||||
# triple storage but provides better legal coverage with time-stamped archival copies.
|
||||
# Uses ocrmypdf with Ghostscript for the conversion.
|
||||
ENABLE_PDFA_CONVERSION=false
|
||||
# PDF/A format variant: 1 = PDF/A-1b, 2 = PDF/A-2b (default), 3 = PDF/A-3b
|
||||
PDFA_FORMAT=2
|
||||
# Upload original-file PDF/A variant to all configured storage providers
|
||||
PDFA_UPLOAD_ORIGINAL=false
|
||||
# Upload processed-file PDF/A variant to all configured storage providers
|
||||
PDFA_UPLOAD_PROCESSED=false
|
||||
# Subfolder name appended to each provider's folder for PDF/A uploads
|
||||
# e.g. if Dropbox folder is '/Documents' this puts PDF/A files into '/Documents/pdfa'
|
||||
PDFA_UPLOAD_FOLDER=pdfa
|
||||
# Google Drive folder ID for PDF/A uploads (uses folder IDs, not paths)
|
||||
# Leave empty to use the same folder as regular uploads
|
||||
GOOGLE_DRIVE_PDFA_FOLDER_ID=
|
||||
# RFC 3161 timestamping of PDF/A files (creates .tsr proof-of-existence files)
|
||||
PDFA_TIMESTAMP_ENABLED=false
|
||||
# Timestamp Authority URL (default: FreeTSA, a free RFC 3161 TSA)
|
||||
PDFA_TIMESTAMP_URL=https://freetsa.org/tsr
|
||||
# Model used to generate text embeddings for document similarity.
|
||||
# Must be supported by your OpenAI-compatible API endpoint.
|
||||
EMBEDDING_MODEL=text-embedding-3-small
|
||||
# Maximum tokens to send to the embedding model. Set below the model's
|
||||
# context window (e.g. 8000 for an 8192-token model).
|
||||
EMBEDDING_MAX_TOKENS=8000
|
||||
|
||||
# **Support / Help Center – Zammad Integration**
|
||||
# Base URL of your Zammad instance (required for chat and ticket form).
|
||||
# ZAMMAD_URL=https://zammad.example.com
|
||||
# Show a live-chat widget on the Help Center page (requires an online Zammad agent).
|
||||
# ZAMMAD_CHAT_ENABLED=false
|
||||
# Zammad chat topic ID (see Zammad → Channels → Chat → Topics).
|
||||
# ZAMMAD_CHAT_ID=1
|
||||
# Show a "Submit a Ticket" feedback form on the Help Center page.
|
||||
# ZAMMAD_FORM_ENABLED=false
|
||||
# Support e-mail address displayed on the Help Center page.
|
||||
# SUPPORT_EMAIL=support@example.com
|
||||
|
||||
# **Observability – Sentry Error & Performance Monitoring**
|
||||
# Sentry DSN – obtain from https://sentry.io (Project → Settings → Client Keys).
|
||||
# Leave commented out (or set to empty) to disable Sentry entirely.
|
||||
# SENTRY_DSN=https://<key>@o<org>.ingest.sentry.io/<project>
|
||||
#
|
||||
# Environment label shown in the Sentry dashboard (e.g. development / staging / production).
|
||||
# SENTRY_ENVIRONMENT=production
|
||||
#
|
||||
# Fraction of requests to capture for performance tracing (0.0–1.0).
|
||||
# 0.0 disables tracing; 1.0 captures every request. Default: 0.1 (10 %).
|
||||
# SENTRY_TRACES_SAMPLE_RATE=0.1
|
||||
#
|
||||
# Fraction of profiled transactions to send to Sentry (0.0–1.0).
|
||||
# Profiling is only active when SENTRY_TRACES_SAMPLE_RATE > 0. Default: 0.0 (disabled).
|
||||
# SENTRY_PROFILES_SAMPLE_RATE=0.0
|
||||
#
|
||||
# Attach PII (IP addresses, user agents) to Sentry events.
|
||||
# Disable (default) to stay GDPR/CCPA compliant.
|
||||
# SENTRY_SEND_DEFAULT_PII=false
|
||||
|
||||
# **Mobile App – Push Notifications**
|
||||
# Push notifications are delivered via Expo's push notification service
|
||||
# (https://expo.dev/notifications) which routes to APNs (iOS) and FCM (Android).
|
||||
# No additional credentials are required on the server side.
|
||||
# The mobile app registers its Expo push token via POST /api/mobile/register-device.
|
||||
#
|
||||
# To use native FCM/APNs directly (without Expo relay), replace the
|
||||
# send_expo_push_notification function in app/utils/push_notification.py.
|
||||
|
||||
@@ -163,6 +163,13 @@ pytest --tb=short -q
|
||||
- Keep JavaScript minimal - prefer server-side rendering
|
||||
- Follow existing template structure and patterns
|
||||
|
||||
### Internationalization (i18n) & Localization (l10n)
|
||||
- **Always** use the `_("key")` helper in Jinja2 templates and `translate("key", locale)` in Python for every user-visible string — never hardcode UI text.
|
||||
- **Only add new keys to `frontend/translations/en.json`** — that is the one and only file you must touch when introducing new UI strings.
|
||||
- Do **not** manually edit any non-English translation file (`de.json`, `fr.json`, etc.). An external automation script syncs all other language files from `en.json` automatically.
|
||||
- Key naming convention: `<section>.<descriptor>` in snake_case, e.g. `language.search_placeholder`, `nav.help`, `common.cancel`.
|
||||
- The `test_all_languages_have_same_keys` check has been intentionally removed — key completeness across locales is enforced by the external sync script, not by the test suite.
|
||||
|
||||
### Testing
|
||||
- Write tests in `tests/` directory, mirroring `app/` structure
|
||||
- Use pytest markers: `@pytest.mark.unit`, `@pytest.mark.integration`, etc.
|
||||
|
||||
+61
-198
@@ -2,15 +2,10 @@ name: CI Pipeline
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- develop
|
||||
tags:
|
||||
- 'v*'
|
||||
- '[0-9]+.*'
|
||||
branches: [main, develop]
|
||||
tags: ['v*', '[0-9]+.*']
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
branches: [main]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -22,272 +17,149 @@ concurrency:
|
||||
|
||||
env:
|
||||
IMAGE_NAME: christianlouis/docuelevate
|
||||
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
|
||||
|
||||
jobs:
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 1: Ruff Lint & Format (runs first to catch style issues early)
|
||||
# Stage 1: Static Analysis (Fast Fail Gates)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
lint:
|
||||
name: Ruff Lint & Format
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
|
||||
cache: 'pip'
|
||||
- name: Install Ruff
|
||||
run: pip install ruff
|
||||
|
||||
- name: Show Ruff version (debug)
|
||||
run: ruff --version
|
||||
|
||||
- name: Check for merge conflict markers
|
||||
run: |
|
||||
if git grep -rn -E '^(<{7} |>{7} |={7}$)' -- '.'; then
|
||||
echo "ERROR: Merge conflict markers found in tracked files."
|
||||
echo "ERROR: Merge conflict markers found."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Run Ruff Lint (check)
|
||||
# ruff check can --fix locally, but CI should only check (no modifications)
|
||||
run: ruff check app/ tests/
|
||||
|
||||
- name: Run Ruff Format check
|
||||
# ruff format only supports --check; do not pass --fix here
|
||||
run: ruff format --check app/ tests/
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 1b: HTML Accessibility Lint (catches a11y regressions early)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
- run: ruff check app/ tests/
|
||||
- run: ruff format --check app/ tests/
|
||||
|
||||
html-lint:
|
||||
name: HTML Accessibility Lint
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
|
||||
- name: Install djLint
|
||||
run: pip install djlint>=1.36.0
|
||||
|
||||
- name: Lint HTML templates for accessibility
|
||||
run: djlint frontend/templates/ --lint
|
||||
cache: 'pip'
|
||||
- run: pip install djlint>=1.36.0
|
||||
- run: djlint frontend/templates/ --lint
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 2a: Dependency Vulnerability Scan (runs in parallel with lint)
|
||||
# Stage 2: Parallel Heavy Lifters (Consolidated for Efficiency)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
mypy:
|
||||
name: Mypy Type Check
|
||||
runs-on: ubuntu-latest
|
||||
needs: [lint]
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
cache: 'pip'
|
||||
- name: Install Dependencies
|
||||
run: pip install -r requirements-dev.txt
|
||||
- run: mypy app/
|
||||
|
||||
dependency-scan:
|
||||
name: Dependency Vulnerability Scan
|
||||
name: Dependency Scan
|
||||
runs-on: ubuntu-latest
|
||||
needs: [lint]
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
cache: 'pip'
|
||||
- run: pip install pip-audit>=2.7.0
|
||||
- run: pip-audit -r requirements.txt --desc on
|
||||
|
||||
- name: Install pip-audit
|
||||
run: pip install pip-audit>=2.7.0
|
||||
|
||||
- name: Run pip-audit on production dependencies
|
||||
run: pip-audit -r requirements.txt --desc on
|
||||
|
||||
- name: Run pip-audit on dev dependencies
|
||||
run: pip-audit -r requirements-dev.txt --desc on
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 2b: Quick Tests (unit + basic integration — fast fail gate)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
test-quick:
|
||||
name: Quick Tests
|
||||
run-tests:
|
||||
name: Execute All Tests (Quick + Integration)
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
needs: [lint, html-lint, dependency-scan]
|
||||
needs: [lint]
|
||||
services:
|
||||
redis:
|
||||
image: redis:7
|
||||
ports:
|
||||
- 6379:6379
|
||||
options: >-
|
||||
--health-cmd "redis-cli ping"
|
||||
--health-interval 10s
|
||||
--health-timeout 5s
|
||||
--health-retries 5
|
||||
ports: ["6379:6379"]
|
||||
options: --health-cmd "redis-cli ping" --health-interval 10s --health-timeout 5s --health-retries 5
|
||||
rabbitmq:
|
||||
image: rabbitmq:3-management
|
||||
ports: ["5672:5672", "15672:15672"]
|
||||
options: --health-cmd "rabbitmq-diagnostics -q ping" --health-interval 10s --health-timeout 5s --health-retries 5
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
cache: 'pip'
|
||||
|
||||
- name: Install Dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-dev.txt
|
||||
|
||||
- name: Run Quick Tests
|
||||
- name: Run Tests
|
||||
run: >
|
||||
pytest tests/ -v
|
||||
--timeout=120
|
||||
--cov=app --cov-report=xml --cov-report=term
|
||||
pytest tests/ -v --timeout=300
|
||||
--cov=app --cov-report=xml:coverage.xml
|
||||
--junitxml=junit.xml -o junit_family=legacy
|
||||
-m "not e2e and not requires_docker and not requires_external and not slow"
|
||||
-m "not e2e"
|
||||
|
||||
- name: Upload coverage reports to Codecov
|
||||
if: ${{ !cancelled() }}
|
||||
- name: Upload Unified Coverage to Codecov
|
||||
if: always()
|
||||
uses: codecov/codecov-action@v5
|
||||
with:
|
||||
token: ${{ secrets.CODECOV_TOKEN }}
|
||||
files: ./coverage.xml
|
||||
fail_ci_if_error: false
|
||||
|
||||
- name: Upload test results to Codecov
|
||||
if: ${{ !cancelled() }}
|
||||
uses: codecov/codecov-action@v5
|
||||
with:
|
||||
token: ${{ secrets.CODECOV_TOKEN }}
|
||||
files: ./junit.xml
|
||||
report_type: test_results
|
||||
fail_ci_if_error: false
|
||||
|
||||
- name: Upload test artifacts
|
||||
if: ${{ !cancelled() }}
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: test-results-quick
|
||||
path: |
|
||||
junit.xml
|
||||
coverage.xml
|
||||
fail_ci_if_error: true
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 2c: Integration Tests (Docker containers, external services)
|
||||
# Stage 3: Build & Push (Quality Gate)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
test-integration:
|
||||
name: Integration Tests
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
needs: [test-quick] # Only run after quick tests pass (fail early)
|
||||
services:
|
||||
redis:
|
||||
image: redis:7
|
||||
ports:
|
||||
- 6379:6379
|
||||
options: >-
|
||||
--health-cmd "redis-cli ping"
|
||||
--health-interval 10s
|
||||
--health-timeout 5s
|
||||
--health-retries 5
|
||||
rabbitmq:
|
||||
image: rabbitmq:3-management
|
||||
ports:
|
||||
- 5672:5672
|
||||
- 15672:15672
|
||||
options: >-
|
||||
--health-cmd "rabbitmq-diagnostics -q ping"
|
||||
--health-interval 10s
|
||||
--health-timeout 5s
|
||||
--health-retries 5
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
|
||||
- name: Install Dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-dev.txt
|
||||
|
||||
- name: Run Integration Tests
|
||||
run: >
|
||||
pytest tests/ -v
|
||||
--timeout=300
|
||||
--junitxml=junit-integration.xml -o junit_family=legacy
|
||||
-m "(requires_docker or requires_external or slow) and not e2e"
|
||||
|
||||
- name: Upload integration test results
|
||||
if: ${{ !cancelled() }}
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: test-results-integration
|
||||
path: junit-integration.xml
|
||||
|
||||
mypy:
|
||||
name: Mypy
|
||||
runs-on: ubuntu-latest
|
||||
needs: [lint, html-lint, dependency-scan] # Wait for lint, HTML a11y lint, and dependency scan before running type checks
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
|
||||
- name: Install Dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-dev.txt
|
||||
|
||||
- name: Run Mypy
|
||||
run: mypy app/
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 3: Build & Push Docker Image (only after all Stage 2 jobs pass)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
build:
|
||||
name: Build & Push Docker Image
|
||||
runs-on: ubuntu-latest
|
||||
needs: [test-quick, test-integration, lint, html-lint, mypy, dependency-scan]
|
||||
needs: [run-tests, mypy, dependency-scan, html-lint]
|
||||
if: github.event_name == 'push'
|
||||
|
||||
steps:
|
||||
- name: Checkout Code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Generate Build Metadata
|
||||
run: |
|
||||
chmod +x scripts/generate_build_metadata.sh
|
||||
./scripts/generate_build_metadata.sh
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Log in to Docker Hub
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
username: ${{ secrets.DOCKER_USERNAME }}
|
||||
password: ${{ secrets.DOCKER_PASSWORD }}
|
||||
|
||||
- name: Log in to GitHub Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Extract metadata for tags
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
@@ -299,48 +171,40 @@ jobs:
|
||||
type=ref,event=branch
|
||||
type=sha,prefix={{branch}}-
|
||||
type=semver,pattern={{version}}
|
||||
type=semver,pattern={{major}}.{{minor}}
|
||||
type=raw,value=latest,enable={{is_default_branch}}
|
||||
|
||||
- name: Build and Push Docker Image
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: Dockerfile
|
||||
platforms: linux/amd64
|
||||
push: true
|
||||
sbom: true
|
||||
provenance: mode=max
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
sbom: true
|
||||
provenance: mode=max
|
||||
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
# Stage 4: Update preprod K8s manifest (ArgoCD GitOps, only on main)
|
||||
# Stage 4: GitOps Update
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
update-k8s-manifest:
|
||||
name: Update Preprod K8s Manifest
|
||||
runs-on: ubuntu-latest
|
||||
needs: [build]
|
||||
if: github.ref == 'refs/heads/main' && github.event_name == 'push'
|
||||
|
||||
steps:
|
||||
- name: Compute image tag
|
||||
id: tag
|
||||
run: |
|
||||
SHORT_SHA=$(echo "${{ github.sha }}" | cut -c1-7)
|
||||
echo "image=ghcr.io/${{ github.repository_owner }}/docuelevate:main-${SHORT_SHA}" >> "$GITHUB_OUTPUT"
|
||||
echo "tag=main-${SHORT_SHA}" >> "$GITHUB_OUTPUT"
|
||||
|
||||
echo "image=ghcr.io/${{ github.repository_owner }}/docuelevate:main-${SHORT_SHA}" >> "$GITHUB_OUTPUT"
|
||||
- name: Checkout k8s-cluster-state
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
repository: christianlouis/k8s-cluster-state
|
||||
token: ${{ secrets.GH_PAT }}
|
||||
path: k8s-cluster-state
|
||||
|
||||
- name: Update image tag in preprod manifest
|
||||
uses: mikefarah/yq@v4.44.6
|
||||
env:
|
||||
@@ -349,7 +213,6 @@ jobs:
|
||||
cmd: |
|
||||
yq -i '(.. | select(tag == "!!str") | select(test("^(ghcr\\.io/christianlouis/docuelevate|christianlouis/docuelevate):"))) = strenv(IMAGE)' \
|
||||
k8s-cluster-state/apps/docuelevate/preprod/docuelevate-stack.yaml
|
||||
|
||||
- name: Commit and push
|
||||
run: |
|
||||
cd k8s-cluster-state
|
||||
@@ -357,7 +220,7 @@ jobs:
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
git add apps/docuelevate/preprod/docuelevate-stack.yaml
|
||||
if git diff --staged --quiet; then
|
||||
echo "No changes to commit -- image tag already up to date"
|
||||
echo "No changes to commit"
|
||||
else
|
||||
git commit -m "chore(preprod): update docuelevate image to ${{ steps.tag.outputs.tag }}"
|
||||
git push
|
||||
|
||||
@@ -8,6 +8,9 @@ on:
|
||||
schedule:
|
||||
- cron: '37 1 * * 1'
|
||||
|
||||
env:
|
||||
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
|
||||
|
||||
jobs:
|
||||
analyze:
|
||||
name: Analyze (${{ matrix.language }})
|
||||
@@ -24,6 +27,8 @@ jobs:
|
||||
include:
|
||||
- language: actions
|
||||
build-mode: none
|
||||
- language: javascript
|
||||
build-mode: none
|
||||
- language: javascript-typescript
|
||||
build-mode: none
|
||||
- language: python
|
||||
|
||||
@@ -12,6 +12,9 @@ permissions:
|
||||
pull-requests: write
|
||||
packages: write
|
||||
|
||||
env:
|
||||
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
|
||||
|
||||
jobs:
|
||||
release:
|
||||
name: Semantic Release
|
||||
|
||||
@@ -16,6 +16,9 @@ permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
env:
|
||||
FORCE_JAVASCRIPT_ACTIONS_TO_NODE24: true
|
||||
|
||||
jobs:
|
||||
ruff-auto-fix:
|
||||
name: Auto-fix Ruff Issues
|
||||
|
||||
@@ -171,6 +171,7 @@ venv.bak/
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
/docs_build
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
|
||||
@@ -0,0 +1,4 @@
|
||||
## 2024-05-24 - SSRF in WebDAV connection test
|
||||
**Vulnerability:** The `_test_webdav_connection` function had a custom SSRF check that failed to resolve DNS names, allowing attackers to bypass the check by providing a domain that resolves to an internal IP (e.g., `127.0.0.1`).
|
||||
**Learning:** DNS resolution is required for robust SSRF protection when validating URLs provided by users.
|
||||
**Prevention:** Use a centralized `is_private_ip` function (now in `app/utils/network.py`) that resolves the hostname to its IPs and checks if any are private.
|
||||
+1
-1
@@ -1 +1 @@
|
||||
2026-03-01T17:23:25Z
|
||||
2026-03-15T21:39:26Z
|
||||
|
||||
+2355
File diff suppressed because it is too large
Load Diff
+21
@@ -7,6 +7,22 @@ WORKDIR /app
|
||||
COPY requirements.txt /app/
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# ── Documentation build stage ───────────────────────────────────────────────
|
||||
FROM python:3.14.1-slim AS docs-builder
|
||||
|
||||
WORKDIR /docs
|
||||
|
||||
# Install MkDocs Material and its dependencies
|
||||
COPY docs/requirements.txt /docs/requirements.txt
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# Copy documentation sources
|
||||
COPY docs /docs/docs
|
||||
COPY mkdocs.yml /docs/mkdocs.yml
|
||||
|
||||
# Build the static documentation site
|
||||
RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build
|
||||
|
||||
# Second stage for the actual runtime
|
||||
FROM python:3.14.3-slim
|
||||
|
||||
@@ -33,6 +49,8 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
# Copy application code
|
||||
COPY ./app /app/app
|
||||
COPY ./frontend /app/frontend
|
||||
COPY ./migrations /app/migrations
|
||||
COPY ./alembic.ini /app/alembic.ini
|
||||
COPY ./LICENSE /app/LICENSE
|
||||
|
||||
# Copy build metadata files (generated at build time)
|
||||
@@ -41,6 +59,9 @@ COPY ./BUILD_DATE /app/BUILD_DATE
|
||||
COPY ./GIT_SHA /app/GIT_SHA
|
||||
COPY ./RUNTIME_INFO /app/RUNTIME_INFO
|
||||
|
||||
# Copy the pre-built MkDocs documentation site (served at /help)
|
||||
COPY --from=docs-builder /docs/docs_build /app/docs_build
|
||||
|
||||
# Create runtime_info directory
|
||||
RUN mkdir -p /app/runtime_info
|
||||
|
||||
|
||||
@@ -6,6 +6,19 @@ WORKDIR /app
|
||||
COPY requirements.txt /app/
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# ── Documentation build stage ───────────────────────────────────────────────
|
||||
FROM python:3.14.1-slim AS docs-builder
|
||||
|
||||
WORKDIR /docs
|
||||
|
||||
COPY docs/requirements.txt /docs/requirements.txt
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
COPY docs /docs/docs
|
||||
COPY mkdocs.yml /docs/mkdocs.yml
|
||||
|
||||
RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build
|
||||
|
||||
FROM python:3.14.1-slim
|
||||
|
||||
WORKDIR /app
|
||||
@@ -27,10 +40,15 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
|
||||
COPY ./app /app/app
|
||||
COPY ./frontend /app/frontend
|
||||
COPY ./migrations /app/migrations
|
||||
COPY ./alembic.ini /app/alembic.ini
|
||||
COPY ./LICENSE /app/LICENSE
|
||||
COPY ./VERSION /app/VERSION
|
||||
COPY ./BUILD_DATE /app/BUILD_DATE
|
||||
|
||||
# Copy the pre-built MkDocs documentation site (served at /help)
|
||||
COPY --from=docs-builder /docs/docs_build /app/docs_build
|
||||
|
||||
# Local fallbacks for build metadata
|
||||
RUN echo "local" > /app/GIT_SHA \
|
||||
&& echo "local" > /app/RUNTIME_INFO
|
||||
|
||||
@@ -24,121 +24,154 @@
|
||||
</div>
|
||||
|
||||
<div align="center">
|
||||
<a href="https://www.docuelevate.org"><img src="frontend/static/hero.png" alt="DocuElevate Logo" width="80%" /></a>
|
||||
<a href="https://www.docuelevate.org"><img src="frontend/static/hero.png" alt="DocuElevate Hero" width="80%" /></a>
|
||||
</div>
|
||||
|
||||
## Overview
|
||||
|
||||
DocuElevate automates the handling, extraction, and processing of documents using a variety of services, including:
|
||||
DocuElevate is an intelligent document processing system that automates the ingestion, OCR, AI-powered metadata extraction, and distribution of documents. It supports a wide range of AI providers, OCR engines, and cloud storage destinations out of the box.
|
||||
|
||||
- **AI Provider** (pluggable – OpenAI, Anthropic, Gemini, Ollama, OpenRouter, Portkey, and more) for metadata extraction and text refinement.
|
||||
- **Dropbox**, **Nextcloud**, and **Google Drive** for file storage and uploads.
|
||||
- **Paperless NGX** for document indexing and management.
|
||||
- **Azure Document Intelligence** for OCR on PDFs.
|
||||
- **Gotenberg** for file-to-PDF conversions.
|
||||
- **Authentik** for authentication and user management.
|
||||
**Key capabilities:**
|
||||
|
||||
It is designed for flexibility and configurability through environment variables, making it easily customizable for different workflows. The system can fetch documents from multiple IMAP mailboxes, process them (OCR, metadata extraction, PDF conversion), and store them in the desired destinations.
|
||||
- **AI-Powered Metadata Extraction** — pluggable AI providers including OpenAI, Anthropic Claude, Google Gemini, Ollama (local), OpenRouter, Portkey, and Azure OpenAI via LiteLLM
|
||||
- **Multi-Engine OCR** — Azure Document Intelligence, Tesseract, EasyOCR, Mistral OCR, Google Cloud Document AI, and AWS Textract with configurable merge strategies
|
||||
- **12 Storage Destinations** — Dropbox, Google Drive, OneDrive, Amazon S3, Nextcloud, WebDAV, FTP, SFTP, iCloud Drive, Email (SMTP), Paperless-ngx, and Rclone
|
||||
- **Multi-Channel Ingestion** — web upload, browser extension, mobile app, CLI, REST API, IMAP email, and watched folders (local, cloud, FTP/SFTP)
|
||||
- **Processing Pipelines** — customizable multi-step workflows with conditional routing rules
|
||||
- **Full-Text Search** — powered by Meilisearch for instant document discovery
|
||||
- **Multi-User with SSO** — local accounts, OAuth2/OIDC (Authentik), and social login (Google, Microsoft, Apple, Dropbox)
|
||||
|
||||
The project includes a **UI** for uploading and managing files, and an API documentation page is available at `/docs` (powered by **FastAPI**).
|
||||
|
||||
## Documentation Index
|
||||
|
||||
- [User Guide](docs/UserGuide.md) - How to use DocuElevate
|
||||
- [Browser Extension Guide](docs/BrowserExtension.md) - Install and use the browser extension
|
||||
- [API Documentation](docs/API.md) - API reference
|
||||
- [Deployment Guide](docs/DeploymentGuide.md) - How to deploy DocuElevate
|
||||
- [Configuration Guide](docs/ConfigurationGuide.md) - Available configuration options
|
||||
- [Build Metadata](docs/BuildMetadata.md) - Automated version and build information
|
||||
- [CI/CD Tools Guide](docs/CIToolsGuide.md) - CI/CD pipeline and tool documentation
|
||||
- [CI Workflow Guide](docs/CIWorkflow.md) - Detailed workflow documentation
|
||||
- [Development Guide](CONTRIBUTING.md) - How to contribute to DocuElevate
|
||||
- [Troubleshooting](docs/Troubleshooting.md) - Common issues and solutions
|
||||
The project ships with a web UI, a REST + GraphQL API, a CLI tool, a native mobile app (iOS & Android), a browser extension, and Helm charts for Kubernetes deployment.
|
||||
|
||||
## Screenshots
|
||||
|
||||
<div align="center">
|
||||
<img src="docs/upload-view.png" alt="DocuElevate Upload Interface" width="80%" />
|
||||
<p><em>Upload interface for adding new documents</em></p>
|
||||
<p><em>Upload interface — drag-and-drop file upload with real-time progress</em></p>
|
||||
|
||||
<img src="docs/files-view.png" alt="DocuElevate Files View" width="80%" />
|
||||
<p><em>Files view with processed documents and metadata</em></p>
|
||||
<p><em>Files view — processed documents with AI-extracted metadata</em></p>
|
||||
|
||||
<img src="docs/status-view.png" alt="DocuElevate Status View" width="80%" />
|
||||
<p><em>Status view — system health and service monitoring</em></p>
|
||||
</div>
|
||||
|
||||
> **Note:** Screenshots may not reflect the very latest UI. For the most current look, visit [docuelevate.org](https://www.docuelevate.org).
|
||||
|
||||
## Workflow Process
|
||||
|
||||
DocuElevate follows a streamlined document processing workflow:
|
||||
## Workflow
|
||||
|
||||
<div align="center">
|
||||
<img src="docs/workflow-diagram.png" alt="DocuElevate Workflow" width="90%" />
|
||||
</div>
|
||||
|
||||
### Document Ingestion
|
||||
Documents enter DocuElevate through four possible channels:
|
||||
1. **Web Upload**: Users manually upload files via the web interface
|
||||
2. **Browser Extension**: Send files directly from your browser with one click
|
||||
3. **Email Attachments**: Automatic polling of configured IMAP mailboxes (supports multiple accounts)
|
||||
4. **API**: Direct programmatic uploads via the REST API
|
||||
### Ingestion
|
||||
|
||||
Documents enter DocuElevate through multiple channels:
|
||||
|
||||
| Channel | Description |
|
||||
|---------|-------------|
|
||||
| **Web Upload** | Drag-and-drop interface with real-time progress (up to 1 GB per file) |
|
||||
| **Browser Extension** | Clip web pages or send files from Chrome, Firefox, or Edge |
|
||||
| **Mobile App** | Capture documents with the device camera or upload from the photo library |
|
||||
| **CLI** | Batch uploads and scripted workflows via the `docuelevate` command-line tool |
|
||||
| **REST API** | Programmatic uploads with full API-token authentication |
|
||||
| **Email (IMAP)** | Automatic polling of multiple mailboxes with attachment filtering |
|
||||
| **Watched Folders** | Monitor local paths, FTP, SFTP, S3, Dropbox, Google Drive, OneDrive, Nextcloud, or WebDAV for new files |
|
||||
|
||||
### Processing Pipeline
|
||||
Every document goes through the following steps:
|
||||
1. **PDF Conversion**: Non-PDF files are converted to PDF format using Gotenberg
|
||||
2. **OCR Processing**: Azure Document Intelligence extracts text from images/scans
|
||||
3. **Metadata Extraction**: The configured AI provider analyzes document content to identify:
|
||||
- Document type (invoice, receipt, contract, etc.)
|
||||
- Key entities (dates, names, amounts, account numbers)
|
||||
- Important data points specific to the document type
|
||||
4. **Enrichment**: Metadata is attached to the document in a structured format
|
||||
|
||||
Each document passes through a configurable set of steps:
|
||||
|
||||
1. **PDF Conversion** — Non-PDF files are converted using Gotenberg, with optional PDF/A archival conversion
|
||||
2. **OCR** — Text extraction via one or more OCR engines (Azure, Tesseract, EasyOCR, Mistral, Google Document AI, AWS Textract) with configurable merge strategies
|
||||
3. **AI Metadata Extraction** — The configured AI provider classifies the document and extracts structured metadata (type, dates, amounts, entities)
|
||||
4. **Enrichment** — Metadata is embedded into the PDF and stored alongside the document
|
||||
5. **Embedding Generation** — Vector embeddings for similarity search and duplicate detection
|
||||
|
||||
Steps can be customized using **Pipelines** and **Routing Rules** for conditional processing.
|
||||
|
||||
### Distribution
|
||||
Processed documents with their metadata can be automatically sent to:
|
||||
- **Dropbox**: For cloud storage and sharing
|
||||
- **Nextcloud**: For self-hosted file storage
|
||||
- **Google Drive**: For Google Workspace integration
|
||||
- **Paperless-NGX**: For advanced document management with search capabilities
|
||||
|
||||
Users can choose to send documents to any combination of these destinations through configuration settings or manual selection.
|
||||
Processed documents are distributed to any combination of configured destinations:
|
||||
|
||||
| Destination | Type |
|
||||
|------------|------|
|
||||
| **Dropbox** | Cloud storage |
|
||||
| **Google Drive** | Cloud storage |
|
||||
| **OneDrive** | Cloud storage |
|
||||
| **Amazon S3** | Object storage |
|
||||
| **Nextcloud** | Self-hosted cloud |
|
||||
| **WebDAV** | Protocol-based |
|
||||
| **FTP / SFTP** | File transfer |
|
||||
| **iCloud Drive** | Apple cloud |
|
||||
| **Email (SMTP)** | Send as attachment |
|
||||
| **Paperless-ngx** | Document management system |
|
||||
| **Rclone** | 70+ cloud providers via Rclone |
|
||||
|
||||
## Features
|
||||
|
||||
- **Intuitive File Upload**:
|
||||
- Drag-and-drop file upload on both Upload and Files pages—upload anywhere on the Files page
|
||||
- Real-time upload progress with validation
|
||||
- Support for PDF, Office documents, images, and more (up to 500MB per file)
|
||||
- **Browser Extension**:
|
||||
- Send files directly from your browser to DocuElevate with one click
|
||||
- Compatible with Chrome, Firefox, Edge, and other Chromium-based browsers
|
||||
- Context menu integration for quick access
|
||||
- See [Browser Extension Guide](docs/BrowserExtension.md) for installation and usage
|
||||
- **Document Upload & Storage**:
|
||||
- Manual uploads (via API or UI) to Dropbox, Nextcloud, Google Drive, or Paperless
|
||||
- **OCR Processing (Azure)**:
|
||||
- Extract text from scanned PDFs using Azure Document Intelligence
|
||||
- **Metadata Extraction (AI Provider)**:
|
||||
- Use any supported AI provider (OpenAI, Anthropic, Gemini, Ollama, etc.) to classify, label, or otherwise enrich the text with structured metadata
|
||||
- **PDF Conversion (Gotenberg)**:
|
||||
- Convert non-PDF attachments (e.g., Word docs, images) into PDFs
|
||||
- **Document Management (Paperless NGX)**:
|
||||
- Store processed documents and metadata in a Paperless NGX instance
|
||||
- **IMAP Integration**:
|
||||
- Fetch documents from multiple mailboxes (including Gmail) and automatically enqueue them for processing
|
||||
- **Authentication**:
|
||||
- Secure access to the system using **Authentik** for OAuth2-based login
|
||||
### Document Processing
|
||||
- **Multi-engine OCR** with quality checks and configurable merge strategies (AI merge, longest, primary)
|
||||
- **AI metadata extraction** using any supported provider (OpenAI, Anthropic, Gemini, Ollama, OpenRouter, Portkey, Azure OpenAI)
|
||||
- **PDF conversion** via Gotenberg with optional PDF/A archival format
|
||||
- **Duplicate detection** — exact (SHA-256) and near-duplicate (content similarity with vector embeddings)
|
||||
- **Customizable pipelines** — define multi-step processing workflows with conditional routing rules
|
||||
|
||||
## Frameworks Used
|
||||
### Document Management
|
||||
- **Full-text search** powered by Meilisearch with saved searches
|
||||
- **File detail view** with metadata, text preview, processing history, and similarity analysis
|
||||
- **Shared links** for public document access with expiration controls
|
||||
- **Bulk operations** — reprocess, delete, or reassign documents in batch
|
||||
|
||||
- **FastAPI**: High-performance web framework for APIs.
|
||||
- **Celery**: Task queue for asynchronous processing.
|
||||
- **Redis**: Message broker and result backend.
|
||||
- **SQLAlchemy**: ORM for database interactions.
|
||||
- **Tailwind CSS**: Utility-first CSS framework.
|
||||
- **Docker**: Containerization for easy deployment.
|
||||
### Multi-Channel Ingestion
|
||||
- **Web UI** — drag-and-drop upload with real-time progress
|
||||
- **Browser extension** — clip web pages or send files from Chrome, Firefox, Edge ([guide](docs/BrowserExtension.md))
|
||||
- **Mobile app** — iOS and Android with camera capture, push notifications, and SSO ([guide](docs/MobileApp.md))
|
||||
- **CLI tool** — batch uploads, downloads, search, and API-token management ([guide](docs/CLIGuide.md))
|
||||
- **REST API & GraphQL** — full programmatic access with Swagger documentation at `/docs`
|
||||
- **IMAP email** — poll multiple mailboxes with attachment filtering and auto-processing
|
||||
- **Watched folders** — local filesystem, FTP, SFTP, and cloud storage providers
|
||||
|
||||
### Administration
|
||||
- **Multi-user mode** with per-user document isolation and ownership
|
||||
- **Subscription & billing** — Stripe integration with configurable plans and quotas
|
||||
- **Scheduled jobs** — IMAP polling, watched folder scans, automated backups, uptime monitoring
|
||||
- **Audit logging** with SIEM integration support
|
||||
- **Compliance templates** — GDPR, HIPAA, SOC 2
|
||||
- **Admin dashboard** — user management, queue monitoring, credential management, backup/restore
|
||||
|
||||
### Authentication & Security
|
||||
- **Local accounts** with self-service registration and password reset
|
||||
- **OAuth2/OIDC** via Authentik or any OIDC provider
|
||||
- **Social login** — Google, Microsoft, Apple, Dropbox
|
||||
- **API tokens** for CLI, mobile, and automation access
|
||||
- **Security headers** — HSTS, CSP, X-Frame-Options, X-Content-Type-Options
|
||||
- **Rate limiting** with configurable per-endpoint controls
|
||||
|
||||
### Notifications
|
||||
- **100+ notification backends** via Apprise — Discord, Telegram, Slack, Microsoft Teams, Email, webhooks, and more
|
||||
- **Configurable events** — task failures, credential issues, file processed, user signup, payment issues
|
||||
- **In-app notification inbox** with per-user preferences
|
||||
- **Webhooks** — push events to external systems with HMAC signature verification and retry
|
||||
|
||||
## Tech Stack
|
||||
|
||||
| Component | Technology |
|
||||
|-----------|-----------|
|
||||
| **Backend** | FastAPI, Celery, Redis, SQLAlchemy, Alembic |
|
||||
| **Frontend** | Jinja2, Tailwind CSS |
|
||||
| **Search** | Meilisearch |
|
||||
| **Mobile** | React Native (Expo) — iOS & Android |
|
||||
| **AI** | LiteLLM (OpenAI, Anthropic, Gemini, Ollama, OpenRouter, Portkey) |
|
||||
| **OCR** | Azure Document Intelligence, Tesseract, EasyOCR, Mistral, Google Doc AI, AWS Textract |
|
||||
| **PDF** | Gotenberg, pypdf |
|
||||
| **Auth** | Authlib (OAuth2/OIDC), MSAL, social providers |
|
||||
| **Infrastructure** | Docker, Docker Compose, Helm/Kubernetes |
|
||||
| **Docs** | MkDocs Material |
|
||||
|
||||
## Quick Start
|
||||
|
||||
For detailed installation and deployment instructions, please refer to the [Deployment Guide](docs/DeploymentGuide.md).
|
||||
For detailed installation and deployment instructions, see the [Deployment Guide](docs/DeploymentGuide.md).
|
||||
|
||||
```bash
|
||||
# Clone the repository
|
||||
@@ -147,20 +180,96 @@ cd DocuElevate
|
||||
|
||||
# Configure environment variables
|
||||
cp .env.demo .env
|
||||
# Edit .env with your settings
|
||||
# Edit .env with your settings (see Configuration Guide for all options)
|
||||
|
||||
# Run with Docker Compose
|
||||
docker-compose up -d
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
The API will be available at **`http://localhost:8000`**, and the API documentation is available at **`http://localhost:8000/docs`**.
|
||||
The web UI is available at **`http://localhost:8000`** and the interactive API documentation at **`http://localhost:8000/docs`**.
|
||||
|
||||
### Kubernetes / Helm
|
||||
|
||||
```bash
|
||||
helm repo add docuelevate https://christianlouis.github.io/DocuElevate
|
||||
helm install docuelevate docuelevate/docuelevate -f values.yaml
|
||||
```
|
||||
|
||||
See the [Kubernetes Deployment Guide](docs/KubernetesDeployment.md) for full details.
|
||||
|
||||
## Documentation
|
||||
|
||||
### Getting Started
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [Setup Wizard](docs/SetupWizard.md) | Interactive first-run setup |
|
||||
| [User Guide](docs/UserGuide.md) | How to use DocuElevate |
|
||||
| [Browser Extension](docs/BrowserExtension.md) | Install and use the browser extension |
|
||||
| [Mobile App](docs/MobileApp.md) | iOS and Android mobile app |
|
||||
| [CLI Guide](docs/CLIGuide.md) | Command-line tool for automation |
|
||||
|
||||
### How-To Guides
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [How-To Overview](docs/HowToGuides.md) | Index of all how-to guides |
|
||||
| [Email Ingestion](docs/howto/EmailIngestion.md) | Set up IMAP email polling |
|
||||
| [Watched Folder](docs/howto/WatchedFolderSetup.md) | Monitor local or remote folders |
|
||||
| [Mobile Scanning](docs/howto/MobileScanning.md) | Scan documents with your phone |
|
||||
|
||||
### Reference
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [API Documentation](docs/API.md) | REST & GraphQL API reference |
|
||||
| [Configuration Guide](docs/ConfigurationGuide.md) | All environment variables |
|
||||
| [Configuration Master](docs/ConfigurationMaster.md) | Configuration overview |
|
||||
| [Settings Management](docs/SettingsManagement.md) | Runtime settings UI |
|
||||
|
||||
### Deployment & Operations
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [Deployment Guide](docs/DeploymentGuide.md) | Docker Compose deployment |
|
||||
| [Kubernetes / Helm](docs/KubernetesDeployment.md) | Kubernetes deployment with Helm charts |
|
||||
| [Production Readiness](docs/ProductionReadiness.md) | Checklist for production environments |
|
||||
| [Database Configuration](docs/DatabaseConfiguration.md) | Database setup and migration |
|
||||
| [Backup & Restore](docs/ConfigurationGuide.md#backup--restore) | Automated backup configuration |
|
||||
|
||||
### Storage Integration Setup
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [Dropbox](docs/DropboxSetup.md) | Dropbox OAuth setup |
|
||||
| [Google Drive](docs/GoogleDriveSetup.md) | Google Drive service account / OAuth |
|
||||
| [OneDrive](docs/OneDriveSetup.md) | Microsoft OneDrive setup |
|
||||
| [Amazon S3](docs/AmazonS3Setup.md) | S3 bucket configuration |
|
||||
| [Authentication](docs/AuthenticationSetup.md) | OAuth2, OIDC, and social login |
|
||||
| [Notifications](docs/NotificationsSetup.md) | Notification backend setup |
|
||||
|
||||
### Security & Compliance
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [Credential Rotation](docs/CredentialRotationGuide.md) | Rotate secrets safely |
|
||||
| [Licensing Compliance](docs/LicensingCompliance.md) | Dependency licenses |
|
||||
| [Privacy & GDPR](docs/PrivacyCompliance.md) | Privacy compliance |
|
||||
|
||||
### Development
|
||||
|
||||
| Guide | Description |
|
||||
|-------|-------------|
|
||||
| [Contributing](CONTRIBUTING.md) | Code style, commits, and PR process |
|
||||
| [Troubleshooting](docs/Troubleshooting.md) | Common issues and solutions |
|
||||
| [Configuration Troubleshooting](docs/ConfigurationTroubleshooting.md) | Configuration-specific issues |
|
||||
| [Build Metadata](docs/BuildMetadata.md) | Version and build information |
|
||||
| [Internationalization](docs/InternationalizationGuide.md) | Translation and localization |
|
||||
|
||||
## Development & Testing
|
||||
|
||||
### Running Tests
|
||||
|
||||
DocuElevate includes comprehensive test coverage. To run tests:
|
||||
|
||||
```bash
|
||||
# Install development dependencies
|
||||
pip install -r requirements-dev.txt
|
||||
@@ -175,21 +284,21 @@ pytest --cov=app --cov-report=term-missing
|
||||
pytest -m unit
|
||||
```
|
||||
|
||||
Tests are automatically configured with the necessary environment variables - **no manual setup required!**
|
||||
Tests are automatically configured with the necessary environment variables — **no manual setup required!**
|
||||
|
||||
For detailed testing information, including integration tests with Docker and authentication testing, see the [Contributing Guide](CONTRIBUTING.md#running-tests).
|
||||
For detailed testing information, see the [Contributing Guide](CONTRIBUTING.md#running-tests).
|
||||
|
||||
### Contributing
|
||||
|
||||
We welcome contributions! Please see [CONTRIBUTING.md](CONTRIBUTING.md) for:
|
||||
- Code style guidelines
|
||||
- Code style guidelines (Ruff for formatting and linting)
|
||||
- Commit message format (Conventional Commits)
|
||||
- Testing requirements
|
||||
- Pull request process
|
||||
|
||||
## License
|
||||
|
||||
This project is licensed under the Apache License 2.0 - see the [LICENSE](LICENSE) file for details.
|
||||
This project is licensed under the Apache License 2.0 — see the [LICENSE](LICENSE) file for details.
|
||||
|
||||
## Third-Party Software
|
||||
|
||||
@@ -214,13 +323,10 @@ The following is a summary of the licenses used by our direct dependencies:
|
||||
| Uvicorn | BSD |
|
||||
| SQLAlchemy | MIT |
|
||||
| Pydantic | MIT |
|
||||
| openai | MIT |
|
||||
| litellm | MIT |
|
||||
| pypdf | BSD |
|
||||
| Requests | Apache 2.0 |
|
||||
| puremagic | MIT |
|
||||
| filetype | MIT |
|
||||
| Dropbox | MIT |
|
||||
| Dropbox SDK | MIT |
|
||||
| Azure AI Document Intelligence | MIT |
|
||||
| Authlib | BSD |
|
||||
| Starlette | BSD |
|
||||
@@ -229,15 +335,15 @@ The following is a summary of the licenses used by our direct dependencies:
|
||||
| Microsoft Graph Core | MIT |
|
||||
| MSAL | MIT |
|
||||
| Boto3 | Apache 2.0 |
|
||||
| Paramiko | LGPL-2.1|
|
||||
| Paramiko | LGPL-2.1 |
|
||||
| Apprise | MIT |
|
||||
| Redis | BSD |
|
||||
| Gotenberg | MIT |
|
||||
| Redis (py) | BSD |
|
||||
| Gotenberg Client | MIT |
|
||||
| Meilisearch | MIT |
|
||||
|
||||
For a comprehensive list of all dependencies and their licenses, run:
|
||||
|
||||
```
|
||||
```bash
|
||||
pip install pip-licenses
|
||||
pip-licenses
|
||||
|
||||
```
|
||||
|
||||
+22
-9
@@ -7,7 +7,20 @@
|
||||
|
||||
DocuElevate aims to be the premier open-source intelligent document processing platform, providing seamless integration with cloud storage providers, advanced AI-powered metadata extraction, and enterprise-grade security and scalability.
|
||||
|
||||
## Current Status (v0.5.0)
|
||||
## Release Naming
|
||||
|
||||
Each major milestone release carries a codename to anchor key project moments. These names appear in the status dashboard, build metadata, and changelog. For details, see [docs/ReleaseNaming.md](docs/ReleaseNaming.md).
|
||||
|
||||
| Version Range | Codename | Theme |
|
||||
|---------------|---------------|--------------------------------------------------|
|
||||
| 0.5.x | **Foundation** | Core platform, multi-provider storage, AI, UI |
|
||||
| 0.6.x | **Clarity** | Enhanced search, filtering, UI/UX improvements |
|
||||
| 0.7.x | **Conductor** | Workflow automation, pipelines, rule-based logic |
|
||||
| 1.0.x | **Summit** | Enterprise features, multi-tenancy, RBAC |
|
||||
| 1.1.x | **Bridge** | Collaboration, sharing, analytics |
|
||||
| 2.0.x | **Horizon** | On-premise AI, platform expansion |
|
||||
|
||||
## Current Status (v0.5.0 "Foundation")
|
||||
|
||||
### Core Features ✅
|
||||
- Multi-provider document storage (Dropbox, Google Drive, OneDrive, Nextcloud, S3, etc.)
|
||||
@@ -23,7 +36,7 @@ DocuElevate aims to be the premier open-source intelligent document processing p
|
||||
- Celery-based async task processing
|
||||
- OAuth2 authentication via Authentik with admin group support
|
||||
|
||||
## Short-term Goals (Q1-Q2 2026) - v0.4.x to v0.5.x
|
||||
## Short-term Goals (Q1-Q2 2026) - v0.4.x to v0.5.x "Foundation"
|
||||
|
||||
### Quality & Stability 🎯
|
||||
- **Test Coverage** (High Priority)
|
||||
@@ -53,7 +66,7 @@ DocuElevate aims to be the premier open-source intelligent document processing p
|
||||
- [x] Integrate Docker builds with releases
|
||||
|
||||
### Features - v0.4.0
|
||||
- **Enhanced Search & Filtering**
|
||||
- **Enhanced Search & Filtering** → _preparing for v0.6.0 "Clarity"_
|
||||
- [ ] Full-text search across documents
|
||||
- [ ] Advanced filtering by metadata, tags, date ranges
|
||||
- [ ] Saved search queries
|
||||
@@ -67,8 +80,8 @@ DocuElevate aims to be the premier open-source intelligent document processing p
|
||||
- [ ] Progress indicators for long-running tasks
|
||||
- [ ] Real-time notifications via WebSocket
|
||||
|
||||
### Features - v0.5.0
|
||||
- **Workflow Automation**
|
||||
### Features - v0.5.0 "Foundation"
|
||||
- **Workflow Automation** → _evolving into v0.7.0 "Conductor"_
|
||||
- [ ] Custom processing pipelines
|
||||
- [ ] Conditional routing based on document type
|
||||
- [ ] Scheduled batch processing
|
||||
@@ -82,9 +95,9 @@ DocuElevate aims to be the premier open-source intelligent document processing p
|
||||
- [ ] Automatic duplicate detection
|
||||
- [ ] Intelligent document splitting
|
||||
|
||||
## Medium-term Goals (Q3-Q4 2026) - v1.0.x
|
||||
## Medium-term Goals (Q3-Q4 2026) - v1.0.x "Summit"
|
||||
|
||||
### Enterprise Features - v1.0.0
|
||||
### Enterprise Features - v1.0.0 "Summit"
|
||||
- **Multi-tenancy**
|
||||
- [ ] Organization/team management
|
||||
- [ ] Role-based access control (RBAC)
|
||||
@@ -106,7 +119,7 @@ DocuElevate aims to be the premier open-source intelligent document processing p
|
||||
- [ ] Custom webhook receivers
|
||||
- [ ] GraphQL API
|
||||
|
||||
### Features - v1.1.0
|
||||
### Features - v1.1.0 "Bridge"
|
||||
- **Collaboration**
|
||||
- [ ] Document sharing with expiring links
|
||||
- [ ] Comments and annotations
|
||||
@@ -121,7 +134,7 @@ DocuElevate aims to be the premier open-source intelligent document processing p
|
||||
- [ ] Cost analysis per provider
|
||||
- [ ] Export reports (PDF, CSV, Excel)
|
||||
|
||||
## Long-term Goals (2027+) - v2.0+
|
||||
## Long-term Goals (2027+) - v2.0+ "Horizon"
|
||||
|
||||
### Strategic Initiatives
|
||||
- **On-Premise AI Models**
|
||||
|
||||
+6
-6
@@ -1,10 +1,10 @@
|
||||
DocuElevate Build Information
|
||||
==============================
|
||||
Version: 0.67.2
|
||||
Build Date: 2026-03-01T17:23:25Z
|
||||
Git Commit: 0134ed37d5c602faf5b10cc6a7229263ba2f6aa1
|
||||
Git Short SHA: 0134ed3
|
||||
Version: 0.145.2
|
||||
Build Date: 2026-03-15T21:39:26Z
|
||||
Git Commit: 237af31f5fe598cfc3c08f2bbba79b3d0925787e
|
||||
Git Short SHA: 237af31
|
||||
Git Branch: main
|
||||
Commit Date: 2026-03-01T18:23:06+01:00
|
||||
Build Timestamp: 2026-03-01T17:23:25Z
|
||||
Commit Date: 2026-03-15T22:39:04+01:00
|
||||
Build Timestamp: 2026-03-15T21:39:26Z
|
||||
==============================
|
||||
|
||||
@@ -861,3 +861,21 @@ Security headers implementation is complete and production-ready. The middleware
|
||||
---
|
||||
|
||||
**Next Audit Due:** 2026-05-07 (Quarterly)
|
||||
|
||||
## Per-User IMAP Account Passwords (Added 2026-03-08)
|
||||
|
||||
### Known Limitation: Plain-text Password Storage
|
||||
|
||||
IMAP account passwords in the `user_imap_accounts` table are stored in plain text in the database.
|
||||
|
||||
**Risk:** Anyone with direct database access (DBA, backup access) can read IMAP credentials for all users.
|
||||
|
||||
**Mitigations in place:**
|
||||
- Database itself should be protected with appropriate OS-level file permissions (SQLite) or network ACLs (PostgreSQL/MySQL).
|
||||
- Passwords are never returned in API responses (the `_to_response` serialiser omits them).
|
||||
- Only the account owner can read or update their own accounts (ownership enforced at the API layer).
|
||||
- Passwords are never logged.
|
||||
|
||||
**Future improvement:** Encrypt IMAP passwords at rest using `cryptography.fernet` (symmetric encryption with the app's `SESSION_SECRET` as key material). This is tracked as a TODO item in `app/api/imap_accounts.py` and should be implemented before this feature is used in high-security environments.
|
||||
|
||||
**Recommended admin action:** Use app-specific passwords (Gmail, Outlook) rather than account passwords where possible, so that compromised IMAP credentials can be revoked without affecting the user's primary account.
|
||||
|
||||
@@ -6,23 +6,48 @@ import logging
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from app.api.admin_users import router as admin_users_router
|
||||
from app.api.api_tokens import router as api_tokens_router
|
||||
from app.api.audit_logs import router as audit_logs_router
|
||||
from app.api.azure import router as azure_router
|
||||
from app.api.backup import router as backup_router
|
||||
from app.api.billing import router as billing_router
|
||||
from app.api.compliance import router as compliance_router
|
||||
from app.api.database import router as database_router
|
||||
from app.api.diagnostic import router as diagnostic_router
|
||||
from app.api.dropbox import router as dropbox_router
|
||||
from app.api.duplicates import router as duplicates_router
|
||||
from app.api.files import router as files_router
|
||||
from app.api.google_drive import router as google_drive_router
|
||||
from app.api.i18n import router as i18n_router
|
||||
from app.api.imap_accounts import router as imap_accounts_router
|
||||
from app.api.imap_profiles import router as imap_profiles_router
|
||||
from app.api.integrations import router as integrations_router
|
||||
from app.api.logs import router as logs_router
|
||||
from app.api.mobile import router as mobile_router
|
||||
from app.api.notifications import router as notifications_router
|
||||
from app.api.onboarding import router as onboarding_router
|
||||
from app.api.onedrive import router as onedrive_router
|
||||
from app.api.openai import router as openai_router
|
||||
from app.api.pipelines import router as pipelines_router
|
||||
from app.api.plans import router as plans_router
|
||||
from app.api.process import router as process_router
|
||||
from app.api.profile import router as profile_router
|
||||
from app.api.queue import router as queue_router
|
||||
from app.api.routing_rules import router as routing_rules_router
|
||||
from app.api.saved_searches import router as saved_searches_router
|
||||
from app.api.scheduled_jobs import router as scheduled_jobs_router
|
||||
from app.api.search import router as search_router
|
||||
from app.api.settings import router as settings_router
|
||||
from app.api.shared_links import public_router as shared_links_public_router
|
||||
from app.api.shared_links import router as shared_links_router
|
||||
from app.api.similarity import router as similarity_router
|
||||
from app.api.subscriptions import router as subscriptions_router
|
||||
from app.api.url_upload import router as url_upload_router
|
||||
|
||||
# Import all the individual routers
|
||||
from app.api.user import router as user_router
|
||||
from app.api.webhooks import router as webhooks_router
|
||||
|
||||
# Set up logging
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -31,7 +56,10 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
# Include all the routers
|
||||
router.include_router(admin_users_router)
|
||||
router.include_router(api_tokens_router)
|
||||
router.include_router(user_router)
|
||||
router.include_router(backup_router)
|
||||
router.include_router(files_router)
|
||||
router.include_router(process_router)
|
||||
router.include_router(diagnostic_router)
|
||||
@@ -46,3 +74,25 @@ router.include_router(url_upload_router)
|
||||
router.include_router(search_router)
|
||||
router.include_router(queue_router)
|
||||
router.include_router(saved_searches_router)
|
||||
router.include_router(similarity_router)
|
||||
router.include_router(shared_links_router)
|
||||
router.include_router(shared_links_public_router)
|
||||
router.include_router(duplicates_router)
|
||||
router.include_router(webhooks_router)
|
||||
router.include_router(database_router)
|
||||
router.include_router(subscriptions_router)
|
||||
router.include_router(plans_router)
|
||||
router.include_router(onboarding_router)
|
||||
router.include_router(billing_router)
|
||||
router.include_router(pipelines_router)
|
||||
router.include_router(profile_router)
|
||||
router.include_router(routing_rules_router)
|
||||
router.include_router(imap_accounts_router)
|
||||
router.include_router(imap_profiles_router)
|
||||
router.include_router(integrations_router)
|
||||
router.include_router(notifications_router)
|
||||
router.include_router(scheduled_jobs_router)
|
||||
router.include_router(audit_logs_router)
|
||||
router.include_router(i18n_router)
|
||||
router.include_router(mobile_router)
|
||||
router.include_router(compliance_router)
|
||||
|
||||
@@ -0,0 +1,668 @@
|
||||
"""API endpoints for admin user management.
|
||||
|
||||
Provides CRUD operations for user profiles and aggregate statistics so that
|
||||
administrators can inspect, configure, and manage users in multi-user mode.
|
||||
Also provides endpoints for admins to create and manage local (email/password)
|
||||
user accounts directly, without requiring email verification.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import FileRecord, LocalUser, UserProfile
|
||||
from app.utils.local_auth import generate_token, hash_password, send_password_reset_email
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/admin/users", tags=["admin-users"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin. Raises 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class UserProfileUpsert(BaseModel):
|
||||
"""Body for creating or updating a user profile."""
|
||||
|
||||
display_name: str | None = Field(default=None, max_length=255, description="Human-readable display name")
|
||||
daily_upload_limit: int | None = Field(
|
||||
default=None, ge=0, description="Per-user daily upload cap; null = use global default"
|
||||
)
|
||||
notes: str | None = Field(default=None, max_length=4096, description="Admin notes about this user")
|
||||
is_blocked: bool = Field(default=False, description="Block this user from uploading")
|
||||
subscription_tier: str | None = Field(
|
||||
default="free",
|
||||
description="Subscription tier: free | starter | professional | business",
|
||||
)
|
||||
subscription_billing_cycle: str = Field(default="monthly", pattern="^(monthly|yearly)$")
|
||||
subscription_period_start: datetime | None = None
|
||||
allow_overage: bool = False
|
||||
is_complimentary: bool = Field(
|
||||
default=False,
|
||||
description="When True the user is on a complimentary (uncharged) plan — they keep all tier "
|
||||
"quota benefits but are never billed via Stripe.",
|
||||
)
|
||||
|
||||
|
||||
class PaymentIssueBody(BaseModel):
|
||||
"""Body for reporting a payment issue for a user."""
|
||||
|
||||
issue: str = Field(..., min_length=1, max_length=2048, description="Description of the payment issue")
|
||||
|
||||
|
||||
class UserProfileResponse(BaseModel):
|
||||
"""Response schema for a user profile record."""
|
||||
|
||||
id: int
|
||||
user_id: str
|
||||
display_name: str | None
|
||||
daily_upload_limit: int | None
|
||||
notes: str | None
|
||||
is_blocked: bool
|
||||
subscription_tier: str | None
|
||||
subscription_billing_cycle: str
|
||||
subscription_period_start: str | None
|
||||
allow_overage: bool
|
||||
is_complimentary: bool
|
||||
created_at: str | None
|
||||
updated_at: str | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class UserSummary(BaseModel):
|
||||
"""Per-user summary combining profile data with document statistics."""
|
||||
|
||||
user_id: str
|
||||
display_name: str | None
|
||||
daily_upload_limit: int | None
|
||||
notes: str | None
|
||||
is_blocked: bool
|
||||
subscription_tier: str | None
|
||||
subscription_billing_cycle: str | None
|
||||
subscription_period_start: str | None
|
||||
allow_overage: bool
|
||||
is_complimentary: bool
|
||||
profile_id: int | None
|
||||
document_count: int
|
||||
last_upload: str | None
|
||||
|
||||
|
||||
class LocalUserCreate(BaseModel):
|
||||
"""Body for admin-creating a local (email/password) user account."""
|
||||
|
||||
email: str = Field(..., max_length=255, description="Email address for the new user")
|
||||
username: str = Field(..., min_length=3, max_length=64, pattern=r"^[a-zA-Z0-9_-]+$")
|
||||
display_name: str | None = Field(default=None, max_length=255)
|
||||
password: str = Field(..., min_length=8, max_length=128)
|
||||
is_admin: bool = Field(default=False, description="Grant admin privileges")
|
||||
|
||||
|
||||
class LocalUserUpdate(BaseModel):
|
||||
"""Body for admin-updating a local (email/password) user account."""
|
||||
|
||||
email: str | None = Field(default=None, max_length=255, description="New email address")
|
||||
display_name: str | None = Field(default=None, max_length=255, description="New display name")
|
||||
is_admin: bool | None = Field(default=None, description="Grant or revoke admin privileges")
|
||||
is_active: bool | None = Field(default=None, description="Activate or deactivate the account")
|
||||
|
||||
|
||||
class LocalUserSetPassword(BaseModel):
|
||||
"""Body for admin setting a temporary password for a local user."""
|
||||
|
||||
password: str = Field(..., min_length=8, max_length=128, description="New temporary password")
|
||||
|
||||
|
||||
class LocalUserResponse(BaseModel):
|
||||
"""Summary of a local user account."""
|
||||
|
||||
id: int
|
||||
email: str
|
||||
username: str
|
||||
display_name: str | None
|
||||
is_active: bool
|
||||
is_admin: bool
|
||||
created_at: str | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_or_none(db: Session, user_id: str) -> UserProfile | None:
|
||||
"""Return the UserProfile row for *user_id*, or None if it doesn't exist."""
|
||||
return db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
|
||||
|
||||
def _profile_to_dict(profile: UserProfile) -> dict[str, Any]:
|
||||
return {
|
||||
"id": profile.id,
|
||||
"user_id": profile.user_id,
|
||||
"display_name": profile.display_name,
|
||||
"daily_upload_limit": profile.daily_upload_limit,
|
||||
"notes": profile.notes,
|
||||
"is_blocked": profile.is_blocked,
|
||||
"subscription_tier": profile.subscription_tier or "free",
|
||||
"subscription_billing_cycle": profile.subscription_billing_cycle or "monthly",
|
||||
"subscription_period_start": profile.subscription_period_start.isoformat()
|
||||
if profile.subscription_period_start
|
||||
else None,
|
||||
"allow_overage": bool(profile.allow_overage),
|
||||
"is_complimentary": bool(profile.is_complimentary),
|
||||
"created_at": profile.created_at.isoformat() if profile.created_at else None,
|
||||
"updated_at": profile.updated_at.isoformat() if profile.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", summary="List all known users with statistics")
|
||||
def list_users(
|
||||
db: DbSession,
|
||||
_admin: AdminUser,
|
||||
q: str = Query("", description="Filter by user_id substring (case-insensitive)"),
|
||||
page: int = Query(1, ge=1, description="Page number"),
|
||||
per_page: int = Query(25, ge=1, le=100, description="Items per page"),
|
||||
) -> dict[str, Any]:
|
||||
"""Return every distinct user_id that has at least one document or an explicit profile,
|
||||
enriched with aggregate document statistics and the admin-managed profile.
|
||||
|
||||
Supports substring filtering (``q``) and pagination.
|
||||
"""
|
||||
# 1. Collect every distinct owner_id from documents
|
||||
doc_stats_query = (
|
||||
db.query(
|
||||
FileRecord.owner_id.label("user_id"),
|
||||
func.count(FileRecord.id).label("doc_count"),
|
||||
func.max(FileRecord.created_at).label("last_upload"),
|
||||
)
|
||||
.filter(FileRecord.owner_id.isnot(None))
|
||||
.group_by(FileRecord.owner_id)
|
||||
)
|
||||
|
||||
# 2. Collect all user_ids that have explicit profiles (may not have docs yet)
|
||||
profile_query = db.query(UserProfile)
|
||||
|
||||
# Build a unified set of user_ids
|
||||
doc_rows = {row.user_id: row for row in doc_stats_query.all()}
|
||||
profile_rows = {p.user_id: p for p in profile_query.all()}
|
||||
|
||||
all_user_ids = set(doc_rows.keys()) | set(profile_rows.keys())
|
||||
|
||||
# Apply optional substring filter
|
||||
if q.strip():
|
||||
q_lower = q.strip().lower()
|
||||
all_user_ids = {uid for uid in all_user_ids if q_lower in uid.lower()}
|
||||
|
||||
# Sort and paginate
|
||||
sorted_ids = sorted(all_user_ids)
|
||||
total = len(sorted_ids)
|
||||
start = (page - 1) * per_page
|
||||
page_ids = sorted_ids[start : start + per_page]
|
||||
|
||||
users: list[dict[str, Any]] = []
|
||||
for uid in page_ids:
|
||||
doc_row = doc_rows.get(uid)
|
||||
profile = profile_rows.get(uid)
|
||||
users.append(
|
||||
{
|
||||
"user_id": uid,
|
||||
"display_name": profile.display_name if profile else None,
|
||||
"daily_upload_limit": profile.daily_upload_limit if profile else None,
|
||||
"notes": profile.notes if profile else None,
|
||||
"is_blocked": profile.is_blocked if profile else False,
|
||||
"subscription_tier": (profile.subscription_tier or "free") if profile else "free",
|
||||
"subscription_billing_cycle": (profile.subscription_billing_cycle or "monthly")
|
||||
if profile
|
||||
else "monthly",
|
||||
"subscription_period_start": profile.subscription_period_start.isoformat()
|
||||
if (profile and profile.subscription_period_start)
|
||||
else None,
|
||||
"allow_overage": bool(profile.allow_overage) if profile else False,
|
||||
"is_complimentary": bool(profile.is_complimentary) if profile else False,
|
||||
"profile_id": profile.id if profile else None,
|
||||
"document_count": doc_row.doc_count if doc_row else 0,
|
||||
"last_upload": doc_row.last_upload.isoformat() if (doc_row and doc_row.last_upload) else None,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"users": users,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"per_page": per_page,
|
||||
"pages": max(1, (total + per_page - 1) // per_page),
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Local user management (admin-only)
|
||||
# ---------------------------------------------------------------------------
|
||||
# NOTE: These routes MUST be defined before /{user_id:path} to avoid being
|
||||
# swallowed by the catch-all path parameter.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/local", summary="List all local (email/password) user accounts")
|
||||
def list_local_users(db: DbSession, _admin: AdminUser) -> list[dict[str, Any]]:
|
||||
"""Return every local user account with basic metadata."""
|
||||
users = db.query(LocalUser).order_by(LocalUser.created_at.desc()).all()
|
||||
return [
|
||||
{
|
||||
"id": u.id,
|
||||
"email": u.email,
|
||||
"username": u.username,
|
||||
"display_name": u.display_name,
|
||||
"is_active": u.is_active,
|
||||
"is_admin": u.is_admin,
|
||||
"created_at": u.created_at.isoformat() if u.created_at else None,
|
||||
}
|
||||
for u in users
|
||||
]
|
||||
|
||||
|
||||
@router.post("/local", status_code=status.HTTP_201_CREATED, summary="Create a local user account")
|
||||
def create_local_user(body: LocalUserCreate, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Create a new local (email/password) user account.
|
||||
|
||||
The account is immediately active — no email verification is required when
|
||||
created by an administrator. A matching UserProfile row is also created.
|
||||
|
||||
Raises:
|
||||
409: Email or username already registered.
|
||||
"""
|
||||
if db.query(LocalUser).filter(LocalUser.email == body.email).first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Email already registered.")
|
||||
if db.query(LocalUser).filter(LocalUser.username == body.username).first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Username already taken.")
|
||||
|
||||
user = LocalUser(
|
||||
email=body.email,
|
||||
username=body.username,
|
||||
display_name=body.display_name,
|
||||
hashed_password=hash_password(body.password),
|
||||
is_active=True,
|
||||
is_admin=body.is_admin,
|
||||
)
|
||||
db.add(user)
|
||||
|
||||
# Ensure a UserProfile exists for the new user
|
||||
if not db.query(UserProfile).filter(UserProfile.user_id == body.email).first():
|
||||
db.add(UserProfile(user_id=body.email, display_name=body.display_name or body.username))
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Admin created local user account: %s", body.email)
|
||||
return {
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
"username": user.username,
|
||||
"display_name": user.display_name,
|
||||
"is_active": user.is_active,
|
||||
"is_admin": user.is_admin,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/local/{local_user_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="Delete a local user account",
|
||||
)
|
||||
def delete_local_user(local_user_id: int, db: DbSession, _admin: AdminUser) -> None:
|
||||
"""Delete a local user account by its numeric ID.
|
||||
|
||||
The associated UserProfile is also removed. Documents owned by this user
|
||||
are **not** deleted.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.id == local_user_id).first()
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Local user not found.")
|
||||
|
||||
# Remove associated profile if present
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user.email).first()
|
||||
if profile:
|
||||
db.delete(profile)
|
||||
|
||||
try:
|
||||
db.delete(user)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Admin deleted local user account: %s", user.email)
|
||||
|
||||
|
||||
@router.patch("/local/{local_user_id}", summary="Update a local user account")
|
||||
def update_local_user(local_user_id: int, body: LocalUserUpdate, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Update the email address, display name, admin flag, or active status of a local user account.
|
||||
|
||||
Only fields explicitly provided (non-None) are modified. If the email is changed
|
||||
the associated UserProfile row is also updated to keep ``user_id`` in sync.
|
||||
|
||||
Raises:
|
||||
404: Local user not found.
|
||||
409: The new email is already taken by another account.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.id == local_user_id).first()
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Local user not found.")
|
||||
|
||||
old_email = user.email
|
||||
|
||||
if body.email is not None and body.email != user.email:
|
||||
if db.query(LocalUser).filter(LocalUser.email == body.email, LocalUser.id != local_user_id).first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Email already registered.")
|
||||
user.email = body.email
|
||||
|
||||
if body.display_name is not None:
|
||||
# Normalise empty string to None so that clearing the field removes the display name
|
||||
user.display_name = body.display_name or None
|
||||
|
||||
if body.is_admin is not None:
|
||||
user.is_admin = body.is_admin
|
||||
|
||||
if body.is_active is not None:
|
||||
user.is_active = body.is_active
|
||||
|
||||
try:
|
||||
db.flush()
|
||||
# Keep UserProfile.user_id in sync when email changes
|
||||
if body.email is not None and body.email != old_email:
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == old_email).first()
|
||||
if profile:
|
||||
profile.user_id = body.email
|
||||
db.commit()
|
||||
db.refresh(user)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Admin updated local user %s (id=%d)", user.email, user.id)
|
||||
return {
|
||||
"id": user.id,
|
||||
"email": user.email,
|
||||
"username": user.username,
|
||||
"display_name": user.display_name,
|
||||
"is_active": user.is_active,
|
||||
"is_admin": user.is_admin,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/local/{local_user_id}/send-password-reset",
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Send a password reset email to a local user",
|
||||
)
|
||||
def admin_send_password_reset(local_user_id: int, request: Request, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Generate a password reset token and email the reset link to the local user.
|
||||
|
||||
This is a last-resort tool for admins to help users who are locked out.
|
||||
Returns ``{"sent": true}`` on success and ``{"sent": false, "reason": "..."}`` when
|
||||
SMTP is not configured or sending fails.
|
||||
|
||||
Raises:
|
||||
404: Local user not found.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.id == local_user_id).first()
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Local user not found.")
|
||||
|
||||
if not settings.email_host:
|
||||
logger.warning("Admin requested password reset for %s but SMTP is not configured", user.email)
|
||||
return {"sent": False, "reason": "SMTP is not configured on this server."}
|
||||
|
||||
token = generate_token()
|
||||
user.password_reset_token = token
|
||||
user.password_reset_sent_at = datetime.now(tz=timezone.utc)
|
||||
db.commit()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
try:
|
||||
send_password_reset_email(user.email, user.username, token, base_url)
|
||||
except Exception as exc:
|
||||
logger.warning("Admin-triggered password reset email failed for %s: %s", user.email, exc)
|
||||
return {"sent": False, "reason": str(exc)}
|
||||
|
||||
logger.info("[SECURITY] ADMIN_PASSWORD_RESET_EMAIL user=%s admin=%s", user.email, _admin.get("email", "unknown"))
|
||||
return {"sent": True, "email": user.email}
|
||||
|
||||
|
||||
@router.post(
|
||||
"/local/{local_user_id}/set-password",
|
||||
status_code=status.HTTP_200_OK,
|
||||
summary="Set a temporary password for a local user account",
|
||||
)
|
||||
def admin_set_password(
|
||||
local_user_id: int, body: LocalUserSetPassword, db: DbSession, _admin: AdminUser
|
||||
) -> dict[str, Any]:
|
||||
"""Directly set a new password for a local user without requiring an email token.
|
||||
|
||||
Use this as a last resort when email delivery is unavailable. The user
|
||||
should be advised to change their password after logging in.
|
||||
|
||||
Raises:
|
||||
404: Local user not found.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.id == local_user_id).first()
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Local user not found.")
|
||||
|
||||
user.hashed_password = hash_password(body.password)
|
||||
# Clear any outstanding reset tokens and activate the account so the user
|
||||
# can log in immediately after an admin sets their password.
|
||||
user.password_reset_token = None
|
||||
user.password_reset_sent_at = None
|
||||
user.is_active = True
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("[SECURITY] ADMIN_SET_PASSWORD user=%s admin=%s", user.email, _admin.get("email", "unknown"))
|
||||
return {"updated": True, "email": user.email}
|
||||
|
||||
|
||||
@router.get("/{user_id:path}", summary="Get details for a single user")
|
||||
def get_user(user_id: str, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Return profile and document statistics for a specific user."""
|
||||
doc_count = db.query(func.count(FileRecord.id)).filter(FileRecord.owner_id == user_id).scalar() or 0
|
||||
last_row = (
|
||||
db.query(FileRecord.created_at)
|
||||
.filter(FileRecord.owner_id == user_id)
|
||||
.order_by(FileRecord.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
last_upload = last_row[0].isoformat() if last_row and last_row[0] else None
|
||||
|
||||
profile = _get_or_none(db, user_id)
|
||||
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"display_name": profile.display_name if profile else None,
|
||||
"daily_upload_limit": profile.daily_upload_limit if profile else None,
|
||||
"notes": profile.notes if profile else None,
|
||||
"is_blocked": profile.is_blocked if profile else False,
|
||||
"subscription_tier": (profile.subscription_tier or "free") if profile else "free",
|
||||
"subscription_billing_cycle": (profile.subscription_billing_cycle or "monthly") if profile else "monthly",
|
||||
"subscription_period_start": profile.subscription_period_start.isoformat()
|
||||
if (profile and profile.subscription_period_start)
|
||||
else None,
|
||||
"allow_overage": bool(profile.allow_overage) if profile else False,
|
||||
"is_complimentary": bool(profile.is_complimentary) if profile else False,
|
||||
"profile_id": profile.id if profile else None,
|
||||
"document_count": doc_count,
|
||||
"last_upload": last_upload,
|
||||
"profile": _profile_to_dict(profile) if profile else None,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/{user_id:path}", summary="Create or update a user profile")
|
||||
def upsert_user_profile(
|
||||
user_id: str,
|
||||
body: UserProfileUpsert,
|
||||
db: DbSession,
|
||||
_admin: AdminUser,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new profile or update an existing one for *user_id*.
|
||||
|
||||
Returns the persisted profile.
|
||||
"""
|
||||
profile = _get_or_none(db, user_id)
|
||||
if profile is None:
|
||||
profile = UserProfile(user_id=user_id)
|
||||
db.add(profile)
|
||||
|
||||
old_tier = (profile.subscription_tier or "free") if profile.id else None # None means brand-new profile
|
||||
profile.display_name = body.display_name
|
||||
profile.daily_upload_limit = body.daily_upload_limit
|
||||
profile.notes = body.notes
|
||||
profile.is_blocked = body.is_blocked
|
||||
profile.subscription_billing_cycle = body.subscription_billing_cycle
|
||||
profile.subscription_period_start = body.subscription_period_start
|
||||
profile.allow_overage = body.allow_overage
|
||||
profile.is_complimentary = body.is_complimentary
|
||||
tier_changed = False
|
||||
new_tier: str | None = None
|
||||
if body.subscription_tier is not None:
|
||||
from app.utils.subscription import TIERS
|
||||
|
||||
if body.subscription_tier not in TIERS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Invalid subscription_tier '{body.subscription_tier}'. Valid values: {list(TIERS.keys())}",
|
||||
)
|
||||
# Detect a real change only for existing profiles (old_tier is not None)
|
||||
if old_tier is not None and old_tier != body.subscription_tier:
|
||||
tier_changed = True
|
||||
new_tier = body.subscription_tier
|
||||
profile.subscription_tier = body.subscription_tier
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Admin upserted profile for user %s", user_id)
|
||||
|
||||
# Notify admins and fire webhook when plan is changed by an admin
|
||||
if tier_changed and new_tier is not None:
|
||||
try:
|
||||
from app.utils.notification import notify_plan_changed
|
||||
from app.utils.webhook import dispatch_webhook_event
|
||||
|
||||
notify_plan_changed(user_id, old_tier=old_tier, new_tier=new_tier, changed_by="admin") # type: ignore[arg-type]
|
||||
dispatch_webhook_event(
|
||||
"user.plan_changed",
|
||||
{
|
||||
"user_id": user_id,
|
||||
"old_tier": old_tier,
|
||||
"new_tier": new_tier,
|
||||
"billing_cycle": body.subscription_billing_cycle,
|
||||
"changed_by": "admin",
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to send plan-change notification/webhook for user %s", user_id)
|
||||
|
||||
return _profile_to_dict(profile)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/{user_id:path}/payment-issue", status_code=status.HTTP_200_OK, summary="Report a payment issue for a user"
|
||||
)
|
||||
def report_payment_issue(user_id: str, body: PaymentIssueBody, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Notify admins and fire a webhook for a payment issue reported against *user_id*.
|
||||
|
||||
The user profile must exist. Use this endpoint when a payment processor
|
||||
webhook or manual review identifies a billing problem (e.g. failed charge,
|
||||
expired card, disputed transaction).
|
||||
|
||||
Returns the user profile dict alongside an acknowledgement flag.
|
||||
"""
|
||||
profile = _get_or_none(db, user_id)
|
||||
if not profile:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User profile not found")
|
||||
|
||||
logger.warning("Payment issue reported for user %s: %s", user_id, body.issue)
|
||||
|
||||
try:
|
||||
from app.utils.notification import notify_payment_issue
|
||||
from app.utils.webhook import dispatch_webhook_event
|
||||
|
||||
notify_payment_issue(user_id, issue=body.issue)
|
||||
dispatch_webhook_event(
|
||||
"user.payment_issue",
|
||||
{
|
||||
"user_id": user_id,
|
||||
"issue": body.issue,
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to send payment-issue notification/webhook for user %s", user_id)
|
||||
|
||||
return {"acknowledged": True, "user_id": user_id, "profile": _profile_to_dict(profile)}
|
||||
|
||||
|
||||
@router.delete("/{user_id:path}", status_code=status.HTTP_204_NO_CONTENT, summary="Delete a user profile")
|
||||
def delete_user_profile(user_id: str, db: DbSession, _admin: AdminUser) -> None:
|
||||
"""Delete the admin-managed profile for *user_id*.
|
||||
|
||||
Documents owned by this user are **not** removed; only the profile record
|
||||
is deleted. To reassign or purge documents use the files API.
|
||||
"""
|
||||
profile = _get_or_none(db, user_id)
|
||||
if not profile:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User profile not found")
|
||||
|
||||
try:
|
||||
db.delete(profile)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Admin deleted profile for user %s", user_id)
|
||||
@@ -0,0 +1,224 @@
|
||||
"""API endpoints for managing personal API tokens.
|
||||
|
||||
Provides CRUD operations so users can create, list, and revoke tokens
|
||||
that grant programmatic access to the DocuElevate API (e.g. webhook
|
||||
uploads, scripted integrations).
|
||||
|
||||
Tokens use ``secrets.token_urlsafe`` from the Python standard library
|
||||
(no extra dependencies) and are prefixed with ``de_`` for easy
|
||||
identification. Only a PBKDF2-HMAC-SHA256 hash is persisted; the
|
||||
plaintext is returned exactly once at creation time.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import ApiToken
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/api-tokens", tags=["api-tokens"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Prefix prepended to every generated token for easy identification.
|
||||
TOKEN_PREFIX = "de_"
|
||||
#: Number of random bytes for the token body (32 → 43 URL-safe chars).
|
||||
TOKEN_BYTES = 32
|
||||
#: PBKDF2 iteration count for hashing API tokens.
|
||||
TOKEN_HASH_ITERATIONS = 100_000
|
||||
#: PBKDF2 salt for API token hashing (not secret, but fixed for determinism).
|
||||
TOKEN_HASH_SALT = b"api-token-v1"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def generate_api_token() -> str:
|
||||
"""Generate a new API token with the ``de_`` prefix.
|
||||
|
||||
Returns:
|
||||
A URL-safe random token string, e.g. ``de_Ab3xY…``.
|
||||
"""
|
||||
return TOKEN_PREFIX + secrets.token_urlsafe(TOKEN_BYTES)
|
||||
|
||||
|
||||
def hash_token(token: str) -> str:
|
||||
"""Return a PBKDF2-HMAC-SHA256 hex digest of *token*.
|
||||
|
||||
Args:
|
||||
token: The plaintext API token.
|
||||
|
||||
Returns:
|
||||
64-character lowercase hex string.
|
||||
"""
|
||||
dk = hashlib.pbkdf2_hmac(
|
||||
"sha256",
|
||||
token.encode("utf-8"),
|
||||
TOKEN_HASH_SALT,
|
||||
TOKEN_HASH_ITERATIONS,
|
||||
)
|
||||
return dk.hex()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TokenCreate(BaseModel):
|
||||
"""Schema for creating a new API token."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255, description="Human-readable label for the token")
|
||||
|
||||
|
||||
class TokenResponse(BaseModel):
|
||||
"""Schema returned when listing tokens (plaintext is never included)."""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
token_prefix: str
|
||||
is_active: bool
|
||||
last_used_at: datetime | None
|
||||
last_used_ip: str | None
|
||||
created_at: datetime | None
|
||||
revoked_at: datetime | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class TokenCreatedResponse(TokenResponse):
|
||||
"""Schema returned once at creation time — includes the full plaintext token."""
|
||||
|
||||
token: str = Field(..., description="The full API token. Store it securely — it will not be shown again.")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED, response_model=TokenCreatedResponse)
|
||||
async def create_token(
|
||||
body: TokenCreate,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new personal API token.
|
||||
|
||||
The full token is returned **only once** in the response. Subsequent
|
||||
``GET`` requests will only show the prefix for identification.
|
||||
"""
|
||||
plaintext = generate_api_token()
|
||||
token_hash_value = hash_token(plaintext)
|
||||
prefix = plaintext[:12] # "de_" prefix + 9 random chars = 12 chars total
|
||||
|
||||
db_token = ApiToken(
|
||||
owner_id=owner_id,
|
||||
name=body.name,
|
||||
token_hash=token_hash_value,
|
||||
token_prefix=prefix,
|
||||
)
|
||||
try:
|
||||
db.add(db_token)
|
||||
db.commit()
|
||||
db.refresh(db_token)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("API token created: id=%s owner=%s name=%r", db_token.id, owner_id, body.name)
|
||||
|
||||
return {
|
||||
"id": db_token.id,
|
||||
"name": db_token.name,
|
||||
"token_prefix": db_token.token_prefix,
|
||||
"is_active": db_token.is_active,
|
||||
"last_used_at": db_token.last_used_at,
|
||||
"last_used_ip": db_token.last_used_ip,
|
||||
"created_at": db_token.created_at,
|
||||
"revoked_at": db_token.revoked_at,
|
||||
"token": plaintext,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/", response_model=list[TokenResponse])
|
||||
async def list_tokens(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all API tokens for the authenticated user."""
|
||||
tokens = db.query(ApiToken).filter(ApiToken.owner_id == owner_id).order_by(ApiToken.created_at.desc()).all()
|
||||
return [
|
||||
{
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"token_prefix": t.token_prefix,
|
||||
"is_active": t.is_active,
|
||||
"last_used_at": t.last_used_at,
|
||||
"last_used_ip": t.last_used_ip,
|
||||
"created_at": t.created_at,
|
||||
"revoked_at": t.revoked_at,
|
||||
}
|
||||
for t in tokens
|
||||
]
|
||||
|
||||
|
||||
@router.delete("/{token_id}", status_code=status.HTTP_200_OK)
|
||||
async def revoke_token(
|
||||
token_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Revoke (soft-delete) an API token.
|
||||
|
||||
The token row is kept for audit purposes but marked inactive with a
|
||||
``revoked_at`` timestamp.
|
||||
"""
|
||||
db_token = db.query(ApiToken).filter(ApiToken.id == token_id, ApiToken.owner_id == owner_id).first()
|
||||
if not db_token:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Token not found")
|
||||
|
||||
if not db_token.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Token is already revoked")
|
||||
|
||||
try:
|
||||
db_token.is_active = False
|
||||
db_token.revoked_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("API token revoked: id=%s owner=%s", token_id, owner_id)
|
||||
return {"detail": "Token revoked"}
|
||||
@@ -0,0 +1,117 @@
|
||||
"""
|
||||
Audit log REST API endpoints.
|
||||
|
||||
Provides read-only access to the comprehensive audit log for admin users.
|
||||
Events are append-only — there are no update or delete endpoints.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.utils.audit_service import count_events, query_events
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
@router.get("/audit-logs")
|
||||
@require_login
|
||||
async def list_audit_logs(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
action: Annotated[str | None, Query(description="Filter by action (exact match)")] = None,
|
||||
user: Annotated[str | None, Query(description="Filter by username")] = None,
|
||||
resource_type: Annotated[str | None, Query(description="Filter by resource type")] = None,
|
||||
severity: Annotated[str | None, Query(description="Filter by severity level")] = None,
|
||||
since: Annotated[datetime | None, Query(description="Only events at or after this ISO-8601 timestamp")] = None,
|
||||
until: Annotated[datetime | None, Query(description="Only events at or before this ISO-8601 timestamp")] = None,
|
||||
limit: Annotated[int, Query(ge=1, le=500, description="Max rows to return")] = 50,
|
||||
offset: Annotated[int, Query(ge=0, description="Rows to skip for pagination")] = 0,
|
||||
) -> dict[str, Any]:
|
||||
"""Return audit log entries with optional filtering and pagination.
|
||||
|
||||
Requires authentication. Returns events in reverse chronological order.
|
||||
"""
|
||||
entries = query_events(
|
||||
db,
|
||||
action=action,
|
||||
user=user,
|
||||
resource_type=resource_type,
|
||||
severity=severity,
|
||||
since=since,
|
||||
until=until,
|
||||
limit=limit,
|
||||
offset=offset,
|
||||
)
|
||||
total = count_events(
|
||||
db,
|
||||
action=action,
|
||||
user=user,
|
||||
resource_type=resource_type,
|
||||
severity=severity,
|
||||
since=since,
|
||||
until=until,
|
||||
)
|
||||
return {
|
||||
"items": [_serialize(e) for e in entries],
|
||||
"total": total,
|
||||
"limit": limit,
|
||||
"offset": offset,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/audit-logs/actions")
|
||||
@require_login
|
||||
async def list_distinct_actions(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
) -> list[str]:
|
||||
"""Return the distinct action values present in the audit log."""
|
||||
from app.models import AuditLog
|
||||
|
||||
rows = db.query(AuditLog.action).distinct().order_by(AuditLog.action).all()
|
||||
return [r[0] for r in rows]
|
||||
|
||||
|
||||
@router.get("/audit-logs/users")
|
||||
@require_login
|
||||
async def list_distinct_users(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
) -> list[str]:
|
||||
"""Return the distinct user values present in the audit log."""
|
||||
from app.models import AuditLog
|
||||
|
||||
rows = db.query(AuditLog.user).distinct().order_by(AuditLog.user).all()
|
||||
return [r[0] for r in rows]
|
||||
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
|
||||
def _serialize(entry) -> dict[str, Any]:
|
||||
"""Convert an AuditLog row to a JSON-safe dict."""
|
||||
import json as _json
|
||||
|
||||
return {
|
||||
"id": entry.id,
|
||||
"timestamp": entry.timestamp.isoformat() if entry.timestamp else None,
|
||||
"user": entry.user,
|
||||
"action": entry.action,
|
||||
"resource_type": entry.resource_type,
|
||||
"resource_id": entry.resource_id,
|
||||
"ip_address": entry.ip_address,
|
||||
"details": _json.loads(entry.details) if entry.details else None,
|
||||
"severity": entry.severity,
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
"""
|
||||
Backup and restore API endpoints for DocuElevate.
|
||||
|
||||
Provides REST endpoints for:
|
||||
- Listing existing backups
|
||||
- Triggering a manual backup
|
||||
- Downloading a backup archive
|
||||
- Restoring from an uploaded backup file
|
||||
- Deleting a backup record
|
||||
- Running retention cleanup
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile, status
|
||||
from fastapi.responses import FileResponse
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import BackupRecord
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/admin/backup", tags=["backup"])
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin. Raises 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
# Annotated shorthand so FastAPI can resolve and tests can override it.
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/")
|
||||
async def list_backups(
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> list[dict]:
|
||||
"""Return all backup records, newest first."""
|
||||
records = db.query(BackupRecord).order_by(BackupRecord.created_at.desc()).all()
|
||||
return [
|
||||
{
|
||||
"id": r.id,
|
||||
"filename": r.filename,
|
||||
"backup_type": r.backup_type,
|
||||
"size_bytes": r.size_bytes,
|
||||
"checksum": r.checksum,
|
||||
"status": r.status,
|
||||
"local_path": r.local_path,
|
||||
"remote_destination": r.remote_destination,
|
||||
"remote_path": r.remote_path,
|
||||
"created_at": r.created_at.isoformat() if r.created_at else None,
|
||||
"local_available": bool(r.local_path and os.path.exists(r.local_path)),
|
||||
}
|
||||
for r in records
|
||||
]
|
||||
|
||||
|
||||
@router.post("/create")
|
||||
async def trigger_backup(
|
||||
_admin: AdminUser,
|
||||
backup_type: str = "hourly",
|
||||
) -> dict:
|
||||
"""Trigger a manual backup immediately.
|
||||
|
||||
Query parameter ``backup_type`` accepts ``hourly``, ``daily``, or
|
||||
``weekly`` (default: ``hourly``).
|
||||
"""
|
||||
if backup_type not in ("hourly", "daily", "weekly"):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid backup_type")
|
||||
|
||||
from app.tasks.backup_tasks import create_backup
|
||||
|
||||
task = create_backup.delay(backup_type=backup_type)
|
||||
return {"task_id": task.id, "status": "queued", "backup_type": backup_type}
|
||||
|
||||
|
||||
@router.get("/{backup_id}/download")
|
||||
async def download_backup(
|
||||
backup_id: int,
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> FileResponse:
|
||||
"""Stream the backup archive to the client."""
|
||||
rec = db.get(BackupRecord, backup_id)
|
||||
if rec is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Backup not found")
|
||||
if not rec.local_path or not os.path.exists(rec.local_path):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="Local archive file is not available (may have been pruned)",
|
||||
)
|
||||
return FileResponse(
|
||||
path=rec.local_path,
|
||||
filename=rec.filename,
|
||||
media_type="application/gzip",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/restore")
|
||||
async def restore_backup(
|
||||
_admin: AdminUser,
|
||||
file: UploadFile,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""Restore the database from an uploaded gzip-compressed SQL dump.
|
||||
|
||||
**Warning**: This overwrites the current database contents.
|
||||
|
||||
Supported formats (must match the currently configured database backend):
|
||||
|
||||
- ``*.db.gz`` – gzip-compressed SQLite ``.dump()`` SQL script (SQLite backend)
|
||||
- ``*.pgsql.gz`` – gzip-compressed ``pg_dump --format=plain`` output (PostgreSQL backend)
|
||||
- ``*.mysql.gz`` – gzip-compressed ``mysqldump`` output (MySQL / MariaDB backend)
|
||||
"""
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
from app.config import settings as app_settings
|
||||
from app.tasks.backup_tasks import (
|
||||
_archive_ext_for_backend,
|
||||
_db_path,
|
||||
_restore_mysql,
|
||||
_restore_postgresql,
|
||||
_restore_sqlite,
|
||||
)
|
||||
|
||||
url = make_url(app_settings.database_url)
|
||||
backend = url.get_backend_name()
|
||||
expected_ext = _archive_ext_for_backend(backend)
|
||||
|
||||
if not file.filename or not file.filename.endswith(expected_ext):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=(
|
||||
f"Uploaded file must be a '{expected_ext}' backup archive for the current database backend ({backend})."
|
||||
),
|
||||
)
|
||||
|
||||
# Write upload to a temp file
|
||||
with tempfile.NamedTemporaryFile(suffix=expected_ext, delete=False) as tmp:
|
||||
tmp_path = Path(tmp.name)
|
||||
content = await file.read()
|
||||
tmp.write(content)
|
||||
|
||||
try:
|
||||
if backend == "sqlite":
|
||||
db_path = _db_path()
|
||||
if db_path is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Restore is only supported for file-based SQLite databases.",
|
||||
)
|
||||
# Close the application DB session before replacing the file
|
||||
db.close()
|
||||
try:
|
||||
_restore_sqlite(db_path, tmp_path)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
except RuntimeError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
|
||||
elif backend == "postgresql":
|
||||
db.close()
|
||||
try:
|
||||
_restore_postgresql(app_settings.database_url, tmp_path)
|
||||
except FileNotFoundError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"psql binary not found – is PostgreSQL client installed? ({exc})",
|
||||
) from exc
|
||||
except RuntimeError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"PostgreSQL restore failed: {exc}",
|
||||
) from exc
|
||||
|
||||
elif backend == "mysql":
|
||||
db.close()
|
||||
try:
|
||||
_restore_mysql(app_settings.database_url, tmp_path)
|
||||
except FileNotFoundError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"mysql binary not found – is MySQL client installed? ({exc})",
|
||||
) from exc
|
||||
except RuntimeError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"MySQL restore failed: {exc}",
|
||||
) from exc
|
||||
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Database backend '{backend}' does not support restore.",
|
||||
)
|
||||
|
||||
finally:
|
||||
tmp_path.unlink(missing_ok=True)
|
||||
|
||||
logger.info(f"Database restored from uploaded backup: {file.filename}")
|
||||
return {"status": "restored", "filename": file.filename}
|
||||
|
||||
|
||||
@router.delete("/{backup_id}")
|
||||
async def delete_backup(
|
||||
backup_id: int,
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""Delete a backup record (and local file if present)."""
|
||||
rec = db.get(BackupRecord, backup_id)
|
||||
if rec is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Backup not found")
|
||||
|
||||
if rec.local_path and os.path.exists(rec.local_path):
|
||||
try:
|
||||
os.remove(rec.local_path)
|
||||
except OSError as exc:
|
||||
logger.warning(f"Could not remove local backup file {rec.local_path}: {exc}")
|
||||
|
||||
db.delete(rec)
|
||||
db.commit()
|
||||
return {"status": "deleted", "id": backup_id}
|
||||
|
||||
|
||||
@router.post("/cleanup")
|
||||
async def run_cleanup(_admin: AdminUser) -> dict:
|
||||
"""Manually trigger the retention cleanup for all backup tiers."""
|
||||
from app.tasks.backup_tasks import cleanup_old_backups
|
||||
|
||||
task = cleanup_old_backups.delay()
|
||||
return {"task_id": task.id, "status": "queued"}
|
||||
@@ -0,0 +1,619 @@
|
||||
"""Stripe billing integration for DocuElevate.
|
||||
|
||||
Provides three endpoints:
|
||||
- POST /api/billing/create-checkout-session — starts Stripe Checkout for a plan upgrade
|
||||
- POST /api/billing/create-portal-session — opens Stripe Customer Portal (manage/cancel)
|
||||
- POST /api/billing/webhook — handles Stripe webhook events
|
||||
- GET /api/billing/success — success landing page after checkout
|
||||
|
||||
Stripe Python SDK license: MIT (compatible with this project's Apache 2.0 license).
|
||||
|
||||
GDPR: Stripe acts as a data processor under a Data Processing Agreement (DPA).
|
||||
Stripe is SOC 2 Type II certified and supports EU data residency.
|
||||
SOC2: Stripe is SOC 2 Type II certified.
|
||||
EU VAT: Configure Stripe Tax in the Stripe Dashboard for automatic VAT collection.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import pathlib
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
import stripe
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.templating import Jinja2Templates
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import SubscriptionPlan, UserProfile
|
||||
from app.utils.i18n import translate as _translate
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/billing", tags=["billing"])
|
||||
|
||||
_templates_dir = pathlib.Path(__file__).parents[2] / "frontend" / "templates"
|
||||
_templates = Jinja2Templates(directory=str(_templates_dir))
|
||||
_templates.env.globals["_"] = lambda key, **kwargs: _translate(key, "en", **kwargs)
|
||||
|
||||
|
||||
def _get_stripe() -> stripe.StripeClient | None:
|
||||
"""Return a configured Stripe client, or None when not configured."""
|
||||
if not settings.stripe_secret_key:
|
||||
return None
|
||||
return stripe.StripeClient(settings.stripe_secret_key)
|
||||
|
||||
|
||||
def _get_or_create_stripe_customer(
|
||||
client: stripe.StripeClient,
|
||||
db: Session,
|
||||
owner_id: str,
|
||||
email: str | None,
|
||||
name: str | None,
|
||||
) -> str:
|
||||
"""Return the Stripe customer_id for *owner_id*, creating one if needed.
|
||||
|
||||
Args:
|
||||
client: Configured Stripe client.
|
||||
db: Database session.
|
||||
owner_id: Stable user identifier.
|
||||
email: User's email for the Stripe customer record.
|
||||
name: User's display name for the Stripe customer record.
|
||||
|
||||
Returns:
|
||||
The Stripe customer ID string.
|
||||
"""
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == owner_id).first()
|
||||
if profile and profile.stripe_customer_id:
|
||||
return profile.stripe_customer_id
|
||||
|
||||
customer = client.customers.create(
|
||||
params={
|
||||
"email": email or "",
|
||||
"name": name or "",
|
||||
"metadata": {"docuelevate_user_id": owner_id},
|
||||
}
|
||||
)
|
||||
if profile:
|
||||
profile.stripe_customer_id = customer.id
|
||||
db.commit()
|
||||
return customer.id
|
||||
|
||||
|
||||
class CheckoutSessionBody(BaseModel):
|
||||
"""Request body for creating a Stripe Checkout session."""
|
||||
|
||||
plan_id: str
|
||||
billing_cycle: str = "monthly" # "monthly" | "yearly"
|
||||
|
||||
|
||||
class PortalSessionBody(BaseModel):
|
||||
"""Request body for creating a Stripe Customer Portal session."""
|
||||
|
||||
return_url: str | None = None
|
||||
|
||||
|
||||
@router.post("/create-checkout-session", summary="Create a Stripe Checkout session for a plan upgrade")
|
||||
@require_login
|
||||
async def create_checkout_session(
|
||||
request: Request,
|
||||
body: CheckoutSessionBody,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict[str, Any]:
|
||||
"""Create a Stripe Checkout session.
|
||||
|
||||
The client should redirect the user to the returned ``checkout_url``.
|
||||
|
||||
Raises:
|
||||
503: Stripe is not configured.
|
||||
404: Plan not found or has no Stripe price configured.
|
||||
"""
|
||||
client = _get_stripe()
|
||||
if not client:
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing is not configured.")
|
||||
|
||||
plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == body.plan_id).first()
|
||||
if plan is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Plan {body.plan_id!r} not found.")
|
||||
|
||||
price_id = plan.stripe_price_id_yearly if body.billing_cycle == "yearly" else plan.stripe_price_id_monthly
|
||||
if not price_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=(
|
||||
f"Stripe price ID not configured for plan {body.plan_id!r} ({body.billing_cycle}). "
|
||||
"Please set it in the Admin Plan Designer."
|
||||
),
|
||||
)
|
||||
|
||||
user = request.session.get("user") or {}
|
||||
owner_id = get_current_owner_id(request) or user.get("email") or ""
|
||||
email = user.get("email")
|
||||
name = user.get("name")
|
||||
|
||||
customer_id = _get_or_create_stripe_customer(client, db, owner_id, email, name)
|
||||
|
||||
base = str(request.base_url).rstrip("/")
|
||||
success_url = settings.stripe_success_url or f"{base}/api/billing/success"
|
||||
cancel_url = settings.stripe_cancel_url or f"{base}/pricing"
|
||||
|
||||
trial_days = plan.trial_days if plan.trial_days > 0 else None
|
||||
|
||||
session_params: dict[str, Any] = {
|
||||
"customer": customer_id,
|
||||
"mode": "subscription",
|
||||
"line_items": [{"price": price_id, "quantity": 1}],
|
||||
"success_url": success_url + "?session_id={CHECKOUT_SESSION_ID}",
|
||||
"cancel_url": cancel_url,
|
||||
"subscription_data": {
|
||||
"metadata": {
|
||||
"docuelevate_user_id": owner_id,
|
||||
"plan_id": body.plan_id,
|
||||
"billing_cycle": body.billing_cycle,
|
||||
},
|
||||
},
|
||||
"metadata": {"docuelevate_user_id": owner_id, "plan_id": body.plan_id},
|
||||
"allow_promotion_codes": True,
|
||||
"billing_address_collection": "auto",
|
||||
"tax_id_collection": {"enabled": True},
|
||||
"automatic_tax": {"enabled": True},
|
||||
}
|
||||
if trial_days:
|
||||
session_params["subscription_data"]["trial_period_days"] = trial_days
|
||||
|
||||
checkout_session = client.checkout.sessions.create(params=session_params)
|
||||
|
||||
logger.info(
|
||||
"Created Stripe checkout session %s for user %s plan %s",
|
||||
checkout_session.id,
|
||||
owner_id,
|
||||
body.plan_id,
|
||||
)
|
||||
return {"checkout_url": checkout_session.url, "session_id": checkout_session.id}
|
||||
|
||||
|
||||
@router.post("/create-portal-session", summary="Create a Stripe Customer Portal session")
|
||||
@require_login
|
||||
async def create_portal_session(
|
||||
request: Request,
|
||||
body: PortalSessionBody,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict[str, Any]:
|
||||
"""Create a Stripe Customer Portal session for subscription self-management.
|
||||
|
||||
Raises:
|
||||
503: Stripe not configured.
|
||||
404: No Stripe customer found for this user.
|
||||
"""
|
||||
client = _get_stripe()
|
||||
if not client:
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing is not configured.")
|
||||
|
||||
user = request.session.get("user") or {}
|
||||
owner_id = get_current_owner_id(request) or user.get("email") or ""
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == owner_id).first()
|
||||
if not profile or not profile.stripe_customer_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="No billing account found. Please subscribe to a plan first.",
|
||||
)
|
||||
|
||||
base = str(request.base_url).rstrip("/")
|
||||
return_url = body.return_url or f"{base}/subscription"
|
||||
|
||||
portal = client.billing_portal.sessions.create(
|
||||
params={
|
||||
"customer": profile.stripe_customer_id,
|
||||
"return_url": return_url,
|
||||
}
|
||||
)
|
||||
|
||||
logger.info("Created Stripe portal session for user %s", owner_id)
|
||||
return {"portal_url": portal.url}
|
||||
|
||||
|
||||
@router.post("/webhook", include_in_schema=False)
|
||||
async def stripe_webhook(request: Request, db: Session = Depends(get_db)) -> dict[str, str]:
|
||||
"""Handle Stripe webhook events.
|
||||
|
||||
Syncs subscription status to UserProfile.subscription_tier.
|
||||
|
||||
Events handled:
|
||||
|
||||
- ``checkout.session.completed`` — activate subscription after payment
|
||||
- ``customer.subscription.updated`` — sync tier change
|
||||
- ``customer.subscription.deleted`` — downgrade to free on cancellation
|
||||
- ``invoice.payment_failed`` — log failed payment
|
||||
"""
|
||||
if not settings.stripe_secret_key:
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing not configured.")
|
||||
|
||||
payload = await request.body()
|
||||
sig_header = request.headers.get("stripe-signature", "")
|
||||
|
||||
try:
|
||||
if settings.stripe_webhook_secret:
|
||||
event = stripe.Webhook.construct_event(payload, sig_header, settings.stripe_webhook_secret)
|
||||
else:
|
||||
logger.warning(
|
||||
"[SECURITY] STRIPE_WEBHOOK_SECRET is not configured. "
|
||||
"Webhook events are accepted without signature verification. "
|
||||
"Set STRIPE_WEBHOOK_SECRET in production to prevent spoofed events."
|
||||
)
|
||||
event = stripe.Event.construct_from(json.loads(payload), stripe.api_key)
|
||||
except stripe.SignatureVerificationError:
|
||||
logger.warning("[SECURITY] Stripe webhook signature verification failed")
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid webhook signature.")
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to parse Stripe webhook: %s", exc)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid webhook payload.")
|
||||
|
||||
_handle_stripe_event(db, event)
|
||||
return {"status": "ok"}
|
||||
|
||||
|
||||
@router.get("/success", include_in_schema=False)
|
||||
@require_login
|
||||
async def billing_success(request: Request) -> Any:
|
||||
"""Show a success page after a completed Stripe Checkout."""
|
||||
return _templates.TemplateResponse("billing_success.html", {"request": request})
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Admin: Stripe status + sync helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> None:
|
||||
"""Raise 403 if the current session user is not an admin."""
|
||||
user = request.session.get("user") or {}
|
||||
if not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required.")
|
||||
|
||||
|
||||
@router.get("/stripe/status", summary="Check Stripe connection and plan sync status (admin only)")
|
||||
@require_login
|
||||
async def stripe_status(request: Request, db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
"""Return Stripe connection health and per-plan price-ID sync status.
|
||||
|
||||
Returns a JSON object with:
|
||||
- ``configured``: whether STRIPE_SECRET_KEY is set
|
||||
- ``connection``: ``"ok"`` or an error string (live/test mode label)
|
||||
- ``mode``: ``"live"`` | ``"test"`` | ``null``
|
||||
- ``plans``: list of plan objects with ``plan_id``, ``name``,
|
||||
``stripe_price_id_monthly``, ``stripe_price_id_yearly``, ``synced``
|
||||
|
||||
Raises:
|
||||
403: Not admin.
|
||||
503: Stripe not configured.
|
||||
"""
|
||||
_require_admin(request)
|
||||
|
||||
if not settings.stripe_secret_key:
|
||||
return {
|
||||
"configured": False,
|
||||
"connection": "not_configured",
|
||||
"mode": None,
|
||||
"plans": [],
|
||||
}
|
||||
|
||||
client = _get_stripe()
|
||||
# Probe Stripe with a lightweight account fetch
|
||||
mode: str | None = None
|
||||
connection_status = "ok"
|
||||
try:
|
||||
account = client.accounts.retrieve("me") # type: ignore[arg-type]
|
||||
livemode = getattr(account, "livemode", None)
|
||||
if livemode is True:
|
||||
mode = "live"
|
||||
elif livemode is False:
|
||||
mode = "test"
|
||||
else:
|
||||
mode = "test" if settings.stripe_secret_key.startswith("sk_test_") else "live"
|
||||
except Exception:
|
||||
logger.exception("Stripe connection check failed")
|
||||
connection_status = "error"
|
||||
mode = "test" if settings.stripe_secret_key.startswith("sk_test_") else "live"
|
||||
|
||||
plans = db.query(SubscriptionPlan).order_by(SubscriptionPlan.sort_order).all()
|
||||
plan_statuses = []
|
||||
for plan in plans:
|
||||
has_monthly = bool(plan.stripe_price_id_monthly)
|
||||
has_yearly = bool(plan.stripe_price_id_yearly)
|
||||
is_paid = plan.price_monthly > 0 or plan.price_yearly > 0
|
||||
synced = (not is_paid) or (has_monthly and (not plan.price_yearly or has_yearly))
|
||||
plan_statuses.append(
|
||||
{
|
||||
"plan_id": plan.plan_id,
|
||||
"name": plan.name,
|
||||
"price_monthly": plan.price_monthly,
|
||||
"price_yearly": plan.price_yearly,
|
||||
"stripe_price_id_monthly": plan.stripe_price_id_monthly,
|
||||
"stripe_price_id_yearly": plan.stripe_price_id_yearly,
|
||||
"synced": synced,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"configured": True,
|
||||
"connection": connection_status,
|
||||
"mode": mode,
|
||||
"webhook_secret_configured": bool(settings.stripe_webhook_secret),
|
||||
"plans": plan_statuses,
|
||||
"webhook_endpoint": str(request.base_url).rstrip("/") + "/api/billing/webhook",
|
||||
}
|
||||
|
||||
|
||||
@router.post("/stripe/sync-plans", summary="Auto-create Stripe products and prices for all plans (admin only)")
|
||||
@require_login
|
||||
async def stripe_sync_plans(request: Request, db: Session = Depends(get_db)) -> dict[str, Any]:
|
||||
"""Create Stripe Product + Price objects for every paid plan that is missing them.
|
||||
|
||||
For each paid plan (``price_monthly > 0``) that lacks a ``stripe_price_id_monthly``,
|
||||
this endpoint:
|
||||
|
||||
1. Creates a Stripe *Product* named after the plan.
|
||||
2. Creates a Stripe *Price* for the monthly amount.
|
||||
3. Optionally creates a yearly Price if ``price_yearly > 0``.
|
||||
4. Persists the resulting ``price_id`` values back into ``SubscriptionPlan``.
|
||||
|
||||
Already-synced plans (those that already have ``stripe_price_id_monthly``) are
|
||||
skipped — existing prices in Stripe are never modified.
|
||||
|
||||
Raises:
|
||||
403: Not admin.
|
||||
503: Stripe not configured.
|
||||
"""
|
||||
_require_admin(request)
|
||||
|
||||
client = _get_stripe()
|
||||
if not client:
|
||||
raise HTTPException(status_code=status.HTTP_503_SERVICE_UNAVAILABLE, detail="Billing is not configured.")
|
||||
|
||||
plans = db.query(SubscriptionPlan).order_by(SubscriptionPlan.sort_order).all()
|
||||
results: list[dict[str, Any]] = []
|
||||
|
||||
for plan in plans:
|
||||
is_paid = plan.price_monthly > 0 or plan.price_yearly > 0
|
||||
if not is_paid:
|
||||
results.append({"plan_id": plan.plan_id, "name": plan.name, "status": "skipped_free"})
|
||||
continue
|
||||
|
||||
already_has_monthly = bool(plan.stripe_price_id_monthly)
|
||||
already_has_yearly = bool(plan.stripe_price_id_yearly)
|
||||
|
||||
if already_has_monthly and (not plan.price_yearly or already_has_yearly):
|
||||
results.append({"plan_id": plan.plan_id, "name": plan.name, "status": "already_synced"})
|
||||
continue
|
||||
|
||||
try:
|
||||
# Create (or look up) the Stripe Product for this plan
|
||||
product = client.products.create(
|
||||
params={
|
||||
"name": str(plan.name),
|
||||
"metadata": {"docuelevate_plan_id": plan.plan_id},
|
||||
}
|
||||
)
|
||||
|
||||
changed = False
|
||||
|
||||
# Monthly price
|
||||
if not already_has_monthly and plan.price_monthly > 0:
|
||||
monthly_price = client.prices.create(
|
||||
params={
|
||||
"product": product.id,
|
||||
"unit_amount": int(round(plan.price_monthly * 100)),
|
||||
"currency": "usd",
|
||||
"recurring": {"interval": "month"},
|
||||
"metadata": {"docuelevate_plan_id": plan.plan_id, "billing_cycle": "monthly"},
|
||||
}
|
||||
)
|
||||
plan.stripe_price_id_monthly = monthly_price.id
|
||||
changed = True
|
||||
|
||||
# Yearly price
|
||||
if not already_has_yearly and plan.price_yearly > 0:
|
||||
yearly_price = client.prices.create(
|
||||
params={
|
||||
"product": product.id,
|
||||
"unit_amount": int(round(plan.price_yearly * 100)),
|
||||
"currency": "usd",
|
||||
"recurring": {"interval": "year"},
|
||||
"metadata": {"docuelevate_plan_id": plan.plan_id, "billing_cycle": "yearly"},
|
||||
}
|
||||
)
|
||||
plan.stripe_price_id_yearly = yearly_price.id
|
||||
changed = True
|
||||
|
||||
if changed:
|
||||
db.commit()
|
||||
logger.info(
|
||||
"Stripe sync: created product/prices for plan %s (product %s)",
|
||||
plan.plan_id,
|
||||
product.id,
|
||||
)
|
||||
|
||||
results.append(
|
||||
{
|
||||
"plan_id": plan.plan_id,
|
||||
"name": plan.name,
|
||||
"status": "created",
|
||||
"stripe_price_id_monthly": plan.stripe_price_id_monthly,
|
||||
"stripe_price_id_yearly": plan.stripe_price_id_yearly,
|
||||
}
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.error("Stripe sync failed for plan %s: %s", plan.plan_id, exc)
|
||||
results.append(
|
||||
{
|
||||
"plan_id": plan.plan_id,
|
||||
"name": str(plan.name),
|
||||
"status": "error",
|
||||
"detail": str(exc),
|
||||
}
|
||||
)
|
||||
|
||||
return {"results": results}
|
||||
|
||||
|
||||
def _handle_stripe_event(db: Session, event: Any) -> None:
|
||||
"""Dispatch Stripe event to the appropriate handler.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
event: Parsed Stripe event object.
|
||||
"""
|
||||
etype = event.get("type", "") if isinstance(event, dict) else getattr(event, "type", "")
|
||||
data_obj = (
|
||||
event.get("data", {}).get("object", {})
|
||||
if isinstance(event, dict)
|
||||
else getattr(getattr(event, "data", None), "object", {})
|
||||
)
|
||||
|
||||
if etype == "checkout.session.completed":
|
||||
_on_checkout_completed(db, data_obj)
|
||||
elif etype == "customer.subscription.updated":
|
||||
_on_subscription_updated(db, data_obj)
|
||||
elif etype == "customer.subscription.deleted":
|
||||
_on_subscription_deleted(db, data_obj)
|
||||
elif etype == "invoice.payment_failed":
|
||||
customer_id = data_obj.get("customer", "") if isinstance(data_obj, dict) else getattr(data_obj, "customer", "")
|
||||
logger.warning("Stripe invoice payment failed for customer %s", customer_id)
|
||||
else:
|
||||
logger.debug("Unhandled Stripe event type: %s", etype)
|
||||
|
||||
|
||||
def _resolve_user_id_from_customer(db: Session, customer_id: str) -> str | None:
|
||||
"""Look up the DocuElevate user_id for a Stripe customer_id.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
customer_id: Stripe customer ID.
|
||||
|
||||
Returns:
|
||||
The matching ``UserProfile.user_id``, or ``None`` if not found.
|
||||
"""
|
||||
profile = db.query(UserProfile).filter(UserProfile.stripe_customer_id == customer_id).first()
|
||||
return profile.user_id if profile else None
|
||||
|
||||
|
||||
def _resolve_plan_id_from_price(db: Session, price_id: str) -> str | None:
|
||||
"""Map a Stripe price_id to a DocuElevate plan_id via SubscriptionPlan.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
price_id: Stripe price ID.
|
||||
|
||||
Returns:
|
||||
The matching ``SubscriptionPlan.plan_id``, or ``None`` if not found.
|
||||
"""
|
||||
plan = (
|
||||
db.query(SubscriptionPlan)
|
||||
.filter(
|
||||
(SubscriptionPlan.stripe_price_id_monthly == price_id)
|
||||
| (SubscriptionPlan.stripe_price_id_yearly == price_id)
|
||||
)
|
||||
.first()
|
||||
)
|
||||
return plan.plan_id if plan else None
|
||||
|
||||
|
||||
def _on_checkout_completed(db: Session, data: Any) -> None:
|
||||
"""Activate a subscription after a successful checkout.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
data: Stripe ``checkout.session`` object.
|
||||
"""
|
||||
meta = data.get("metadata") or {} if isinstance(data, dict) else getattr(data, "metadata", {}) or {}
|
||||
user_id = meta.get("docuelevate_user_id") if isinstance(meta, dict) else getattr(meta, "docuelevate_user_id", None)
|
||||
plan_id = meta.get("plan_id") if isinstance(meta, dict) else getattr(meta, "plan_id", None)
|
||||
billing_cycle = (
|
||||
meta.get("billing_cycle", "monthly") if isinstance(meta, dict) else getattr(meta, "billing_cycle", "monthly")
|
||||
)
|
||||
if not user_id:
|
||||
return
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile and plan_id:
|
||||
profile.subscription_tier = plan_id
|
||||
profile.subscription_billing_cycle = billing_cycle
|
||||
profile.subscription_period_start = datetime.now(tz=timezone.utc)
|
||||
customer_id = data.get("customer", "") if isinstance(data, dict) else getattr(data, "customer", "")
|
||||
if customer_id:
|
||||
profile.stripe_customer_id = customer_id
|
||||
db.commit()
|
||||
logger.info("Activated plan %s/%s after checkout", plan_id, billing_cycle)
|
||||
|
||||
|
||||
def _on_subscription_updated(db: Session, data: Any) -> None:
|
||||
"""Sync tier change when a subscription is updated.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
data: Stripe ``customer.subscription`` object.
|
||||
"""
|
||||
customer_id = data.get("customer", "") if isinstance(data, dict) else getattr(data, "customer", "")
|
||||
user_id = _resolve_user_id_from_customer(db, customer_id)
|
||||
if not user_id:
|
||||
return
|
||||
|
||||
items_data = data.get("items") or {} if isinstance(data, dict) else getattr(data, "items", None) or {}
|
||||
items = items_data.get("data") or [] if isinstance(items_data, dict) else getattr(items_data, "data", []) or []
|
||||
if not items:
|
||||
return
|
||||
|
||||
first_item = items[0]
|
||||
price_obj = (
|
||||
first_item.get("price") or {} if isinstance(first_item, dict) else getattr(first_item, "price", {}) or {}
|
||||
)
|
||||
price_id = price_obj.get("id") if isinstance(price_obj, dict) else getattr(price_obj, "id", None)
|
||||
if not price_id:
|
||||
return
|
||||
|
||||
plan_id = _resolve_plan_id_from_price(db, price_id)
|
||||
if not plan_id:
|
||||
logger.warning("Unknown Stripe price_id %s on subscription.updated", price_id)
|
||||
return
|
||||
|
||||
recurring = (
|
||||
price_obj.get("recurring", {}) if isinstance(price_obj, dict) else getattr(price_obj, "recurring", {}) or {}
|
||||
)
|
||||
interval = (
|
||||
recurring.get("interval", "month") if isinstance(recurring, dict) else getattr(recurring, "interval", "month")
|
||||
)
|
||||
billing_cycle = "yearly" if interval == "year" else "monthly"
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile:
|
||||
profile.subscription_tier = plan_id
|
||||
profile.subscription_billing_cycle = billing_cycle
|
||||
db.commit()
|
||||
logger.info("Updated subscription to %s/%s", plan_id, billing_cycle)
|
||||
|
||||
|
||||
def _on_subscription_deleted(db: Session, data: Any) -> None:
|
||||
"""Downgrade user to free tier after subscription cancellation.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
data: Stripe ``customer.subscription`` object.
|
||||
"""
|
||||
customer_id = data.get("customer", "") if isinstance(data, dict) else getattr(data, "customer", "")
|
||||
user_id = _resolve_user_id_from_customer(db, customer_id)
|
||||
if not user_id:
|
||||
return
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile:
|
||||
profile.subscription_tier = "free"
|
||||
profile.subscription_billing_cycle = "monthly"
|
||||
db.commit()
|
||||
logger.info("Downgraded user %s to free tier after subscription cancellation", user_id)
|
||||
@@ -0,0 +1,183 @@
|
||||
"""API endpoints for managing compliance templates (GDPR, HIPAA, SOC2).
|
||||
|
||||
All endpoints require admin privileges.
|
||||
|
||||
Available routes:
|
||||
GET /api/compliance/templates – list all compliance templates
|
||||
GET /api/compliance/templates/{name} – get a single template with checks
|
||||
POST /api/compliance/templates/{name}/apply – one-click apply a template
|
||||
GET /api/compliance/templates/{name}/status – evaluate compliance status
|
||||
GET /api/compliance/summary – overall compliance dashboard data
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.utils.compliance_service import (
|
||||
COMPLIANCE_TEMPLATES,
|
||||
apply_template,
|
||||
evaluate_template_status,
|
||||
get_all_templates,
|
||||
get_compliance_summary,
|
||||
get_template_by_name,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/compliance", tags=["compliance"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Authorisation helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin; raises HTTP 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic response models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class CheckResult(BaseModel):
|
||||
"""Individual compliance check result."""
|
||||
|
||||
key: str
|
||||
label: str
|
||||
description: str
|
||||
expected: str
|
||||
actual: str
|
||||
passing: bool
|
||||
|
||||
|
||||
class TemplateStatusResponse(BaseModel):
|
||||
"""Status evaluation for a compliance template."""
|
||||
|
||||
status: str
|
||||
total: int
|
||||
passed: int
|
||||
failed: int
|
||||
check_results: list[CheckResult]
|
||||
|
||||
|
||||
class TemplateResponse(BaseModel):
|
||||
"""Full compliance template representation."""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
display_name: str
|
||||
description: str | None
|
||||
enabled: bool
|
||||
status: str
|
||||
applied_at: str | None
|
||||
applied_by: str | None
|
||||
settings: dict[str, str]
|
||||
checks: list[dict[str, Any]]
|
||||
check_count: int
|
||||
|
||||
|
||||
class ApplyResponse(BaseModel):
|
||||
"""Result of applying a compliance template."""
|
||||
|
||||
success: bool
|
||||
template: str | None = None
|
||||
applied_settings: dict[str, str] | None = None
|
||||
errors: list[str] | None = None
|
||||
error: str | None = None
|
||||
status: TemplateStatusResponse | None = None
|
||||
|
||||
|
||||
class SummaryTemplateResponse(BaseModel):
|
||||
"""Per-template summary for the compliance dashboard."""
|
||||
|
||||
name: str
|
||||
display_name: str
|
||||
enabled: bool
|
||||
status: str
|
||||
total: int
|
||||
passed: int
|
||||
failed: int
|
||||
applied_at: str | None
|
||||
applied_by: str | None
|
||||
|
||||
|
||||
class ComplianceSummaryResponse(BaseModel):
|
||||
"""Overall compliance dashboard summary."""
|
||||
|
||||
overall_status: str
|
||||
total_checks: int
|
||||
total_passed: int
|
||||
total_failed: int
|
||||
templates: list[SummaryTemplateResponse]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/templates", response_model=list[TemplateResponse])
|
||||
async def list_templates(db: DbSession, admin: AdminUser) -> list[dict[str, Any]]:
|
||||
"""List all compliance templates with their current status."""
|
||||
return get_all_templates(db)
|
||||
|
||||
|
||||
@router.get("/templates/{name}", response_model=TemplateResponse)
|
||||
async def get_template(name: str, db: DbSession, admin: AdminUser) -> dict[str, Any]:
|
||||
"""Get a single compliance template by name."""
|
||||
templates = get_all_templates(db)
|
||||
for t in templates:
|
||||
if t["name"] == name:
|
||||
return t
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Template '{name}' not found")
|
||||
|
||||
|
||||
@router.post("/templates/{name}/apply", response_model=ApplyResponse)
|
||||
async def apply_compliance_template(name: str, db: DbSession, admin: AdminUser) -> dict[str, Any]:
|
||||
"""Apply a compliance template (one-click).
|
||||
|
||||
Writes all template settings to the database and evaluates the resulting
|
||||
compliance status.
|
||||
"""
|
||||
if name not in COMPLIANCE_TEMPLATES:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Template '{name}' not found")
|
||||
|
||||
template = get_template_by_name(db, name)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Template '{name}' not found")
|
||||
|
||||
admin_email = admin.get("email", "admin")
|
||||
result = apply_template(db, name, applied_by=admin_email)
|
||||
if not result.get("success") and result.get("error"):
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=result["error"])
|
||||
return result
|
||||
|
||||
|
||||
@router.get("/templates/{name}/status", response_model=TemplateStatusResponse)
|
||||
async def get_template_status(name: str, db: DbSession, admin: AdminUser) -> dict[str, Any]:
|
||||
"""Evaluate the live compliance status of a template."""
|
||||
template = get_template_by_name(db, name)
|
||||
if template is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Template '{name}' not found")
|
||||
return evaluate_template_status(db, name)
|
||||
|
||||
|
||||
@router.get("/summary", response_model=ComplianceSummaryResponse)
|
||||
async def compliance_summary(db: DbSession, admin: AdminUser) -> dict[str, Any]:
|
||||
"""Overall compliance dashboard summary across all templates."""
|
||||
return get_compliance_summary(db)
|
||||
@@ -0,0 +1,170 @@
|
||||
"""
|
||||
API endpoints for the database configuration wizard and migration tool.
|
||||
|
||||
Provides REST endpoints for:
|
||||
- Testing database connections
|
||||
- Building connection strings from form components
|
||||
- Previewing and executing data migrations between databases
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.utils.db_migrate import migrate_data, preview_migration
|
||||
from app.utils.db_wizard import (
|
||||
build_connection_string,
|
||||
get_supported_backends,
|
||||
parse_connection_string,
|
||||
test_connection,
|
||||
validate_url_format,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/database", tags=["database"])
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin. Raises 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request / Response models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ConnectionStringRequest(BaseModel):
|
||||
"""Request body for building a connection string."""
|
||||
|
||||
backend: str = Field(..., description="Database backend: sqlite, postgresql, mysql")
|
||||
host: str = Field("", description="Database server hostname")
|
||||
port: int | None = Field(None, description="Database server port")
|
||||
database: str = Field("", description="Database name")
|
||||
username: str = Field("", description="Authentication username")
|
||||
password: str = Field("", description="Authentication password")
|
||||
ssl_mode: str = Field("", description="SSL mode (e.g. require, verify-full)")
|
||||
extra_options: str = Field("", description="Additional query-string options")
|
||||
sqlite_path: str = Field("", description="File path for SQLite databases")
|
||||
|
||||
|
||||
class TestConnectionRequest(BaseModel):
|
||||
"""Request body for testing a database connection."""
|
||||
|
||||
url: str = Field(..., description="Full SQLAlchemy connection URL to test")
|
||||
|
||||
|
||||
class MigrateRequest(BaseModel):
|
||||
"""Request body for data migration."""
|
||||
|
||||
source_url: str = Field(..., description="Source database connection URL")
|
||||
target_url: str = Field(..., description="Target database connection URL")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/backends")
|
||||
async def list_backends() -> list[dict]:
|
||||
"""List all supported database backends with metadata."""
|
||||
return get_supported_backends()
|
||||
|
||||
|
||||
@router.post("/build-url")
|
||||
async def build_url(body: ConnectionStringRequest, request: Request) -> dict:
|
||||
"""Build a SQLAlchemy connection string from individual components.
|
||||
|
||||
Returns the assembled URL string.
|
||||
"""
|
||||
_require_admin(request)
|
||||
try:
|
||||
url = build_connection_string(
|
||||
backend=body.backend,
|
||||
host=body.host,
|
||||
port=body.port,
|
||||
database=body.database,
|
||||
username=body.username,
|
||||
password=body.password,
|
||||
ssl_mode=body.ssl_mode,
|
||||
extra_options=body.extra_options,
|
||||
sqlite_path=body.sqlite_path,
|
||||
)
|
||||
return {"url": url}
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
|
||||
|
||||
|
||||
@router.post("/parse-url")
|
||||
async def parse_url(body: TestConnectionRequest, request: Request) -> dict:
|
||||
"""Parse a connection string into its components."""
|
||||
_require_admin(request)
|
||||
return parse_connection_string(body.url)
|
||||
|
||||
|
||||
@router.post("/validate-url")
|
||||
async def validate_url(body: TestConnectionRequest, request: Request) -> dict:
|
||||
"""Validate a connection string format without connecting."""
|
||||
_require_admin(request)
|
||||
return validate_url_format(body.url)
|
||||
|
||||
|
||||
@router.post("/test-connection")
|
||||
async def test_db_connection(body: TestConnectionRequest, request: Request) -> dict:
|
||||
"""Test connectivity to a database and return status info.
|
||||
|
||||
This creates a temporary engine, executes ``SELECT 1``, and disposes
|
||||
of the engine. It does **not** modify any global application state.
|
||||
"""
|
||||
_require_admin(request)
|
||||
return test_connection(body.url)
|
||||
|
||||
|
||||
@router.post("/preview-migration")
|
||||
async def preview_db_migration(body: TestConnectionRequest, request: Request) -> dict:
|
||||
"""Preview what a migration from the given source would include.
|
||||
|
||||
Returns a table-by-table row count without actually copying data.
|
||||
"""
|
||||
_require_admin(request)
|
||||
return preview_migration(body.url)
|
||||
|
||||
|
||||
@router.post("/migrate")
|
||||
async def execute_migration(body: MigrateRequest, request: Request) -> dict:
|
||||
"""Execute a full data migration from source to target database.
|
||||
|
||||
**Warning:** This copies all data from the source database into the
|
||||
target. The target schema is created from the current application
|
||||
models. Existing data in the target is **not** deleted first — use
|
||||
on an empty target database.
|
||||
"""
|
||||
_require_admin(request)
|
||||
|
||||
# Validate both URLs first
|
||||
src_check = validate_url_format(body.source_url)
|
||||
if not src_check.get("valid"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid source URL: {src_check.get('error', 'unknown')}",
|
||||
)
|
||||
tgt_check = validate_url_format(body.target_url)
|
||||
if not tgt_check.get("valid"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid target URL: {tgt_check.get('error', 'unknown')}",
|
||||
)
|
||||
|
||||
result = migrate_data(body.source_url, body.target_url)
|
||||
if not result["success"]:
|
||||
error_summary = "; ".join(result.get("errors", ["Unknown error"]))
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Migration completed with errors: {error_summary}",
|
||||
)
|
||||
return result
|
||||
@@ -2,19 +2,112 @@
|
||||
Diagnostic API endpoints
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
|
||||
import redis as redis_lib
|
||||
from fastapi import APIRouter, Request
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy import text
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import engine
|
||||
|
||||
# Set up logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_REDIS_URL = "redis://localhost:6379/0"
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
@router.get("/diagnostic/health")
|
||||
@require_login
|
||||
async def health_check(request: Request):
|
||||
"""
|
||||
System health endpoint for monitoring tools (Grafana, Uptime Kuma, etc.).
|
||||
|
||||
Checks database connectivity and Redis availability and returns a
|
||||
machine-readable summary that monitoring systems can scrape.
|
||||
|
||||
**Authentication:** Required (no-op when AUTH_ENABLED=False)
|
||||
|
||||
**Response (200 OK) – all subsystems healthy:**
|
||||
```json
|
||||
{
|
||||
"status": "healthy",
|
||||
"version": "1.2.3",
|
||||
"timestamp": "2024-01-15T10:30:00+00:00",
|
||||
"checks": {
|
||||
"database": {"status": "ok"},
|
||||
"redis": {"status": "ok"}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Response (200 OK) – one or more subsystems degraded:**
|
||||
```json
|
||||
{
|
||||
"status": "degraded",
|
||||
"version": "1.2.3",
|
||||
"timestamp": "2024-01-15T10:30:00+00:00",
|
||||
"checks": {
|
||||
"database": {"status": "ok"},
|
||||
"redis": {"status": "error", "detail": "Connection refused"}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
The outer ``status`` field is always one of:
|
||||
- ``"healthy"`` – all checks passed
|
||||
- ``"degraded"`` – at least one non-critical check failed
|
||||
- ``"unhealthy"`` – a critical check failed (currently: database)
|
||||
"""
|
||||
timestamp = datetime.datetime.now(datetime.timezone.utc).isoformat()
|
||||
checks: dict[str, dict[str, str]] = {}
|
||||
|
||||
# ── Database check ─────────────────────────────────────────────────────
|
||||
db_ok = False
|
||||
try:
|
||||
with engine.connect() as conn:
|
||||
conn.execute(text("SELECT 1"))
|
||||
checks["database"] = {"status": "ok"}
|
||||
db_ok = True
|
||||
except Exception as exc:
|
||||
logger.warning("Health check: database probe failed: %s", exc)
|
||||
checks["database"] = {"status": "error", "detail": str(exc)}
|
||||
|
||||
# ── Redis check ────────────────────────────────────────────────────────
|
||||
try:
|
||||
redis_url = settings.redis_url or _DEFAULT_REDIS_URL
|
||||
r = redis_lib.from_url(redis_url, socket_connect_timeout=2, socket_timeout=2)
|
||||
r.ping()
|
||||
checks["redis"] = {"status": "ok"}
|
||||
except Exception as exc:
|
||||
logger.warning("Health check: Redis probe failed: %s", exc)
|
||||
checks["redis"] = {"status": "error", "detail": str(exc)}
|
||||
|
||||
# ── Overall status ─────────────────────────────────────────────────────
|
||||
if not db_ok:
|
||||
overall = "unhealthy"
|
||||
elif any(v.get("status") != "ok" for v in checks.values()):
|
||||
overall = "degraded"
|
||||
else:
|
||||
overall = "healthy"
|
||||
|
||||
http_status = 503 if overall == "unhealthy" else 200
|
||||
|
||||
payload = {
|
||||
"status": overall,
|
||||
"version": settings.version,
|
||||
"timestamp": timestamp,
|
||||
"checks": checks,
|
||||
}
|
||||
|
||||
return JSONResponse(content=payload, status_code=http_status)
|
||||
|
||||
|
||||
@router.post("/diagnostic/test-notification")
|
||||
@require_login
|
||||
async def test_notification(request: Request):
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
"""Duplicate document detection and management API endpoints.
|
||||
|
||||
Provides endpoints for listing all duplicate groups (exact SHA-256 duplicates) and
|
||||
for retrieving both exact and near-duplicate matches for a specific document.
|
||||
|
||||
Near-duplicate detection is powered by the same text-embedding cosine-similarity
|
||||
engine used by the ``/api/files/{id}/similar`` endpoint
|
||||
(see ``app/utils/similarity.py``).
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import FileRecord
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
@router.get("/duplicates")
|
||||
@require_login
|
||||
def list_duplicate_groups(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
page: int = Query(1, ge=1, description="Page number"),
|
||||
per_page: int = Query(25, ge=1, le=200, description="Items per page"),
|
||||
):
|
||||
"""List all groups of exact-duplicate documents (same SHA-256 hash).
|
||||
|
||||
Returns one entry per duplicate group showing the original document and all
|
||||
files that were detected as copies of it. Groups are sorted by descending
|
||||
duplicate count.
|
||||
|
||||
Example:
|
||||
```
|
||||
GET /api/duplicates
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"groups": [
|
||||
{
|
||||
"filehash": "abc123...",
|
||||
"original": {"id": 1, "original_filename": "invoice.pdf", ...},
|
||||
"duplicates": [{"id": 5, "original_filename": "invoice_copy.pdf", ...}],
|
||||
"duplicate_count": 1
|
||||
}
|
||||
],
|
||||
"total_groups": 1,
|
||||
"total_duplicate_files": 1,
|
||||
"pagination": {...}
|
||||
}
|
||||
```
|
||||
"""
|
||||
# Find all hashes that have at least one duplicate record
|
||||
dup_hashes_query = db.query(FileRecord.filehash).filter(FileRecord.is_duplicate.is_(True)).distinct()
|
||||
total_groups = dup_hashes_query.count()
|
||||
|
||||
# Paginate hash groups
|
||||
offset = (page - 1) * per_page
|
||||
dup_hashes = [row.filehash for row in dup_hashes_query.offset(offset).limit(per_page).all()]
|
||||
|
||||
groups = []
|
||||
total_duplicate_files = 0
|
||||
|
||||
for filehash in dup_hashes:
|
||||
# Find the original (non-duplicate) record with this hash
|
||||
original = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(False))
|
||||
.order_by(FileRecord.id.asc())
|
||||
.first()
|
||||
)
|
||||
|
||||
# Find all duplicate records for this hash
|
||||
duplicates = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(True))
|
||||
.order_by(FileRecord.id.asc())
|
||||
.all()
|
||||
)
|
||||
|
||||
total_duplicate_files += len(duplicates)
|
||||
|
||||
groups.append(
|
||||
{
|
||||
"filehash": filehash,
|
||||
"original": _file_record_to_dict(original) if original else None,
|
||||
"duplicates": [_file_record_to_dict(d) for d in duplicates],
|
||||
"duplicate_count": len(duplicates),
|
||||
}
|
||||
)
|
||||
|
||||
total_pages = (total_groups + per_page - 1) // per_page if total_groups > 0 else 1
|
||||
|
||||
return {
|
||||
"groups": groups,
|
||||
"total_groups": total_groups,
|
||||
"total_duplicate_files": total_duplicate_files,
|
||||
"pagination": {
|
||||
"page": page,
|
||||
"per_page": per_page,
|
||||
"total": total_groups,
|
||||
"pages": total_pages,
|
||||
"next": str(request.url.include_query_params(page=page + 1)) if page < total_pages else None,
|
||||
"previous": str(request.url.include_query_params(page=page - 1)) if page > 1 else None,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/duplicates")
|
||||
@require_login
|
||||
def get_file_duplicates(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
near_duplicate_limit: int = Query(5, ge=1, le=20, description="Maximum near-duplicates to return"),
|
||||
near_duplicate_threshold: float = Query(
|
||||
-1.0,
|
||||
ge=-1.0,
|
||||
le=1.0,
|
||||
description="Minimum similarity score for near-duplicates; -1 uses the configured default",
|
||||
),
|
||||
):
|
||||
"""Get exact and near-duplicate documents for the specified file.
|
||||
|
||||
**Exact duplicates** share the same SHA-256 hash.
|
||||
**Near-duplicates** have a text-embedding cosine similarity score ≥
|
||||
``NEAR_DUPLICATE_THRESHOLD`` (configurable; default 0.85).
|
||||
|
||||
Near-duplicate detection requires OCR text to be available for both the
|
||||
target file and candidate files. Files without OCR text are excluded.
|
||||
|
||||
Example:
|
||||
```
|
||||
GET /api/files/42/duplicates
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"file_id": 42,
|
||||
"exact_duplicates": [
|
||||
{"id": 7, "original_filename": "invoice.pdf", "is_duplicate": true, "duplicate_of_id": 42, ...}
|
||||
],
|
||||
"near_duplicates": [
|
||||
{"file_id": 15, "original_filename": "invoice_jan.pdf", "similarity_score": 0.92, ...}
|
||||
],
|
||||
"near_duplicate_threshold": 0.85
|
||||
}
|
||||
```
|
||||
"""
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
# --- Exact duplicates ---
|
||||
# Case 1: This file is the original — find all records that are duplicates of it
|
||||
exact_duplicates_of_this = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.filehash == file_record.filehash, FileRecord.id != file_id)
|
||||
.order_by(FileRecord.id.asc())
|
||||
.all()
|
||||
)
|
||||
|
||||
# Case 2: This file itself is a duplicate — find the original
|
||||
is_self_duplicate = file_record.is_duplicate
|
||||
duplicate_of_original: FileRecord | None = None
|
||||
if is_self_duplicate and file_record.duplicate_of_id:
|
||||
duplicate_of_original = db.query(FileRecord).filter(FileRecord.id == file_record.duplicate_of_id).first()
|
||||
|
||||
exact_duplicate_dicts = [_file_record_to_dict(f) for f in exact_duplicates_of_this]
|
||||
|
||||
# --- Near-duplicates (embedding-based) ---
|
||||
effective_threshold = (
|
||||
near_duplicate_threshold if near_duplicate_threshold >= 0.0 else settings.near_duplicate_threshold
|
||||
)
|
||||
|
||||
near_duplicates: list[dict] = []
|
||||
if file_record.ocr_text and file_record.ocr_text.strip():
|
||||
try:
|
||||
from app.utils.similarity import find_similar_documents
|
||||
|
||||
near_duplicates = find_similar_documents(
|
||||
db,
|
||||
file_id,
|
||||
limit=near_duplicate_limit,
|
||||
threshold=effective_threshold,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"Near-duplicate detection failed for file {file_id}: {e}")
|
||||
near_duplicates = []
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"is_duplicate": is_self_duplicate,
|
||||
"duplicate_of": _file_record_to_dict(duplicate_of_original) if duplicate_of_original else None,
|
||||
"exact_duplicates": exact_duplicate_dicts,
|
||||
"near_duplicates": near_duplicates,
|
||||
"near_duplicate_threshold": effective_threshold,
|
||||
"exact_duplicate_count": len(exact_duplicate_dicts),
|
||||
"near_duplicate_count": len(near_duplicates),
|
||||
}
|
||||
|
||||
|
||||
def _file_record_to_dict(file_record: FileRecord | None) -> dict | None:
|
||||
"""Serialise a ``FileRecord`` to a plain dict for JSON responses."""
|
||||
if file_record is None:
|
||||
return None
|
||||
return {
|
||||
"id": file_record.id,
|
||||
"original_filename": file_record.original_filename,
|
||||
"filehash": file_record.filehash,
|
||||
"file_size": file_record.file_size,
|
||||
"mime_type": file_record.mime_type,
|
||||
"is_duplicate": file_record.is_duplicate,
|
||||
"duplicate_of_id": file_record.duplicate_of_id,
|
||||
"document_title": file_record.document_title,
|
||||
"created_at": file_record.created_at.isoformat() if file_record.created_at else None,
|
||||
}
|
||||
+292
-20
@@ -11,7 +11,7 @@ import zipfile
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile, status
|
||||
from fastapi.responses import StreamingResponse
|
||||
from sqlalchemy import asc, desc
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -23,10 +23,12 @@ from app.models import FileProcessingStep, FileRecord, ProcessingLog
|
||||
from app.tasks.convert_to_pdf import convert_to_pdf
|
||||
from app.tasks.process_document import process_document
|
||||
from app.utils.allowed_types import ALLOWED_EXTENSIONS, ALLOWED_MIME_TYPES, IMAGE_MIME_TYPES
|
||||
from app.utils.file_operations import hash_file
|
||||
from app.utils.file_queries import apply_status_filter
|
||||
from app.utils.file_status import get_files_processing_status
|
||||
from app.utils.filename_utils import sanitize_filename
|
||||
from app.utils.input_validation import validate_search_query, validate_sort_field, validate_sort_order
|
||||
from app.utils.user_scope import apply_owner_filter, get_current_owner_id
|
||||
|
||||
# Set up logging
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -49,7 +51,7 @@ def list_files_api(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
page: int = Query(1, ge=1, description="Page number"),
|
||||
per_page: int = Query(50, ge=1, le=200, description="Items per page"),
|
||||
per_page: int = Query(25, ge=1, le=200, description="Items per page"),
|
||||
sort_by: str = Query(
|
||||
"created_at",
|
||||
description="Sort field: id, original_filename, file_size, mime_type, created_at, status",
|
||||
@@ -69,7 +71,7 @@ def list_files_api(
|
||||
|
||||
Query Parameters:
|
||||
- page: Page number (default: 1)
|
||||
- per_page: Items per page (default: 50, max: 200)
|
||||
- per_page: Items per page (default: 25, max: 200)
|
||||
- sort_by: Field to sort by (default: created_at)
|
||||
- sort_order: asc or desc (default: desc)
|
||||
- search: Search in filename
|
||||
@@ -85,9 +87,11 @@ def list_files_api(
|
||||
"files": [...],
|
||||
"pagination": {
|
||||
"page": 1,
|
||||
"per_page": 50,
|
||||
"total_items": 150,
|
||||
"total_pages": 3
|
||||
"per_page": 25,
|
||||
"total": 150,
|
||||
"pages": 6,
|
||||
"next": "http://host/api/files?page=2",
|
||||
"previous": null
|
||||
}
|
||||
}
|
||||
"""
|
||||
@@ -96,8 +100,9 @@ def list_files_api(
|
||||
validate_sort_order(sort_order)
|
||||
search = validate_search_query(search)
|
||||
|
||||
# Start with base query
|
||||
# Start with base query, scoped to the current user in multi-user mode
|
||||
query = db.query(FileRecord)
|
||||
query = apply_owner_filter(query, request)
|
||||
|
||||
# Apply search filter
|
||||
if search:
|
||||
@@ -205,13 +210,19 @@ def list_files_api(
|
||||
# Calculate pagination info
|
||||
total_pages = (total_items + per_page - 1) // per_page
|
||||
|
||||
# Build next / previous page URLs by replacing the page query parameter
|
||||
next_url = str(request.url.include_query_params(page=page + 1)) if page < total_pages else None
|
||||
previous_url = str(request.url.include_query_params(page=page - 1)) if page > 1 else None
|
||||
|
||||
return {
|
||||
"files": result,
|
||||
"pagination": {
|
||||
"page": page,
|
||||
"per_page": per_page,
|
||||
"total_items": total_items,
|
||||
"total_pages": total_pages,
|
||||
"total": total_items,
|
||||
"pages": total_pages,
|
||||
"next": next_url,
|
||||
"previous": previous_url,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -232,8 +243,10 @@ def get_file_details(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
Get detailed information about a specific file including processing history.
|
||||
"""
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
# Find the file record, scoped to the current user in multi-user mode
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found")
|
||||
@@ -291,8 +304,10 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
|
||||
raise HTTPException(status_code=403, detail="File deletion is disabled in the configuration")
|
||||
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
# Find the file record, scoped to the current user in multi-user mode
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
file_record = query.first()
|
||||
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found")
|
||||
@@ -1204,7 +1219,7 @@ def download_file(
|
||||
|
||||
@router.post("/ui-upload")
|
||||
@require_login
|
||||
async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
async def ui_upload(request: Request, db: DbSession, file: UploadFile = File(...)):
|
||||
"""Endpoint to accept a user-uploaded file and enqueue it for processing."""
|
||||
workdir = settings.workdir
|
||||
|
||||
@@ -1241,6 +1256,23 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
# Store both the safe original name and the unique name
|
||||
target_path = os.path.join(workdir, target_filename)
|
||||
|
||||
# Determine the owner_id for multi-user document isolation
|
||||
upload_owner_id = get_current_owner_id(request) if settings.multi_user_enabled else None
|
||||
|
||||
# Enforce subscription tier upload quotas (multi-user mode only) BEFORE writing the file
|
||||
# so that users who have exceeded their quota do not waste bandwidth or disk I/O.
|
||||
if settings.multi_user_enabled and upload_owner_id:
|
||||
from app.utils.subscription import QuotaExceeded, check_upload_allowed, get_user_tier_id
|
||||
|
||||
tier_id = get_user_tier_id(db, upload_owner_id)
|
||||
try:
|
||||
check_upload_allowed(db, upload_owner_id, tier_id)
|
||||
except QuotaExceeded as qe:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_402_PAYMENT_REQUIRED,
|
||||
detail=str(qe),
|
||||
)
|
||||
|
||||
# Read file in chunks to avoid loading the entire body into memory at once,
|
||||
# enforcing the size limit during the read so memory usage stays bounded.
|
||||
try:
|
||||
@@ -1301,7 +1333,7 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
task_ids = []
|
||||
for split_file in split_files:
|
||||
split_filename = os.path.basename(split_file)
|
||||
task = process_document.delay(split_file, original_filename=split_filename)
|
||||
task = process_document.delay(split_file, original_filename=split_filename, owner_id=upload_owner_id)
|
||||
task_ids.append(task.id)
|
||||
logger.info(f"Enqueued split PDF part for processing: {split_file}")
|
||||
|
||||
@@ -1324,7 +1356,7 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
|
||||
if is_pdf and not should_split:
|
||||
# If it's a PDF, process directly
|
||||
task = process_document.delay(target_path, original_filename=safe_filename)
|
||||
task = process_document.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
|
||||
logger.info(f"Enqueued PDF for processing: {target_path}")
|
||||
elif mime_type in IMAGE_MIME_TYPES or file_ext in {
|
||||
".jpg",
|
||||
@@ -1338,20 +1370,260 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
".svg",
|
||||
}:
|
||||
# If it's an image, convert to PDF first
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
|
||||
logger.info(f"Enqueued image for PDF conversion: {target_path}")
|
||||
elif mime_type in ALLOWED_MIME_TYPES or file_ext in ALLOWED_EXTENSIONS:
|
||||
# Office document, HTML, Markdown, or other Gotenberg-supported format
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
|
||||
logger.info(f"Enqueued document for PDF conversion: {target_path}")
|
||||
else:
|
||||
# For any other file type, attempt conversion but log a warning
|
||||
logger.warning(f"Unsupported MIME type {mime_type} for {target_path}, attempting conversion")
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename)
|
||||
task = convert_to_pdf.delay(target_path, original_filename=safe_filename, owner_id=upload_owner_id)
|
||||
|
||||
return {
|
||||
# Check for exact duplicates (same SHA-256 hash) before returning.
|
||||
# This gives the caller an immediate warning without waiting for the pipeline.
|
||||
# Only performed when deduplication is enabled in settings.
|
||||
exact_duplicate_warning = None
|
||||
if settings.enable_deduplication:
|
||||
try:
|
||||
filehash = hash_file(target_path)
|
||||
existing = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.filehash == filehash, FileRecord.is_duplicate.is_(False))
|
||||
.order_by(FileRecord.id.asc())
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
exact_duplicate_warning = {
|
||||
"duplicate_type": "exact",
|
||||
"original_file_id": existing.id,
|
||||
"original_filename": existing.original_filename,
|
||||
"message": (
|
||||
"This file appears to be an exact duplicate of an already-processed document. "
|
||||
"It will still be queued but will be flagged as a duplicate."
|
||||
),
|
||||
}
|
||||
logger.info(f"Exact duplicate detected on upload: '{safe_filename}' matches file ID {existing.id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Duplicate check failed for uploaded file '{safe_filename}': {e}")
|
||||
|
||||
response: dict = {
|
||||
"task_id": task.id,
|
||||
"status": "queued",
|
||||
"original_filename": safe_filename,
|
||||
"stored_filename": target_filename,
|
||||
}
|
||||
if exact_duplicate_warning:
|
||||
response["duplicate_warning"] = exact_duplicate_warning
|
||||
return response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Document ownership / claim endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/files/{file_id}/claim")
|
||||
@require_login
|
||||
def claim_file(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
Claim an unowned document for the current user.
|
||||
|
||||
Only documents with ``owner_id IS NULL`` can be claimed. The requesting
|
||||
user's identifier is written into ``owner_id``. In single-user mode
|
||||
the endpoint is a no-op (returns the file unchanged).
|
||||
"""
|
||||
if not settings.multi_user_enabled:
|
||||
raise HTTPException(status_code=400, detail="Multi-user mode is not enabled")
|
||||
|
||||
owner_id = get_current_owner_id(request)
|
||||
if owner_id is None:
|
||||
raise HTTPException(status_code=401, detail="Authentication required to claim a document")
|
||||
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=404, detail=f"File record with ID {file_id} not found")
|
||||
|
||||
if file_record.owner_id is not None:
|
||||
if file_record.owner_id == owner_id:
|
||||
return {"status": "already_owned", "message": "You already own this document", "file_id": file_id}
|
||||
raise HTTPException(status_code=403, detail="This document is already owned by another user")
|
||||
|
||||
file_record.owner_id = owner_id
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.exception(f"Error claiming file {file_id}: {e}")
|
||||
raise HTTPException(status_code=500, detail="Failed to claim document")
|
||||
|
||||
logger.info(f"File {file_id} claimed by user '{owner_id}'")
|
||||
return {"status": "success", "message": "Document claimed successfully", "file_id": file_id, "owner_id": owner_id}
|
||||
|
||||
|
||||
@router.post("/files/bulk-claim")
|
||||
@require_login
|
||||
def bulk_claim_files(request: Request, file_ids: list[int], db: DbSession):
|
||||
"""
|
||||
Claim multiple unowned documents for the current user.
|
||||
|
||||
Only documents with ``owner_id IS NULL`` will be claimed. Documents
|
||||
already owned (by anyone) are skipped and reported in ``skipped``.
|
||||
"""
|
||||
if not settings.multi_user_enabled:
|
||||
raise HTTPException(status_code=400, detail="Multi-user mode is not enabled")
|
||||
|
||||
owner_id = get_current_owner_id(request)
|
||||
if owner_id is None:
|
||||
raise HTTPException(status_code=401, detail="Authentication required to claim documents")
|
||||
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
|
||||
claimed = []
|
||||
skipped = []
|
||||
for rec in file_records:
|
||||
if rec.owner_id is None:
|
||||
rec.owner_id = owner_id
|
||||
claimed.append(rec.id)
|
||||
else:
|
||||
skipped.append({"file_id": rec.id, "reason": "already owned"})
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.exception(f"Error during bulk claim: {e}")
|
||||
raise HTTPException(status_code=500, detail="Failed to claim documents")
|
||||
|
||||
logger.info(f"Bulk claim by '{owner_id}': claimed={claimed}, skipped={[s['file_id'] for s in skipped]}")
|
||||
return {
|
||||
"status": "success",
|
||||
"claimed_count": len(claimed),
|
||||
"claimed_ids": claimed,
|
||||
"skipped": skipped,
|
||||
"owner_id": owner_id,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/files/assign-owner")
|
||||
@require_login
|
||||
def assign_owner(request: Request, db: DbSession, owner_id: str = Query(...), file_ids: list[int] | None = None):
|
||||
"""
|
||||
Admin-only: assign an owner to documents.
|
||||
|
||||
If ``file_ids`` is provided, only those files are updated. If omitted,
|
||||
**all** currently unowned documents (``owner_id IS NULL``) are assigned
|
||||
to the given ``owner_id``.
|
||||
"""
|
||||
if not settings.multi_user_enabled:
|
||||
raise HTTPException(status_code=400, detail="Multi-user mode is not enabled")
|
||||
|
||||
user = request.session.get("user")
|
||||
if not isinstance(user, dict) or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=403, detail="Only admins can assign document owners")
|
||||
|
||||
if not owner_id or not owner_id.strip():
|
||||
raise HTTPException(status_code=422, detail="owner_id must be a non-empty string")
|
||||
owner_id = owner_id.strip()
|
||||
|
||||
if file_ids is not None:
|
||||
# Assign to specific files
|
||||
updated = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.id.in_(file_ids))
|
||||
.update({FileRecord.owner_id: owner_id}, synchronize_session="fetch")
|
||||
)
|
||||
else:
|
||||
# Assign to all currently unowned documents
|
||||
updated = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.owner_id.is_(None))
|
||||
.update({FileRecord.owner_id: owner_id}, synchronize_session="fetch")
|
||||
)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.exception(f"Error assigning owner: {e}")
|
||||
raise HTTPException(status_code=500, detail="Failed to assign owner")
|
||||
|
||||
admin_name = get_current_owner_id(request) or "admin"
|
||||
logger.info(f"Admin '{admin_name}' assigned owner_id='{owner_id}' to {updated} file(s)")
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Assigned owner to {updated} document(s)",
|
||||
"updated_count": updated,
|
||||
"owner_id": owner_id,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pipeline assignment
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/files/{file_id}/assign-pipeline")
|
||||
@require_login
|
||||
def assign_pipeline_to_file(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
pipeline_id: int | None = None,
|
||||
):
|
||||
"""Assign (or remove) a processing pipeline from a file.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The file to update.
|
||||
|
||||
Query / Body Parameters:
|
||||
pipeline_id: The pipeline to assign. Pass ``null`` or omit to clear the
|
||||
assignment (the system default will be used for future processing).
|
||||
|
||||
Returns:
|
||||
A summary dict with the file_id and updated pipeline_id.
|
||||
|
||||
Raises:
|
||||
HTTPException 404: If the file or pipeline does not exist / is not
|
||||
accessible to the current user.
|
||||
"""
|
||||
from app.auth import get_current_user, get_current_user_id
|
||||
from app.models import Pipeline
|
||||
|
||||
user = get_current_user(request)
|
||||
user_id: str = get_current_user_id(request)
|
||||
|
||||
is_admin_user = bool(user and user.get("is_admin"))
|
||||
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
# Non-admins may only update files they own (or unowned files in single-user mode)
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not is_admin_user and file_record.owner_id is not None and file_record.owner_id != owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
if pipeline_id is not None:
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
# Check access: users can only assign their own pipelines or system pipelines (owner_id=None)
|
||||
if not is_admin_user and pipeline.owner_id is not None and pipeline.owner_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
file_record.pipeline_id = pipeline_id
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(file_record)
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to assign pipeline to file id={file_id}: {exc}")
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to assign pipeline")
|
||||
|
||||
logger.info(f"Pipeline {pipeline_id!r} assigned to file id={file_id}")
|
||||
return {"file_id": file_id, "pipeline_id": file_record.pipeline_id}
|
||||
|
||||
@@ -0,0 +1,431 @@
|
||||
"""
|
||||
GraphQL API endpoint for DocuElevate.
|
||||
|
||||
Provides a flexible query interface alongside the existing REST API.
|
||||
Schema covers: documents, pipelines, settings, and users.
|
||||
|
||||
Endpoint: /graphql
|
||||
GraphiQL playground: /graphql (via browser)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any
|
||||
|
||||
import strawberry
|
||||
from fastapi import Depends, Request
|
||||
from sqlalchemy.orm import Session
|
||||
from strawberry.fastapi import GraphQLRouter
|
||||
|
||||
from app.auth import get_current_user
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import ApplicationSettings, FileRecord, Pipeline, PipelineStep, UserProfile
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Strawberry types
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class DocumentType:
|
||||
"""A processed document stored in the system."""
|
||||
|
||||
id: int
|
||||
owner_id: str | None
|
||||
original_filename: str | None
|
||||
local_filename: str
|
||||
file_size: int
|
||||
mime_type: str | None
|
||||
document_title: str | None
|
||||
is_duplicate: bool
|
||||
ocr_quality_score: int | None
|
||||
pipeline_id: int | None
|
||||
created_at: datetime | None
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class PipelineStepType:
|
||||
"""A single step within a processing pipeline."""
|
||||
|
||||
id: int
|
||||
pipeline_id: int
|
||||
position: int
|
||||
step_type: str
|
||||
label: str | None
|
||||
enabled: bool
|
||||
created_at: datetime | None
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class PipelineType:
|
||||
"""A processing pipeline with its ordered steps."""
|
||||
|
||||
id: int
|
||||
owner_id: str | None
|
||||
name: str
|
||||
description: str | None
|
||||
is_default: bool
|
||||
is_active: bool
|
||||
steps: list[PipelineStepType]
|
||||
created_at: datetime | None
|
||||
updated_at: datetime | None
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class SettingType:
|
||||
"""An application configuration setting stored in the database."""
|
||||
|
||||
id: int
|
||||
key: str
|
||||
value: str | None
|
||||
created_at: datetime | None
|
||||
updated_at: datetime | None
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class UserType:
|
||||
"""A user profile in the system."""
|
||||
|
||||
id: int
|
||||
user_id: str
|
||||
display_name: str | None
|
||||
is_blocked: bool
|
||||
subscription_tier: str | None
|
||||
onboarding_completed: bool
|
||||
created_at: datetime | None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Conversion helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _document_from_record(rec: FileRecord) -> DocumentType:
|
||||
return DocumentType(
|
||||
id=rec.id,
|
||||
owner_id=rec.owner_id,
|
||||
original_filename=rec.original_filename,
|
||||
local_filename=rec.local_filename,
|
||||
file_size=rec.file_size,
|
||||
mime_type=rec.mime_type,
|
||||
document_title=rec.document_title,
|
||||
is_duplicate=rec.is_duplicate,
|
||||
ocr_quality_score=rec.ocr_quality_score,
|
||||
pipeline_id=rec.pipeline_id,
|
||||
created_at=rec.created_at,
|
||||
)
|
||||
|
||||
|
||||
def _pipeline_step_from_record(step: PipelineStep) -> PipelineStepType:
|
||||
return PipelineStepType(
|
||||
id=step.id,
|
||||
pipeline_id=step.pipeline_id,
|
||||
position=step.position,
|
||||
step_type=step.step_type,
|
||||
label=step.label,
|
||||
enabled=step.enabled,
|
||||
created_at=step.created_at,
|
||||
)
|
||||
|
||||
|
||||
def _pipeline_from_record(pipeline: Pipeline, db: Session) -> PipelineType:
|
||||
steps = db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline.id).order_by(PipelineStep.position).all()
|
||||
return PipelineType(
|
||||
id=pipeline.id,
|
||||
owner_id=pipeline.owner_id,
|
||||
name=pipeline.name,
|
||||
description=pipeline.description,
|
||||
is_default=pipeline.is_default,
|
||||
is_active=pipeline.is_active,
|
||||
steps=[_pipeline_step_from_record(s) for s in steps],
|
||||
created_at=pipeline.created_at,
|
||||
updated_at=pipeline.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _setting_from_record(setting: ApplicationSettings) -> SettingType:
|
||||
return SettingType(
|
||||
id=setting.id,
|
||||
key=setting.key,
|
||||
value=setting.value,
|
||||
created_at=setting.created_at,
|
||||
updated_at=setting.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _user_from_profile(profile: UserProfile) -> UserType:
|
||||
return UserType(
|
||||
id=profile.id,
|
||||
user_id=profile.user_id,
|
||||
display_name=profile.display_name,
|
||||
is_blocked=profile.is_blocked,
|
||||
subscription_tier=profile.subscription_tier,
|
||||
onboarding_completed=profile.onboarding_completed,
|
||||
created_at=profile.created_at,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Context helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Keys that contain sensitive data and must never be returned via GraphQL
|
||||
_SENSITIVE_SETTING_KEYS: frozenset[str] = frozenset(
|
||||
{
|
||||
"openai_api_key",
|
||||
"azure_ai_key",
|
||||
"session_secret",
|
||||
"database_url",
|
||||
"redis_url",
|
||||
"dropbox_app_secret",
|
||||
"dropbox_refresh_token",
|
||||
"google_drive_credentials_json",
|
||||
"onedrive_client_secret",
|
||||
"onedrive_refresh_token",
|
||||
"smtp_password",
|
||||
"nextcloud_password",
|
||||
"s3_secret_access_key",
|
||||
"ftp_password",
|
||||
"sftp_password",
|
||||
"webdav_password",
|
||||
"stripe_secret_key",
|
||||
"stripe_webhook_secret",
|
||||
"sentry_dsn",
|
||||
"social_auth_google_client_secret",
|
||||
"social_auth_microsoft_client_secret",
|
||||
"social_auth_apple_private_key",
|
||||
"social_auth_dropbox_app_secret",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _get_current_user_id(user: dict[str, Any] | None) -> str | None:
|
||||
"""Extract the stable user identifier from the user dict."""
|
||||
if not user:
|
||||
return None
|
||||
return user.get("preferred_username") or user.get("email") or user.get("id") or None
|
||||
|
||||
|
||||
def _get_db_and_user(info: strawberry.types.Info) -> tuple[Session, dict[str, Any] | None]:
|
||||
"""Extract the database session and current user from the Strawberry context."""
|
||||
db: Session = info.context["db"]
|
||||
user: dict[str, Any] | None = info.context.get("user")
|
||||
return db, user
|
||||
|
||||
|
||||
def _require_auth(user: dict[str, Any] | None) -> None:
|
||||
"""Raise an error when authentication is enabled and no valid user is present."""
|
||||
if settings.auth_enabled and not user:
|
||||
raise strawberry.exceptions.StrawberryGraphQLError("Authentication required")
|
||||
|
||||
|
||||
def _require_admin(user: dict[str, Any] | None) -> None:
|
||||
"""Raise an error when the current user is not an admin.
|
||||
|
||||
When ``auth_enabled`` is *False* (single-user / development mode) all
|
||||
callers are implicitly treated as administrators.
|
||||
"""
|
||||
if not settings.auth_enabled:
|
||||
# Single-user mode: no auth, treat caller as admin
|
||||
return
|
||||
_require_auth(user)
|
||||
if not (user and user.get("is_admin")):
|
||||
raise strawberry.exceptions.StrawberryGraphQLError("Admin access required")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Query resolvers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@strawberry.type
|
||||
class Query:
|
||||
"""Root query type for the DocuElevate GraphQL API."""
|
||||
|
||||
@strawberry.field(description="List documents, optionally filtered by owner.")
|
||||
def documents(
|
||||
self,
|
||||
info: strawberry.types.Info,
|
||||
owner_id: str | None = None,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> list[DocumentType]:
|
||||
"""Return a paginated list of documents.
|
||||
|
||||
When *auth_enabled* the caller must be authenticated. Non-admin users
|
||||
receive only their own documents; admins may query any *owner_id*.
|
||||
"""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_auth(user)
|
||||
|
||||
limit = max(1, min(limit, 100))
|
||||
offset = max(0, offset)
|
||||
|
||||
query = db.query(FileRecord)
|
||||
|
||||
if settings.auth_enabled and user:
|
||||
is_admin = user.get("is_admin", False)
|
||||
current_user_id = _get_current_user_id(user)
|
||||
if not is_admin:
|
||||
# Non-admins can only see their own documents
|
||||
query = query.filter(FileRecord.owner_id == current_user_id)
|
||||
elif owner_id:
|
||||
query = query.filter(FileRecord.owner_id == owner_id)
|
||||
elif owner_id:
|
||||
query = query.filter(FileRecord.owner_id == owner_id)
|
||||
|
||||
records = query.order_by(FileRecord.created_at.desc()).offset(offset).limit(limit).all()
|
||||
return [_document_from_record(r) for r in records]
|
||||
|
||||
@strawberry.field(description="Fetch a single document by ID.")
|
||||
def document(self, info: strawberry.types.Info, id: int) -> DocumentType | None:
|
||||
"""Return one document by its primary key, or *null* if not found."""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_auth(user)
|
||||
|
||||
rec = db.query(FileRecord).filter(FileRecord.id == id).first()
|
||||
if rec is None:
|
||||
return None
|
||||
|
||||
if settings.auth_enabled and user:
|
||||
is_admin = user.get("is_admin", False)
|
||||
current_user_id = _get_current_user_id(user)
|
||||
if not is_admin and rec.owner_id != current_user_id:
|
||||
return None
|
||||
|
||||
return _document_from_record(rec)
|
||||
|
||||
@strawberry.field(description="List processing pipelines.")
|
||||
def pipelines(
|
||||
self,
|
||||
info: strawberry.types.Info,
|
||||
owner_id: str | None = None,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> list[PipelineType]:
|
||||
"""Return a paginated list of pipelines."""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_auth(user)
|
||||
|
||||
limit = max(1, min(limit, 100))
|
||||
offset = max(0, offset)
|
||||
|
||||
query = db.query(Pipeline)
|
||||
|
||||
if settings.auth_enabled and user:
|
||||
is_admin = user.get("is_admin", False)
|
||||
current_user_id = _get_current_user_id(user)
|
||||
if not is_admin:
|
||||
query = query.filter((Pipeline.owner_id == current_user_id) | (Pipeline.owner_id.is_(None)))
|
||||
elif owner_id:
|
||||
query = query.filter(Pipeline.owner_id == owner_id)
|
||||
elif owner_id:
|
||||
query = query.filter(Pipeline.owner_id == owner_id)
|
||||
|
||||
rows = query.order_by(Pipeline.id).offset(offset).limit(limit).all()
|
||||
return [_pipeline_from_record(p, db) for p in rows]
|
||||
|
||||
@strawberry.field(description="Fetch a single pipeline by ID.")
|
||||
def pipeline(self, info: strawberry.types.Info, id: int) -> PipelineType | None:
|
||||
"""Return one pipeline by its primary key, or *null* if not found."""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_auth(user)
|
||||
|
||||
row = db.query(Pipeline).filter(Pipeline.id == id).first()
|
||||
if row is None:
|
||||
return None
|
||||
|
||||
if settings.auth_enabled and user:
|
||||
is_admin = user.get("is_admin", False)
|
||||
current_user_id = _get_current_user_id(user)
|
||||
if not is_admin and row.owner_id is not None and row.owner_id != current_user_id:
|
||||
return None
|
||||
|
||||
return _pipeline_from_record(row, db)
|
||||
|
||||
@strawberry.field(description="List non-sensitive application settings (admin only).")
|
||||
def settings(
|
||||
self,
|
||||
info: strawberry.types.Info,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
) -> list[SettingType]:
|
||||
"""Return application settings stored in the database.
|
||||
|
||||
Sensitive keys (API secrets, passwords, etc.) are automatically
|
||||
excluded. Requires admin privileges when auth is enabled.
|
||||
"""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_admin(user)
|
||||
|
||||
limit = max(1, min(limit, 200))
|
||||
offset = max(0, offset)
|
||||
|
||||
rows = (
|
||||
db.query(ApplicationSettings)
|
||||
.filter(ApplicationSettings.key.notin_(_SENSITIVE_SETTING_KEYS))
|
||||
.order_by(ApplicationSettings.key)
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [_setting_from_record(r) for r in rows]
|
||||
|
||||
@strawberry.field(description="List user profiles (admin only).")
|
||||
def users(
|
||||
self,
|
||||
info: strawberry.types.Info,
|
||||
limit: int = 20,
|
||||
offset: int = 0,
|
||||
) -> list[UserType]:
|
||||
"""Return a paginated list of user profiles. Requires admin privileges."""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_admin(user)
|
||||
|
||||
limit = max(1, min(limit, 100))
|
||||
offset = max(0, offset)
|
||||
|
||||
rows = db.query(UserProfile).order_by(UserProfile.user_id).offset(offset).limit(limit).all()
|
||||
return [_user_from_profile(r) for r in rows]
|
||||
|
||||
@strawberry.field(description="Fetch a user profile by user_id (admin only).")
|
||||
def user(self, info: strawberry.types.Info, user_id: str) -> UserType | None:
|
||||
"""Return one user profile by *user_id*, or *null* if not found."""
|
||||
db, user = _get_db_and_user(info)
|
||||
_require_admin(user)
|
||||
|
||||
row = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
return _user_from_profile(row) if row else None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Schema and router
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
schema = strawberry.Schema(query=Query)
|
||||
|
||||
|
||||
async def get_graphql_context(
|
||||
request: Request,
|
||||
db: Annotated[Session, Depends(get_db)],
|
||||
) -> dict[str, Any]:
|
||||
"""Build the per-request context injected into every resolver."""
|
||||
try:
|
||||
user = get_current_user(request)
|
||||
except Exception:
|
||||
logger.debug("Could not resolve current user for GraphQL context", exc_info=True)
|
||||
user = None
|
||||
return {"request": request, "db": db, "user": user}
|
||||
|
||||
|
||||
graphql_router = GraphQLRouter(
|
||||
schema,
|
||||
context_getter=get_graphql_context,
|
||||
graphql_ide="graphiql",
|
||||
)
|
||||
+136
@@ -0,0 +1,136 @@
|
||||
"""API endpoints for internationalization (i18n).
|
||||
|
||||
Provides endpoints for:
|
||||
* Listing available languages
|
||||
* Getting/setting user language preference (persisted in session + cookie + DB)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, Depends, Request, Response
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import UserProfile
|
||||
from app.utils.i18n import (
|
||||
DEFAULT_LANGUAGE,
|
||||
SUPPORTED_LANGUAGE_CODES,
|
||||
SUPPORTED_LANGUAGES,
|
||||
detect_language,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/i18n", tags=["i18n"])
|
||||
|
||||
|
||||
class LanguageInfo(BaseModel):
|
||||
"""Schema for a supported language."""
|
||||
|
||||
code: str
|
||||
name: str
|
||||
native: str
|
||||
flag: str
|
||||
|
||||
|
||||
class LanguageListResponse(BaseModel):
|
||||
"""Response for the list-languages endpoint."""
|
||||
|
||||
languages: list[LanguageInfo]
|
||||
current: str
|
||||
default: str
|
||||
|
||||
|
||||
class SetLanguageRequest(BaseModel):
|
||||
"""Request body for setting the preferred language."""
|
||||
|
||||
language: str
|
||||
|
||||
|
||||
class SetLanguageResponse(BaseModel):
|
||||
"""Response after changing the language."""
|
||||
|
||||
language: str
|
||||
message: str
|
||||
|
||||
|
||||
@router.get("/languages", response_model=LanguageListResponse)
|
||||
async def list_languages(request: Request) -> LanguageListResponse:
|
||||
"""Return all supported UI languages and the current active language."""
|
||||
current = detect_language(request)
|
||||
return LanguageListResponse(
|
||||
languages=[LanguageInfo(**lang) for lang in SUPPORTED_LANGUAGES],
|
||||
current=current,
|
||||
default=DEFAULT_LANGUAGE,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/language", response_model=SetLanguageResponse)
|
||||
async def set_language(
|
||||
body: SetLanguageRequest,
|
||||
request: Request,
|
||||
response: Response,
|
||||
db: Session = Depends(get_db),
|
||||
) -> SetLanguageResponse:
|
||||
"""Set the preferred UI language.
|
||||
|
||||
Persists the choice in:
|
||||
1. The server-side session
|
||||
2. A ``docuelevate_lang`` cookie (30-day expiry)
|
||||
3. The ``UserProfile.preferred_language`` column (if authenticated)
|
||||
"""
|
||||
lang = body.language.lower().strip()
|
||||
if lang not in SUPPORTED_LANGUAGE_CODES:
|
||||
lang = DEFAULT_LANGUAGE
|
||||
|
||||
# 1. Session
|
||||
if hasattr(request, "session"):
|
||||
request.session["preferred_language"] = lang
|
||||
|
||||
# 2. Cookie (30 days)
|
||||
response.set_cookie(
|
||||
key="docuelevate_lang",
|
||||
value=lang,
|
||||
max_age=30 * 24 * 60 * 60,
|
||||
httponly=False,
|
||||
samesite="lax",
|
||||
)
|
||||
|
||||
# 3. Database (if user is authenticated)
|
||||
_persist_language_to_profile(request, db, lang)
|
||||
|
||||
language_name = next(
|
||||
(entry["native"] for entry in SUPPORTED_LANGUAGES if entry["code"] == lang),
|
||||
lang,
|
||||
)
|
||||
logger.info("Language preference set to '%s'", lang)
|
||||
return SetLanguageResponse(
|
||||
language=lang,
|
||||
message=f"Language changed to {language_name}",
|
||||
)
|
||||
|
||||
|
||||
def _persist_language_to_profile(request: Request, db: Session, lang: str) -> None:
|
||||
"""Write language preference to the UserProfile row, if the user is logged in."""
|
||||
user_id: str | None = None
|
||||
if hasattr(request, "session"):
|
||||
user = request.session.get("user")
|
||||
if isinstance(user, dict):
|
||||
user_id = user.get("preferred_username") or user.get("email") or user.get("id")
|
||||
elif isinstance(user, str):
|
||||
user_id = user
|
||||
|
||||
if not user_id:
|
||||
return
|
||||
|
||||
try:
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile:
|
||||
profile.preferred_language = lang # type: ignore[attr-defined]
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.debug("Could not persist language preference for user_id=%s", user_id)
|
||||
@@ -0,0 +1,371 @@
|
||||
"""API endpoints for managing per-user IMAP ingestion accounts.
|
||||
|
||||
Provides CRUD operations for a user's IMAP accounts, quota enforcement
|
||||
against their subscription plan's ``max_mailboxes`` limit, and a
|
||||
test-connection endpoint so users can verify credentials before saving.
|
||||
"""
|
||||
|
||||
import imaplib
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import UserImapAccount
|
||||
from app.utils.encryption import decrypt_value, encrypt_value
|
||||
from app.utils.subscription import get_tier, get_user_tier_id
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/imap-accounts", tags=["imap-accounts"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if owner_id is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Quota helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_FREE_TIER_ID = "free"
|
||||
|
||||
|
||||
def _get_max_mailboxes(tier: dict[str, Any]) -> int | None:
|
||||
"""Return the maximum number of IMAP accounts allowed by *tier*.
|
||||
|
||||
Returns:
|
||||
``None`` — unlimited (paid tiers with ``max_mailboxes == 0``)
|
||||
``0`` — no mailboxes allowed (free tier)
|
||||
positive — the configured limit
|
||||
"""
|
||||
tier_id: str = tier.get("id", _FREE_TIER_ID)
|
||||
max_mb: int = tier.get("max_mailboxes", 0)
|
||||
|
||||
# Free tier: 0 means "no access" (not "unlimited")
|
||||
if tier_id == _FREE_TIER_ID:
|
||||
return 0
|
||||
|
||||
# Paid tiers: 0 means unlimited
|
||||
if max_mb == 0:
|
||||
return None
|
||||
|
||||
return max_mb
|
||||
|
||||
|
||||
def _check_quota(db: Session, owner_id: str) -> None:
|
||||
"""Raise 403 if the user has reached their IMAP account quota."""
|
||||
tier_id = get_user_tier_id(db, owner_id)
|
||||
tier = get_tier(tier_id, db)
|
||||
max_mb = _get_max_mailboxes(tier)
|
||||
|
||||
if max_mb == 0:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=("Your current plan does not include email ingestion. Upgrade to a paid plan to add IMAP accounts."),
|
||||
)
|
||||
|
||||
if max_mb is not None:
|
||||
current_count = db.query(UserImapAccount).filter(UserImapAccount.owner_id == owner_id).count()
|
||||
if current_count >= max_mb:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=(
|
||||
f"You have reached your plan limit of {max_mb} IMAP account(s). "
|
||||
"Please delete an existing account or upgrade your plan."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ImapAccountCreate(BaseModel):
|
||||
"""Schema for creating a new IMAP account."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255, description="Human-readable label")
|
||||
host: str = Field(..., min_length=1, max_length=255, description="IMAP server hostname")
|
||||
port: int = Field(default=993, ge=1, le=65535, description="IMAP server port")
|
||||
username: str = Field(..., min_length=1, max_length=255, description="IMAP login username")
|
||||
password: str = Field(..., min_length=1, max_length=1024, description="IMAP login password")
|
||||
use_ssl: bool = Field(default=True, description="Use SSL/TLS connection")
|
||||
delete_after_process: bool = Field(default=False, description="Delete emails from mailbox after processing")
|
||||
is_active: bool = Field(default=True, description="Whether to poll this mailbox")
|
||||
profile_id: int | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ID of the ImapIngestionProfile that controls which attachment types to ingest. "
|
||||
"Null inherits the global imap_attachment_filter setting."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ImapAccountUpdate(BaseModel):
|
||||
"""Schema for updating an existing IMAP account (all fields optional)."""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
host: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
port: int | None = Field(default=None, ge=1, le=65535)
|
||||
username: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
password: str | None = Field(default=None, min_length=1, max_length=1024)
|
||||
use_ssl: bool | None = None
|
||||
delete_after_process: bool | None = None
|
||||
is_active: bool | None = None
|
||||
profile_id: int | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"ID of the ImapIngestionProfile to use. "
|
||||
"Explicitly sending null clears the override (falls back to global setting)."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class ImapTestRequest(BaseModel):
|
||||
"""Schema for testing an IMAP connection without saving it."""
|
||||
|
||||
host: str = Field(..., min_length=1, max_length=255)
|
||||
port: int = Field(default=993, ge=1, le=65535)
|
||||
username: str = Field(..., min_length=1, max_length=255)
|
||||
password: str = Field(..., min_length=1, max_length=1024)
|
||||
use_ssl: bool = Field(default=True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialisation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _to_response(acct: UserImapAccount) -> dict[str, Any]:
|
||||
"""Serialize a ``UserImapAccount`` row to a response dict.
|
||||
|
||||
Passwords are never included in responses.
|
||||
"""
|
||||
return {
|
||||
"id": acct.id,
|
||||
"owner_id": acct.owner_id,
|
||||
"name": acct.name,
|
||||
"host": acct.host,
|
||||
"port": acct.port,
|
||||
"username": acct.username,
|
||||
"use_ssl": acct.use_ssl,
|
||||
"delete_after_process": acct.delete_after_process,
|
||||
"is_active": acct.is_active,
|
||||
"profile_id": acct.profile_id,
|
||||
"last_checked_at": acct.last_checked_at.isoformat() if acct.last_checked_at else None,
|
||||
"last_error": acct.last_error,
|
||||
"created_at": acct.created_at.isoformat() if acct.created_at else None,
|
||||
"updated_at": acct.updated_at.isoformat() if acct.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Connection test helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _test_imap_connection(host: str, port: int, username: str, password: str, use_ssl: bool) -> dict[str, Any]:
|
||||
"""Attempt to connect and log in to the IMAP server.
|
||||
|
||||
Returns a dict with ``{"success": bool, "message": str}``.
|
||||
"""
|
||||
try:
|
||||
if use_ssl:
|
||||
mail = imaplib.IMAP4_SSL(host, port)
|
||||
else:
|
||||
mail = imaplib.IMAP4(host, port)
|
||||
|
||||
mail.login(username, password)
|
||||
mail.logout()
|
||||
return {"success": True, "message": "Connection successful"}
|
||||
except OSError as exc:
|
||||
logger.warning("IMAP network error for %s@%s: %s", username, host, exc)
|
||||
return {"success": False, "message": f"Connection error: {exc}"}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("IMAP error for %s@%s: %s", username, host, exc)
|
||||
return {"success": False, "message": f"IMAP error: {exc}"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", summary="List IMAP accounts for the current user")
|
||||
def list_imap_accounts(request: Request, db: DbSession, owner_id: CurrentOwner) -> list[dict[str, Any]]:
|
||||
"""Return all IMAP accounts belonging to the authenticated user."""
|
||||
accounts = db.query(UserImapAccount).filter(UserImapAccount.owner_id == owner_id).order_by(UserImapAccount.id).all()
|
||||
return [_to_response(a) for a in accounts]
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED, summary="Create a new IMAP account")
|
||||
def create_imap_account(
|
||||
request: Request, body: ImapAccountCreate, db: DbSession, owner_id: CurrentOwner
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new IMAP ingestion account for the current user.
|
||||
|
||||
Quota is enforced against the user's subscription plan's ``max_mailboxes``
|
||||
limit before the account is persisted.
|
||||
"""
|
||||
_check_quota(db, owner_id)
|
||||
|
||||
acct = UserImapAccount(
|
||||
owner_id=owner_id,
|
||||
name=body.name,
|
||||
host=body.host,
|
||||
port=body.port,
|
||||
username=body.username,
|
||||
password=encrypt_value(body.password),
|
||||
use_ssl=body.use_ssl,
|
||||
delete_after_process=body.delete_after_process,
|
||||
is_active=body.is_active,
|
||||
profile_id=body.profile_id,
|
||||
)
|
||||
try:
|
||||
db.add(acct)
|
||||
db.commit()
|
||||
db.refresh(acct)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s created IMAP account %d (%s)", owner_id, acct.id, body.host)
|
||||
return _to_response(acct)
|
||||
|
||||
|
||||
@router.get("/{account_id}", summary="Get a single IMAP account")
|
||||
def get_imap_account(account_id: int, request: Request, db: DbSession, owner_id: CurrentOwner) -> dict[str, Any]:
|
||||
"""Return a single IMAP account by ID (must belong to the current user)."""
|
||||
acct = (
|
||||
db.query(UserImapAccount).filter(UserImapAccount.id == account_id, UserImapAccount.owner_id == owner_id).first()
|
||||
)
|
||||
if not acct:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="IMAP account not found")
|
||||
return _to_response(acct)
|
||||
|
||||
|
||||
@router.put("/{account_id}", summary="Update an IMAP account")
|
||||
def update_imap_account(
|
||||
account_id: int,
|
||||
request: Request,
|
||||
body: ImapAccountUpdate,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Update an existing IMAP account. Only provided fields are changed."""
|
||||
acct = (
|
||||
db.query(UserImapAccount).filter(UserImapAccount.id == account_id, UserImapAccount.owner_id == owner_id).first()
|
||||
)
|
||||
if not acct:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="IMAP account not found")
|
||||
|
||||
if body.name is not None:
|
||||
acct.name = body.name
|
||||
if body.host is not None:
|
||||
acct.host = body.host
|
||||
if body.port is not None:
|
||||
acct.port = body.port
|
||||
if body.username is not None:
|
||||
acct.username = body.username
|
||||
if body.password is not None:
|
||||
acct.password = encrypt_value(body.password)
|
||||
if body.use_ssl is not None:
|
||||
acct.use_ssl = body.use_ssl
|
||||
if body.delete_after_process is not None:
|
||||
acct.delete_after_process = body.delete_after_process
|
||||
if body.is_active is not None:
|
||||
acct.is_active = body.is_active
|
||||
# profile_id: update whenever the field is explicitly present in the request payload
|
||||
# (including sending null to clear the override).
|
||||
if "profile_id" in body.model_fields_set:
|
||||
acct.profile_id = body.profile_id
|
||||
|
||||
# Reset last_error so the next poll gives a fresh result
|
||||
acct.last_error = None
|
||||
acct.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(acct)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s updated IMAP account %d", owner_id, account_id)
|
||||
return _to_response(acct)
|
||||
|
||||
|
||||
@router.delete("/{account_id}", status_code=status.HTTP_204_NO_CONTENT, summary="Delete an IMAP account")
|
||||
def delete_imap_account(account_id: int, request: Request, db: DbSession, owner_id: CurrentOwner) -> None:
|
||||
"""Delete an IMAP account permanently."""
|
||||
acct = (
|
||||
db.query(UserImapAccount).filter(UserImapAccount.id == account_id, UserImapAccount.owner_id == owner_id).first()
|
||||
)
|
||||
if not acct:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="IMAP account not found")
|
||||
|
||||
try:
|
||||
db.delete(acct)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s deleted IMAP account %d", owner_id, account_id)
|
||||
|
||||
|
||||
@router.post("/{account_id}/test", summary="Test an existing IMAP account's connection")
|
||||
def test_saved_imap_account(account_id: int, request: Request, db: DbSession, owner_id: CurrentOwner) -> dict[str, Any]:
|
||||
"""Test the connection for an already-saved IMAP account."""
|
||||
acct = (
|
||||
db.query(UserImapAccount).filter(UserImapAccount.id == account_id, UserImapAccount.owner_id == owner_id).first()
|
||||
)
|
||||
if not acct:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="IMAP account not found")
|
||||
|
||||
return _test_imap_connection(acct.host, acct.port, acct.username, decrypt_value(acct.password), acct.use_ssl)
|
||||
|
||||
|
||||
@router.post("/test", summary="Test an IMAP connection without saving")
|
||||
def test_imap_connection(request: Request, body: ImapTestRequest, owner_id: CurrentOwner) -> dict[str, Any]:
|
||||
"""Test IMAP credentials without persisting anything.
|
||||
|
||||
Useful for the "Test connection" button in the UI before the user saves
|
||||
a new account.
|
||||
"""
|
||||
return _test_imap_connection(body.host, body.port, body.username, body.password, body.use_ssl)
|
||||
|
||||
|
||||
@router.get("/quota/", summary="Get IMAP account quota information for the current user")
|
||||
def get_imap_quota(request: Request, db: DbSession, owner_id: CurrentOwner) -> dict[str, Any]:
|
||||
"""Return the user's current IMAP account usage vs. their plan quota."""
|
||||
tier_id = get_user_tier_id(db, owner_id)
|
||||
tier = get_tier(tier_id, db)
|
||||
max_mb = _get_max_mailboxes(tier)
|
||||
current_count = db.query(UserImapAccount).filter(UserImapAccount.owner_id == owner_id).count()
|
||||
|
||||
return {
|
||||
"current_count": current_count,
|
||||
"max_mailboxes": max_mb, # None = unlimited, 0 = not allowed
|
||||
"can_add": max_mb is None or (max_mb > 0 and current_count < max_mb),
|
||||
"tier_id": tier_id,
|
||||
"tier_name": tier.get("name", tier_id),
|
||||
}
|
||||
@@ -0,0 +1,257 @@
|
||||
"""API endpoints for managing IMAP ingestion profiles.
|
||||
|
||||
Ingestion profiles allow fine-grained control over which attachment types are
|
||||
accepted when ingesting emails via IMAP. Each profile carries a list of enabled
|
||||
file-type categories (e.g. ``["pdf", "office", "images"]``) drawn from the
|
||||
canonical set defined in :mod:`app.utils.allowed_types`.
|
||||
|
||||
Built-in system profiles (``is_builtin=True``) are read-only and cannot be
|
||||
deleted or modified. Users may create their own profiles which are private to
|
||||
their ``owner_id``. System-level global profiles (``owner_id=None``) are visible
|
||||
to all users but can only be created by administrators.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import ImapIngestionProfile
|
||||
from app.utils.allowed_types import FILE_TYPE_CATEGORIES
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/imap-profiles", tags=["imap-profiles"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if owner_id is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_VALID_CATEGORIES = set(FILE_TYPE_CATEGORIES.keys())
|
||||
|
||||
|
||||
class ImapProfileCreate(BaseModel):
|
||||
"""Schema for creating a new ingestion profile."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255, description="Human-readable profile name")
|
||||
description: str | None = Field(default=None, description="Optional description")
|
||||
allowed_categories: list[str] = Field(
|
||||
...,
|
||||
min_length=1,
|
||||
description=(f"List of enabled file-type category keys. Valid values: {sorted(_VALID_CATEGORIES)}"),
|
||||
)
|
||||
|
||||
|
||||
class ImapProfileUpdate(BaseModel):
|
||||
"""Schema for updating an existing profile (all fields optional)."""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
description: str | None = None
|
||||
allowed_categories: list[str] | None = Field(default=None, min_length=1)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Validation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _validate_categories(categories: list[str]) -> list[str]:
|
||||
"""Raise 422 if any category key is unknown; return the cleaned list."""
|
||||
unknown = [c for c in categories if c not in _VALID_CATEGORIES]
|
||||
if unknown:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Unknown category key(s): {unknown}. Valid keys: {sorted(_VALID_CATEGORIES)}",
|
||||
)
|
||||
# Deduplicate while preserving order
|
||||
seen: set[str] = set()
|
||||
result: list[str] = []
|
||||
for cat in categories:
|
||||
if cat not in seen:
|
||||
seen.add(cat)
|
||||
result.append(cat)
|
||||
return result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialisation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _to_response(profile: ImapIngestionProfile) -> dict[str, Any]:
|
||||
"""Serialize a profile row to a response dict."""
|
||||
try:
|
||||
categories = json.loads(profile.allowed_categories)
|
||||
except (ValueError, TypeError):
|
||||
categories = []
|
||||
|
||||
# Enrich categories with display metadata
|
||||
categories_detail = [
|
||||
{
|
||||
"key": cat,
|
||||
"label": FILE_TYPE_CATEGORIES[cat]["label"] if cat in FILE_TYPE_CATEGORIES else cat,
|
||||
"description": FILE_TYPE_CATEGORIES[cat]["description"] if cat in FILE_TYPE_CATEGORIES else "",
|
||||
}
|
||||
for cat in categories
|
||||
]
|
||||
|
||||
return {
|
||||
"id": profile.id,
|
||||
"name": profile.name,
|
||||
"description": profile.description,
|
||||
"owner_id": profile.owner_id,
|
||||
"allowed_categories": categories,
|
||||
"categories_detail": categories_detail,
|
||||
"is_builtin": profile.is_builtin,
|
||||
"created_at": profile.created_at.isoformat() if profile.created_at else None,
|
||||
"updated_at": profile.updated_at.isoformat() if profile.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/categories", summary="List available file-type categories")
|
||||
def list_categories(request: Request, owner_id: CurrentOwner) -> list[dict[str, Any]]:
|
||||
"""Return the full list of file-type categories that can be used in profiles."""
|
||||
return [
|
||||
{
|
||||
"key": key,
|
||||
"label": info["label"],
|
||||
"description": info["description"],
|
||||
}
|
||||
for key, info in FILE_TYPE_CATEGORIES.items()
|
||||
]
|
||||
|
||||
|
||||
@router.get("/", summary="List ingestion profiles visible to the current user")
|
||||
def list_profiles(request: Request, db: DbSession, owner_id: CurrentOwner) -> list[dict[str, Any]]:
|
||||
"""Return all profiles: system-global (owner_id=NULL) and the user's own profiles."""
|
||||
profiles = (
|
||||
db.query(ImapIngestionProfile)
|
||||
.filter(
|
||||
# SQLAlchemy requires `== None` for IS NULL comparison in ORM filters
|
||||
(ImapIngestionProfile.owner_id == None) | (ImapIngestionProfile.owner_id == owner_id) # noqa: E711
|
||||
)
|
||||
.order_by(ImapIngestionProfile.is_builtin.desc(), ImapIngestionProfile.id)
|
||||
.all()
|
||||
)
|
||||
return [_to_response(p) for p in profiles]
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED, summary="Create a new ingestion profile")
|
||||
def create_profile(request: Request, body: ImapProfileCreate, db: DbSession, owner_id: CurrentOwner) -> dict[str, Any]:
|
||||
"""Create a new ingestion profile owned by the current user."""
|
||||
categories = _validate_categories(body.allowed_categories)
|
||||
|
||||
profile = ImapIngestionProfile(
|
||||
name=body.name,
|
||||
description=body.description,
|
||||
owner_id=owner_id,
|
||||
allowed_categories=json.dumps(categories),
|
||||
is_builtin=False,
|
||||
)
|
||||
try:
|
||||
db.add(profile)
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s created IMAP ingestion profile %d ('%s')", owner_id, profile.id, body.name)
|
||||
return _to_response(profile)
|
||||
|
||||
|
||||
@router.get("/{profile_id}", summary="Get a single ingestion profile")
|
||||
def get_profile(profile_id: int, request: Request, db: DbSession, owner_id: CurrentOwner) -> dict[str, Any]:
|
||||
"""Return a single profile by ID. Only the owner or system profiles are accessible."""
|
||||
profile = db.query(ImapIngestionProfile).filter(ImapIngestionProfile.id == profile_id).first()
|
||||
if not profile or (profile.owner_id is not None and profile.owner_id != owner_id):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Ingestion profile not found")
|
||||
return _to_response(profile)
|
||||
|
||||
|
||||
@router.put("/{profile_id}", summary="Update an ingestion profile")
|
||||
def update_profile(
|
||||
profile_id: int,
|
||||
request: Request,
|
||||
body: ImapProfileUpdate,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Update an existing ingestion profile. Built-in profiles cannot be modified."""
|
||||
profile = db.query(ImapIngestionProfile).filter(ImapIngestionProfile.id == profile_id).first()
|
||||
if not profile or (profile.owner_id is not None and profile.owner_id != owner_id):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Ingestion profile not found")
|
||||
if profile.is_builtin:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Built-in profiles cannot be modified.",
|
||||
)
|
||||
|
||||
if body.name is not None:
|
||||
profile.name = body.name
|
||||
if "description" in body.model_fields_set:
|
||||
profile.description = body.description
|
||||
if body.allowed_categories is not None:
|
||||
categories = _validate_categories(body.allowed_categories)
|
||||
profile.allowed_categories = json.dumps(categories)
|
||||
|
||||
profile.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s updated IMAP ingestion profile %d", owner_id, profile_id)
|
||||
return _to_response(profile)
|
||||
|
||||
|
||||
@router.delete("/{profile_id}", status_code=status.HTTP_204_NO_CONTENT, summary="Delete an ingestion profile")
|
||||
def delete_profile(profile_id: int, request: Request, db: DbSession, owner_id: CurrentOwner) -> None:
|
||||
"""Delete an ingestion profile. Built-in profiles cannot be deleted."""
|
||||
profile = db.query(ImapIngestionProfile).filter(ImapIngestionProfile.id == profile_id).first()
|
||||
if not profile or (profile.owner_id is not None and profile.owner_id != owner_id):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Ingestion profile not found")
|
||||
if profile.is_builtin:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Built-in profiles cannot be deleted.",
|
||||
)
|
||||
|
||||
try:
|
||||
db.delete(profile)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s deleted IMAP ingestion profile %d", owner_id, profile_id)
|
||||
@@ -0,0 +1,683 @@
|
||||
"""API endpoints for managing per-user integrations (sources and destinations).
|
||||
|
||||
Provides CRUD operations for :class:`~app.models.UserIntegration` records.
|
||||
Each record represents one ingestion source (e.g. IMAP, Watch Folder) or
|
||||
storage destination (e.g. S3, Dropbox, Google Drive) configured by a user.
|
||||
|
||||
Sensitive credentials are encrypted at rest using Fernet symmetric encryption
|
||||
(keyed from ``SESSION_SECRET``) via :mod:`app.utils.encryption`. Credential
|
||||
values are **never** returned in API responses.
|
||||
|
||||
Subscription quota enforcement
|
||||
------------------------------
|
||||
On creation, the endpoint checks the user's subscription tier limits:
|
||||
|
||||
* **Destinations** — ``max_storage_destinations`` from the plan.
|
||||
* **Sources (IMAP)** — ``max_mailboxes`` from the plan.
|
||||
|
||||
Exceeding the quota returns HTTP 403 with an actionable error message.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import IntegrationDirection, IntegrationType, UserIntegration
|
||||
from app.utils.encryption import decrypt_value, encrypt_value
|
||||
from app.utils.subscription import get_tier, get_user_tier_id
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/integrations", tags=["integrations"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if owner_id is None:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Quota helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_FREE_TIER_ID = "free"
|
||||
|
||||
# Source types that consume the mailbox quota
|
||||
_MAILBOX_SOURCE_TYPES = {IntegrationType.IMAP}
|
||||
|
||||
|
||||
def _get_max_destinations(tier: dict[str, Any]) -> int | None:
|
||||
"""Return the maximum number of storage destinations allowed by *tier*.
|
||||
|
||||
Returns:
|
||||
``None`` — unlimited (paid tiers with ``max_storage_destinations == 0``)
|
||||
positive — the configured limit
|
||||
"""
|
||||
tier_id: str = tier.get("id", _FREE_TIER_ID)
|
||||
max_dest: int = tier.get("max_storage_destinations", 0)
|
||||
|
||||
# Free tier: the value itself is the limit (e.g. 1)
|
||||
if tier_id == _FREE_TIER_ID:
|
||||
return max_dest if max_dest > 0 else 1 # safe default
|
||||
|
||||
# Paid tiers: 0 means unlimited
|
||||
if max_dest == 0:
|
||||
return None
|
||||
|
||||
return max_dest
|
||||
|
||||
|
||||
def _get_max_sources(tier: dict[str, Any]) -> int | None:
|
||||
"""Return the maximum number of IMAP source integrations allowed by *tier*.
|
||||
|
||||
Returns:
|
||||
``None`` — unlimited (paid tiers with ``max_mailboxes == 0``)
|
||||
``0`` — no mailboxes allowed (free tier)
|
||||
positive — the configured limit
|
||||
"""
|
||||
tier_id: str = tier.get("id", _FREE_TIER_ID)
|
||||
max_mb: int = tier.get("max_mailboxes", 0)
|
||||
|
||||
# Free tier: 0 means "no access" (not "unlimited")
|
||||
if tier_id == _FREE_TIER_ID:
|
||||
return 0
|
||||
|
||||
# Paid tiers: 0 means unlimited
|
||||
if max_mb == 0:
|
||||
return None
|
||||
|
||||
return max_mb
|
||||
|
||||
|
||||
def _check_quota(db: Session, owner_id: str, direction: str, integration_type: str) -> None:
|
||||
"""Raise 403 if the user has reached their integration quota.
|
||||
|
||||
Quota rules:
|
||||
* DESTINATION integrations are limited by ``max_storage_destinations``.
|
||||
* SOURCE integrations of type IMAP are limited by ``max_mailboxes``.
|
||||
* Other SOURCE types (WATCH_FOLDER, WEBHOOK) are not quota-limited yet.
|
||||
"""
|
||||
tier_id = get_user_tier_id(db, owner_id)
|
||||
tier = get_tier(tier_id, db)
|
||||
|
||||
if direction == IntegrationDirection.DESTINATION:
|
||||
max_dest = _get_max_destinations(tier)
|
||||
if max_dest is not None:
|
||||
current_count = (
|
||||
db.query(UserIntegration)
|
||||
.filter(
|
||||
UserIntegration.owner_id == owner_id,
|
||||
UserIntegration.direction == IntegrationDirection.DESTINATION,
|
||||
)
|
||||
.count()
|
||||
)
|
||||
if current_count >= max_dest:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=(
|
||||
f"You have reached your plan limit of {max_dest} storage destination(s). "
|
||||
"Please remove an existing destination or upgrade your plan."
|
||||
),
|
||||
)
|
||||
|
||||
elif direction == IntegrationDirection.SOURCE and integration_type in _MAILBOX_SOURCE_TYPES:
|
||||
max_src = _get_max_sources(tier)
|
||||
|
||||
if max_src == 0:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Your current plan does not include email ingestion. Upgrade to a paid plan to add IMAP sources.",
|
||||
)
|
||||
|
||||
if max_src is not None:
|
||||
current_count = (
|
||||
db.query(UserIntegration)
|
||||
.filter(
|
||||
UserIntegration.owner_id == owner_id,
|
||||
UserIntegration.direction == IntegrationDirection.SOURCE,
|
||||
UserIntegration.integration_type.in_(list(_MAILBOX_SOURCE_TYPES)),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
if current_count >= max_src:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail=(
|
||||
f"You have reached your plan limit of {max_src} IMAP source(s). "
|
||||
"Please remove an existing source or upgrade your plan."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_VALID_DIRECTIONS = IntegrationDirection.ALL
|
||||
_VALID_TYPES = IntegrationType.ALL
|
||||
|
||||
|
||||
class IntegrationCreate(BaseModel):
|
||||
"""Schema for creating a new integration."""
|
||||
|
||||
direction: str = Field(..., description="'SOURCE' or 'DESTINATION'")
|
||||
integration_type: str = Field(..., description="Integration type (e.g. 'IMAP', 'S3', 'DROPBOX')")
|
||||
name: str = Field(..., min_length=1, max_length=255, description="Human-readable label")
|
||||
config: dict[str, Any] | None = Field(default=None, description="Non-sensitive configuration (JSON object)")
|
||||
credentials: dict[str, Any] | None = Field(
|
||||
default=None, description="Sensitive credentials (JSON object, encrypted at rest)"
|
||||
)
|
||||
is_active: bool = Field(default=True, description="Whether the integration is active")
|
||||
|
||||
|
||||
class IntegrationUpdate(BaseModel):
|
||||
"""Schema for updating an existing integration (all fields optional)."""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
config: dict[str, Any] | None = None
|
||||
credentials: dict[str, Any] | None = None
|
||||
is_active: bool | None = None
|
||||
|
||||
|
||||
class IntegrationTestRequest(BaseModel):
|
||||
"""Schema for testing an integration connection without saving it."""
|
||||
|
||||
integration_type: str = Field(..., description="Integration type (e.g. 'IMAP', 'S3', 'DROPBOX')")
|
||||
config: dict[str, Any] | None = Field(default=None, description="Non-sensitive configuration")
|
||||
credentials: dict[str, Any] | None = Field(default=None, description="Credentials for the connection test")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Validation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _validate_direction(direction: str) -> None:
|
||||
"""Raise 400 if *direction* is not a known value."""
|
||||
if direction not in _VALID_DIRECTIONS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid direction '{direction}'. Must be one of: {sorted(_VALID_DIRECTIONS)}",
|
||||
)
|
||||
|
||||
|
||||
def _validate_integration_type(integration_type: str) -> None:
|
||||
"""Raise 400 if *integration_type* is not a known value."""
|
||||
if integration_type not in _VALID_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid integration_type '{integration_type}'. Must be one of: {sorted(_VALID_TYPES)}",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Serialisation helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _to_response(integration: UserIntegration) -> dict[str, Any]:
|
||||
"""Serialise a :class:`UserIntegration` row to a response dict.
|
||||
|
||||
Credentials are **never** included; only a boolean flag indicating
|
||||
whether credentials have been configured is returned.
|
||||
"""
|
||||
config_data: dict[str, Any] | None = None
|
||||
if integration.config:
|
||||
try:
|
||||
config_data = json.loads(integration.config)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
config_data = None
|
||||
|
||||
return {
|
||||
"id": integration.id,
|
||||
"owner_id": integration.owner_id,
|
||||
"direction": integration.direction,
|
||||
"integration_type": integration.integration_type,
|
||||
"name": integration.name,
|
||||
"config": config_data,
|
||||
"has_credentials": bool(integration.credentials),
|
||||
"is_active": integration.is_active,
|
||||
"last_used_at": integration.last_used_at.isoformat() if integration.last_used_at else None,
|
||||
"last_error": integration.last_error,
|
||||
"created_at": integration.created_at.isoformat() if integration.created_at else None,
|
||||
"updated_at": integration.updated_at.isoformat() if integration.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _encode_credentials(credentials: dict[str, Any] | None) -> str | None:
|
||||
"""Serialise *credentials* dict to an encrypted JSON string for storage."""
|
||||
if not credentials:
|
||||
return None
|
||||
plaintext = json.dumps(credentials)
|
||||
return encrypt_value(plaintext)
|
||||
|
||||
|
||||
def _decode_credentials(stored: str | None) -> dict[str, Any] | None:
|
||||
"""Decrypt and deserialise stored credentials back to a dict.
|
||||
|
||||
Returns ``None`` when *stored* is empty or cannot be decoded.
|
||||
"""
|
||||
if not stored:
|
||||
return None
|
||||
plaintext = decrypt_value(stored)
|
||||
if not plaintext:
|
||||
return None
|
||||
try:
|
||||
return json.loads(plaintext)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.error("Failed to decode credentials JSON after decryption")
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", summary="List integrations for the current user")
|
||||
def list_integrations(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
direction: str | None = None,
|
||||
integration_type: str | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return all integrations belonging to the authenticated user.
|
||||
|
||||
Optional query-string filters:
|
||||
|
||||
- ``direction`` — ``SOURCE`` or ``DESTINATION``
|
||||
- ``integration_type`` — e.g. ``IMAP``, ``S3``, ``DROPBOX``
|
||||
"""
|
||||
query = db.query(UserIntegration).filter(UserIntegration.owner_id == owner_id)
|
||||
|
||||
if direction is not None:
|
||||
_validate_direction(direction)
|
||||
query = query.filter(UserIntegration.direction == direction)
|
||||
|
||||
if integration_type is not None:
|
||||
_validate_integration_type(integration_type)
|
||||
query = query.filter(UserIntegration.integration_type == integration_type)
|
||||
|
||||
integrations = query.order_by(UserIntegration.id).all()
|
||||
return [_to_response(i) for i in integrations]
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED, summary="Create a new integration")
|
||||
def create_integration(
|
||||
request: Request,
|
||||
body: IntegrationCreate,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new source or destination integration for the current user.
|
||||
|
||||
``credentials`` are encrypted at rest using Fernet symmetric encryption
|
||||
before being persisted and are **never** returned in API responses.
|
||||
|
||||
Quota is enforced against the user's subscription plan before the
|
||||
integration is persisted.
|
||||
"""
|
||||
_validate_direction(body.direction)
|
||||
_validate_integration_type(body.integration_type)
|
||||
|
||||
_check_quota(db, owner_id, body.direction, body.integration_type)
|
||||
|
||||
integration = UserIntegration(
|
||||
owner_id=owner_id,
|
||||
direction=body.direction,
|
||||
integration_type=body.integration_type,
|
||||
name=body.name,
|
||||
config=json.dumps(body.config) if body.config is not None else None,
|
||||
credentials=_encode_credentials(body.credentials),
|
||||
is_active=body.is_active,
|
||||
)
|
||||
|
||||
try:
|
||||
db.add(integration)
|
||||
db.commit()
|
||||
db.refresh(integration)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info(
|
||||
"User %s created %s integration %d (%s)",
|
||||
owner_id,
|
||||
body.direction,
|
||||
integration.id,
|
||||
body.integration_type,
|
||||
)
|
||||
return _to_response(integration)
|
||||
|
||||
|
||||
@router.get("/{integration_id}", summary="Get a single integration")
|
||||
def get_integration(
|
||||
integration_id: int,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Return a single integration by ID (must belong to the current user)."""
|
||||
integration = (
|
||||
db.query(UserIntegration)
|
||||
.filter(UserIntegration.id == integration_id, UserIntegration.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not integration:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Integration not found")
|
||||
return _to_response(integration)
|
||||
|
||||
|
||||
@router.put("/{integration_id}", summary="Update an integration")
|
||||
def update_integration(
|
||||
integration_id: int,
|
||||
request: Request,
|
||||
body: IntegrationUpdate,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Update an existing integration. Only provided fields are changed.
|
||||
|
||||
When ``credentials`` is supplied the stored value is replaced in full
|
||||
with the freshly encrypted version of the new credentials dict.
|
||||
"""
|
||||
integration = (
|
||||
db.query(UserIntegration)
|
||||
.filter(UserIntegration.id == integration_id, UserIntegration.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not integration:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Integration not found")
|
||||
|
||||
if body.name is not None:
|
||||
integration.name = body.name
|
||||
if body.config is not None:
|
||||
integration.config = json.dumps(body.config)
|
||||
if body.credentials is not None:
|
||||
integration.credentials = _encode_credentials(body.credentials)
|
||||
if body.is_active is not None:
|
||||
integration.is_active = body.is_active
|
||||
|
||||
# Reset last_error so the next operation gives a fresh result
|
||||
integration.last_error = None
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(integration)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s updated integration %d", owner_id, integration_id)
|
||||
return _to_response(integration)
|
||||
|
||||
|
||||
@router.delete("/{integration_id}", status_code=status.HTTP_204_NO_CONTENT, summary="Delete an integration")
|
||||
def delete_integration(
|
||||
integration_id: int,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> None:
|
||||
"""Delete an integration permanently."""
|
||||
integration = (
|
||||
db.query(UserIntegration)
|
||||
.filter(UserIntegration.id == integration_id, UserIntegration.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not integration:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Integration not found")
|
||||
|
||||
try:
|
||||
db.delete(integration)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("User %s deleted integration %d", owner_id, integration_id)
|
||||
|
||||
|
||||
@router.get("/{integration_id}/credentials", summary="Retrieve decrypted credentials for an integration")
|
||||
def get_integration_credentials(
|
||||
integration_id: int,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Return the decrypted credentials dict for a saved integration.
|
||||
|
||||
This endpoint is intended for internal use by background tasks that need
|
||||
to authenticate with a third-party service. Treat the response as
|
||||
sensitive — it contains plaintext secrets.
|
||||
"""
|
||||
integration = (
|
||||
db.query(UserIntegration)
|
||||
.filter(UserIntegration.id == integration_id, UserIntegration.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not integration:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Integration not found")
|
||||
|
||||
credentials = _decode_credentials(integration.credentials)
|
||||
return {"credentials": credentials or {}}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Connection test helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _test_imap_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Test an IMAP connection using the provided config and credentials."""
|
||||
import imaplib
|
||||
|
||||
cfg = config or {}
|
||||
creds = credentials or {}
|
||||
host = cfg.get("host", "")
|
||||
port = int(cfg.get("port", 993))
|
||||
username = cfg.get("username", "")
|
||||
password = creds.get("password", "")
|
||||
use_ssl = cfg.get("use_ssl", True)
|
||||
|
||||
if not host or not username or not password:
|
||||
return {"success": False, "message": "Missing required fields: host, username, and password"}
|
||||
|
||||
try:
|
||||
if use_ssl:
|
||||
mail = imaplib.IMAP4_SSL(host, port)
|
||||
else:
|
||||
mail = imaplib.IMAP4(host, port)
|
||||
mail.login(username, password)
|
||||
mail.logout()
|
||||
return {"success": True, "message": "IMAP connection successful"}
|
||||
except OSError as exc:
|
||||
logger.warning("IMAP network error for %s@%s: %s", username, host, exc)
|
||||
return {"success": False, "message": "IMAP connection failed — check host, port, and network connectivity"}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("IMAP error for %s@%s: %s", username, host, exc)
|
||||
return {"success": False, "message": "IMAP authentication or connection failed"}
|
||||
|
||||
|
||||
def _test_s3_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Test an S3 connection by calling HeadBucket."""
|
||||
try:
|
||||
import boto3
|
||||
from botocore.exceptions import BotoCoreError, ClientError
|
||||
except ImportError:
|
||||
return {"success": False, "message": "boto3 is not installed"}
|
||||
|
||||
cfg = config or {}
|
||||
creds = credentials or {}
|
||||
bucket = cfg.get("bucket", "")
|
||||
region = cfg.get("region", "us-east-1")
|
||||
|
||||
if not bucket:
|
||||
return {"success": False, "message": "Missing required field: bucket"}
|
||||
|
||||
try:
|
||||
client = boto3.client(
|
||||
"s3",
|
||||
region_name=region,
|
||||
aws_access_key_id=creds.get("access_key_id", ""),
|
||||
aws_secret_access_key=creds.get("secret_access_key", ""),
|
||||
endpoint_url=cfg.get("endpoint_url"),
|
||||
)
|
||||
client.head_bucket(Bucket=bucket)
|
||||
return {"success": True, "message": f"S3 bucket '{bucket}' is accessible"}
|
||||
except (BotoCoreError, ClientError) as exc:
|
||||
logger.warning("S3 connection error for bucket '%s': %s", bucket, exc)
|
||||
return {"success": False, "message": "S3 connection failed — check bucket name, region, and credentials"}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("S3 unexpected error for bucket '%s': %s", bucket, exc)
|
||||
return {"success": False, "message": "S3 connection failed"}
|
||||
|
||||
|
||||
def _test_webdav_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Test a WebDAV/Nextcloud connection by issuing an HTTP PROPFIND."""
|
||||
import urllib.request
|
||||
|
||||
cfg = config or {}
|
||||
creds = credentials or {}
|
||||
url = cfg.get("url", "")
|
||||
username = creds.get("username", "")
|
||||
password = creds.get("password", "")
|
||||
|
||||
if not url:
|
||||
return {"success": False, "message": "Missing required field: url"}
|
||||
|
||||
# Only allow http/https to prevent file:// or other custom scheme attacks
|
||||
from urllib.parse import urlparse
|
||||
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in ("http", "https"):
|
||||
return {"success": False, "message": "URL must use http or https scheme"}
|
||||
|
||||
# Block requests to private/internal IPs to prevent SSRF
|
||||
hostname = parsed.hostname or ""
|
||||
if hostname:
|
||||
from app.utils.network import is_private_ip
|
||||
|
||||
if is_private_ip(hostname):
|
||||
return {"success": False, "message": "URLs pointing to internal or private networks are not allowed"}
|
||||
|
||||
try:
|
||||
import base64
|
||||
|
||||
req = urllib.request.Request(url, method="PROPFIND") # noqa: S310
|
||||
if username and password:
|
||||
token = base64.b64encode(f"{username}:{password}".encode()).decode()
|
||||
req.add_header("Authorization", f"Basic {token}")
|
||||
req.add_header("Depth", "0")
|
||||
with urllib.request.urlopen(req, timeout=10) as resp: # noqa: S310
|
||||
if resp.status < 400:
|
||||
return {"success": True, "message": "WebDAV connection successful"}
|
||||
return {"success": False, "message": f"WebDAV returned HTTP {resp.status}"}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("WebDAV connection error for %s: %s", hostname, exc)
|
||||
return {"success": False, "message": "WebDAV connection failed — check URL and credentials"}
|
||||
|
||||
|
||||
_CONNECTION_TESTERS: dict[str, Any] = {
|
||||
IntegrationType.IMAP: _test_imap_connection,
|
||||
IntegrationType.S3: _test_s3_connection,
|
||||
IntegrationType.WEBDAV: _test_webdav_connection,
|
||||
IntegrationType.NEXTCLOUD: _test_webdav_connection,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test & quota endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/test", summary="Test an integration connection without saving")
|
||||
def test_integration_connection(
|
||||
request: Request,
|
||||
body: IntegrationTestRequest,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Test integration credentials without persisting anything.
|
||||
|
||||
Useful for the "Test connection" button in the UI before the user saves
|
||||
a new integration. Returns ``{"success": bool, "message": str}``.
|
||||
"""
|
||||
_validate_integration_type(body.integration_type)
|
||||
|
||||
tester = _CONNECTION_TESTERS.get(body.integration_type)
|
||||
if tester is None:
|
||||
return {
|
||||
"success": False,
|
||||
"message": f"Connection testing is not yet supported for '{body.integration_type}'. "
|
||||
"The integration can still be saved and will be validated on first use.",
|
||||
}
|
||||
|
||||
return tester(body.config, body.credentials)
|
||||
|
||||
|
||||
@router.get("/quota/", summary="Get integration quota information for the current user")
|
||||
def get_integration_quota(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
owner_id: CurrentOwner,
|
||||
) -> dict[str, Any]:
|
||||
"""Return the user's current integration usage vs. their plan quota.
|
||||
|
||||
Includes separate counts for destinations and IMAP sources.
|
||||
"""
|
||||
tier_id = get_user_tier_id(db, owner_id)
|
||||
tier = get_tier(tier_id, db)
|
||||
|
||||
max_dest = _get_max_destinations(tier)
|
||||
max_src = _get_max_sources(tier)
|
||||
|
||||
dest_count = (
|
||||
db.query(UserIntegration)
|
||||
.filter(
|
||||
UserIntegration.owner_id == owner_id,
|
||||
UserIntegration.direction == IntegrationDirection.DESTINATION,
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
src_count = (
|
||||
db.query(UserIntegration)
|
||||
.filter(
|
||||
UserIntegration.owner_id == owner_id,
|
||||
UserIntegration.direction == IntegrationDirection.SOURCE,
|
||||
UserIntegration.integration_type.in_(list(_MAILBOX_SOURCE_TYPES)),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
return {
|
||||
"tier_id": tier_id,
|
||||
"tier_name": tier.get("name", tier_id),
|
||||
"destinations": {
|
||||
"current_count": dest_count,
|
||||
"max_allowed": max_dest,
|
||||
"can_add": max_dest is None or dest_count < max_dest,
|
||||
},
|
||||
"sources": {
|
||||
"current_count": src_count,
|
||||
"max_allowed": max_src,
|
||||
"can_add": max_src is None or (max_src > 0 and src_count < max_src),
|
||||
},
|
||||
}
|
||||
@@ -0,0 +1,393 @@
|
||||
"""Local user authentication API — signup, email verification, password reset.
|
||||
|
||||
Provides the REST endpoints and page routes for the self-registration flow:
|
||||
|
||||
- GET /signup — signup page (HTML)
|
||||
- POST /api/auth/signup — create account + send verification email
|
||||
- GET /verify-email — activate account from email link (redirect)
|
||||
- GET /verify-email-sent — confirmation landing page (HTML)
|
||||
- POST /api/auth/resend-verification — re-send verification email
|
||||
- POST /api/auth/request-password-reset — start password reset
|
||||
- POST /api/auth/reset-password — set new password using token
|
||||
- GET /reset-password — password reset form page (HTML)
|
||||
"""
|
||||
|
||||
import logging
|
||||
import pathlib
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.templating import Jinja2Templates
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.responses import RedirectResponse
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import LocalUser, UserProfile
|
||||
from app.utils.i18n import translate as _translate
|
||||
from app.utils.local_auth import (
|
||||
build_session_user,
|
||||
generate_token,
|
||||
hash_password,
|
||||
is_token_expired,
|
||||
send_forgot_username_email,
|
||||
send_password_reset_email,
|
||||
send_verification_email,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["local-auth"])
|
||||
|
||||
_templates_dir = pathlib.Path(__file__).parents[2] / "frontend" / "templates"
|
||||
templates = Jinja2Templates(directory=str(_templates_dir))
|
||||
templates.env.globals["_"] = lambda key, **kwargs: _translate(key, "en", **kwargs)
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SignupBody(BaseModel):
|
||||
"""Body for the signup endpoint."""
|
||||
|
||||
email: str = Field(..., max_length=255)
|
||||
username: str = Field(..., min_length=3, max_length=64, pattern=r"^[a-zA-Z0-9_-]+$")
|
||||
display_name: str | None = Field(default=None, max_length=255)
|
||||
password: str = Field(..., min_length=8, max_length=128)
|
||||
password_confirm: str
|
||||
|
||||
|
||||
class ResendVerificationBody(BaseModel):
|
||||
"""Body for the resend-verification endpoint."""
|
||||
|
||||
email: str
|
||||
|
||||
|
||||
class PasswordResetRequestBody(BaseModel):
|
||||
"""Body for the request-password-reset endpoint."""
|
||||
|
||||
email: str
|
||||
|
||||
|
||||
class PasswordResetBody(BaseModel):
|
||||
"""Body for the reset-password endpoint."""
|
||||
|
||||
token: str
|
||||
new_password: str = Field(..., min_length=8, max_length=128)
|
||||
new_password_confirm: str
|
||||
|
||||
|
||||
class ForgotUsernameBody(BaseModel):
|
||||
"""Body for the forgot-username endpoint."""
|
||||
|
||||
email: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Page routes (return HTML)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/signup", include_in_schema=False)
|
||||
async def signup_page(request: Request) -> Any:
|
||||
"""Render the signup page, or redirect to login when multi-user / signup is disabled."""
|
||||
if not settings.multi_user_enabled:
|
||||
return RedirectResponse(url="/login?error=Multi-user+mode+is+not+enabled", status_code=302)
|
||||
if not settings.allow_local_signup:
|
||||
return RedirectResponse(url="/login?error=Registration+is+not+enabled", status_code=302)
|
||||
return templates.TemplateResponse(
|
||||
"signup.html",
|
||||
{
|
||||
"request": request,
|
||||
"csrf_token": getattr(request.state, "csrf_token", ""),
|
||||
"app_version": settings.version,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/verify-email-sent", include_in_schema=False)
|
||||
async def verify_email_sent_page(request: Request) -> Any:
|
||||
"""Render the verify-email-sent confirmation page."""
|
||||
return templates.TemplateResponse("verify_email_sent.html", {"request": request})
|
||||
|
||||
|
||||
@router.get("/forgot-username", include_in_schema=False)
|
||||
async def forgot_username_page(request: Request) -> Any:
|
||||
"""Render the forgot-username page where users can request a username reminder email."""
|
||||
return templates.TemplateResponse(
|
||||
"forgot_username.html",
|
||||
{
|
||||
"request": request,
|
||||
"csrf_token": getattr(request.state, "csrf_token", ""),
|
||||
"app_version": settings.version,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/forgot-password", include_in_schema=False)
|
||||
async def forgot_password_page(request: Request) -> Any:
|
||||
"""Render the forgot-password page where users can request a reset email."""
|
||||
return templates.TemplateResponse(
|
||||
"forgot_password.html",
|
||||
{
|
||||
"request": request,
|
||||
"csrf_token": getattr(request.state, "csrf_token", ""),
|
||||
"app_version": settings.version,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/reset-password", include_in_schema=False)
|
||||
async def reset_password_page(request: Request) -> Any:
|
||||
"""Render the password reset form page."""
|
||||
token = request.query_params.get("token", "")
|
||||
return templates.TemplateResponse(
|
||||
"password_reset_form.html",
|
||||
{
|
||||
"request": request,
|
||||
"token": token,
|
||||
"csrf_token": getattr(request.state, "csrf_token", ""),
|
||||
"app_version": settings.version,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# API endpoints (return JSON or redirect)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/api/auth/signup", status_code=status.HTTP_201_CREATED)
|
||||
async def signup(request: Request, body: SignupBody, db: DbSession) -> dict[str, str | bool]:
|
||||
"""Create a new local user account.
|
||||
|
||||
When SMTP is configured the account is inactive until the user clicks the
|
||||
verification link sent to their email. When SMTP is **not** configured the
|
||||
account is activated immediately so that deployments without email can still
|
||||
use the self-registration flow.
|
||||
|
||||
Both ``MULTI_USER_ENABLED`` and ``ALLOW_LOCAL_SIGNUP`` must be ``True``.
|
||||
|
||||
Raises:
|
||||
403: Multi-user mode or local signup is disabled.
|
||||
422: Passwords do not match.
|
||||
409: Email or username already registered.
|
||||
"""
|
||||
if not settings.multi_user_enabled:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Multi-user mode is not enabled.")
|
||||
if not settings.allow_local_signup:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Registration is not enabled.")
|
||||
if body.password != body.password_confirm:
|
||||
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail="Passwords do not match.")
|
||||
|
||||
if db.query(LocalUser).filter(LocalUser.email == body.email).first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Email already registered.")
|
||||
if db.query(LocalUser).filter(LocalUser.username == body.username).first():
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Username already taken.")
|
||||
|
||||
smtp_configured = bool(settings.email_host)
|
||||
|
||||
if smtp_configured:
|
||||
token = generate_token()
|
||||
user = LocalUser(
|
||||
email=body.email,
|
||||
username=body.username,
|
||||
display_name=body.display_name,
|
||||
hashed_password=hash_password(body.password),
|
||||
is_active=False,
|
||||
email_verification_token=token,
|
||||
email_verification_sent_at=datetime.now(tz=timezone.utc),
|
||||
)
|
||||
else:
|
||||
# No SMTP configured — activate the account immediately.
|
||||
token = None
|
||||
user = LocalUser(
|
||||
email=body.email,
|
||||
username=body.username,
|
||||
display_name=body.display_name,
|
||||
hashed_password=hash_password(body.password),
|
||||
is_active=True,
|
||||
)
|
||||
|
||||
db.add(user)
|
||||
|
||||
profile = UserProfile(
|
||||
user_id=body.email,
|
||||
display_name=body.display_name or body.username,
|
||||
)
|
||||
db.add(profile)
|
||||
|
||||
# Flush to the DB so constraint violations (duplicate key etc.) surface NOW,
|
||||
# before we attempt to send the email. We do NOT commit yet — the commit only
|
||||
# happens after the email is sent successfully so that a failed email leaves
|
||||
# no orphan records in the database.
|
||||
try:
|
||||
db.flush()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
if smtp_configured and token:
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
try:
|
||||
send_verification_email(body.email, body.username, token, base_url)
|
||||
except Exception as exc:
|
||||
# Email failed — roll back so no unverifiable user row persists.
|
||||
# The user can simply try registering again once SMTP is fixed.
|
||||
db.rollback()
|
||||
logger.warning("Signup email failed for %s: %s", body.email, exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=(
|
||||
"Failed to send verification email. Please check that SMTP is correctly configured and try again."
|
||||
),
|
||||
) from exc
|
||||
|
||||
db.commit()
|
||||
logger.info("New local user registered: %s", body.email)
|
||||
|
||||
if smtp_configured:
|
||||
return {"message": "Verification email sent. Please check your inbox.", "email_verification_required": True}
|
||||
return {"message": "Account created successfully. You can now log in.", "email_verification_required": False}
|
||||
|
||||
|
||||
@router.get("/verify-email", include_in_schema=False)
|
||||
async def verify_email(request: Request, db: DbSession) -> Any:
|
||||
"""Activate a local user account from the email verification link.
|
||||
|
||||
Redirects to the login page on failure, or to onboarding/upload on success.
|
||||
"""
|
||||
token = request.query_params.get("token", "")
|
||||
user = db.query(LocalUser).filter(LocalUser.email_verification_token == token).first()
|
||||
|
||||
if not user:
|
||||
return RedirectResponse(
|
||||
url="/login?error=Invalid+or+expired+verification+link",
|
||||
status_code=302,
|
||||
)
|
||||
if is_token_expired(user.email_verification_sent_at):
|
||||
return RedirectResponse(
|
||||
url="/login?error=Verification+link+has+expired.+Please+request+a+new+one",
|
||||
status_code=302,
|
||||
)
|
||||
|
||||
user.is_active = True
|
||||
user.email_verification_token = None
|
||||
user.email_verification_sent_at = None
|
||||
|
||||
# Ensure profile exists
|
||||
if not db.query(UserProfile).filter(UserProfile.user_id == user.email).first():
|
||||
db.add(UserProfile(user_id=user.email, display_name=user.display_name or user.username))
|
||||
|
||||
db.commit()
|
||||
|
||||
request.session["user"] = build_session_user(user)
|
||||
logger.info("[SECURITY] EMAIL_VERIFIED user=%s", user.email)
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user.email).first()
|
||||
if profile and not profile.onboarding_completed:
|
||||
post_onboarding = request.session.pop("redirect_after_login", "/upload")
|
||||
request.session["post_onboarding_redirect"] = post_onboarding
|
||||
return RedirectResponse(url="/onboarding", status_code=302)
|
||||
return RedirectResponse(url="/upload", status_code=302)
|
||||
|
||||
|
||||
@router.post("/api/auth/resend-verification")
|
||||
async def resend_verification(request: Request, body: ResendVerificationBody, db: DbSession) -> dict[str, str]:
|
||||
"""Re-send the verification email for a pending account.
|
||||
|
||||
Always returns 200 to avoid leaking whether an email is registered.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.email == body.email).first()
|
||||
if not user or user.is_active:
|
||||
return {"message": "Verification email resent if account exists."}
|
||||
|
||||
token = generate_token()
|
||||
user.email_verification_token = token
|
||||
user.email_verification_sent_at = datetime.now(tz=timezone.utc)
|
||||
db.commit()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
try:
|
||||
send_verification_email(user.email, user.username, token, base_url)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to resend verification email to %s: %s", user.email, exc)
|
||||
|
||||
return {"message": "Verification email resent if account exists."}
|
||||
|
||||
|
||||
@router.post("/api/auth/request-password-reset")
|
||||
async def request_password_reset(request: Request, body: PasswordResetRequestBody, db: DbSession) -> dict[str, str]:
|
||||
"""Send a password reset email.
|
||||
|
||||
Always returns 200 to avoid leaking whether an email is registered.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.email == body.email).first()
|
||||
if not user:
|
||||
return {"message": "Password reset email sent if account exists."}
|
||||
|
||||
token = generate_token()
|
||||
user.password_reset_token = token
|
||||
user.password_reset_sent_at = datetime.now(tz=timezone.utc)
|
||||
db.commit()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
try:
|
||||
send_password_reset_email(user.email, user.username, token, base_url)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to send password reset email to %s: %s", user.email, exc)
|
||||
|
||||
return {"message": "Password reset email sent if account exists."}
|
||||
|
||||
|
||||
@router.post("/api/auth/reset-password")
|
||||
async def reset_password(body: PasswordResetBody, db: DbSession) -> dict[str, str]:
|
||||
"""Set a new password using a valid reset token.
|
||||
|
||||
Raises:
|
||||
400: Token is invalid or expired.
|
||||
422: Passwords do not match.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.password_reset_token == body.token).first()
|
||||
if not user or is_token_expired(user.password_reset_sent_at):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid or expired reset token.",
|
||||
)
|
||||
if body.new_password != body.new_password_confirm:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="Passwords do not match.",
|
||||
)
|
||||
|
||||
user.hashed_password = hash_password(body.new_password)
|
||||
user.password_reset_token = None
|
||||
user.password_reset_sent_at = None
|
||||
# Activate the account in case it was still pending email verification.
|
||||
# A valid password-reset token proves control of the registered email address.
|
||||
user.is_active = True
|
||||
db.commit()
|
||||
|
||||
logger.info("[SECURITY] PASSWORD_RESET_SUCCESS user=%s", user.email)
|
||||
return {"message": "Password updated successfully."}
|
||||
|
||||
|
||||
@router.post("/api/auth/forgot-username")
|
||||
async def forgot_username(body: ForgotUsernameBody, db: DbSession) -> dict[str, str]:
|
||||
"""Send a username reminder email.
|
||||
|
||||
Always returns 200 to avoid leaking whether an email is registered.
|
||||
"""
|
||||
user = db.query(LocalUser).filter(LocalUser.email == body.email).first()
|
||||
if user:
|
||||
try:
|
||||
send_forgot_username_email(user.email, user.username)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to send forgot-username email to %s: %s", user.email, exc)
|
||||
|
||||
return {"message": "Username reminder sent if account exists."}
|
||||
@@ -0,0 +1,347 @@
|
||||
"""Mobile app API endpoints.
|
||||
|
||||
Provides endpoints specifically designed for the DocuElevate native mobile
|
||||
app (iOS / Android via React Native / Expo):
|
||||
|
||||
* ``POST /mobile/generate-token`` – exchange an active session for a
|
||||
long-lived API token that the mobile app stores securely. The token is
|
||||
auto-named "Mobile App – <device_name>" and is identical to regular API
|
||||
tokens (Bearer auth works everywhere).
|
||||
|
||||
* ``POST /mobile/register-device`` – register a push-notification device
|
||||
token (Expo push token) so the user receives push notifications when
|
||||
documents finish processing.
|
||||
|
||||
* ``GET /mobile/devices`` – list registered devices for the current user.
|
||||
|
||||
* ``DELETE /mobile/devices/{device_id}`` – deactivate a device.
|
||||
|
||||
* ``GET /mobile/whoami`` – lightweight profile endpoint for the mobile app
|
||||
to verify authentication state.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.api_tokens import generate_api_token, hash_token
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import ApiToken, MobileDevice
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/mobile", tags=["mobile"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request / Response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class GenerateTokenRequest(BaseModel):
|
||||
"""Request body for auto-generating a mobile app token."""
|
||||
|
||||
device_name: str = Field(
|
||||
default="Mobile App",
|
||||
min_length=1,
|
||||
max_length=120,
|
||||
description="Human-readable device name used to label the token.",
|
||||
)
|
||||
|
||||
|
||||
class GenerateTokenResponse(BaseModel):
|
||||
"""Response containing the one-time-visible API token."""
|
||||
|
||||
token: str
|
||||
token_id: int
|
||||
name: str
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class RegisterDeviceRequest(BaseModel):
|
||||
"""Request body for registering a push-notification device token."""
|
||||
|
||||
push_token: str = Field(
|
||||
min_length=1,
|
||||
max_length=512,
|
||||
description="Expo push token (ExponentPushToken[…]) obtained from the mobile app.",
|
||||
)
|
||||
device_name: str | None = Field(
|
||||
default=None,
|
||||
max_length=255,
|
||||
description="Optional human-readable device name (e.g. 'John's iPhone').",
|
||||
)
|
||||
platform: str = Field(
|
||||
default="ios",
|
||||
description="Device platform: 'ios', 'android', or 'web'.",
|
||||
)
|
||||
|
||||
|
||||
class DeviceResponse(BaseModel):
|
||||
"""Serialised MobileDevice record."""
|
||||
|
||||
id: int
|
||||
device_name: str | None
|
||||
platform: str
|
||||
push_token_preview: str
|
||||
is_active: bool
|
||||
created_at: datetime
|
||||
last_seen_at: datetime | None
|
||||
|
||||
|
||||
class WhoAmIResponse(BaseModel):
|
||||
"""Lightweight profile response for the mobile app."""
|
||||
|
||||
owner_id: str
|
||||
display_name: str | None
|
||||
email: str | None
|
||||
avatar_url: str | None
|
||||
is_admin: bool
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _device_to_response(device: MobileDevice) -> dict[str, Any]:
|
||||
"""Convert a MobileDevice ORM object to a serialisable dict."""
|
||||
# Show only first 20 chars of the push token for security.
|
||||
token_preview = device.push_token[:20] + "…" if len(device.push_token) > 20 else device.push_token
|
||||
return {
|
||||
"id": device.id,
|
||||
"device_name": device.device_name,
|
||||
"platform": device.platform,
|
||||
"push_token_preview": token_preview,
|
||||
"is_active": device.is_active,
|
||||
"created_at": device.created_at,
|
||||
"last_seen_at": device.last_seen_at,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/generate-token", status_code=status.HTTP_201_CREATED, response_model=GenerateTokenResponse)
|
||||
@require_login
|
||||
async def generate_mobile_token(
|
||||
request: Request,
|
||||
body: GenerateTokenRequest,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Generate a long-lived API token for the mobile app.
|
||||
|
||||
The mobile app calls this endpoint immediately after SSO login to obtain
|
||||
a Bearer token it can store in the secure keychain. The returned token
|
||||
is functionally identical to manually-created API tokens and works with
|
||||
every authenticated endpoint.
|
||||
|
||||
The token is shown **exactly once** in the response; subsequent requests
|
||||
show only the prefix for identification.
|
||||
"""
|
||||
token_name = f"Mobile App – {body.device_name}"
|
||||
plaintext = generate_api_token()
|
||||
token_hash_value = hash_token(plaintext)
|
||||
prefix = plaintext[:12]
|
||||
|
||||
db_token = ApiToken(
|
||||
owner_id=owner_id,
|
||||
name=token_name,
|
||||
token_hash=token_hash_value,
|
||||
token_prefix=prefix,
|
||||
)
|
||||
try:
|
||||
db.add(db_token)
|
||||
db.commit()
|
||||
db.refresh(db_token)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create mobile API token for owner_id=%s", owner_id)
|
||||
raise
|
||||
|
||||
logger.info("Mobile API token created: id=%s owner=%s device=%r", db_token.id, owner_id, body.device_name)
|
||||
|
||||
return {
|
||||
"token": plaintext,
|
||||
"token_id": db_token.id,
|
||||
"name": token_name,
|
||||
"created_at": db_token.created_at,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/register-device", status_code=status.HTTP_201_CREATED, response_model=DeviceResponse)
|
||||
@require_login
|
||||
async def register_device(
|
||||
request: Request,
|
||||
body: RegisterDeviceRequest,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Register or refresh a push-notification device token.
|
||||
|
||||
If the same ``push_token`` is already registered for this user the
|
||||
record is reactivated and ``last_seen_at`` is updated rather than
|
||||
creating a duplicate.
|
||||
"""
|
||||
platform = body.platform.lower()
|
||||
if platform not in {"ios", "android", "web"}:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="platform must be one of: ios, android, web",
|
||||
)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
# Upsert: reuse existing record if the token is already known.
|
||||
existing = (
|
||||
db.query(MobileDevice)
|
||||
.filter(MobileDevice.owner_id == owner_id, MobileDevice.push_token == body.push_token)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
existing.is_active = True
|
||||
existing.last_seen_at = now
|
||||
if body.device_name:
|
||||
existing.device_name = body.device_name
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(existing)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("Mobile device refreshed: id=%s owner=%s", existing.id, owner_id)
|
||||
return _device_to_response(existing)
|
||||
|
||||
device = MobileDevice(
|
||||
owner_id=owner_id,
|
||||
device_name=body.device_name,
|
||||
platform=platform,
|
||||
push_token=body.push_token,
|
||||
is_active=True,
|
||||
last_seen_at=now,
|
||||
)
|
||||
try:
|
||||
db.add(device)
|
||||
db.commit()
|
||||
db.refresh(device)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to register mobile device for owner_id=%s", owner_id)
|
||||
raise
|
||||
|
||||
logger.info("Mobile device registered: id=%s owner=%s platform=%s", device.id, owner_id, platform)
|
||||
return _device_to_response(device)
|
||||
|
||||
|
||||
@router.get("/devices", response_model=list[DeviceResponse])
|
||||
@require_login
|
||||
async def list_devices(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all registered push-notification devices for the current user."""
|
||||
devices = (
|
||||
db.query(MobileDevice).filter(MobileDevice.owner_id == owner_id).order_by(MobileDevice.created_at.desc()).all()
|
||||
)
|
||||
return [_device_to_response(d) for d in devices]
|
||||
|
||||
|
||||
@router.delete("/devices/{device_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
async def deactivate_device(
|
||||
request: Request,
|
||||
device_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> None:
|
||||
"""Deactivate a push-notification device registration.
|
||||
|
||||
The device record is kept for audit purposes but will no longer receive
|
||||
push notifications.
|
||||
"""
|
||||
device = db.get(MobileDevice, device_id)
|
||||
if not device or device.owner_id != owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Device not found")
|
||||
|
||||
device.is_active = False
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Mobile device deactivated: id=%s owner=%s", device_id, owner_id)
|
||||
|
||||
|
||||
@router.get("/whoami", response_model=WhoAmIResponse)
|
||||
@require_login
|
||||
async def whoami(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Return basic profile information for the authenticated user.
|
||||
|
||||
The mobile app calls this after token exchange to populate the user
|
||||
profile screen and verify that the stored token is still valid.
|
||||
"""
|
||||
from app.auth import get_gravatar_url
|
||||
from app.models import LocalUser, UserProfile
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == owner_id).first()
|
||||
local_user = db.query(LocalUser).filter(LocalUser.email == owner_id).first()
|
||||
|
||||
display_name: str | None = None
|
||||
email: str | None = None
|
||||
avatar_url: str | None = None
|
||||
is_admin = False
|
||||
|
||||
if profile:
|
||||
display_name = profile.display_name
|
||||
|
||||
if local_user:
|
||||
email = local_user.email
|
||||
is_admin = bool(local_user.is_admin)
|
||||
if not display_name and local_user.display_name:
|
||||
display_name = local_user.display_name
|
||||
elif "@" in owner_id:
|
||||
# SSO users commonly have their email as owner_id
|
||||
email = owner_id
|
||||
|
||||
if email:
|
||||
avatar_url = get_gravatar_url(email)
|
||||
|
||||
return {
|
||||
"owner_id": owner_id,
|
||||
"display_name": display_name,
|
||||
"email": email,
|
||||
"avatar_url": avatar_url,
|
||||
"is_admin": is_admin,
|
||||
}
|
||||
@@ -0,0 +1,484 @@
|
||||
"""API endpoints for per-user notification targets, preferences, and in-app inbox.
|
||||
|
||||
Users can define notification targets (email via SMTP, webhook via HTTP POST)
|
||||
and configure which document events trigger which targets. In-app notifications
|
||||
are always created and surfaced via the bell icon / inbox endpoints.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import InAppNotification, UserNotificationPreference, UserNotificationTarget
|
||||
from app.utils.user_notification import USER_EVENT_LABELS
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/user-notifications", tags=["user-notifications"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper (mirrors api_tokens.py pattern)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
VALID_CHANNEL_TYPES = {"email", "webhook"}
|
||||
VALID_EVENT_TYPES = set(USER_EVENT_LABELS.keys())
|
||||
|
||||
|
||||
class NotificationTargetCreate(BaseModel):
|
||||
"""Schema for creating a new notification target."""
|
||||
|
||||
channel_type: str = Field(..., pattern="^(email|webhook)$")
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
config: dict[str, Any] = Field(default_factory=dict)
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
class NotificationTargetUpdate(BaseModel):
|
||||
"""Schema for updating an existing notification target."""
|
||||
|
||||
name: str | None = Field(None, min_length=1, max_length=255)
|
||||
config: dict[str, Any] | None = None
|
||||
is_active: bool | None = None
|
||||
|
||||
|
||||
class PreferenceItem(BaseModel):
|
||||
"""A single preference toggle for one event+channel combination."""
|
||||
|
||||
is_enabled: bool
|
||||
target_id: int | None = None
|
||||
|
||||
|
||||
class PreferenceItemFull(BaseModel):
|
||||
"""Full preference item including event and channel type (used in bulk update)."""
|
||||
|
||||
event_type: str
|
||||
channel_type: str
|
||||
is_enabled: bool
|
||||
target_id: int | None = None
|
||||
|
||||
|
||||
class PreferencesUpdate(BaseModel):
|
||||
"""Bulk preferences update payload — a flat list of preference items."""
|
||||
|
||||
preferences: list[PreferenceItemFull]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _mask_email_config(config: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Return a copy of an email config dict with the password masked."""
|
||||
masked = dict(config)
|
||||
if masked.get("smtp_password"):
|
||||
masked["smtp_password"] = "****"
|
||||
return masked
|
||||
|
||||
|
||||
def _target_to_dict(target: UserNotificationTarget) -> dict[str, Any]:
|
||||
"""Serialize a UserNotificationTarget to a response dict, masking secrets."""
|
||||
config: dict[str, Any] = {}
|
||||
if target.config:
|
||||
try:
|
||||
config = json.loads(target.config)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
config = {}
|
||||
|
||||
if target.channel_type == "email":
|
||||
config = _mask_email_config(config)
|
||||
|
||||
return {
|
||||
"id": target.id,
|
||||
"channel_type": target.channel_type,
|
||||
"name": target.name,
|
||||
"config": config,
|
||||
"is_active": target.is_active,
|
||||
"created_at": target.created_at,
|
||||
"updated_at": target.updated_at,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inbox endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/inbox")
|
||||
async def list_inbox(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List in-app notifications for the authenticated user, newest first."""
|
||||
notifications = (
|
||||
db.query(InAppNotification)
|
||||
.filter(InAppNotification.owner_id == owner_id)
|
||||
.order_by(InAppNotification.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": n.id,
|
||||
"event_type": n.event_type,
|
||||
"title": n.title,
|
||||
"message": n.message,
|
||||
"is_read": n.is_read,
|
||||
"file_id": n.file_id,
|
||||
"created_at": n.created_at,
|
||||
}
|
||||
for n in notifications
|
||||
]
|
||||
|
||||
|
||||
@router.get("/inbox/unread-count")
|
||||
async def unread_count(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, int]:
|
||||
"""Return the number of unread in-app notifications."""
|
||||
count = (
|
||||
db.query(InAppNotification)
|
||||
.filter(InAppNotification.owner_id == owner_id, InAppNotification.is_read == False) # noqa: E712
|
||||
.count()
|
||||
)
|
||||
return {"count": count}
|
||||
|
||||
|
||||
@router.post("/inbox/{notification_id}/read", status_code=status.HTTP_200_OK)
|
||||
async def mark_read(
|
||||
notification_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Mark a single in-app notification as read."""
|
||||
notif = (
|
||||
db.query(InAppNotification)
|
||||
.filter(InAppNotification.id == notification_id, InAppNotification.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not notif:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Notification not found")
|
||||
try:
|
||||
notif.is_read = True
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
return {"detail": "Marked as read"}
|
||||
|
||||
|
||||
@router.post("/inbox/read-all", status_code=status.HTTP_200_OK)
|
||||
async def mark_all_read(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Mark all in-app notifications as read for the authenticated user."""
|
||||
try:
|
||||
db.query(InAppNotification).filter(
|
||||
InAppNotification.owner_id == owner_id,
|
||||
InAppNotification.is_read == False, # noqa: E712
|
||||
).update({"is_read": True})
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
return {"detail": "All notifications marked as read"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Notification target endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/targets")
|
||||
async def list_targets(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all notification targets for the authenticated user."""
|
||||
targets = (
|
||||
db.query(UserNotificationTarget)
|
||||
.filter(UserNotificationTarget.owner_id == owner_id)
|
||||
.order_by(UserNotificationTarget.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_target_to_dict(t) for t in targets]
|
||||
|
||||
|
||||
@router.post("/targets", status_code=status.HTTP_201_CREATED)
|
||||
async def create_target(
|
||||
body: NotificationTargetCreate,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new notification target (email or webhook)."""
|
||||
target = UserNotificationTarget(
|
||||
owner_id=owner_id,
|
||||
channel_type=body.channel_type,
|
||||
name=body.name,
|
||||
config=json.dumps(body.config),
|
||||
is_active=body.is_active,
|
||||
)
|
||||
try:
|
||||
db.add(target)
|
||||
db.commit()
|
||||
db.refresh(target)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Notification target created: id=%s owner=%s type=%s", target.id, owner_id, body.channel_type)
|
||||
return _target_to_dict(target)
|
||||
|
||||
|
||||
@router.put("/targets/{target_id}", status_code=status.HTTP_200_OK)
|
||||
async def update_target(
|
||||
target_id: int,
|
||||
body: NotificationTargetUpdate,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Update an existing notification target."""
|
||||
target = (
|
||||
db.query(UserNotificationTarget)
|
||||
.filter(UserNotificationTarget.id == target_id, UserNotificationTarget.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not target:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Target not found")
|
||||
|
||||
try:
|
||||
if body.name is not None:
|
||||
target.name = body.name
|
||||
if body.config is not None:
|
||||
# Merge new config over existing, preserving masked password field if unchanged
|
||||
existing_config: dict[str, Any] = {}
|
||||
if target.config:
|
||||
try:
|
||||
existing_config = json.loads(target.config)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
existing_config = {}
|
||||
merged = dict(existing_config)
|
||||
for k, v in body.config.items():
|
||||
# Skip writing back a masked password placeholder
|
||||
if k == "smtp_password" and v == "****":
|
||||
continue
|
||||
merged[k] = v
|
||||
target.config = json.dumps(merged)
|
||||
if body.is_active is not None:
|
||||
target.is_active = body.is_active
|
||||
db.commit()
|
||||
db.refresh(target)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Notification target updated: id=%s owner=%s", target_id, owner_id)
|
||||
return _target_to_dict(target)
|
||||
|
||||
|
||||
@router.delete("/targets/{target_id}", status_code=status.HTTP_200_OK)
|
||||
async def delete_target(
|
||||
target_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Delete a notification target and its associated preferences."""
|
||||
target = (
|
||||
db.query(UserNotificationTarget)
|
||||
.filter(UserNotificationTarget.id == target_id, UserNotificationTarget.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not target:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Target not found")
|
||||
|
||||
try:
|
||||
# Remove any preferences that reference this target
|
||||
db.query(UserNotificationPreference).filter(
|
||||
UserNotificationPreference.owner_id == owner_id,
|
||||
UserNotificationPreference.target_id == target_id,
|
||||
).delete()
|
||||
db.delete(target)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Notification target deleted: id=%s owner=%s", target_id, owner_id)
|
||||
return {"detail": "Target deleted"}
|
||||
|
||||
|
||||
@router.post("/targets/{target_id}/test", status_code=status.HTTP_200_OK)
|
||||
async def test_target(
|
||||
target_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Send a test notification to the specified target."""
|
||||
target = (
|
||||
db.query(UserNotificationTarget)
|
||||
.filter(UserNotificationTarget.id == target_id, UserNotificationTarget.owner_id == owner_id)
|
||||
.first()
|
||||
)
|
||||
if not target:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Target not found")
|
||||
|
||||
config: dict[str, Any] = {}
|
||||
if target.config:
|
||||
try:
|
||||
config = json.loads(target.config)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
config = {}
|
||||
|
||||
title = "DocuElevate Test Notification"
|
||||
message = f"This is a test notification from DocuElevate for target '{target.name}'."
|
||||
|
||||
if target.channel_type == "email":
|
||||
from app.utils.user_notification import _send_email_notification
|
||||
|
||||
ok = _send_email_notification(config, title, message)
|
||||
elif target.channel_type == "webhook":
|
||||
from app.utils.user_notification import _send_webhook_notification
|
||||
|
||||
ok = _send_webhook_notification(config, "test", title, message)
|
||||
else:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Unknown channel type")
|
||||
|
||||
if not ok:
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail="Failed to send test notification")
|
||||
|
||||
return {"detail": "Test notification sent"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Preferences endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/preferences")
|
||||
async def get_preferences(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Return all notification preferences for the authenticated user.
|
||||
|
||||
Response structure:
|
||||
{
|
||||
"event_types": ["document.processed", "document.failed"],
|
||||
"event_labels": {"document.processed": "Document Processed", ...},
|
||||
"preferences": {
|
||||
"document.processed": {
|
||||
"in_app": {"is_enabled": true, "target_id": null},
|
||||
"email": {"is_enabled": false, "target_id": 1},
|
||||
...
|
||||
}
|
||||
}
|
||||
}
|
||||
"""
|
||||
prefs = db.query(UserNotificationPreference).filter(UserNotificationPreference.owner_id == owner_id).all()
|
||||
|
||||
# Build nested dict: event_type -> channel_type -> {is_enabled, target_id}
|
||||
result: dict[str, dict[str, dict[str, Any]]] = {}
|
||||
for pref in prefs:
|
||||
result.setdefault(pref.event_type, {})[pref.channel_type] = {
|
||||
"is_enabled": pref.is_enabled,
|
||||
"target_id": pref.target_id,
|
||||
}
|
||||
|
||||
return {
|
||||
"event_types": list(USER_EVENT_LABELS.keys()),
|
||||
"event_labels": USER_EVENT_LABELS,
|
||||
"preferences": result,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/preferences", status_code=status.HTTP_200_OK)
|
||||
async def update_preferences(
|
||||
body: PreferencesUpdate,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Bulk upsert notification preferences for the authenticated user.
|
||||
|
||||
Validates that any referenced target_id belongs to the requesting user.
|
||||
"""
|
||||
# Collect all target IDs referenced in the payload for ownership validation
|
||||
referenced_target_ids: set[int] = set()
|
||||
for item in body.preferences:
|
||||
if item.target_id is not None:
|
||||
referenced_target_ids.add(item.target_id)
|
||||
|
||||
if referenced_target_ids:
|
||||
owned_ids = {
|
||||
row.id
|
||||
for row in db.query(UserNotificationTarget.id)
|
||||
.filter(
|
||||
UserNotificationTarget.owner_id == owner_id,
|
||||
UserNotificationTarget.id.in_(referenced_target_ids),
|
||||
)
|
||||
.all()
|
||||
}
|
||||
invalid = referenced_target_ids - owned_ids
|
||||
if invalid:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid or inaccessible target_id(s): {sorted(invalid)}",
|
||||
)
|
||||
|
||||
try:
|
||||
for item in body.preferences:
|
||||
existing = (
|
||||
db.query(UserNotificationPreference)
|
||||
.filter(
|
||||
UserNotificationPreference.owner_id == owner_id,
|
||||
UserNotificationPreference.event_type == item.event_type,
|
||||
UserNotificationPreference.channel_type == item.channel_type,
|
||||
UserNotificationPreference.target_id == item.target_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
existing.is_enabled = item.is_enabled
|
||||
else:
|
||||
db.add(
|
||||
UserNotificationPreference(
|
||||
owner_id=owner_id,
|
||||
event_type=item.event_type,
|
||||
channel_type=item.channel_type,
|
||||
target_id=item.target_id,
|
||||
is_enabled=item.is_enabled,
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Notification preferences updated for owner=%s", owner_id)
|
||||
return {"detail": "Preferences updated"}
|
||||
@@ -0,0 +1,253 @@
|
||||
"""API endpoints for the user onboarding wizard.
|
||||
|
||||
Provides a REST interface for the multi-step onboarding flow, allowing
|
||||
authenticated users to set their profile, choose a subscription plan,
|
||||
select a storage destination, and mark onboarding as complete.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import UserProfile
|
||||
from app.utils.subscription import TIERS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/onboarding", tags=["onboarding"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_current_user_id(request: Request) -> str:
|
||||
"""Extract the stable user_id from the session using the same priority as _ensure_user_profile.
|
||||
|
||||
Priority: sub → preferred_username → email → id.
|
||||
|
||||
Raises:
|
||||
HTTPException: 401 if the user is not authenticated.
|
||||
"""
|
||||
user = request.session.get("user")
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
user_id = user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id")
|
||||
if not user_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return user_id
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProfileBody(BaseModel):
|
||||
"""Body for the profile step of the onboarding wizard."""
|
||||
|
||||
display_name: str | None = Field(default=None, max_length=255)
|
||||
contact_email: str | None = Field(default=None, max_length=255)
|
||||
|
||||
|
||||
class PlanBody(BaseModel):
|
||||
"""Body for the plan step of the onboarding wizard."""
|
||||
|
||||
subscription_tier: str
|
||||
billing_cycle: str = Field(pattern="^(monthly|yearly)$")
|
||||
|
||||
|
||||
class StorageBody(BaseModel):
|
||||
"""Body for the storage step of the onboarding wizard."""
|
||||
|
||||
preferred_destination: str | None = Field(default=None, max_length=50)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _profile_to_dict(profile: UserProfile) -> dict[str, Any]:
|
||||
"""Serialize a UserProfile to a plain dict for API responses."""
|
||||
return {
|
||||
"user_id": profile.user_id,
|
||||
"display_name": profile.display_name,
|
||||
"contact_email": profile.contact_email,
|
||||
"subscription_tier": profile.subscription_tier or "free",
|
||||
"subscription_billing_cycle": profile.subscription_billing_cycle or "monthly",
|
||||
"preferred_destination": profile.preferred_destination,
|
||||
"onboarding_completed": bool(profile.onboarding_completed),
|
||||
"onboarding_completed_at": profile.onboarding_completed_at.isoformat()
|
||||
if profile.onboarding_completed_at
|
||||
else None,
|
||||
}
|
||||
|
||||
|
||||
def _get_or_create_profile(db: Session, user_id: str) -> UserProfile:
|
||||
"""Return the UserProfile for *user_id*, creating one if it does not exist."""
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile is None:
|
||||
profile = UserProfile(user_id=user_id)
|
||||
db.add(profile)
|
||||
db.flush()
|
||||
return profile
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/status", summary="Get onboarding status for the current user")
|
||||
def get_onboarding_status(request: Request, db: DbSession) -> dict[str, Any]:
|
||||
"""Return whether onboarding has been completed and the current step.
|
||||
|
||||
The ``step`` field is a best-effort estimate: 1 for brand-new profiles,
|
||||
further along when partial data has already been saved.
|
||||
"""
|
||||
user_id = _get_current_user_id(request)
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
|
||||
if profile is None:
|
||||
return {"completed": False, "step": 1, "profile": None}
|
||||
|
||||
# Derive a sensible current step from saved data so the wizard can resume.
|
||||
step = 1
|
||||
if profile.display_name or profile.contact_email:
|
||||
step = 2
|
||||
if profile.subscription_tier and profile.subscription_tier != "free":
|
||||
step = 3
|
||||
if profile.preferred_destination:
|
||||
step = 4
|
||||
if profile.onboarding_completed:
|
||||
step = 5
|
||||
|
||||
return {
|
||||
"completed": bool(profile.onboarding_completed),
|
||||
"step": step,
|
||||
"profile": _profile_to_dict(profile),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/profile", summary="Save profile step during onboarding")
|
||||
def save_profile(request: Request, body: ProfileBody, db: DbSession) -> dict[str, Any]:
|
||||
"""Persist the user's display name and contact email from the profile step."""
|
||||
user_id = _get_current_user_id(request)
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
|
||||
if body.display_name is not None:
|
||||
profile.display_name = body.display_name
|
||||
if body.contact_email is not None:
|
||||
profile.contact_email = body.contact_email
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Onboarding: saved profile for user %s", user_id)
|
||||
return _profile_to_dict(profile)
|
||||
|
||||
|
||||
@router.post("/plan", summary="Save plan selection during onboarding")
|
||||
def save_plan(request: Request, body: PlanBody, db: DbSession) -> dict[str, Any]:
|
||||
"""Persist the chosen subscription tier and billing cycle from the plan step.
|
||||
|
||||
Raises:
|
||||
HTTPException: 422 if the tier is not a recognised value.
|
||||
"""
|
||||
user_id = _get_current_user_id(request)
|
||||
|
||||
if body.subscription_tier not in TIERS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Invalid subscription_tier '{body.subscription_tier}'. Valid values: {list(TIERS.keys())}",
|
||||
)
|
||||
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
old_tier = profile.subscription_tier or "free"
|
||||
profile.subscription_tier = body.subscription_tier
|
||||
profile.subscription_billing_cycle = body.billing_cycle
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Onboarding: saved plan %s/%s", body.subscription_tier, body.billing_cycle)
|
||||
|
||||
# Notify admins and fire webhook when the plan actually changes
|
||||
if old_tier != body.subscription_tier:
|
||||
try:
|
||||
from app.utils.notification import notify_plan_changed
|
||||
from app.utils.webhook import dispatch_webhook_event
|
||||
|
||||
notify_plan_changed(user_id, old_tier=old_tier, new_tier=body.subscription_tier, changed_by="user")
|
||||
dispatch_webhook_event(
|
||||
"user.plan_changed",
|
||||
{
|
||||
"user_id": user_id,
|
||||
"old_tier": old_tier,
|
||||
"new_tier": body.subscription_tier,
|
||||
"billing_cycle": body.billing_cycle,
|
||||
"changed_by": "user",
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to send plan-change notification/webhook for user %s", user_id)
|
||||
|
||||
return _profile_to_dict(profile)
|
||||
|
||||
|
||||
@router.post("/storage", summary="Save storage preference during onboarding")
|
||||
def save_storage(request: Request, body: StorageBody, db: DbSession) -> dict[str, Any]:
|
||||
"""Persist the user's preferred storage destination from the storage step."""
|
||||
user_id = _get_current_user_id(request)
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
profile.preferred_destination = body.preferred_destination
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Onboarding: saved storage preference '%s' for user %s", body.preferred_destination, user_id)
|
||||
return _profile_to_dict(profile)
|
||||
|
||||
|
||||
@router.post("/complete", summary="Mark onboarding as completed")
|
||||
def complete_onboarding(request: Request, db: DbSession) -> dict[str, Any]:
|
||||
"""Set onboarding_completed=True, record the completion timestamp, and return the post-onboarding redirect URL.
|
||||
|
||||
The redirect URL is read from ``request.session["post_onboarding_redirect"]`` (stored by
|
||||
``oauth_callback`` when it reroutes a first-time user to the wizard) and defaults to
|
||||
``/upload`` when the session key is absent.
|
||||
"""
|
||||
user_id = _get_current_user_id(request)
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
profile.onboarding_completed = True
|
||||
profile.onboarding_completed_at = datetime.now(tz=timezone.utc)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
redirect_url = request.session.pop("post_onboarding_redirect", "/upload")
|
||||
logger.info("Onboarding: completed for user %s, redirecting to %s", user_id, redirect_url)
|
||||
return {"success": True, "redirect_url": redirect_url}
|
||||
@@ -0,0 +1,927 @@
|
||||
"""
|
||||
Pipelines API endpoints.
|
||||
|
||||
Provides full CRUD for processing pipelines and their steps. Pipelines are
|
||||
user-specific: regular users can only manage their own pipelines, while admins
|
||||
can also create and manage *system default* pipelines (owner_id = NULL) that
|
||||
are visible to all users.
|
||||
|
||||
Built-in step types are exposed via GET /api/pipelines/step-types so that UIs
|
||||
can render the correct configuration form without hard-coding the catalogue.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import get_current_user, get_current_user_id, require_login
|
||||
from app.database import get_db
|
||||
from app.models import Pipeline, PipelineStep
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/pipelines", tags=["pipelines"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Built-in step type catalogue
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
PIPELINE_STEP_TYPES: dict[str, dict[str, Any]] = {
|
||||
"convert_to_pdf": {
|
||||
"label": "Convert to PDF",
|
||||
"description": "Convert non-PDF documents to PDF format using Gotenberg.",
|
||||
"config_schema": {},
|
||||
},
|
||||
"check_duplicates": {
|
||||
"label": "Check for Duplicates",
|
||||
"description": "Compare file hash against existing documents to detect duplicates.",
|
||||
"config_schema": {},
|
||||
},
|
||||
"ocr": {
|
||||
"label": "OCR Processing",
|
||||
"description": "Extract text using Azure Document Intelligence or local Tesseract.",
|
||||
"config_schema": {
|
||||
"force_cloud_ocr": {
|
||||
"type": "boolean",
|
||||
"default": False,
|
||||
"description": "Always use cloud OCR even if the PDF already has embedded text.",
|
||||
},
|
||||
"ocr_language": {
|
||||
"type": "select",
|
||||
"default": "auto",
|
||||
"description": (
|
||||
"Language(s) used for OCR text extraction. Applies to Tesseract and EasyOCR "
|
||||
"providers; Azure and Mistral perform auto-detection by default. "
|
||||
"Use Tesseract codes such as 'eng', 'deu', or 'eng+deu' for multi-language "
|
||||
"documents. 'auto' falls back to the global system setting."
|
||||
),
|
||||
"options": [
|
||||
{"value": "auto", "label": "Auto (use system default)"},
|
||||
{"value": "ara", "label": "Arabic"},
|
||||
{"value": "chi_sim", "label": "Chinese (Simplified)"},
|
||||
{"value": "chi_tra", "label": "Chinese (Traditional)"},
|
||||
{"value": "ces", "label": "Czech"},
|
||||
{"value": "dan", "label": "Danish"},
|
||||
{"value": "nld", "label": "Dutch"},
|
||||
{"value": "eng", "label": "English"},
|
||||
{"value": "fin", "label": "Finnish"},
|
||||
{"value": "fra", "label": "French"},
|
||||
{"value": "deu", "label": "German"},
|
||||
{"value": "ell", "label": "Greek"},
|
||||
{"value": "heb", "label": "Hebrew"},
|
||||
{"value": "hin", "label": "Hindi"},
|
||||
{"value": "hun", "label": "Hungarian"},
|
||||
{"value": "ita", "label": "Italian"},
|
||||
{"value": "jpn", "label": "Japanese"},
|
||||
{"value": "kor", "label": "Korean"},
|
||||
{"value": "nor", "label": "Norwegian"},
|
||||
{"value": "pol", "label": "Polish"},
|
||||
{"value": "por", "label": "Portuguese"},
|
||||
{"value": "ron", "label": "Romanian"},
|
||||
{"value": "rus", "label": "Russian"},
|
||||
{"value": "spa", "label": "Spanish"},
|
||||
{"value": "swe", "label": "Swedish"},
|
||||
{"value": "tha", "label": "Thai"},
|
||||
{"value": "tur", "label": "Turkish"},
|
||||
{"value": "ukr", "label": "Ukrainian"},
|
||||
{"value": "vie", "label": "Vietnamese"},
|
||||
],
|
||||
},
|
||||
},
|
||||
},
|
||||
"extract_metadata": {
|
||||
"label": "Metadata Extraction",
|
||||
"description": "Extract structured metadata (document type, sender, recipient, tags) using AI.",
|
||||
"config_schema": {},
|
||||
},
|
||||
"embed_metadata": {
|
||||
"label": "Embed Metadata into PDF",
|
||||
"description": "Write the extracted metadata into the PDF document properties.",
|
||||
"config_schema": {},
|
||||
},
|
||||
"compute_embedding": {
|
||||
"label": "Compute Text Embedding",
|
||||
"description": "Compute semantic text embeddings for full-text and similarity search.",
|
||||
"config_schema": {},
|
||||
},
|
||||
"send_to_destinations": {
|
||||
"label": "Send to Storage Destinations",
|
||||
"description": "Upload the processed document to all configured storage destinations.",
|
||||
"config_schema": {},
|
||||
},
|
||||
"classify": {
|
||||
"label": "Document Classification",
|
||||
"description": "Classify the document type using AI without full metadata extraction.",
|
||||
"config_schema": {},
|
||||
},
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
MAX_STEPS_PER_PIPELINE = 50
|
||||
MAX_NAME_LENGTH = 255
|
||||
|
||||
|
||||
def _get_user_id(request: Request) -> str:
|
||||
"""Return a stable user identifier from the session.
|
||||
|
||||
Delegates to :func:`app.auth.get_current_user_id` so the same fallback
|
||||
logic ("anonymous") is used consistently throughout the application.
|
||||
"""
|
||||
return get_current_user_id(request)
|
||||
|
||||
|
||||
def _is_admin(request: Request) -> bool:
|
||||
"""Return True if the current session user is an admin."""
|
||||
user = get_current_user(request)
|
||||
return bool(user and user.get("is_admin"))
|
||||
|
||||
|
||||
def _can_access_pipeline(pipeline: Pipeline, user_id: str, admin: bool) -> bool:
|
||||
"""Return True if the user may read or write this pipeline."""
|
||||
# System pipelines (owner_id=NULL) are readable by everyone; only admins can write
|
||||
if pipeline.owner_id is None:
|
||||
return True
|
||||
# Own pipeline
|
||||
return pipeline.owner_id == user_id or admin
|
||||
|
||||
|
||||
def _can_write_pipeline(pipeline: Pipeline, user_id: str, admin: bool) -> bool:
|
||||
"""Return True if the user may create/update/delete this pipeline."""
|
||||
if pipeline.owner_id is None:
|
||||
return admin
|
||||
return pipeline.owner_id == user_id or admin
|
||||
|
||||
|
||||
def _serialize_step(step: PipelineStep) -> dict[str, Any]:
|
||||
return {
|
||||
"id": step.id,
|
||||
"pipeline_id": step.pipeline_id,
|
||||
"position": step.position,
|
||||
"step_type": step.step_type,
|
||||
"label": step.label,
|
||||
"config": json.loads(step.config) if step.config else {},
|
||||
"enabled": step.enabled,
|
||||
"created_at": step.created_at.isoformat() if step.created_at else None,
|
||||
"updated_at": step.updated_at.isoformat() if step.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _serialize_pipeline(pipeline: Pipeline, include_steps: bool = False, db: Session | None = None) -> dict[str, Any]:
|
||||
data: dict[str, Any] = {
|
||||
"id": pipeline.id,
|
||||
"owner_id": pipeline.owner_id,
|
||||
"name": pipeline.name,
|
||||
"description": pipeline.description,
|
||||
"is_default": pipeline.is_default,
|
||||
"is_active": pipeline.is_active,
|
||||
"created_at": pipeline.created_at.isoformat() if pipeline.created_at else None,
|
||||
"updated_at": pipeline.updated_at.isoformat() if pipeline.updated_at else None,
|
||||
}
|
||||
if include_steps and db is not None:
|
||||
steps = (
|
||||
db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline.id).order_by(PipelineStep.position).all()
|
||||
)
|
||||
data["steps"] = [_serialize_step(s) for s in steps]
|
||||
return data
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PipelineCreate(BaseModel):
|
||||
"""Body for creating a pipeline."""
|
||||
|
||||
name: str = Field(..., max_length=MAX_NAME_LENGTH, description="Human-readable pipeline name")
|
||||
description: str | None = Field(default=None, max_length=4096)
|
||||
is_default: bool = Field(default=False)
|
||||
is_active: bool = Field(default=True)
|
||||
|
||||
|
||||
class PipelineUpdate(BaseModel):
|
||||
"""Body for updating a pipeline (all fields optional)."""
|
||||
|
||||
name: str | None = Field(default=None, max_length=MAX_NAME_LENGTH)
|
||||
description: str | None = Field(default=None, max_length=4096)
|
||||
is_default: bool | None = None
|
||||
is_active: bool | None = None
|
||||
|
||||
|
||||
class PipelineStepCreate(BaseModel):
|
||||
"""Body for adding a step to a pipeline."""
|
||||
|
||||
step_type: str = Field(..., description="One of the recognised step type keys")
|
||||
label: str | None = Field(default=None, max_length=MAX_NAME_LENGTH)
|
||||
config: dict[str, Any] = Field(default_factory=dict)
|
||||
enabled: bool = Field(default=True)
|
||||
position: int | None = Field(default=None, ge=0, description="Insertion position; appended at end if omitted")
|
||||
|
||||
|
||||
class PipelineStepUpdate(BaseModel):
|
||||
"""Body for updating a pipeline step (all fields optional)."""
|
||||
|
||||
step_type: str | None = None
|
||||
label: str | None = Field(default=None, max_length=MAX_NAME_LENGTH)
|
||||
config: dict[str, Any] | None = None
|
||||
enabled: bool | None = None
|
||||
position: int | None = Field(default=None, ge=0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Step-types catalogue endpoint (no auth required — it's public metadata)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/step-types")
|
||||
def list_step_types() -> dict[str, Any]:
|
||||
"""Return the catalogue of built-in pipeline step types.
|
||||
|
||||
Returns:
|
||||
A mapping of step_type key → metadata (label, description, config_schema).
|
||||
"""
|
||||
return PIPELINE_STEP_TYPES
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pipeline CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("")
|
||||
@require_login
|
||||
def list_pipelines(request: Request, db: DbSession) -> list[dict[str, Any]]:
|
||||
"""List pipelines visible to the current user.
|
||||
|
||||
Regular users see: their own pipelines + system pipelines (owner_id=NULL).
|
||||
Admins see: all pipelines from all users.
|
||||
|
||||
Returns:
|
||||
A list of pipeline objects (without steps — use GET /pipelines/{id} for steps).
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
if admin:
|
||||
pipelines = db.query(Pipeline).order_by(Pipeline.owner_id.nullsfirst(), Pipeline.name).all()
|
||||
else:
|
||||
pipelines = (
|
||||
db.query(Pipeline)
|
||||
.filter((Pipeline.owner_id == user_id) | (Pipeline.owner_id.is_(None)))
|
||||
.order_by(Pipeline.owner_id.nullsfirst(), Pipeline.name)
|
||||
.all()
|
||||
)
|
||||
|
||||
return [_serialize_pipeline(p) for p in pipelines]
|
||||
|
||||
|
||||
@router.post("", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
def create_pipeline(request: Request, db: DbSession, body: PipelineCreate) -> dict[str, Any]:
|
||||
"""Create a new pipeline for the current user.
|
||||
|
||||
Admins can create system default pipelines by passing ``owner_id=null``
|
||||
via the body — however, that is handled implicitly: to create a system
|
||||
pipeline, call ``POST /api/admin/pipelines`` (admin endpoint) instead.
|
||||
Regular users always get their own user_id as owner.
|
||||
|
||||
Returns:
|
||||
The created pipeline object.
|
||||
|
||||
Raises:
|
||||
HTTPException 409: If a pipeline with the same name already exists for this owner.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
name = body.name.strip() if body.name else ""
|
||||
if not name:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="name is required",
|
||||
)
|
||||
|
||||
# Enforce unique name per owner
|
||||
existing = db.query(Pipeline).filter(Pipeline.owner_id == user_id, Pipeline.name == name).first()
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"A pipeline named '{name}' already exists",
|
||||
)
|
||||
|
||||
# If this pipeline is marked as default, unset the existing default for this user
|
||||
if body.is_default:
|
||||
_unset_default(db, user_id)
|
||||
|
||||
pipeline = Pipeline(
|
||||
owner_id=user_id,
|
||||
name=name,
|
||||
description=body.description,
|
||||
is_default=body.is_default,
|
||||
is_active=body.is_active,
|
||||
)
|
||||
try:
|
||||
db.add(pipeline)
|
||||
db.commit()
|
||||
db.refresh(pipeline)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create pipeline user=%s", user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to create pipeline",
|
||||
)
|
||||
|
||||
logger.info("Pipeline created: id=%s, owner=%s, name=%r", pipeline.id, user_id, name)
|
||||
return _serialize_pipeline(pipeline)
|
||||
|
||||
|
||||
@router.get("/{pipeline_id}")
|
||||
@require_login
|
||||
def get_pipeline(pipeline_id: int, request: Request, db: DbSession) -> dict[str, Any]:
|
||||
"""Return a single pipeline with its steps.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The ID of the pipeline.
|
||||
|
||||
Returns:
|
||||
The pipeline object including its ordered steps.
|
||||
|
||||
Raises:
|
||||
HTTPException 404: If the pipeline does not exist or is not accessible.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
return _serialize_pipeline(pipeline, include_steps=True, db=db)
|
||||
|
||||
|
||||
@router.put("/{pipeline_id}")
|
||||
@require_login
|
||||
def update_pipeline(pipeline_id: int, request: Request, db: DbSession, body: PipelineUpdate) -> dict[str, Any]:
|
||||
"""Update a pipeline's metadata.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The ID of the pipeline to update.
|
||||
|
||||
Returns:
|
||||
The updated pipeline object.
|
||||
|
||||
Raises:
|
||||
HTTPException 403: If the caller does not own this pipeline.
|
||||
HTTPException 404: If the pipeline does not exist.
|
||||
HTTPException 409: If the new name conflicts with an existing pipeline.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
if not _can_write_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this pipeline")
|
||||
|
||||
if body.name is not None:
|
||||
new_name = body.name.strip()
|
||||
if not new_name:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="name must not be empty",
|
||||
)
|
||||
if new_name != pipeline.name:
|
||||
conflict = (
|
||||
db.query(Pipeline)
|
||||
.filter(Pipeline.owner_id == pipeline.owner_id, Pipeline.name == new_name, Pipeline.id != pipeline_id)
|
||||
.first()
|
||||
)
|
||||
if conflict:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"A pipeline named '{new_name}' already exists",
|
||||
)
|
||||
pipeline.name = new_name
|
||||
|
||||
if body.description is not None:
|
||||
pipeline.description = body.description
|
||||
|
||||
if body.is_active is not None:
|
||||
pipeline.is_active = body.is_active
|
||||
|
||||
if body.is_default is not None:
|
||||
if body.is_default and not pipeline.is_default:
|
||||
_unset_default(db, pipeline.owner_id)
|
||||
pipeline.is_default = body.is_default
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(pipeline)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to update pipeline id=%s", pipeline_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update pipeline",
|
||||
)
|
||||
|
||||
logger.info("Pipeline updated: id=%s, user=%s", pipeline_id, user_id)
|
||||
return _serialize_pipeline(pipeline, include_steps=True, db=db)
|
||||
|
||||
|
||||
@router.delete("/{pipeline_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
def delete_pipeline(pipeline_id: int, request: Request, db: DbSession) -> None:
|
||||
"""Delete a pipeline and all its steps.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The ID of the pipeline to delete.
|
||||
|
||||
Raises:
|
||||
HTTPException 403: If the caller does not own this pipeline.
|
||||
HTTPException 404: If the pipeline does not exist.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
if not _can_write_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot delete this pipeline")
|
||||
|
||||
try:
|
||||
db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline_id).delete()
|
||||
db.delete(pipeline)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to delete pipeline id=%s", pipeline_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to delete pipeline",
|
||||
)
|
||||
|
||||
logger.info("Pipeline deleted: id=%s, user=%s", pipeline_id, user_id)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Admin-only: create system (owner_id=NULL) pipeline
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/admin/system", status_code=status.HTTP_201_CREATED, tags=["admin-pipelines"])
|
||||
@require_login
|
||||
def create_system_pipeline(request: Request, db: DbSession, body: PipelineCreate) -> dict[str, Any]:
|
||||
"""Create a system-level (owner_id=NULL) default pipeline. Admin only.
|
||||
|
||||
System pipelines are visible to all users and can be set as the global
|
||||
default. Only admins may create them.
|
||||
|
||||
Returns:
|
||||
The created system pipeline.
|
||||
|
||||
Raises:
|
||||
HTTPException 403: If the caller is not an admin.
|
||||
HTTPException 409: If a system pipeline with the same name already exists.
|
||||
"""
|
||||
if not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
|
||||
name = body.name.strip() if body.name else ""
|
||||
if not name:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="name is required",
|
||||
)
|
||||
|
||||
existing = db.query(Pipeline).filter(Pipeline.owner_id.is_(None), Pipeline.name == name).first()
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"A system pipeline named '{name}' already exists",
|
||||
)
|
||||
|
||||
if body.is_default:
|
||||
_unset_default(db, None)
|
||||
|
||||
pipeline = Pipeline(
|
||||
owner_id=None,
|
||||
name=name,
|
||||
description=body.description,
|
||||
is_default=body.is_default,
|
||||
is_active=body.is_active,
|
||||
)
|
||||
try:
|
||||
db.add(pipeline)
|
||||
db.commit()
|
||||
db.refresh(pipeline)
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to create system pipeline: {exc}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to create system pipeline",
|
||||
)
|
||||
|
||||
logger.info(f"System pipeline created: id={pipeline.id}, name={name!r}")
|
||||
return _serialize_pipeline(pipeline)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Step management
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/{pipeline_id}/steps", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
def add_step(pipeline_id: int, request: Request, db: DbSession, body: PipelineStepCreate) -> dict[str, Any]:
|
||||
"""Add a step to a pipeline.
|
||||
|
||||
Steps are automatically appended at the end unless an explicit ``position``
|
||||
is supplied. All existing steps at or after the insertion position are
|
||||
shifted forward by one.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The pipeline to add the step to.
|
||||
|
||||
Returns:
|
||||
The created step object.
|
||||
|
||||
Raises:
|
||||
HTTPException 403: If the caller cannot modify this pipeline.
|
||||
HTTPException 404: If the pipeline does not exist.
|
||||
HTTPException 422: If the step_type is not recognised.
|
||||
HTTPException 409: If the maximum number of steps per pipeline is reached.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
if not _can_write_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this pipeline")
|
||||
|
||||
if body.step_type not in PIPELINE_STEP_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Unknown step type '{body.step_type}'. Valid types: {sorted(PIPELINE_STEP_TYPES)}",
|
||||
)
|
||||
|
||||
current_count = db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline_id).count()
|
||||
if current_count >= MAX_STEPS_PER_PIPELINE:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"Maximum of {MAX_STEPS_PER_PIPELINE} steps per pipeline reached",
|
||||
)
|
||||
|
||||
# Determine insertion position
|
||||
if body.position is None:
|
||||
max_pos = (
|
||||
db.query(PipelineStep.position)
|
||||
.filter(PipelineStep.pipeline_id == pipeline_id)
|
||||
.order_by(PipelineStep.position.desc())
|
||||
.first()
|
||||
)
|
||||
insert_pos = (max_pos[0] + 1) if max_pos else 0
|
||||
else:
|
||||
insert_pos = body.position
|
||||
# Shift existing steps
|
||||
steps_to_shift = (
|
||||
db.query(PipelineStep)
|
||||
.filter(PipelineStep.pipeline_id == pipeline_id, PipelineStep.position >= insert_pos)
|
||||
.all()
|
||||
)
|
||||
for s in steps_to_shift:
|
||||
s.position += 1
|
||||
|
||||
step = PipelineStep(
|
||||
pipeline_id=pipeline_id,
|
||||
position=insert_pos,
|
||||
step_type=body.step_type,
|
||||
label=body.label,
|
||||
config=json.dumps(body.config) if body.config else None,
|
||||
enabled=body.enabled,
|
||||
)
|
||||
try:
|
||||
db.add(step)
|
||||
db.commit()
|
||||
db.refresh(step)
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to add step to pipeline id={pipeline_id}: {exc}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to add step",
|
||||
)
|
||||
|
||||
logger.info(f"Step added: pipeline={pipeline_id}, step_type={body.step_type!r}, pos={insert_pos}")
|
||||
return _serialize_step(step)
|
||||
|
||||
|
||||
@router.put("/{pipeline_id}/steps/reorder")
|
||||
@require_login
|
||||
def reorder_steps(
|
||||
pipeline_id: int,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
step_ids: list[int] = Body(..., description="Ordered list of step IDs representing the new order"),
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Replace the step order for a pipeline.
|
||||
|
||||
Provide a complete ordered list of *all* step IDs. Their ``position``
|
||||
values will be reassigned 0, 1, 2, … in the given order.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The pipeline whose steps are being reordered.
|
||||
|
||||
Returns:
|
||||
The updated, ordered list of step objects.
|
||||
|
||||
Raises:
|
||||
HTTPException 422: If the provided list does not contain exactly the
|
||||
current set of step IDs for this pipeline.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
if not _can_write_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this pipeline")
|
||||
|
||||
existing_steps = db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline_id).all()
|
||||
existing_ids = {s.id for s in existing_steps}
|
||||
|
||||
if set(step_ids) != existing_ids or len(step_ids) != len(existing_ids):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="step_ids must contain exactly the current step IDs for this pipeline",
|
||||
)
|
||||
|
||||
step_map = {s.id: s for s in existing_steps}
|
||||
for pos, sid in enumerate(step_ids):
|
||||
step_map[sid].position = pos
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to reorder steps for pipeline id={pipeline_id}: {exc}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to reorder steps",
|
||||
)
|
||||
|
||||
updated = (
|
||||
db.query(PipelineStep).filter(PipelineStep.pipeline_id == pipeline_id).order_by(PipelineStep.position).all()
|
||||
)
|
||||
return [_serialize_step(s) for s in updated]
|
||||
|
||||
|
||||
@router.put("/{pipeline_id}/steps/{step_id}")
|
||||
@require_login
|
||||
def update_step(
|
||||
pipeline_id: int, step_id: int, request: Request, db: DbSession, body: PipelineStepUpdate
|
||||
) -> dict[str, Any]:
|
||||
"""Update an existing pipeline step.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The owning pipeline.
|
||||
step_id: The step to update.
|
||||
|
||||
Returns:
|
||||
The updated step object.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
if not _can_write_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this pipeline")
|
||||
|
||||
step = db.query(PipelineStep).filter(PipelineStep.id == step_id, PipelineStep.pipeline_id == pipeline_id).first()
|
||||
if not step:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Step not found")
|
||||
|
||||
if body.step_type is not None:
|
||||
if body.step_type not in PIPELINE_STEP_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Unknown step type '{body.step_type}'",
|
||||
)
|
||||
step.step_type = body.step_type
|
||||
|
||||
if body.label is not None:
|
||||
step.label = body.label
|
||||
|
||||
if body.config is not None:
|
||||
step.config = json.dumps(body.config)
|
||||
|
||||
if body.enabled is not None:
|
||||
step.enabled = body.enabled
|
||||
|
||||
if body.position is not None and body.position != step.position:
|
||||
old_pos = step.position
|
||||
new_pos = body.position
|
||||
if new_pos > old_pos:
|
||||
# Moving down: shift intervening steps up
|
||||
db.query(PipelineStep).filter(
|
||||
PipelineStep.pipeline_id == pipeline_id,
|
||||
PipelineStep.position > old_pos,
|
||||
PipelineStep.position <= new_pos,
|
||||
PipelineStep.id != step_id,
|
||||
).update({"position": PipelineStep.position - 1})
|
||||
else:
|
||||
# Moving up: shift intervening steps down
|
||||
db.query(PipelineStep).filter(
|
||||
PipelineStep.pipeline_id == pipeline_id,
|
||||
PipelineStep.position >= new_pos,
|
||||
PipelineStep.position < old_pos,
|
||||
PipelineStep.id != step_id,
|
||||
).update({"position": PipelineStep.position + 1})
|
||||
step.position = new_pos
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(step)
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to update step id={step_id}: {exc}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update step",
|
||||
)
|
||||
|
||||
return _serialize_step(step)
|
||||
|
||||
|
||||
@router.delete("/{pipeline_id}/steps/{step_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
def delete_step(pipeline_id: int, step_id: int, request: Request, db: DbSession) -> None:
|
||||
"""Delete a step from a pipeline.
|
||||
|
||||
Path Parameters:
|
||||
pipeline_id: The owning pipeline.
|
||||
step_id: The step to delete.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == pipeline_id).first()
|
||||
if not pipeline or not _can_access_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Pipeline not found")
|
||||
|
||||
if not _can_write_pipeline(pipeline, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this pipeline")
|
||||
|
||||
step = db.query(PipelineStep).filter(PipelineStep.id == step_id, PipelineStep.pipeline_id == pipeline_id).first()
|
||||
if not step:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Step not found")
|
||||
|
||||
deleted_pos = step.position
|
||||
try:
|
||||
db.delete(step)
|
||||
# Compact remaining step positions
|
||||
db.query(PipelineStep).filter(
|
||||
PipelineStep.pipeline_id == pipeline_id,
|
||||
PipelineStep.position > deleted_pos,
|
||||
).update({"position": PipelineStep.position - 1})
|
||||
db.commit()
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to delete step id={step_id}: {exc}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to delete step",
|
||||
)
|
||||
|
||||
logger.info(f"Step deleted: id={step_id}, pipeline={pipeline_id}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper: unset default flag for an owner
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _unset_default(db: Session, owner_id: str | None) -> None:
|
||||
"""Clear the is_default flag on all pipelines for the given owner."""
|
||||
if owner_id is None:
|
||||
db.query(Pipeline).filter(Pipeline.owner_id.is_(None), Pipeline.is_default.is_(True)).update(
|
||||
{"is_default": False}
|
||||
)
|
||||
else:
|
||||
db.query(Pipeline).filter(Pipeline.owner_id == owner_id, Pipeline.is_default.is_(True)).update(
|
||||
{"is_default": False}
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Default system pipeline seeding
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# The steps that make up the standard document-processing workflow. The order
|
||||
# here mirrors what the existing Celery-based pipeline executes for every
|
||||
# uploaded file.
|
||||
_DEFAULT_PIPELINE_STEPS: list[tuple[str, str]] = [
|
||||
("convert_to_pdf", "Convert to PDF"),
|
||||
("check_duplicates", "Check for Duplicates"),
|
||||
("ocr", "OCR Processing"),
|
||||
("extract_metadata", "Extract Metadata"),
|
||||
("embed_metadata", "Embed Metadata into PDF"),
|
||||
("compute_embedding", "Compute Text Embedding"),
|
||||
("send_to_destinations", "Send to Storage Destinations"),
|
||||
]
|
||||
|
||||
#: Human-readable name shown in the management UI for the auto-seeded pipeline.
|
||||
DEFAULT_PIPELINE_NAME = "Standard Processing Pipeline"
|
||||
|
||||
|
||||
def seed_default_pipeline(db: Session) -> int:
|
||||
"""Ensure a system-owned default pipeline exists in the database.
|
||||
|
||||
This function is idempotent — it is a no-op when any system pipeline
|
||||
(``owner_id IS NULL``) already exists. It is intended to be called once
|
||||
at application startup (in ``app.main.lifespan``) so that the pipeline
|
||||
management UI always shows the default workflow that mirrors the existing
|
||||
Celery-based processing steps.
|
||||
|
||||
The created pipeline:
|
||||
|
||||
* ``owner_id = None`` — owned by the system, visible to all users
|
||||
* ``is_default = True`` — selected automatically for new documents
|
||||
* Steps (in order): convert_to_pdf → check_duplicates → ocr →
|
||||
extract_metadata → embed_metadata → compute_embedding →
|
||||
send_to_destinations
|
||||
|
||||
Args:
|
||||
db: An active SQLAlchemy session.
|
||||
|
||||
Returns:
|
||||
``1`` if a new pipeline was created, ``0`` if one already existed.
|
||||
"""
|
||||
try:
|
||||
if db.query(Pipeline).filter(Pipeline.owner_id.is_(None)).count() > 0:
|
||||
return 0
|
||||
except Exception:
|
||||
# Table may not exist yet during the very first migration run.
|
||||
return 0
|
||||
|
||||
pipeline = Pipeline(
|
||||
owner_id=None,
|
||||
name=DEFAULT_PIPELINE_NAME,
|
||||
description=(
|
||||
"The standard document processing workflow: PDF conversion, "
|
||||
"duplicate detection, OCR, metadata extraction and embedding, "
|
||||
"semantic embeddings, and final distribution to storage destinations."
|
||||
),
|
||||
is_default=True,
|
||||
is_active=True,
|
||||
)
|
||||
db.add(pipeline)
|
||||
try:
|
||||
db.flush() # Assign pipeline.id without committing yet
|
||||
except Exception as exc: # pragma: no cover
|
||||
db.rollback()
|
||||
logger.error(f"Failed to create default pipeline: {exc}")
|
||||
return 0
|
||||
|
||||
for pos, (step_type, label) in enumerate(_DEFAULT_PIPELINE_STEPS):
|
||||
db.add(
|
||||
PipelineStep(
|
||||
pipeline_id=pipeline.id,
|
||||
position=pos,
|
||||
step_type=step_type,
|
||||
label=label,
|
||||
enabled=True,
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
logger.info("Seeded default system pipeline: '%s' (id=%d)", DEFAULT_PIPELINE_NAME, pipeline.id)
|
||||
except Exception as exc: # pragma: no cover
|
||||
db.rollback()
|
||||
logger.error(f"Failed to seed default pipeline steps: {exc}")
|
||||
return 0
|
||||
|
||||
return 1
|
||||
@@ -0,0 +1,279 @@
|
||||
"""REST API for subscription plan CRUD.
|
||||
|
||||
Endpoints:
|
||||
GET /api/plans/ — list active plans (public)
|
||||
GET /api/plans/admin — list all plans inc. inactive (admin only)
|
||||
POST /api/plans/ — create plan (admin only)
|
||||
GET /api/plans/{plan_id} — get single active plan (public)
|
||||
PUT /api/plans/{plan_id} — update plan (admin only)
|
||||
DELETE /api/plans/{plan_id} — delete plan (admin only)
|
||||
POST /api/plans/seed — seed default plans (admin only)
|
||||
POST /api/plans/reorder — set sort_order for multiple plans (admin only)
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import SubscriptionPlan
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/plans", tags=["plans"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper (admin-only)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin. Raises 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class PlanUpsert(BaseModel):
|
||||
"""Body for creating or updating a subscription plan."""
|
||||
|
||||
name: str
|
||||
tagline: str | None = None
|
||||
price_monthly: float = 0.0
|
||||
price_yearly: float = 0.0
|
||||
trial_days: int = 0
|
||||
lifetime_file_limit: int = 0
|
||||
daily_upload_limit: int = 0
|
||||
monthly_upload_limit: int = 0
|
||||
max_storage_destinations: int = 0
|
||||
max_ocr_pages_monthly: int = 0
|
||||
max_file_size_mb: int = 0
|
||||
max_mailboxes: int = 0
|
||||
overage_percent: int = Field(default=20, ge=0, le=200)
|
||||
allow_overage_billing: bool = False
|
||||
overage_price_per_doc: float | None = None
|
||||
overage_price_per_ocr_page: float | None = None
|
||||
is_active: bool = True
|
||||
is_highlighted: bool = False
|
||||
badge_text: str | None = None
|
||||
cta_text: str = "Get started"
|
||||
sort_order: int = 0
|
||||
features: list[str] = []
|
||||
api_access: bool = False
|
||||
stripe_price_id_monthly: str | None = None
|
||||
stripe_price_id_yearly: str | None = None
|
||||
|
||||
|
||||
class ReorderBody(BaseModel):
|
||||
"""Body for reordering plans."""
|
||||
|
||||
order: list[str]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _plan_to_response(plan: SubscriptionPlan) -> dict[str, Any]:
|
||||
features: list[str] = []
|
||||
if plan.features:
|
||||
try:
|
||||
features = json.loads(plan.features)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
features = []
|
||||
return {
|
||||
"id": plan.id,
|
||||
"plan_id": plan.plan_id,
|
||||
"name": plan.name,
|
||||
"tagline": plan.tagline,
|
||||
"price_monthly": plan.price_monthly,
|
||||
"price_yearly": plan.price_yearly,
|
||||
"trial_days": plan.trial_days,
|
||||
"lifetime_file_limit": plan.lifetime_file_limit,
|
||||
"daily_upload_limit": plan.daily_upload_limit,
|
||||
"monthly_upload_limit": plan.monthly_upload_limit,
|
||||
"max_storage_destinations": plan.max_storage_destinations,
|
||||
"max_ocr_pages_monthly": plan.max_ocr_pages_monthly,
|
||||
"max_file_size_mb": plan.max_file_size_mb,
|
||||
"max_mailboxes": plan.max_mailboxes,
|
||||
"overage_percent": plan.overage_percent,
|
||||
"allow_overage_billing": plan.allow_overage_billing,
|
||||
"overage_price_per_doc": plan.overage_price_per_doc,
|
||||
"overage_price_per_ocr_page": plan.overage_price_per_ocr_page,
|
||||
"is_active": plan.is_active,
|
||||
"is_highlighted": plan.is_highlighted,
|
||||
"badge_text": plan.badge_text,
|
||||
"cta_text": plan.cta_text,
|
||||
"sort_order": plan.sort_order,
|
||||
"features": features,
|
||||
"api_access": plan.api_access,
|
||||
"stripe_price_id_monthly": plan.stripe_price_id_monthly,
|
||||
"stripe_price_id_yearly": plan.stripe_price_id_yearly,
|
||||
"created_at": plan.created_at.isoformat() if plan.created_at else None,
|
||||
"updated_at": plan.updated_at.isoformat() if plan.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _apply_body(plan: SubscriptionPlan, body: PlanUpsert) -> None:
|
||||
"""Apply PlanUpsert fields onto a SubscriptionPlan ORM object."""
|
||||
plan.name = body.name
|
||||
plan.tagline = body.tagline
|
||||
plan.price_monthly = body.price_monthly
|
||||
plan.price_yearly = body.price_yearly
|
||||
plan.trial_days = body.trial_days
|
||||
plan.lifetime_file_limit = body.lifetime_file_limit
|
||||
plan.daily_upload_limit = body.daily_upload_limit
|
||||
plan.monthly_upload_limit = body.monthly_upload_limit
|
||||
plan.max_storage_destinations = body.max_storage_destinations
|
||||
plan.max_ocr_pages_monthly = body.max_ocr_pages_monthly
|
||||
plan.max_file_size_mb = body.max_file_size_mb
|
||||
plan.max_mailboxes = body.max_mailboxes
|
||||
plan.overage_percent = body.overage_percent
|
||||
plan.allow_overage_billing = body.allow_overage_billing
|
||||
plan.overage_price_per_doc = body.overage_price_per_doc
|
||||
plan.overage_price_per_ocr_page = body.overage_price_per_ocr_page
|
||||
plan.is_active = body.is_active
|
||||
plan.is_highlighted = body.is_highlighted
|
||||
plan.badge_text = body.badge_text
|
||||
plan.cta_text = body.cta_text
|
||||
plan.sort_order = body.sort_order
|
||||
plan.features = json.dumps(body.features)
|
||||
plan.api_access = body.api_access
|
||||
plan.stripe_price_id_monthly = body.stripe_price_id_monthly or None
|
||||
plan.stripe_price_id_yearly = body.stripe_price_id_yearly or None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", summary="List active plans (public)")
|
||||
def list_active_plans(db: DbSession) -> dict[str, Any]:
|
||||
"""Return all active plans in sort order. Public endpoint — no auth required."""
|
||||
plans = (
|
||||
db.query(SubscriptionPlan)
|
||||
.filter(SubscriptionPlan.is_active.is_(True))
|
||||
.order_by(SubscriptionPlan.sort_order)
|
||||
.all()
|
||||
)
|
||||
return {"plans": [_plan_to_response(p) for p in plans]}
|
||||
|
||||
|
||||
@router.get("/admin", summary="List all plans including inactive (admin only)")
|
||||
def list_all_plans(db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Return all plans (active and inactive) in sort order. Admin only."""
|
||||
plans = db.query(SubscriptionPlan).order_by(SubscriptionPlan.sort_order).all()
|
||||
return {"plans": [_plan_to_response(p) for p in plans]}
|
||||
|
||||
|
||||
@router.post("/seed", summary="Seed default plans (admin only)", status_code=status.HTTP_200_OK)
|
||||
def seed_plans(db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Seed the subscription_plans table from TIER_DEFAULTS. No-op if plans already exist."""
|
||||
from app.utils.subscription import seed_default_plans
|
||||
|
||||
inserted = seed_default_plans(db)
|
||||
return {"inserted": inserted, "message": f"Seeded {inserted} default plan(s)."}
|
||||
|
||||
|
||||
@router.post("/reorder", summary="Reorder plans (admin only)")
|
||||
def reorder_plans(body: ReorderBody, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Update sort_order for each plan_id in *body.order* (position = index in list)."""
|
||||
updated = 0
|
||||
for sort_order, plan_id in enumerate(body.order):
|
||||
plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == plan_id).first()
|
||||
if plan:
|
||||
plan.sort_order = sort_order
|
||||
updated += 1
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to reorder plans")
|
||||
return {"updated": updated}
|
||||
|
||||
|
||||
@router.post("/", summary="Create a new plan (admin only)", status_code=status.HTTP_201_CREATED)
|
||||
def create_plan(plan_id: str, body: PlanUpsert, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Create a new subscription plan with the given *plan_id* slug."""
|
||||
existing = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == plan_id).first()
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"Plan '{plan_id}' already exists.",
|
||||
)
|
||||
plan = SubscriptionPlan(plan_id=plan_id)
|
||||
_apply_body(plan, body)
|
||||
db.add(plan)
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(plan)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("Admin created subscription plan '%s'", plan_id)
|
||||
return _plan_to_response(plan)
|
||||
|
||||
|
||||
@router.get("/{plan_id}", summary="Get a single active plan (public)")
|
||||
def get_plan(plan_id: str, db: DbSession) -> dict[str, Any]:
|
||||
"""Return a single active plan by plan_id. Public endpoint."""
|
||||
plan = (
|
||||
db.query(SubscriptionPlan)
|
||||
.filter(
|
||||
SubscriptionPlan.plan_id == plan_id,
|
||||
SubscriptionPlan.is_active.is_(True),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not plan:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Plan '{plan_id}' not found.")
|
||||
return _plan_to_response(plan)
|
||||
|
||||
|
||||
@router.put("/{plan_id}", summary="Update an existing plan (admin only)")
|
||||
def update_plan(plan_id: str, body: PlanUpsert, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Update an existing subscription plan. Admin only."""
|
||||
plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == plan_id).first()
|
||||
if not plan:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Plan '{plan_id}' not found.")
|
||||
_apply_body(plan, body)
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(plan)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("Admin updated subscription plan '%s'", plan_id)
|
||||
return _plan_to_response(plan)
|
||||
|
||||
|
||||
@router.delete("/{plan_id}", summary="Delete a plan (admin only)", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def delete_plan(plan_id: str, db: DbSession, _admin: AdminUser) -> None:
|
||||
"""Delete a subscription plan. Admin only."""
|
||||
plan = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id == plan_id).first()
|
||||
if not plan:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=f"Plan '{plan_id}' not found.")
|
||||
try:
|
||||
db.delete(plan)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("Admin deleted subscription plan '%s'", plan_id)
|
||||
@@ -0,0 +1,340 @@
|
||||
"""User self-service profile API.
|
||||
|
||||
Provides endpoints for the authenticated user to view and update their own
|
||||
profile settings without requiring admin access.
|
||||
|
||||
Routes:
|
||||
GET /api/profile — read current user's profile
|
||||
PATCH /api/profile — update display name, language, theme
|
||||
POST /api/profile/avatar — upload a new profile picture (JPEG/PNG/GIF/WebP, max 2 MB)
|
||||
DELETE /api/profile/avatar — remove custom avatar (reverts to Gravatar)
|
||||
POST /api/profile/change-password — change password (local-auth users only)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import logging
|
||||
from hashlib import md5
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, UploadFile, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import LocalUser, UserProfile
|
||||
from app.utils.i18n import SUPPORTED_LANGUAGE_CODES
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/profile", tags=["profile"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# Maximum avatar upload size: 2 MB
|
||||
_MAX_AVATAR_BYTES = 2 * 1024 * 1024
|
||||
|
||||
# Allowed MIME types for avatar uploads
|
||||
_ALLOWED_AVATAR_TYPES = {"image/jpeg", "image/png", "image/gif", "image/webp"}
|
||||
|
||||
# Valid theme values
|
||||
_VALID_THEMES = {"light", "dark", "system"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_user_id(request: Request) -> str:
|
||||
"""Return the stable user identifier from the session.
|
||||
|
||||
Raises HTTP 401 if no user is logged in.
|
||||
"""
|
||||
user = request.session.get("user")
|
||||
if not user or not isinstance(user, dict):
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
uid = user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id")
|
||||
if not uid:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Cannot determine user identity")
|
||||
return uid
|
||||
|
||||
|
||||
def _gravatar_url(email: str | None) -> str:
|
||||
"""Generate a Gravatar URL for *email*, falling back to identicon."""
|
||||
if not email:
|
||||
return "https://www.gravatar.com/avatar/?d=identicon"
|
||||
# MD5 used for Gravatar URL generation only — not for security
|
||||
h = md5(email.strip().lower().encode(), usedforsecurity=False).hexdigest()
|
||||
return f"https://www.gravatar.com/avatar/{h}?d=identicon"
|
||||
|
||||
|
||||
def _get_or_create_profile(db: Session, user_id: str) -> UserProfile:
|
||||
"""Return the UserProfile for *user_id*, creating a stub if one doesn't exist."""
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile is None:
|
||||
profile = UserProfile(user_id=user_id)
|
||||
db.add(profile)
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
return profile
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ProfileResponse(BaseModel):
|
||||
"""Response body for GET /api/profile."""
|
||||
|
||||
user_id: str
|
||||
display_name: str | None
|
||||
contact_email: str | None
|
||||
preferred_language: str | None
|
||||
preferred_theme: str | None
|
||||
avatar_url: str
|
||||
"""Gravatar URL or ``data:`` URI for a custom uploaded avatar."""
|
||||
is_local_user: bool
|
||||
"""True when the account was created via local email/password sign-up."""
|
||||
|
||||
|
||||
class ProfileUpdateRequest(BaseModel):
|
||||
"""Request body for PATCH /api/profile."""
|
||||
|
||||
display_name: str | None = Field(default=None, max_length=255, description="Human-readable display name")
|
||||
contact_email: str | None = Field(default=None, max_length=255, description="Contact / notification e-mail")
|
||||
preferred_language: str | None = Field(default=None, description="ISO 639-1 language code, e.g. 'en', 'de'")
|
||||
preferred_theme: str | None = Field(default=None, description="Colour scheme: 'light', 'dark', or 'system'")
|
||||
|
||||
|
||||
class ChangePasswordRequest(BaseModel):
|
||||
"""Request body for POST /api/profile/change-password."""
|
||||
|
||||
current_password: str = Field(..., min_length=1, max_length=128)
|
||||
new_password: str = Field(..., min_length=8, max_length=128)
|
||||
new_password_confirm: str = Field(..., min_length=8, max_length=128)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("", response_model=ProfileResponse)
|
||||
@require_login
|
||||
async def get_profile(request: Request, db: DbSession) -> ProfileResponse:
|
||||
"""Return the current user's profile settings."""
|
||||
user_id = _get_user_id(request)
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
|
||||
session_user = request.session.get("user", {})
|
||||
email = session_user.get("email") if isinstance(session_user, dict) else None
|
||||
|
||||
# Determine avatar: prefer stored data, fall back to Gravatar
|
||||
avatar_url = profile.avatar_data if profile.avatar_data else _gravatar_url(email) # type: ignore[attr-defined]
|
||||
|
||||
# Check whether this is a local (email/password) account
|
||||
is_local = db.query(LocalUser).filter(LocalUser.username == user_id).first() is not None
|
||||
|
||||
return ProfileResponse(
|
||||
user_id=user_id,
|
||||
display_name=profile.display_name, # type: ignore[arg-type]
|
||||
contact_email=profile.contact_email, # type: ignore[arg-type]
|
||||
preferred_language=profile.preferred_language, # type: ignore[arg-type]
|
||||
preferred_theme=profile.preferred_theme, # type: ignore[arg-type]
|
||||
avatar_url=avatar_url,
|
||||
is_local_user=is_local,
|
||||
)
|
||||
|
||||
|
||||
@router.patch("", response_model=ProfileResponse)
|
||||
@require_login
|
||||
async def update_profile(
|
||||
body: ProfileUpdateRequest, request: Request, response: Response, db: DbSession
|
||||
) -> ProfileResponse:
|
||||
"""Update the current user's editable profile settings."""
|
||||
user_id = _get_user_id(request)
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
|
||||
# Validate language code
|
||||
if body.preferred_language is not None:
|
||||
lang = body.preferred_language.lower().strip()
|
||||
if lang and lang not in SUPPORTED_LANGUAGE_CODES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Unsupported language code: {lang}",
|
||||
)
|
||||
profile.preferred_language = lang or None # type: ignore[assignment]
|
||||
|
||||
# Keep session and cookie in sync so detect_language() picks up
|
||||
# the new preference immediately (without a DB round-trip).
|
||||
if hasattr(request, "session"):
|
||||
if lang:
|
||||
request.session["preferred_language"] = lang
|
||||
else:
|
||||
request.session.pop("preferred_language", None)
|
||||
if lang:
|
||||
response.set_cookie(
|
||||
key="docuelevate_lang",
|
||||
value=lang,
|
||||
max_age=30 * 24 * 60 * 60,
|
||||
httponly=False,
|
||||
samesite="lax",
|
||||
)
|
||||
else:
|
||||
response.delete_cookie(key="docuelevate_lang")
|
||||
|
||||
# Validate theme
|
||||
if body.preferred_theme is not None:
|
||||
theme = body.preferred_theme.lower().strip()
|
||||
if theme and theme not in _VALID_THEMES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Invalid theme: {theme}. Must be one of: {', '.join(sorted(_VALID_THEMES))}",
|
||||
)
|
||||
profile.preferred_theme = theme or None # type: ignore[assignment]
|
||||
|
||||
if body.display_name is not None:
|
||||
profile.display_name = body.display_name.strip() or None # type: ignore[assignment]
|
||||
|
||||
if body.contact_email is not None:
|
||||
profile.contact_email = body.contact_email.strip() or None # type: ignore[assignment]
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(profile)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
session_user = request.session.get("user", {})
|
||||
email = session_user.get("email") if isinstance(session_user, dict) else None
|
||||
avatar_url = profile.avatar_data if profile.avatar_data else _gravatar_url(email) # type: ignore[attr-defined]
|
||||
is_local = db.query(LocalUser).filter(LocalUser.username == user_id).first() is not None
|
||||
|
||||
return ProfileResponse(
|
||||
user_id=user_id,
|
||||
display_name=profile.display_name, # type: ignore[arg-type]
|
||||
contact_email=profile.contact_email, # type: ignore[arg-type]
|
||||
preferred_language=profile.preferred_language, # type: ignore[arg-type]
|
||||
preferred_theme=profile.preferred_theme, # type: ignore[arg-type]
|
||||
avatar_url=avatar_url,
|
||||
is_local_user=is_local,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/avatar", status_code=status.HTTP_200_OK)
|
||||
@require_login
|
||||
async def upload_avatar(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
file: UploadFile = File(..., description="Profile picture (JPEG, PNG, GIF or WebP; max 2 MB)"),
|
||||
) -> dict:
|
||||
"""Upload a new profile picture.
|
||||
|
||||
The image is stored as a base64-encoded data URL in ``UserProfile.avatar_data``.
|
||||
Accepts JPEG, PNG, GIF, or WebP files up to 2 MB.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
content_type = (file.content_type or "").lower()
|
||||
if content_type not in _ALLOWED_AVATAR_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
|
||||
detail=f"Unsupported image type '{content_type}'. Allowed: JPEG, PNG, GIF, WebP.",
|
||||
)
|
||||
|
||||
# Check declared size first (available when the client sends a Content-Length header)
|
||||
if file.size is not None and file.size > _MAX_AVATAR_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail="Avatar image must be 2 MB or smaller.",
|
||||
)
|
||||
|
||||
# Read up to one byte past the limit so we can detect oversized uploads
|
||||
raw = await file.read(_MAX_AVATAR_BYTES + 1)
|
||||
if len(raw) > _MAX_AVATAR_BYTES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
|
||||
detail="Avatar image must be 2 MB or smaller.",
|
||||
)
|
||||
|
||||
b64 = base64.b64encode(raw).decode("ascii")
|
||||
data_url = f"data:{content_type};base64,{b64}"
|
||||
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
profile.avatar_data = data_url # type: ignore[assignment]
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
return {"avatar_url": data_url}
|
||||
|
||||
|
||||
@router.delete("/avatar", status_code=status.HTTP_200_OK)
|
||||
@require_login
|
||||
async def delete_avatar(request: Request, db: DbSession) -> dict:
|
||||
"""Remove the custom avatar and revert to the Gravatar fallback."""
|
||||
user_id = _get_user_id(request)
|
||||
profile = _get_or_create_profile(db, user_id)
|
||||
profile.avatar_data = None # type: ignore[assignment]
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
session_user = request.session.get("user", {})
|
||||
email = session_user.get("email") if isinstance(session_user, dict) else None
|
||||
return {"avatar_url": _gravatar_url(email)}
|
||||
|
||||
|
||||
@router.post("/change-password", status_code=status.HTTP_200_OK)
|
||||
@require_login
|
||||
async def change_password(body: ChangePasswordRequest, request: Request, db: DbSession) -> dict:
|
||||
"""Change the password for local (email/password) accounts.
|
||||
|
||||
Raises 403 if the account is not a local account or the current password is wrong.
|
||||
Raises 422 if the new passwords do not match.
|
||||
"""
|
||||
from app.utils.local_auth import hash_password, verify_password
|
||||
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
local_user = db.query(LocalUser).filter(LocalUser.username == user_id).first()
|
||||
if local_user is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Password change is only available for local accounts.",
|
||||
)
|
||||
|
||||
if not verify_password(body.current_password, local_user.hashed_password):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Current password is incorrect.",
|
||||
)
|
||||
|
||||
if body.new_password != body.new_password_confirm:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="New passwords do not match.",
|
||||
)
|
||||
|
||||
local_user.hashed_password = hash_password(body.new_password)
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Password changed for local user: %s", user_id)
|
||||
return {"detail": "Password changed successfully."}
|
||||
@@ -0,0 +1,468 @@
|
||||
"""Routing rules API endpoints.
|
||||
|
||||
Provides full CRUD for pipeline routing rules that conditionally assign
|
||||
documents to pipelines based on document properties (file type, category,
|
||||
metadata fields, size, etc.).
|
||||
|
||||
Rules are evaluated in ascending ``position`` order. The first rule whose
|
||||
condition matches wins and routes the document to the specified target
|
||||
pipeline. If no rule matches, the caller falls back to the owner's (or
|
||||
system) default pipeline.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import Pipeline, PipelineRoutingRule
|
||||
from app.utils.routing_engine import (
|
||||
BUILTIN_FIELDS,
|
||||
VALID_OPERATORS,
|
||||
_evaluate_condition,
|
||||
_resolve_field,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/routing-rules", tags=["routing-rules"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
MAX_RULES_PER_OWNER = 100
|
||||
MAX_NAME_LENGTH = 255
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_user_id(request: Request) -> str:
|
||||
"""Return the authenticated user identifier."""
|
||||
user = getattr(request.state, "user", None)
|
||||
if user:
|
||||
if isinstance(user, dict):
|
||||
return user.get("sub", user.get("email", "anonymous"))
|
||||
return getattr(user, "sub", getattr(user, "email", "anonymous"))
|
||||
return "anonymous"
|
||||
|
||||
|
||||
def _is_admin(request: Request) -> bool:
|
||||
"""Return ``True`` when the current user has admin privileges."""
|
||||
user = getattr(request.state, "user", None)
|
||||
if not user:
|
||||
return False
|
||||
groups = user.get("groups", []) if isinstance(user, dict) else getattr(user, "groups", [])
|
||||
return "admin" in groups
|
||||
|
||||
|
||||
def _can_access_rule(rule: PipelineRoutingRule, user_id: str, admin: bool) -> bool:
|
||||
"""Check whether the user is allowed to read this rule."""
|
||||
if admin:
|
||||
return True
|
||||
return rule.owner_id == user_id
|
||||
|
||||
|
||||
def _can_write_rule(rule: PipelineRoutingRule, user_id: str, admin: bool) -> bool:
|
||||
"""Check whether the user is allowed to modify this rule."""
|
||||
if rule.owner_id is None:
|
||||
return admin
|
||||
return rule.owner_id == user_id
|
||||
|
||||
|
||||
def _validate_field(field: str) -> None:
|
||||
"""Raise 422 if the field name is invalid."""
|
||||
if field in BUILTIN_FIELDS:
|
||||
return
|
||||
if field.startswith("metadata.") and len(field) > len("metadata."):
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=(
|
||||
f"Invalid field '{field}'. "
|
||||
f"Valid built-in fields: {sorted(BUILTIN_FIELDS)}. "
|
||||
"For AI metadata, use 'metadata.<key>'."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _validate_operator(operator: str) -> None:
|
||||
"""Raise 422 if the operator is not recognised."""
|
||||
if operator not in VALID_OPERATORS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Invalid operator '{operator}'. Valid operators: {sorted(VALID_OPERATORS)}",
|
||||
)
|
||||
|
||||
|
||||
def _serialize_rule(rule: PipelineRoutingRule) -> dict[str, Any]:
|
||||
"""Serialize a routing rule to a JSON-compatible dict."""
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"position": rule.position,
|
||||
"field": rule.field,
|
||||
"operator": rule.operator,
|
||||
"value": rule.value,
|
||||
"target_pipeline_id": rule.target_pipeline_id,
|
||||
"is_active": rule.is_active,
|
||||
"created_at": rule.created_at.isoformat() if rule.created_at else None,
|
||||
"updated_at": rule.updated_at.isoformat() if rule.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic request models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RoutingRuleCreate(BaseModel):
|
||||
"""Request body for creating a routing rule."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=MAX_NAME_LENGTH)
|
||||
field: str = Field(..., min_length=1, max_length=255)
|
||||
operator: str = Field(..., min_length=1, max_length=50)
|
||||
value: str = Field(..., max_length=1024)
|
||||
target_pipeline_id: int
|
||||
position: int | None = None
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
class RoutingRuleUpdate(BaseModel):
|
||||
"""Request body for updating a routing rule."""
|
||||
|
||||
name: str | None = Field(None, min_length=1, max_length=MAX_NAME_LENGTH)
|
||||
field: str | None = Field(None, min_length=1, max_length=255)
|
||||
operator: str | None = Field(None, min_length=1, max_length=50)
|
||||
value: str | None = Field(None, max_length=1024)
|
||||
target_pipeline_id: int | None = None
|
||||
position: int | None = None
|
||||
is_active: bool | None = None
|
||||
|
||||
|
||||
class RoutingRuleEvaluateRequest(BaseModel):
|
||||
"""Request body for dry-run rule evaluation."""
|
||||
|
||||
file_type: str | None = None
|
||||
filename: str | None = None
|
||||
size: int | None = None
|
||||
document_type: str | None = None
|
||||
metadata: dict[str, Any] | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("")
|
||||
@require_login
|
||||
def list_routing_rules(request: Request, db: DbSession) -> list[dict[str, Any]]:
|
||||
"""List all routing rules accessible by the current user.
|
||||
|
||||
Returns the user's own rules plus any system-wide rules (``owner_id=NULL``).
|
||||
Rules are sorted by position.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
rules = (
|
||||
db.query(PipelineRoutingRule)
|
||||
.filter((PipelineRoutingRule.owner_id == user_id) | (PipelineRoutingRule.owner_id.is_(None)))
|
||||
.order_by(
|
||||
PipelineRoutingRule.owner_id.is_(None).asc(),
|
||||
PipelineRoutingRule.position.asc(),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
return [_serialize_rule(r) for r in rules]
|
||||
|
||||
|
||||
@router.post("", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
def create_routing_rule(request: Request, db: DbSession, body: RoutingRuleCreate) -> dict[str, Any]:
|
||||
"""Create a new routing rule for the current user.
|
||||
|
||||
Returns:
|
||||
The created routing rule.
|
||||
|
||||
Raises:
|
||||
HTTPException 422: If the field or operator is invalid.
|
||||
HTTPException 404: If the target pipeline does not exist.
|
||||
HTTPException 409: If the maximum number of rules is reached.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
_validate_field(body.field)
|
||||
_validate_operator(body.operator)
|
||||
|
||||
# Verify target pipeline exists and is accessible.
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == body.target_pipeline_id).first()
|
||||
if not pipeline:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Target pipeline {body.target_pipeline_id} not found",
|
||||
)
|
||||
|
||||
# Enforce per-owner limit.
|
||||
count = db.query(PipelineRoutingRule).filter(PipelineRoutingRule.owner_id == user_id).count()
|
||||
if count >= MAX_RULES_PER_OWNER:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"Maximum of {MAX_RULES_PER_OWNER} routing rules per user reached",
|
||||
)
|
||||
|
||||
# Auto-assign position if not specified.
|
||||
position = body.position
|
||||
if position is None:
|
||||
max_pos = (
|
||||
db.query(PipelineRoutingRule.position)
|
||||
.filter(PipelineRoutingRule.owner_id == user_id)
|
||||
.order_by(PipelineRoutingRule.position.desc())
|
||||
.first()
|
||||
)
|
||||
position = (max_pos[0] + 1) if max_pos else 0
|
||||
|
||||
rule = PipelineRoutingRule(
|
||||
owner_id=user_id,
|
||||
name=body.name.strip(),
|
||||
position=position,
|
||||
field=body.field,
|
||||
operator=body.operator,
|
||||
value=body.value,
|
||||
target_pipeline_id=body.target_pipeline_id,
|
||||
is_active=body.is_active,
|
||||
)
|
||||
|
||||
try:
|
||||
db.add(rule)
|
||||
db.commit()
|
||||
db.refresh(rule)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create routing rule for user=%s", user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to create routing rule",
|
||||
)
|
||||
|
||||
logger.info("Routing rule created: id=%s, user=%s", rule.id, user_id)
|
||||
return _serialize_rule(rule)
|
||||
|
||||
|
||||
@router.get("/operators")
|
||||
def list_operators() -> dict[str, Any]:
|
||||
"""Return the list of supported operators and fields.
|
||||
|
||||
This is a public endpoint (no auth required) so that UIs can populate
|
||||
dropdowns without hard-coding the catalogue.
|
||||
"""
|
||||
return {
|
||||
"operators": sorted(VALID_OPERATORS),
|
||||
"builtin_fields": sorted(BUILTIN_FIELDS),
|
||||
"metadata_prefix": "metadata.",
|
||||
}
|
||||
|
||||
|
||||
@router.post("/evaluate")
|
||||
@require_login
|
||||
def evaluate_rules(request: Request, db: DbSession, body: RoutingRuleEvaluateRequest) -> dict[str, Any]:
|
||||
"""Dry-run rule evaluation against the provided document properties.
|
||||
|
||||
Returns the first matching rule and target pipeline (if any), or
|
||||
indicates that no rule matched (default pipeline will be used).
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
doc_props: dict[str, Any] = {
|
||||
"file_type": body.file_type,
|
||||
"filename": body.filename,
|
||||
"size": body.size,
|
||||
"document_type": body.document_type,
|
||||
"metadata": body.metadata or {},
|
||||
}
|
||||
|
||||
rules = (
|
||||
db.query(PipelineRoutingRule)
|
||||
.filter(
|
||||
PipelineRoutingRule.is_active.is_(True),
|
||||
(PipelineRoutingRule.owner_id == user_id) | (PipelineRoutingRule.owner_id.is_(None)),
|
||||
)
|
||||
.order_by(
|
||||
PipelineRoutingRule.owner_id.is_(None).asc(),
|
||||
PipelineRoutingRule.position.asc(),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
for rule in rules:
|
||||
actual = _resolve_field(rule.field, doc_props)
|
||||
if _evaluate_condition(actual, rule.operator, rule.value):
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == rule.target_pipeline_id).first()
|
||||
return {
|
||||
"matched": True,
|
||||
"rule": _serialize_rule(rule),
|
||||
"target_pipeline": {
|
||||
"id": pipeline.id,
|
||||
"name": pipeline.name,
|
||||
"is_active": pipeline.is_active,
|
||||
}
|
||||
if pipeline
|
||||
else None,
|
||||
}
|
||||
|
||||
return {"matched": False, "rule": None, "target_pipeline": None}
|
||||
|
||||
|
||||
@router.put("/reorder")
|
||||
@require_login
|
||||
def reorder_routing_rules(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
rule_ids: list[int] = Body(..., embed=True),
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Reorder the caller's routing rules.
|
||||
|
||||
Expects a JSON body ``{"rule_ids": [3, 1, 2]}`` where the list
|
||||
contains the IDs of the caller's rules in the desired order.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
rules = (
|
||||
db.query(PipelineRoutingRule)
|
||||
.filter(PipelineRoutingRule.owner_id == user_id, PipelineRoutingRule.id.in_(rule_ids))
|
||||
.all()
|
||||
)
|
||||
|
||||
rule_map = {r.id: r for r in rules}
|
||||
|
||||
if len(rule_map) != len(rule_ids) or set(rule_map.keys()) != set(rule_ids):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="rule_ids must contain exactly the IDs of your routing rules",
|
||||
)
|
||||
|
||||
for pos, rid in enumerate(rule_ids):
|
||||
rule_map[rid].position = pos
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to reorder routing rules for user=%s", user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to reorder routing rules",
|
||||
)
|
||||
|
||||
ordered = sorted(rules, key=lambda r: r.position)
|
||||
return [_serialize_rule(r) for r in ordered]
|
||||
|
||||
|
||||
@router.get("/{rule_id}")
|
||||
@require_login
|
||||
def get_routing_rule(rule_id: int, request: Request, db: DbSession) -> dict[str, Any]:
|
||||
"""Return a single routing rule by ID."""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
rule = db.query(PipelineRoutingRule).filter(PipelineRoutingRule.id == rule_id).first()
|
||||
if not rule or not _can_access_rule(rule, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Routing rule not found")
|
||||
|
||||
return _serialize_rule(rule)
|
||||
|
||||
|
||||
@router.put("/{rule_id}")
|
||||
@require_login
|
||||
def update_routing_rule(rule_id: int, request: Request, db: DbSession, body: RoutingRuleUpdate) -> dict[str, Any]:
|
||||
"""Update a routing rule.
|
||||
|
||||
Only the fields present in the request body are updated.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
rule = db.query(PipelineRoutingRule).filter(PipelineRoutingRule.id == rule_id).first()
|
||||
if not rule or not _can_access_rule(rule, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Routing rule not found")
|
||||
|
||||
if not _can_write_rule(rule, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this rule")
|
||||
|
||||
if body.field is not None:
|
||||
_validate_field(body.field)
|
||||
rule.field = body.field
|
||||
|
||||
if body.operator is not None:
|
||||
_validate_operator(body.operator)
|
||||
rule.operator = body.operator
|
||||
|
||||
if body.value is not None:
|
||||
rule.value = body.value
|
||||
|
||||
if body.name is not None:
|
||||
rule.name = body.name.strip()
|
||||
|
||||
if body.target_pipeline_id is not None:
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == body.target_pipeline_id).first()
|
||||
if not pipeline:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Target pipeline {body.target_pipeline_id} not found",
|
||||
)
|
||||
rule.target_pipeline_id = body.target_pipeline_id
|
||||
|
||||
if body.position is not None:
|
||||
rule.position = body.position
|
||||
|
||||
if body.is_active is not None:
|
||||
rule.is_active = body.is_active
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(rule)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to update routing rule id=%s", rule_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update routing rule",
|
||||
)
|
||||
|
||||
logger.info("Routing rule updated: id=%s, user=%s", rule_id, user_id)
|
||||
return _serialize_rule(rule)
|
||||
|
||||
|
||||
@router.delete("/{rule_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
def delete_routing_rule(rule_id: int, request: Request, db: DbSession) -> None:
|
||||
"""Delete a routing rule."""
|
||||
user_id = _get_user_id(request)
|
||||
admin = _is_admin(request)
|
||||
|
||||
rule = db.query(PipelineRoutingRule).filter(PipelineRoutingRule.id == rule_id).first()
|
||||
if not rule or not _can_access_rule(rule, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Routing rule not found")
|
||||
|
||||
if not _can_write_rule(rule, user_id, admin):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this rule")
|
||||
|
||||
try:
|
||||
db.delete(rule)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to delete routing rule id=%s", rule_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to delete routing rule",
|
||||
)
|
||||
|
||||
logger.info("Routing rule deleted: id=%s, user=%s", rule_id, user_id)
|
||||
@@ -190,15 +190,15 @@ def create_saved_search(
|
||||
db.add(saved_search)
|
||||
db.commit()
|
||||
db.refresh(saved_search)
|
||||
except Exception as exc:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to create saved search for user={user_id}: {exc}")
|
||||
logger.exception("Failed to create saved search for user=%s", user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to save search",
|
||||
)
|
||||
|
||||
logger.info(f"Saved search created: user={user_id}, name={name!r}")
|
||||
logger.info("Saved search created: user=%s, name=%r", user_id, name)
|
||||
return _serialize_saved_search(saved_search)
|
||||
|
||||
|
||||
@@ -263,15 +263,15 @@ def update_saved_search(
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(saved_search)
|
||||
except Exception as exc:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to update saved search id={search_id}, user={user_id}: {exc}")
|
||||
logger.exception("Failed to update saved search id=%s, user=%s", search_id, user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update saved search",
|
||||
)
|
||||
|
||||
logger.info(f"Saved search updated: id={search_id}, user={user_id}")
|
||||
logger.info("Saved search updated: id=%s, user=%s", search_id, user_id)
|
||||
return _serialize_saved_search(saved_search)
|
||||
|
||||
|
||||
@@ -294,12 +294,12 @@ def delete_saved_search(search_id: int, request: Request, db: DbSession):
|
||||
try:
|
||||
db.delete(saved_search)
|
||||
db.commit()
|
||||
except Exception as exc:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to delete saved search id={search_id}, user={user_id}: {exc}")
|
||||
logger.exception("Failed to delete saved search id=%s, user=%s", search_id, user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to delete saved search",
|
||||
)
|
||||
|
||||
logger.info(f"Saved search deleted: id={search_id}, user={user_id}")
|
||||
logger.info("Saved search deleted: id=%s, user=%s", search_id, user_id)
|
||||
|
||||
@@ -0,0 +1,350 @@
|
||||
"""
|
||||
Admin API endpoints for managing scheduled batch processing jobs.
|
||||
|
||||
All endpoints require admin privileges (checked via session ``is_admin`` flag).
|
||||
|
||||
Available routes:
|
||||
GET /api/admin/scheduled-jobs – list all scheduled jobs
|
||||
PATCH /api/admin/scheduled-jobs/{id} – update schedule / enable-disable
|
||||
POST /api/admin/scheduled-jobs/{id}/run-now – trigger a job immediately
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import ScheduledJob
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/admin/scheduled-jobs", tags=["admin-scheduled-jobs"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Authorisation helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin; raises HTTP 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ScheduledJobResponse(BaseModel):
|
||||
"""Read model for a scheduled job."""
|
||||
|
||||
id: int
|
||||
name: str
|
||||
display_name: str
|
||||
description: str | None
|
||||
task_name: str
|
||||
enabled: bool
|
||||
schedule_type: str
|
||||
cron_minute: str
|
||||
cron_hour: str
|
||||
cron_day_of_week: str
|
||||
cron_day_of_month: str
|
||||
cron_month_of_year: str
|
||||
interval_seconds: int | None
|
||||
last_run_at: datetime | None
|
||||
last_run_status: str | None
|
||||
last_run_detail: str | None
|
||||
created_at: datetime | None
|
||||
updated_at: datetime | None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ScheduledJobUpdate(BaseModel):
|
||||
"""Writable fields for a scheduled job update (all optional)."""
|
||||
|
||||
enabled: bool | None = Field(None, description="Whether the job is active")
|
||||
schedule_type: str | None = Field(None, pattern="^(cron|interval)$", description="'cron' or 'interval'")
|
||||
cron_minute: str | None = Field(None, max_length=50)
|
||||
cron_hour: str | None = Field(None, max_length=50)
|
||||
cron_day_of_week: str | None = Field(None, max_length=50)
|
||||
cron_day_of_month: str | None = Field(None, max_length=50)
|
||||
cron_month_of_year: str | None = Field(None, max_length=50)
|
||||
interval_seconds: int | None = Field(None, ge=60, description="Interval in seconds (min 60)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Default job definitions – seeded into the DB on first startup
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
DEFAULT_JOBS: list[dict[str, Any]] = [
|
||||
{
|
||||
"name": "process-new-documents",
|
||||
"display_name": "Process New Documents",
|
||||
"description": (
|
||||
"Scans for documents that have been uploaded but never processed "
|
||||
"and queues them through the full processing pipeline. "
|
||||
"Runs hourly by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.process_new_documents",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "0",
|
||||
"cron_hour": "*/1",
|
||||
"cron_day_of_week": "*",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "reprocess-failed-documents",
|
||||
"display_name": "Reprocess Failed Documents",
|
||||
"description": (
|
||||
"Finds documents whose last processing attempt failed and re-queues "
|
||||
"them for reprocessing. Only picks up files that are not currently "
|
||||
"being processed. Runs every 6 hours by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.reprocess_failed_documents",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "30",
|
||||
"cron_hour": "*/6",
|
||||
"cron_day_of_week": "*",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "cleanup-temp-files",
|
||||
"display_name": "Clean Up Temporary Files",
|
||||
"description": (
|
||||
"Removes stale files from the workdir/tmp directory. "
|
||||
"Only files older than 24 hours that are not referenced by any active "
|
||||
"processing job are deleted. Runs daily at 03:30 UTC by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.cleanup_temp_files",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "30",
|
||||
"cron_hour": "3",
|
||||
"cron_day_of_week": "*",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "expire-shared-links",
|
||||
"display_name": "Expire Stale Shared Links",
|
||||
"description": (
|
||||
"Marks shared document links as inactive when their expiry time has passed. "
|
||||
"Access is already blocked at request time, but this task keeps the "
|
||||
"management UI counts accurate. Runs daily at 01:00 UTC by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.expire_shared_links",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "0",
|
||||
"cron_hour": "1",
|
||||
"cron_day_of_week": "*",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "prune-processing-logs",
|
||||
"display_name": "Prune Old Processing Logs",
|
||||
"description": (
|
||||
"Deletes processing log entries and settings audit log entries older than "
|
||||
"30 days to prevent unbounded database growth. "
|
||||
"Runs weekly on Sunday at 04:00 UTC by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.prune_processing_logs",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "0",
|
||||
"cron_hour": "4",
|
||||
"cron_day_of_week": "0",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "prune-old-notifications",
|
||||
"display_name": "Prune Old Notifications",
|
||||
"description": (
|
||||
"Deletes read in-app notifications older than 30 days. "
|
||||
"Unread notifications are never deleted. "
|
||||
"Runs weekly on Sunday at 04:30 UTC by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.prune_old_notifications",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "30",
|
||||
"cron_hour": "4",
|
||||
"cron_day_of_week": "0",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "backfill-missing-metadata",
|
||||
"display_name": "Backfill Missing AI Metadata",
|
||||
"description": (
|
||||
"Re-triggers AI metadata extraction for documents that have extracted "
|
||||
"text but no AI metadata yet (e.g., processed before an AI provider "
|
||||
"was configured). Processes up to 50 documents per run. "
|
||||
"Runs every 6 hours by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.backfill_missing_metadata",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "0",
|
||||
"cron_hour": "*/6",
|
||||
"cron_day_of_week": "*",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
{
|
||||
"name": "sync-search-index",
|
||||
"display_name": "Sync Search Index",
|
||||
"description": (
|
||||
"Indexes documents that have OCR text or AI metadata but are missing "
|
||||
"from the Meilisearch search index. Useful after enabling search on "
|
||||
"an existing installation or after an index rebuild. "
|
||||
"Processes up to 100 documents per run. "
|
||||
"Runs hourly by default."
|
||||
),
|
||||
"task_name": "app.tasks.batch_tasks.sync_search_index",
|
||||
"enabled": True,
|
||||
"schedule_type": "cron",
|
||||
"cron_minute": "15",
|
||||
"cron_hour": "*/1",
|
||||
"cron_day_of_week": "*",
|
||||
"cron_day_of_month": "*",
|
||||
"cron_month_of_year": "*",
|
||||
"interval_seconds": None,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def seed_default_scheduled_jobs(db: Session) -> None:
|
||||
"""
|
||||
Insert the built-in scheduled jobs if they do not already exist.
|
||||
|
||||
Called from the FastAPI lifespan handler so the records are available
|
||||
immediately after the first startup.
|
||||
"""
|
||||
for job_data in DEFAULT_JOBS:
|
||||
existing = db.query(ScheduledJob).filter(ScheduledJob.name == job_data["name"]).first()
|
||||
if existing is None:
|
||||
db.add(ScheduledJob(**job_data))
|
||||
try:
|
||||
db.commit()
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.error("Failed to seed default scheduled jobs: %s", exc)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("", response_model=list[ScheduledJobResponse])
|
||||
def list_scheduled_jobs(request: Request, db: DbSession, _admin: AdminUser) -> list[ScheduledJobResponse]:
|
||||
"""
|
||||
Return all scheduled jobs ordered by display name.
|
||||
|
||||
Requires admin privileges.
|
||||
"""
|
||||
jobs = db.query(ScheduledJob).order_by(ScheduledJob.display_name).all()
|
||||
return jobs # type: ignore[return-value]
|
||||
|
||||
|
||||
@router.patch("/{job_id}", response_model=ScheduledJobResponse)
|
||||
def update_scheduled_job(
|
||||
job_id: int,
|
||||
payload: ScheduledJobUpdate,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
_admin: AdminUser,
|
||||
) -> ScheduledJobResponse:
|
||||
"""
|
||||
Update schedule configuration or enabled state for a job.
|
||||
|
||||
Only the fields included in the request body are modified.
|
||||
Changes to the Celery Beat schedule take effect after the worker restarts.
|
||||
|
||||
Requires admin privileges.
|
||||
"""
|
||||
job = db.query(ScheduledJob).filter(ScheduledJob.id == job_id).first()
|
||||
if job is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Scheduled job not found")
|
||||
|
||||
update_data = payload.model_dump(exclude_none=True)
|
||||
if not update_data:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="No fields to update")
|
||||
|
||||
for field, value in update_data.items():
|
||||
setattr(job, field, value)
|
||||
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(job)
|
||||
except Exception as exc:
|
||||
db.rollback()
|
||||
logger.error("Failed to update scheduled job %s: %s", job_id, exc)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update scheduled job",
|
||||
) from exc
|
||||
|
||||
logger.info("Admin updated scheduled job %s (id=%s): %s", job.name, job_id, update_data)
|
||||
return job # type: ignore[return-value]
|
||||
|
||||
|
||||
@router.post("/{job_id}/run-now")
|
||||
def run_scheduled_job_now(
|
||||
job_id: int,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
_admin: AdminUser,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Immediately dispatch the Celery task for the given scheduled job.
|
||||
|
||||
The task is sent to the default queue; its result is tracked asynchronously
|
||||
via the ``last_run_at`` / ``last_run_status`` fields updated by the task
|
||||
itself.
|
||||
|
||||
Requires admin privileges.
|
||||
"""
|
||||
job = db.query(ScheduledJob).filter(ScheduledJob.id == job_id).first()
|
||||
if job is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Scheduled job not found")
|
||||
|
||||
from app.celery_app import celery as celery_app
|
||||
|
||||
task = celery_app.send_task(job.task_name)
|
||||
logger.info("Admin triggered scheduled job %s (id=%s) manually, task_id=%s", job.name, job_id, task.id)
|
||||
|
||||
return {
|
||||
"status": "dispatched",
|
||||
"job_id": job_id,
|
||||
"job_name": job.name,
|
||||
"task_id": task.id,
|
||||
}
|
||||
@@ -446,6 +446,42 @@ async def install_ocr_languages(request: Request, admin: AdminUser):
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{key}/suggestions")
|
||||
async def get_setting_suggestions(
|
||||
key: str,
|
||||
request: Request,
|
||||
q: str = "",
|
||||
limit: int = 10,
|
||||
):
|
||||
"""
|
||||
Return autocomplete suggestions for a setting key.
|
||||
|
||||
Fetches values dynamically from cloud SDKs, installed tools, or
|
||||
curated static lists depending on the setting. Results are filtered
|
||||
by case-insensitive substring match on the ``q`` parameter.
|
||||
|
||||
This endpoint does **not** require admin privileges so that the
|
||||
autocomplete widget works for any authenticated user viewing settings.
|
||||
"""
|
||||
from app.utils.suggestion_providers import SUGGESTION_PROVIDERS, get_suggestions # noqa: PLC0415
|
||||
|
||||
if key not in SUGGESTION_PROVIDERS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"No suggestions available for setting '{key}'",
|
||||
)
|
||||
|
||||
try:
|
||||
suggestions = get_suggestions(key, query=q, limit=max(1, min(limit, 50)))
|
||||
return {"key": key, "suggestions": suggestions}
|
||||
except Exception as e:
|
||||
logger.error(f"Error fetching suggestions for {key}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to fetch suggestions",
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{key}/history")
|
||||
async def get_key_history(key: str, request: Request, db: DbSession, admin: AdminUser):
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,500 @@
|
||||
"""API endpoints for document sharing via expiring links.
|
||||
|
||||
Authenticated users can create time-limited or view-limited shareable
|
||||
links for their documents. Each link has a cryptographically random
|
||||
token that forms a public ``/share/<token>`` URL. Optional password
|
||||
protection is supported; only a PBKDF2-HMAC-SHA256 hash is stored.
|
||||
|
||||
Public consumers access files through the ``/share/<token>/download``
|
||||
and ``/share/<token>/info`` endpoints — no authentication required.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from fastapi.responses import FileResponse
|
||||
from pydantic import BaseModel, Field, field_validator
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import FileRecord, SharedLink
|
||||
from app.utils.user_scope import apply_owner_filter, get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/shared-links", tags=["shared-links"])
|
||||
public_router = APIRouter(tags=["shared-links-public"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: PBKDF2 iteration count — matches OWASP 2023 recommendation for PBKDF2-HMAC-SHA256.
|
||||
_PWD_HASH_ITERATIONS = 600_000
|
||||
#: Length of the random per-password salt in bytes (128-bit entropy).
|
||||
_PWD_SALT_BYTES = 16
|
||||
|
||||
# Valid expiry durations (in hours) presented in the UI.
|
||||
EXPIRY_OPTIONS: dict[str, int] = {
|
||||
"1h": 1,
|
||||
"6h": 6,
|
||||
"12h": 12,
|
||||
"24h": 24,
|
||||
"3d": 72,
|
||||
"7d": 168,
|
||||
"14d": 336,
|
||||
"30d": 720,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Return the current user's owner ID, raising 401 if unauthenticated."""
|
||||
owner_id = get_current_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Not authenticated")
|
||||
return owner_id
|
||||
|
||||
|
||||
CurrentOwner = Annotated[str, Depends(_get_owner_id)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _generate_token() -> str:
|
||||
"""Generate a 43-character URL-safe random token."""
|
||||
return secrets.token_urlsafe(32)
|
||||
|
||||
|
||||
def _hash_password(password: str) -> str:
|
||||
"""Hash *password* with PBKDF2-HMAC-SHA256 and a random per-password salt.
|
||||
|
||||
The returned string uses the format ``{salt_hex}:{dk_hex}`` so that
|
||||
both the salt and the digest can be recovered from a single column.
|
||||
|
||||
Args:
|
||||
password: Plaintext password string.
|
||||
|
||||
Returns:
|
||||
String in the form ``<32-char salt hex>:<64-char digest hex>``,
|
||||
totalling 97 characters (well within the 128-char column limit).
|
||||
"""
|
||||
salt = secrets.token_bytes(_PWD_SALT_BYTES)
|
||||
dk = hashlib.pbkdf2_hmac(
|
||||
"sha256",
|
||||
password.encode("utf-8"),
|
||||
salt,
|
||||
_PWD_HASH_ITERATIONS,
|
||||
)
|
||||
return f"{salt.hex()}:{dk.hex()}"
|
||||
|
||||
|
||||
def _verify_password(password: str, stored_hash: str) -> bool:
|
||||
"""Verify *password* against a hash produced by :func:`_hash_password`.
|
||||
|
||||
Uses constant-time comparison to prevent timing attacks.
|
||||
|
||||
Args:
|
||||
password: Plaintext password to check.
|
||||
stored_hash: The value previously returned by :func:`_hash_password`.
|
||||
|
||||
Returns:
|
||||
``True`` if *password* matches, ``False`` otherwise.
|
||||
"""
|
||||
try:
|
||||
salt_hex, dk_hex = stored_hash.split(":", 1)
|
||||
salt = bytes.fromhex(salt_hex)
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
dk = hashlib.pbkdf2_hmac(
|
||||
"sha256",
|
||||
password.encode("utf-8"),
|
||||
salt,
|
||||
_PWD_HASH_ITERATIONS,
|
||||
)
|
||||
return secrets.compare_digest(dk.hex(), dk_hex)
|
||||
|
||||
|
||||
def _is_link_valid(link: SharedLink) -> bool:
|
||||
"""Return True when *link* is active, unexpired, and within view limit."""
|
||||
if not link.is_active:
|
||||
return False
|
||||
now = datetime.now(timezone.utc)
|
||||
if link.expires_at is not None:
|
||||
exp = link.expires_at
|
||||
if exp.tzinfo is None:
|
||||
exp = exp.replace(tzinfo=timezone.utc)
|
||||
if now > exp:
|
||||
return False
|
||||
if link.max_views is not None and link.view_count >= link.max_views:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _resolve_file_path(file_record: FileRecord) -> str | None:
|
||||
"""Return the best available file path for *file_record*.
|
||||
|
||||
Checks processed path first, then original, then local (tmp) path.
|
||||
Returns ``None`` when no file exists on disk.
|
||||
"""
|
||||
from app.config import settings
|
||||
|
||||
workdir = os.path.realpath(settings.workdir)
|
||||
candidates = [
|
||||
file_record.processed_file_path,
|
||||
file_record.original_file_path,
|
||||
file_record.local_filename,
|
||||
]
|
||||
for path in candidates:
|
||||
if not path:
|
||||
continue
|
||||
# Guard against path traversal in DB values.
|
||||
real = os.path.realpath(path)
|
||||
if not real.startswith(workdir + os.sep) and real != workdir:
|
||||
logger.warning("Shared link file path outside workdir rejected: %s", path)
|
||||
continue
|
||||
if os.path.exists(real):
|
||||
return real
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SharedLinkCreate(BaseModel):
|
||||
"""Schema for creating a new shared link."""
|
||||
|
||||
file_id: int = Field(..., description="ID of the file to share")
|
||||
expires_in_hours: int | None = Field(
|
||||
None,
|
||||
ge=1,
|
||||
le=720,
|
||||
description="Expiry in hours (1–720). NULL means the link never expires.",
|
||||
)
|
||||
max_views: int | None = Field(
|
||||
None,
|
||||
ge=1,
|
||||
le=10_000,
|
||||
description="Maximum number of downloads/views. NULL means unlimited.",
|
||||
)
|
||||
password: str | None = Field(
|
||||
None,
|
||||
min_length=1,
|
||||
max_length=128,
|
||||
description="Optional password protecting the link.",
|
||||
)
|
||||
label: str | None = Field(
|
||||
None,
|
||||
max_length=255,
|
||||
description="Optional human-readable label for the link.",
|
||||
)
|
||||
|
||||
@field_validator("expires_in_hours")
|
||||
@classmethod
|
||||
def validate_expiry(cls, v: int | None) -> int | None:
|
||||
if v is not None and v not in range(1, 721):
|
||||
raise ValueError("expires_in_hours must be between 1 and 720")
|
||||
return v
|
||||
|
||||
|
||||
class SharedLinkResponse(BaseModel):
|
||||
"""Shared link info returned to the authenticated owner."""
|
||||
|
||||
id: int
|
||||
token: str
|
||||
file_id: int
|
||||
label: str | None
|
||||
expires_at: datetime | None
|
||||
max_views: int | None
|
||||
view_count: int
|
||||
has_password: bool
|
||||
is_active: bool
|
||||
created_at: datetime | None
|
||||
revoked_at: datetime | None
|
||||
# Filled in by the endpoint, not stored in DB.
|
||||
share_url: str = ""
|
||||
original_filename: str | None = None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class SharedLinkInfoResponse(BaseModel):
|
||||
"""Public metadata about a shared link (used on the share landing page)."""
|
||||
|
||||
token: str
|
||||
label: str | None
|
||||
original_filename: str | None
|
||||
expires_at: datetime | None
|
||||
max_views: int | None
|
||||
view_count: int
|
||||
has_password: bool
|
||||
is_valid: bool
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Private (authenticated) endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED, response_model=SharedLinkResponse)
|
||||
async def create_shared_link(
|
||||
body: SharedLinkCreate,
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new shareable link for a document.
|
||||
|
||||
The caller must own the file (or be in single-user mode).
|
||||
Returns the full link metadata including the generated token.
|
||||
"""
|
||||
# Verify the file exists and belongs to the caller.
|
||||
q = db.query(FileRecord).filter(FileRecord.id == body.file_id)
|
||||
q = apply_owner_filter(q, request)
|
||||
file_record = q.first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
token = _generate_token()
|
||||
expires_at = None
|
||||
if body.expires_in_hours is not None:
|
||||
expires_at = datetime.now(timezone.utc).replace(microsecond=0)
|
||||
from datetime import timedelta
|
||||
|
||||
expires_at = expires_at + timedelta(hours=body.expires_in_hours)
|
||||
|
||||
password_hash = _hash_password(body.password) if body.password else None
|
||||
|
||||
db_link = SharedLink(
|
||||
token=token,
|
||||
file_id=body.file_id,
|
||||
owner_id=owner_id,
|
||||
label=body.label,
|
||||
expires_at=expires_at,
|
||||
max_views=body.max_views,
|
||||
view_count=0,
|
||||
password_hash=password_hash,
|
||||
)
|
||||
try:
|
||||
db.add(db_link)
|
||||
db.commit()
|
||||
db.refresh(db_link)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Shared link created: id=%s owner=%s file_id=%s", db_link.id, owner_id, body.file_id)
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
return _link_to_dict(db_link, base_url, file_record.original_filename)
|
||||
|
||||
|
||||
@router.get("/", response_model=list[SharedLinkResponse])
|
||||
async def list_shared_links(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
active_only: bool = Query(False, description="When true, only return active (non-revoked) links"),
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List all shared links created by the authenticated user."""
|
||||
q = db.query(SharedLink).filter(SharedLink.owner_id == owner_id)
|
||||
if active_only:
|
||||
q = q.filter(SharedLink.is_active.is_(True))
|
||||
links = q.order_by(SharedLink.created_at.desc()).all()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
result = []
|
||||
for link in links:
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == link.file_id).first()
|
||||
filename = file_record.original_filename if file_record else None
|
||||
result.append(_link_to_dict(link, base_url, filename))
|
||||
return result
|
||||
|
||||
|
||||
@router.delete("/{link_id}", status_code=status.HTTP_200_OK)
|
||||
async def revoke_shared_link(
|
||||
link_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Revoke (soft-delete) a shared link.
|
||||
|
||||
The record is kept for audit purposes but the link immediately
|
||||
stops working for recipients.
|
||||
"""
|
||||
db_link = db.query(SharedLink).filter(SharedLink.id == link_id, SharedLink.owner_id == owner_id).first()
|
||||
if not db_link:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Shared link not found")
|
||||
|
||||
if not db_link.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Link is already revoked")
|
||||
|
||||
try:
|
||||
db_link.is_active = False
|
||||
db_link.revoked_at = datetime.now(timezone.utc)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Shared link revoked: id=%s owner=%s", link_id, owner_id)
|
||||
return {"detail": "Link revoked"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public endpoints (no authentication required)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@public_router.get("/share/{token}/info", response_model=SharedLinkInfoResponse)
|
||||
def get_shared_link_info(
|
||||
token: str,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Return public metadata about a shared link.
|
||||
|
||||
Used by the share landing page to decide whether to show a password
|
||||
prompt or a direct download button. Never returns sensitive data.
|
||||
"""
|
||||
link = db.query(SharedLink).filter(SharedLink.token == token).first()
|
||||
if not link:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Link not found")
|
||||
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == link.file_id).first()
|
||||
filename = file_record.original_filename if file_record else None
|
||||
|
||||
return {
|
||||
"token": link.token,
|
||||
"label": link.label,
|
||||
"original_filename": filename,
|
||||
"expires_at": link.expires_at,
|
||||
"max_views": link.max_views,
|
||||
"view_count": link.view_count,
|
||||
"has_password": link.password_hash is not None,
|
||||
"is_valid": _is_link_valid(link),
|
||||
}
|
||||
|
||||
|
||||
@public_router.get("/share/{token}/download")
|
||||
def download_via_shared_link(
|
||||
token: str,
|
||||
db: DbSession,
|
||||
) -> FileResponse:
|
||||
"""Download a file via a shared link that does NOT require a password.
|
||||
|
||||
For password-protected links use ``POST /api/share/{token}/download``
|
||||
with ``{"password": "<value>"}`` in the JSON body instead.
|
||||
|
||||
Increments the view counter and validates expiry / view limit before
|
||||
serving the file.
|
||||
"""
|
||||
return _serve_shared_file(token, db, password=None)
|
||||
|
||||
|
||||
class PasswordBody(BaseModel):
|
||||
"""Request body for password-protected shared link downloads."""
|
||||
|
||||
password: str = Field(..., min_length=1, max_length=128, description="Password for the shared link")
|
||||
|
||||
|
||||
@public_router.post("/share/{token}/download")
|
||||
def download_via_shared_link_with_password(
|
||||
token: str,
|
||||
body: PasswordBody,
|
||||
db: DbSession,
|
||||
) -> FileResponse:
|
||||
"""Download a password-protected file via a shared link.
|
||||
|
||||
Accepts the password in the JSON request body to avoid it appearing in
|
||||
server access logs, browser history, or ``Referer`` headers.
|
||||
"""
|
||||
return _serve_shared_file(token, db, password=body.password)
|
||||
|
||||
|
||||
def _serve_shared_file(token: str, db: Session, password: str | None) -> FileResponse:
|
||||
"""Core download logic shared by the GET and POST download endpoints."""
|
||||
link = db.query(SharedLink).filter(SharedLink.token == token).first()
|
||||
if not link:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Link not found or expired")
|
||||
|
||||
if not _is_link_valid(link):
|
||||
raise HTTPException(status_code=status.HTTP_410_GONE, detail="Link has expired or reached its view limit")
|
||||
|
||||
# Password check
|
||||
if link.password_hash is not None:
|
||||
if not password:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="This link requires a password",
|
||||
)
|
||||
if not _verify_password(password, link.password_hash):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Incorrect password")
|
||||
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == link.file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
file_path = _resolve_file_path(file_record)
|
||||
if not file_path:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not available on disk")
|
||||
|
||||
# Increment view count — fail the request if this cannot be persisted so
|
||||
# that view-limited links are not bypassed during temporary DB outages.
|
||||
try:
|
||||
link.view_count = (link.view_count or 0) + 1
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.error("Failed to increment view_count for shared link id=%s — aborting download", link.id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Service temporarily unavailable. Please try again.",
|
||||
)
|
||||
|
||||
return FileResponse(
|
||||
path=file_path,
|
||||
media_type=file_record.mime_type or "application/octet-stream",
|
||||
headers={
|
||||
"Content-Disposition": f'attachment; filename="{file_record.original_filename or "document"}"',
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helper
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _link_to_dict(link: SharedLink, base_url: str, original_filename: str | None) -> dict[str, Any]:
|
||||
"""Serialise a ``SharedLink`` ORM row to a plain dict."""
|
||||
return {
|
||||
"id": link.id,
|
||||
"token": link.token,
|
||||
"file_id": link.file_id,
|
||||
"label": link.label,
|
||||
"expires_at": link.expires_at,
|
||||
"max_views": link.max_views,
|
||||
"view_count": link.view_count,
|
||||
"has_password": link.password_hash is not None,
|
||||
"is_active": link.is_active,
|
||||
"created_at": link.created_at,
|
||||
"revoked_at": link.revoked_at,
|
||||
"share_url": f"{base_url}/share/{link.token}",
|
||||
"original_filename": original_filename,
|
||||
}
|
||||
@@ -0,0 +1,473 @@
|
||||
"""Document similarity API endpoints.
|
||||
|
||||
Provides endpoints to find documents similar to a given file based on
|
||||
text embeddings and cosine similarity scoring, plus debug/diagnostic
|
||||
endpoints for inspecting and triggering embedding computation.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import FileRecord
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/similar")
|
||||
@require_login
|
||||
def get_similar_documents(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
limit: int = Query(5, ge=1, le=20, description="Maximum number of similar documents to return"),
|
||||
threshold: float = Query(0.3, ge=0.0, le=1.0, description="Minimum similarity score (0–1)"),
|
||||
):
|
||||
"""Find documents similar to the specified file.
|
||||
|
||||
Uses text embeddings generated from OCR-extracted text and cosine
|
||||
similarity to rank documents by relevance. Similarity scores range
|
||||
from 0 (completely different) to 1 (identical content).
|
||||
|
||||
Embeddings are generated on first access and cached for subsequent
|
||||
requests. Documents without OCR text are excluded.
|
||||
|
||||
Query Parameters:
|
||||
- limit: Maximum results to return (default: 5, max: 20)
|
||||
- threshold: Minimum similarity score to include (default: 0.3)
|
||||
|
||||
Example:
|
||||
```
|
||||
GET /api/files/42/similar?limit=5&threshold=0.5
|
||||
```
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"file_id": 42,
|
||||
"similar_documents": [
|
||||
{
|
||||
"file_id": 15,
|
||||
"original_filename": "Invoice_2026-01.pdf",
|
||||
"document_title": "January Invoice",
|
||||
"similarity_score": 0.8934,
|
||||
"mime_type": "application/pdf",
|
||||
"created_at": "2026-01-15T10:30:00+00:00"
|
||||
}
|
||||
],
|
||||
"count": 1
|
||||
}
|
||||
```
|
||||
"""
|
||||
# Verify the file exists
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
if not file_record.ocr_text or not file_record.ocr_text.strip():
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"similar_documents": [],
|
||||
"count": 0,
|
||||
"message": "No OCR text available for similarity comparison",
|
||||
}
|
||||
|
||||
# Check whether an embedding has been computed yet
|
||||
if not file_record.embedding:
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"similar_documents": [],
|
||||
"count": 0,
|
||||
"message": (
|
||||
"Embedding not yet computed for this file. "
|
||||
"It will be generated automatically during processing or via the backfill task. "
|
||||
"You can also trigger it manually with POST /api/files/{file_id}/compute-embedding."
|
||||
),
|
||||
}
|
||||
|
||||
try:
|
||||
from app.utils.similarity import find_similar_documents
|
||||
|
||||
similar = find_similar_documents(db, file_id, limit=limit, threshold=threshold)
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"similar_documents": similar,
|
||||
"count": len(similar),
|
||||
}
|
||||
except Exception as e:
|
||||
logger.error(f"Error finding similar documents for file {file_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to compute document similarity",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Debug / diagnostic endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/embedding-status")
|
||||
@require_login
|
||||
def get_embedding_status(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
):
|
||||
"""Return the embedding status for a single file.
|
||||
|
||||
Useful for debugging whether the embedding has been computed
|
||||
and cached for a given document.
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"file_id": 42,
|
||||
"has_embedding": true,
|
||||
"embedding_dimensions": 1536,
|
||||
"has_ocr_text": true,
|
||||
"ocr_text_length": 4200,
|
||||
"embedding_model": "text-embedding-3-small"
|
||||
}
|
||||
```
|
||||
"""
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
has_embedding = False
|
||||
embedding_dimensions = None
|
||||
if file_record.embedding:
|
||||
try:
|
||||
parsed = json.loads(file_record.embedding)
|
||||
has_embedding = True
|
||||
embedding_dimensions = len(parsed)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
has_ocr_text = bool(file_record.ocr_text and file_record.ocr_text.strip())
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"has_embedding": has_embedding,
|
||||
"embedding_dimensions": embedding_dimensions,
|
||||
"has_ocr_text": has_ocr_text,
|
||||
"ocr_text_length": len(file_record.ocr_text) if file_record.ocr_text else 0,
|
||||
"embedding_model": settings.embedding_model,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/files/{file_id}/compute-embedding")
|
||||
@require_login
|
||||
def trigger_compute_embedding(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
):
|
||||
"""Trigger embedding computation for a single file.
|
||||
|
||||
If the file already has a cached embedding it will be recomputed.
|
||||
The computation happens synchronously so the caller receives the
|
||||
result immediately.
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"file_id": 42,
|
||||
"status": "success",
|
||||
"embedding_dimensions": 1536
|
||||
}
|
||||
```
|
||||
"""
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
if not file_record.ocr_text or not file_record.ocr_text.strip():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="File has no OCR text — cannot generate embedding",
|
||||
)
|
||||
|
||||
try:
|
||||
from app.utils.similarity import generate_embedding
|
||||
|
||||
# Clear cached embedding to force recomputation
|
||||
file_record.embedding = None
|
||||
db.flush()
|
||||
|
||||
embedding = generate_embedding(file_record.ocr_text)
|
||||
file_record.embedding = json.dumps(embedding)
|
||||
db.commit()
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"status": "success",
|
||||
"embedding_dimensions": len(embedding),
|
||||
}
|
||||
except Exception as e:
|
||||
db.rollback()
|
||||
logger.error(f"Failed to compute embedding for file {file_id}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Embedding computation failed: {e}",
|
||||
)
|
||||
|
||||
|
||||
@router.get("/diagnostic/embeddings")
|
||||
@require_login
|
||||
def get_embeddings_overview(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
):
|
||||
"""Return an overview of embedding status across all files.
|
||||
|
||||
Provides aggregate counts as well as a per-file breakdown so an
|
||||
administrator can quickly identify documents that are missing
|
||||
embeddings.
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"total_files": 120,
|
||||
"files_with_ocr_text": 95,
|
||||
"files_with_embedding": 42,
|
||||
"files_missing_embedding": 53,
|
||||
"embedding_model": "text-embedding-3-small",
|
||||
"files": [
|
||||
{
|
||||
"file_id": 1,
|
||||
"original_filename": "invoice.pdf",
|
||||
"has_ocr_text": true,
|
||||
"has_embedding": true,
|
||||
"embedding_dimensions": 1536
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
"""
|
||||
# Use column-only query to avoid loading full ORM objects into memory
|
||||
all_files = (
|
||||
db.query(
|
||||
FileRecord.id,
|
||||
FileRecord.original_filename,
|
||||
FileRecord.ocr_text,
|
||||
FileRecord.embedding,
|
||||
)
|
||||
.order_by(FileRecord.id.desc())
|
||||
.all()
|
||||
)
|
||||
|
||||
files_info = []
|
||||
total_with_ocr = 0
|
||||
total_with_embedding = 0
|
||||
|
||||
for f in all_files:
|
||||
has_ocr = bool(f.ocr_text and f.ocr_text.strip())
|
||||
has_emb = False
|
||||
emb_dims = None
|
||||
|
||||
if f.embedding:
|
||||
try:
|
||||
parsed = json.loads(f.embedding)
|
||||
has_emb = True
|
||||
emb_dims = len(parsed)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
if has_ocr:
|
||||
total_with_ocr += 1
|
||||
if has_emb:
|
||||
total_with_embedding += 1
|
||||
|
||||
files_info.append(
|
||||
{
|
||||
"file_id": f.id,
|
||||
"original_filename": f.original_filename,
|
||||
"has_ocr_text": has_ocr,
|
||||
"has_embedding": has_emb,
|
||||
"embedding_dimensions": emb_dims,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"total_files": len(all_files),
|
||||
"files_with_ocr_text": total_with_ocr,
|
||||
"files_with_embedding": total_with_embedding,
|
||||
"files_missing_embedding": total_with_ocr - total_with_embedding,
|
||||
"embedding_model": settings.embedding_model,
|
||||
"files": files_info,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/diagnostic/compute-all-embeddings")
|
||||
@require_login
|
||||
def trigger_compute_all_embeddings(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
):
|
||||
"""Queue embedding computation for all files that have OCR text but no embedding.
|
||||
|
||||
Each file is processed as a separate Celery task so the endpoint
|
||||
returns immediately.
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"status": "queued",
|
||||
"files_queued": 53
|
||||
}
|
||||
```
|
||||
"""
|
||||
candidates = (
|
||||
db.query(FileRecord)
|
||||
.filter(
|
||||
FileRecord.ocr_text.isnot(None),
|
||||
FileRecord.ocr_text != "",
|
||||
(FileRecord.embedding.is_(None)) | (FileRecord.embedding == ""),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
|
||||
queued = 0
|
||||
for f in candidates:
|
||||
try:
|
||||
from app.tasks.compute_embedding import compute_document_embedding
|
||||
|
||||
compute_document_embedding.delay(f.id)
|
||||
queued += 1
|
||||
except Exception as e:
|
||||
logger.warning(f"Could not queue embedding for file {f.id}: {e}")
|
||||
|
||||
return {
|
||||
"status": "queued",
|
||||
"files_queued": queued,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/similarity/pairs")
|
||||
@require_login
|
||||
def get_similarity_pairs(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
threshold: float = Query(0.7, ge=0.0, le=1.0, description="Minimum similarity score for a pair"),
|
||||
limit: int = Query(50, ge=1, le=200, description="Maximum number of pairs to return"),
|
||||
page: int = Query(1, ge=1, description="Page number"),
|
||||
):
|
||||
"""Return pairs of documents with high similarity across the entire corpus.
|
||||
|
||||
Unlike the per-file ``/files/{id}/similar`` endpoint, this scans every
|
||||
document that has a pre-computed embedding and returns **all** pairs
|
||||
whose cosine similarity exceeds ``threshold``, sorted by descending
|
||||
score.
|
||||
|
||||
To keep memory bounded the query loads only the columns needed for
|
||||
scoring and streams results in chunks.
|
||||
|
||||
Response:
|
||||
```json
|
||||
{
|
||||
"pairs": [
|
||||
{
|
||||
"file_a": {"file_id": 1, "original_filename": "invoice_jan.pdf", ...},
|
||||
"file_b": {"file_id": 5, "original_filename": "invoice_feb.pdf", ...},
|
||||
"similarity_score": 0.94
|
||||
}
|
||||
],
|
||||
"total_pairs": 12,
|
||||
"threshold": 0.7,
|
||||
"page": 1,
|
||||
"pages": 1,
|
||||
"embedding_coverage": {"total_files": 120, "files_with_embedding": 95}
|
||||
}
|
||||
```
|
||||
"""
|
||||
from app.utils.similarity import cosine_similarity
|
||||
|
||||
# Load all files that have embeddings (columns only for efficiency)
|
||||
rows = (
|
||||
db.query(
|
||||
FileRecord.id,
|
||||
FileRecord.original_filename,
|
||||
FileRecord.document_title,
|
||||
FileRecord.mime_type,
|
||||
FileRecord.created_at,
|
||||
FileRecord.embedding,
|
||||
)
|
||||
.filter(
|
||||
FileRecord.embedding.isnot(None),
|
||||
FileRecord.embedding != "",
|
||||
)
|
||||
.order_by(FileRecord.id)
|
||||
.all()
|
||||
)
|
||||
|
||||
# Parse embeddings upfront
|
||||
parsed: list[tuple] = []
|
||||
for row in rows:
|
||||
try:
|
||||
vec = json.loads(row.embedding)
|
||||
parsed.append((row, vec))
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
continue
|
||||
|
||||
# Pairwise comparison (triangle: i < j avoids duplicating A↔B / B↔A)
|
||||
all_pairs: list[dict] = []
|
||||
for i in range(len(parsed)):
|
||||
row_a, vec_a = parsed[i]
|
||||
for j in range(i + 1, len(parsed)):
|
||||
row_b, vec_b = parsed[j]
|
||||
score = cosine_similarity(vec_a, vec_b)
|
||||
if score >= threshold:
|
||||
all_pairs.append(
|
||||
{
|
||||
"file_a": _row_to_dict(row_a),
|
||||
"file_b": _row_to_dict(row_b),
|
||||
"similarity_score": round(score, 4),
|
||||
}
|
||||
)
|
||||
|
||||
# Sort by score descending
|
||||
all_pairs.sort(key=lambda p: p["similarity_score"], reverse=True)
|
||||
|
||||
total_pairs = len(all_pairs)
|
||||
total_pages = max(1, (total_pairs + limit - 1) // limit)
|
||||
offset = (page - 1) * limit
|
||||
page_pairs = all_pairs[offset : offset + limit]
|
||||
|
||||
total_files = db.query(FileRecord).count()
|
||||
|
||||
return {
|
||||
"pairs": page_pairs,
|
||||
"total_pairs": total_pairs,
|
||||
"threshold": threshold,
|
||||
"page": page,
|
||||
"pages": total_pages,
|
||||
"per_page": limit,
|
||||
"embedding_coverage": {
|
||||
"total_files": total_files,
|
||||
"files_with_embedding": len(parsed),
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _row_to_dict(row) -> dict:
|
||||
"""Serialise a column-only query row to a dict for JSON responses."""
|
||||
return {
|
||||
"file_id": row.id,
|
||||
"original_filename": row.original_filename,
|
||||
"document_title": row.document_title,
|
||||
"mime_type": row.mime_type,
|
||||
"created_at": row.created_at.isoformat() if row.created_at else None,
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
"""API endpoints for subscription tiers and usage statistics.
|
||||
|
||||
Public endpoints:
|
||||
GET /api/subscriptions/tiers — list all available plans
|
||||
GET /api/subscriptions/my — current user's plan + usage (auth required)
|
||||
POST /api/subscriptions/change — request a plan change (auth required)
|
||||
DELETE /api/subscriptions/change — cancel a pending plan change (auth required)
|
||||
GET /api/subscriptions/platform — platform-wide stats (admin only)
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.admin_users import _require_admin
|
||||
from app.database import get_db
|
||||
from app.utils.subscription import (
|
||||
TIER_ORDER,
|
||||
TIERS,
|
||||
SubscriptionChangeError,
|
||||
apply_pending_subscription_changes,
|
||||
cancel_pending_subscription_change,
|
||||
get_all_tiers,
|
||||
get_tier,
|
||||
get_user_tier_id,
|
||||
get_user_usage,
|
||||
request_subscription_change,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/subscriptions", tags=["subscriptions"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request / response models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SubscriptionChangeRequest(BaseModel):
|
||||
"""Request body for a subscription plan change."""
|
||||
|
||||
plan_id: str
|
||||
billing_cycle: str = "monthly" # "monthly" | "yearly"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_owner_id(request: Request) -> str:
|
||||
"""Extract the authenticated user's owner_id from the session."""
|
||||
user = request.session.get("user") or {}
|
||||
return user.get("username") or user.get("email") or user.get("sub") or ""
|
||||
|
||||
|
||||
def _require_authenticated(request: Request) -> str:
|
||||
"""Return the owner_id or raise 401."""
|
||||
owner_id = _get_owner_id(request)
|
||||
if not owner_id:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Authentication required")
|
||||
return owner_id
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/tiers", summary="List all subscription tiers")
|
||||
def list_tiers() -> dict[str, Any]:
|
||||
"""Return the full list of subscription plans in display order."""
|
||||
return {
|
||||
"tiers": get_all_tiers(),
|
||||
"order": TIER_ORDER,
|
||||
"default": "free",
|
||||
}
|
||||
|
||||
|
||||
@router.get("/my", summary="Get current user's subscription and usage")
|
||||
def my_subscription(request: Request, db: DbSession) -> dict[str, Any]:
|
||||
"""Return the authenticated user's subscription tier and current usage counts.
|
||||
|
||||
Also applies any pending subscription changes that have become due.
|
||||
"""
|
||||
from app.config import settings
|
||||
from app.models import UserProfile
|
||||
|
||||
user = request.session.get("user")
|
||||
|
||||
if not settings.multi_user_enabled:
|
||||
# In single-user mode there is no concept of a subscription plan
|
||||
return {
|
||||
"multi_user_mode": False,
|
||||
"tier": TIERS["business"], # unrestricted
|
||||
"usage": None,
|
||||
}
|
||||
|
||||
if not user:
|
||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Authentication required")
|
||||
|
||||
owner_id: str = user.get("username") or user.get("email") or user.get("sub") or ""
|
||||
|
||||
# Apply any pending change that has become due
|
||||
apply_pending_subscription_changes(db, owner_id)
|
||||
|
||||
tier_id = get_user_tier_id(db, owner_id)
|
||||
tier = get_tier(tier_id, db)
|
||||
usage = get_user_usage(db, owner_id)
|
||||
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == owner_id).first()
|
||||
pending_tier_id: str | None = profile.subscription_change_pending_tier if profile else None
|
||||
pending_date: str | None = (
|
||||
profile.subscription_change_pending_date.isoformat()
|
||||
if profile and profile.subscription_change_pending_date
|
||||
else None
|
||||
)
|
||||
period_start: str | None = (
|
||||
profile.subscription_period_start.isoformat() if profile and profile.subscription_period_start else None
|
||||
)
|
||||
|
||||
return {
|
||||
"multi_user_mode": True,
|
||||
"owner_id": owner_id,
|
||||
"tier": tier,
|
||||
"usage": usage,
|
||||
"period_start": period_start,
|
||||
"pending_change": (
|
||||
{
|
||||
"tier_id": pending_tier_id,
|
||||
"tier": get_tier(pending_tier_id, db),
|
||||
"effective_date": pending_date,
|
||||
}
|
||||
if pending_tier_id
|
||||
else None
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/change", summary="Request a subscription plan change", status_code=status.HTTP_200_OK)
|
||||
def change_subscription(request: Request, body: SubscriptionChangeRequest, db: DbSession) -> dict[str, Any]:
|
||||
"""Request a subscription tier change.
|
||||
|
||||
**Upgrades** (moving to a higher-ranked plan) take effect immediately.
|
||||
|
||||
**Downgrades** (moving to a lower-ranked plan) are scheduled for the end
|
||||
of the current billing period to prevent gaming. The user keeps their
|
||||
current plan benefits until the scheduled date.
|
||||
|
||||
Requesting the currently active tier while a downgrade is pending cancels
|
||||
that pending change.
|
||||
"""
|
||||
from app.config import settings
|
||||
|
||||
if not settings.multi_user_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Subscription management is not available in single-user mode.",
|
||||
)
|
||||
|
||||
owner_id = _require_authenticated(request)
|
||||
|
||||
try:
|
||||
result = request_subscription_change(db, owner_id, body.plan_id, body.billing_cycle)
|
||||
except SubscriptionChangeError as exc:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@router.delete("/change", summary="Cancel a pending subscription change", status_code=status.HTTP_200_OK)
|
||||
def cancel_subscription_change(request: Request, db: DbSession) -> dict[str, Any]:
|
||||
"""Cancel a scheduled future subscription change.
|
||||
|
||||
Only downgrades can be pending; upgrades always take effect immediately.
|
||||
Returns 404 when there is no pending change to cancel.
|
||||
"""
|
||||
from app.config import settings
|
||||
|
||||
if not settings.multi_user_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Subscription management is not available in single-user mode.",
|
||||
)
|
||||
|
||||
owner_id = _require_authenticated(request)
|
||||
|
||||
cancelled = cancel_pending_subscription_change(db, owner_id)
|
||||
if not cancelled:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="No pending subscription change found.")
|
||||
|
||||
return {"cancelled": True, "message": "Your pending subscription change has been cancelled."}
|
||||
|
||||
|
||||
@router.get("/platform", summary="Platform-wide usage statistics (admin only)")
|
||||
def platform_stats(request: Request, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Return aggregate statistics across all users and tiers (admin only)."""
|
||||
from app.models import FileRecord, UserProfile
|
||||
|
||||
today = datetime.now(timezone.utc).date()
|
||||
|
||||
# Total files
|
||||
total_files: int = db.query(func.count(FileRecord.id)).scalar() or 0
|
||||
|
||||
# Files today
|
||||
files_today: int = (
|
||||
db.query(func.count(FileRecord.id)).filter(func.date(FileRecord.created_at) == today).scalar() or 0
|
||||
)
|
||||
|
||||
# Files this month
|
||||
files_this_month: int = (
|
||||
db.query(func.count(FileRecord.id))
|
||||
.filter(func.strftime("%Y-%m", FileRecord.created_at) == today.strftime("%Y-%m"))
|
||||
.scalar()
|
||||
or 0
|
||||
)
|
||||
|
||||
# Files with OCR text (proxy for pages OCRed — approximation)
|
||||
files_with_ocr: int = db.query(func.count(FileRecord.id)).filter(FileRecord.ocr_text.isnot(None)).scalar() or 0
|
||||
|
||||
# Unique active users (ever uploaded)
|
||||
unique_users: int = (
|
||||
db.query(func.count(func.distinct(FileRecord.owner_id))).filter(FileRecord.owner_id.isnot(None)).scalar() or 0
|
||||
)
|
||||
|
||||
# Users per subscription tier
|
||||
profiles = (
|
||||
db.query(UserProfile.subscription_tier, func.count(UserProfile.id))
|
||||
.group_by(UserProfile.subscription_tier)
|
||||
.all()
|
||||
)
|
||||
tier_distribution: dict[str, int] = {row[0] or "free": row[1] for row in profiles}
|
||||
|
||||
# Fill in zeros for tiers with no users
|
||||
for tid in TIER_ORDER:
|
||||
tier_distribution.setdefault(tid, 0)
|
||||
|
||||
return {
|
||||
"files": {
|
||||
"total": total_files,
|
||||
"today": files_today,
|
||||
"this_month": files_this_month,
|
||||
"with_ocr": files_with_ocr,
|
||||
},
|
||||
"users": {
|
||||
"unique_uploaders": unique_users,
|
||||
"tier_distribution": tier_distribution,
|
||||
},
|
||||
"generated_at": datetime.now(timezone.utc).isoformat(),
|
||||
}
|
||||
+1
-32
@@ -2,7 +2,6 @@
|
||||
API endpoint for processing files from URLs
|
||||
"""
|
||||
|
||||
import ipaddress
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
@@ -19,6 +18,7 @@ from app.config import settings
|
||||
from app.tasks.process_document import process_document
|
||||
from app.utils.allowed_types import ALLOWED_MIME_TYPES
|
||||
from app.utils.filename_utils import sanitize_filename
|
||||
from app.utils.network import is_private_ip
|
||||
|
||||
# Set up logging
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -42,37 +42,6 @@ class URLUploadRequest(BaseModel):
|
||||
return v
|
||||
|
||||
|
||||
def is_private_ip(hostname: str) -> bool:
|
||||
"""
|
||||
Check if a hostname resolves to a private/internal IP address.
|
||||
Protects against SSRF attacks by blocking access to internal networks.
|
||||
"""
|
||||
try:
|
||||
# Try to parse as IP address directly
|
||||
ip = ipaddress.ip_address(hostname)
|
||||
return ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved
|
||||
except ValueError:
|
||||
# Not a direct IP, try to resolve hostname
|
||||
try:
|
||||
import socket
|
||||
|
||||
# Get all IP addresses for this hostname
|
||||
addr_info = socket.getaddrinfo(hostname, None)
|
||||
for info in addr_info:
|
||||
ip_str = info[4][0]
|
||||
ip = ipaddress.ip_address(ip_str)
|
||||
# Block if ANY resolved IP is private/internal
|
||||
if ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved:
|
||||
return True
|
||||
return False
|
||||
except (socket.gaierror, socket.error):
|
||||
# Cannot resolve - allow for testing/development
|
||||
# In production, DNS should work properly
|
||||
# Log this for debugging
|
||||
logger.warning(f"Could not resolve hostname: {hostname}")
|
||||
return False # Changed from True to False to allow external domains in tests
|
||||
|
||||
|
||||
def validate_url_safety(url: str) -> None:
|
||||
"""
|
||||
Validate that URL is safe to fetch (SSRF protection).
|
||||
|
||||
+56
-7
@@ -4,16 +4,25 @@ User-related API endpoints
|
||||
|
||||
import logging
|
||||
from hashlib import md5
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import FileRecord, UserProfile
|
||||
|
||||
# Set up logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
async def whoami_handler(request: Request):
|
||||
|
||||
async def whoami_handler(request: Request, db: Session):
|
||||
"""
|
||||
Returns user info if logged in, else 401.
|
||||
"""
|
||||
@@ -32,17 +41,57 @@ async def whoami_handler(request: Request):
|
||||
|
||||
# Add the gravatar URL to the user object instead of creating a new response
|
||||
user_response = user.copy() # Create a copy to avoid modifying the session
|
||||
user_response["picture"] = gravatar_url
|
||||
|
||||
# Check if the user has a custom avatar stored in their profile
|
||||
user_id = user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id")
|
||||
if user_id:
|
||||
try:
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if profile and profile.avatar_data:
|
||||
user_response["picture"] = profile.avatar_data
|
||||
else:
|
||||
user_response["picture"] = gravatar_url
|
||||
except Exception:
|
||||
user_response["picture"] = gravatar_url
|
||||
else:
|
||||
user_response["picture"] = gravatar_url
|
||||
|
||||
return user_response
|
||||
|
||||
|
||||
# Register the same handler under two different paths
|
||||
@router.get("/whoami")
|
||||
async def whoami(request: Request):
|
||||
return await whoami_handler(request)
|
||||
async def whoami(request: Request, db: DbSession):
|
||||
return await whoami_handler(request, db)
|
||||
|
||||
|
||||
@router.get("/auth/whoami")
|
||||
async def auth_whoami(request: Request):
|
||||
return await whoami_handler(request)
|
||||
async def auth_whoami(request: Request, db: DbSession):
|
||||
return await whoami_handler(request, db)
|
||||
|
||||
|
||||
@router.get("/users/search")
|
||||
@require_login
|
||||
def search_known_users(
|
||||
db: DbSession,
|
||||
q: str = Query("", description="Substring to match against known owner IDs"),
|
||||
limit: int = Query(5, ge=1, le=20, description="Maximum number of results"),
|
||||
):
|
||||
"""
|
||||
Search known user identifiers (owner_ids) from existing documents.
|
||||
|
||||
Returns distinct ``owner_id`` values from the files table that contain
|
||||
the query string as a case-insensitive substring. Results are limited
|
||||
to at most ``limit`` entries (default 5).
|
||||
|
||||
This powers the autocomplete widget on the settings page for the
|
||||
``default_owner_id`` field.
|
||||
"""
|
||||
base_query = db.query(FileRecord.owner_id).filter(FileRecord.owner_id.isnot(None)).distinct()
|
||||
|
||||
if q.strip():
|
||||
base_query = base_query.filter(func.lower(FileRecord.owner_id).contains(q.strip().lower()))
|
||||
|
||||
results = base_query.order_by(FileRecord.owner_id).limit(limit).all()
|
||||
|
||||
return {"users": [row[0] for row in results]}
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
"""API endpoints for managing webhook configurations.
|
||||
|
||||
Provides CRUD operations for webhook configs that notify external systems
|
||||
when document events occur (``document.uploaded``, ``document.processed``,
|
||||
``document.failed``).
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.database import get_db
|
||||
from app.models import WebhookConfig
|
||||
from app.utils.webhook import VALID_EVENTS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/webhooks", tags=["webhooks"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper (reuse the pattern from settings API)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_admin(request: Request) -> dict:
|
||||
"""Ensure the caller is an admin. Raises 403 otherwise."""
|
||||
user = request.session.get("user")
|
||||
if not user or not user.get("is_admin"):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Admin access required")
|
||||
return user
|
||||
|
||||
|
||||
AdminUser = Annotated[dict, Depends(_require_admin)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class WebhookCreate(BaseModel):
|
||||
"""Schema for creating a new webhook configuration."""
|
||||
|
||||
url: str = Field(..., min_length=1, max_length=2048, description="Target URL for webhook delivery")
|
||||
secret: str | None = Field(default=None, max_length=512, description="Shared secret for HMAC-SHA256 signatures")
|
||||
events: list[str] = Field(..., min_length=1, description="List of events to subscribe to")
|
||||
is_active: bool = Field(default=True, description="Whether the webhook is active")
|
||||
description: str | None = Field(default=None, max_length=500, description="Optional human-readable description")
|
||||
|
||||
|
||||
class WebhookUpdate(BaseModel):
|
||||
"""Schema for updating an existing webhook configuration."""
|
||||
|
||||
url: str | None = Field(default=None, min_length=1, max_length=2048)
|
||||
secret: str | None = Field(default=None, max_length=512)
|
||||
events: list[str] | None = Field(default=None, min_length=1)
|
||||
is_active: bool | None = None
|
||||
description: str | None = Field(default=None, max_length=500)
|
||||
|
||||
|
||||
class WebhookResponse(BaseModel):
|
||||
"""Schema returned to clients (secret is never exposed)."""
|
||||
|
||||
id: int
|
||||
url: str
|
||||
events: list[str]
|
||||
is_active: bool
|
||||
description: str | None
|
||||
has_secret: bool
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _validate_events(events: list[str]) -> None:
|
||||
"""Raise 422 if any event name is not recognised."""
|
||||
invalid = set(events) - VALID_EVENTS
|
||||
if invalid:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Invalid event(s): {', '.join(sorted(invalid))}. Valid events: {', '.join(sorted(VALID_EVENTS))}",
|
||||
)
|
||||
|
||||
|
||||
def _to_response(cfg: WebhookConfig) -> dict[str, Any]:
|
||||
"""Convert a DB model instance to a response dict."""
|
||||
try:
|
||||
events = json.loads(cfg.events)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
events = []
|
||||
return {
|
||||
"id": cfg.id,
|
||||
"url": cfg.url,
|
||||
"events": events,
|
||||
"is_active": cfg.is_active,
|
||||
"description": cfg.description,
|
||||
"has_secret": cfg.secret is not None and len(cfg.secret) > 0,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", summary="List all webhook configurations")
|
||||
def list_webhooks(db: DbSession, _admin: AdminUser) -> list[dict[str, Any]]:
|
||||
"""Return all webhook configurations. Secrets are never included."""
|
||||
configs = db.query(WebhookConfig).order_by(WebhookConfig.id).all()
|
||||
return [_to_response(c) for c in configs]
|
||||
|
||||
|
||||
@router.get("/{webhook_id}", summary="Get a single webhook configuration")
|
||||
def get_webhook(webhook_id: int, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Return a single webhook configuration by ID."""
|
||||
cfg = db.query(WebhookConfig).filter(WebhookConfig.id == webhook_id).first()
|
||||
if not cfg:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Webhook not found")
|
||||
return _to_response(cfg)
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED, summary="Create a webhook configuration")
|
||||
def create_webhook(body: WebhookCreate, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Create a new webhook configuration."""
|
||||
_validate_events(body.events)
|
||||
|
||||
cfg = WebhookConfig(
|
||||
url=body.url,
|
||||
secret=body.secret,
|
||||
events=json.dumps(sorted(body.events)),
|
||||
is_active=body.is_active,
|
||||
description=body.description,
|
||||
)
|
||||
try:
|
||||
db.add(cfg)
|
||||
db.commit()
|
||||
db.refresh(cfg)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Webhook %d created for events %s", cfg.id, body.events)
|
||||
return _to_response(cfg)
|
||||
|
||||
|
||||
@router.put("/{webhook_id}", summary="Update a webhook configuration")
|
||||
def update_webhook(webhook_id: int, body: WebhookUpdate, db: DbSession, _admin: AdminUser) -> dict[str, Any]:
|
||||
"""Update an existing webhook configuration. Only supplied fields are changed."""
|
||||
cfg = db.query(WebhookConfig).filter(WebhookConfig.id == webhook_id).first()
|
||||
if not cfg:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Webhook not found")
|
||||
|
||||
if body.url is not None:
|
||||
cfg.url = body.url
|
||||
if body.secret is not None:
|
||||
cfg.secret = body.secret
|
||||
if body.events is not None:
|
||||
_validate_events(body.events)
|
||||
cfg.events = json.dumps(sorted(body.events))
|
||||
if body.is_active is not None:
|
||||
cfg.is_active = body.is_active
|
||||
if body.description is not None:
|
||||
cfg.description = body.description
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(cfg)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Webhook %d updated", cfg.id)
|
||||
return _to_response(cfg)
|
||||
|
||||
|
||||
@router.delete("/{webhook_id}", status_code=status.HTTP_204_NO_CONTENT, summary="Delete a webhook configuration")
|
||||
def delete_webhook(webhook_id: int, db: DbSession, _admin: AdminUser) -> None:
|
||||
"""Delete a webhook configuration."""
|
||||
cfg = db.query(WebhookConfig).filter(WebhookConfig.id == webhook_id).first()
|
||||
if not cfg:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Webhook not found")
|
||||
|
||||
try:
|
||||
db.delete(cfg)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Webhook %d deleted", cfg.id)
|
||||
|
||||
|
||||
@router.get("/events/", summary="List valid webhook event types")
|
||||
def list_events(_admin: AdminUser) -> list[str]:
|
||||
"""Return the list of valid event types that can be subscribed to."""
|
||||
return sorted(VALID_EVENTS)
|
||||
+802
-24
@@ -2,29 +2,49 @@ import hashlib
|
||||
import inspect
|
||||
import logging
|
||||
import pathlib
|
||||
from datetime import datetime, timezone
|
||||
from functools import wraps
|
||||
from urllib.parse import urlencode, urlparse
|
||||
|
||||
from authlib.integrations.starlette_client import OAuth
|
||||
from fastapi import APIRouter, Request, status
|
||||
from fastapi import APIRouter, Depends, Request, status
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi.templating import Jinja2Templates
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
from starlette.responses import RedirectResponse
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.middleware.audit_log import get_client_ip
|
||||
|
||||
oauth = OAuth()
|
||||
# Conditional imports: only used when multi_user_enabled=True. Imported here at
|
||||
# module level (not inside auth()) so they don't incur repeated import overhead.
|
||||
# Guards at call-sites ensure they are never *called* in single-user mode.
|
||||
from app.models import LocalUser as _LocalUser
|
||||
from app.models import UserProfile as _UserProfile
|
||||
from app.utils.i18n import translate as _translate
|
||||
from app.utils.local_auth import build_session_user as _build_session_user
|
||||
from app.utils.local_auth import verify_password as _verify_password
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
oauth = OAuth()
|
||||
|
||||
AUTH_ENABLED = settings.auth_enabled
|
||||
|
||||
# Set up templates for authentication
|
||||
templates_dir = pathlib.Path(__file__).parents[1] / "frontend" / "templates"
|
||||
templates = Jinja2Templates(directory=str(templates_dir))
|
||||
templates.env.globals["_"] = lambda key, **kwargs: _translate(key, "en", **kwargs)
|
||||
|
||||
# Configure OAuth provider if credentials are provided
|
||||
OAUTH_CONFIGURED = False
|
||||
OAUTH_PROVIDER_NAME = "Single Sign-On"
|
||||
|
||||
# Social login providers that are enabled and registered
|
||||
SOCIAL_PROVIDERS: dict[str, dict[str, str]] = {}
|
||||
|
||||
if AUTH_ENABLED and settings.authentik_client_id and settings.authentik_client_secret:
|
||||
oauth.register(
|
||||
name="authentik",
|
||||
@@ -36,27 +56,191 @@ if AUTH_ENABLED and settings.authentik_client_id and settings.authentik_client_s
|
||||
OAUTH_CONFIGURED = True
|
||||
OAUTH_PROVIDER_NAME = settings.oauth_provider_name or "Authentik SSO"
|
||||
|
||||
# --- Social Login Providers ---------------------------------------------------
|
||||
if AUTH_ENABLED and settings.social_auth_google_enabled:
|
||||
if settings.social_auth_google_client_id and settings.social_auth_google_client_secret:
|
||||
oauth.register(
|
||||
name="google",
|
||||
client_id=settings.social_auth_google_client_id,
|
||||
client_secret=settings.social_auth_google_client_secret,
|
||||
server_metadata_url="https://accounts.google.com/.well-known/openid-configuration",
|
||||
client_kwargs={"scope": "openid profile email"},
|
||||
)
|
||||
SOCIAL_PROVIDERS["google"] = {"name": "Google", "icon": "fab fa-google", "color": "red"}
|
||||
logger.info("Social login provider registered: Google")
|
||||
else:
|
||||
logger.warning("SOCIAL_AUTH_GOOGLE_ENABLED=true but client ID/secret not configured")
|
||||
|
||||
if AUTH_ENABLED and settings.social_auth_microsoft_enabled:
|
||||
if settings.social_auth_microsoft_client_id and settings.social_auth_microsoft_client_secret:
|
||||
tenant = settings.social_auth_microsoft_tenant or "common"
|
||||
oauth.register(
|
||||
name="microsoft",
|
||||
client_id=settings.social_auth_microsoft_client_id,
|
||||
client_secret=settings.social_auth_microsoft_client_secret,
|
||||
server_metadata_url=f"https://login.microsoftonline.com/{tenant}/v2.0/.well-known/openid-configuration",
|
||||
client_kwargs={"scope": "openid profile email"},
|
||||
)
|
||||
SOCIAL_PROVIDERS["microsoft"] = {"name": "Microsoft", "icon": "fab fa-microsoft", "color": "blue"}
|
||||
logger.info("Social login provider registered: Microsoft (tenant=%s)", tenant)
|
||||
else:
|
||||
logger.warning("SOCIAL_AUTH_MICROSOFT_ENABLED=true but client ID/secret not configured")
|
||||
|
||||
if AUTH_ENABLED and settings.social_auth_apple_enabled:
|
||||
if settings.social_auth_apple_client_id and settings.social_auth_apple_team_id:
|
||||
oauth.register(
|
||||
name="apple",
|
||||
client_id=settings.social_auth_apple_client_id,
|
||||
server_metadata_url="https://appleid.apple.com/.well-known/openid-configuration",
|
||||
client_kwargs={
|
||||
"scope": "openid name email",
|
||||
"response_mode": "form_post",
|
||||
},
|
||||
)
|
||||
SOCIAL_PROVIDERS["apple"] = {"name": "Apple", "icon": "fab fa-apple", "color": "gray"}
|
||||
logger.info("Social login provider registered: Apple")
|
||||
else:
|
||||
logger.warning("SOCIAL_AUTH_APPLE_ENABLED=true but client ID/team ID not configured")
|
||||
|
||||
if AUTH_ENABLED and settings.social_auth_dropbox_enabled:
|
||||
if settings.social_auth_dropbox_client_id and settings.social_auth_dropbox_client_secret:
|
||||
oauth.register(
|
||||
name="dropbox",
|
||||
client_id=settings.social_auth_dropbox_client_id,
|
||||
client_secret=settings.social_auth_dropbox_client_secret,
|
||||
authorize_url="https://www.dropbox.com/oauth2/authorize",
|
||||
access_token_url="https://api.dropboxapi.com/oauth2/token",
|
||||
userinfo_endpoint="https://api.dropboxapi.com/2/users/get_current_account",
|
||||
client_kwargs={"token_endpoint_auth_method": "client_secret_post"},
|
||||
)
|
||||
SOCIAL_PROVIDERS["dropbox"] = {"name": "Dropbox", "icon": "fab fa-dropbox", "color": "blue"}
|
||||
logger.info("Social login provider registered: Dropbox")
|
||||
else:
|
||||
logger.warning("SOCIAL_AUTH_DROPBOX_ENABLED=true but client ID/secret not configured")
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def get_current_user(request: Request):
|
||||
# Check for Bearer token auth first (API tokens)
|
||||
api_user = getattr(request.state, "api_token_user", None)
|
||||
if isinstance(api_user, dict):
|
||||
return api_user
|
||||
return request.session.get("user")
|
||||
|
||||
|
||||
def _resolve_bearer_user(request: Request, db: Session) -> dict | None:
|
||||
"""Resolve a user from a Bearer API token in the Authorization header.
|
||||
|
||||
If the header is present and the token is valid, updates usage tracking
|
||||
(last_used_at, last_used_ip) and returns a synthetic user dict compatible
|
||||
with the session user format.
|
||||
|
||||
Returns:
|
||||
A user dict or ``None`` if no valid Bearer token is present.
|
||||
"""
|
||||
auth_header = request.headers.get("authorization", "")
|
||||
if not isinstance(auth_header, str) or not auth_header.startswith("Bearer "):
|
||||
return None
|
||||
|
||||
raw_token = auth_header[7:]
|
||||
if not raw_token or not isinstance(raw_token, str):
|
||||
return None
|
||||
|
||||
from app.api.api_tokens import hash_token
|
||||
from app.models import ApiToken
|
||||
|
||||
token_hash = hash_token(raw_token)
|
||||
db_token = db.query(ApiToken).filter(ApiToken.token_hash == token_hash, ApiToken.is_active.is_(True)).first()
|
||||
if db_token is None:
|
||||
return None
|
||||
|
||||
# Update usage tracking
|
||||
try:
|
||||
db_token.last_used_at = datetime.now(timezone.utc)
|
||||
# Extract client IP (respect X-Forwarded-For from reverse proxy)
|
||||
client_ip = request.headers.get("x-forwarded-for", "").split(",")[0].strip()
|
||||
if not client_ip and request.client:
|
||||
client_ip = request.client.host
|
||||
db_token.last_used_ip = client_ip or None
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.debug("Failed to update API token usage tracking for token_id=%s", db_token.id)
|
||||
|
||||
# Build a synthetic user dict that mimics the session user format
|
||||
return {
|
||||
"id": db_token.owner_id,
|
||||
"email": db_token.owner_id,
|
||||
"preferred_username": db_token.owner_id,
|
||||
"is_admin": False,
|
||||
"_api_token_id": db_token.id,
|
||||
}
|
||||
|
||||
|
||||
def get_current_user_id(request: Request) -> str:
|
||||
"""Return a stable string identifier for the authenticated user.
|
||||
|
||||
Falls back to ``"anonymous"`` when no user is in the session (e.g. when
|
||||
AUTH_ENABLED=False in single-user mode). The returned value is consistent
|
||||
with how pipeline and file ownership is stored in the database.
|
||||
|
||||
Priority order: ``preferred_username`` → ``email`` → ``id`` → ``"anonymous"``.
|
||||
|
||||
Args:
|
||||
request: The current FastAPI request.
|
||||
|
||||
Returns:
|
||||
A non-empty string identifying the current user.
|
||||
"""
|
||||
user = get_current_user(request)
|
||||
if not user or not isinstance(user, dict):
|
||||
return "anonymous"
|
||||
return user.get("preferred_username") or user.get("email") or user.get("id") or "anonymous"
|
||||
|
||||
|
||||
def require_login(func):
|
||||
if not AUTH_ENABLED:
|
||||
return func # no-op
|
||||
|
||||
@wraps(func)
|
||||
async def wrapper(request: Request, *args, **kwargs):
|
||||
if not request.session.get("user"):
|
||||
request.session["redirect_after_login"] = str(request.url)
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND)
|
||||
# Check if the wrapped function is a coroutine function
|
||||
if inspect.iscoroutinefunction(func):
|
||||
return await func(request, *args, **kwargs)
|
||||
else:
|
||||
return func(request, *args, **kwargs)
|
||||
# Check session auth first
|
||||
if request.session.get("user"):
|
||||
if inspect.iscoroutinefunction(func):
|
||||
return await func(*args, request=request, **kwargs)
|
||||
else:
|
||||
return func(*args, request=request, **kwargs)
|
||||
|
||||
# Fall back to Bearer token auth for API endpoints
|
||||
url_path = urlparse(str(request.url)).path
|
||||
if url_path.startswith("/api/"):
|
||||
try:
|
||||
from app.database import SessionLocal
|
||||
|
||||
db = SessionLocal()
|
||||
try:
|
||||
api_user = _resolve_bearer_user(request, db)
|
||||
finally:
|
||||
db.close()
|
||||
except Exception:
|
||||
api_user = None
|
||||
|
||||
if api_user:
|
||||
request.state.api_token_user = api_user
|
||||
if inspect.iscoroutinefunction(func):
|
||||
return await func(*args, request=request, **kwargs)
|
||||
else:
|
||||
return func(*args, request=request, **kwargs)
|
||||
|
||||
return JSONResponse(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
content={"error": "Not authenticated"},
|
||||
)
|
||||
|
||||
# Non-API endpoint with no session — redirect to login
|
||||
request.session["redirect_after_login"] = str(request.url)
|
||||
return RedirectResponse(url="/login", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
return wrapper
|
||||
|
||||
@@ -70,7 +254,39 @@ def get_gravatar_url(email):
|
||||
|
||||
|
||||
async def login(request: Request):
|
||||
"""Show login page with appropriate authentication options"""
|
||||
"""Show login page with appropriate authentication options."""
|
||||
# Persist the mobile deep-link redirect URI in the session so it survives
|
||||
# the OAuth provider round-trip and is available when auth completes.
|
||||
# Accepted schemes:
|
||||
# • "docuelevate://" — production / EAS builds (custom app scheme)
|
||||
# • "exp://" — Expo Go development client
|
||||
# Only custom (non-HTTP) schemes are accepted to prevent open-redirect abuse.
|
||||
_MOBILE_ALLOWED_SCHEMES = ("docuelevate://", "exp://")
|
||||
if request.query_params.get("mobile") == "1":
|
||||
redirect_uri = request.query_params.get("redirect_uri", "")
|
||||
logger.debug(
|
||||
"[MOBILE] Login page opened with mobile=1: redirect_uri=%r client_ip=%s",
|
||||
redirect_uri,
|
||||
get_client_ip(request),
|
||||
)
|
||||
if any(redirect_uri.startswith(s) for s in _MOBILE_ALLOWED_SCHEMES):
|
||||
request.session["mobile_redirect_uri"] = redirect_uri
|
||||
logger.info(
|
||||
"[MOBILE] Mobile redirect URI stored in session: %r",
|
||||
redirect_uri,
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"[MOBILE] Rejected redirect_uri with disallowed scheme: %r (allowed: %s)",
|
||||
redirect_uri,
|
||||
", ".join(_MOBILE_ALLOWED_SCHEMES),
|
||||
)
|
||||
else:
|
||||
logger.debug(
|
||||
"[MOBILE] Login page opened without mobile=1 (standard browser flow) client_ip=%s",
|
||||
get_client_ip(request),
|
||||
)
|
||||
|
||||
return templates.TemplateResponse(
|
||||
"login.html",
|
||||
{
|
||||
@@ -79,8 +295,11 @@ async def login(request: Request):
|
||||
"message": request.query_params.get("message"),
|
||||
"show_oauth": OAUTH_CONFIGURED,
|
||||
"oauth_provider_name": OAUTH_PROVIDER_NAME,
|
||||
"app_version": settings.version, # Changed from app_version to version
|
||||
"social_providers": SOCIAL_PROVIDERS,
|
||||
"app_version": settings.version,
|
||||
"csrf_token": getattr(request.state, "csrf_token", ""),
|
||||
# "Create account" link is only shown when multi-user mode AND local signup are both enabled
|
||||
"allow_signup": settings.multi_user_enabled and settings.allow_local_signup,
|
||||
},
|
||||
)
|
||||
|
||||
@@ -94,7 +313,252 @@ async def oauth_login(request: Request):
|
||||
return await oauth.authentik.authorize_redirect(request, redirect_uri)
|
||||
|
||||
|
||||
async def oauth_callback(request: Request):
|
||||
async def social_login(request: Request, provider: str):
|
||||
"""Initiate a social login flow for the given provider.
|
||||
|
||||
Args:
|
||||
request: The current FastAPI request.
|
||||
provider: One of the registered social provider keys (google, microsoft, apple, dropbox).
|
||||
|
||||
Returns:
|
||||
A redirect to the provider's authorization page, or back to /login on error.
|
||||
"""
|
||||
if provider not in SOCIAL_PROVIDERS:
|
||||
return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
redirect_uri = request.url_for("social_callback", provider=provider)
|
||||
oauth_client = getattr(oauth, provider, None)
|
||||
if oauth_client is None:
|
||||
return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
return await oauth_client.authorize_redirect(request, redirect_uri)
|
||||
|
||||
|
||||
def _normalize_social_userinfo(provider: str, token: dict, raw_userinfo: dict | None) -> dict:
|
||||
"""Normalize the userinfo payload from different social providers into a common format.
|
||||
|
||||
Returns a dict with keys: sub, email, name, preferred_username, picture.
|
||||
|
||||
Args:
|
||||
provider: The social provider key (google, microsoft, apple, dropbox).
|
||||
token: The OAuth token response from the provider. Included for future
|
||||
provider-specific claim extraction (e.g. ``id_token`` claims).
|
||||
raw_userinfo: The raw userinfo dict (may be None for providers without standard OIDC userinfo).
|
||||
|
||||
Returns:
|
||||
A normalized user-data dict compatible with the session user format.
|
||||
"""
|
||||
userinfo: dict = raw_userinfo or {}
|
||||
|
||||
if provider == "dropbox":
|
||||
# Dropbox returns a non-standard userinfo response
|
||||
email = userinfo.get("email", "")
|
||||
name_info = userinfo.get("name", {})
|
||||
display_name = name_info.get("display_name", "") if isinstance(name_info, dict) else str(name_info)
|
||||
return {
|
||||
"sub": userinfo.get("account_id", email),
|
||||
"email": email,
|
||||
"name": display_name,
|
||||
"preferred_username": email,
|
||||
"picture": userinfo.get("profile_photo_url", ""),
|
||||
}
|
||||
|
||||
# Standard OIDC providers (Google, Microsoft, Apple)
|
||||
return {
|
||||
"sub": userinfo.get("sub", ""),
|
||||
"email": userinfo.get("email", ""),
|
||||
"name": userinfo.get("name", ""),
|
||||
"preferred_username": userinfo.get("email", ""),
|
||||
"picture": userinfo.get("picture", ""),
|
||||
}
|
||||
|
||||
|
||||
async def social_callback(request: Request, provider: str, db: Session = Depends(get_db)):
|
||||
"""Handle the OAuth callback from a social login provider.
|
||||
|
||||
After the user authorizes with the social provider, this endpoint exchanges
|
||||
the authorization code for tokens, extracts user information, creates or
|
||||
updates the user profile, and establishes a session.
|
||||
|
||||
Args:
|
||||
request: The current FastAPI request.
|
||||
provider: One of the registered social provider keys.
|
||||
db: Database session (injected).
|
||||
|
||||
Returns:
|
||||
A redirect to the user's original destination or the upload page.
|
||||
"""
|
||||
if provider not in SOCIAL_PROVIDERS:
|
||||
return RedirectResponse(url="/login?error=Unknown+social+provider", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
oauth_client = getattr(oauth, provider, None)
|
||||
if oauth_client is None:
|
||||
return RedirectResponse(url="/login?error=Provider+not+configured", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
try:
|
||||
token = await oauth_client.authorize_access_token(request)
|
||||
|
||||
# Try standard OIDC userinfo first, fall back to token-embedded userinfo
|
||||
raw_userinfo = token.get("userinfo")
|
||||
if not raw_userinfo:
|
||||
try:
|
||||
resp = await oauth_client.userinfo(token=token)
|
||||
raw_userinfo = resp if isinstance(resp, dict) else resp.json() if hasattr(resp, "json") else {}
|
||||
except Exception:
|
||||
raw_userinfo = {}
|
||||
|
||||
user_data = _normalize_social_userinfo(provider, token, raw_userinfo)
|
||||
|
||||
if not user_data.get("email"):
|
||||
return RedirectResponse(
|
||||
url="/login?error=Could+not+retrieve+email+from+provider",
|
||||
status_code=status.HTTP_302_FOUND,
|
||||
)
|
||||
|
||||
# Add Gravatar if no picture provided
|
||||
if not user_data.get("picture") and user_data.get("email"):
|
||||
user_data["picture"] = get_gravatar_url(user_data["email"])
|
||||
|
||||
# Tag the login source for audit/debugging
|
||||
user_data["auth_provider"] = provider
|
||||
|
||||
# Social login users are never admin by default (admin must be granted
|
||||
# via the Authentik/OIDC admin group or manually in the admin panel)
|
||||
user_data["is_admin"] = False
|
||||
|
||||
request.session["user"] = user_data
|
||||
|
||||
# Auto-create or update UserProfile
|
||||
_ensure_user_profile(db, user_data, is_admin=False)
|
||||
|
||||
provider_name = SOCIAL_PROVIDERS[provider]["name"]
|
||||
logger.info(
|
||||
"[SECURITY] SOCIAL_LOGIN_SUCCESS provider=%s user=%s", provider_name, user_data.get("email", "unknown")
|
||||
)
|
||||
|
||||
# Redirect first-time users to onboarding
|
||||
user_id = (
|
||||
user_data.get("sub") or user_data.get("preferred_username") or user_data.get("email") or user_data.get("id")
|
||||
)
|
||||
|
||||
# Mobile app flow: issue an inline API token and redirect back to the app.
|
||||
mobile_resp = _create_mobile_redirect(request, db)
|
||||
if mobile_resp:
|
||||
return mobile_resp
|
||||
|
||||
if user_id:
|
||||
profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first()
|
||||
if profile and not profile.onboarding_completed:
|
||||
post_onboarding = request.session.pop("redirect_after_login", "/upload")
|
||||
request.session["post_onboarding_redirect"] = post_onboarding
|
||||
return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
redirect_url = request.session.pop("redirect_after_login", "/upload")
|
||||
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
|
||||
except Exception as e:
|
||||
logger.warning("[SECURITY] SOCIAL_LOGIN_FAILURE provider=%s error=%s", provider, type(e).__name__)
|
||||
return RedirectResponse(
|
||||
url="/login?error=Social+login+failed.+Please+try+again.", status_code=status.HTTP_302_FOUND
|
||||
)
|
||||
|
||||
|
||||
def _ensure_user_profile(db: Session, user_data: dict, is_admin: bool = False) -> None:
|
||||
"""Create or update a UserProfile row for *user_data*.
|
||||
|
||||
Uses the same identifier priority as ``get_current_owner_id`` (sub →
|
||||
preferred_username → email → id) so that the profile's ``user_id`` matches
|
||||
``FileRecord.owner_id`` for every document the user uploads.
|
||||
|
||||
For regular users, an existing profile is left unchanged so that
|
||||
admin-managed settings (tier, limits, etc.) are preserved across logins.
|
||||
|
||||
For admin users (*is_admin=True*) the following rules apply:
|
||||
- If no profile exists: one is created with the highest subscription tier,
|
||||
``is_complimentary=True``, and ``onboarding_completed=True`` so that
|
||||
admins skip the first-time setup wizard.
|
||||
- If a profile already exists: ``is_complimentary`` is set to ``True``
|
||||
and, when the current tier is ``"free"``, the tier is upgraded to the
|
||||
highest available plan. Other admin-managed settings are left intact.
|
||||
|
||||
Args:
|
||||
db: Active database session.
|
||||
user_data: Mapping of user attributes as returned by the OAuth provider
|
||||
or built by :func:`app.utils.local_auth.build_session_user`.
|
||||
is_admin: When ``True``, apply admin-specific defaults on first login
|
||||
and ensure the complimentary flag is always set.
|
||||
"""
|
||||
from app.models import UserProfile
|
||||
from app.utils.subscription import TIER_ORDER
|
||||
|
||||
highest_tier = TIER_ORDER[-1]
|
||||
|
||||
user_id = (
|
||||
user_data.get("sub") or user_data.get("preferred_username") or user_data.get("email") or user_data.get("id")
|
||||
)
|
||||
if not user_id:
|
||||
logger.warning("Cannot create UserProfile: no stable user identifier in OAuth userinfo")
|
||||
return
|
||||
|
||||
try:
|
||||
existing = db.query(UserProfile).filter(UserProfile.user_id == user_id).first()
|
||||
if existing is None:
|
||||
display_name = user_data.get("name") or user_data.get("preferred_username") or user_data.get("email")
|
||||
email = user_data.get("email")
|
||||
profile = UserProfile(
|
||||
user_id=user_id,
|
||||
display_name=display_name,
|
||||
subscription_tier=highest_tier if is_admin else "free",
|
||||
is_complimentary=is_admin,
|
||||
onboarding_completed=is_admin,
|
||||
)
|
||||
db.add(profile)
|
||||
db.commit()
|
||||
logger.info(
|
||||
"Auto-created UserProfile for user_id=%s (admin=%s, tier=%s)",
|
||||
user_id,
|
||||
is_admin,
|
||||
highest_tier if is_admin else "free",
|
||||
)
|
||||
# Notify admins and fire webhook for new (non-admin) user signup
|
||||
if not is_admin:
|
||||
try:
|
||||
from app.utils.notification import notify_user_signup
|
||||
from app.utils.webhook import dispatch_webhook_event
|
||||
|
||||
notify_user_signup(user_id, display_name=display_name, email=email)
|
||||
dispatch_webhook_event(
|
||||
"user.signup",
|
||||
{
|
||||
"user_id": user_id,
|
||||
"display_name": display_name,
|
||||
"email": email,
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Failed to send signup notification/webhook for user_id=%s", user_id)
|
||||
elif is_admin:
|
||||
# Ensure existing admin profiles always have complimentary flag set.
|
||||
# Also upgrade from free tier to highest if still on default.
|
||||
changed = False
|
||||
if not existing.is_complimentary:
|
||||
existing.is_complimentary = True
|
||||
changed = True
|
||||
if (existing.subscription_tier or "free") == "free":
|
||||
existing.subscription_tier = highest_tier
|
||||
changed = True
|
||||
if changed:
|
||||
db.commit()
|
||||
logger.info(
|
||||
"Updated admin UserProfile for user_id=%s (complimentary=True, tier=%s)",
|
||||
user_id,
|
||||
existing.subscription_tier,
|
||||
)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to auto-create/update UserProfile for user_id=%s", user_id)
|
||||
|
||||
|
||||
async def oauth_callback(request: Request, db: Session = Depends(get_db)):
|
||||
"""Handle OAuth callback from provider"""
|
||||
try:
|
||||
token = await oauth.authentik.authorize_access_token(request)
|
||||
@@ -126,8 +590,42 @@ async def oauth_callback(request: Request):
|
||||
|
||||
request.session["user"] = user_data
|
||||
|
||||
# Auto-create or update UserProfile so the user appears in admin user management
|
||||
_ensure_user_profile(db, user_data, is_admin=is_admin)
|
||||
|
||||
# Log the successful authentication
|
||||
logger.info(f"[SECURITY] OAUTH_LOGIN_SUCCESS user={user_data.get('email', 'unknown')} admin={is_admin}")
|
||||
logger.info("[SECURITY] OAUTH_LOGIN_SUCCESS user=%s admin=%s", user_data.get("email", "unknown"), is_admin)
|
||||
_record_login_event(
|
||||
db,
|
||||
request,
|
||||
user_data.get("email") or user_data.get("preferred_username") or "unknown",
|
||||
success=True,
|
||||
method="oauth",
|
||||
)
|
||||
|
||||
# Redirect first-time users to onboarding
|
||||
user_id = (
|
||||
user_data.get("sub") or user_data.get("preferred_username") or user_data.get("email") or user_data.get("id")
|
||||
)
|
||||
|
||||
# Mobile app flow: issue an inline API token and redirect back to the app.
|
||||
# This check runs before onboarding so native-app users are never sent
|
||||
# to the web-based onboarding wizard.
|
||||
logger.debug(
|
||||
"[MOBILE] oauth_callback: checking for mobile redirect (session has mobile_redirect_uri=%s)",
|
||||
"mobile_redirect_uri" in request.session,
|
||||
)
|
||||
mobile_resp = _create_mobile_redirect(request, db)
|
||||
if mobile_resp:
|
||||
logger.info("[MOBILE] oauth_callback: returning mobile redirect response")
|
||||
return mobile_resp
|
||||
|
||||
if user_id:
|
||||
profile = db.query(_UserProfile).filter(_UserProfile.user_id == user_id).first()
|
||||
if profile and not profile.onboarding_completed:
|
||||
post_onboarding = request.session.pop("redirect_after_login", "/upload")
|
||||
request.session["post_onboarding_redirect"] = post_onboarding
|
||||
return RedirectResponse(url="/onboarding", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
# Redirect to original destination or default
|
||||
redirect_url = request.session.pop("redirect_after_login", "/upload")
|
||||
@@ -137,15 +635,259 @@ async def oauth_callback(request: Request):
|
||||
return RedirectResponse(url=f"/login?error=Authentication+failed:+{str(e)}", status_code=status.HTTP_302_FOUND)
|
||||
|
||||
|
||||
async def auth(request: Request):
|
||||
"""Handle local username/password authentication"""
|
||||
def _record_login_event(
|
||||
db: Session,
|
||||
request: Request,
|
||||
username: str,
|
||||
*,
|
||||
success: bool,
|
||||
method: str = "local",
|
||||
detail: str | None = None,
|
||||
) -> None:
|
||||
"""Write a login or login-failure audit event to the database.
|
||||
|
||||
Failures are silently swallowed so that an audit-service error never
|
||||
prevents a legitimate login or surfaces an unrelated 500 error to the user.
|
||||
|
||||
Args:
|
||||
db: Active database session.
|
||||
request: The current HTTP request (used to extract the client IP).
|
||||
username: The username that attempted authentication.
|
||||
success: ``True`` for a successful login, ``False`` for a failure.
|
||||
method: Authentication method, e.g. ``"local"`` or ``"oauth"``.
|
||||
detail: Optional extra context for failures (e.g. ``"wrong_password"``).
|
||||
"""
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
action = "login" if success else "login.failure"
|
||||
details: dict = {"method": method}
|
||||
if detail:
|
||||
details["reason"] = detail
|
||||
record_event(
|
||||
db,
|
||||
action=action,
|
||||
user=username,
|
||||
resource_type="session",
|
||||
ip_address=get_client_ip(request),
|
||||
details=details,
|
||||
severity="info" if success else "warning",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write login audit event for user=%s", username, exc_info=True)
|
||||
|
||||
|
||||
def _create_mobile_redirect(request: Request, db: Session) -> RedirectResponse | None:
|
||||
"""Generate a mobile API token and return a redirect to the mobile app.
|
||||
|
||||
If ``mobile_redirect_uri`` is stored in the session (set when the login
|
||||
page was opened with ``?mobile=1&redirect_uri=docuelevate://...``), this
|
||||
function creates a long-lived API token, appends it as a ``?token=``
|
||||
query parameter to the redirect URI, and returns the redirect so that
|
||||
``WebBrowser.openAuthSessionAsync`` in the Expo app intercepts the
|
||||
deep link and stores the token.
|
||||
|
||||
Returns ``None`` when the request is not part of a mobile SSO flow.
|
||||
|
||||
Args:
|
||||
request: The current FastAPI request. The ``user`` dict must already
|
||||
be stored in ``request.session`` before calling this function.
|
||||
db: Active database session used to persist the new API token.
|
||||
|
||||
Returns:
|
||||
A ``RedirectResponse`` to the deep-link URI with ``?token=<plaintext>``,
|
||||
or ``None`` if no mobile redirect URI is pending.
|
||||
"""
|
||||
mobile_redirect_uri = request.session.pop("mobile_redirect_uri", None)
|
||||
if not mobile_redirect_uri:
|
||||
logger.debug("[MOBILE] _create_mobile_redirect: no mobile_redirect_uri in session — skipping mobile flow")
|
||||
return None
|
||||
|
||||
logger.info(
|
||||
"[MOBILE] _create_mobile_redirect: mobile flow detected, redirect_uri=%r",
|
||||
mobile_redirect_uri,
|
||||
)
|
||||
|
||||
user = request.session.get("user") or {}
|
||||
owner_id = user.get("sub") or user.get("preferred_username") or user.get("email") or user.get("id")
|
||||
logger.debug(
|
||||
"[MOBILE] Resolving owner_id from session user: sub=%r preferred_username=%r email=%r id=%r → owner_id=%r",
|
||||
user.get("sub"),
|
||||
user.get("preferred_username"),
|
||||
user.get("email"),
|
||||
user.get("id"),
|
||||
owner_id,
|
||||
)
|
||||
if not owner_id:
|
||||
logger.warning("Mobile SSO redirect requested but no owner_id could be resolved from session")
|
||||
return None
|
||||
|
||||
# Lazy imports to avoid circular dependency via app.api.__init__
|
||||
from app.api.api_tokens import generate_api_token, hash_token
|
||||
from app.models import ApiToken as _ApiToken
|
||||
|
||||
plaintext = generate_api_token()
|
||||
token_hash_value = hash_token(plaintext)
|
||||
prefix = plaintext[:12]
|
||||
|
||||
db_token = _ApiToken(
|
||||
owner_id=owner_id,
|
||||
name="Mobile App",
|
||||
token_hash=token_hash_value,
|
||||
token_prefix=prefix,
|
||||
)
|
||||
try:
|
||||
db.add(db_token)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create mobile API token for owner_id=%s", owner_id)
|
||||
return None
|
||||
|
||||
# Safely append the token as a query parameter, preserving any existing params.
|
||||
separator = "&" if "?" in mobile_redirect_uri else "?"
|
||||
redirect_url = f"{mobile_redirect_uri}{separator}{urlencode({'token': plaintext})}"
|
||||
|
||||
# Log the full redirect URL at DEBUG so it's visible when debug logging is enabled.
|
||||
# At INFO level, log a sanitised version (scheme + host only, token prefix only)
|
||||
# so the plaintext token is never written to persistent info logs.
|
||||
parsed = urlparse(redirect_url)
|
||||
existing_params = f"&{parsed.query.replace(f'token={plaintext}', '')}" if parsed.query else ""
|
||||
sanitised_url = (
|
||||
f"{parsed.scheme}://{parsed.netloc}{parsed.path}?token={prefix}…[redacted]{existing_params.rstrip('&')}"
|
||||
)
|
||||
logger.info(
|
||||
"[MOBILE] MOBILE_SSO_TOKEN_ISSUED owner=%s token_id=%s redirect_target=%s",
|
||||
owner_id,
|
||||
db_token.id,
|
||||
sanitised_url,
|
||||
)
|
||||
logger.debug(
|
||||
"[MOBILE] Full redirect URL being sent to client: %s",
|
||||
redirect_url,
|
||||
)
|
||||
return RedirectResponse(url=redirect_url, status_code=status.HTTP_302_FOUND)
|
||||
|
||||
|
||||
async def auth(request: Request, db: Session = Depends(get_db)):
|
||||
"""Handle local username/password authentication.
|
||||
|
||||
In multi-user mode (``MULTI_USER_ENABLED=True``) local registered users are
|
||||
checked first; if no matching LocalUser is found the request falls through to
|
||||
the single admin-credential check so that single-user deployments continue to
|
||||
work without any database involvement.
|
||||
|
||||
In single-user mode (``MULTI_USER_ENABLED=False``, the default) the LocalUser
|
||||
table is never queried — only the configured ADMIN_USERNAME / ADMIN_PASSWORD
|
||||
are accepted, preserving full backward compatibility.
|
||||
"""
|
||||
form_data = await request.form()
|
||||
username = form_data.get("username")
|
||||
password = form_data.get("password")
|
||||
|
||||
if username == settings.admin_username and password == settings.admin_password:
|
||||
# Create user session
|
||||
request.session["user"] = {
|
||||
logger.debug(
|
||||
"[AUTH] Login attempt: username=%r password_provided=%s multi_user_enabled=%s",
|
||||
username,
|
||||
bool(password),
|
||||
settings.multi_user_enabled,
|
||||
)
|
||||
|
||||
# --- LocalUser check (multi-user mode only) ---
|
||||
if settings.multi_user_enabled:
|
||||
if not username:
|
||||
logger.warning(
|
||||
"[AUTH] LOGIN_FAILURE reason=empty_username multi_user_enabled=%s form_keys=%s content_type=%s",
|
||||
settings.multi_user_enabled,
|
||||
list(form_data.keys()),
|
||||
request.headers.get("content-type", "<missing>"),
|
||||
)
|
||||
return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302)
|
||||
|
||||
username_lower = username.lower()
|
||||
local_user = (
|
||||
db.query(_LocalUser)
|
||||
.filter(
|
||||
(func.lower(_LocalUser.username) == username_lower) | (func.lower(_LocalUser.email) == username_lower)
|
||||
)
|
||||
.first()
|
||||
)
|
||||
logger.debug(
|
||||
"[AUTH] LocalUser lookup: username=%r found=%s",
|
||||
username,
|
||||
local_user is not None,
|
||||
)
|
||||
if local_user is not None:
|
||||
if not local_user.is_active:
|
||||
logger.warning(
|
||||
"[SECURITY] LOCAL_LOGIN_UNVERIFIED user=%s is_active=%s",
|
||||
username,
|
||||
local_user.is_active,
|
||||
)
|
||||
_record_login_event(db, request, username, success=False, detail="account_not_verified")
|
||||
return RedirectResponse(
|
||||
url="/login?error=Please+verify+your+email+address+before+logging+in",
|
||||
status_code=302,
|
||||
)
|
||||
pw_ok = _verify_password(password or "", local_user.hashed_password)
|
||||
logger.debug(
|
||||
"[AUTH] Password verification: user=%s ok=%s",
|
||||
username,
|
||||
pw_ok,
|
||||
)
|
||||
if not pw_ok:
|
||||
logger.warning("[SECURITY] LOCAL_LOGIN_FAILURE reason=wrong_password user=%s", username)
|
||||
_record_login_event(db, request, username, success=False, detail="wrong_password")
|
||||
return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302)
|
||||
user_data = _build_session_user(local_user)
|
||||
request.session["user"] = user_data
|
||||
logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", local_user.email)
|
||||
_record_login_event(db, request, local_user.email, success=True)
|
||||
_ensure_user_profile(db, user_data, is_admin=bool(local_user.is_admin))
|
||||
|
||||
# Mobile app flow: issue an inline API token and redirect back to the app.
|
||||
logger.debug(
|
||||
"[MOBILE] local auth: checking for mobile redirect (session has mobile_redirect_uri=%s)",
|
||||
"mobile_redirect_uri" in request.session,
|
||||
)
|
||||
mobile_resp = _create_mobile_redirect(request, db)
|
||||
if mobile_resp:
|
||||
logger.info("[MOBILE] local auth: returning mobile redirect response")
|
||||
return mobile_resp
|
||||
|
||||
profile = db.query(_UserProfile).filter(_UserProfile.user_id == local_user.email).first()
|
||||
if profile and not profile.onboarding_completed:
|
||||
post_onboarding = request.session.pop("redirect_after_login", "/upload")
|
||||
request.session["post_onboarding_redirect"] = post_onboarding
|
||||
return RedirectResponse(url="/onboarding", status_code=302)
|
||||
redirect_url = request.session.pop("redirect_after_login", "/upload")
|
||||
return RedirectResponse(url=redirect_url, status_code=302)
|
||||
|
||||
# Local user not found; fall through to admin credential check below.
|
||||
logger.debug(
|
||||
"[AUTH] No LocalUser matched username=%r; falling through to admin credential check",
|
||||
username,
|
||||
)
|
||||
|
||||
# --- Admin credentials (always available as a fallback / single-user mode) ---
|
||||
# Guard: only attempt the match when credentials are actually configured.
|
||||
# Without this guard, Python's `None == None` would be True when neither
|
||||
# ADMIN_USERNAME nor ADMIN_PASSWORD is set, allowing any request that omits
|
||||
# those form fields to be authenticated as an admin — creating a phantom
|
||||
# "None@local.docuelevate" admin profile with full privileges.
|
||||
admin_configured = bool(settings.admin_username and settings.admin_password)
|
||||
logger.debug(
|
||||
"[AUTH] Admin credential check: admin_configured=%s username_match=%s multi_user_enabled=%s",
|
||||
admin_configured,
|
||||
(username or "").lower() == settings.admin_username.lower() if admin_configured else False,
|
||||
settings.multi_user_enabled,
|
||||
)
|
||||
if (
|
||||
settings.admin_username
|
||||
and settings.admin_password
|
||||
and (username or "").lower() == settings.admin_username.lower()
|
||||
and password == settings.admin_password
|
||||
):
|
||||
admin_user_data = {
|
||||
"id": "admin",
|
||||
"name": "Administrator",
|
||||
"email": f"{username}@local.docuelevate",
|
||||
@@ -153,22 +895,56 @@ async def auth(request: Request):
|
||||
"picture": "/static/images/default-avatar.svg",
|
||||
"is_admin": True,
|
||||
}
|
||||
logger.info(f"[SECURITY] LOCAL_LOGIN_SUCCESS user={username}")
|
||||
# Redirect to original destination or default
|
||||
request.session["user"] = admin_user_data
|
||||
logger.info("[SECURITY] LOCAL_LOGIN_SUCCESS user=%s", username)
|
||||
_record_login_event(db, request, username, success=True)
|
||||
_ensure_user_profile(db, admin_user_data, is_admin=True)
|
||||
|
||||
# Mobile app flow: issue an inline API token and redirect back to the app.
|
||||
logger.debug(
|
||||
"[MOBILE] admin auth: checking for mobile redirect (session has mobile_redirect_uri=%s)",
|
||||
"mobile_redirect_uri" in request.session,
|
||||
)
|
||||
mobile_resp = _create_mobile_redirect(request, db)
|
||||
if mobile_resp:
|
||||
logger.info("[MOBILE] admin auth: returning mobile redirect response")
|
||||
return mobile_resp
|
||||
|
||||
redirect_url = request.session.pop("redirect_after_login", "/upload")
|
||||
return RedirectResponse(url=redirect_url, status_code=302)
|
||||
else:
|
||||
logger.warning(f"[SECURITY] LOCAL_LOGIN_FAILURE user={username}")
|
||||
logger.warning(
|
||||
"[SECURITY] LOCAL_LOGIN_FAILURE reason=no_match user=%r "
|
||||
"multi_user_enabled=%s admin_configured=%s form_empty=%s",
|
||||
username,
|
||||
settings.multi_user_enabled,
|
||||
admin_configured,
|
||||
not username and not password,
|
||||
)
|
||||
_record_login_event(db, request, username or "anonymous", success=False, detail="invalid_credentials")
|
||||
return RedirectResponse(url="/login?error=Invalid+username+or+password", status_code=302)
|
||||
|
||||
|
||||
async def logout(request: Request):
|
||||
async def logout(request: Request, db: Session = Depends(get_db)):
|
||||
"""Handle user logout"""
|
||||
user = request.session.get("user")
|
||||
username = "unknown"
|
||||
if isinstance(user, dict):
|
||||
username = user.get("preferred_username") or user.get("email") or "unknown"
|
||||
logger.info(f"[SECURITY] LOGOUT user={username}")
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="logout",
|
||||
user=username,
|
||||
resource_type="session",
|
||||
ip_address=get_client_ip(request),
|
||||
severity="info",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write logout audit event for user=%s", username, exc_info=True)
|
||||
request.session.pop("user", None)
|
||||
return RedirectResponse(url="/login?message=You+have+been+logged+out+successfully", status_code=302)
|
||||
|
||||
@@ -177,6 +953,8 @@ if AUTH_ENABLED:
|
||||
router.add_api_route("/login", login, methods=["GET"])
|
||||
router.add_api_route("/oauth-login", oauth_login, methods=["GET"])
|
||||
router.add_api_route("/oauth-callback", oauth_callback, methods=["GET"])
|
||||
router.add_api_route("/social-login/{provider}", social_login, methods=["GET"])
|
||||
router.add_api_route("/social-callback/{provider}", social_callback, methods=["GET"])
|
||||
router.add_api_route("/auth", auth, methods=["POST"])
|
||||
router.add_api_route("/logout", logout, methods=["GET"])
|
||||
|
||||
@@ -185,7 +963,7 @@ if AUTH_ENABLED:
|
||||
@require_login
|
||||
async def whoami(request: Request):
|
||||
"""API endpoint to get current user information"""
|
||||
user = request.session.get("user")
|
||||
user = get_current_user(request)
|
||||
return user or {"error": "Not authenticated"}
|
||||
|
||||
|
||||
|
||||
+9
-1
@@ -1,7 +1,7 @@
|
||||
# app/celery_app.py
|
||||
|
||||
from celery import Celery
|
||||
from celery.signals import task_failure
|
||||
from celery.signals import task_failure, worker_ready
|
||||
|
||||
from app.config import settings
|
||||
|
||||
@@ -22,6 +22,14 @@ celery.conf.task_routes = {
|
||||
}
|
||||
|
||||
|
||||
@worker_ready.connect
|
||||
def init_sentry_on_worker_ready(**kwargs):
|
||||
"""Initialise Sentry SDK in the Celery worker process."""
|
||||
from app.utils.sentry import init_sentry
|
||||
|
||||
init_sentry(integrations_extra=["celery"])
|
||||
|
||||
|
||||
@task_failure.connect
|
||||
def task_failure_handler(
|
||||
sender=None, task_id=None, exception=None, args=None, kwargs=None, traceback=None, einfo=None, **kw
|
||||
|
||||
+151
-9
@@ -1,5 +1,7 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import logging
|
||||
|
||||
from celery.schedules import crontab
|
||||
|
||||
# Ensure tasks are loaded
|
||||
@@ -8,36 +10,59 @@ from app import tasks # noqa: F401 - Imports app/tasks.py so Celery can registe
|
||||
# Import the shared Celery instance
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.backup_tasks import cleanup_old_backups, create_backup # noqa: F401
|
||||
from app.tasks.batch_tasks import ( # noqa: F401
|
||||
backfill_missing_metadata,
|
||||
cleanup_temp_files,
|
||||
expire_shared_links,
|
||||
process_new_documents,
|
||||
prune_old_notifications,
|
||||
prune_processing_logs,
|
||||
reprocess_failed_documents,
|
||||
sync_search_index,
|
||||
)
|
||||
from app.tasks.check_credentials import check_credentials
|
||||
from app.tasks.compute_embedding import backfill_missing_embeddings, compute_document_embedding # noqa: F401
|
||||
from app.tasks.convert_to_pdf import convert_to_pdf # noqa: F401
|
||||
from app.tasks.convert_to_pdfa import convert_to_pdfa # noqa: F401
|
||||
from app.tasks.embed_metadata_into_pdf import embed_metadata_into_pdf # noqa: F401
|
||||
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt # noqa: F401
|
||||
from app.tasks.finalize_document_storage import finalize_document_storage # noqa: F401
|
||||
from app.tasks.imap_tasks import pull_all_inboxes # noqa: F401
|
||||
from app.tasks.monitor_stalled_steps import monitor_stalled_steps # noqa: F401
|
||||
|
||||
# **Ensure all tasks are imported before Celery starts**
|
||||
from app.tasks.process_document import process_document # noqa: F401
|
||||
from app.tasks.process_with_azure_document_intelligence import process_with_azure_document_intelligence # noqa: F401
|
||||
from app.tasks.process_with_ocr import process_with_ocr # noqa: F401
|
||||
from app.tasks.refine_text_with_gpt import refine_text_with_gpt # noqa: F401
|
||||
from app.tasks.rotate_pdf_pages import rotate_pdf_pages # noqa: F401
|
||||
from app.tasks.send_to_all import send_to_all_destinations # noqa: F401
|
||||
from app.tasks.subscription_tasks import apply_pending_subscription_changes_all # noqa: F401
|
||||
|
||||
# Import new send tasks
|
||||
from app.tasks.upload_to_dropbox import upload_to_dropbox # noqa: F401
|
||||
from app.tasks.upload_to_email import upload_to_email # noqa: F401
|
||||
from app.tasks.upload_to_ftp import upload_to_ftp # noqa: F401
|
||||
from app.tasks.upload_to_google_drive import upload_to_google_drive # noqa: F401
|
||||
from app.tasks.upload_to_icloud import upload_to_icloud # noqa: F401
|
||||
from app.tasks.upload_to_nextcloud import upload_to_nextcloud # noqa: F401
|
||||
from app.tasks.upload_to_onedrive import upload_to_onedrive # noqa: F401
|
||||
from app.tasks.upload_to_paperless import upload_to_paperless # noqa: F401
|
||||
from app.tasks.upload_to_s3 import upload_to_s3 # noqa: F401
|
||||
from app.tasks.upload_to_sftp import upload_to_sftp # noqa: F401
|
||||
from app.tasks.upload_to_user_integration import upload_to_user_integration # noqa: F401
|
||||
from app.tasks.upload_to_webdav import upload_to_webdav # noqa: F401
|
||||
from app.tasks.upload_with_rclone import send_to_all_rclone_destinations, upload_with_rclone # noqa: F401
|
||||
from app.tasks.uptime_kuma_tasks import ping_uptime_kuma # noqa: F401
|
||||
from app.tasks.watch_folder_tasks import scan_all_watch_folders # noqa: F401
|
||||
from app.tasks.webhook_tasks import deliver_webhook_task # noqa: F401
|
||||
|
||||
# Register the settings reload signal handler so workers pick up config changes
|
||||
from app.utils.settings_sync import register_settings_reload_signal
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
register_settings_reload_signal()
|
||||
|
||||
celery.conf.task_routes = {
|
||||
@@ -54,15 +79,13 @@ def test_task():
|
||||
check_credentials.apply_async(countdown=10) # Run 10 seconds after worker starts
|
||||
|
||||
celery.conf.beat_schedule = {
|
||||
"poll-inboxes-every-minute": (
|
||||
{
|
||||
"task": "app.tasks.imap_tasks.pull_all_inboxes",
|
||||
"schedule": crontab(minute="*/1"), # every 1 minute
|
||||
"options": {"expires": 55}, # Ensure tasks don't pile up
|
||||
}
|
||||
if (settings.imap1_host or settings.imap2_host)
|
||||
else None
|
||||
),
|
||||
# IMAP polling — always enabled because per-user IMAP integrations may
|
||||
# exist in the database even when no system-level IMAP hosts are configured.
|
||||
"poll-inboxes-every-minute": {
|
||||
"task": "app.tasks.imap_tasks.pull_all_inboxes",
|
||||
"schedule": crontab(minute="*/1"), # every 1 minute
|
||||
"options": {"expires": 55}, # Ensure tasks don't pile up
|
||||
},
|
||||
# Add Uptime Kuma ping task if configured
|
||||
"ping-uptime-kuma": (
|
||||
{
|
||||
@@ -91,7 +114,126 @@ celery.conf.beat_schedule = {
|
||||
"schedule": crontab(minute="*/1"), # Every minute
|
||||
"options": {"expires": 55}, # Must complete within 55 seconds
|
||||
},
|
||||
# Watch folder scanning — always enabled because per-user WATCH_FOLDER
|
||||
# integrations may exist in the database even when no system-level watch
|
||||
# folder settings are configured.
|
||||
# Schedule is controlled by WATCH_FOLDER_POLL_INTERVAL (default: 1 minute).
|
||||
"scan-watch-folders": {
|
||||
"task": "app.tasks.watch_folder_tasks.scan_all_watch_folders",
|
||||
"schedule": crontab(minute=f"*/{max(1, settings.watch_folder_poll_interval)}"),
|
||||
"options": {"expires": 55},
|
||||
},
|
||||
# Backfill embeddings for files that were processed before the
|
||||
# embedding pipeline was enabled, or where the embedding task failed.
|
||||
"backfill-missing-embeddings": {
|
||||
"task": "backfill_missing_embeddings",
|
||||
"schedule": crontab(minute="*/5"), # Every 5 minutes
|
||||
"options": {"expires": 240}, # 4 minutes expiry
|
||||
},
|
||||
# Apply scheduled subscription downgrades daily at 00:05 UTC
|
||||
"apply-pending-subscription-changes": {
|
||||
"task": "app.tasks.subscription_tasks.apply_pending_subscription_changes_all",
|
||||
"schedule": crontab(hour="0", minute="5"), # 00:05 UTC daily
|
||||
"options": {"expires": 3600},
|
||||
},
|
||||
# ── Database backup tasks ──────────────────────────────────────────────
|
||||
# Hourly backup (kept for 4 days)
|
||||
"backup-hourly": (
|
||||
{
|
||||
"task": "app.tasks.backup_tasks.create_backup",
|
||||
"schedule": crontab(minute="0"), # top of every hour
|
||||
"kwargs": {"backup_type": "hourly"},
|
||||
"options": {"expires": 3300},
|
||||
}
|
||||
if settings.backup_enabled
|
||||
else None
|
||||
),
|
||||
# Daily backup (kept for 3 weeks) – runs at 02:30 UTC
|
||||
"backup-daily": (
|
||||
{
|
||||
"task": "app.tasks.backup_tasks.create_backup",
|
||||
"schedule": crontab(hour="2", minute="30"),
|
||||
"kwargs": {"backup_type": "daily"},
|
||||
"options": {"expires": 3600},
|
||||
}
|
||||
if settings.backup_enabled
|
||||
else None
|
||||
),
|
||||
# Weekly backup (kept for 13 weeks) – runs every Sunday at 03:00 UTC
|
||||
"backup-weekly": (
|
||||
{
|
||||
"task": "app.tasks.backup_tasks.create_backup",
|
||||
"schedule": crontab(hour="3", minute="0", day_of_week="0"),
|
||||
"kwargs": {"backup_type": "weekly"},
|
||||
"options": {"expires": 3600},
|
||||
}
|
||||
if settings.backup_enabled
|
||||
else None
|
||||
),
|
||||
}
|
||||
|
||||
# Remove None entries from beat_schedule
|
||||
celery.conf.beat_schedule = {k: v for k, v in celery.conf.beat_schedule.items() if v is not None}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Load admin-managed scheduled jobs from the database
|
||||
# ---------------------------------------------------------------------------
|
||||
# These jobs are defined in the ``scheduled_jobs`` table (seeded by
|
||||
# ``app.api.scheduled_jobs.seed_default_scheduled_jobs``) and can be
|
||||
# enabled/disabled and rescheduled via the admin UI at /admin/scheduled-jobs.
|
||||
# The schedule is read once at worker startup; changes take effect after
|
||||
# the worker is restarted.
|
||||
|
||||
|
||||
def _load_db_scheduled_jobs() -> None:
|
||||
"""
|
||||
Extend ``celery.conf.beat_schedule`` with entries from the ``scheduled_jobs``
|
||||
database table.
|
||||
|
||||
Only rows with ``enabled=True`` are added. Rows whose ``name`` key
|
||||
already exists in the static schedule (defined above) are skipped so
|
||||
that hardcoded entries cannot be overridden accidentally.
|
||||
|
||||
Failures are logged as warnings and do not prevent the worker from
|
||||
starting.
|
||||
"""
|
||||
try:
|
||||
from app.database import SessionLocal
|
||||
from app.models import ScheduledJob
|
||||
|
||||
with SessionLocal() as db:
|
||||
jobs = db.query(ScheduledJob).filter(ScheduledJob.enabled.is_(True)).all()
|
||||
|
||||
added = 0
|
||||
for job in jobs:
|
||||
if job.name in celery.conf.beat_schedule:
|
||||
# Static entry takes precedence; skip silently.
|
||||
continue
|
||||
|
||||
if job.schedule_type == "interval" and job.interval_seconds:
|
||||
from celery.schedules import schedule as interval_schedule
|
||||
|
||||
sched = interval_schedule(run_every=job.interval_seconds)
|
||||
else:
|
||||
# Default to cron.
|
||||
sched = crontab(
|
||||
minute=job.cron_minute,
|
||||
hour=job.cron_hour,
|
||||
day_of_week=job.cron_day_of_week,
|
||||
day_of_month=job.cron_day_of_month,
|
||||
month_of_year=job.cron_month_of_year,
|
||||
)
|
||||
|
||||
celery.conf.beat_schedule[job.name] = {
|
||||
"task": job.task_name,
|
||||
"schedule": sched,
|
||||
"options": {"expires": 3600},
|
||||
}
|
||||
added += 1
|
||||
|
||||
logger.info("Loaded %d scheduled job(s) from database into Celery Beat.", added)
|
||||
except Exception as exc:
|
||||
logger.warning("Could not load scheduled jobs from database: %s", exc)
|
||||
|
||||
|
||||
_load_db_scheduled_jobs()
|
||||
|
||||
+680
@@ -0,0 +1,680 @@
|
||||
"""DocuElevate command-line interface.
|
||||
|
||||
Provides a pipe-friendly CLI for scripting and automation against the
|
||||
DocuElevate REST API. Authentication is via personal API tokens (the
|
||||
same tokens managed at ``/api-tokens`` in the web UI).
|
||||
|
||||
Usage::
|
||||
|
||||
docuelevate --url http://my-instance --token de_xxx list
|
||||
DOCUELEVATE_URL=http://my-instance DOCUELEVATE_API_TOKEN=de_xxx docuelevate list
|
||||
|
||||
Commands
|
||||
--------
|
||||
upload Upload one or more local files for processing.
|
||||
download Download a processed (or original) file by ID.
|
||||
search Full-text search across all documents.
|
||||
list List documents with optional filtering.
|
||||
token Sub-commands: create / list / revoke API tokens.
|
||||
"""
|
||||
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import unquote
|
||||
|
||||
import click
|
||||
import requests
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Environment-variable defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
ENV_URL = "DOCUELEVATE_URL"
|
||||
ENV_TOKEN = "DOCUELEVATE_API_TOKEN"
|
||||
|
||||
_DEFAULT_URL = "http://localhost:8000"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _build_headers(token: str) -> dict[str, str]:
|
||||
"""Return Authorization headers for the given API token."""
|
||||
return {"Authorization": f"Bearer {token}"}
|
||||
|
||||
|
||||
def _api(
|
||||
method: str,
|
||||
base_url: str,
|
||||
path: str,
|
||||
token: str,
|
||||
timeout: int = 60,
|
||||
**kwargs: Any,
|
||||
) -> requests.Response:
|
||||
"""Make an authenticated API request and return the response.
|
||||
|
||||
Args:
|
||||
method: HTTP method (GET, POST, DELETE, …).
|
||||
base_url: The base URL of the DocuElevate instance.
|
||||
path: API path starting with ``/``.
|
||||
token: Plaintext API token.
|
||||
timeout: Request timeout in seconds (default: 60).
|
||||
**kwargs: Extra keyword arguments forwarded to :func:`requests.request`.
|
||||
|
||||
Returns:
|
||||
The :class:`requests.Response` object.
|
||||
|
||||
Raises:
|
||||
click.ClickException: On network errors.
|
||||
"""
|
||||
url = base_url.rstrip("/") + path
|
||||
headers = _build_headers(token)
|
||||
try:
|
||||
resp = requests.request(method, url, headers=headers, timeout=timeout, **kwargs)
|
||||
except requests.ConnectionError as exc:
|
||||
raise click.ClickException(f"Could not connect to {base_url}: {exc}") from exc
|
||||
except requests.Timeout as exc:
|
||||
raise click.ClickException(f"Request timed out: {exc}") from exc
|
||||
return resp
|
||||
|
||||
|
||||
def _require_ok(resp: requests.Response) -> dict[str, Any] | list[Any]:
|
||||
"""Assert a successful HTTP response and return parsed JSON.
|
||||
|
||||
Args:
|
||||
resp: The response to check.
|
||||
|
||||
Returns:
|
||||
Parsed JSON payload.
|
||||
|
||||
Raises:
|
||||
click.ClickException: If the response status indicates an error.
|
||||
"""
|
||||
if resp.status_code >= 400:
|
||||
try:
|
||||
detail = resp.json().get("detail", resp.text)
|
||||
except Exception:
|
||||
detail = resp.text
|
||||
raise click.ClickException(f"API error {resp.status_code}: {detail}")
|
||||
try:
|
||||
return resp.json()
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _output(data: Any, fmt: str) -> None:
|
||||
"""Write *data* to stdout in the requested format.
|
||||
|
||||
Args:
|
||||
data: The value to serialise (dict, list, or primitive).
|
||||
fmt: Either ``"json"`` (machine-readable) or ``"table"`` (human-readable).
|
||||
"""
|
||||
if fmt == "json":
|
||||
click.echo(json.dumps(data, indent=2, default=str))
|
||||
else:
|
||||
_print_table(data)
|
||||
|
||||
|
||||
def _print_table(data: Any) -> None:
|
||||
"""Pretty-print a list of dicts as a fixed-width table.
|
||||
|
||||
Falls back to JSON if the data is not a homogeneous list of dicts.
|
||||
|
||||
Args:
|
||||
data: Data to render.
|
||||
"""
|
||||
if isinstance(data, dict):
|
||||
# Single-object output — print as key: value pairs
|
||||
for key, value in data.items():
|
||||
click.echo(f" {key}: {value}")
|
||||
return
|
||||
|
||||
if not isinstance(data, list) or not data:
|
||||
click.echo(json.dumps(data, indent=2, default=str))
|
||||
return
|
||||
|
||||
if not isinstance(data[0], dict):
|
||||
for item in data:
|
||||
click.echo(str(item))
|
||||
return
|
||||
|
||||
# Determine column widths
|
||||
keys = list(data[0].keys())
|
||||
widths: dict[str, int] = {k: len(k) for k in keys}
|
||||
for row in data:
|
||||
for k in keys:
|
||||
widths[k] = max(widths[k], len(str(row.get(k, ""))))
|
||||
|
||||
header = " ".join(k.upper().ljust(widths[k]) for k in keys)
|
||||
separator = " ".join("-" * widths[k] for k in keys)
|
||||
click.echo(header)
|
||||
click.echo(separator)
|
||||
for row in data:
|
||||
click.echo(" ".join(str(row.get(k, "")).ljust(widths[k]) for k in keys))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Root command group
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@click.group(context_settings={"help_option_names": ["-h", "--help"]})
|
||||
@click.option(
|
||||
"--url",
|
||||
envvar=ENV_URL,
|
||||
default=_DEFAULT_URL,
|
||||
show_default=True,
|
||||
show_envvar=True,
|
||||
help="Base URL of the DocuElevate instance.",
|
||||
metavar="URL",
|
||||
)
|
||||
@click.option(
|
||||
"--token",
|
||||
envvar=ENV_TOKEN,
|
||||
default=None,
|
||||
show_envvar=True,
|
||||
help="API token (de_…). Required for all commands except help.",
|
||||
metavar="TOKEN",
|
||||
)
|
||||
@click.option(
|
||||
"--format",
|
||||
"fmt",
|
||||
type=click.Choice(["table", "json"], case_sensitive=False),
|
||||
default="table",
|
||||
show_default=True,
|
||||
help="Output format. Use 'json' for machine-readable / pipe-friendly output.",
|
||||
)
|
||||
@click.option(
|
||||
"--timeout",
|
||||
default=60,
|
||||
show_default=True,
|
||||
envvar="DOCUELEVATE_TIMEOUT",
|
||||
show_envvar=True,
|
||||
type=int,
|
||||
help="HTTP request timeout in seconds.",
|
||||
)
|
||||
@click.version_option(package_name="docuelevate", prog_name="docuelevate")
|
||||
@click.pass_context
|
||||
def cli(ctx: click.Context, url: str, token: str | None, fmt: str, timeout: int) -> None:
|
||||
"""DocuElevate CLI — interact with DocuElevate from the command line.
|
||||
|
||||
Configure the target instance and credentials via options or environment
|
||||
variables:
|
||||
|
||||
\b
|
||||
DOCUELEVATE_URL Base URL of the instance (default: http://localhost:8000)
|
||||
DOCUELEVATE_API_TOKEN Personal API token (de_…)
|
||||
DOCUELEVATE_TIMEOUT HTTP request timeout in seconds (default: 60)
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
# Upload a file
|
||||
docuelevate --token de_xxx upload report.pdf
|
||||
|
||||
\b
|
||||
# List files as JSON for further processing
|
||||
docuelevate --token de_xxx --format json list | jq '.[].original_filename'
|
||||
|
||||
\b
|
||||
# Search for invoices
|
||||
docuelevate --token de_xxx search "invoice amazon"
|
||||
"""
|
||||
ctx.ensure_object(dict)
|
||||
ctx.obj["url"] = url
|
||||
ctx.obj["token"] = token
|
||||
ctx.obj["fmt"] = fmt
|
||||
ctx.obj["timeout"] = timeout
|
||||
|
||||
|
||||
def _get_token(ctx: click.Context) -> str:
|
||||
"""Return the token from context, raising ClickException if absent.
|
||||
|
||||
Args:
|
||||
ctx: The current Click context.
|
||||
|
||||
Returns:
|
||||
The API token string.
|
||||
|
||||
Raises:
|
||||
click.ClickException: If no token has been provided.
|
||||
"""
|
||||
token = ctx.obj.get("token")
|
||||
if not token:
|
||||
raise click.ClickException(f"No API token provided. Use --token or set the {ENV_TOKEN} environment variable.")
|
||||
return token
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# list command
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@cli.command("list")
|
||||
@click.option("--page", default=1, show_default=True, help="Page number.")
|
||||
@click.option("--per-page", default=25, show_default=True, help="Items per page (max 200).")
|
||||
@click.option("--search", default=None, help="Filter by filename substring.")
|
||||
@click.option("--mime-type", default=None, help="Filter by MIME type (e.g. application/pdf).")
|
||||
@click.option("--status", "file_status", default=None, help="Filter by status: pending, processing, completed, failed.")
|
||||
@click.option("--sort-by", default="created_at", show_default=True, help="Sort field.")
|
||||
@click.option("--sort-order", type=click.Choice(["asc", "desc"]), default="desc", show_default=True)
|
||||
@click.pass_context
|
||||
def list_files(
|
||||
ctx: click.Context,
|
||||
page: int,
|
||||
per_page: int,
|
||||
search: str | None,
|
||||
mime_type: str | None,
|
||||
file_status: str | None,
|
||||
sort_by: str,
|
||||
sort_order: str,
|
||||
) -> None:
|
||||
"""List documents stored in DocuElevate.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate list
|
||||
docuelevate list --status completed --per-page 10
|
||||
docuelevate --format json list | jq '.[].original_filename'
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
fmt: str = ctx.obj["fmt"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
params: dict[str, Any] = {
|
||||
"page": page,
|
||||
"per_page": per_page,
|
||||
"sort_by": sort_by,
|
||||
"sort_order": sort_order,
|
||||
}
|
||||
if search:
|
||||
params["search"] = search
|
||||
if mime_type:
|
||||
params["mime_type"] = mime_type
|
||||
if file_status:
|
||||
params["status"] = file_status
|
||||
|
||||
resp = _api("GET", url, "/api/files", token, timeout=timeout, params=params)
|
||||
payload = _require_ok(resp)
|
||||
|
||||
# Extract the list from the paginated response
|
||||
files: list[dict[str, Any]] = payload.get("files", payload) if isinstance(payload, dict) else payload # type: ignore[assignment]
|
||||
pagination: dict[str, Any] = payload.get("pagination", {}) if isinstance(payload, dict) else {}
|
||||
|
||||
if fmt == "json":
|
||||
_output(files, fmt)
|
||||
else:
|
||||
# Trim fields for readable table
|
||||
rows = [
|
||||
{
|
||||
"id": f.get("id"),
|
||||
"filename": f.get("original_filename"),
|
||||
"size": f.get("file_size"),
|
||||
"status": f.get("status"),
|
||||
"created_at": str(f.get("created_at", ""))[:19],
|
||||
}
|
||||
for f in files
|
||||
]
|
||||
_output(rows, fmt)
|
||||
if pagination:
|
||||
click.echo(f"\nPage {pagination.get('page')}/{pagination.get('pages')} ({pagination.get('total')} total)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# upload command
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@cli.command("upload")
|
||||
@click.argument("files", nargs=-1, required=True, type=click.Path(exists=True, readable=True))
|
||||
@click.option(
|
||||
"--batch-size",
|
||||
default=5,
|
||||
show_default=True,
|
||||
help="Maximum number of concurrent uploads (sequential when 1).",
|
||||
)
|
||||
@click.pass_context
|
||||
def upload_files(ctx: click.Context, files: tuple[str, ...], batch_size: int) -> None:
|
||||
"""Upload one or more local files for processing.
|
||||
|
||||
Supports glob patterns and multiple arguments for batch uploads.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate upload report.pdf
|
||||
docuelevate upload *.pdf invoice_*.png
|
||||
docuelevate upload --batch-size 3 /scans/*.pdf
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
fmt: str = ctx.obj["fmt"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
results: list[dict[str, Any]] = []
|
||||
failed = 0
|
||||
|
||||
for i, file_path in enumerate(files, 1):
|
||||
path = Path(file_path)
|
||||
click.echo(f"[{i}/{len(files)}] Uploading {path.name}…", err=True)
|
||||
try:
|
||||
with path.open("rb") as fh:
|
||||
resp = _api(
|
||||
"POST",
|
||||
url,
|
||||
"/api/ui-upload",
|
||||
token,
|
||||
timeout=timeout,
|
||||
files={"file": (path.name, fh)},
|
||||
)
|
||||
if resp.status_code >= 400:
|
||||
try:
|
||||
detail = resp.json().get("detail", resp.text)
|
||||
except Exception:
|
||||
detail = resp.text
|
||||
click.echo(f" ERROR {resp.status_code}: {detail}", err=True)
|
||||
results.append({"file": path.name, "status": "error", "detail": detail})
|
||||
failed += 1
|
||||
else:
|
||||
data = resp.json()
|
||||
results.append({"file": path.name, "status": "queued", **data})
|
||||
click.echo(f" OK task_id={data.get('task_id', '?')}", err=True)
|
||||
except click.ClickException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
click.echo(f" ERROR: {exc}", err=True)
|
||||
results.append({"file": path.name, "status": "error", "detail": str(exc)})
|
||||
failed += 1
|
||||
|
||||
_output(results, fmt)
|
||||
|
||||
if failed:
|
||||
click.echo(f"\n{failed}/{len(files)} upload(s) failed.", err=True)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# download command
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@cli.command("download")
|
||||
@click.argument("file_id", type=int)
|
||||
@click.option(
|
||||
"--output",
|
||||
"-o",
|
||||
default=None,
|
||||
help="Destination file path. Defaults to the server-provided filename in the current directory.",
|
||||
type=click.Path(),
|
||||
)
|
||||
@click.option(
|
||||
"--version",
|
||||
type=click.Choice(["processed", "original"]),
|
||||
default="processed",
|
||||
show_default=True,
|
||||
help="Which version to download.",
|
||||
)
|
||||
@click.pass_context
|
||||
def download_file(ctx: click.Context, file_id: int, output: str | None, version: str) -> None:
|
||||
"""Download a file by its numeric ID.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate download 42
|
||||
docuelevate download 42 --version original -o /tmp/orig.pdf
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
resp = _api(
|
||||
"GET",
|
||||
url,
|
||||
f"/api/files/{file_id}/download",
|
||||
token,
|
||||
timeout=timeout,
|
||||
params={"version": version},
|
||||
stream=True,
|
||||
)
|
||||
_require_ok(resp)
|
||||
|
||||
# Determine output filename
|
||||
if output:
|
||||
dest = Path(output)
|
||||
else:
|
||||
content_disp = resp.headers.get("content-disposition", "")
|
||||
filename = f"file_{file_id}"
|
||||
for raw_part in content_disp.split(";"):
|
||||
clean = raw_part.strip()
|
||||
if clean.startswith("filename="):
|
||||
filename = clean[len("filename=") :].strip('"').strip("'")
|
||||
break
|
||||
if clean.startswith("filename*="):
|
||||
raw = clean[len("filename*=") :]
|
||||
if raw.upper().startswith("UTF-8''"):
|
||||
filename = unquote(raw[7:])
|
||||
break
|
||||
dest = Path(filename)
|
||||
|
||||
with dest.open("wb") as fh:
|
||||
for chunk in resp.iter_content(chunk_size=65536):
|
||||
fh.write(chunk)
|
||||
|
||||
click.echo(f"Downloaded {dest} ({dest.stat().st_size} bytes)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# search command
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@cli.command("search")
|
||||
@click.argument("query")
|
||||
@click.option("--mime-type", default=None, help="Filter by MIME type.")
|
||||
@click.option("--document-type", default=None, help="Filter by document type (e.g. Invoice).")
|
||||
@click.option("--tags", default=None, help="Filter by tag.")
|
||||
@click.option("--language", default=None, help="Filter by language code (e.g. en, de).")
|
||||
@click.option("--page", default=1, show_default=True)
|
||||
@click.option("--per-page", default=20, show_default=True, help="Results per page (max 100).")
|
||||
@click.pass_context
|
||||
def search(
|
||||
ctx: click.Context,
|
||||
query: str,
|
||||
mime_type: str | None,
|
||||
document_type: str | None,
|
||||
tags: str | None,
|
||||
language: str | None,
|
||||
page: int,
|
||||
per_page: int,
|
||||
) -> None:
|
||||
"""Full-text search across all documents.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate search "invoice amazon"
|
||||
docuelevate search "contract" --document-type Contract --language en
|
||||
docuelevate --format json search "receipt" | jq '.[].file_id'
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
fmt: str = ctx.obj["fmt"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
params: dict[str, Any] = {"q": query, "page": page, "per_page": per_page}
|
||||
if mime_type:
|
||||
params["mime_type"] = mime_type
|
||||
if document_type:
|
||||
params["document_type"] = document_type
|
||||
if tags:
|
||||
params["tags"] = tags
|
||||
if language:
|
||||
params["language"] = language
|
||||
|
||||
resp = _api("GET", url, "/api/search", token, timeout=timeout, params=params)
|
||||
payload = _require_ok(resp)
|
||||
|
||||
results: list[dict[str, Any]] = (
|
||||
payload.get("results", payload) if isinstance(payload, dict) else payload # type: ignore[assignment]
|
||||
)
|
||||
total: int = payload.get("total", len(results)) if isinstance(payload, dict) else len(results)
|
||||
pages: int = payload.get("pages", 1) if isinstance(payload, dict) else 1
|
||||
|
||||
if fmt == "json":
|
||||
_output(results, fmt)
|
||||
else:
|
||||
rows = [
|
||||
{
|
||||
"file_id": r.get("file_id"),
|
||||
"filename": r.get("original_filename"),
|
||||
"type": r.get("document_type"),
|
||||
"tags": ",".join(r.get("tags") or []),
|
||||
}
|
||||
for r in results
|
||||
]
|
||||
_output(rows, fmt)
|
||||
click.echo(f"\nPage {page}/{pages} ({total} total results)")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# token sub-group
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@cli.group("token")
|
||||
@click.pass_context
|
||||
def token_group(ctx: click.Context) -> None:
|
||||
"""Manage personal API tokens.
|
||||
|
||||
Tokens can be created, listed, and revoked. Token rotation is achieved
|
||||
by creating a new token before revoking the old one.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate token create "CI Pipeline"
|
||||
docuelevate token list
|
||||
docuelevate token revoke 3
|
||||
"""
|
||||
|
||||
|
||||
@token_group.command("create")
|
||||
@click.argument("name")
|
||||
@click.pass_context
|
||||
def token_create(ctx: click.Context, name: str) -> None:
|
||||
"""Create a new personal API token.
|
||||
|
||||
The full token value is printed exactly once. Store it securely.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate token create "My script"
|
||||
docuelevate --format json token create "CI" | jq -r '.token'
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
fmt: str = ctx.obj["fmt"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
resp = _api("POST", url, "/api/api-tokens/", token, timeout=timeout, json={"name": name})
|
||||
payload = _require_ok(resp)
|
||||
|
||||
if fmt == "json":
|
||||
_output(payload, fmt)
|
||||
else:
|
||||
if not isinstance(payload, dict):
|
||||
raise click.ClickException("Unexpected API response format.")
|
||||
click.echo("Token created successfully:")
|
||||
click.echo(f" ID: {payload.get('id')}")
|
||||
click.echo(f" Name: {payload.get('name')}")
|
||||
click.echo(f" Prefix: {payload.get('token_prefix')}")
|
||||
click.echo(f" Token: {payload.get('token')}")
|
||||
click.echo()
|
||||
click.echo("Store this token securely — it will not be shown again.", err=True)
|
||||
|
||||
|
||||
@token_group.command("list")
|
||||
@click.pass_context
|
||||
def token_list(ctx: click.Context) -> None:
|
||||
"""List all your API tokens (active and revoked).
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate token list
|
||||
docuelevate --format json token list | jq '.[] | select(.is_active)'
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
fmt: str = ctx.obj["fmt"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
resp = _api("GET", url, "/api/api-tokens/", token, timeout=timeout)
|
||||
payload = _require_ok(resp)
|
||||
|
||||
if fmt == "json":
|
||||
_output(payload, fmt)
|
||||
else:
|
||||
if not isinstance(payload, list):
|
||||
raise click.ClickException("Unexpected API response format.")
|
||||
rows = [
|
||||
{
|
||||
"id": t.get("id"),
|
||||
"name": t.get("name"),
|
||||
"prefix": t.get("token_prefix"),
|
||||
"active": t.get("is_active"),
|
||||
"last_used": str(t.get("last_used_at") or "never")[:19],
|
||||
"created": str(t.get("created_at") or "")[:19],
|
||||
}
|
||||
for t in payload
|
||||
]
|
||||
_output(rows, fmt)
|
||||
|
||||
|
||||
@token_group.command("revoke")
|
||||
@click.argument("token_id", type=int)
|
||||
@click.option("--yes", "-y", is_flag=True, help="Skip confirmation prompt.")
|
||||
@click.pass_context
|
||||
def token_revoke(ctx: click.Context, token_id: int, yes: bool) -> None:
|
||||
"""Revoke an API token by its numeric ID.
|
||||
|
||||
The token is soft-deleted (kept for audit) but immediately invalidated.
|
||||
|
||||
Examples:
|
||||
|
||||
\b
|
||||
docuelevate token revoke 3
|
||||
docuelevate token revoke 3 --yes
|
||||
"""
|
||||
token = _get_token(ctx)
|
||||
url: str = ctx.obj["url"]
|
||||
timeout: int = ctx.obj["timeout"]
|
||||
|
||||
if not yes:
|
||||
click.confirm(f"Revoke token {token_id}?", abort=True)
|
||||
|
||||
resp = _api("DELETE", url, f"/api/api-tokens/{token_id}", token, timeout=timeout)
|
||||
_require_ok(resp)
|
||||
click.echo(f"Token {token_id} revoked.")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Entry point
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Entry point for the ``docuelevate`` console script."""
|
||||
cli(auto_envvar_prefix="DOCUELEVATE") # type: ignore[call-arg]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+714
-2
@@ -1,5 +1,6 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import json
|
||||
import os
|
||||
from typing import Any, List, Optional, Union
|
||||
|
||||
@@ -48,18 +49,30 @@ class Settings(BaseSettings):
|
||||
debug: bool = False # Default to False
|
||||
|
||||
# Making Dropbox optional
|
||||
dropbox_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable Dropbox as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
dropbox_app_key: Optional[str] = None
|
||||
dropbox_app_secret: Optional[str] = None
|
||||
dropbox_folder: Optional[str] = None
|
||||
dropbox_refresh_token: Optional[str] = None
|
||||
|
||||
# Making Nextcloud optional
|
||||
nextcloud_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable Nextcloud as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
nextcloud_upload_url: Optional[str] = None
|
||||
nextcloud_username: Optional[str] = None
|
||||
nextcloud_password: Optional[str] = None
|
||||
nextcloud_folder: Optional[str] = None
|
||||
|
||||
# Making Paperless optional
|
||||
paperless_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable Paperless-ngx as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
paperless_ngx_api_token: Optional[str] = None
|
||||
paperless_host: Optional[str] = None
|
||||
paperless_custom_field_absender: Optional[str] = None # Name of the "absender" custom field in Paperless
|
||||
@@ -115,12 +128,297 @@ class Settings(BaseSettings):
|
||||
session_secret: Optional[str] = None
|
||||
admin_group_name: str = "admin"
|
||||
|
||||
# Authentik
|
||||
# Multi-user settings
|
||||
multi_user_enabled: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Enable multi-user mode with individual document spaces per user. "
|
||||
"When enabled, each authenticated user sees only their own documents, "
|
||||
"uploads, and search results. Shared settings (AI, OCR) remain global. "
|
||||
"Requires auth_enabled=True. Default: False (single-user/shared mode)."
|
||||
),
|
||||
)
|
||||
default_daily_upload_limit: int = Field(
|
||||
default=0,
|
||||
description=(
|
||||
"Default maximum number of document uploads allowed per user per day "
|
||||
"in multi-user mode. Set to 0 for unlimited. "
|
||||
"Individual user limits can override this default. Default: 0 (unlimited)."
|
||||
),
|
||||
)
|
||||
unowned_docs_visible_to_all: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"In multi-user mode, controls whether documents without an owner (owner_id is NULL) "
|
||||
"are visible to all authenticated users. When True, unowned documents appear in every "
|
||||
"user's file list alongside their own files. When False, only admins can see unowned "
|
||||
"documents. Default: True."
|
||||
),
|
||||
)
|
||||
default_owner_id: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"When set, automatically assigns this owner ID to newly ingested documents that would "
|
||||
"otherwise have no owner (e.g. documents from IMAP, API without session, or legacy imports). "
|
||||
"Use the admin /api/files/assign-owner endpoint to bulk-assign existing unclaimed documents. "
|
||||
"Default: None (documents remain unowned until claimed)."
|
||||
),
|
||||
)
|
||||
|
||||
subscription_overage_percent: int = Field(
|
||||
default=20,
|
||||
ge=0,
|
||||
le=200,
|
||||
description=(
|
||||
"Soft-limit overage buffer in percent (0–200). The announced monthly quota is "
|
||||
"increased by this percentage for actual enforcement. E.g. 20 means a 150-doc/month "
|
||||
"plan enforces at 180 docs (150 × 1.20). Set 0 to enforce exactly at the announced "
|
||||
"limit. Per-plan overage_percent (set in Plan Designer) overrides this global default. "
|
||||
"Default: 20."
|
||||
),
|
||||
)
|
||||
|
||||
# Authentik / Generic OIDC
|
||||
authentik_client_id: Optional[str] = None
|
||||
authentik_client_secret: Optional[str] = None
|
||||
authentik_config_url: Optional[str] = None
|
||||
oauth_provider_name: Optional[str] = None # Name to display for the OAuth provider
|
||||
|
||||
# Social Login Providers
|
||||
# Google OAuth2
|
||||
social_auth_google_enabled: bool = False
|
||||
social_auth_google_client_id: Optional[str] = None
|
||||
social_auth_google_client_secret: Optional[str] = None
|
||||
|
||||
# Microsoft OAuth2 (Azure AD / Microsoft Entra ID)
|
||||
social_auth_microsoft_enabled: bool = False
|
||||
social_auth_microsoft_client_id: Optional[str] = None
|
||||
social_auth_microsoft_client_secret: Optional[str] = None
|
||||
social_auth_microsoft_tenant: str = Field(
|
||||
default="common",
|
||||
description=(
|
||||
"Azure AD tenant ID or one of 'common', 'organizations', 'consumers'. "
|
||||
"Use 'common' to allow any Microsoft account and any Azure AD org. "
|
||||
"Use a specific tenant ID (GUID) to restrict to a single organization. "
|
||||
"Default: common."
|
||||
),
|
||||
)
|
||||
|
||||
# Apple Sign-In
|
||||
social_auth_apple_enabled: bool = False
|
||||
social_auth_apple_client_id: Optional[str] = None
|
||||
social_auth_apple_team_id: Optional[str] = None
|
||||
social_auth_apple_key_id: Optional[str] = None
|
||||
social_auth_apple_private_key: Optional[str] = None
|
||||
|
||||
# Dropbox OAuth2
|
||||
social_auth_dropbox_enabled: bool = False
|
||||
social_auth_dropbox_client_id: Optional[str] = None
|
||||
social_auth_dropbox_client_secret: Optional[str] = None
|
||||
|
||||
# Local user signup
|
||||
allow_local_signup: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Allow users to self-register with email and password. "
|
||||
"Has no effect unless MULTI_USER_ENABLED is also True. "
|
||||
"Requires SMTP to be configured so verification emails can be sent. "
|
||||
"Default: False (registration disabled — admin creates users manually)."
|
||||
),
|
||||
)
|
||||
|
||||
# Stripe billing
|
||||
stripe_secret_key: Optional[str] = None
|
||||
stripe_publishable_key: Optional[str] = None
|
||||
stripe_webhook_secret: Optional[str] = None
|
||||
stripe_success_url: Optional[str] = None # e.g. https://app.example.com/billing/success
|
||||
stripe_cancel_url: Optional[str] = None # e.g. https://app.example.com/pricing
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Watch Folder Ingestion
|
||||
# ---------------------------------------------------------------------------
|
||||
# Local filesystem watch folders (comma-separated list of absolute paths).
|
||||
# DocuElevate will poll each path for new files and automatically ingest them.
|
||||
# Works with any mounted path, including SMB/CIFS (via system mount), NFS, etc.
|
||||
# Example: /watchfolders/scanner,/mnt/shared/inbox
|
||||
watch_folders: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Comma-separated list of local filesystem paths (absolute) that DocuElevate will "
|
||||
"poll for new files to ingest. Each file found is enqueued for document processing. "
|
||||
"Works with any mounted path including SMB/CIFS (mounted via system) and NFS. "
|
||||
"Example: /watchfolders/scanner,/mnt/shared/inbox"
|
||||
),
|
||||
)
|
||||
watch_folder_poll_interval: int = Field(
|
||||
default=1,
|
||||
description=("Poll interval in minutes for local watch folder scanning. Default: 1 minute."),
|
||||
)
|
||||
watch_folder_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Delete files from local watch folders after they have been successfully enqueued "
|
||||
"for processing. When False (default), files are left in place and tracked via a "
|
||||
"cache file to avoid re-ingesting them."
|
||||
),
|
||||
)
|
||||
|
||||
# FTP Ingest / Watch Folder
|
||||
# Uses the existing FTP credentials (ftp_host, ftp_username, ftp_password) to poll
|
||||
# a source folder on the FTP server for new files to ingest.
|
||||
ftp_ingest_folder: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"FTP folder path to monitor for new files to ingest. "
|
||||
"Uses the existing FTP connection settings (FTP_HOST, FTP_USERNAME, FTP_PASSWORD). "
|
||||
"When set, DocuElevate will periodically poll this folder and download new files for processing."
|
||||
),
|
||||
)
|
||||
ftp_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable FTP watch folder ingestion. Requires FTP_INGEST_FOLDER and FTP connection settings.",
|
||||
)
|
||||
ftp_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Delete files from the FTP ingest folder after they have been successfully downloaded "
|
||||
"and enqueued for processing. Default: False (files are left in place)."
|
||||
),
|
||||
)
|
||||
|
||||
# SFTP Ingest / Watch Folder
|
||||
# Uses the existing SFTP credentials (sftp_host, sftp_username, sftp_password/sftp_private_key)
|
||||
# to poll a source folder on the SFTP server for new files to ingest.
|
||||
sftp_ingest_folder: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"SFTP folder path to monitor for new files to ingest. "
|
||||
"Uses the existing SFTP connection settings (SFTP_HOST, SFTP_USERNAME, SFTP_PASSWORD/SFTP_PRIVATE_KEY). "
|
||||
"When set, DocuElevate will periodically poll this folder and download new files for processing."
|
||||
),
|
||||
)
|
||||
sftp_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable SFTP watch folder ingestion. Requires SFTP_INGEST_FOLDER and SFTP connection settings.",
|
||||
)
|
||||
sftp_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Delete files from the SFTP ingest folder after they have been successfully downloaded "
|
||||
"and enqueued for processing. Default: False (files are left in place)."
|
||||
),
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cloud Provider Watch Folders
|
||||
# ---------------------------------------------------------------------------
|
||||
# Each cloud provider has three settings:
|
||||
# <provider>_ingest_enabled — enable the watch-folder for this provider
|
||||
# <provider>_ingest_folder — the remote path / folder ID to poll
|
||||
# <provider>_ingest_delete_after_process — delete from cloud after download
|
||||
|
||||
# Dropbox ingest — reuses existing Dropbox OAuth credentials
|
||||
dropbox_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable Dropbox watch folder ingestion. Requires Dropbox OAuth credentials.",
|
||||
)
|
||||
dropbox_ingest_folder: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Dropbox folder path to poll for new files to ingest (e.g. /Inbox/Scanner). "
|
||||
"Uses the existing Dropbox OAuth credentials."
|
||||
),
|
||||
)
|
||||
dropbox_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description="Delete files from Dropbox ingest folder after download and enqueue.",
|
||||
)
|
||||
|
||||
# Google Drive ingest — reuses existing Google Drive credentials
|
||||
google_drive_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable Google Drive watch folder ingestion. Requires Google Drive credentials.",
|
||||
)
|
||||
google_drive_ingest_folder_id: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Google Drive folder ID to poll for new files to ingest. "
|
||||
"Uses the existing Google Drive service-account or OAuth credentials."
|
||||
),
|
||||
)
|
||||
google_drive_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description="Delete files from Google Drive ingest folder after download and enqueue.",
|
||||
)
|
||||
|
||||
# OneDrive ingest — reuses existing OneDrive MSAL credentials
|
||||
onedrive_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable OneDrive watch folder ingestion. Requires OneDrive MSAL credentials.",
|
||||
)
|
||||
onedrive_ingest_folder_path: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"OneDrive folder path to poll for new files to ingest (e.g. /Inbox/Scanner). "
|
||||
"Uses the existing OneDrive client credentials."
|
||||
),
|
||||
)
|
||||
onedrive_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description="Delete files from OneDrive ingest folder after download and enqueue.",
|
||||
)
|
||||
|
||||
# Nextcloud ingest — reuses existing Nextcloud WebDAV credentials
|
||||
nextcloud_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable Nextcloud watch folder ingestion. Requires Nextcloud WebDAV credentials.",
|
||||
)
|
||||
nextcloud_ingest_folder: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Nextcloud folder path to poll for new files to ingest (e.g. /Scans/Inbox). "
|
||||
"Uses the existing Nextcloud upload URL and credentials."
|
||||
),
|
||||
)
|
||||
nextcloud_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description="Delete files from Nextcloud ingest folder after download and enqueue.",
|
||||
)
|
||||
|
||||
# S3 ingest — reuses existing AWS/S3 credentials
|
||||
s3_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable Amazon S3 watch folder (prefix) ingestion. Requires S3 credentials.",
|
||||
)
|
||||
s3_ingest_prefix: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"S3 key prefix to poll for new objects to ingest (e.g. inbox/scanner/). "
|
||||
"Uses the existing S3 bucket and AWS credentials."
|
||||
),
|
||||
)
|
||||
s3_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description="Delete objects from S3 ingest prefix after download and enqueue.",
|
||||
)
|
||||
|
||||
# WebDAV ingest — reuses existing WebDAV credentials
|
||||
webdav_ingest_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable WebDAV watch folder ingestion. Requires WebDAV URL and credentials.",
|
||||
)
|
||||
webdav_ingest_folder: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"WebDAV folder path to poll for new files to ingest (e.g. /remote.php/webdav/Inbox). "
|
||||
"Uses the existing WebDAV URL and credentials."
|
||||
),
|
||||
)
|
||||
webdav_ingest_delete_after_process: bool = Field(
|
||||
default=False,
|
||||
description="Delete files from WebDAV ingest folder after download and enqueue.",
|
||||
)
|
||||
|
||||
# IMAP 1
|
||||
imap1_host: Optional[str] = None
|
||||
imap1_port: Optional[int] = 993
|
||||
@@ -140,6 +438,10 @@ class Settings(BaseSettings):
|
||||
imap2_delete_after_process: bool = False
|
||||
|
||||
# Google Drive settings
|
||||
google_drive_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable Google Drive as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
google_drive_credentials_json: Optional[str] = ""
|
||||
google_drive_folder_id: Optional[str] = ""
|
||||
google_drive_delegate_to: Optional[str] = "" # Optional delegated user email
|
||||
@@ -151,6 +453,10 @@ class Settings(BaseSettings):
|
||||
google_drive_refresh_token: Optional[str] = ""
|
||||
|
||||
# WebDAV settings
|
||||
webdav_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable WebDAV as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
webdav_url: Optional[str] = None
|
||||
webdav_username: Optional[str] = None
|
||||
webdav_password: Optional[str] = None
|
||||
@@ -158,6 +464,10 @@ class Settings(BaseSettings):
|
||||
webdav_verify_ssl: bool = True
|
||||
|
||||
# FTP settings
|
||||
ftp_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable FTP as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
ftp_host: Optional[str] = None
|
||||
ftp_port: Optional[int] = 21
|
||||
ftp_username: Optional[str] = None
|
||||
@@ -167,6 +477,10 @@ class Settings(BaseSettings):
|
||||
ftp_allow_plaintext: bool = True # Default to allowing plaintext fallback
|
||||
|
||||
# SFTP settings
|
||||
sftp_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable SFTP as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
sftp_host: Optional[str] = None
|
||||
sftp_port: Optional[int] = 22
|
||||
sftp_username: Optional[str] = None
|
||||
@@ -178,7 +492,7 @@ class Settings(BaseSettings):
|
||||
# In development/testing, set to True to disable verification (not recommended)
|
||||
sftp_disable_host_key_verification: bool = False # Default enforces host key verification
|
||||
|
||||
# Email settings
|
||||
# Email settings (shared SMTP – used for password reset, verification emails, etc.)
|
||||
email_host: Optional[str] = None
|
||||
email_port: Optional[int] = 587
|
||||
email_username: Optional[str] = None
|
||||
@@ -187,7 +501,24 @@ class Settings(BaseSettings):
|
||||
email_sender: Optional[str] = None # From address, defaults to email_username if not set
|
||||
email_default_recipient: Optional[str] = None
|
||||
|
||||
# Email destination settings (dedicated SMTP for document delivery – decoupled from shared email above)
|
||||
dest_email_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable Email as an upload destination. Set to False to disable document delivery via email even when credentials are configured.",
|
||||
)
|
||||
dest_email_host: Optional[str] = None
|
||||
dest_email_port: Optional[int] = 587
|
||||
dest_email_username: Optional[str] = None
|
||||
dest_email_password: Optional[str] = None
|
||||
dest_email_use_tls: bool = True
|
||||
dest_email_sender: Optional[str] = None # From address for delivered documents
|
||||
dest_email_default_recipient: Optional[str] = None # Fallback recipient for document delivery
|
||||
|
||||
# OneDrive settings
|
||||
onedrive_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable OneDrive as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
onedrive_client_id: Optional[str] = None
|
||||
onedrive_client_secret: Optional[str] = None
|
||||
onedrive_tenant_id: Optional[str] = "common" # Default to "common" for personal accounts
|
||||
@@ -195,6 +526,10 @@ class Settings(BaseSettings):
|
||||
onedrive_folder_path: Optional[str] = None
|
||||
|
||||
# AWS S3 settings
|
||||
s3_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable Amazon S3 as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
aws_access_key_id: Optional[str] = None
|
||||
aws_secret_access_key: Optional[str] = None
|
||||
aws_region: Optional[str] = "us-east-1" # Default region
|
||||
@@ -203,6 +538,16 @@ class Settings(BaseSettings):
|
||||
s3_storage_class: Optional[str] = "STANDARD" # Default storage class
|
||||
s3_acl: Optional[str] = "private" # Default ACL
|
||||
|
||||
# iCloud Drive settings
|
||||
icloud_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable iCloud Drive as an upload destination. Set to False to disable uploads even when credentials are configured.",
|
||||
)
|
||||
icloud_username: Optional[str] = None # Apple ID email address
|
||||
icloud_password: Optional[str] = None # App-specific password (required for 2FA accounts)
|
||||
icloud_folder: Optional[str] = None # Target folder path in iCloud Drive (e.g. "Documents/Uploads")
|
||||
icloud_cookie_directory: Optional[str] = None # Directory for session cookies (default: ~/.pyicloud)
|
||||
|
||||
# Uptime Kuma settings
|
||||
uptime_kuma_url: Optional[str] = None
|
||||
uptime_kuma_ping_interval: int = 5 # Default ping interval in minutes
|
||||
@@ -221,6 +566,77 @@ class Settings(BaseSettings):
|
||||
|
||||
# Feature flags
|
||||
allow_file_delete: bool = True # Default to allowing file deletion from database
|
||||
compliance_enabled: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Enable the compliance templates dashboard (GDPR, HIPAA, SOC 2). "
|
||||
"When enabled, admins can view compliance status and apply "
|
||||
"pre-built regulatory configurations. Default: True."
|
||||
),
|
||||
)
|
||||
|
||||
# PDF/A archival conversion settings
|
||||
enable_pdfa_conversion: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Enable PDF/A archival variant generation. When enabled, PDF/A copies of both the "
|
||||
"original ingested file and the processed file are created and saved alongside the "
|
||||
"standard copies. Uses ocrmypdf with Ghostscript for the conversion. "
|
||||
"This may double or triple storage but provides better legal coverage. Default: False."
|
||||
),
|
||||
)
|
||||
pdfa_format: str = Field(
|
||||
default="2",
|
||||
description=(
|
||||
"PDF/A format variant to produce. Passed to ocrmypdf --output-type pdfa-N. "
|
||||
"Valid values: '1' (PDF/A-1b), '2' (PDF/A-2b), '3' (PDF/A-3b). Default: '2'."
|
||||
),
|
||||
)
|
||||
pdfa_upload_original: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Upload the original-file PDF/A variant to all configured storage providers. "
|
||||
"Files are placed in the provider's folder + PDFA_UPLOAD_FOLDER subfolder. Default: False."
|
||||
),
|
||||
)
|
||||
pdfa_upload_processed: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Upload the processed-file PDF/A variant to all configured storage providers. "
|
||||
"Files are placed in the provider's folder + PDFA_UPLOAD_FOLDER subfolder. Default: False."
|
||||
),
|
||||
)
|
||||
pdfa_upload_folder: str = Field(
|
||||
default="pdfa",
|
||||
description=(
|
||||
"Subfolder name appended to each storage provider's configured folder for PDF/A uploads. "
|
||||
"For example if Dropbox folder is '/Documents' and this is 'pdfa', PDF/A files go to "
|
||||
"'/Documents/pdfa'. Set to empty string to upload into the same folder. Default: 'pdfa'."
|
||||
),
|
||||
)
|
||||
google_drive_pdfa_folder_id: str = Field(
|
||||
default="",
|
||||
description=(
|
||||
"Google Drive folder ID for PDF/A uploads. Since Google Drive uses IDs not paths, "
|
||||
"this must be set separately. If empty, uses the standard google_drive_folder_id."
|
||||
),
|
||||
)
|
||||
pdfa_timestamp_enabled: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Enable RFC 3161 timestamping of PDF/A files via a Timestamp Authority (TSA). "
|
||||
"Creates a .tsr file alongside each PDF/A file for legal proof of existence. "
|
||||
"Requires openssl binary on PATH. Default: False."
|
||||
),
|
||||
)
|
||||
pdfa_timestamp_url: str = Field(
|
||||
default="https://freetsa.org/tsr",
|
||||
description=(
|
||||
"URL of the RFC 3161 Timestamp Authority. Default: FreeTSA (https://freetsa.org/tsr). "
|
||||
"Other options: GlobalSign, DigiStamp, or any RFC 3161-compliant TSA."
|
||||
),
|
||||
)
|
||||
|
||||
imap_readonly_mode: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
@@ -231,6 +647,17 @@ class Settings(BaseSettings):
|
||||
),
|
||||
)
|
||||
|
||||
imap_attachment_filter: str = Field(
|
||||
default="documents_only",
|
||||
description=(
|
||||
"Controls which attachment types are ingested from IMAP emails. "
|
||||
"Accepted values: "
|
||||
"'documents_only' – ingest only PDFs and office files (Word, Excel, PowerPoint, ODT, etc.); "
|
||||
"'all' – ingest all supported file types including images. "
|
||||
"This is the global default; individual user IMAP accounts can override it."
|
||||
),
|
||||
)
|
||||
|
||||
# Batch processing settings
|
||||
processall_throttle_threshold: int = Field(
|
||||
default=20,
|
||||
@@ -272,6 +699,66 @@ class Settings(BaseSettings):
|
||||
default=True,
|
||||
description="Send notifications when files are successfully processed",
|
||||
)
|
||||
notify_on_user_signup: bool = Field(
|
||||
default=True,
|
||||
description="Send admin notifications when a new user signs up",
|
||||
)
|
||||
notify_on_plan_change: bool = Field(
|
||||
default=True,
|
||||
description="Send admin notifications when a user changes their subscription plan",
|
||||
)
|
||||
notify_on_payment_issue: bool = Field(
|
||||
default=True,
|
||||
description="Send admin notifications when a payment issue is reported for a user",
|
||||
)
|
||||
|
||||
# Webhook settings
|
||||
webhook_enabled: bool = Field(
|
||||
default=True,
|
||||
description="Enable webhook delivery for document events",
|
||||
)
|
||||
|
||||
# ── Backup / restore settings ──────────────────────────────────────────────
|
||||
backup_enabled: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Enable automatic scheduled database backups. "
|
||||
"When enabled, hourly, daily, and weekly backups are created automatically. Default: True."
|
||||
),
|
||||
)
|
||||
backup_dir: Optional[str] = Field(
|
||||
default=None,
|
||||
description=("Directory where local backup archives are stored. Defaults to <workdir>/backups when not set."),
|
||||
)
|
||||
# Remote destination: one of s3, dropbox, google_drive, onedrive, nextcloud,
|
||||
# webdav, ftp, sftp, email, or empty/None for local-only.
|
||||
backup_remote_destination: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Storage provider to upload remote backup copies to. "
|
||||
"Accepted values: s3, dropbox, google_drive, onedrive, nextcloud, webdav, ftp, sftp, email. "
|
||||
"Leave empty to keep backups local only."
|
||||
),
|
||||
)
|
||||
backup_remote_folder: str = Field(
|
||||
default="backups",
|
||||
description=(
|
||||
"Sub-folder / key prefix used when uploading backup archives to the remote destination. Default: 'backups'."
|
||||
),
|
||||
)
|
||||
# Retention counts (number of snapshots to keep per tier)
|
||||
backup_retain_hourly: int = Field(
|
||||
default=96,
|
||||
description="Number of hourly backups to retain (default 96 = 4 days × 24 h).",
|
||||
)
|
||||
backup_retain_daily: int = Field(
|
||||
default=21,
|
||||
description="Number of daily backups to retain (default 21 = 3 weeks).",
|
||||
)
|
||||
backup_retain_weekly: int = Field(
|
||||
default=13,
|
||||
description="Number of weekly backups to retain (default 13 ≈ 3 months / 91 days).",
|
||||
)
|
||||
|
||||
# File upload size limits (for security - see SECURITY_AUDIT.md)
|
||||
max_upload_size: int = Field(
|
||||
@@ -309,6 +796,36 @@ class Settings(BaseSettings):
|
||||
" If False, the check is still performed but not displayed. Default: True."
|
||||
),
|
||||
)
|
||||
near_duplicate_threshold: float = Field(
|
||||
default=0.85,
|
||||
description=(
|
||||
"Minimum cosine similarity score (0–1) between two documents' text embeddings to consider "
|
||||
"them near-duplicates. Higher values require closer content matches. Default: 0.85."
|
||||
),
|
||||
)
|
||||
embedding_model: str = Field(
|
||||
default="text-embedding-3-small",
|
||||
description=(
|
||||
"Model name used for generating text embeddings via the OpenAI-compatible API. "
|
||||
"Embeddings drive the document similarity feature. Default: text-embedding-3-small."
|
||||
),
|
||||
)
|
||||
embedding_max_tokens: int = Field(
|
||||
default=8000,
|
||||
description=(
|
||||
"Maximum number of tokens to send to the embedding model. "
|
||||
"Text is truncated to approximately this many tokens (using a "
|
||||
"conservative 3-chars-per-token estimate) before calling the API. "
|
||||
"Set this below the model's context window (e.g. 8000 for an 8192-token model)."
|
||||
),
|
||||
)
|
||||
embedding_backfill_batch_size: int = Field(
|
||||
default=50,
|
||||
description=(
|
||||
"Maximum number of files to queue for embedding computation per "
|
||||
"backfill run. Keeps the worker and embedding API load bounded."
|
||||
),
|
||||
)
|
||||
|
||||
# Text quality check - AI-based assessment of embedded PDF text
|
||||
enable_text_quality_check: bool = Field(
|
||||
@@ -339,6 +856,32 @@ class Settings(BaseSettings):
|
||||
),
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task retry settings (see app/tasks/retry_config.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
task_retry_max_retries: int = Field(
|
||||
default=3,
|
||||
description=("Maximum number of automatic retry attempts for failed Celery tasks. Default: 3."),
|
||||
)
|
||||
task_retry_delays: Union[List[int], str] = Field(
|
||||
default_factory=lambda: [60, 300, 900],
|
||||
description=(
|
||||
"Comma-separated list of retry countdown values in seconds. "
|
||||
"Each value is the delay before the corresponding retry attempt. "
|
||||
"If a task fails more times than entries in this list, the last delay "
|
||||
"is doubled for each additional attempt. "
|
||||
"Default: 60,300,900 (1 min, 5 min, 15 min)."
|
||||
),
|
||||
)
|
||||
task_retry_jitter: bool = Field(
|
||||
default=True,
|
||||
description=(
|
||||
"Apply ±20 % random jitter to retry countdowns to prevent "
|
||||
"thundering-herd problems when many tasks fail simultaneously. "
|
||||
"Default: True (enabled)."
|
||||
),
|
||||
)
|
||||
|
||||
# Processing step timeout - prevents files from getting stuck in "in_progress" state
|
||||
step_timeout: int = Field(
|
||||
default=600,
|
||||
@@ -407,6 +950,49 @@ class Settings(BaseSettings):
|
||||
),
|
||||
)
|
||||
|
||||
# SIEM / External Audit Log Forwarding
|
||||
# Forward audit events to external SIEM systems for centralised monitoring.
|
||||
audit_siem_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Enable forwarding of audit events to an external SIEM system.",
|
||||
)
|
||||
audit_siem_transport: str = Field(
|
||||
default="syslog",
|
||||
description=(
|
||||
"Transport used to forward audit events. "
|
||||
"Options: 'syslog' (RFC 5424 over UDP/TCP), 'http' (JSON POST to a webhook URL, "
|
||||
"compatible with Splunk HEC, Logstash HTTP input, Grafana Loki, etc.)."
|
||||
),
|
||||
)
|
||||
audit_siem_syslog_host: str = Field(
|
||||
default="localhost",
|
||||
description="Hostname or IP of the syslog receiver.",
|
||||
)
|
||||
audit_siem_syslog_port: int = Field(
|
||||
default=514,
|
||||
description="Port of the syslog receiver.",
|
||||
)
|
||||
audit_siem_syslog_protocol: str = Field(
|
||||
default="udp",
|
||||
description="Protocol for syslog transport: 'udp' or 'tcp'.",
|
||||
)
|
||||
audit_siem_http_url: str = Field(
|
||||
default="",
|
||||
description=(
|
||||
"HTTP endpoint URL for SIEM webhook delivery. "
|
||||
"Supports Splunk HEC (https://splunk:8088/services/collector/event), "
|
||||
"Logstash HTTP input, Grafana Loki push API, or any JSON-accepting endpoint."
|
||||
),
|
||||
)
|
||||
audit_siem_http_token: str = Field(
|
||||
default="",
|
||||
description="Bearer / HEC token included in the Authorization header of SIEM HTTP requests.",
|
||||
)
|
||||
audit_siem_http_custom_headers: str = Field(
|
||||
default="",
|
||||
description="Comma-separated 'Key:Value' pairs of extra headers for SIEM HTTP requests.",
|
||||
)
|
||||
|
||||
# UI / Appearance
|
||||
ui_default_color_scheme: str = Field(
|
||||
default="system",
|
||||
@@ -472,6 +1058,77 @@ class Settings(BaseSettings):
|
||||
description="Allowed request headers for CORS. Use ['*'] to allow all headers.",
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Support / Help Center – Zammad integration
|
||||
# ---------------------------------------------------------------------------
|
||||
zammad_url: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Base URL of your Zammad instance (e.g. https://zammad.example.com). "
|
||||
"Required for the chat widget and feedback form on the Help Center page."
|
||||
),
|
||||
)
|
||||
zammad_chat_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Show the Zammad live-chat widget on the Help Center page.",
|
||||
)
|
||||
zammad_chat_id: int = Field(
|
||||
default=1,
|
||||
description="Zammad chat topic ID to use for the live-chat widget.",
|
||||
)
|
||||
zammad_form_enabled: bool = Field(
|
||||
default=False,
|
||||
description="Show the Zammad feedback / ticket form on the Help Center page.",
|
||||
)
|
||||
support_email: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Support e-mail address displayed on the Help Center page.",
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Observability – Sentry error & performance monitoring
|
||||
# ---------------------------------------------------------------------------
|
||||
sentry_dsn: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Sentry Data Source Name (DSN). When set, error reporting and "
|
||||
"performance tracing are enabled automatically. Leave blank (or unset) "
|
||||
"to disable Sentry entirely."
|
||||
),
|
||||
)
|
||||
sentry_environment: str = Field(
|
||||
default="production",
|
||||
description=(
|
||||
"Environment tag sent to Sentry (e.g. 'development', 'staging', 'production'). "
|
||||
"Helps you filter events in the Sentry dashboard."
|
||||
),
|
||||
)
|
||||
sentry_traces_sample_rate: float = Field(
|
||||
default=0.1,
|
||||
description=(
|
||||
"Fraction of transactions to capture for performance monitoring (0.0–1.0). "
|
||||
"Set to 0.0 to disable tracing, 1.0 to capture every transaction. "
|
||||
"Values above 0 may increase Sentry quota usage."
|
||||
),
|
||||
)
|
||||
sentry_profiles_sample_rate: float = Field(
|
||||
default=0.0,
|
||||
description=(
|
||||
"Fraction of profiled transactions to send to Sentry (0.0–1.0). "
|
||||
"Profiling is only active when traces_sample_rate > 0. "
|
||||
"Defaults to 0.0 (disabled) to minimise overhead."
|
||||
),
|
||||
)
|
||||
sentry_send_default_pii: bool = Field(
|
||||
default=False,
|
||||
description=(
|
||||
"Whether to attach personally identifiable information (PII) – such as "
|
||||
"IP addresses and user agents – to Sentry events. Disabled by default "
|
||||
"for privacy compliance (GDPR / CCPA). Enable only if your Sentry "
|
||||
"project is configured to handle PII."
|
||||
),
|
||||
)
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def strip_outer_quotes(cls, data: Any) -> Any:
|
||||
@@ -527,6 +1184,15 @@ class Settings(BaseSettings):
|
||||
return []
|
||||
return v
|
||||
|
||||
@field_validator("task_retry_delays", mode="before")
|
||||
@classmethod
|
||||
def parse_task_retry_delays(cls, v: str | list) -> list[int]:
|
||||
"""Parse task retry delays from comma-separated string or list of ints."""
|
||||
if isinstance(v, str):
|
||||
parts = [p.strip() for p in v.split(",") if p.strip()]
|
||||
return [int(p) for p in parts]
|
||||
return [int(item) for item in v]
|
||||
|
||||
@field_validator("session_secret")
|
||||
@classmethod
|
||||
def validate_session_secret(cls, v: str | None, info: object) -> str | None:
|
||||
@@ -599,5 +1265,51 @@ class Settings(BaseSettings):
|
||||
# Return basic info if file not found
|
||||
return f"Version: {self.version}\nBuild Date: {self.build_date}\nGit SHA: {self.git_sha}"
|
||||
|
||||
@property
|
||||
def release_name(self) -> str | None:
|
||||
"""Get the release codename for the current version from release_names.json.
|
||||
|
||||
Looks up the current version's minor version prefix (e.g., '0.5' for '0.5.3')
|
||||
in release_names.json to find the associated codename. Returns None if no
|
||||
codename is defined for the current version.
|
||||
|
||||
Returns:
|
||||
The release codename string, or None if not found.
|
||||
"""
|
||||
version = self.version
|
||||
if not version or version == "unknown":
|
||||
return None
|
||||
|
||||
release_names_file = os.path.join(os.path.dirname(os.path.dirname(__file__)), "release_names.json")
|
||||
if not os.path.exists(release_names_file):
|
||||
return None
|
||||
|
||||
try:
|
||||
with open(release_names_file, "r") as f:
|
||||
data = json.load(f)
|
||||
|
||||
releases = data.get("releases", {})
|
||||
|
||||
# Try exact version match first (e.g., "0.5.0")
|
||||
if version in releases:
|
||||
return releases[version].get("codename")
|
||||
|
||||
# Try minor version prefix (e.g., "0.5" for "0.5.3")
|
||||
parts = version.split(".")
|
||||
if len(parts) >= 2:
|
||||
minor_prefix = f"{parts[0]}.{parts[1]}"
|
||||
if minor_prefix in releases:
|
||||
return releases[minor_prefix].get("codename")
|
||||
|
||||
# Try major version prefix (e.g., "1" for "1.0.0")
|
||||
if len(parts) >= 1:
|
||||
major_prefix = parts[0]
|
||||
if major_prefix in releases:
|
||||
return releases[major_prefix].get("codename")
|
||||
|
||||
return None
|
||||
except (json.JSONDecodeError, KeyError, IndexError):
|
||||
return None
|
||||
|
||||
|
||||
settings = Settings()
|
||||
|
||||
+27
-8
@@ -26,8 +26,14 @@ SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
|
||||
def init_db() -> None:
|
||||
"""
|
||||
Ensures the SQLite database file and its parent directory exist (if using sqlite).
|
||||
Then runs Base.metadata.create_all(bind=engine) to initialize tables and
|
||||
applies any pending Alembic migrations.
|
||||
Then initializes tables and applies any pending Alembic migrations:
|
||||
|
||||
- Fresh/legacy databases (no ``alembic_version`` table): creates all tables via
|
||||
``Base.metadata.create_all()``, then stamps the Alembic version to ``head``.
|
||||
- Alembic-tracked databases (``alembic_version`` present): skips ``create_all()``
|
||||
and applies pending migrations via ``alembic upgrade head``. Skipping
|
||||
``create_all()`` prevents an ``OperationalError`` when the ORM model defines a
|
||||
table (e.g. ``webhook_configs``) that a pending migration also tries to create.
|
||||
"""
|
||||
# 1. Parse the DB URL to see if it's sqlite
|
||||
url = make_url(DB_URL)
|
||||
@@ -47,12 +53,19 @@ def init_db() -> None:
|
||||
logger.info(f"Creating new SQLite database file at {database_path}")
|
||||
open(database_path, "a").close()
|
||||
|
||||
# 5. Now create tables if they don't exist yet
|
||||
# 5. Create tables only for fresh/legacy databases not yet tracked by Alembic.
|
||||
# For Alembic-tracked databases, skip create_all to avoid conflicts where
|
||||
# the ORM model would create a table (e.g. webhook_configs) that a pending
|
||||
# Alembic migration also tries to create, causing an OperationalError.
|
||||
try:
|
||||
Base.metadata.create_all(bind=engine)
|
||||
logger.info("Database initialization complete (tables created if not exist).")
|
||||
from sqlalchemy import inspect
|
||||
|
||||
# 6. Run Alembic migrations for existing databases
|
||||
table_names = inspect(engine).get_table_names()
|
||||
if "alembic_version" not in table_names:
|
||||
Base.metadata.create_all(bind=engine)
|
||||
logger.info("Database initialization complete (tables created if not exist).")
|
||||
|
||||
# 6. Run Alembic migrations (stamps fresh/legacy DBs to head, upgrades tracked DBs)
|
||||
_run_alembic_upgrade(engine)
|
||||
except exc.SQLAlchemyError as e:
|
||||
logger.error(f"Error initializing database: {e}")
|
||||
@@ -197,8 +210,10 @@ def _run_schema_migrations(engine: Any) -> None:
|
||||
if unique_filehash_indexes:
|
||||
logger.info("Migrating files: dropping unique index on 'filehash'")
|
||||
with engine.begin() as conn:
|
||||
preparer = conn.dialect.identifier_preparer
|
||||
for index in unique_filehash_indexes:
|
||||
conn.execute(text(f"DROP INDEX IF EXISTS {index['name']}"))
|
||||
quoted_idx = preparer.quote(index["name"])
|
||||
conn.execute(text(f"DROP INDEX IF EXISTS {quoted_idx}"))
|
||||
logger.info("Migration complete: unique index on 'filehash' removed")
|
||||
except Exception as exc:
|
||||
logger.warning(f"Skipping filehash unique index drop: {exc}")
|
||||
@@ -250,12 +265,16 @@ def _ensure_indexes(engine: Any, inspector: Any) -> None:
|
||||
table_names = inspector.get_table_names()
|
||||
columns_by_table: dict[str, set[str]] = {}
|
||||
with engine.begin() as conn:
|
||||
preparer = conn.dialect.identifier_preparer
|
||||
for idx_name, table, column in _PERF_INDEXES:
|
||||
if table in table_names:
|
||||
if table not in columns_by_table:
|
||||
columns_by_table[table] = {col["name"] for col in inspector.get_columns(table)}
|
||||
if column in columns_by_table[table]:
|
||||
conn.execute(text(f"CREATE INDEX IF NOT EXISTS {idx_name} ON {table} ({column})"))
|
||||
quoted_idx = preparer.quote(idx_name)
|
||||
quoted_table = preparer.quote(table)
|
||||
quoted_col = preparer.quote(column)
|
||||
conn.execute(text(f"CREATE INDEX IF NOT EXISTS {quoted_idx} ON {quoted_table} ({quoted_col})"))
|
||||
|
||||
logger.info("Performance indexes ensured")
|
||||
|
||||
|
||||
+103
-7
@@ -16,6 +16,8 @@ from starlette.middleware.trustedhost import TrustedHostMiddleware
|
||||
from uvicorn.middleware.proxy_headers import ProxyHeadersMiddleware
|
||||
|
||||
from app.api import router as api_router
|
||||
from app.api.graphql_api import graphql_router
|
||||
from app.api.local_auth import router as local_auth_router
|
||||
from app.auth import router as auth_router
|
||||
from app.config import settings
|
||||
from app.database import init_db
|
||||
@@ -26,6 +28,7 @@ from app.middleware.request_size_limit import RequestSizeLimitMiddleware
|
||||
from app.middleware.security_headers import SecurityHeadersMiddleware
|
||||
from app.utils.config_validator import check_all_configs
|
||||
from app.utils.notification import init_apprise, notify_shutdown, notify_startup
|
||||
from app.utils.sentry import init_sentry
|
||||
|
||||
# Import the routers - now using views directly instead of frontend
|
||||
from app.views import router as frontend_router
|
||||
@@ -69,6 +72,10 @@ async def lifespan(app: FastAPI):
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# Initialize Sentry after DB settings are loaded so that values configured
|
||||
# via the database UI (e.g. SENTRY_DSN) are respected in addition to env vars.
|
||||
init_sentry()
|
||||
|
||||
# Ensure OCR language data is available (background download, non-blocking)
|
||||
from app.utils.ocr_language_manager import ensure_ocr_languages_async
|
||||
|
||||
@@ -99,6 +106,61 @@ async def lifespan(app: FastAPI):
|
||||
# Send startup notification
|
||||
notify_startup()
|
||||
|
||||
# Seed default subscription plans if none exist
|
||||
try:
|
||||
from app.database import SessionLocal as _SessionLocal
|
||||
from app.utils.subscription import seed_default_plans as _seed_plans
|
||||
|
||||
_db_seed = _SessionLocal()
|
||||
try:
|
||||
_seed_plans(_db_seed)
|
||||
finally:
|
||||
_db_seed.close()
|
||||
except Exception:
|
||||
logging.debug("Subscription plan seeding skipped — DB may not be ready yet") # noqa: S110
|
||||
|
||||
# Seed the default system pipeline (mirrors the current hardcoded processing
|
||||
# workflow) so it is immediately visible in the Pipelines management UI.
|
||||
try:
|
||||
from app.api.pipelines import seed_default_pipeline as _seed_pipeline
|
||||
from app.database import SessionLocal as _SessionLocal # noqa: F811 (re-import for clarity)
|
||||
|
||||
_db_pipeline = _SessionLocal()
|
||||
try:
|
||||
_seed_pipeline(_db_pipeline)
|
||||
finally:
|
||||
_db_pipeline.close()
|
||||
except Exception:
|
||||
logging.debug("Default pipeline seeding skipped — DB may not be ready yet") # noqa: S110
|
||||
|
||||
# Seed the default scheduled batch processing jobs so they appear in the
|
||||
# admin UI (/admin/scheduled-jobs) on first startup.
|
||||
try:
|
||||
from app.api.scheduled_jobs import seed_default_scheduled_jobs as _seed_jobs
|
||||
from app.database import SessionLocal as _SessionLocal # noqa: F811
|
||||
|
||||
_db_jobs = _SessionLocal()
|
||||
try:
|
||||
_seed_jobs(_db_jobs)
|
||||
finally:
|
||||
_db_jobs.close()
|
||||
except Exception:
|
||||
logging.debug("Scheduled jobs seeding skipped — DB may not be ready yet") # noqa: S110
|
||||
|
||||
# Seed the built-in compliance templates (GDPR, HIPAA, SOC2) so they
|
||||
# are available in the admin compliance dashboard on first startup.
|
||||
try:
|
||||
from app.database import SessionLocal as _SessionLocal # noqa: F811
|
||||
from app.utils.compliance_service import seed_compliance_templates as _seed_compliance
|
||||
|
||||
_db_compliance = _SessionLocal()
|
||||
try:
|
||||
_seed_compliance(_db_compliance)
|
||||
finally:
|
||||
_db_compliance.close()
|
||||
except Exception:
|
||||
logging.debug("Compliance template seeding skipped — DB may not be ready yet") # noqa: S110
|
||||
|
||||
# Application is now running
|
||||
yield
|
||||
|
||||
@@ -109,7 +171,12 @@ async def lifespan(app: FastAPI):
|
||||
notify_shutdown()
|
||||
|
||||
|
||||
app = FastAPI(title="DocuElevate", lifespan=lifespan)
|
||||
app = FastAPI(
|
||||
title="DocuElevate",
|
||||
lifespan=lifespan,
|
||||
docs_url="/admin/api-docs",
|
||||
redoc_url="/admin/api-redoc",
|
||||
)
|
||||
|
||||
# Initialize rate limiter and attach to app state
|
||||
limiter = create_limiter(redis_url=settings.redis_url, enabled=settings.rate_limiting_enabled)
|
||||
@@ -174,8 +241,36 @@ if os.path.exists(static_dir):
|
||||
else:
|
||||
print(f"WARNING: Static directory not found at {static_dir}. Static files will not be served.")
|
||||
|
||||
# Mount the built MkDocs developer documentation at /developer-docs/
|
||||
# These docs target administrators and developers, not end-users.
|
||||
# The user-facing Help Center is served by the /help view instead.
|
||||
# The docs are pre-built into docs_build/ during the Docker image build.
|
||||
# When running locally, run `mkdocs build` from the repo root first.
|
||||
docs_build_dir = pathlib.Path(__file__).parents[1] / "docs_build"
|
||||
if os.path.exists(docs_build_dir):
|
||||
app.mount("/developer-docs", StaticFiles(directory=str(docs_build_dir), html=True), name="developer_docs")
|
||||
else:
|
||||
print(f"INFO: Developer docs not found at {docs_build_dir}. Run 'mkdocs build' to generate them.")
|
||||
|
||||
|
||||
# Custom exception handlers that return JSON for API routes and HTML for frontend routes
|
||||
# These use their own separate templates instance so that patches in tests on individual
|
||||
# view modules do not affect the error handler rendering.
|
||||
_error_templates_dir = pathlib.Path(__file__).parents[1] / "frontend" / "templates"
|
||||
_error_templates = Jinja2Templates(directory=str(_error_templates_dir))
|
||||
# Register the i18n translate helper as a global so error templates can use {{ _("key") }}.
|
||||
# Error pages use the default language (English); request-specific locale is not needed here.
|
||||
from app.utils.i18n import SUPPORTED_LANGUAGES as _SUPPORTED_LANGUAGES # noqa: E402
|
||||
from app.utils.i18n import get_suggested_languages as _get_suggested_languages # noqa: E402
|
||||
from app.utils.i18n import translate as _translate_fn # noqa: E402
|
||||
|
||||
_error_templates.env.globals["_"] = lambda key, **kwargs: _translate_fn(key, "en", **kwargs)
|
||||
_error_templates.env.globals["min"] = min
|
||||
_error_templates.env.globals["max"] = max
|
||||
_error_templates.env.globals["supported_languages"] = _SUPPORTED_LANGUAGES
|
||||
_error_templates.env.globals["suggested_languages"] = _get_suggested_languages("en", "")
|
||||
|
||||
|
||||
@app.exception_handler(HTTPException)
|
||||
async def http_exception_handler(request: Request, exc: HTTPException):
|
||||
"""
|
||||
@@ -187,15 +282,15 @@ async def http_exception_handler(request: Request, exc: HTTPException):
|
||||
return JSONResponse(status_code=exc.status_code, content={"detail": exc.detail})
|
||||
|
||||
# For frontend routes, return appropriate HTML templates
|
||||
templates = Jinja2Templates(directory=str(static_dir.parent / "templates"))
|
||||
|
||||
# Handle 404 errors with a custom template
|
||||
if exc.status_code == 404:
|
||||
return templates.TemplateResponse("404.html", {"request": request}, status_code=status.HTTP_404_NOT_FOUND)
|
||||
return _error_templates.TemplateResponse(
|
||||
"404.html", {"request": request}, status_code=status.HTTP_404_NOT_FOUND
|
||||
)
|
||||
|
||||
# For other HTTP errors, we could create specific templates or use a generic one
|
||||
# For now, return a simple error page
|
||||
return templates.TemplateResponse(
|
||||
return _error_templates.TemplateResponse(
|
||||
"404.html", # Reuse 404 template for other errors, or create a generic error template
|
||||
{"request": request},
|
||||
status_code=exc.status_code,
|
||||
@@ -216,8 +311,7 @@ async def custom_500_handler(request: Request, exc: Exception):
|
||||
)
|
||||
|
||||
# Serve the 500 template for non-API routes
|
||||
templates = Jinja2Templates(directory=str(static_dir.parent / "templates"))
|
||||
return templates.TemplateResponse(
|
||||
return _error_templates.TemplateResponse(
|
||||
"500.html",
|
||||
{"request": request, "exc": exc},
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
@@ -233,4 +327,6 @@ def test_500():
|
||||
app.include_router(frontend_router)
|
||||
app.include_router(files_router) # Explicitly include the files router
|
||||
app.include_router(auth_router)
|
||||
app.include_router(local_auth_router)
|
||||
app.include_router(api_router, prefix="/api")
|
||||
app.include_router(graphql_router, prefix="/graphql")
|
||||
|
||||
@@ -109,6 +109,13 @@ class CSRFMiddleware(BaseHTTPMiddleware):
|
||||
|
||||
# Validate for state-changing methods on non-exempt paths.
|
||||
if request.method in CSRF_PROTECTED_METHODS and request.url.path not in CSRF_EXEMPT_PATHS:
|
||||
# Bearer-authenticated requests (API tokens) are exempt from CSRF
|
||||
# because the token itself acts as proof of intent — it cannot be
|
||||
# injected by a cross-site request from a browser.
|
||||
auth_header = request.headers.get("authorization", "")
|
||||
if auth_header.startswith("Bearer "):
|
||||
return await call_next(request)
|
||||
|
||||
submitted_token = await self._get_submitted_token(request)
|
||||
if not submitted_token or not secrets.compare_digest(csrf_token, submitted_token):
|
||||
logger.warning(f"[SECURITY] CSRF_VALIDATION_FAILED method={request.method} path={request.url.path}")
|
||||
@@ -148,10 +155,26 @@ class CSRFMiddleware(BaseHTTPMiddleware):
|
||||
|
||||
# 2. For URL-encoded form bodies only (plain HTML form submissions).
|
||||
content_type = request.headers.get("content-type", "")
|
||||
logger.debug("CSRF: content_type=%r method=%s path=%s", content_type, request.method, request.url.path)
|
||||
if "application/x-www-form-urlencoded" in content_type:
|
||||
try:
|
||||
# Cache the raw body bytes before parsing the form. Starlette's
|
||||
# BaseHTTPMiddleware uses _CachedRequest.wrapped_receive to relay
|
||||
# the body to downstream handlers. When form() is called it
|
||||
# internally uses stream() which sets _stream_consumed=True but
|
||||
# does NOT populate _body. wrapped_receive then sees a consumed
|
||||
# stream and forwards an empty body, so the auth endpoint gets
|
||||
# form_keys=[]. Calling body() first stores the bytes in _body;
|
||||
# wrapped_receive detects this and replays the real body to any
|
||||
# downstream handler (e.g. the /auth endpoint).
|
||||
await request.body()
|
||||
form = await request.form()
|
||||
token = form.get("csrf_token")
|
||||
logger.debug(
|
||||
"CSRF: form_keys=%s csrf_token_present=%s",
|
||||
list(form.keys()),
|
||||
bool(token),
|
||||
)
|
||||
if token:
|
||||
return str(token)
|
||||
except Exception as exc:
|
||||
|
||||
+879
-1
@@ -1,11 +1,13 @@
|
||||
# app/models.py
|
||||
|
||||
from sqlalchemy import Boolean, Column, DateTime, ForeignKey, Integer, String, Text, UniqueConstraint, func
|
||||
from sqlalchemy import Boolean, Column, DateTime, Float, ForeignKey, Integer, String, Text, UniqueConstraint, func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
# Foreign key constants
|
||||
_FILES_ID_FK = "files.id"
|
||||
_PIPELINES_ID_FK = "pipelines.id"
|
||||
_ROUTING_RULES_TABLE = "pipeline_routing_rules"
|
||||
|
||||
|
||||
class DocumentMetadata(Base):
|
||||
@@ -24,6 +26,11 @@ class FileRecord(Base):
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Owner identifier for multi-user mode.
|
||||
# Stores the user's unique identifier (e.g. email or OAuth sub claim).
|
||||
# NULL means the file belongs to the shared/global space (single-user mode).
|
||||
owner_id = Column(String, nullable=True, index=True)
|
||||
|
||||
# Hash of the file content (e.g. SHA-256)
|
||||
# Note: duplicates are allowed so filehash is not unique
|
||||
filehash = Column(String, index=True, nullable=False)
|
||||
@@ -67,6 +74,16 @@ class FileRecord(Base):
|
||||
# Human-readable document title from AI metadata
|
||||
document_title = Column(String, nullable=True)
|
||||
|
||||
# PDF/A archival variant paths (generated when ENABLE_PDFA_CONVERSION is True)
|
||||
original_pdfa_path = Column(String, nullable=True) # PDF/A copy of the original ingested file
|
||||
processed_pdfa_path = Column(String, nullable=True) # PDF/A copy of the processed file
|
||||
|
||||
# Pre-computed text embedding vector stored as JSON array of floats
|
||||
embedding = Column(Text, nullable=True)
|
||||
|
||||
# Processing pipeline assigned to this file (NULL = use system default)
|
||||
pipeline_id = Column(Integer, ForeignKey(_PIPELINES_ID_FK), nullable=True, index=True)
|
||||
|
||||
# Timestamp when we inserted this record
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now(), index=True)
|
||||
|
||||
@@ -130,6 +147,27 @@ class SettingsAuditLog(Base):
|
||||
action = Column(String, nullable=False) # "update" or "delete"
|
||||
|
||||
|
||||
class AuditLog(Base):
|
||||
"""Comprehensive audit log for compliance tracking.
|
||||
|
||||
Records all significant actions: login/logout, document CRUD, settings
|
||||
changes, and administrative operations. Rows are append-only; the API
|
||||
and service layer never update or delete entries.
|
||||
"""
|
||||
|
||||
__tablename__ = "audit_logs"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
timestamp = Column(DateTime(timezone=True), server_default=func.now(), nullable=False, index=True)
|
||||
user = Column(String, nullable=False, index=True) # Username or "anonymous" / "system"
|
||||
action = Column(String, nullable=False, index=True) # e.g. "login", "document.create", "settings.update"
|
||||
resource_type = Column(String, nullable=True, index=True) # e.g. "document", "user", "settings"
|
||||
resource_id = Column(String, nullable=True) # ID of the affected resource
|
||||
ip_address = Column(String, nullable=True) # Client IP address
|
||||
details = Column(Text, nullable=True) # JSON-encoded extra context
|
||||
severity = Column(String(16), nullable=False, server_default="info") # info / warning / error / critical
|
||||
|
||||
|
||||
class SavedSearch(Base):
|
||||
"""User-defined saved search filters for quick access to frequently used filter combinations."""
|
||||
|
||||
@@ -143,3 +181,843 @@ class SavedSearch(Base):
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
__table_args__ = (UniqueConstraint("user_id", "name", name="unique_user_search_name"),)
|
||||
|
||||
|
||||
class WebhookConfig(Base):
|
||||
"""Webhook configuration for notifying external systems of document events."""
|
||||
|
||||
__tablename__ = "webhook_configs"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
url = Column(String, nullable=False) # Target URL for webhook delivery
|
||||
secret = Column(String, nullable=True) # Shared secret for HMAC-SHA256 signature
|
||||
events = Column(Text, nullable=False) # JSON list of subscribed events
|
||||
is_active = Column(Boolean, default=True, nullable=False) # Whether the webhook is active
|
||||
description = Column(String, nullable=True) # Optional human-readable description
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class LocalUser(Base):
|
||||
"""A locally-registered user authenticated by email and bcrypt password.
|
||||
|
||||
Created during the self-registration flow when ``allow_local_signup`` is
|
||||
enabled. The account is inactive (``is_active=False``) until the user
|
||||
clicks the verification link sent to their email address.
|
||||
"""
|
||||
|
||||
__tablename__ = "local_users"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
email = Column(String(255), unique=True, nullable=False, index=True)
|
||||
username = Column(String(64), unique=True, nullable=False, index=True)
|
||||
display_name = Column(String(255), nullable=True)
|
||||
hashed_password = Column(String(255), nullable=False)
|
||||
is_active = Column(Boolean, nullable=False, default=False, server_default="0")
|
||||
is_admin = Column(Boolean, nullable=False, default=False, server_default="0")
|
||||
email_verification_token = Column(String(128), nullable=True)
|
||||
email_verification_sent_at = Column(DateTime(timezone=True), nullable=True)
|
||||
password_reset_token = Column(String(128), nullable=True)
|
||||
password_reset_sent_at = Column(DateTime(timezone=True), nullable=True)
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class UserProfile(Base):
|
||||
"""Per-user profile for admin-managed settings in multi-user mode.
|
||||
|
||||
Each row corresponds to one authenticated user (identified by their
|
||||
``user_id``, which matches ``FileRecord.owner_id``). The admin can
|
||||
create or update profiles to override global defaults such as the
|
||||
daily upload limit and to attach notes or block a user.
|
||||
"""
|
||||
|
||||
__tablename__ = "user_profiles"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Stable user identifier — matches FileRecord.owner_id (OAuth sub / email / username)
|
||||
user_id = Column(String, unique=True, nullable=False, index=True)
|
||||
|
||||
# Optional human-readable display name set by the admin
|
||||
display_name = Column(String, nullable=True)
|
||||
|
||||
# Per-user daily upload limit; NULL means "use global default"
|
||||
daily_upload_limit = Column(Integer, nullable=True)
|
||||
|
||||
# Admin-only free-text notes about this user
|
||||
notes = Column(Text, nullable=True)
|
||||
|
||||
# When True the user is prevented from uploading new documents
|
||||
is_blocked = Column(Boolean, default=False, nullable=False)
|
||||
|
||||
# Subscription tier: "free" | "starter" | "professional" | "business"
|
||||
# NULL is treated as "free" by the subscription utility.
|
||||
subscription_tier = Column(String(50), nullable=True, default="free")
|
||||
|
||||
# Billing cycle and overage settings (added in migration 016)
|
||||
subscription_billing_cycle = Column(String(10), nullable=False, default="monthly", server_default="monthly")
|
||||
subscription_period_start = Column(DateTime(timezone=True), nullable=True)
|
||||
allow_overage = Column(Boolean, nullable=False, default=False, server_default="0")
|
||||
|
||||
# Pending subscription change (added in migration 020_add_subscription_change_pending)
|
||||
# When a user requests a downgrade, the new tier is stored here and the
|
||||
# change is applied on `subscription_change_pending_date`. Upgrades are
|
||||
# applied immediately and these fields are left NULL.
|
||||
subscription_change_pending_tier = Column(String(50), nullable=True)
|
||||
subscription_change_pending_date = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# When True, the user is on a complimentary (uncharged) plan — they keep all tier
|
||||
# quota benefits but are never billed via Stripe. Automatically set for admin users.
|
||||
is_complimentary = Column(Boolean, nullable=False, default=False, server_default="0")
|
||||
|
||||
# Onboarding tracking (added in migration 017)
|
||||
onboarding_completed = Column(Boolean, nullable=False, default=False, server_default="0")
|
||||
onboarding_completed_at = Column(DateTime(timezone=True), nullable=True)
|
||||
contact_email = Column(String(255), nullable=True)
|
||||
preferred_destination = Column(String(50), nullable=True)
|
||||
stripe_customer_id = Column(String(64), nullable=True)
|
||||
|
||||
# UI language preference for i18n (ISO 639-1 code, e.g. "en", "de", "fr")
|
||||
# NULL means "auto-detect from browser Accept-Language header"
|
||||
preferred_language = Column(String(10), nullable=True)
|
||||
|
||||
# UI colour scheme preference: "light" | "dark" | "system" (NULL = "system")
|
||||
preferred_theme = Column(String(10), nullable=True)
|
||||
|
||||
# Custom profile avatar stored as a base64 data-URL (e.g. "data:image/png;base64,...")
|
||||
# NULL means use the Gravatar fallback derived from the user's e-mail address.
|
||||
avatar_data = Column(Text, nullable=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class SubscriptionPlan(Base):
|
||||
"""Dynamically configurable subscription plan stored in the database.
|
||||
|
||||
Plans are shown on the public /pricing page and assigned to users via
|
||||
UserProfile.subscription_tier (which stores plan_id). On first start the
|
||||
four default plans are seeded from TIER_DEFAULTS in app/utils/subscription.py.
|
||||
"""
|
||||
|
||||
__tablename__ = "subscription_plans"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
plan_id = Column(String(50), unique=True, nullable=False, index=True)
|
||||
name = Column(String(100), nullable=False)
|
||||
tagline = Column(String(255), nullable=True)
|
||||
|
||||
# Pricing
|
||||
price_monthly = Column(Float, nullable=False, default=0.0)
|
||||
price_yearly = Column(Float, nullable=False, default=0.0)
|
||||
trial_days = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# Volume limits (0 = unlimited)
|
||||
lifetime_file_limit = Column(Integer, nullable=False, default=0)
|
||||
daily_upload_limit = Column(Integer, nullable=False, default=0)
|
||||
monthly_upload_limit = Column(Integer, nullable=False, default=0)
|
||||
max_storage_destinations = Column(Integer, nullable=False, default=0)
|
||||
max_ocr_pages_monthly = Column(Integer, nullable=False, default=0)
|
||||
max_file_size_mb = Column(Integer, nullable=False, default=0)
|
||||
max_mailboxes = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# Overage configuration
|
||||
overage_percent = Column(Integer, nullable=False, default=20)
|
||||
allow_overage_billing = Column(Boolean, nullable=False, default=False)
|
||||
overage_price_per_doc = Column(Float, nullable=True)
|
||||
overage_price_per_ocr_page = Column(Float, nullable=True)
|
||||
|
||||
# Display / marketing
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
is_highlighted = Column(Boolean, nullable=False, default=False)
|
||||
badge_text = Column(String(50), nullable=True)
|
||||
cta_text = Column(String(100), nullable=False, default="Get started")
|
||||
sort_order = Column(Integer, nullable=False, default=0)
|
||||
features = Column(Text, nullable=True) # JSON-encoded list[str]
|
||||
api_access = Column(Boolean, nullable=False, default=False)
|
||||
stripe_price_id_monthly = Column(String(128), nullable=True)
|
||||
stripe_price_id_yearly = Column(String(128), nullable=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class Pipeline(Base):
|
||||
"""User-defined processing pipeline: an ordered set of steps.
|
||||
|
||||
Pipelines are user-specific. A pipeline with ``owner_id = NULL`` is a
|
||||
*system default* pipeline that only admins may create. Regular users
|
||||
create pipelines under their own ``owner_id``. When a file has no
|
||||
explicit pipeline assigned, the active system default is used.
|
||||
"""
|
||||
|
||||
__tablename__ = "pipelines"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Owner of this pipeline. NULL = system/admin pipeline visible to everyone.
|
||||
owner_id = Column(String, nullable=True, index=True)
|
||||
|
||||
# Human-readable name (unique per owner)
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# Optional description
|
||||
description = Column(Text, nullable=True)
|
||||
|
||||
# When True this pipeline is the default for new files belonging to the owner
|
||||
# (or the global default when owner_id is NULL). Only one pipeline per
|
||||
# owner may be active default at a time — enforced at the application level.
|
||||
is_default = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
# Soft-disable without deleting
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class PipelineStep(Base):
|
||||
"""A single step in a processing pipeline.
|
||||
|
||||
Steps are executed in ascending ``position`` order. Each step has a
|
||||
``step_type`` that maps to a built-in processing action and an optional
|
||||
``config`` JSON blob with step-specific parameters.
|
||||
"""
|
||||
|
||||
__tablename__ = "pipeline_steps"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
pipeline_id = Column(Integer, ForeignKey(_PIPELINES_ID_FK), nullable=False, index=True)
|
||||
|
||||
# Execution order within the pipeline (lower = earlier)
|
||||
position = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# One of the recognised step types (see PIPELINE_STEP_TYPES in pipelines.py)
|
||||
step_type = Column(String(100), nullable=False)
|
||||
|
||||
# Optional human-readable label override (defaults to step_type label)
|
||||
label = Column(String(255), nullable=True)
|
||||
|
||||
# JSON-encoded step-specific configuration dict
|
||||
config = Column(Text, nullable=True)
|
||||
|
||||
# When False this step is skipped during execution
|
||||
enabled = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class ImapIngestionProfile(Base):
|
||||
"""Named ingestion profile controlling which attachment types are accepted from IMAP emails.
|
||||
|
||||
Profiles group file-type categories (e.g. "pdf", "office", "images") so users
|
||||
can precisely control what gets ingested from each mailbox.
|
||||
|
||||
System-provided built-in profiles (``is_builtin=True``) are seeded by the
|
||||
migration and cannot be deleted or renamed. Users may create their own profiles
|
||||
(``owner_id`` set to their identifier) or rely on the global system profiles
|
||||
(``owner_id=None``).
|
||||
|
||||
``allowed_categories`` stores a JSON list of category strings, e.g.::
|
||||
|
||||
'["pdf", "office", "opendocument", "text", "web"]'
|
||||
|
||||
Valid category names are defined in ``app.utils.allowed_types.FILE_TYPE_CATEGORIES``.
|
||||
"""
|
||||
|
||||
__tablename__ = "imap_ingestion_profiles"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Human-readable profile name (e.g. "Documents Only", "Documents + Images")
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# Optional description shown in the UI
|
||||
description = Column(Text, nullable=True)
|
||||
|
||||
# Owner of this profile. NULL = global/system profile available to all users.
|
||||
owner_id = Column(String, nullable=True, index=True)
|
||||
|
||||
# JSON-encoded list of enabled category keys. Example: '["pdf","office","text"]'
|
||||
# See FILE_TYPE_CATEGORIES in app/utils/allowed_types.py for valid values.
|
||||
allowed_categories = Column(Text, nullable=False, default='["pdf","office","opendocument","text","web"]')
|
||||
|
||||
# Built-in system profiles that cannot be deleted or modified via the API.
|
||||
is_builtin = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class UserImapAccount(Base):
|
||||
"""Per-user IMAP ingestion account.
|
||||
|
||||
Each row represents one IMAP mailbox that a user wants DocuElevate to
|
||||
poll for document attachments. The periodic ``pull_all_inboxes`` Celery
|
||||
task iterates over all active accounts and processes any new emails.
|
||||
|
||||
Quota enforcement: the user's subscription plan's ``max_mailboxes`` field
|
||||
controls how many accounts a user may configure (0 = unlimited for paid
|
||||
plans; free tier is not permitted any accounts).
|
||||
"""
|
||||
|
||||
__tablename__ = "user_imap_accounts"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Stable owner identifier — matches FileRecord.owner_id
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# Human-readable label chosen by the user (e.g. "Work Gmail", "Scanner mailbox")
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# IMAP connection settings
|
||||
host = Column(String(255), nullable=False)
|
||||
port = Column(Integer, nullable=False, default=993)
|
||||
username = Column(String(255), nullable=False)
|
||||
# Password stored encrypted using Fernet symmetric encryption via
|
||||
# app.utils.encryption.encrypt_value / decrypt_value (keyed from SESSION_SECRET).
|
||||
# New records are always encrypted; legacy plaintext records are transparently
|
||||
# handled by decrypt_value which returns the value unchanged when no "enc:" prefix
|
||||
# is present.
|
||||
password = Column(String(1024), nullable=False)
|
||||
use_ssl = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Processing options
|
||||
# When True, emails are deleted from the mailbox after their attachments are processed
|
||||
delete_after_process = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
# Optional reference to an ImapIngestionProfile.
|
||||
# NULL means "use the global imap_attachment_filter setting" (system default).
|
||||
profile_id = Column(Integer, ForeignKey("imap_ingestion_profiles.id"), nullable=True)
|
||||
|
||||
# When False the account is not polled by the periodic task (but not deleted)
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Last time this mailbox was successfully polled
|
||||
last_checked_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# Last error message if the most recent poll failed (NULL = last poll succeeded)
|
||||
last_error = Column(Text, nullable=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class BackupRecord(Base):
|
||||
"""Tracks database backup files and their retention metadata.
|
||||
|
||||
Each row represents one backup archive (a gzipped SQLite dump).
|
||||
``backup_type`` classifies the backup for retention purposes:
|
||||
- ``hourly`` – kept for up to 4 days (96 snapshots)
|
||||
- ``daily`` – kept for up to 3 weeks (21 snapshots)
|
||||
- ``weekly`` – kept for up to 13 weeks (≈ 90 days)
|
||||
``local_path`` is the full filesystem path of the local copy (``None``
|
||||
once pruned). ``remote_destination`` and ``remote_path`` describe the
|
||||
remote copy when one has been uploaded to a storage provider or e-mailed.
|
||||
"""
|
||||
|
||||
__tablename__ = "backup_records"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Human-readable archive filename (e.g. backup_hourly_2026-03-07T12-00-00.db.gz)
|
||||
filename = Column(String(255), nullable=False, unique=True)
|
||||
|
||||
# Full path on the local filesystem (may be NULL for remote-only backups)
|
||||
local_path = Column(String(1024), nullable=True)
|
||||
|
||||
# Classification used by the retention policy
|
||||
backup_type = Column(String(20), nullable=False, index=True) # hourly | daily | weekly
|
||||
|
||||
# Archive size in bytes (0 if unknown)
|
||||
size_bytes = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# Checksum of the archive for integrity verification (SHA-256 hex)
|
||||
checksum = Column(String(64), nullable=True)
|
||||
|
||||
# Whether the backup was successfully created
|
||||
status = Column(String(20), nullable=False, default="ok") # ok | failed
|
||||
|
||||
# Storage destination where a remote copy was uploaded (e.g. "s3", "dropbox", "email")
|
||||
remote_destination = Column(String(50), nullable=True)
|
||||
|
||||
# Path / key of the remote copy (bucket key, folder path, etc.)
|
||||
remote_path = Column(String(1024), nullable=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now(), index=True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Integration direction / type constants
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class IntegrationDirection:
|
||||
"""Direction of data flow for a UserIntegration."""
|
||||
|
||||
SOURCE = "SOURCE"
|
||||
DESTINATION = "DESTINATION"
|
||||
|
||||
ALL = {SOURCE, DESTINATION}
|
||||
|
||||
|
||||
class IntegrationType:
|
||||
"""Supported integration types for UserIntegration."""
|
||||
|
||||
# Source integrations (ingestion)
|
||||
IMAP = "IMAP"
|
||||
WATCH_FOLDER = "WATCH_FOLDER"
|
||||
WEBHOOK = "WEBHOOK"
|
||||
|
||||
# Destination integrations (storage / output)
|
||||
S3 = "S3"
|
||||
DROPBOX = "DROPBOX"
|
||||
GOOGLE_DRIVE = "GOOGLE_DRIVE"
|
||||
ONEDRIVE = "ONEDRIVE"
|
||||
WEBDAV = "WEBDAV"
|
||||
NEXTCLOUD = "NEXTCLOUD"
|
||||
FTP = "FTP"
|
||||
SFTP = "SFTP"
|
||||
EMAIL = "EMAIL"
|
||||
PAPERLESS = "PAPERLESS"
|
||||
RCLONE = "RCLONE"
|
||||
ICLOUD = "ICLOUD"
|
||||
|
||||
ALL = {
|
||||
IMAP,
|
||||
WATCH_FOLDER,
|
||||
WEBHOOK,
|
||||
S3,
|
||||
DROPBOX,
|
||||
GOOGLE_DRIVE,
|
||||
ONEDRIVE,
|
||||
WEBDAV,
|
||||
NEXTCLOUD,
|
||||
FTP,
|
||||
SFTP,
|
||||
EMAIL,
|
||||
PAPERLESS,
|
||||
RCLONE,
|
||||
ICLOUD,
|
||||
}
|
||||
|
||||
|
||||
class UserIntegration(Base):
|
||||
"""Generic per-user integration record (source or destination).
|
||||
|
||||
Replaces ad-hoc per-integration-type tables with a single, extensible
|
||||
model that supports any combination of ingestion sources and storage
|
||||
destinations without schema changes when new integrations are added.
|
||||
|
||||
``config`` holds non-sensitive connection settings as a JSON string
|
||||
(e.g. host, port, bucket name, folder path).
|
||||
|
||||
``credentials`` holds sensitive secrets (passwords, tokens, API keys)
|
||||
as a Fernet-encrypted JSON string. Always use
|
||||
``app.utils.encryption.encrypt_value`` / ``decrypt_value`` when
|
||||
writing / reading this field.
|
||||
|
||||
Example config + credentials shapes by integration type:
|
||||
|
||||
IMAP:
|
||||
config = {"host": "imap.example.com", "port": 993,
|
||||
"username": "user@example.com", "use_ssl": true,
|
||||
"delete_after_process": false,
|
||||
"gmail_apply_labels": true}
|
||||
credentials = {"password": "secret"}
|
||||
|
||||
WATCH_FOLDER (local):
|
||||
config = {"source_type": "local",
|
||||
"folder_path": "/data/inbox",
|
||||
"delete_after_process": false}
|
||||
|
||||
WATCH_FOLDER (s3):
|
||||
config = {"source_type": "s3", "bucket": "my-bucket",
|
||||
"region": "us-east-1", "prefix": "inbox/",
|
||||
"endpoint_url": null, "delete_after_process": false}
|
||||
credentials = {"access_key_id": "AKI…", "secret_access_key": "…"}
|
||||
|
||||
WATCH_FOLDER (dropbox):
|
||||
config = {"source_type": "dropbox",
|
||||
"folder_path": "/Inbox/Scanner",
|
||||
"delete_after_process": false}
|
||||
credentials = {"refresh_token": "…", "app_key": "…",
|
||||
"app_secret": "…"}
|
||||
|
||||
WATCH_FOLDER (google_drive):
|
||||
config = {"source_type": "google_drive",
|
||||
"folder_id": "1abc…",
|
||||
"delete_after_process": false}
|
||||
credentials = {"credentials_json": "{…service-account…}"}
|
||||
|
||||
WATCH_FOLDER (onedrive):
|
||||
config = {"source_type": "onedrive",
|
||||
"folder_path": "/Documents/Inbox",
|
||||
"delete_after_process": false}
|
||||
credentials = {"refresh_token": "…", "client_id": "…",
|
||||
"client_secret": "…"}
|
||||
|
||||
WATCH_FOLDER (nextcloud):
|
||||
config = {"source_type": "nextcloud",
|
||||
"url": "https://cloud.example.com",
|
||||
"folder_path": "/Documents/Inbox",
|
||||
"delete_after_process": false}
|
||||
credentials = {"username": "user", "password": "secret"}
|
||||
|
||||
WATCH_FOLDER (webdav):
|
||||
config = {"source_type": "webdav",
|
||||
"url": "https://webdav.example.com/dav/",
|
||||
"folder_path": "/remote.php/webdav/Inbox",
|
||||
"delete_after_process": false}
|
||||
credentials = {"username": "user", "password": "secret"}
|
||||
|
||||
S3:
|
||||
config = {"bucket": "my-bucket", "region": "us-east-1",
|
||||
"endpoint_url": null, "folder_prefix": ""}
|
||||
credentials = {"access_key_id": "AKI…", "secret_access_key": "…"}
|
||||
|
||||
DROPBOX:
|
||||
config = {"folder": "/DocuElevate"}
|
||||
credentials = {"refresh_token": "…", "app_key": "…",
|
||||
"app_secret": "…"}
|
||||
|
||||
GOOGLE_DRIVE:
|
||||
config = {"folder_id": "1abc…"}
|
||||
credentials = {"credentials_json": "{…service-account or OAuth…}"}
|
||||
|
||||
WEBDAV / NEXTCLOUD:
|
||||
config = {"url": "https://cloud.example.com/dav/",
|
||||
"folder": "/Documents"}
|
||||
credentials = {"username": "user", "password": "secret"}
|
||||
"""
|
||||
|
||||
__tablename__ = "user_integrations"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Stable owner identifier — matches FileRecord.owner_id / UserImapAccount.owner_id
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# "SOURCE" or "DESTINATION" (see IntegrationDirection)
|
||||
direction = Column(String(20), nullable=False, index=True)
|
||||
|
||||
# One of the IntegrationType constants (e.g. "IMAP", "S3", "DROPBOX")
|
||||
integration_type = Column(String(50), nullable=False, index=True)
|
||||
|
||||
# Human-readable label chosen by the user (e.g. "Work Gmail", "S3 Archive")
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# Non-sensitive connection configuration (JSON string)
|
||||
config = Column(Text, nullable=True)
|
||||
|
||||
# Sensitive credentials — always stored encrypted via encrypt_value()
|
||||
credentials = Column(Text, nullable=True)
|
||||
|
||||
# When False the integration is not polled / used by background tasks
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Timestamp of the last successful use of this integration
|
||||
last_used_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# Last error message if the most recent operation failed (NULL = last op succeeded)
|
||||
last_error = Column(Text, nullable=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class ApiToken(Base):
|
||||
"""Personal API token for programmatic access.
|
||||
|
||||
Users can create multiple tokens, each with a human-readable name.
|
||||
Only the SHA-256 hash of the token is stored; the plaintext is shown
|
||||
exactly once at creation time. A short prefix (first 8 chars) is
|
||||
persisted for easy identification in the UI.
|
||||
|
||||
Usage tracking records the timestamp and IP address of the most
|
||||
recent request that used the token.
|
||||
"""
|
||||
|
||||
__tablename__ = "api_tokens"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Stable owner identifier — matches FileRecord.owner_id / UserIntegration.owner_id
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# Human-readable label chosen by the user (e.g. "CI Pipeline", "Webhook Upload")
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# SHA-256 hex digest of the full token value
|
||||
token_hash = Column(String(64), nullable=False, unique=True, index=True)
|
||||
|
||||
# First 12 characters of the token for display (e.g. "de_Ab3xY7kL…")
|
||||
token_prefix = Column(String(16), nullable=False)
|
||||
|
||||
# Usage tracking
|
||||
last_used_at = Column(DateTime(timezone=True), nullable=True)
|
||||
last_used_ip = Column(String(45), nullable=True) # IPv6 max length
|
||||
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
revoked_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class SharedLink(Base):
|
||||
"""Shareable, time-limited or view-limited document link.
|
||||
|
||||
A ``SharedLink`` grants unauthenticated access to one ``FileRecord``
|
||||
via a cryptographically random URL token. The link may optionally
|
||||
expire after a given datetime, be limited to a fixed number of views,
|
||||
and require a password. Only a PBKDF2-HMAC-SHA256 hash of the
|
||||
password is stored.
|
||||
|
||||
Owners can view and revoke their active links from the management UI.
|
||||
"""
|
||||
|
||||
__tablename__ = "shared_links"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Unique URL-safe token — forms the public /share/<token> URL.
|
||||
token = Column(String(64), nullable=False, unique=True, index=True)
|
||||
|
||||
# File this link grants access to.
|
||||
file_id = Column(Integer, ForeignKey(_FILES_ID_FK), nullable=False, index=True)
|
||||
|
||||
# Owner who created the link (matches FileRecord.owner_id).
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# Optional human-readable description chosen by the creator.
|
||||
label = Column(String(255), nullable=True)
|
||||
|
||||
# Time-based expiration (NULL = never expires).
|
||||
expires_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
# View-count limit (NULL = unlimited).
|
||||
max_views = Column(Integer, nullable=True)
|
||||
|
||||
# Cumulative view count (incremented on every successful access).
|
||||
view_count = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# Optional password protection — stores PBKDF2-HMAC-SHA256 hex digest.
|
||||
password_hash = Column(String(128), nullable=True)
|
||||
|
||||
# Whether the link is still valid (set to False to revoke immediately).
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
revoked_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class UserNotificationTarget(Base):
|
||||
"""Per-user notification target (email or webhook channel)."""
|
||||
|
||||
__tablename__ = "user_notification_targets"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
channel_type = Column(String(20), nullable=False) # "email" or "webhook"
|
||||
name = Column(String(255), nullable=False) # Human-readable label
|
||||
config = Column(Text, nullable=True) # JSON: smtp config or webhook url
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class UserNotificationPreference(Base):
|
||||
"""Mapping: which user events trigger which notification channel."""
|
||||
|
||||
__tablename__ = "user_notification_preferences"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
event_type = Column(String(50), nullable=False) # "document.processed", "document.failed"
|
||||
channel_type = Column(String(20), nullable=False) # "in_app", "email", "webhook"
|
||||
target_id = Column(Integer, nullable=True) # NULL = in_app, else UserNotificationTarget.id
|
||||
is_enabled = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
__table_args__ = (UniqueConstraint("owner_id", "event_type", "channel_type", "target_id"),)
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class InAppNotification(Base):
|
||||
"""In-app notification record for the bell icon / inbox."""
|
||||
|
||||
__tablename__ = "in_app_notifications"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
event_type = Column(String(50), nullable=False) # "document.processed", "document.failed"
|
||||
title = Column(String(255), nullable=False)
|
||||
message = Column(Text, nullable=True)
|
||||
is_read = Column(Boolean, nullable=False, default=False, index=True)
|
||||
file_id = Column(Integer, nullable=True) # Optional link to FileRecord
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now(), index=True)
|
||||
|
||||
|
||||
class ScheduledJob(Base):
|
||||
"""
|
||||
Configuration record for an admin-managed scheduled batch processing job.
|
||||
|
||||
Each row represents one recurring job entry. The Celery Beat schedule is
|
||||
built from these rows at worker startup; changes take effect after the
|
||||
worker process is restarted.
|
||||
|
||||
Schedule types
|
||||
--------------
|
||||
- ``"cron"`` – standard cron expression fields (minute/hour/…)
|
||||
- ``"interval"`` – fixed interval in seconds (e.g. 3600 for hourly)
|
||||
"""
|
||||
|
||||
__tablename__ = "scheduled_jobs"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Unique machine-readable key used as the Celery Beat schedule entry name.
|
||||
name = Column(String(100), unique=True, nullable=False, index=True)
|
||||
|
||||
# Human-readable display name shown in the admin UI.
|
||||
display_name = Column(String(255), nullable=False)
|
||||
|
||||
# Short description of what the job does.
|
||||
description = Column(Text, nullable=True)
|
||||
|
||||
# Fully-qualified Celery task name, e.g. "app.tasks.batch_tasks.process_new_documents".
|
||||
task_name = Column(String(255), nullable=False)
|
||||
|
||||
# Whether the job is active. Inactive jobs are excluded from the beat schedule.
|
||||
enabled = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Schedule type: "cron" or "interval".
|
||||
schedule_type = Column(String(20), nullable=False, default="cron")
|
||||
|
||||
# --- Cron fields (used when schedule_type == "cron") ---
|
||||
cron_minute = Column(String(50), nullable=False, default="0")
|
||||
cron_hour = Column(String(50), nullable=False, default="*")
|
||||
cron_day_of_week = Column(String(50), nullable=False, default="*")
|
||||
cron_day_of_month = Column(String(50), nullable=False, default="*")
|
||||
cron_month_of_year = Column(String(50), nullable=False, default="*")
|
||||
|
||||
# --- Interval field (used when schedule_type == "interval") ---
|
||||
# Interval in seconds; e.g. 3600 = hourly, 86400 = daily.
|
||||
interval_seconds = Column(Integer, nullable=True)
|
||||
|
||||
# Timestamps populated by the worker after each execution.
|
||||
last_run_at = Column(DateTime(timezone=True), nullable=True)
|
||||
last_run_status = Column(String(20), nullable=True) # "success", "failed", "running"
|
||||
last_run_detail = Column(Text, nullable=True) # Brief result summary or error
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class MobileDevice(Base):
|
||||
"""Registered mobile device for push notifications.
|
||||
|
||||
Stores the push token (Expo push token, FCM token, or APNs token) for a
|
||||
specific user device so that document-processing events can be forwarded
|
||||
as push notifications to the native mobile app.
|
||||
"""
|
||||
|
||||
__tablename__ = "mobile_devices"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# User that owns this device registration.
|
||||
owner_id = Column(String, nullable=False, index=True)
|
||||
|
||||
# Human-readable name the user gave this device (e.g. "John's iPhone").
|
||||
device_name = Column(String(255), nullable=True)
|
||||
|
||||
# Platform: "ios", "android", or "web".
|
||||
platform = Column(String(20), nullable=False, default="ios")
|
||||
|
||||
# Expo push token (ExponentPushToken[…]) or raw FCM/APNs token.
|
||||
push_token = Column(String(512), nullable=False)
|
||||
|
||||
# Whether push notifications are enabled for this device.
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
# Timestamps.
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
last_seen_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
__table_args__ = (UniqueConstraint("owner_id", "push_token", name="uq_mobile_device_owner_token"),)
|
||||
|
||||
|
||||
class ComplianceTemplate(Base):
|
||||
"""Pre-built compliance configuration templates (GDPR, HIPAA, SOC2).
|
||||
|
||||
Each row represents an applied compliance template. The ``settings_json``
|
||||
column stores the concrete setting key/value pairs that were written when
|
||||
the template was applied. ``status`` tracks the current compliance posture.
|
||||
"""
|
||||
|
||||
__tablename__ = "compliance_templates"
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
name = Column(String(50), unique=True, nullable=False, index=True) # GDPR, HIPAA, SOC2
|
||||
display_name = Column(String(100), nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
settings_json = Column(Text, nullable=False, default="{}") # JSON of applied settings
|
||||
enabled = Column(Boolean, nullable=False, default=False)
|
||||
status = Column(String(20), nullable=False, default="not_applied") # not_applied, compliant, partial, non_compliant
|
||||
applied_at = Column(DateTime(timezone=True), nullable=True)
|
||||
applied_by = Column(String(255), nullable=True)
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
|
||||
class PipelineRoutingRule(Base):
|
||||
"""Conditional routing rule that assigns documents to pipelines.
|
||||
|
||||
Rules are evaluated in ascending ``position`` order for a given owner.
|
||||
The first rule whose condition matches the document properties wins and
|
||||
the document is routed to ``target_pipeline_id``. If no rule matches,
|
||||
the caller falls back to the owner's (or system) default pipeline.
|
||||
|
||||
Supported fields:
|
||||
file_type, document_type, category, filename, size, and any key
|
||||
inside the AI-extracted metadata JSON (prefixed ``metadata.``).
|
||||
|
||||
Supported operators:
|
||||
equals, not_equals, contains, not_contains, regex, gt, lt, gte, lte.
|
||||
"""
|
||||
|
||||
__tablename__ = _ROUTING_RULES_TABLE
|
||||
|
||||
id = Column(Integer, primary_key=True, index=True)
|
||||
|
||||
# Owner of this rule. NULL = system-wide rule (admin only).
|
||||
owner_id = Column(String, nullable=True, index=True)
|
||||
|
||||
# Human-readable label for the rule.
|
||||
name = Column(String(255), nullable=False)
|
||||
|
||||
# Evaluation order (lower = earlier). First matching rule wins.
|
||||
position = Column(Integer, nullable=False, default=0)
|
||||
|
||||
# The document property to evaluate.
|
||||
# Built-in: file_type, document_type, category, filename, size.
|
||||
# For AI metadata fields, use the "metadata.<key>" prefix.
|
||||
field = Column(String(255), nullable=False)
|
||||
|
||||
# Comparison operator.
|
||||
operator = Column(String(50), nullable=False)
|
||||
|
||||
# Value to compare against (always stored as text; cast as needed).
|
||||
value = Column(String(1024), nullable=False)
|
||||
|
||||
# Target pipeline when the condition matches.
|
||||
target_pipeline_id = Column(Integer, ForeignKey(_PIPELINES_ID_FK), nullable=False, index=True)
|
||||
|
||||
# Soft-disable without deleting.
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
|
||||
created_at = Column(DateTime(timezone=True), server_default=func.now())
|
||||
updated_at = Column(DateTime(timezone=True), server_default=func.now(), onupdate=func.now())
|
||||
|
||||
@@ -0,0 +1,741 @@
|
||||
"""
|
||||
Backup and restore tasks for DocuElevate.
|
||||
|
||||
Retention strategy
|
||||
------------------
|
||||
- **hourly** backups – retained for 4 days (``backup_retain_hourly``, default 96)
|
||||
- **daily** backups – retained for 3 weeks (``backup_retain_daily``, default 21)
|
||||
- **weekly** backups – retained for 13 weeks (``backup_retain_weekly``, default 13)
|
||||
|
||||
Three separate Celery-beat entries call ``create_backup`` with the appropriate
|
||||
``backup_type`` argument:
|
||||
- every hour → ``create_backup("hourly")``
|
||||
- every day → ``create_backup("daily")``
|
||||
- every week → ``create_backup("weekly")``
|
||||
|
||||
After each backup is created ``_apply_retention`` prunes old local backups for
|
||||
that tier. Remote copies are pruned by ``_prune_remote_backups`` which mirrors
|
||||
the same retention limits.
|
||||
|
||||
Supported database backends
|
||||
----------------------------
|
||||
- **SQLite** – dumped via Python's built-in ``sqlite3.iterdump()``; archive extension ``.db.gz``
|
||||
- **PostgreSQL** – dumped via ``pg_dump --format=plain``; archive extension ``.pgsql.gz``
|
||||
- **MySQL / MariaDB** – dumped via ``mysqldump --single-transaction``; archive extension ``.mysql.gz``
|
||||
"""
|
||||
|
||||
import gzip
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import BackupRecord
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_BACKUP_TYPE_RETAIN: dict[str, str] = {
|
||||
"hourly": "backup_retain_hourly",
|
||||
"daily": "backup_retain_daily",
|
||||
"weekly": "backup_retain_weekly",
|
||||
}
|
||||
|
||||
#: Map of backend name → archive file extension.
|
||||
_BACKEND_EXTENSIONS: dict[str, str] = {
|
||||
"sqlite": ".db.gz",
|
||||
"postgresql": ".pgsql.gz",
|
||||
"mysql": ".mysql.gz",
|
||||
}
|
||||
|
||||
|
||||
def _backup_dir() -> Path:
|
||||
"""Return (and create) the local backup directory."""
|
||||
raw = getattr(settings, "backup_dir", None) or os.path.join(settings.workdir, "backups")
|
||||
path = Path(raw)
|
||||
path.mkdir(parents=True, exist_ok=True)
|
||||
return path
|
||||
|
||||
|
||||
def _db_backend() -> str:
|
||||
"""Return the database backend name (e.g. ``'sqlite'``, ``'postgresql'``, ``'mysql'``)."""
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
url = make_url(settings.database_url)
|
||||
return url.get_backend_name()
|
||||
|
||||
|
||||
def _db_path() -> Path | None:
|
||||
"""Return the SQLite database file path, or None for non-SQLite databases."""
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
url = make_url(settings.database_url)
|
||||
if url.get_backend_name() != "sqlite":
|
||||
return None
|
||||
db = url.database
|
||||
if not db or db == ":memory:":
|
||||
return None
|
||||
return Path(db)
|
||||
|
||||
|
||||
def _archive_ext_for_backend(backend: str) -> str:
|
||||
"""Return the archive file extension for the given database backend.
|
||||
|
||||
Args:
|
||||
backend: Backend name as returned by
|
||||
``sqlalchemy.engine.url.URL.get_backend_name()`` (e.g. ``'sqlite'``).
|
||||
|
||||
Returns:
|
||||
File extension string including the leading dot, e.g. ``'.db.gz'``.
|
||||
Falls back to ``'.sql.gz'`` for unknown backends.
|
||||
"""
|
||||
return _BACKEND_EXTENSIONS.get(backend, ".sql.gz")
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
"""Return the SHA-256 hex digest of *path*."""
|
||||
h = hashlib.sha256()
|
||||
with open(path, "rb") as fh:
|
||||
for chunk in iter(lambda: fh.read(65536), b""):
|
||||
h.update(chunk)
|
||||
return h.hexdigest()
|
||||
|
||||
|
||||
def _dump_sqlite(db_path: Path, dest: Path) -> None:
|
||||
"""Write a gzip-compressed SQL dump of *db_path* to *dest*."""
|
||||
import sqlite3
|
||||
|
||||
conn = sqlite3.connect(str(db_path))
|
||||
try:
|
||||
with gzip.open(str(dest), "wt", encoding="utf-8") as gz:
|
||||
for line in conn.iterdump():
|
||||
gz.write(line + "\n")
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _dump_postgresql(db_url: str, dest: Path) -> None:
|
||||
"""Write a gzip-compressed ``pg_dump`` of the PostgreSQL database to *dest*.
|
||||
|
||||
Uses ``PGPASSWORD`` environment variable so the password is never exposed on
|
||||
the process command line.
|
||||
|
||||
Args:
|
||||
db_url: Full SQLAlchemy database URL (e.g. ``postgresql://user:pass@host/db``).
|
||||
dest: Destination path for the ``.pgsql.gz`` archive.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If ``pg_dump`` exits with a non-zero return code.
|
||||
FileNotFoundError: If the ``pg_dump`` binary is not found.
|
||||
"""
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
url = make_url(db_url)
|
||||
env = os.environ.copy()
|
||||
if url.password:
|
||||
env["PGPASSWORD"] = str(url.password)
|
||||
|
||||
# Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL),
|
||||
# not from user-controlled input. shell=False (the default when passing a list) is used
|
||||
# so there is no shell interpretation of the argument values.
|
||||
cmd: list[str] = ["pg_dump", "--format=plain", "--no-password"]
|
||||
if url.host:
|
||||
cmd.extend(["-h", url.host])
|
||||
if url.port:
|
||||
cmd.extend(["-p", str(url.port)])
|
||||
if url.username:
|
||||
cmd.extend(["-U", url.username])
|
||||
if url.database:
|
||||
cmd.append(url.database)
|
||||
|
||||
with gzip.open(str(dest), "wb") as gz:
|
||||
proc = subprocess.Popen( # noqa: S603
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
env=env,
|
||||
)
|
||||
stdout = proc.stdout
|
||||
if stdout is None: # pragma: no cover – guaranteed by stdout=PIPE
|
||||
raise RuntimeError("pg_dump produced no stdout pipe")
|
||||
try:
|
||||
while True:
|
||||
chunk = stdout.read(65536)
|
||||
if not chunk:
|
||||
break
|
||||
gz.write(chunk)
|
||||
finally:
|
||||
stdout.close()
|
||||
stderr_bytes = proc.stderr.read() if proc.stderr else b""
|
||||
proc.wait()
|
||||
|
||||
if proc.returncode != 0:
|
||||
dest.unlink(missing_ok=True)
|
||||
raise RuntimeError(
|
||||
f"pg_dump exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}"
|
||||
)
|
||||
|
||||
|
||||
def _dump_mysql(db_url: str, dest: Path) -> None:
|
||||
"""Write a gzip-compressed ``mysqldump`` of the MySQL database to *dest*.
|
||||
|
||||
Uses the ``MYSQL_PWD`` environment variable so the password is never exposed
|
||||
on the process command line.
|
||||
|
||||
Args:
|
||||
db_url: Full SQLAlchemy database URL
|
||||
(e.g. ``mysql+pymysql://user:pass@host/db``).
|
||||
dest: Destination path for the ``.mysql.gz`` archive.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If ``mysqldump`` exits with a non-zero return code.
|
||||
FileNotFoundError: If the ``mysqldump`` binary is not found.
|
||||
"""
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
url = make_url(db_url)
|
||||
env = os.environ.copy()
|
||||
if url.password:
|
||||
env["MYSQL_PWD"] = str(url.password)
|
||||
|
||||
# Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL).
|
||||
# shell=False (list form) prevents shell interpretation of argument values.
|
||||
cmd: list[str] = ["mysqldump", "--single-transaction", "--routines", "--triggers"]
|
||||
if url.host:
|
||||
cmd.extend(["-h", url.host])
|
||||
if url.port:
|
||||
cmd.extend(["-P", str(url.port)])
|
||||
if url.username:
|
||||
cmd.extend(["-u", url.username])
|
||||
if url.database:
|
||||
cmd.append(url.database)
|
||||
|
||||
with gzip.open(str(dest), "wb") as gz:
|
||||
proc = subprocess.Popen( # noqa: S603
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
env=env,
|
||||
)
|
||||
stdout = proc.stdout
|
||||
if stdout is None: # pragma: no cover – guaranteed by stdout=PIPE
|
||||
raise RuntimeError("mysqldump produced no stdout pipe")
|
||||
try:
|
||||
while True:
|
||||
chunk = stdout.read(65536)
|
||||
if not chunk:
|
||||
break
|
||||
gz.write(chunk)
|
||||
finally:
|
||||
stdout.close()
|
||||
stderr_bytes = proc.stderr.read() if proc.stderr else b""
|
||||
proc.wait()
|
||||
|
||||
if proc.returncode != 0:
|
||||
dest.unlink(missing_ok=True)
|
||||
raise RuntimeError(
|
||||
f"mysqldump exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}"
|
||||
)
|
||||
|
||||
|
||||
def _restore_sqlite(db_path: Path, archive_path: Path) -> None:
|
||||
"""Restore a SQLite database from a gzip-compressed SQL dump archive.
|
||||
|
||||
Validates the SQL by replaying it on an in-memory database before touching
|
||||
the live file. Saves a ``<db_path>.pre_restore`` rollback copy first.
|
||||
|
||||
Args:
|
||||
db_path: Path to the live SQLite database file to overwrite.
|
||||
archive_path: Path to the ``.db.gz`` gzip-compressed SQL dump.
|
||||
|
||||
Raises:
|
||||
ValueError: If the archive cannot be decompressed or contains invalid SQL.
|
||||
RuntimeError: If writing the restored database fails.
|
||||
"""
|
||||
import shutil
|
||||
import sqlite3
|
||||
|
||||
# Decompress and read SQL statements
|
||||
try:
|
||||
with gzip.open(str(archive_path), "rt", encoding="utf-8") as gz:
|
||||
sql_script = gz.read()
|
||||
except Exception as exc:
|
||||
raise ValueError(f"Failed to decompress backup file: {exc}") from exc
|
||||
|
||||
# Validate by replaying on an in-memory database
|
||||
try:
|
||||
mem_conn = sqlite3.connect(":memory:")
|
||||
mem_conn.executescript(sql_script)
|
||||
mem_conn.close()
|
||||
except sqlite3.Error as exc:
|
||||
raise ValueError(f"Backup file contains invalid SQL: {exc}") from exc
|
||||
|
||||
# Preserve the current DB before overwriting
|
||||
bak = str(db_path) + ".pre_restore"
|
||||
try:
|
||||
shutil.copy2(str(db_path), bak)
|
||||
except OSError as exc:
|
||||
logger.warning(f"Could not create pre-restore backup at {bak}: {exc}")
|
||||
|
||||
try:
|
||||
restore_conn = sqlite3.connect(str(db_path))
|
||||
restore_conn.executescript(sql_script)
|
||||
restore_conn.close()
|
||||
except sqlite3.Error as exc:
|
||||
# Attempt rollback to the pre-restore copy
|
||||
try:
|
||||
if os.path.exists(bak):
|
||||
shutil.copy2(bak, str(db_path))
|
||||
except OSError as rollback_exc:
|
||||
logger.error(f"Rollback failed; database may be corrupted: {rollback_exc}")
|
||||
raise RuntimeError(f"SQLite restore failed: {exc}") from exc
|
||||
|
||||
|
||||
def _restore_postgresql(db_url: str, archive_path: Path) -> None:
|
||||
"""Restore a PostgreSQL database from a gzip-compressed SQL dump archive.
|
||||
|
||||
Pipes the decompressed dump to ``psql``. Uses ``PGPASSWORD`` so the
|
||||
password is never exposed on the process command line.
|
||||
|
||||
Args:
|
||||
db_url: Full SQLAlchemy database URL.
|
||||
archive_path: Path to the ``.pgsql.gz`` gzip-compressed ``pg_dump`` archive.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If ``psql`` exits with a non-zero return code.
|
||||
FileNotFoundError: If the ``psql`` binary is not found.
|
||||
"""
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
url = make_url(db_url)
|
||||
env = os.environ.copy()
|
||||
if url.password:
|
||||
env["PGPASSWORD"] = str(url.password)
|
||||
|
||||
# Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL).
|
||||
# shell=False (list form) prevents shell interpretation of argument values.
|
||||
cmd: list[str] = ["psql", "--no-password"]
|
||||
if url.host:
|
||||
cmd.extend(["-h", url.host])
|
||||
if url.port:
|
||||
cmd.extend(["-p", str(url.port)])
|
||||
if url.username:
|
||||
cmd.extend(["-U", url.username])
|
||||
if url.database:
|
||||
cmd.append(url.database)
|
||||
|
||||
with gzip.open(str(archive_path), "rb") as gz:
|
||||
proc = subprocess.Popen( # noqa: S603
|
||||
cmd,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
env=env,
|
||||
)
|
||||
_, stderr_bytes = proc.communicate(input=gz.read())
|
||||
|
||||
if proc.returncode != 0:
|
||||
raise RuntimeError(f"psql exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}")
|
||||
|
||||
|
||||
def _restore_mysql(db_url: str, archive_path: Path) -> None:
|
||||
"""Restore a MySQL database from a gzip-compressed SQL dump archive.
|
||||
|
||||
Pipes the decompressed dump to ``mysql``. Uses the ``MYSQL_PWD``
|
||||
environment variable so the password is never exposed on the command line.
|
||||
|
||||
Args:
|
||||
db_url: Full SQLAlchemy database URL.
|
||||
archive_path: Path to the ``.mysql.gz`` gzip-compressed ``mysqldump`` archive.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If ``mysql`` exits with a non-zero return code.
|
||||
FileNotFoundError: If the ``mysql`` binary is not found.
|
||||
"""
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
url = make_url(db_url)
|
||||
env = os.environ.copy()
|
||||
if url.password:
|
||||
env["MYSQL_PWD"] = str(url.password)
|
||||
|
||||
# Command arguments are built from the SQLAlchemy URL (admin-configured DATABASE_URL).
|
||||
# shell=False (list form) prevents shell interpretation of argument values.
|
||||
cmd: list[str] = ["mysql"]
|
||||
if url.host:
|
||||
cmd.extend(["-h", url.host])
|
||||
if url.port:
|
||||
cmd.extend(["-P", str(url.port)])
|
||||
if url.username:
|
||||
cmd.extend(["-u", url.username])
|
||||
if url.database:
|
||||
cmd.append(url.database)
|
||||
|
||||
with gzip.open(str(archive_path), "rb") as gz:
|
||||
proc = subprocess.Popen( # noqa: S603
|
||||
cmd,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
env=env,
|
||||
)
|
||||
_, stderr_bytes = proc.communicate(input=gz.read())
|
||||
|
||||
if proc.returncode != 0:
|
||||
raise RuntimeError(f"mysql exited with code {proc.returncode}: {stderr_bytes.decode(errors='replace').strip()}")
|
||||
|
||||
|
||||
def _apply_retention(backup_type: str, db: object) -> None:
|
||||
"""Delete local backups beyond the retention limit for *backup_type*.
|
||||
|
||||
Args:
|
||||
backup_type: One of ``hourly``, ``daily``, ``weekly``.
|
||||
db: Active SQLAlchemy session.
|
||||
"""
|
||||
retain_attr = _BACKUP_TYPE_RETAIN.get(backup_type, "backup_retain_hourly")
|
||||
retain = int(getattr(settings, retain_attr, 96))
|
||||
|
||||
# Query ALL records for this tier (with or without a local file) so that
|
||||
# remote-only and already-pruned records still count toward the retention window.
|
||||
records = (
|
||||
db.query(BackupRecord)
|
||||
.filter(BackupRecord.backup_type == backup_type)
|
||||
.order_by(BackupRecord.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
|
||||
to_prune = records[retain:]
|
||||
for rec in to_prune:
|
||||
if rec.local_path and os.path.exists(rec.local_path):
|
||||
try:
|
||||
os.remove(rec.local_path)
|
||||
logger.info(f"Pruned local backup: {rec.local_path}")
|
||||
except OSError as exc:
|
||||
logger.warning(f"Failed to remove local backup {rec.local_path}: {exc}")
|
||||
rec.local_path = None
|
||||
# If no remote copy either, delete the record entirely
|
||||
if not rec.remote_path:
|
||||
db.delete(rec)
|
||||
|
||||
db.commit()
|
||||
|
||||
|
||||
def _prune_remote_backups(backup_type: str, db: object) -> None:
|
||||
"""Prune remote backup records beyond the retention limit.
|
||||
|
||||
The actual remote deletion is best-effort (logged but not fatal).
|
||||
|
||||
Args:
|
||||
backup_type: One of ``hourly``, ``daily``, ``weekly``.
|
||||
db: Active SQLAlchemy session.
|
||||
"""
|
||||
retain_attr = _BACKUP_TYPE_RETAIN.get(backup_type, "backup_retain_hourly")
|
||||
retain = int(getattr(settings, retain_attr, 96))
|
||||
|
||||
# Query ALL records for this tier so that already-pruned local records
|
||||
# still count toward the retention window.
|
||||
records = (
|
||||
db.query(BackupRecord)
|
||||
.filter(BackupRecord.backup_type == backup_type)
|
||||
.order_by(BackupRecord.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
|
||||
to_prune = [r for r in records[retain:] if r.remote_path]
|
||||
for rec in to_prune:
|
||||
_delete_remote_copy(rec)
|
||||
rec.remote_path = None
|
||||
rec.remote_destination = None
|
||||
if not rec.local_path:
|
||||
db.delete(rec)
|
||||
|
||||
db.commit()
|
||||
|
||||
|
||||
def _delete_remote_copy(rec: BackupRecord) -> None: # noqa: C901
|
||||
"""Best-effort deletion of the remote copy described by *rec*."""
|
||||
dest = rec.remote_destination
|
||||
remote_path = rec.remote_path
|
||||
if not dest or not remote_path:
|
||||
return
|
||||
|
||||
try:
|
||||
if dest == "s3":
|
||||
import boto3
|
||||
|
||||
s3 = boto3.client(
|
||||
"s3",
|
||||
region_name=settings.aws_region,
|
||||
aws_access_key_id=settings.aws_access_key_id,
|
||||
aws_secret_access_key=settings.aws_secret_access_key,
|
||||
)
|
||||
s3.delete_object(Bucket=settings.s3_bucket_name, Key=remote_path)
|
||||
logger.info(f"Deleted remote S3 backup: s3://{settings.s3_bucket_name}/{remote_path}")
|
||||
|
||||
elif dest == "dropbox":
|
||||
import dropbox as dbx_module
|
||||
|
||||
dbx = dbx_module.Dropbox(settings.dropbox_refresh_token)
|
||||
dbx.files_delete_v2(remote_path)
|
||||
logger.info(f"Deleted remote Dropbox backup: {remote_path}")
|
||||
|
||||
elif dest in ("ftp", "sftp", "nextcloud", "webdav", "google_drive", "onedrive", "email"):
|
||||
# For other providers best-effort is logged only – deletion not implemented yet.
|
||||
logger.debug(f"Remote deletion not implemented for destination '{dest}', skipping {remote_path}")
|
||||
|
||||
except Exception as exc:
|
||||
logger.warning(f"Failed to delete remote backup {remote_path} from {dest}: {exc}")
|
||||
|
||||
|
||||
def _upload_remote(archive_path: Path, filename: str) -> tuple[str, str] | None: # noqa: C901
|
||||
"""Upload *archive_path* to the configured remote destination.
|
||||
|
||||
Returns:
|
||||
``(destination, remote_path)`` on success, ``None`` on failure or when
|
||||
no remote destination is configured.
|
||||
"""
|
||||
dest = getattr(settings, "backup_remote_destination", None)
|
||||
if not dest:
|
||||
return None
|
||||
|
||||
remote_folder = getattr(settings, "backup_remote_folder", "backups") or "backups"
|
||||
remote_key = f"{remote_folder}/{filename}"
|
||||
|
||||
try:
|
||||
if dest == "s3":
|
||||
import boto3
|
||||
|
||||
s3 = boto3.client(
|
||||
"s3",
|
||||
region_name=settings.aws_region,
|
||||
aws_access_key_id=settings.aws_access_key_id,
|
||||
aws_secret_access_key=settings.aws_secret_access_key,
|
||||
)
|
||||
with open(archive_path, "rb") as fh:
|
||||
s3.upload_fileobj(fh, settings.s3_bucket_name, remote_key)
|
||||
logger.info(f"Uploaded backup to S3: s3://{settings.s3_bucket_name}/{remote_key}")
|
||||
return (dest, remote_key)
|
||||
|
||||
elif dest == "dropbox":
|
||||
import dropbox as dbx_module
|
||||
|
||||
dbx = dbx_module.Dropbox(settings.dropbox_refresh_token)
|
||||
dropbox_path = f"/{remote_key}"
|
||||
with open(archive_path, "rb") as fh:
|
||||
dbx.files_upload(fh.read(), dropbox_path, mode=dbx_module.files.WriteMode("overwrite"))
|
||||
logger.info(f"Uploaded backup to Dropbox: {dropbox_path}")
|
||||
return (dest, dropbox_path)
|
||||
|
||||
elif dest == "email":
|
||||
_email_backup(archive_path, filename)
|
||||
return (dest, f"email:{filename}")
|
||||
|
||||
elif dest == "nextcloud":
|
||||
import requests
|
||||
|
||||
url = f"{settings.nextcloud_upload_url}/{remote_key}"
|
||||
with open(archive_path, "rb") as fh:
|
||||
resp = requests.put(
|
||||
url,
|
||||
data=fh,
|
||||
auth=(settings.nextcloud_username, settings.nextcloud_password),
|
||||
timeout=120,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
logger.info(f"Uploaded backup to Nextcloud: {url}")
|
||||
return (dest, url)
|
||||
|
||||
elif dest == "webdav":
|
||||
import requests
|
||||
|
||||
url = f"{settings.webdav_url}/{remote_key}"
|
||||
with open(archive_path, "rb") as fh:
|
||||
resp = requests.put(
|
||||
url,
|
||||
data=fh,
|
||||
auth=(settings.webdav_username, settings.webdav_password),
|
||||
verify=settings.webdav_verify_ssl,
|
||||
timeout=120,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
logger.info(f"Uploaded backup to WebDAV: {url}")
|
||||
return (dest, url)
|
||||
|
||||
else:
|
||||
logger.warning(f"Backup remote destination '{dest}' upload not implemented; keeping local only.")
|
||||
return None
|
||||
|
||||
except Exception as exc:
|
||||
logger.error(f"Failed to upload backup to {dest}: {exc}", exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def _email_backup(archive_path: Path, filename: str) -> None:
|
||||
"""Send *archive_path* as an e-mail attachment to the default recipient."""
|
||||
import smtplib
|
||||
from email.mime.application import MIMEApplication
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
|
||||
recipient = settings.email_default_recipient
|
||||
if not recipient:
|
||||
raise ValueError("email_default_recipient is not configured")
|
||||
|
||||
msg = MIMEMultipart()
|
||||
msg["Subject"] = f"[DocuElevate] Database backup – {filename}"
|
||||
msg["From"] = settings.email_sender or settings.email_username or "docuelevate@localhost"
|
||||
msg["To"] = recipient
|
||||
|
||||
body = MIMEText(f"Automated database backup from DocuElevate.\n\nFile: {filename}\n", "plain")
|
||||
msg.attach(body)
|
||||
|
||||
with open(archive_path, "rb") as fh:
|
||||
part = MIMEApplication(fh.read(), Name=filename)
|
||||
part["Content-Disposition"] = f'attachment; filename="{filename}"'
|
||||
msg.attach(part)
|
||||
|
||||
with smtplib.SMTP(settings.email_host, settings.email_port, timeout=60) as server:
|
||||
if settings.email_use_tls:
|
||||
server.starttls()
|
||||
if settings.email_username and settings.email_password:
|
||||
server.login(settings.email_username, settings.email_password)
|
||||
server.sendmail(msg["From"], [recipient], msg.as_string())
|
||||
|
||||
logger.info(f"Backup e-mailed to {recipient}: {filename}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public Celery tasks
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.backup_tasks.create_backup", bind=True)
|
||||
def create_backup(self, backup_type: str = "hourly") -> dict:
|
||||
"""Create a database backup archive and apply retention.
|
||||
|
||||
Supports SQLite (``.db.gz``), PostgreSQL (``.pgsql.gz``), and
|
||||
MySQL / MariaDB (``.mysql.gz``) databases. The native dump tool for the
|
||||
configured backend (``sqlite3``, ``pg_dump``, or ``mysqldump``) must be
|
||||
available on the worker's ``PATH``.
|
||||
|
||||
Args:
|
||||
backup_type: ``"hourly"``, ``"daily"``, or ``"weekly"``.
|
||||
|
||||
Returns:
|
||||
A dict with ``filename``, ``size_bytes``, and ``status``.
|
||||
"""
|
||||
if backup_type not in _BACKUP_TYPE_RETAIN:
|
||||
backup_type = "hourly"
|
||||
|
||||
if not getattr(settings, "backup_enabled", True):
|
||||
logger.debug("Backup is disabled; skipping create_backup task.")
|
||||
return {"status": "disabled"}
|
||||
|
||||
backend = _db_backend()
|
||||
ext = _archive_ext_for_backend(backend)
|
||||
|
||||
ts = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H-%M-%S")
|
||||
filename = f"backup_{backup_type}_{ts}{ext}"
|
||||
archive_path = _backup_dir() / filename
|
||||
|
||||
# SQLite: verify the database file exists before attempting to dump it
|
||||
db_path: Path | None = None
|
||||
if backend == "sqlite":
|
||||
db_path = _db_path()
|
||||
if db_path is None:
|
||||
logger.warning("Backup task skipped: in-memory SQLite databases are not supported.")
|
||||
return {"status": "unsupported_db"}
|
||||
if not db_path.exists():
|
||||
logger.error(f"Database file not found: {db_path}")
|
||||
return {"status": "error", "detail": f"DB file missing: {db_path}"}
|
||||
elif backend not in ("postgresql", "mysql"):
|
||||
logger.warning(f"Backup task skipped: unsupported database backend '{backend}'.")
|
||||
return {"status": "unsupported_db"}
|
||||
|
||||
status = "ok"
|
||||
checksum: str | None = None
|
||||
size_bytes = 0
|
||||
remote_destination: str | None = None
|
||||
remote_path: str | None = None
|
||||
|
||||
try:
|
||||
if backend == "sqlite":
|
||||
# db_path is guaranteed non-None: we returned early if it were None
|
||||
if db_path is None: # pragma: no cover
|
||||
return {"status": "error", "detail": "db_path unexpectedly None"}
|
||||
_dump_sqlite(db_path, archive_path)
|
||||
elif backend == "postgresql":
|
||||
_dump_postgresql(settings.database_url, archive_path)
|
||||
elif backend == "mysql":
|
||||
_dump_mysql(settings.database_url, archive_path)
|
||||
size_bytes = archive_path.stat().st_size
|
||||
checksum = _sha256(archive_path)
|
||||
logger.info(f"Created {backup_type} backup: {archive_path} ({size_bytes:,} bytes)")
|
||||
except Exception as exc:
|
||||
logger.error(f"Failed to create backup archive {filename}: {exc}", exc_info=True)
|
||||
status = "failed"
|
||||
# Record the failure so it is visible in the dashboard
|
||||
with SessionLocal() as db:
|
||||
rec = BackupRecord(
|
||||
filename=filename,
|
||||
local_path=None,
|
||||
backup_type=backup_type,
|
||||
size_bytes=0,
|
||||
checksum=None,
|
||||
status="failed",
|
||||
)
|
||||
db.add(rec)
|
||||
db.commit()
|
||||
return {"status": "error", "detail": str(exc)}
|
||||
|
||||
# Optional remote upload
|
||||
result = _upload_remote(archive_path, filename)
|
||||
if result:
|
||||
remote_destination, remote_path = result
|
||||
|
||||
with SessionLocal() as db:
|
||||
rec = BackupRecord(
|
||||
filename=filename,
|
||||
local_path=str(archive_path),
|
||||
backup_type=backup_type,
|
||||
size_bytes=size_bytes,
|
||||
checksum=checksum,
|
||||
status=status,
|
||||
remote_destination=remote_destination,
|
||||
remote_path=remote_path,
|
||||
)
|
||||
db.add(rec)
|
||||
db.commit()
|
||||
|
||||
# Apply retention policy for this tier
|
||||
_apply_retention(backup_type, db)
|
||||
if remote_destination:
|
||||
_prune_remote_backups(backup_type, db)
|
||||
|
||||
return {
|
||||
"filename": filename,
|
||||
"size_bytes": size_bytes,
|
||||
"status": status,
|
||||
"remote_destination": remote_destination,
|
||||
}
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.backup_tasks.cleanup_old_backups")
|
||||
def cleanup_old_backups() -> dict:
|
||||
"""Manually trigger retention clean-up for all backup tiers.
|
||||
|
||||
This is also called automatically after each ``create_backup`` run.
|
||||
"""
|
||||
with SessionLocal() as db:
|
||||
for btype in ("hourly", "daily", "weekly"):
|
||||
_apply_retention(btype, db)
|
||||
_prune_remote_backups(btype, db)
|
||||
return {"status": "ok"}
|
||||
@@ -0,0 +1,669 @@
|
||||
"""
|
||||
Scheduled batch processing tasks for DocuElevate.
|
||||
|
||||
This module provides Celery tasks that can be scheduled via Celery Beat
|
||||
and managed through the admin UI (``/admin/scheduled-jobs``):
|
||||
|
||||
Core batch jobs
|
||||
---------------
|
||||
- ``process_new_documents`` – Queue documents that have never been processed.
|
||||
- ``reprocess_failed_documents`` – Re-queue documents whose processing failed.
|
||||
- ``cleanup_temp_files`` – Remove stale files from the ``workdir/tmp`` directory.
|
||||
|
||||
Maintenance / housekeeping jobs
|
||||
--------------------------------
|
||||
- ``expire_shared_links`` – Auto-revoke SharedLinks whose ``expires_at`` has passed.
|
||||
- ``prune_processing_logs`` – Delete old rows from ``processing_logs`` and
|
||||
``settings_audit_log`` to prevent unbounded table growth.
|
||||
- ``prune_old_notifications`` – Delete old read ``in_app_notifications`` rows.
|
||||
- ``backfill_missing_metadata`` – Re-trigger AI metadata extraction for completed files
|
||||
that have OCR text but no ``ai_metadata``.
|
||||
- ``sync_search_index`` – Index documents in Meilisearch that have OCR text /
|
||||
metadata but are not yet in the search index.
|
||||
|
||||
Each task records its execution result back to the ``ScheduledJob`` table so
|
||||
the admin UI can display last-run times and statuses.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import (
|
||||
FileProcessingStep,
|
||||
FileRecord,
|
||||
InAppNotification,
|
||||
ProcessingLog,
|
||||
ScheduledJob,
|
||||
SettingsAuditLog,
|
||||
SharedLink,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_PROCESSING_STEPS = {
|
||||
"create_file_record",
|
||||
"check_text",
|
||||
"extract_text",
|
||||
"process_with_ocr",
|
||||
"process_with_azure_document_intelligence",
|
||||
"extract_metadata_with_gpt",
|
||||
"embed_metadata_into_pdf",
|
||||
"finalize_document_storage",
|
||||
"send_to_all_destinations",
|
||||
}
|
||||
|
||||
|
||||
def _update_job_status(job_name: str, status: str, detail: str) -> None:
|
||||
"""Persist run status back to the ScheduledJob row for display in the UI."""
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
job = db.query(ScheduledJob).filter(ScheduledJob.name == job_name).first()
|
||||
if job:
|
||||
job.last_run_at = datetime.now(timezone.utc)
|
||||
job.last_run_status = status
|
||||
job.last_run_detail = detail
|
||||
db.commit()
|
||||
except Exception as exc: # pragma: no cover – best-effort status update
|
||||
logger.warning("Could not update ScheduledJob status for %s: %s", job_name, exc)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: process new (unprocessed) documents
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.process_new_documents")
|
||||
def process_new_documents() -> dict:
|
||||
"""
|
||||
Queue all documents that have never been processed.
|
||||
|
||||
A document is considered *new* when it has no ``FileProcessingStep`` rows
|
||||
that match the core pipeline steps. The task loads each qualifying
|
||||
``FileRecord``, verifies that the original file still exists on disk, and
|
||||
dispatches ``process_document`` for each one.
|
||||
|
||||
Returns a summary dict with ``queued`` and ``skipped`` counts.
|
||||
"""
|
||||
from app.tasks.process_document import process_document # avoid circular import
|
||||
|
||||
job_name = "process-new-documents"
|
||||
logger.info("[batch] Starting process_new_documents task")
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
# Files that already have at least one processing step recorded.
|
||||
processed_file_ids = (
|
||||
db.query(FileProcessingStep.file_id)
|
||||
.filter(FileProcessingStep.step_name.in_(_PROCESSING_STEPS))
|
||||
.distinct()
|
||||
.subquery()
|
||||
)
|
||||
|
||||
# Candidate files: non-duplicate records with no processing steps yet.
|
||||
candidates = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.is_duplicate.is_(False))
|
||||
.filter(~FileRecord.id.in_(db.query(processed_file_ids.c.file_id)))
|
||||
.all()
|
||||
)
|
||||
|
||||
queued = 0
|
||||
skipped = 0
|
||||
for record in candidates:
|
||||
if not record.local_filename or not os.path.exists(record.local_filename):
|
||||
logger.warning(
|
||||
"[batch] Skipping file_id=%s — local file not found: %s",
|
||||
record.id,
|
||||
record.local_filename,
|
||||
)
|
||||
skipped += 1
|
||||
continue
|
||||
process_document.delay(
|
||||
record.local_filename,
|
||||
original_filename=record.original_filename,
|
||||
file_id=record.id,
|
||||
owner_id=record.owner_id,
|
||||
)
|
||||
queued += 1
|
||||
|
||||
detail = f"Queued {queued} document(s) for processing; skipped {skipped} (file not on disk)."
|
||||
logger.info("[batch] process_new_documents: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"queued": queued, "skipped": skipped}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] process_new_documents failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"queued": 0, "skipped": 0, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: reprocess failed documents
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.reprocess_failed_documents")
|
||||
def reprocess_failed_documents() -> dict:
|
||||
"""
|
||||
Re-queue documents whose most-recent processing attempt failed.
|
||||
|
||||
Only files that have at least one ``FileProcessingStep`` with
|
||||
``status == "failure"`` **and** no currently ``in_progress`` steps are
|
||||
selected so that actively-running jobs are not interrupted.
|
||||
|
||||
Returns a summary dict with ``queued`` and ``skipped`` counts.
|
||||
"""
|
||||
from app.tasks.process_document import process_document # avoid circular import
|
||||
|
||||
job_name = "reprocess-failed-documents"
|
||||
logger.info("[batch] Starting reprocess_failed_documents task")
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
# Files with at least one failed step.
|
||||
failed_file_ids = (
|
||||
db.query(FileProcessingStep.file_id)
|
||||
.filter(FileProcessingStep.step_name.in_(_PROCESSING_STEPS))
|
||||
.filter(FileProcessingStep.status == "failure")
|
||||
.distinct()
|
||||
.subquery()
|
||||
)
|
||||
|
||||
# Exclude files that are currently being processed.
|
||||
in_progress_file_ids = (
|
||||
db.query(FileProcessingStep.file_id)
|
||||
.filter(FileProcessingStep.status == "in_progress")
|
||||
.distinct()
|
||||
.subquery()
|
||||
)
|
||||
|
||||
candidates = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.is_duplicate.is_(False))
|
||||
.filter(FileRecord.id.in_(db.query(failed_file_ids.c.file_id)))
|
||||
.filter(~FileRecord.id.in_(db.query(in_progress_file_ids.c.file_id)))
|
||||
.all()
|
||||
)
|
||||
|
||||
queued = 0
|
||||
skipped = 0
|
||||
for record in candidates:
|
||||
if not record.local_filename or not os.path.exists(record.local_filename):
|
||||
logger.warning(
|
||||
"[batch] Skipping file_id=%s — local file not found: %s",
|
||||
record.id,
|
||||
record.local_filename,
|
||||
)
|
||||
skipped += 1
|
||||
continue
|
||||
process_document.delay(
|
||||
record.local_filename,
|
||||
original_filename=record.original_filename,
|
||||
file_id=record.id,
|
||||
owner_id=record.owner_id,
|
||||
)
|
||||
queued += 1
|
||||
|
||||
detail = f"Re-queued {queued} failed document(s); skipped {skipped} (file not on disk)."
|
||||
logger.info("[batch] reprocess_failed_documents: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"queued": queued, "skipped": skipped}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] reprocess_failed_documents failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"queued": 0, "skipped": 0, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: clean up temporary files
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Files in ``workdir/tmp`` that are older than this threshold are deleted.
|
||||
_TEMP_FILE_MAX_AGE_HOURS: int = 24
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.cleanup_temp_files")
|
||||
def cleanup_temp_files(max_age_hours: int = _TEMP_FILE_MAX_AGE_HOURS) -> dict:
|
||||
"""
|
||||
Delete stale files from the ``workdir/tmp`` directory.
|
||||
|
||||
A file is considered stale when **both** of the following are true:
|
||||
|
||||
1. Its modification time is older than *max_age_hours* (default 24 h).
|
||||
2. No ``FileRecord.local_filename`` points to it **or** the file is not
|
||||
referenced by any active in-progress processing step.
|
||||
|
||||
This prevents accidental deletion of files that are being actively
|
||||
processed by the pipeline.
|
||||
|
||||
Args:
|
||||
max_age_hours: Minimum age (in hours) before a temp file is eligible
|
||||
for deletion. Defaults to 24.
|
||||
|
||||
Returns:
|
||||
A summary dict with ``deleted`` and ``skipped`` counts.
|
||||
"""
|
||||
job_name = "cleanup-temp-files"
|
||||
logger.info("[batch] Starting cleanup_temp_files (max_age_hours=%s)", max_age_hours)
|
||||
|
||||
tmp_dir = Path(settings.workdir) / "tmp"
|
||||
if not tmp_dir.exists():
|
||||
detail = "workdir/tmp does not exist; nothing to clean."
|
||||
logger.info("[batch] cleanup_temp_files: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"deleted": 0, "skipped": 0}
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(hours=max_age_hours)
|
||||
|
||||
deleted = 0
|
||||
skipped = 0
|
||||
errors = 0
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
# Collect filenames actively referenced by in-progress processing steps.
|
||||
in_progress_filenames: set[str] = set()
|
||||
in_progress_records = (
|
||||
db.query(FileRecord.local_filename)
|
||||
.join(FileProcessingStep, FileProcessingStep.file_id == FileRecord.id)
|
||||
.filter(FileProcessingStep.status == "in_progress")
|
||||
.distinct()
|
||||
.all()
|
||||
)
|
||||
for row in in_progress_records:
|
||||
if row.local_filename:
|
||||
in_progress_filenames.add(os.path.basename(row.local_filename))
|
||||
|
||||
# Also collect all filenames referenced by FileRecord.local_filename
|
||||
# that point into workdir/tmp (files still in the tmp pipeline).
|
||||
active_tmp_filenames: set[str] = set()
|
||||
tmp_dir_str = str(tmp_dir.resolve())
|
||||
active_records = (
|
||||
db.query(FileRecord.local_filename).filter(FileRecord.local_filename.like(f"{tmp_dir_str}%")).all()
|
||||
)
|
||||
for row in active_records:
|
||||
if row.local_filename:
|
||||
active_tmp_filenames.add(os.path.basename(row.local_filename))
|
||||
|
||||
protected_basenames = in_progress_filenames | active_tmp_filenames
|
||||
|
||||
for entry in tmp_dir.iterdir():
|
||||
if not entry.is_file():
|
||||
continue
|
||||
|
||||
# Check modification time.
|
||||
try:
|
||||
mtime = datetime.fromtimestamp(entry.stat().st_mtime, tz=timezone.utc)
|
||||
except OSError: # pragma: no cover – only reachable if file vanishes between iterdir() and stat()
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
if mtime >= cutoff:
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
if entry.name in protected_basenames:
|
||||
logger.debug("[batch] cleanup_temp_files: keeping protected file %s", entry.name)
|
||||
skipped += 1
|
||||
continue
|
||||
|
||||
try:
|
||||
entry.unlink()
|
||||
logger.debug("[batch] cleanup_temp_files: deleted %s", entry)
|
||||
deleted += 1
|
||||
except OSError as exc:
|
||||
logger.warning("[batch] cleanup_temp_files: could not delete %s: %s", entry, exc)
|
||||
errors += 1
|
||||
|
||||
detail = f"Deleted {deleted} stale temp file(s); skipped {skipped} (too new or protected); {errors} error(s)."
|
||||
status = "failed" if errors and not deleted else "success"
|
||||
logger.info("[batch] cleanup_temp_files: %s", detail)
|
||||
_update_job_status(job_name, status, detail)
|
||||
return {"deleted": deleted, "skipped": skipped, "errors": errors}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] cleanup_temp_files failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"deleted": 0, "skipped": 0, "errors": 1, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: expire stale shared links
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.expire_shared_links")
|
||||
def expire_shared_links() -> dict:
|
||||
"""
|
||||
Auto-revoke SharedLinks whose ``expires_at`` timestamp has passed.
|
||||
|
||||
The ``_is_link_valid`` helper in the shared-links API already blocks
|
||||
access at request time, but the database rows remain flagged as
|
||||
``is_active=True``. This task sweeps those rows and sets
|
||||
``is_active=False`` + ``revoked_at`` so the management UI reflects
|
||||
the true state and counts are accurate.
|
||||
|
||||
Returns a summary dict with ``revoked`` count.
|
||||
"""
|
||||
job_name = "expire-shared-links"
|
||||
logger.info("[batch] Starting expire_shared_links task")
|
||||
|
||||
try:
|
||||
now = datetime.now(timezone.utc)
|
||||
with SessionLocal() as db:
|
||||
stale = (
|
||||
db.query(SharedLink)
|
||||
.filter(
|
||||
SharedLink.is_active.is_(True),
|
||||
SharedLink.expires_at.isnot(None),
|
||||
SharedLink.expires_at < now,
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for link in stale:
|
||||
link.is_active = False
|
||||
link.revoked_at = now
|
||||
db.commit()
|
||||
revoked = len(stale)
|
||||
|
||||
detail = f"Revoked {revoked} expired shared link(s)."
|
||||
logger.info("[batch] expire_shared_links: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"revoked": revoked}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] expire_shared_links failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"revoked": 0, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: prune old processing logs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Default retention period for processing logs and audit log rows.
|
||||
_LOG_RETENTION_DAYS: int = 30
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.prune_processing_logs")
|
||||
def prune_processing_logs(retention_days: int = _LOG_RETENTION_DAYS) -> dict:
|
||||
"""
|
||||
Delete ``processing_logs`` and ``settings_audit_log`` rows older than
|
||||
*retention_days* (default 30) to prevent unbounded table growth.
|
||||
|
||||
Rows for the most recent *retention_days* days are kept so that recent
|
||||
activity is still visible in the logs/audit UI.
|
||||
|
||||
Args:
|
||||
retention_days: Number of days of history to keep (default 30).
|
||||
|
||||
Returns:
|
||||
A summary dict with ``processing_logs_deleted`` and
|
||||
``audit_log_deleted`` counts.
|
||||
"""
|
||||
job_name = "prune-processing-logs"
|
||||
logger.info("[batch] Starting prune_processing_logs (retention_days=%s)", retention_days)
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=retention_days)
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
pl_deleted = db.query(ProcessingLog).filter(ProcessingLog.timestamp < cutoff).delete()
|
||||
al_deleted = db.query(SettingsAuditLog).filter(SettingsAuditLog.changed_at < cutoff).delete()
|
||||
db.commit()
|
||||
|
||||
detail = (
|
||||
f"Deleted {pl_deleted} processing log row(s) and "
|
||||
f"{al_deleted} settings audit log row(s) older than {retention_days} days."
|
||||
)
|
||||
logger.info("[batch] prune_processing_logs: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"processing_logs_deleted": pl_deleted, "audit_log_deleted": al_deleted}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] prune_processing_logs failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"processing_logs_deleted": 0, "audit_log_deleted": 0, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: prune old in-app notifications
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Default retention period for read notifications.
|
||||
_NOTIFICATION_RETENTION_DAYS: int = 30
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.prune_old_notifications")
|
||||
def prune_old_notifications(retention_days: int = _NOTIFICATION_RETENTION_DAYS) -> dict:
|
||||
"""
|
||||
Delete ``in_app_notifications`` rows that are already read and older than
|
||||
*retention_days* days (default 30) to prevent unbounded table growth.
|
||||
|
||||
Unread notifications are always kept regardless of age so users do not
|
||||
miss important alerts.
|
||||
|
||||
Args:
|
||||
retention_days: Number of days of read-notification history to keep
|
||||
(default 30).
|
||||
|
||||
Returns:
|
||||
A summary dict with ``deleted`` count.
|
||||
"""
|
||||
job_name = "prune-old-notifications"
|
||||
logger.info("[batch] Starting prune_old_notifications (retention_days=%s)", retention_days)
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=retention_days)
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
deleted = (
|
||||
db.query(InAppNotification)
|
||||
.filter(
|
||||
InAppNotification.is_read.is_(True),
|
||||
InAppNotification.created_at < cutoff,
|
||||
)
|
||||
.delete()
|
||||
)
|
||||
db.commit()
|
||||
|
||||
detail = f"Deleted {deleted} old read notification(s) older than {retention_days} days."
|
||||
logger.info("[batch] prune_old_notifications: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"deleted": deleted}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] prune_old_notifications failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"deleted": 0, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: backfill missing AI metadata
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Maximum number of files to process per backfill run.
|
||||
_METADATA_BACKFILL_BATCH_SIZE: int = 50
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.backfill_missing_metadata")
|
||||
def backfill_missing_metadata(batch_size: int = _METADATA_BACKFILL_BATCH_SIZE) -> dict:
|
||||
"""
|
||||
Re-trigger AI metadata extraction for documents that have OCR text but
|
||||
no ``ai_metadata``.
|
||||
|
||||
This handles the common case where a document was processed before the AI
|
||||
metadata extraction step was configured (e.g., before an OpenAI API key
|
||||
was added), or where the extraction previously failed.
|
||||
|
||||
Only files that are **not** currently in-progress and have non-empty
|
||||
``ocr_text`` are selected. A configurable *batch_size* caps the number
|
||||
of tasks queued per run to avoid overwhelming the AI provider.
|
||||
|
||||
Args:
|
||||
batch_size: Maximum number of files to queue per run (default 50).
|
||||
|
||||
Returns:
|
||||
A summary dict with ``queued`` count.
|
||||
"""
|
||||
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt # avoid circular import
|
||||
|
||||
job_name = "backfill-missing-metadata"
|
||||
logger.info("[batch] Starting backfill_missing_metadata (batch_size=%s)", batch_size)
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
# Files currently being processed — skip them.
|
||||
in_progress_file_ids = (
|
||||
db.query(FileProcessingStep.file_id)
|
||||
.filter(FileProcessingStep.status == "in_progress")
|
||||
.distinct()
|
||||
.subquery()
|
||||
)
|
||||
|
||||
candidates = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.is_duplicate.is_(False))
|
||||
.filter(FileRecord.ocr_text.isnot(None))
|
||||
.filter(FileRecord.ocr_text != "")
|
||||
.filter((FileRecord.ai_metadata.is_(None)) | (FileRecord.ai_metadata == ""))
|
||||
.filter(~FileRecord.id.in_(db.query(in_progress_file_ids.c.file_id)))
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
queued = 0
|
||||
for record in candidates:
|
||||
filename = record.local_filename or record.original_filename or f"file_{record.id}"
|
||||
extract_metadata_with_gpt.delay(
|
||||
filename,
|
||||
record.ocr_text,
|
||||
file_id=record.id,
|
||||
)
|
||||
queued += 1
|
||||
|
||||
detail = f"Queued {queued} document(s) for AI metadata backfill."
|
||||
logger.info("[batch] backfill_missing_metadata: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"queued": queued}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] backfill_missing_metadata failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"queued": 0, "error": str(exc)}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task: sync Meilisearch search index
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Maximum documents to index per sync run.
|
||||
_SEARCH_SYNC_BATCH_SIZE: int = 100
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.batch_tasks.sync_search_index")
|
||||
def sync_search_index(batch_size: int = _SEARCH_SYNC_BATCH_SIZE) -> dict:
|
||||
"""
|
||||
Index documents in Meilisearch that have OCR text or AI metadata but are
|
||||
not yet present in the search index.
|
||||
|
||||
This is useful after:
|
||||
- Enabling Meilisearch for the first time on an existing installation.
|
||||
- Recovering from a Meilisearch index wipe or migration.
|
||||
- Documents processed before search indexing was added to the pipeline.
|
||||
|
||||
The task queries the Meilisearch index for existing document IDs, then
|
||||
finds ``FileRecord`` rows that have processable content (``ocr_text`` or
|
||||
``ai_metadata``) but are absent from the index, and re-indexes them.
|
||||
|
||||
A configurable *batch_size* caps the number of documents indexed per run.
|
||||
|
||||
Args:
|
||||
batch_size: Maximum number of documents to index per run (default 100).
|
||||
|
||||
Returns:
|
||||
A summary dict with ``indexed`` and ``skipped`` counts.
|
||||
"""
|
||||
from app.utils.meilisearch_client import get_meilisearch_client, index_document
|
||||
|
||||
job_name = "sync-search-index"
|
||||
logger.info("[batch] Starting sync_search_index (batch_size=%s)", batch_size)
|
||||
|
||||
client = get_meilisearch_client()
|
||||
if client is None:
|
||||
detail = "Meilisearch is not configured; skipping search index sync."
|
||||
logger.info("[batch] sync_search_index: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"indexed": 0, "skipped": 0, "reason": "meilisearch_not_configured"}
|
||||
|
||||
try:
|
||||
# Fetch the set of file_ids already in the Meilisearch index.
|
||||
index = client.get_index(settings.meilisearch_index_name)
|
||||
# Fetch up to 10 000 IDs — sufficient to determine gaps for most installs.
|
||||
existing_result = index.get_documents({"fields": ["file_id"], "limit": 10000})
|
||||
existing_ids: set[int] = {doc["file_id"] for doc in existing_result.results if "file_id" in doc}
|
||||
except Exception as exc:
|
||||
detail = f"Error fetching existing Meilisearch IDs: {exc}"
|
||||
logger.error("[batch] sync_search_index: %s", detail)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"indexed": 0, "skipped": 0, "error": str(exc)}
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
# Files with indexable content that are not already in the index.
|
||||
candidates = (
|
||||
db.query(FileRecord)
|
||||
.filter(FileRecord.is_duplicate.is_(False))
|
||||
.filter(
|
||||
(FileRecord.ocr_text.isnot(None) & (FileRecord.ocr_text != ""))
|
||||
| (FileRecord.ai_metadata.isnot(None) & (FileRecord.ai_metadata != ""))
|
||||
)
|
||||
.filter(~FileRecord.id.in_(existing_ids) if existing_ids else True) # type: ignore[arg-type]
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
indexed = 0
|
||||
skipped = 0
|
||||
for record in candidates:
|
||||
metadata: dict = {}
|
||||
if record.ai_metadata:
|
||||
try:
|
||||
metadata = json.loads(record.ai_metadata)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
pass
|
||||
|
||||
success = index_document(record, record.ocr_text or "", metadata)
|
||||
if success:
|
||||
indexed += 1
|
||||
else:
|
||||
skipped += 1
|
||||
|
||||
detail = f"Indexed {indexed} document(s) into Meilisearch; {skipped} skipped (indexing error)."
|
||||
logger.info("[batch] sync_search_index: %s", detail)
|
||||
_update_job_status(job_name, "success", detail)
|
||||
return {"indexed": indexed, "skipped": skipped}
|
||||
|
||||
except Exception as exc:
|
||||
detail = f"Error: {exc}"
|
||||
logger.error("[batch] sync_search_index failed: %s", exc, exc_info=True)
|
||||
_update_job_status(job_name, "failed", detail)
|
||||
return {"indexed": 0, "skipped": 0, "error": str(exc)}
|
||||
@@ -0,0 +1,174 @@
|
||||
"""Celery task for pre-computing document text embeddings.
|
||||
|
||||
Runs after document processing to ensure embeddings are available for
|
||||
the similarity feature without requiring a user to trigger them on first
|
||||
access.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.step_manager import update_step_status
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True, name="compute_document_embedding")
|
||||
def compute_document_embedding(self, file_id: int) -> dict:
|
||||
"""Compute and cache the text embedding for a single document.
|
||||
|
||||
Skips silently when the file has no OCR text or already has a cached
|
||||
embedding. The result is stored in ``FileRecord.embedding`` for
|
||||
subsequent similarity queries.
|
||||
|
||||
Args:
|
||||
file_id: Primary key of the :class:`~app.models.FileRecord`.
|
||||
|
||||
Returns:
|
||||
A dict with ``status`` (``"success"`` / ``"skipped"`` / ``"error"``)
|
||||
and optional ``detail`` message.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info("[%s] Computing embedding for file %s", task_id, file_id)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"compute_embedding",
|
||||
"in_progress",
|
||||
f"Computing text embedding for file {file_id}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
with SessionLocal() as db:
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not file_record:
|
||||
logger.warning("[%s] File %s not found, skipping embedding", task_id, file_id)
|
||||
return {"status": "skipped", "detail": "File not found"}
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
update_step_status(db, file_id, "compute_embedding", "in_progress", started_at=now)
|
||||
|
||||
# Already has a cached embedding – nothing to do
|
||||
if file_record.embedding:
|
||||
logger.info("[%s] File %s already has a cached embedding", task_id, file_id)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"compute_embedding",
|
||||
"success",
|
||||
"Embedding already cached",
|
||||
file_id=file_id,
|
||||
)
|
||||
update_step_status(db, file_id, "compute_embedding", "success", completed_at=now)
|
||||
return {"status": "skipped", "detail": "Embedding already cached"}
|
||||
|
||||
if not file_record.ocr_text or not file_record.ocr_text.strip():
|
||||
logger.info("[%s] File %s has no OCR text, skipping embedding", task_id, file_id)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"compute_embedding",
|
||||
"skipped",
|
||||
"No OCR text available",
|
||||
file_id=file_id,
|
||||
)
|
||||
update_step_status(db, file_id, "compute_embedding", "skipped", completed_at=now)
|
||||
return {"status": "skipped", "detail": "No OCR text available"}
|
||||
|
||||
try:
|
||||
from app.utils.similarity import compute_and_store_embedding
|
||||
|
||||
embedding = compute_and_store_embedding(db, file_record)
|
||||
completed = datetime.now(timezone.utc)
|
||||
if embedding:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"compute_embedding",
|
||||
"success",
|
||||
f"Embedding computed ({len(embedding)} dimensions)",
|
||||
file_id=file_id,
|
||||
)
|
||||
update_step_status(db, file_id, "compute_embedding", "success", completed_at=completed)
|
||||
return {
|
||||
"status": "success",
|
||||
"detail": f"Embedding computed ({len(embedding)} dimensions)",
|
||||
}
|
||||
else:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"compute_embedding",
|
||||
"failure",
|
||||
"Embedding computation returned None",
|
||||
file_id=file_id,
|
||||
)
|
||||
update_step_status(
|
||||
db,
|
||||
file_id,
|
||||
"compute_embedding",
|
||||
"failure",
|
||||
error_message="Embedding computation returned None",
|
||||
completed_at=completed,
|
||||
)
|
||||
return {"status": "error", "detail": "Embedding computation returned None"}
|
||||
except Exception as exc:
|
||||
logger.exception("[%s] Embedding computation failed for file %s: %s", task_id, file_id, exc)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"compute_embedding",
|
||||
"failure",
|
||||
f"Exception: {exc}",
|
||||
file_id=file_id,
|
||||
)
|
||||
update_step_status(
|
||||
db,
|
||||
file_id,
|
||||
"compute_embedding",
|
||||
"failure",
|
||||
error_message=str(exc),
|
||||
completed_at=datetime.now(timezone.utc),
|
||||
)
|
||||
return {"status": "error", "detail": str(exc)}
|
||||
|
||||
|
||||
@celery.task(bind=True, name="backfill_missing_embeddings")
|
||||
def backfill_missing_embeddings(self) -> dict:
|
||||
"""Periodic task that computes embeddings for documents that lack them.
|
||||
|
||||
Iterates over all ``FileRecord`` rows that have OCR text but no
|
||||
cached embedding and queues a :func:`compute_document_embedding`
|
||||
task for each one. A configurable ``batch_size`` caps the number
|
||||
of tasks queued per run to avoid overwhelming the worker or the
|
||||
embedding API.
|
||||
|
||||
Returns:
|
||||
A dict with the number of tasks ``queued``.
|
||||
"""
|
||||
batch_size = settings.embedding_backfill_batch_size
|
||||
task_id = self.request.id
|
||||
logger.info("[%s] Backfill: scanning for files missing embeddings (batch_size=%d)", task_id, batch_size)
|
||||
|
||||
with SessionLocal() as db:
|
||||
candidates = (
|
||||
db.query(FileRecord.id)
|
||||
.filter(
|
||||
FileRecord.ocr_text.isnot(None),
|
||||
FileRecord.ocr_text != "",
|
||||
(FileRecord.embedding.is_(None)) | (FileRecord.embedding == ""),
|
||||
)
|
||||
.limit(batch_size)
|
||||
.all()
|
||||
)
|
||||
|
||||
queued = 0
|
||||
for (file_id,) in candidates:
|
||||
try:
|
||||
compute_document_embedding.delay(file_id)
|
||||
queued += 1
|
||||
except Exception as exc:
|
||||
logger.warning("[%s] Could not queue embedding for file %s: %s", task_id, file_id, exc)
|
||||
|
||||
logger.info("[%s] Backfill: queued %d embedding tasks", task_id, queued)
|
||||
return {"queued": queued}
|
||||
@@ -124,7 +124,9 @@ def _build_filename(file_path: str, original_filename: Optional[str], file_ext:
|
||||
|
||||
|
||||
@shared_task(bind=True)
|
||||
def convert_to_pdf(self, file_path: str, original_filename: Optional[str] = None) -> Optional[str]:
|
||||
def convert_to_pdf(
|
||||
self, file_path: str, original_filename: Optional[str] = None, owner_id: Optional[str] = None
|
||||
) -> Optional[str]:
|
||||
"""
|
||||
Converts a file to PDF using Gotenberg's API.
|
||||
Determines the appropriate Gotenberg endpoint based on the file's MIME type.
|
||||
@@ -133,6 +135,7 @@ def convert_to_pdf(self, file_path: str, original_filename: Optional[str] = None
|
||||
Args:
|
||||
file_path: Path to the file to convert
|
||||
original_filename: Optional original filename (if different from path basename)
|
||||
owner_id: Optional user identifier forwarded to process_document for multi-user mode.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Starting PDF conversion: {file_path}")
|
||||
@@ -332,9 +335,9 @@ def convert_to_pdf(self, file_path: str, original_filename: Optional[str] = None
|
||||
# Change extension to .pdf for the original filename
|
||||
original_base = os.path.splitext(original_filename)[0]
|
||||
pdf_original_filename = f"{original_base}.pdf"
|
||||
process_document.delay(converted_file_path, original_filename=pdf_original_filename)
|
||||
process_document.delay(converted_file_path, original_filename=pdf_original_filename, owner_id=owner_id)
|
||||
else:
|
||||
process_document.delay(converted_file_path)
|
||||
process_document.delay(converted_file_path, owner_id=owner_id)
|
||||
|
||||
return converted_file_path
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,453 @@
|
||||
"""PDF/A archival conversion task.
|
||||
|
||||
Converts PDF files to PDF/A format using ocrmypdf (which relies on Ghostscript
|
||||
internally). Two variants are produced when enabled:
|
||||
|
||||
1. **Original PDF/A** – an archival copy of the ingested file, providing a
|
||||
time-stamped record of the document as it was upon ingestion.
|
||||
2. **Processed PDF/A** – an archival copy of the processed file with embedded
|
||||
metadata.
|
||||
|
||||
Both are saved under ``workdir/pdfa/`` and referenced in the database via
|
||||
``FileRecord.original_pdfa_path`` and ``FileRecord.processed_pdfa_path``.
|
||||
|
||||
When ``PDFA_TIMESTAMP_ENABLED`` is True, each PDF/A file also gets an RFC 3161
|
||||
timestamp response (``.tsr``) from a configurable Timestamp Authority (default:
|
||||
FreeTSA). This provides cryptographic proof of the file's existence at a given
|
||||
point in time.
|
||||
|
||||
.. note::
|
||||
|
||||
PDF/A conversion may alter font rendering (especially OCR text overlays
|
||||
produced by Microsoft Azure Document Intelligence). This is expected –
|
||||
the PDF/A copies are parallel archival variants, not replacements.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
import requests as http_requests
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils import get_unique_filepath_with_counter, log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Subdirectory structure under workdir for PDF/A copies
|
||||
PDFA_ORIGINAL_SUBDIR = os.path.join("pdfa", "original")
|
||||
PDFA_PROCESSED_SUBDIR = os.path.join("pdfa", "processed")
|
||||
|
||||
|
||||
def _convert_pdf_to_pdfa(input_path: str, output_path: str, pdfa_format: str = "2") -> bool:
|
||||
"""Convert a PDF file to PDF/A using ocrmypdf.
|
||||
|
||||
Uses ``ocrmypdf --skip-text --output-type pdfa-N`` so that existing text
|
||||
layers are preserved (not re-OCR'd) while the output is converted to
|
||||
PDF/A via Ghostscript.
|
||||
|
||||
Args:
|
||||
input_path: Absolute path to the source PDF file.
|
||||
output_path: Absolute path for the PDF/A output file.
|
||||
pdfa_format: PDF/A variant ('1', '2', or '3'). Defaults to '2' for PDF/A-2b.
|
||||
|
||||
Returns:
|
||||
True if conversion succeeded, False otherwise.
|
||||
"""
|
||||
# Validate format to prevent argument injection via output-type
|
||||
if pdfa_format not in ("1", "2", "3"):
|
||||
logger.error(f"[convert_to_pdfa] Invalid pdfa_format: {pdfa_format}")
|
||||
return False
|
||||
|
||||
ocrmypdf_bin = shutil.which("ocrmypdf")
|
||||
if not ocrmypdf_bin:
|
||||
logger.error("[convert_to_pdfa] ocrmypdf binary not found on PATH")
|
||||
return False
|
||||
|
||||
output_type = f"pdfa-{pdfa_format}"
|
||||
|
||||
cmd = [
|
||||
ocrmypdf_bin,
|
||||
"--skip-text",
|
||||
"--output-type",
|
||||
output_type,
|
||||
"--quiet",
|
||||
"--invalidate-digital-signatures",
|
||||
input_path,
|
||||
output_path,
|
||||
]
|
||||
|
||||
logger.info(f"[convert_to_pdfa] Running: {' '.join(cmd)}")
|
||||
|
||||
try:
|
||||
proc = subprocess.run(cmd, capture_output=True, text=True, timeout=600, check=False) # noqa: S603
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.warning("[convert_to_pdfa] ocrmypdf timed out after 600s")
|
||||
return False
|
||||
|
||||
if proc.returncode != 0:
|
||||
stderr_snippet = proc.stderr.strip()[:500] if proc.stderr else ""
|
||||
logger.warning(f"[convert_to_pdfa] ocrmypdf exited with code {proc.returncode}: {stderr_snippet}")
|
||||
return False
|
||||
|
||||
logger.info(f"[convert_to_pdfa] PDF/A file written to {output_path}")
|
||||
return True
|
||||
|
||||
|
||||
def _timestamp_file(file_path: str, tsa_url: str) -> str | None:
|
||||
"""Create an RFC 3161 timestamp for a file using a Timestamp Authority.
|
||||
|
||||
Uses ``openssl ts`` to create a timestamp request (TSQ) from the file's
|
||||
SHA-256 hash, submits it to the TSA via HTTP POST, and saves the timestamp
|
||||
response (TSR) alongside the file.
|
||||
|
||||
Args:
|
||||
file_path: Absolute path to the file to timestamp.
|
||||
tsa_url: URL of the RFC 3161 Timestamp Authority.
|
||||
|
||||
Returns:
|
||||
Path to the ``.tsr`` file if successful, None otherwise.
|
||||
"""
|
||||
openssl_bin = shutil.which("openssl")
|
||||
if not openssl_bin:
|
||||
logger.error("[timestamp] openssl binary not found on PATH")
|
||||
return None
|
||||
|
||||
tsr_path = file_path + ".tsr"
|
||||
tsq_path = file_path + ".tsq"
|
||||
|
||||
try:
|
||||
# Step 1: Create timestamp request
|
||||
cmd = [openssl_bin, "ts", "-query", "-data", file_path, "-sha256", "-no_nonce", "-out", tsq_path]
|
||||
proc = subprocess.run(cmd, capture_output=True, text=True, timeout=30, check=False) # noqa: S603
|
||||
if proc.returncode != 0:
|
||||
logger.warning(f"[timestamp] openssl ts -query failed: {proc.stderr.strip()[:200]}")
|
||||
return None
|
||||
|
||||
# Step 2: Submit TSQ to the Timestamp Authority
|
||||
with open(tsq_path, "rb") as f:
|
||||
tsq_data = f.read()
|
||||
|
||||
response = http_requests.post(
|
||||
tsa_url,
|
||||
data=tsq_data,
|
||||
headers={"Content-Type": "application/timestamp-query"},
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.warning(f"[timestamp] TSA returned HTTP {response.status_code} from {tsa_url}")
|
||||
return None
|
||||
|
||||
# Step 3: Save the timestamp response
|
||||
with open(tsr_path, "wb") as f:
|
||||
f.write(response.content)
|
||||
|
||||
logger.info(f"[timestamp] RFC 3161 timestamp saved to {tsr_path}")
|
||||
return tsr_path
|
||||
|
||||
except http_requests.RequestException as e:
|
||||
logger.warning(f"[timestamp] Failed to contact TSA at {tsa_url}: {e}")
|
||||
return None
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.warning("[timestamp] openssl ts timed out")
|
||||
return None
|
||||
finally:
|
||||
# Always clean up the TSQ file
|
||||
if os.path.exists(tsq_path):
|
||||
os.remove(tsq_path)
|
||||
|
||||
|
||||
def _compute_pdfa_folder_overrides() -> dict[str, str]:
|
||||
"""Compute per-provider folder overrides for PDF/A uploads.
|
||||
|
||||
Appends ``settings.pdfa_upload_folder`` to each provider's configured
|
||||
folder. For Google Drive (which uses folder IDs), uses the dedicated
|
||||
``google_drive_pdfa_folder_id`` setting.
|
||||
|
||||
Returns:
|
||||
Dictionary mapping provider names to folder override strings.
|
||||
"""
|
||||
subfolder = settings.pdfa_upload_folder
|
||||
overrides: dict[str, str] = {}
|
||||
|
||||
if not subfolder:
|
||||
return overrides
|
||||
|
||||
# Path-based providers: append subfolder
|
||||
for provider, folder_attr in [
|
||||
("dropbox", "dropbox_folder"),
|
||||
("nextcloud", "nextcloud_folder"),
|
||||
("webdav", "webdav_folder"),
|
||||
("ftp", "ftp_folder"),
|
||||
("sftp", "sftp_folder"),
|
||||
("onedrive", "onedrive_folder_path"),
|
||||
]:
|
||||
base = getattr(settings, folder_attr, "") or ""
|
||||
overrides[provider] = f"{base.rstrip('/')}/{subfolder}" if base else subfolder
|
||||
|
||||
# S3: append subfolder to prefix (trailing slash is required by S3 convention
|
||||
# where "folder" paths are key prefixes, unlike path-based providers above)
|
||||
s3_prefix = getattr(settings, "s3_folder_prefix", "") or ""
|
||||
overrides["s3"] = f"{s3_prefix.rstrip('/')}/{subfolder}/"
|
||||
|
||||
# Google Drive: use dedicated folder ID or fall back to default
|
||||
gdrive_pdfa_id = settings.google_drive_pdfa_folder_id
|
||||
if gdrive_pdfa_id:
|
||||
overrides["google_drive"] = gdrive_pdfa_id
|
||||
|
||||
return overrides
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def convert_to_pdfa(self, file_id: int) -> dict:
|
||||
"""Generate PDF/A archival copies for a processed document.
|
||||
|
||||
Creates PDF/A variants of both the original ingested file and the
|
||||
processed file (with embedded metadata). Files are saved under
|
||||
``workdir/pdfa/original/`` and ``workdir/pdfa/processed/`` respectively.
|
||||
|
||||
When timestamping is enabled, each PDF/A file also gets an RFC 3161
|
||||
``.tsr`` timestamp from the configured TSA.
|
||||
|
||||
Upload of each variant to storage providers is controlled independently
|
||||
by ``pdfa_upload_original`` and ``pdfa_upload_processed``.
|
||||
|
||||
Args:
|
||||
file_id: ID of the FileRecord to create PDF/A copies for.
|
||||
|
||||
Returns:
|
||||
Dictionary with status and file paths.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Starting PDF/A conversion for file_id={file_id}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_to_pdfa",
|
||||
"in_progress",
|
||||
"Starting PDF/A archival conversion",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# Fetch file record
|
||||
with SessionLocal() as db:
|
||||
file_record = db.query(FileRecord).filter_by(id=file_id).first()
|
||||
if not file_record:
|
||||
logger.error(f"[{task_id}] FileRecord {file_id} not found")
|
||||
log_task_progress(task_id, "convert_to_pdfa", "failure", "File record not found", file_id=file_id)
|
||||
return {"error": "File record not found", "file_id": file_id}
|
||||
|
||||
original_path = file_record.original_file_path
|
||||
processed_path = file_record.processed_file_path
|
||||
|
||||
pdfa_format = settings.pdfa_format
|
||||
timestamp_enabled = settings.pdfa_timestamp_enabled
|
||||
timestamp_url = settings.pdfa_timestamp_url
|
||||
results = {}
|
||||
|
||||
# --- Convert original file to PDF/A ---
|
||||
if original_path and os.path.exists(original_path):
|
||||
original_pdfa_dir = os.path.join(settings.workdir, PDFA_ORIGINAL_SUBDIR)
|
||||
os.makedirs(original_pdfa_dir, exist_ok=True)
|
||||
|
||||
base_name = os.path.splitext(os.path.basename(original_path))[0]
|
||||
original_pdfa_path = get_unique_filepath_with_counter(original_pdfa_dir, base_name, ".pdf")
|
||||
|
||||
logger.info(f"[{task_id}] Converting original to PDF/A: {original_path} -> {original_pdfa_path}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_original_to_pdfa",
|
||||
"in_progress",
|
||||
f"Converting original to PDF/A: {os.path.basename(original_path)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
success = _convert_pdf_to_pdfa(original_path, original_pdfa_path, pdfa_format)
|
||||
if success:
|
||||
results["original_pdfa_path"] = original_pdfa_path
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_original_to_pdfa",
|
||||
"success",
|
||||
f"Original PDF/A saved: {os.path.basename(original_pdfa_path)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
# Timestamp the original PDF/A
|
||||
if timestamp_enabled:
|
||||
tsr = _timestamp_file(original_pdfa_path, timestamp_url)
|
||||
if tsr:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"timestamp_original_pdfa",
|
||||
"success",
|
||||
f"Timestamped: {os.path.basename(tsr)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
else:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"timestamp_original_pdfa",
|
||||
"failure",
|
||||
"Failed to timestamp original PDF/A",
|
||||
file_id=file_id,
|
||||
)
|
||||
else:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_original_to_pdfa",
|
||||
"failure",
|
||||
"Failed to convert original to PDF/A",
|
||||
file_id=file_id,
|
||||
)
|
||||
else:
|
||||
logger.warning(f"[{task_id}] Original file not found, skipping original PDF/A conversion")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_original_to_pdfa",
|
||||
"skipped",
|
||||
"Original file not available",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# --- Convert processed file to PDF/A ---
|
||||
if processed_path and os.path.exists(processed_path):
|
||||
processed_pdfa_dir = os.path.join(settings.workdir, PDFA_PROCESSED_SUBDIR)
|
||||
os.makedirs(processed_pdfa_dir, exist_ok=True)
|
||||
|
||||
base_name = os.path.splitext(os.path.basename(processed_path))[0]
|
||||
processed_pdfa_path = get_unique_filepath_with_counter(processed_pdfa_dir, f"{base_name}-PDFA", ".pdf")
|
||||
|
||||
logger.info(f"[{task_id}] Converting processed to PDF/A: {processed_path} -> {processed_pdfa_path}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_processed_to_pdfa",
|
||||
"in_progress",
|
||||
f"Converting processed to PDF/A: {os.path.basename(processed_path)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
success = _convert_pdf_to_pdfa(processed_path, processed_pdfa_path, pdfa_format)
|
||||
if success:
|
||||
results["processed_pdfa_path"] = processed_pdfa_path
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_processed_to_pdfa",
|
||||
"success",
|
||||
f"Processed PDF/A saved: {os.path.basename(processed_pdfa_path)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
# Timestamp the processed PDF/A
|
||||
if timestamp_enabled:
|
||||
tsr = _timestamp_file(processed_pdfa_path, timestamp_url)
|
||||
if tsr:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"timestamp_processed_pdfa",
|
||||
"success",
|
||||
f"Timestamped: {os.path.basename(tsr)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
else:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"timestamp_processed_pdfa",
|
||||
"failure",
|
||||
"Failed to timestamp processed PDF/A",
|
||||
file_id=file_id,
|
||||
)
|
||||
else:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_processed_to_pdfa",
|
||||
"failure",
|
||||
"Failed to convert processed to PDF/A",
|
||||
file_id=file_id,
|
||||
)
|
||||
else:
|
||||
logger.warning(f"[{task_id}] Processed file not found, skipping processed PDF/A conversion")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"convert_processed_to_pdfa",
|
||||
"skipped",
|
||||
"Processed file not available",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# --- Update database with PDF/A paths ---
|
||||
with SessionLocal() as db:
|
||||
file_record = db.query(FileRecord).filter_by(id=file_id).first()
|
||||
if file_record:
|
||||
if "original_pdfa_path" in results:
|
||||
file_record.original_pdfa_path = results["original_pdfa_path"]
|
||||
if "processed_pdfa_path" in results:
|
||||
file_record.processed_pdfa_path = results["processed_pdfa_path"]
|
||||
db.commit()
|
||||
logger.info(f"[{task_id}] Updated database with PDF/A paths")
|
||||
|
||||
# --- Upload PDF/A variants to storage providers ---
|
||||
folder_overrides = _compute_pdfa_folder_overrides()
|
||||
|
||||
if settings.pdfa_upload_original and "original_pdfa_path" in results:
|
||||
from app.tasks.send_to_all import send_to_all_destinations
|
||||
|
||||
logger.info(f"[{task_id}] Uploading original PDF/A to storage providers")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_original_pdfa",
|
||||
"in_progress",
|
||||
"Uploading original PDF/A to storage providers",
|
||||
file_id=file_id,
|
||||
)
|
||||
send_to_all_destinations.delay(
|
||||
results["original_pdfa_path"],
|
||||
True,
|
||||
file_id,
|
||||
folder_overrides=folder_overrides if folder_overrides else None,
|
||||
)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_original_pdfa",
|
||||
"success",
|
||||
"Original PDF/A queued for upload",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
if settings.pdfa_upload_processed and "processed_pdfa_path" in results:
|
||||
from app.tasks.send_to_all import send_to_all_destinations
|
||||
|
||||
logger.info(f"[{task_id}] Uploading processed PDF/A to storage providers")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_processed_pdfa",
|
||||
"in_progress",
|
||||
"Uploading processed PDF/A to storage providers",
|
||||
file_id=file_id,
|
||||
)
|
||||
send_to_all_destinations.delay(
|
||||
results["processed_pdfa_path"],
|
||||
True,
|
||||
file_id,
|
||||
folder_overrides=folder_overrides if folder_overrides else None,
|
||||
)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_processed_pdfa",
|
||||
"success",
|
||||
"Processed PDF/A queued for upload",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# --- Final status ---
|
||||
has_any = bool(results)
|
||||
status = "success" if has_any else "failure"
|
||||
message = (
|
||||
f"PDF/A conversion complete ({len(results)} variant(s) created)" if has_any else "No PDF/A variants created"
|
||||
)
|
||||
log_task_progress(task_id, "convert_to_pdfa", status, message, file_id=file_id)
|
||||
|
||||
return {"status": status, "file_id": file_id, **results}
|
||||
@@ -1,191 +1,192 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
|
||||
# Import the shared Celery instance
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.tasks.embed_metadata_into_pdf import embed_metadata_into_pdf
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.ai_provider import get_ai_provider
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def extract_json_from_text(text):
|
||||
"""
|
||||
Try to extract a JSON object from the text.
|
||||
- First, check for a JSON block inside triple backticks.
|
||||
- If not found, try to extract text from the first '{' to the last '}'.
|
||||
"""
|
||||
pattern = r"```(?:json)?\s*(\{.*?\})\s*```"
|
||||
match = re.search(pattern, text, re.DOTALL)
|
||||
if match:
|
||||
return match.group(1)
|
||||
else:
|
||||
start = text.find("{")
|
||||
end = text.rfind("}")
|
||||
if start != -1 and end != -1 and end > start:
|
||||
return text[start : end + 1]
|
||||
return None
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def extract_metadata_with_gpt(self, filename: str, cleaned_text: str, file_id: int = None):
|
||||
"""
|
||||
Uses OpenAI to classify document metadata.
|
||||
|
||||
Args:
|
||||
filename: Can be either a basename (e.g., "file.pdf") or a full path (e.g., "/workdir/processed/file.pdf")
|
||||
cleaned_text: The extracted text from the document
|
||||
file_id: Optional file ID for tracking
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Starting metadata extraction for: {filename}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"extract_metadata_with_gpt",
|
||||
"in_progress",
|
||||
f"Extracting metadata for {os.path.basename(filename)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# Get file_id from database if not provided
|
||||
if file_id is None:
|
||||
tmp_dir = os.path.join(settings.workdir, "tmp")
|
||||
# Handle both basename and full path
|
||||
if os.path.isabs(filename):
|
||||
file_path = filename
|
||||
else:
|
||||
file_path = os.path.join(tmp_dir, filename)
|
||||
if os.path.exists(file_path):
|
||||
with SessionLocal() as db:
|
||||
file_record = db.query(FileRecord).filter_by(local_filename=file_path).first()
|
||||
if file_record:
|
||||
file_id = file_record.id
|
||||
|
||||
prompt = (
|
||||
"You are a specialized document analyzer trained to extract structured metadata from documents.\n"
|
||||
"Your task is to analyze the given text and return a well-structured JSON object.\n\n"
|
||||
"Extract and return the following fields:\n"
|
||||
"1. **filename**: Machine-readable filename "
|
||||
"(YYYY-MM-DD_DescriptiveTitle, use only letters, numbers, periods, and underscores).\n"
|
||||
'2. **empfaenger**: The recipient, or "Unknown" if not found.\n'
|
||||
'3. **absender**: The sender, or "Unknown" if not found.\n'
|
||||
"4. **correspondent**: The entity or company that issued the document "
|
||||
'(shortest possible name, e.g., "Amazon" instead of "Amazon EU SARL, German branch").\n'
|
||||
"5. **kommunikationsart**: One of [Behoerdlicher_Brief, Rechnung, Kontoauszug, Vertrag, "
|
||||
"Quittung, Privater_Brief, Einladung, Gewerbliche_Korrespondenz, Newsletter, Werbung, Sonstiges].\n"
|
||||
"6. **kommunikationskategorie**: One of [Amtliche_Postbehoerdliche_Dokumente, "
|
||||
"Finanz_und_Vertragsdokumente, Geschaeftliche_Kommunikation, "
|
||||
"Private_Korrespondenz, Sonstige_Informationen].\n"
|
||||
"7. **document_type**: Precise classification (e.g., Invoice, Contract, Information, Unknown).\n"
|
||||
"8. **tags**: A list of up to 4 relevant thematic keywords.\n"
|
||||
'9. **language**: Detected document language (ISO 639-1 code, e.g., "de" or "en").\n'
|
||||
"10. **title**: A human-readable title summarizing the document content.\n"
|
||||
"11. **confidence_score**: A numeric value (0-100) indicating the confidence level "
|
||||
"of the extracted metadata.\n"
|
||||
"12. **reference_number**: Extracted invoice/order/reference number if available.\n"
|
||||
"13. **monetary_amounts**: A list of key monetary values detected in the document.\n\n"
|
||||
"### Important Rules:\n"
|
||||
"- **OCR Correction**: Assume the text has been corrected for OCR errors.\n"
|
||||
"- **Tagging**: Max 4 tags, avoiding generic or overly specific terms.\n"
|
||||
"- **Title**: Concise, no addresses, and contains key identifying features.\n"
|
||||
"- **Date Selection**: Use the most relevant date if multiple are found.\n"
|
||||
"- **Output Language**: Maintain the document's original language.\n\n"
|
||||
f"Extracted text:\n{cleaned_text}\n\n"
|
||||
"Return only valid JSON with no additional commentary.\n"
|
||||
)
|
||||
|
||||
try:
|
||||
logger.info(f"[{task_id}] Sending classification request for {filename}...")
|
||||
log_task_progress(task_id, "call_ai_provider", "in_progress", "Calling AI provider API", file_id=file_id)
|
||||
provider = get_ai_provider()
|
||||
model = settings.ai_model or settings.openai_model
|
||||
content = provider.chat_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": "You are an intelligent document classifier."},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
model=model,
|
||||
temperature=0,
|
||||
)
|
||||
|
||||
logger.info(f"[{task_id}] Raw classification response for {filename}: {content[:200]}...")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"call_ai_provider",
|
||||
"success",
|
||||
"Received AI provider response",
|
||||
file_id=file_id,
|
||||
detail=f"Raw classification response:\n{content}",
|
||||
)
|
||||
|
||||
json_text = extract_json_from_text(content)
|
||||
if not json_text:
|
||||
logger.error(f"[{task_id}] Could not find valid JSON in GPT response for {filename}.")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"extract_metadata_with_gpt",
|
||||
"failure",
|
||||
"Invalid JSON in response",
|
||||
file_id=file_id,
|
||||
detail=f"Could not parse valid JSON from GPT response.\nRaw response:\n{content}",
|
||||
)
|
||||
return {}
|
||||
|
||||
metadata = json.loads(json_text)
|
||||
|
||||
# SECURITY: Validate filename format from GPT to prevent path traversal
|
||||
# The prompt requests filenames with only letters, numbers, periods, and underscores
|
||||
# Enforce this constraint to prevent malicious filenames
|
||||
suggested_filename = metadata.get("filename", "")
|
||||
if suggested_filename:
|
||||
# Check if filename contains only safe characters AND explicitly check for ".."
|
||||
# Defense in depth: While the regex [\w\-\. ]+ already excludes / and \,
|
||||
# we explicitly reject ".." to guard against:
|
||||
# 1. Potential locale-specific \w behavior
|
||||
# 2. Files literally named ".." which are valid but problematic
|
||||
# 3. Future code changes that might relax the regex
|
||||
if not re.match(r"^[\w\-\. ]+$", suggested_filename) or ".." in suggested_filename:
|
||||
logger.warning(f"[{task_id}] Invalid filename format from GPT: '{suggested_filename}', using fallback")
|
||||
# Reset to empty to trigger fallback to original filename
|
||||
metadata["filename"] = ""
|
||||
|
||||
logger.info(f"[{task_id}] Extracted metadata: {metadata}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"parse_metadata",
|
||||
"success",
|
||||
f"Parsed metadata: {list(metadata.keys())}",
|
||||
file_id=file_id,
|
||||
detail=f"Extracted metadata:\n{json.dumps(metadata, ensure_ascii=False, indent=2)}",
|
||||
)
|
||||
|
||||
# Trigger the next step: embedding metadata into the PDF
|
||||
# Pass the filename (can be basename or full path) so embed_metadata_into_pdf can find the file on disk
|
||||
logger.info(f"[{task_id}] Queueing metadata embedding task")
|
||||
log_task_progress(
|
||||
task_id, "extract_metadata_with_gpt", "success", "Metadata extracted, queuing embed task", file_id=file_id
|
||||
)
|
||||
embed_metadata_into_pdf.delay(filename, cleaned_text, metadata, file_id)
|
||||
|
||||
return {"s3_file": os.path.basename(filename), "metadata": metadata}
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"[{task_id}] AI provider classification failed for {filename}: {e}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"extract_metadata_with_gpt",
|
||||
"failure",
|
||||
f"Exception: {str(e)}",
|
||||
file_id=file_id,
|
||||
detail=f"AI provider classification failed for {filename}.\nException: {str(e)}",
|
||||
)
|
||||
return {}
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
|
||||
# Import the shared Celery instance
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.tasks.embed_metadata_into_pdf import embed_metadata_into_pdf
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.ai_provider import get_ai_provider
|
||||
from app.utils.filename_utils import VALID_FILENAME_RE
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def extract_json_from_text(text):
|
||||
"""
|
||||
Try to extract a JSON object from the text.
|
||||
- First, check for a JSON block inside triple backticks.
|
||||
- If not found, try to extract text from the first '{' to the last '}'.
|
||||
"""
|
||||
pattern = r"```(?:json)?\s*(\{.*?\})\s*```"
|
||||
match = re.search(pattern, text, re.DOTALL)
|
||||
if match:
|
||||
return match.group(1)
|
||||
else:
|
||||
start = text.find("{")
|
||||
end = text.rfind("}")
|
||||
if start != -1 and end != -1 and end > start:
|
||||
return text[start : end + 1]
|
||||
return None
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def extract_metadata_with_gpt(self, filename: str, cleaned_text: str, file_id: int = None):
|
||||
"""
|
||||
Uses OpenAI to classify document metadata.
|
||||
|
||||
Args:
|
||||
filename: Can be either a basename (e.g., "file.pdf") or a full path (e.g., "/workdir/processed/file.pdf")
|
||||
cleaned_text: The extracted text from the document
|
||||
file_id: Optional file ID for tracking
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Starting metadata extraction for: {filename}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"extract_metadata_with_gpt",
|
||||
"in_progress",
|
||||
f"Extracting metadata for {os.path.basename(filename)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# Get file_id from database if not provided
|
||||
if file_id is None:
|
||||
tmp_dir = os.path.join(settings.workdir, "tmp")
|
||||
# Handle both basename and full path
|
||||
if os.path.isabs(filename):
|
||||
file_path = filename
|
||||
else:
|
||||
file_path = os.path.join(tmp_dir, filename)
|
||||
if os.path.exists(file_path):
|
||||
with SessionLocal() as db:
|
||||
file_record = db.query(FileRecord).filter_by(local_filename=file_path).first()
|
||||
if file_record:
|
||||
file_id = file_record.id
|
||||
|
||||
prompt = (
|
||||
"You are a specialized document analyzer trained to extract structured metadata from documents.\n"
|
||||
"Your task is to analyze the given text and return a well-structured JSON object.\n\n"
|
||||
"Extract and return the following fields:\n"
|
||||
"1. **filename**: Machine-readable filename "
|
||||
"(YYYY-MM-DD_DescriptiveTitle, use only letters, numbers, spaces, dashes, periods, and underscores).\n"
|
||||
'2. **empfaenger**: The recipient, or "Unknown" if not found.\n'
|
||||
'3. **absender**: The sender, or "Unknown" if not found.\n'
|
||||
"4. **correspondent**: The entity or company that issued the document "
|
||||
'(shortest possible name, e.g., "Amazon" instead of "Amazon EU SARL, German branch").\n'
|
||||
"5. **kommunikationsart**: One of [Behoerdlicher_Brief, Rechnung, Kontoauszug, Vertrag, "
|
||||
"Quittung, Privater_Brief, Einladung, Gewerbliche_Korrespondenz, Newsletter, Werbung, Sonstiges].\n"
|
||||
"6. **kommunikationskategorie**: One of [Amtliche_Postbehoerdliche_Dokumente, "
|
||||
"Finanz_und_Vertragsdokumente, Geschaeftliche_Kommunikation, "
|
||||
"Private_Korrespondenz, Sonstige_Informationen].\n"
|
||||
"7. **document_type**: Precise classification (e.g., Invoice, Contract, Information, Unknown).\n"
|
||||
"8. **tags**: A list of up to 4 relevant thematic keywords.\n"
|
||||
'9. **language**: Detected document language (ISO 639-1 code, e.g., "de" or "en").\n'
|
||||
"10. **title**: A human-readable title summarizing the document content.\n"
|
||||
"11. **confidence_score**: A numeric value (0-100) indicating the confidence level "
|
||||
"of the extracted metadata.\n"
|
||||
"12. **reference_number**: Extracted invoice/order/reference number if available.\n"
|
||||
"13. **monetary_amounts**: A list of key monetary values detected in the document.\n\n"
|
||||
"### Important Rules:\n"
|
||||
"- **OCR Correction**: Assume the text has been corrected for OCR errors.\n"
|
||||
"- **Tagging**: Max 4 tags, avoiding generic or overly specific terms.\n"
|
||||
"- **Title**: Concise, no addresses, and contains key identifying features.\n"
|
||||
"- **Date Selection**: Use the most relevant date if multiple are found.\n"
|
||||
"- **Output Language**: Maintain the document's original language.\n\n"
|
||||
f"Extracted text:\n{cleaned_text}\n\n"
|
||||
"Return only valid JSON with no additional commentary.\n"
|
||||
)
|
||||
|
||||
try:
|
||||
logger.info(f"[{task_id}] Sending classification request for {filename}...")
|
||||
log_task_progress(task_id, "call_ai_provider", "in_progress", "Calling AI provider API", file_id=file_id)
|
||||
provider = get_ai_provider()
|
||||
model = settings.ai_model or settings.openai_model
|
||||
content = provider.chat_completion(
|
||||
messages=[
|
||||
{"role": "system", "content": "You are an intelligent document classifier."},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
model=model,
|
||||
temperature=0,
|
||||
)
|
||||
|
||||
logger.info(f"[{task_id}] Raw classification response for {filename}: {content[:200]}...")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"call_ai_provider",
|
||||
"success",
|
||||
"Received AI provider response",
|
||||
file_id=file_id,
|
||||
detail=f"Raw classification response:\n{content}",
|
||||
)
|
||||
|
||||
json_text = extract_json_from_text(content)
|
||||
if not json_text:
|
||||
logger.error(f"[{task_id}] Could not find valid JSON in GPT response for {filename}.")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"extract_metadata_with_gpt",
|
||||
"failure",
|
||||
"Invalid JSON in response",
|
||||
file_id=file_id,
|
||||
detail=f"Could not parse valid JSON from GPT response.\nRaw response:\n{content}",
|
||||
)
|
||||
return {}
|
||||
|
||||
metadata = json.loads(json_text)
|
||||
|
||||
# SECURITY: Validate filename format from GPT to prevent path traversal
|
||||
# The prompt requests filenames with only letters, numbers, periods, and underscores
|
||||
# Enforce this constraint to prevent malicious filenames
|
||||
suggested_filename = metadata.get("filename", "")
|
||||
if suggested_filename:
|
||||
# Check if filename contains only safe characters AND explicitly check for ".."
|
||||
# Defense in depth: While the regex VALID_FILENAME_PATTERN already excludes / and \,
|
||||
# we explicitly reject ".." to guard against:
|
||||
# 1. Potential locale-specific \w behavior
|
||||
# 2. Files literally named ".." which are valid but problematic
|
||||
# 3. Future code changes that might relax the regex
|
||||
if not VALID_FILENAME_RE.match(suggested_filename) or ".." in suggested_filename:
|
||||
logger.warning(f"[{task_id}] Invalid filename format from GPT: '{suggested_filename}', using fallback")
|
||||
# Reset to empty to trigger fallback to original filename
|
||||
metadata["filename"] = ""
|
||||
|
||||
logger.info(f"[{task_id}] Extracted metadata: {metadata}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"parse_metadata",
|
||||
"success",
|
||||
f"Parsed metadata: {list(metadata.keys())}",
|
||||
file_id=file_id,
|
||||
detail=f"Extracted metadata:\n{json.dumps(metadata, ensure_ascii=False, indent=2)}",
|
||||
)
|
||||
|
||||
# Trigger the next step: embedding metadata into the PDF
|
||||
# Pass the filename (can be basename or full path) so embed_metadata_into_pdf can find the file on disk
|
||||
logger.info(f"[{task_id}] Queueing metadata embedding task")
|
||||
log_task_progress(
|
||||
task_id, "extract_metadata_with_gpt", "success", "Metadata extracted, queuing embed task", file_id=file_id
|
||||
)
|
||||
embed_metadata_into_pdf.delay(filename, cleaned_text, metadata, file_id)
|
||||
|
||||
return {"s3_file": os.path.basename(filename), "metadata": metadata}
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"[{task_id}] AI provider classification failed for {filename}: {e}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"extract_metadata_with_gpt",
|
||||
"failure",
|
||||
f"Exception: {str(e)}",
|
||||
file_id=file_id,
|
||||
detail=f"AI provider classification failed for {filename}.\nException: {str(e)}",
|
||||
)
|
||||
return {}
|
||||
|
||||
@@ -10,8 +10,13 @@ from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
|
||||
# Import the aggregator task and validator
|
||||
from app.tasks.send_to_all import get_configured_services_from_validator, send_to_all_destinations
|
||||
# Import the aggregator tasks and validator
|
||||
from app.tasks.send_to_all import (
|
||||
get_configured_services_from_validator,
|
||||
get_user_destination_count,
|
||||
send_to_all_destinations,
|
||||
send_to_user_destinations,
|
||||
)
|
||||
|
||||
# Import database and logging utils from main
|
||||
from app.utils import log_task_progress
|
||||
@@ -26,13 +31,21 @@ logger = logging.getLogger(__name__)
|
||||
def finalize_document_storage(self, original_file: str, processed_file: str, metadata: dict, file_id: int = None):
|
||||
"""
|
||||
Final storage step after embedding metadata.
|
||||
We will now call 'send_to_all_destinations' to push the final PDF to Dropbox/Nextcloud/Paperless.
|
||||
After uploading, send a notification about the processed file.
|
||||
Routes the processed document to the appropriate destination(s):
|
||||
|
||||
1. If the document has an identified owner and that owner has active
|
||||
DESTINATION UserIntegrations, the file is uploaded to each of those
|
||||
integrations (user-specific routing).
|
||||
2. Otherwise the file is forwarded to the globally-configured destinations
|
||||
via :func:`send_to_all_destinations` (system-wide fallback).
|
||||
|
||||
After queuing uploads, optional PDF/A archival conversion and embedding
|
||||
computation are triggered, and a completion notification is sent.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Finalizing document storage for {processed_file}")
|
||||
|
||||
# 1. Update Database Status (From Main)
|
||||
# 1. Update Database Status
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"finalize_document_storage",
|
||||
@@ -41,7 +54,8 @@ def finalize_document_storage(self, original_file: str, processed_file: str, met
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# Get file_id from database if not provided (fallback logic from Main)
|
||||
# 2. Resolve file_id and owner_id from the database
|
||||
owner_id = None
|
||||
if file_id is None:
|
||||
with SessionLocal() as db:
|
||||
# Only as a last resort, try to find by exact match on local_filename
|
||||
@@ -49,38 +63,73 @@ def finalize_document_storage(self, original_file: str, processed_file: str, met
|
||||
file_record = db.query(FileRecord).filter(FileRecord.local_filename == tmp_path).first()
|
||||
if file_record:
|
||||
file_id = file_record.id
|
||||
owner_id = file_record.owner_id
|
||||
else:
|
||||
with SessionLocal() as db:
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if file_record:
|
||||
owner_id = file_record.owner_id
|
||||
|
||||
# 2. Determine Configured Destinations (From Copilot)
|
||||
# This is needed for the notification message later
|
||||
# 3. Determine configured destinations for notification
|
||||
configured_destinations = []
|
||||
try:
|
||||
configured_services = get_configured_services_from_validator()
|
||||
# Get list of service names that are configured
|
||||
for service_name, is_configured in configured_services.items():
|
||||
if is_configured:
|
||||
# Format service names for display
|
||||
display_name = service_name.replace("_", " ").title()
|
||||
configured_destinations.append(display_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"[WARNING] Could not determine configured destinations: {e}")
|
||||
configured_destinations = ["configured destinations"]
|
||||
|
||||
# 3. Queue Uploads (Merged)
|
||||
# Uses Main branch signature to ensure file_id is passed, but keeps logic structure
|
||||
logger.info(f"[{task_id}] Queueing uploads to all destinations")
|
||||
# 4. Queue Uploads — prefer user-specific destinations when available
|
||||
log_task_progress(
|
||||
task_id, "finalize_document_storage", "success", "Queuing uploads to destinations", file_id=file_id
|
||||
)
|
||||
|
||||
# Note: send_to_all_destinations is asynchronous and queues upload tasks
|
||||
# We pass 'True' (delete_after) and 'file_id' as per Main branch requirements
|
||||
send_to_all_destinations.delay(processed_file, True, file_id)
|
||||
user_dest_count = 0
|
||||
if owner_id:
|
||||
try:
|
||||
user_dest_count = get_user_destination_count(owner_id)
|
||||
except Exception as e:
|
||||
logger.warning("[%s] Could not query user destination count for owner=%s: %s", task_id, owner_id, e)
|
||||
|
||||
# 4. Send Notification (From Copilot)
|
||||
# Note: This notification is sent after processing is complete but while uploads
|
||||
# are being queued.
|
||||
if owner_id and user_dest_count > 0:
|
||||
# User has configured their own destinations → use those exclusively
|
||||
logger.info(
|
||||
"[%s] Routing to %d user-specific destination(s) for owner=%s",
|
||||
task_id,
|
||||
user_dest_count,
|
||||
owner_id,
|
||||
)
|
||||
send_to_user_destinations.delay(processed_file, owner_id, file_id)
|
||||
else:
|
||||
# No user-specific destinations → fall back to global configuration
|
||||
logger.info("[%s] No user-specific destinations found; using global destinations", task_id)
|
||||
send_to_all_destinations.delay(processed_file, True, file_id)
|
||||
|
||||
# 4a. Trigger PDF/A archival conversion if enabled
|
||||
if settings.enable_pdfa_conversion:
|
||||
try:
|
||||
from app.tasks.convert_to_pdfa import convert_to_pdfa
|
||||
|
||||
logger.info(f"[{task_id}] PDF/A conversion enabled, queueing archival conversion")
|
||||
convert_to_pdfa.delay(file_id)
|
||||
except Exception as e:
|
||||
logger.warning(f"[{task_id}] Could not queue PDF/A conversion: {e}")
|
||||
|
||||
# 4b. Queue embedding computation
|
||||
if file_id is not None:
|
||||
try:
|
||||
from app.tasks.compute_embedding import compute_document_embedding
|
||||
|
||||
compute_document_embedding.delay(file_id)
|
||||
logger.info(f"[{task_id}] Queued embedding computation for file {file_id}")
|
||||
except Exception as e:
|
||||
logger.warning(f"[{task_id}] Could not queue embedding task: {e}")
|
||||
|
||||
# 5. Send Notification
|
||||
try:
|
||||
# Get file information
|
||||
file_size = os.path.getsize(processed_file) if os.path.exists(processed_file) else 0
|
||||
filename = os.path.basename(processed_file)
|
||||
|
||||
|
||||
+670
-407
File diff suppressed because it is too large
Load Diff
@@ -1,10 +1,14 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
import shutil
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import pypdf # Upgraded from PyPDF2 to fix CVE-2023-36464
|
||||
from pypdf.errors import PdfReadError
|
||||
@@ -12,7 +16,7 @@ from pypdf.errors import PdfReadError
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.models import FileRecord, Pipeline, PipelineStep
|
||||
from app.tasks.extract_metadata_with_gpt import extract_metadata_with_gpt
|
||||
from app.tasks.process_with_ocr import process_with_ocr
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
@@ -20,12 +24,83 @@ from app.utils import get_unique_filepath_with_counter, hash_file, log_task_prog
|
||||
from app.utils.step_manager import initialize_file_steps
|
||||
from app.utils.text_quality import check_text_quality, detect_pdf_text_source
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_pipeline_ocr_language(db: "Session", file_record: FileRecord, owner_id: str | None) -> str | None:
|
||||
"""Look up the OCR language override from the file's pipeline OCR step config.
|
||||
|
||||
Resolution order:
|
||||
1. Explicit pipeline assigned to the file (``file_record.pipeline_id``).
|
||||
2. User's own default pipeline (``owner_id``, ``is_default=True``).
|
||||
3. System default pipeline (``owner_id=NULL``, ``is_default=True``).
|
||||
|
||||
Returns the ``ocr_language`` value from the pipeline's OCR step config, or
|
||||
``None`` when no override is configured.
|
||||
"""
|
||||
pipeline = None
|
||||
|
||||
if file_record.pipeline_id:
|
||||
pipeline = db.query(Pipeline).filter(Pipeline.id == file_record.pipeline_id).first()
|
||||
|
||||
if pipeline is None and owner_id:
|
||||
pipeline = (
|
||||
db.query(Pipeline)
|
||||
.filter(
|
||||
Pipeline.owner_id == owner_id,
|
||||
Pipeline.is_default.is_(True),
|
||||
Pipeline.is_active.is_(True),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if pipeline is None:
|
||||
pipeline = (
|
||||
db.query(Pipeline)
|
||||
.filter(
|
||||
Pipeline.owner_id.is_(None),
|
||||
Pipeline.is_default.is_(True),
|
||||
Pipeline.is_active.is_(True),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if pipeline is None:
|
||||
return None
|
||||
|
||||
ocr_step = (
|
||||
db.query(PipelineStep)
|
||||
.filter(
|
||||
PipelineStep.pipeline_id == pipeline.id,
|
||||
PipelineStep.step_type == "ocr",
|
||||
PipelineStep.enabled.is_(True),
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if ocr_step is None or not ocr_step.config:
|
||||
return None
|
||||
|
||||
try:
|
||||
step_config = json.loads(ocr_step.config)
|
||||
lang = step_config.get("ocr_language")
|
||||
# "auto" is treated as no override
|
||||
return lang if lang and lang != "auto" else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def process_document(
|
||||
self, original_local_file: str, original_filename: str = None, file_id: int = None, force_cloud_ocr: bool = False
|
||||
self,
|
||||
original_local_file: str,
|
||||
original_filename: str = None,
|
||||
file_id: int = None,
|
||||
force_cloud_ocr: bool = False,
|
||||
owner_id: str = None,
|
||||
):
|
||||
"""
|
||||
Process a document file and trigger appropriate text extraction.
|
||||
@@ -37,6 +112,8 @@ def process_document(
|
||||
detection and reuses the existing record (used for reprocessing).
|
||||
force_cloud_ocr: If True, forces Azure Document Intelligence OCR processing
|
||||
regardless of embedded text quality. Used for re-processing.
|
||||
owner_id: Optional user identifier for multi-user mode. When provided, the
|
||||
created FileRecord is associated with this user.
|
||||
|
||||
Steps:
|
||||
1. Check if we have a FileRecord entry (via SHA-256 hash). If found, skip re-processing.
|
||||
@@ -48,6 +125,11 @@ def process_document(
|
||||
- Otherwise, queue Azure Document Intelligence processing
|
||||
3. If force_cloud_ocr is True, skip local text extraction and use cloud OCR
|
||||
"""
|
||||
# Fall back to the configured default_owner_id when no explicit owner was provided
|
||||
default_owner_id = settings.default_owner_id
|
||||
if owner_id is None and isinstance(default_owner_id, str) and default_owner_id.strip():
|
||||
owner_id = default_owner_id
|
||||
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Starting document processing: {original_local_file}")
|
||||
log_task_progress(
|
||||
@@ -97,6 +179,7 @@ def process_document(
|
||||
)
|
||||
|
||||
# Acquire DB session in the task
|
||||
ocr_language: str | None = None # Pipeline OCR language override resolved inside DB session
|
||||
with SessionLocal() as db:
|
||||
# When file_id is provided, we are reprocessing an existing file.
|
||||
# Skip the duplicate check and reuse the existing record.
|
||||
@@ -141,6 +224,7 @@ def process_document(
|
||||
mime_type=mime_type,
|
||||
is_duplicate=True,
|
||||
duplicate_of_id=existing.id,
|
||||
owner_id=owner_id,
|
||||
)
|
||||
db.add(duplicate_record)
|
||||
db.commit()
|
||||
@@ -191,6 +275,7 @@ def process_document(
|
||||
file_size=file_size,
|
||||
mime_type=mime_type,
|
||||
is_duplicate=False,
|
||||
owner_id=owner_id,
|
||||
)
|
||||
db.add(new_record)
|
||||
db.commit()
|
||||
@@ -291,6 +376,14 @@ def process_document(
|
||||
new_record.local_filename = new_local_path
|
||||
db.commit()
|
||||
|
||||
# Look up pipeline OCR language override before the session closes.
|
||||
# This reads the OCR step config from the file's assigned pipeline (or
|
||||
# the user/system default pipeline) so the language is available when
|
||||
# dispatching process_with_ocr below.
|
||||
ocr_language = _get_pipeline_ocr_language(db, new_record, owner_id)
|
||||
if ocr_language:
|
||||
logger.info(f"[{task_id}] Pipeline OCR language override: {ocr_language!r}")
|
||||
|
||||
# Store file_id before session closes to avoid DetachedInstanceError
|
||||
file_id = new_record.id
|
||||
|
||||
@@ -320,7 +413,7 @@ def process_document(
|
||||
"Queued for forced OCR processing",
|
||||
file_id=file_id,
|
||||
)
|
||||
process_with_ocr.delay(new_filename, file_id)
|
||||
process_with_ocr.delay(new_filename, file_id, language=ocr_language)
|
||||
return {"file": new_local_path, "status": "Queued for forced OCR", "file_id": file_id}
|
||||
|
||||
# If the file is not a PDF, skip embedded text check and convert to PDF first
|
||||
@@ -477,7 +570,7 @@ def process_document(
|
||||
"Queued for OCR (text quality too low)",
|
||||
file_id=file_id,
|
||||
)
|
||||
process_with_ocr.delay(new_filename, file_id, extracted_text)
|
||||
process_with_ocr.delay(new_filename, file_id, extracted_text, language=ocr_language)
|
||||
return {
|
||||
"file": new_local_path,
|
||||
"status": "Queued for OCR (poor embedded text quality)",
|
||||
@@ -550,5 +643,5 @@ def process_document(
|
||||
"Queued for OCR processing",
|
||||
file_id=file_id,
|
||||
)
|
||||
process_with_ocr.delay(new_filename, file_id)
|
||||
process_with_ocr.delay(new_filename, file_id, language=ocr_language)
|
||||
return {"file": new_local_path, "status": "Queued for OCR", "file_id": file_id}
|
||||
|
||||
@@ -9,7 +9,7 @@ from azure.core.credentials import AzureKeyCredential
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import OcrTaskWithRetry
|
||||
from app.tasks.rotate_pdf_pages import rotate_pdf_pages
|
||||
from app.utils import log_task_progress
|
||||
|
||||
@@ -81,7 +81,7 @@ def check_page_rotation(result, filename, task_id=None):
|
||||
return rotation_data
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
@celery.task(base=OcrTaskWithRetry, bind=True)
|
||||
def process_with_azure_document_intelligence(self, filename: str, file_id: int = None):
|
||||
"""
|
||||
Processes a PDF document using Azure Document Intelligence and overlays OCR text onto
|
||||
|
||||
@@ -17,13 +17,12 @@ task with a multi-engine OCR pipeline that:
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import OcrTaskWithRetry
|
||||
from app.tasks.rotate_pdf_pages import rotate_pdf_pages
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.ocr_provider import OCRResult, embed_text_layer, get_ocr_providers, merge_ocr_results
|
||||
@@ -32,8 +31,14 @@ from app.utils.text_quality import TextSource, check_text_quality, compare_text_
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def process_with_ocr(self, filename: str, file_id: Optional[int] = None, original_text: Optional[str] = None):
|
||||
@celery.task(base=OcrTaskWithRetry, bind=True)
|
||||
def process_with_ocr(
|
||||
self,
|
||||
filename: str,
|
||||
file_id: int | None = None,
|
||||
original_text: str | None = None,
|
||||
language: str | None = None,
|
||||
):
|
||||
"""Run the configured OCR providers on *filename* and continue the pipeline.
|
||||
|
||||
When multiple OCR providers are configured the results are merged using the
|
||||
@@ -47,6 +52,10 @@ def process_with_ocr(self, filename: str, file_id: Optional[int] = None, origina
|
||||
filename: Base name of the file inside ``<workdir>/tmp/``.
|
||||
file_id: Optional database record ID passed through to downstream tasks.
|
||||
original_text: Optional original embedded text for head-to-head comparison.
|
||||
language: Optional Tesseract-style language code(s) (e.g. ``"eng+deu"``)
|
||||
to override the global OCR language settings for this specific run.
|
||||
Pass ``None`` or ``"auto"`` to use the global settings. This
|
||||
enables per-pipeline language configuration.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
log_task_progress(
|
||||
@@ -62,7 +71,7 @@ def process_with_ocr(self, filename: str, file_id: Optional[int] = None, origina
|
||||
if not os.path.exists(tmp_file_path):
|
||||
raise FileNotFoundError(f"Local file not found: {tmp_file_path}")
|
||||
|
||||
providers = get_ocr_providers()
|
||||
providers = get_ocr_providers(language=language)
|
||||
provider_names = [p.name for p in providers]
|
||||
logger.info(f"[{task_id}] Running {len(providers)} OCR provider(s): {provider_names}")
|
||||
|
||||
@@ -122,7 +131,12 @@ def process_with_ocr(self, filename: str, file_id: Optional[int] = None, origina
|
||||
# PDF with ocrmypdf to embed an invisible text layer so the output is
|
||||
# selectable/searchable in PDF viewers.
|
||||
if searchable_pdf_path is None:
|
||||
lang = getattr(settings, "tesseract_language", None) or "eng"
|
||||
# Use the per-call language override; fall back to global setting
|
||||
embed_lang = (
|
||||
language
|
||||
if language and language != "auto"
|
||||
else (getattr(settings, "tesseract_language", None) or "eng")
|
||||
)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"embed_text_layer",
|
||||
@@ -130,7 +144,7 @@ def process_with_ocr(self, filename: str, file_id: Optional[int] = None, origina
|
||||
"Embedding searchable text layer into PDF",
|
||||
file_id=file_id,
|
||||
)
|
||||
embedded = embed_text_layer(tmp_file_path, tmp_file_path, language=lang)
|
||||
embedded = embed_text_layer(tmp_file_path, tmp_file_path, language=embed_lang)
|
||||
if embedded:
|
||||
searchable_pdf_path = tmp_file_path
|
||||
log_task_progress(
|
||||
|
||||
+225
-2
@@ -1,9 +1,232 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Retry configuration for Celery tasks with exponential backoff and jitter.
|
||||
|
||||
Provides a :class:`BaseTaskWithRetry` Celery task base class that implements
|
||||
configurable retry logic with exponential backoff and optional ±20 % random
|
||||
jitter. Pre-defined subclasses offer task-type-specific retry policies:
|
||||
|
||||
* :class:`BaseTaskWithRetry` – general default (3 retries: 60 s, 300 s, 900 s)
|
||||
* :class:`OcrTaskWithRetry` – longer waits for OCR / AI API calls
|
||||
* :class:`UploadTaskWithRetry` – standard waits for cloud-storage uploads
|
||||
|
||||
Usage::
|
||||
|
||||
from app.tasks.retry_config import BaseTaskWithRetry, OcrTaskWithRetry
|
||||
|
||||
@celery.task(base=OcrTaskWithRetry, bind=True)
|
||||
def my_ocr_task(self, ...):
|
||||
...
|
||||
"""
|
||||
|
||||
import logging
|
||||
import random
|
||||
from typing import Any
|
||||
|
||||
from celery import Task
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Defaults
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
#: Default per-retry countdowns in seconds (1 min, 5 min, 15 min).
|
||||
DEFAULT_RETRY_DELAYS: list[int] = [60, 300, 900]
|
||||
|
||||
|
||||
def _parse_delay_string(value: str) -> list[int]:
|
||||
"""Parse a comma-separated string of integers into a list.
|
||||
|
||||
Args:
|
||||
value: Comma-separated integer string, e.g. ``"60,300,900"``.
|
||||
|
||||
Returns:
|
||||
Parsed list of integers, e.g. ``[60, 300, 900]``.
|
||||
"""
|
||||
return [int(v.strip()) for v in value.split(",") if v.strip()]
|
||||
|
||||
|
||||
def compute_countdown(
|
||||
retries: int,
|
||||
base_delays: list[int] | None = None,
|
||||
jitter: bool = True,
|
||||
) -> int:
|
||||
"""Compute the countdown in seconds for the next retry attempt.
|
||||
|
||||
Selects the appropriate base delay for the given retry number. When all
|
||||
defined delays are exhausted the last delay is doubled for each additional
|
||||
attempt. An optional ±20 % jitter is then applied to spread retry storms.
|
||||
|
||||
Args:
|
||||
retries: Current retry count (0-based; 0 = first retry attempt).
|
||||
base_delays: Ordered list of base countdown values (in seconds) for
|
||||
each retry attempt. ``None`` uses :data:`DEFAULT_RETRY_DELAYS`.
|
||||
jitter: When ``True``, apply ±20 % random jitter to the countdown.
|
||||
|
||||
Returns:
|
||||
Countdown in seconds (minimum 1 s).
|
||||
|
||||
Examples::
|
||||
|
||||
>>> compute_countdown(0, [60, 300, 900], jitter=False)
|
||||
60
|
||||
>>> compute_countdown(1, [60, 300, 900], jitter=False)
|
||||
300
|
||||
>>> compute_countdown(3, [60, 300, 900], jitter=False) # beyond list
|
||||
1800
|
||||
"""
|
||||
delays = base_delays if base_delays is not None else DEFAULT_RETRY_DELAYS
|
||||
|
||||
if not delays:
|
||||
base = 60
|
||||
elif retries < len(delays):
|
||||
base = delays[retries]
|
||||
else:
|
||||
# Exhausted defined delays – double the last value for each extra attempt.
|
||||
extra = retries - len(delays) + 1
|
||||
base = delays[-1] * (2**extra)
|
||||
|
||||
if jitter:
|
||||
# ±20 % uniform jitter – not cryptographic, S311 is intentional.
|
||||
jitter_factor = 1.0 + random.uniform(-0.2, 0.2) # noqa: S311
|
||||
base = int(base * jitter_factor)
|
||||
|
||||
return max(base, 1)
|
||||
|
||||
|
||||
class BaseTaskWithRetry(Task):
|
||||
"""Celery task base class with exponential backoff and optional jitter.
|
||||
|
||||
Automatically retries on any :class:`Exception` using delays derived from
|
||||
:attr:`retry_delays`. When :attr:`retry_delays` is ``None`` the value is
|
||||
read from ``TASK_RETRY_DELAYS`` (env-var / settings); if that is also
|
||||
unset :data:`DEFAULT_RETRY_DELAYS` (``[60, 300, 900]`` seconds) is used.
|
||||
|
||||
Override class attributes in subclasses to customise per-task-type policy:
|
||||
|
||||
* ``max_retries`` (``int``) – maximum retry attempts; default ``3``.
|
||||
* ``retry_delays`` (``list[int] | None``) – per-retry countdowns in
|
||||
seconds; ``None`` falls back to settings / :data:`DEFAULT_RETRY_DELAYS`.
|
||||
* ``retry_jitter`` (``bool``) – add ±20 % jitter; default ``True``.
|
||||
"""
|
||||
|
||||
#: Retry on any exception raised inside the task body.
|
||||
autoretry_for = (Exception,)
|
||||
retry_kwargs = {"max_retries": 3, "countdown": 10} # 3 retries, 10s delay
|
||||
retry_backoff = True # Exponential backoff
|
||||
|
||||
#: Maximum number of retry attempts.
|
||||
max_retries: int = 3
|
||||
|
||||
#: Pass max_retries through autoretry_for; no countdown override here
|
||||
#: (our retry() method injects the countdown instead).
|
||||
retry_kwargs: dict = {"max_retries": 3}
|
||||
|
||||
#: Per-retry countdown values (seconds). ``None`` → settings / DEFAULT.
|
||||
retry_delays: list[int] | None = None
|
||||
|
||||
#: Apply ±20 % random jitter to prevent thundering-herd problems.
|
||||
retry_jitter: bool = True
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public API
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def retry(
|
||||
self,
|
||||
args: Any = None,
|
||||
kwargs: Any = None,
|
||||
exc: BaseException | None = None,
|
||||
throw: bool = True,
|
||||
eta: Any = None,
|
||||
countdown: int | None = None,
|
||||
max_retries: int | None = None,
|
||||
**options: Any,
|
||||
) -> Any:
|
||||
"""Retry the task, injecting the backoff countdown when not supplied.
|
||||
|
||||
If *countdown* is not explicitly provided (and *eta* is not set) the
|
||||
countdown is computed via :func:`compute_countdown` using this task's
|
||||
:attr:`retry_delays` and :attr:`retry_jitter` settings.
|
||||
"""
|
||||
if countdown is None and eta is None:
|
||||
countdown = compute_countdown(
|
||||
retries=self.request.retries,
|
||||
base_delays=self._effective_retry_delays(),
|
||||
jitter=self.retry_jitter,
|
||||
)
|
||||
logger.debug(
|
||||
"Retry %d/%d for task %s in %d s",
|
||||
self.request.retries + 1,
|
||||
max_retries if max_retries is not None else self.max_retries,
|
||||
self.name,
|
||||
countdown,
|
||||
)
|
||||
|
||||
return super().retry(
|
||||
args=args,
|
||||
kwargs=kwargs,
|
||||
exc=exc,
|
||||
throw=throw,
|
||||
eta=eta,
|
||||
countdown=countdown,
|
||||
max_retries=max_retries,
|
||||
**options,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _effective_retry_delays(self) -> list[int]:
|
||||
"""Return the retry delays to use, with settings-level override support.
|
||||
|
||||
Priority (highest first):
|
||||
|
||||
1. Explicit class-level ``retry_delays`` attribute (not ``None``).
|
||||
2. ``TASK_RETRY_DELAYS`` environment variable / setting.
|
||||
3. :data:`DEFAULT_RETRY_DELAYS` module-level constant.
|
||||
"""
|
||||
if self.retry_delays is not None:
|
||||
return self.retry_delays
|
||||
|
||||
# Lazily read from settings to avoid circular imports at module load.
|
||||
try:
|
||||
from app.config import settings # noqa: PLC0415
|
||||
|
||||
raw = getattr(settings, "task_retry_delays", None)
|
||||
if raw:
|
||||
if isinstance(raw, list):
|
||||
return [int(v) for v in raw]
|
||||
if isinstance(raw, str):
|
||||
return _parse_delay_string(raw)
|
||||
except Exception as exc: # pragma: no cover
|
||||
logger.debug("Could not read task_retry_delays from settings: %s", exc)
|
||||
|
||||
return DEFAULT_RETRY_DELAYS
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Task-type-specific retry policies
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class OcrTaskWithRetry(BaseTaskWithRetry):
|
||||
"""Retry policy for OCR and document-intelligence API tasks.
|
||||
|
||||
Uses longer initial delays to allow transient API rate-limit windows to
|
||||
clear before the next attempt.
|
||||
|
||||
Default: 3 retries at 120 s, 600 s, 1800 s.
|
||||
"""
|
||||
|
||||
retry_delays: list[int] = [120, 600, 1800]
|
||||
|
||||
|
||||
class UploadTaskWithRetry(BaseTaskWithRetry):
|
||||
"""Retry policy for cloud-storage upload tasks.
|
||||
|
||||
Uses the standard default delays (60 s, 300 s, 900 s) which are
|
||||
appropriate for most transient upload failures (network blips, rate
|
||||
limits, temporary service outages).
|
||||
"""
|
||||
|
||||
# Inherits DEFAULT_RETRY_DELAYS via retry_delays = None.
|
||||
|
||||
+194
-15
@@ -6,12 +6,13 @@ import os
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.models import FileRecord, IntegrationDirection, UserIntegration
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.upload_to_dropbox import upload_to_dropbox
|
||||
from app.tasks.upload_to_email import upload_to_email
|
||||
from app.tasks.upload_to_ftp import upload_to_ftp
|
||||
from app.tasks.upload_to_google_drive import upload_to_google_drive
|
||||
from app.tasks.upload_to_icloud import upload_to_icloud
|
||||
from app.tasks.upload_to_nextcloud import upload_to_nextcloud
|
||||
from app.tasks.upload_to_onedrive import upload_to_onedrive
|
||||
from app.tasks.upload_to_paperless import upload_to_paperless
|
||||
@@ -25,18 +26,32 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _should_upload_to_dropbox():
|
||||
return bool(settings.dropbox_app_key and settings.dropbox_app_secret and settings.dropbox_refresh_token)
|
||||
return bool(
|
||||
getattr(settings, "dropbox_enabled", True)
|
||||
and settings.dropbox_app_key
|
||||
and settings.dropbox_app_secret
|
||||
and settings.dropbox_refresh_token
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_nextcloud():
|
||||
return bool(settings.nextcloud_upload_url and settings.nextcloud_username and settings.nextcloud_password)
|
||||
return bool(
|
||||
getattr(settings, "nextcloud_enabled", True)
|
||||
and settings.nextcloud_upload_url
|
||||
and settings.nextcloud_username
|
||||
and settings.nextcloud_password
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_paperless():
|
||||
return bool(settings.paperless_ngx_api_token and settings.paperless_host)
|
||||
return bool(
|
||||
getattr(settings, "paperless_enabled", True) and settings.paperless_ngx_api_token and settings.paperless_host
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_google_drive():
|
||||
if not getattr(settings, "google_drive_enabled", True):
|
||||
return False
|
||||
# Check for OAuth configuration
|
||||
if getattr(settings, "google_drive_use_oauth", False):
|
||||
return bool(
|
||||
@@ -51,36 +66,66 @@ def _should_upload_to_google_drive():
|
||||
|
||||
|
||||
def _should_upload_to_webdav():
|
||||
return bool(settings.webdav_url and settings.webdav_username and settings.webdav_password)
|
||||
return bool(
|
||||
getattr(settings, "webdav_enabled", True)
|
||||
and settings.webdav_url
|
||||
and settings.webdav_username
|
||||
and settings.webdav_password
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_ftp():
|
||||
return bool(settings.ftp_host and settings.ftp_username and settings.ftp_password)
|
||||
return bool(
|
||||
getattr(settings, "ftp_enabled", True) and settings.ftp_host and settings.ftp_username and settings.ftp_password
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_sftp():
|
||||
return bool(settings.sftp_host and settings.sftp_username and (settings.sftp_password or settings.sftp_private_key))
|
||||
return bool(
|
||||
getattr(settings, "sftp_enabled", True)
|
||||
and settings.sftp_host
|
||||
and settings.sftp_username
|
||||
and (settings.sftp_password or settings.sftp_private_key)
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_email():
|
||||
return bool(
|
||||
settings.email_host and settings.email_username and settings.email_password and settings.email_default_recipient
|
||||
getattr(settings, "dest_email_enabled", True)
|
||||
and settings.dest_email_host
|
||||
and settings.dest_email_username
|
||||
and settings.dest_email_password
|
||||
and settings.dest_email_default_recipient
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_onedrive():
|
||||
return bool(settings.onedrive_client_id and settings.onedrive_client_secret and settings.onedrive_refresh_token)
|
||||
return bool(
|
||||
getattr(settings, "onedrive_enabled", True)
|
||||
and settings.onedrive_client_id
|
||||
and settings.onedrive_client_secret
|
||||
and settings.onedrive_refresh_token
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_s3():
|
||||
return bool(settings.s3_bucket_name and settings.aws_access_key_id and settings.aws_secret_access_key)
|
||||
return bool(
|
||||
getattr(settings, "s3_enabled", True)
|
||||
and settings.s3_bucket_name
|
||||
and settings.aws_access_key_id
|
||||
and settings.aws_secret_access_key
|
||||
)
|
||||
|
||||
|
||||
def _should_upload_to_icloud():
|
||||
return bool(getattr(settings, "icloud_enabled", True) and settings.icloud_username and settings.icloud_password)
|
||||
|
||||
|
||||
def get_configured_services_from_validator():
|
||||
"""
|
||||
Use the config validator to determine which services are configured properly.
|
||||
Use the config validator to determine which services are configured and enabled.
|
||||
Returns a dictionary with service names as keys and boolean values indicating
|
||||
whether they're properly configured.
|
||||
whether they're properly configured AND explicitly enabled.
|
||||
"""
|
||||
providers = get_provider_status()
|
||||
|
||||
@@ -95,18 +140,20 @@ def get_configured_services_from_validator():
|
||||
"Email": "email",
|
||||
"OneDrive": "onedrive",
|
||||
"S3 Storage": "s3",
|
||||
"iCloud Drive": "icloud",
|
||||
}
|
||||
|
||||
result = {}
|
||||
for provider_name, internal_name in service_map.items():
|
||||
if provider_name in providers:
|
||||
result[internal_name] = providers[provider_name].get("configured", False)
|
||||
provider = providers[provider_name]
|
||||
result[internal_name] = provider.get("configured", False) and provider.get("enabled", True)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def send_to_all_destinations(self, file_path: str, use_validator=True, file_id: int = None):
|
||||
def send_to_all_destinations(self, file_path: str, use_validator=True, file_id: int = None, folder_overrides=None):
|
||||
"""
|
||||
Distribute a file to all configured storage destinations.
|
||||
|
||||
@@ -115,6 +162,10 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
|
||||
use_validator: Whether to use the config validator to determine enabled services
|
||||
(if False, falls back to individual checks)
|
||||
file_id: Optional file ID to associate with logs
|
||||
folder_overrides: Optional dict mapping provider names to folder override strings.
|
||||
When set, the override is passed to the upload task which uses it
|
||||
instead of the provider's default folder. Example:
|
||||
{"dropbox": "/Documents/pdfa", "s3": "docs/pdfa/"}
|
||||
"""
|
||||
task_id = self.request.id
|
||||
|
||||
@@ -199,6 +250,11 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
|
||||
"should_upload": _should_upload_to_s3,
|
||||
"upload_func": upload_to_s3,
|
||||
},
|
||||
{
|
||||
"name": "icloud",
|
||||
"should_upload": _should_upload_to_icloud,
|
||||
"upload_func": upload_to_icloud,
|
||||
},
|
||||
]
|
||||
|
||||
# Optionally get configuration status from validator
|
||||
@@ -236,7 +292,10 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
|
||||
task_id, f"queue_{service_name}", "in_progress", f"Queueing upload to {service_name}", file_id=file_id
|
||||
)
|
||||
try:
|
||||
task = service["upload_func"].delay(file_path, file_id=file_id)
|
||||
kwargs = {"file_id": file_id}
|
||||
if folder_overrides and service_name in folder_overrides:
|
||||
kwargs["folder_override"] = folder_overrides[service_name]
|
||||
task = service["upload_func"].delay(file_path, **kwargs)
|
||||
results[f"{service_name}_task_id"] = task.id
|
||||
queued_count += 1
|
||||
log_task_progress(
|
||||
@@ -251,3 +310,123 @@ def send_to_all_destinations(self, file_path: str, use_validator=True, file_id:
|
||||
log_task_progress(task_id, "send_to_all_destinations", "success", f"Queued {queued_count} uploads", file_id=file_id)
|
||||
|
||||
return {"status": "Queued", "file_path": file_path, "tasks": results}
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def send_to_user_destinations(self, file_path: str, owner_id: str, file_id: int | None = None):
|
||||
"""Dispatch uploads to all active DESTINATION UserIntegrations for *owner_id*.
|
||||
|
||||
This is the user-specific counterpart of :func:`send_to_all_destinations`.
|
||||
It queries the ``user_integrations`` table for records where:
|
||||
|
||||
* ``owner_id`` matches the document owner,
|
||||
* ``direction == "DESTINATION"``, and
|
||||
* ``is_active == True``.
|
||||
|
||||
One :func:`upload_to_user_integration` Celery task is queued for each
|
||||
matching integration so that uploads proceed asynchronously and
|
||||
independently.
|
||||
|
||||
Args:
|
||||
file_path: Absolute path to the processed document file.
|
||||
owner_id: The stable user identifier from ``FileRecord.owner_id``.
|
||||
file_id: Optional ``FileRecord.id`` used for progress logging.
|
||||
|
||||
Returns:
|
||||
A dict summarising how many integrations were queued.
|
||||
"""
|
||||
from app.tasks.upload_to_user_integration import upload_to_user_integration
|
||||
|
||||
task_id = self.request.id
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not os.path.exists(file_path):
|
||||
error_msg = f"File not found: {file_path}"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(task_id, "send_to_user_destinations", "failure", error_msg, file_id=file_id)
|
||||
raise FileNotFoundError(error_msg)
|
||||
|
||||
logger.info("[%s] Sending %s to user destinations for owner=%s", task_id, filename, owner_id)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"send_to_user_destinations",
|
||||
"in_progress",
|
||||
f"Distributing {filename} to user integrations",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
with SessionLocal() as db:
|
||||
integrations = (
|
||||
db.query(UserIntegration)
|
||||
.filter(
|
||||
UserIntegration.owner_id == owner_id,
|
||||
UserIntegration.direction == IntegrationDirection.DESTINATION,
|
||||
UserIntegration.is_active.is_(True),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
# Snapshot the IDs so we don't keep the session open
|
||||
integration_ids = [(i.id, i.name, i.integration_type) for i in integrations]
|
||||
|
||||
queued = 0
|
||||
task_results: dict[str, str] = {}
|
||||
|
||||
for int_id, int_name, int_type in integration_ids:
|
||||
logger.info("[%s] Queueing upload for integration %d (%s '%s')", task_id, int_id, int_type, int_name)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
f"queue_user_integration_{int_id}",
|
||||
"in_progress",
|
||||
f"Queueing upload to {int_type} '{int_name}'",
|
||||
file_id=file_id,
|
||||
)
|
||||
try:
|
||||
celery_task = upload_to_user_integration.delay(file_path, int_id, file_id)
|
||||
task_results[f"integration_{int_id}_task_id"] = celery_task.id
|
||||
queued += 1
|
||||
log_task_progress(
|
||||
task_id,
|
||||
f"queue_user_integration_{int_id}",
|
||||
"success",
|
||||
f"Queued upload to {int_type} '{int_name}'",
|
||||
file_id=file_id,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
error_msg = str(exc)
|
||||
logger.error("[%s] Failed to queue upload for integration %d: %s", task_id, int_id, error_msg)
|
||||
task_results[f"integration_{int_id}_error"] = error_msg
|
||||
log_task_progress(
|
||||
task_id,
|
||||
f"queue_user_integration_{int_id}",
|
||||
"failure",
|
||||
f"Failed to queue {int_type} '{int_name}': {error_msg}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
logger.info("[%s] Queued %d user-integration upload(s) for owner=%s", task_id, queued, owner_id)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"send_to_user_destinations",
|
||||
"success",
|
||||
f"Queued {queued} user-integration upload(s)",
|
||||
file_id=file_id,
|
||||
)
|
||||
return {"status": "Queued", "file_path": file_path, "queued": queued, "tasks": task_results}
|
||||
|
||||
|
||||
def get_user_destination_count(owner_id: str) -> int:
|
||||
"""Return the number of active DESTINATION integrations for *owner_id*.
|
||||
|
||||
A count of zero means no user-specific destinations are configured and
|
||||
the caller should fall back to the global :func:`send_to_all_destinations`.
|
||||
"""
|
||||
with SessionLocal() as db:
|
||||
return (
|
||||
db.query(UserIntegration)
|
||||
.filter(
|
||||
UserIntegration.owner_id == owner_id,
|
||||
UserIntegration.direction == IntegrationDirection.DESTINATION,
|
||||
UserIntegration.is_active.is_(True),
|
||||
)
|
||||
.count()
|
||||
)
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Celery task to apply pending subscription changes that have become due.
|
||||
|
||||
Runs daily to ensure that scheduled downgrades are applied on time.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from app.celery_app import celery
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(name="app.tasks.subscription_tasks.apply_pending_subscription_changes_all")
|
||||
def apply_pending_subscription_changes_all() -> dict[str, int]:
|
||||
"""Apply all pending subscription changes whose effective date has arrived.
|
||||
|
||||
Iterates over every ``UserProfile`` that has a pending change and calls
|
||||
:func:`app.utils.subscription.apply_pending_subscription_changes` for
|
||||
each one.
|
||||
|
||||
Returns:
|
||||
A dict with ``{"applied": <count>, "checked": <count>}``.
|
||||
"""
|
||||
from app.database import SessionLocal
|
||||
from app.models import UserProfile
|
||||
from app.utils.subscription import apply_pending_subscription_changes
|
||||
|
||||
applied = 0
|
||||
checked = 0
|
||||
db = SessionLocal()
|
||||
try:
|
||||
profiles = (
|
||||
db.query(UserProfile)
|
||||
.filter(
|
||||
UserProfile.subscription_change_pending_tier.isnot(None),
|
||||
UserProfile.subscription_change_pending_date.isnot(None),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
for profile in profiles:
|
||||
checked += 1
|
||||
if apply_pending_subscription_changes(db, profile.user_id):
|
||||
applied += 1
|
||||
except Exception as exc:
|
||||
logger.error("Error in apply_pending_subscription_changes_all: %s", exc)
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
logger.info("apply_pending_subscription_changes_all: checked=%d applied=%d", checked, applied)
|
||||
return {"checked": checked, "applied": applied}
|
||||
@@ -9,7 +9,7 @@ from dropbox.exceptions import ApiError, AuthError
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.filename_utils import extract_remote_path, get_unique_filename
|
||||
|
||||
@@ -102,8 +102,8 @@ def get_dropbox_client():
|
||||
raise
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_dropbox(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_dropbox(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Upload a file to Dropbox.
|
||||
|
||||
@@ -147,7 +147,7 @@ def upload_to_dropbox(self, file_path: str, file_id: int = None):
|
||||
dbx = get_dropbox_client()
|
||||
|
||||
# Calculate remote path based on local file structure
|
||||
remote_base = settings.dropbox_folder or ""
|
||||
remote_base = folder_override if folder_override is not None else (settings.dropbox_folder or "")
|
||||
remote_path = extract_remote_path(file_path, settings.workdir, remote_base)
|
||||
|
||||
# Function to check if file exists in Dropbox
|
||||
|
||||
@@ -15,7 +15,7 @@ from jinja2 import Environment, FileSystemLoader, select_autoescape
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -125,11 +125,11 @@ def attach_logo(msg):
|
||||
def _prepare_recipients(recipients):
|
||||
"""Helper function to prepare email recipients list."""
|
||||
if not recipients:
|
||||
if not settings.email_default_recipient:
|
||||
if not settings.dest_email_default_recipient:
|
||||
error_msg = "No recipients specified and no default recipient configured"
|
||||
logger.error(error_msg)
|
||||
return None, error_msg
|
||||
return [settings.email_default_recipient], None
|
||||
return [settings.dest_email_default_recipient], None
|
||||
elif isinstance(recipients, str):
|
||||
return [recipients], None # Convert single email to list
|
||||
return recipients, None
|
||||
@@ -139,17 +139,17 @@ def _send_email_with_smtp(msg, filename, recipients):
|
||||
"""Helper function to handle SMTP connection and sending."""
|
||||
try:
|
||||
# First try to resolve the hostname
|
||||
socket.gethostbyname(settings.email_host)
|
||||
socket.gethostbyname(settings.dest_email_host)
|
||||
|
||||
# Connect to the SMTP server
|
||||
with smtplib.SMTP(settings.email_host, settings.email_port, timeout=30) as server:
|
||||
with smtplib.SMTP(settings.dest_email_host, settings.dest_email_port, timeout=30) as server:
|
||||
# Use TLS if specified
|
||||
if settings.email_use_tls:
|
||||
if settings.dest_email_use_tls:
|
||||
server.starttls()
|
||||
|
||||
# Login if credentials are provided
|
||||
if settings.email_username and settings.email_password:
|
||||
server.login(settings.email_username, settings.email_password)
|
||||
if settings.dest_email_username and settings.dest_email_password:
|
||||
server.login(settings.dest_email_username, settings.dest_email_password)
|
||||
|
||||
# Send the email
|
||||
server.send_message(msg)
|
||||
@@ -157,16 +157,16 @@ def _send_email_with_smtp(msg, filename, recipients):
|
||||
logger.info(f"Successfully sent {filename} via email to {', '.join(recipients)}")
|
||||
return None
|
||||
except socket.gaierror as e:
|
||||
error_msg = f"Failed to resolve email host: {settings.email_host} - {str(e)}"
|
||||
error_msg = f"Failed to resolve email host: {settings.dest_email_host} - {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return {"status": "Failed", "reason": error_msg, "error": str(e)}
|
||||
except (ConnectionRefusedError, TimeoutError) as e:
|
||||
error_msg = f"Connection error to SMTP server {settings.email_host}:{settings.email_port} - {str(e)}"
|
||||
error_msg = f"Connection error to SMTP server {settings.dest_email_host}:{settings.dest_email_port} - {str(e)}"
|
||||
logger.error(error_msg)
|
||||
return {"status": "Failed", "reason": error_msg, "error": str(e)}
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_email(
|
||||
self,
|
||||
file_path: str,
|
||||
@@ -205,17 +205,17 @@ def upload_to_email(
|
||||
# Extract filename
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
# Check if email settings are configured
|
||||
if not settings.email_host:
|
||||
error_msg = "Email host is not configured"
|
||||
# Check if email destination settings are configured
|
||||
if not settings.dest_email_host:
|
||||
error_msg = "Email destination host is not configured (DEST_EMAIL_HOST)"
|
||||
logger.error(f"[{task_id}] {error_msg}")
|
||||
log_task_progress(task_id, "upload_to_email", "skipped", error_msg, file_id=file_id)
|
||||
return {"status": "Skipped", "reason": error_msg}
|
||||
|
||||
# Log email configuration for debugging
|
||||
logger.debug(
|
||||
f"[{task_id}] Email config - Host: {settings.email_host}, Port: {settings.email_port}, "
|
||||
f"Username: {settings.email_username}, TLS: {settings.email_use_tls}"
|
||||
f"[{task_id}] Email destination config - Host: {settings.dest_email_host}, Port: {settings.dest_email_port}, "
|
||||
f"Username: {settings.dest_email_username}, TLS: {settings.dest_email_use_tls}"
|
||||
)
|
||||
|
||||
# Process recipients
|
||||
@@ -236,7 +236,7 @@ def upload_to_email(
|
||||
try:
|
||||
# Create the email
|
||||
msg = MIMEMultipart("related")
|
||||
msg["From"] = settings.email_sender or settings.email_username
|
||||
msg["From"] = settings.dest_email_sender or settings.dest_email_username
|
||||
msg["To"] = ", ".join(recipients)
|
||||
msg["Subject"] = subject
|
||||
|
||||
|
||||
@@ -8,14 +8,14 @@ import os
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_ftp(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_ftp(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Uploads a file to an FTP server in the configured folder.
|
||||
|
||||
@@ -97,10 +97,11 @@ def upload_to_ftp(self, file_path: str, file_id: int = None):
|
||||
ftp.login(user=settings.ftp_username, passwd=settings.ftp_password)
|
||||
|
||||
# Change to target directory if specified
|
||||
if settings.ftp_folder:
|
||||
ftp_folder_setting = folder_override if folder_override is not None else settings.ftp_folder
|
||||
if ftp_folder_setting:
|
||||
try:
|
||||
# Try to navigate to the directory, create if it doesn't exist
|
||||
ftp_folder = settings.ftp_folder
|
||||
ftp_folder = ftp_folder_setting
|
||||
# Remove leading slash if present
|
||||
if ftp_folder.startswith("/"):
|
||||
ftp_folder = ftp_folder[1:]
|
||||
@@ -138,7 +139,7 @@ def upload_to_ftp(self, file_path: str, file_id: int = None):
|
||||
"status": "Completed",
|
||||
"file": file_path,
|
||||
"ftp_host": settings.ftp_host,
|
||||
"ftp_path": f"{settings.ftp_folder}/{filename}" if settings.ftp_folder else filename,
|
||||
"ftp_path": f"{ftp_folder_setting}/{filename}" if ftp_folder_setting else filename,
|
||||
"used_tls": used_tls,
|
||||
}
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ from googleapiclient.http import MediaFileUpload
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -152,8 +152,10 @@ def truncate_property_value(key, value, max_bytes=100):
|
||||
return str_value
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_google_drive(self, file_path: str, include_metadata=True, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_google_drive(
|
||||
self, file_path: str, include_metadata=True, file_id: int = None, folder_override: str = None
|
||||
):
|
||||
"""
|
||||
Uploads a file to Google Drive in the configured folder with optional metadata.
|
||||
|
||||
@@ -201,8 +203,9 @@ def upload_to_google_drive(self, file_path: str, include_metadata=True, file_id:
|
||||
}
|
||||
|
||||
# If folder ID is specified, set parent folder
|
||||
if settings.google_drive_folder_id:
|
||||
file_metadata["parents"] = [settings.google_drive_folder_id]
|
||||
gdrive_folder_id = folder_override if folder_override is not None else settings.google_drive_folder_id
|
||||
if gdrive_folder_id:
|
||||
file_metadata["parents"] = [gdrive_folder_id]
|
||||
|
||||
# Add custom properties if metadata exists
|
||||
if metadata:
|
||||
|
||||
@@ -0,0 +1,177 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
"""Upload files to Apple iCloud Drive via the pyicloud library.
|
||||
|
||||
This module uses the ``pyicloud`` library to authenticate with Apple's iCloud
|
||||
service and upload files to iCloud Drive. Because Apple does not offer a public
|
||||
REST API for iCloud Drive, this integration relies on the *unofficial*
|
||||
reverse-engineered protocol implemented by ``pyicloud``.
|
||||
|
||||
Requirements
|
||||
~~~~~~~~~~~~
|
||||
* An Apple ID with iCloud Drive enabled.
|
||||
* An **app-specific password** generated at https://appleid.apple.com (required
|
||||
when two-factor authentication is active – which is the default for all modern
|
||||
Apple IDs).
|
||||
* The ``pyicloud`` Python package (``pip install pyicloud``).
|
||||
|
||||
Configuration
|
||||
~~~~~~~~~~~~~
|
||||
Set the following environment variables (or ``app/config.py`` fields):
|
||||
|
||||
* ``ICLOUD_USERNAME`` – Apple ID email address.
|
||||
* ``ICLOUD_PASSWORD`` – App-specific password.
|
||||
* ``ICLOUD_FOLDER`` – Target folder path inside iCloud Drive, using ``/`` as
|
||||
the separator (e.g. ``Documents/Uploads``). The folder is created
|
||||
automatically if it does not exist.
|
||||
* ``ICLOUD_COOKIE_DIRECTORY`` – (Optional) Directory for persisting session
|
||||
cookies so that re-authentication is avoided between task runs. Defaults to
|
||||
``~/.pyicloud``.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _get_icloud_api(
|
||||
username: str,
|
||||
password: str,
|
||||
cookie_directory: str | None = None,
|
||||
):
|
||||
"""Return an authenticated ``PyiCloudService`` instance.
|
||||
|
||||
Args:
|
||||
username: Apple ID email address.
|
||||
password: App-specific password.
|
||||
cookie_directory: Optional directory for session cookies.
|
||||
|
||||
Returns:
|
||||
An authenticated ``PyiCloudService`` instance.
|
||||
|
||||
Raises:
|
||||
ImportError: If ``pyicloud`` is not installed.
|
||||
ValueError: If authentication fails or 2FA is required interactively.
|
||||
"""
|
||||
from pyicloud import PyiCloudService # noqa: S404 – unofficial third-party iCloud client
|
||||
|
||||
kwargs: dict = {}
|
||||
if cookie_directory:
|
||||
kwargs["cookie_directory"] = cookie_directory
|
||||
|
||||
api = PyiCloudService(username, password, **kwargs)
|
||||
|
||||
# If 2SA/2FA is required the user must use an app-specific password instead.
|
||||
if api.requires_2sa or api.requires_2fa:
|
||||
raise ValueError(
|
||||
"iCloud account requires two-factor authentication. "
|
||||
"Please generate an app-specific password at https://appleid.apple.com "
|
||||
"and use it as ICLOUD_PASSWORD."
|
||||
)
|
||||
|
||||
return api
|
||||
|
||||
|
||||
def _navigate_to_folder(drive_root, folder_path: str):
|
||||
"""Navigate into (or create) the folder hierarchy described by *folder_path*.
|
||||
|
||||
Args:
|
||||
drive_root: The iCloud Drive root node (``api.drive``).
|
||||
folder_path: ``/``-separated path such as ``Documents/Uploads``.
|
||||
|
||||
Returns:
|
||||
The drive node representing the target folder.
|
||||
"""
|
||||
node = drive_root
|
||||
if not folder_path:
|
||||
return node
|
||||
|
||||
parts = [p for p in folder_path.strip("/").split("/") if p]
|
||||
for part in parts:
|
||||
children = {child.name: child for child in node.dir()}
|
||||
if part in children:
|
||||
node = children[part]
|
||||
else:
|
||||
# Create the missing folder
|
||||
node = node.mkdir(part)
|
||||
return node
|
||||
|
||||
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_icloud(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""Upload a file to Apple iCloud Drive.
|
||||
|
||||
Args:
|
||||
file_path: Local path to the file to upload.
|
||||
file_id: Optional ``FileRecord.id`` for progress logging.
|
||||
folder_override: If provided, overrides the default ``ICLOUD_FOLDER``
|
||||
setting for this upload.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
logger.info(f"[{task_id}] Starting iCloud Drive upload: {file_path}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_to_icloud",
|
||||
"in_progress",
|
||||
f"Uploading to iCloud Drive: {os.path.basename(file_path)}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Validate inputs
|
||||
# ------------------------------------------------------------------
|
||||
if not os.path.exists(file_path):
|
||||
error_msg = f"File not found: {file_path}"
|
||||
logger.error(f"[{task_id}] {error_msg}")
|
||||
log_task_progress(task_id, "upload_to_icloud", "failure", error_msg, file_id=file_id)
|
||||
raise FileNotFoundError(error_msg)
|
||||
|
||||
if not settings.icloud_username or not settings.icloud_password:
|
||||
error_msg = "iCloud credentials are not configured (ICLOUD_USERNAME / ICLOUD_PASSWORD)"
|
||||
logger.error(f"[{task_id}] {error_msg}")
|
||||
log_task_progress(task_id, "upload_to_icloud", "failure", error_msg, file_id=file_id)
|
||||
raise ValueError(error_msg)
|
||||
|
||||
filename = os.path.basename(file_path)
|
||||
target_folder = folder_override if folder_override is not None else (settings.icloud_folder or "")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Authenticate & upload
|
||||
# ------------------------------------------------------------------
|
||||
try:
|
||||
api = _get_icloud_api(
|
||||
settings.icloud_username,
|
||||
settings.icloud_password,
|
||||
settings.icloud_cookie_directory,
|
||||
)
|
||||
|
||||
folder_node = _navigate_to_folder(api.drive, target_folder)
|
||||
|
||||
with open(file_path, "rb") as fh:
|
||||
folder_node.upload(fh)
|
||||
|
||||
logger.info(f"[{task_id}] Successfully uploaded {filename} to iCloud Drive folder '{target_folder}'")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"upload_to_icloud",
|
||||
"success",
|
||||
f"Uploaded to iCloud Drive: {filename}",
|
||||
file_id=file_id,
|
||||
)
|
||||
return {
|
||||
"status": "Completed",
|
||||
"file": file_path,
|
||||
"icloud_folder": target_folder or "/",
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
error_msg = f"Error uploading {filename} to iCloud Drive: {e}"
|
||||
logger.error(f"[{task_id}] {error_msg}")
|
||||
log_task_progress(task_id, "upload_to_icloud", "failure", error_msg, file_id=file_id)
|
||||
raise RuntimeError(error_msg) from e
|
||||
@@ -8,15 +8,15 @@ from requests.auth import HTTPBasicAuth
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.filename_utils import extract_remote_path, get_unique_filename
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_nextcloud(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_nextcloud(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Upload a file to Nextcloud WebDAV.
|
||||
|
||||
@@ -60,7 +60,9 @@ def upload_to_nextcloud(self, file_path: str, file_id: int = None):
|
||||
webdav_url += "/"
|
||||
|
||||
# Calculate remote path based on local file structure
|
||||
remote_base = getattr(settings, "nextcloud_folder", "") or ""
|
||||
remote_base = (
|
||||
folder_override if folder_override is not None else (getattr(settings, "nextcloud_folder", "") or "")
|
||||
)
|
||||
remote_path = extract_remote_path(file_path, settings.workdir, remote_base)
|
||||
full_url = f"{webdav_url}/{remote_path}"
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import requests
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -209,8 +209,8 @@ def upload_large_file(file_path, upload_url):
|
||||
return response.json()
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_onedrive(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_onedrive(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Uploads a file to OneDrive in the configured folder.
|
||||
|
||||
@@ -248,15 +248,17 @@ def upload_to_onedrive(self, file_path: str, file_id: int = None):
|
||||
# Get access token
|
||||
access_token = get_onedrive_token()
|
||||
|
||||
onedrive_folder = folder_override if folder_override is not None else settings.onedrive_folder_path
|
||||
|
||||
# Create upload session
|
||||
upload_url = create_upload_session(filename, settings.onedrive_folder_path, access_token)
|
||||
upload_url = create_upload_session(filename, onedrive_folder, access_token)
|
||||
|
||||
# Upload the file
|
||||
result = upload_large_file(file_path, upload_url)
|
||||
|
||||
# Log success
|
||||
web_url = result.get("webUrl", "Not available")
|
||||
logger.info(f"[{task_id}] Successfully uploaded {filename} to OneDrive at path {settings.onedrive_folder_path}")
|
||||
logger.info(f"[{task_id}] Successfully uploaded {filename} to OneDrive at path {onedrive_folder}")
|
||||
logger.info(f"[{task_id}] File accessible at: {web_url}")
|
||||
log_task_progress(
|
||||
task_id, "upload_to_onedrive", "success", f"Uploaded to OneDrive: {filename}", file_id=file_id
|
||||
@@ -265,7 +267,7 @@ def upload_to_onedrive(self, file_path: str, file_id: int = None):
|
||||
return {
|
||||
"status": "Completed",
|
||||
"file_path": file_path,
|
||||
"onedrive_path": f"{settings.onedrive_folder_path}/{filename}",
|
||||
"onedrive_path": f"{onedrive_folder}/{filename}",
|
||||
"web_url": web_url,
|
||||
}
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import requests
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -199,7 +199,7 @@ def set_document_custom_fields(doc_id: int, custom_fields: dict, task_id: str) -
|
||||
logger.error(f"[{task_id}] Response: {getattr(exc.response, 'text', '<no response>')}")
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_paperless(self, file_path: str, file_id: int = None):
|
||||
"""
|
||||
Uploads a file to Paperless-ngx and sets custom fields from metadata.
|
||||
|
||||
@@ -8,14 +8,14 @@ from botocore.exceptions import ClientError
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_s3(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_s3(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Uploads a file to Amazon S3 in the configured bucket and folder.
|
||||
|
||||
@@ -61,9 +61,10 @@ def upload_to_s3(self, file_path: str, file_id: int = None):
|
||||
)
|
||||
|
||||
# Construct the S3 key (path within the bucket)
|
||||
if settings.s3_folder_prefix:
|
||||
s3_folder = folder_override if folder_override is not None else settings.s3_folder_prefix
|
||||
if s3_folder:
|
||||
# Ensure folder prefix ends with a slash
|
||||
folder_prefix = settings.s3_folder_prefix
|
||||
folder_prefix = s3_folder
|
||||
if not folder_prefix.endswith("/"):
|
||||
folder_prefix += "/"
|
||||
s3_key = f"{folder_prefix}{filename}"
|
||||
|
||||
@@ -7,15 +7,15 @@ import paramiko
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.filename_utils import extract_remote_path, get_unique_filename, sanitize_filename
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_sftp(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_sftp(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Upload a file to an SFTP server.
|
||||
|
||||
@@ -95,7 +95,7 @@ def upload_to_sftp(self, file_path: str, file_id: int = None):
|
||||
sftp = ssh.open_sftp()
|
||||
|
||||
# Calculate remote path based on local file structure
|
||||
remote_base = settings.sftp_folder or ""
|
||||
remote_base = folder_override if folder_override is not None else (settings.sftp_folder or "")
|
||||
remote_path = extract_remote_path(file_path, settings.workdir, remote_base)
|
||||
|
||||
# Ensure the remote path starts with a slash if the base folder does
|
||||
|
||||
@@ -0,0 +1,763 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Upload dispatcher for user-specific destination integrations.
|
||||
|
||||
This module provides a Celery task that uploads a processed document to a
|
||||
specific ``UserIntegration`` record using that integration's own stored
|
||||
config and decrypted credentials — instead of the global application settings.
|
||||
|
||||
It is the per-destination counterpart of :func:`send_to_all_destinations`
|
||||
and is dispatched by :func:`send_to_user_destinations` once per active
|
||||
DESTINATION integration that belongs to the document's owner.
|
||||
|
||||
Credential shapes per integration type (mirrors the UserIntegration docstring):
|
||||
|
||||
DROPBOX credentials = {"refresh_token", "app_key", "app_secret"}
|
||||
config = {"folder": "/DocuElevate"}
|
||||
|
||||
S3 credentials = {"access_key_id", "secret_access_key"}
|
||||
config = {"bucket", "region", "endpoint_url", "folder_prefix"}
|
||||
|
||||
GOOGLE_DRIVE
|
||||
OAuth credentials = {"client_id", "client_secret", "refresh_token"}
|
||||
config = {"folder_id"}
|
||||
SA credentials = {"credentials_json"}
|
||||
config = {"folder_id"}
|
||||
|
||||
ONEDRIVE credentials = {"client_id", "client_secret", "refresh_token"}
|
||||
config = {"folder_path", "tenant_id"}
|
||||
|
||||
WEBDAV /
|
||||
NEXTCLOUD credentials = {"username", "password"}
|
||||
config = {"url", "folder"}
|
||||
|
||||
FTP credentials = {"password"}
|
||||
config = {"host", "username", "port", "folder", "use_tls"}
|
||||
|
||||
SFTP credentials = {"password"} or {"private_key"}
|
||||
config = {"host", "username", "port", "folder"}
|
||||
|
||||
EMAIL credentials = {"password"}
|
||||
config = {"host", "username", "port", "recipient",
|
||||
"use_tls", "sender_name"}
|
||||
|
||||
PAPERLESS credentials = {"api_token"}
|
||||
config = {"host"}
|
||||
|
||||
RCLONE credentials = {"rclone_conf"} (full rclone config file text)
|
||||
config = {"remote": "myremote:", "folder": "dest/path"}
|
||||
"""
|
||||
|
||||
import ftplib # nosec B402
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import subprocess # nosec B404
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
from urllib.parse import urljoin
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.database import SessionLocal
|
||||
from app.models import IntegrationType, UserIntegration
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils.encryption import decrypt_value
|
||||
from app.utils.logging import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Maximum characters to store in UserIntegration.last_error
|
||||
_MAX_ERROR_LENGTH = 500
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-type upload helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _upload_dropbox(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to Dropbox using per-user OAuth credentials."""
|
||||
import dropbox
|
||||
|
||||
app_key = creds.get("app_key") or ""
|
||||
app_secret = creds.get("app_secret") or ""
|
||||
refresh_token = creds.get("refresh_token") or ""
|
||||
|
||||
if not (app_key and app_secret and refresh_token):
|
||||
raise ValueError("Dropbox integration is missing app_key, app_secret or refresh_token in credentials")
|
||||
|
||||
dbx = dropbox.Dropbox(app_key=app_key, app_secret=app_secret, oauth2_refresh_token=refresh_token)
|
||||
|
||||
remote_folder = cfg.get("folder", "/DocuElevate").rstrip("/")
|
||||
filename = os.path.basename(file_path)
|
||||
remote_path = f"{remote_folder}/{filename}"
|
||||
if not remote_path.startswith("/"):
|
||||
remote_path = "/" + remote_path
|
||||
|
||||
file_size = os.path.getsize(file_path)
|
||||
with open(file_path, "rb") as fh:
|
||||
if file_size > 10 * 1024 * 1024:
|
||||
chunk_size = 4 * 1024 * 1024
|
||||
session_start = dbx.files_upload_session_start(fh.read(chunk_size))
|
||||
cursor = dropbox.files.UploadSessionCursor(session_start.session_id, fh.tell())
|
||||
while fh.tell() < file_size:
|
||||
if (file_size - fh.tell()) <= chunk_size:
|
||||
dbx.files_upload_session_finish(
|
||||
fh.read(chunk_size),
|
||||
cursor,
|
||||
dropbox.files.CommitInfo(path=remote_path, mode=dropbox.files.WriteMode.overwrite),
|
||||
)
|
||||
else:
|
||||
dbx.files_upload_session_append_v2(fh.read(chunk_size), cursor)
|
||||
cursor.offset = fh.tell()
|
||||
else:
|
||||
dbx.files_upload(fh.read(), remote_path, mode=dropbox.files.WriteMode.overwrite)
|
||||
|
||||
logger.info("[%s] Dropbox upload complete: %s", task_id, remote_path)
|
||||
return {"status": "Completed", "dropbox_path": remote_path}
|
||||
|
||||
|
||||
def _upload_s3(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to Amazon S3 (or S3-compatible) using per-user credentials."""
|
||||
import boto3
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
bucket = cfg.get("bucket") or ""
|
||||
region = cfg.get("region") or "us-east-1"
|
||||
endpoint_url = cfg.get("endpoint_url") or None
|
||||
folder_prefix = cfg.get("folder_prefix") or ""
|
||||
storage_class = cfg.get("storage_class") or "STANDARD"
|
||||
|
||||
access_key = creds.get("access_key_id") or ""
|
||||
secret_key = creds.get("secret_access_key") or ""
|
||||
|
||||
if not bucket:
|
||||
raise ValueError("S3 integration is missing bucket in config")
|
||||
if not (access_key and secret_key):
|
||||
raise ValueError("S3 integration is missing access_key_id or secret_access_key in credentials")
|
||||
|
||||
client_kwargs: dict[str, Any] = {
|
||||
"region_name": region,
|
||||
"aws_access_key_id": access_key,
|
||||
"aws_secret_access_key": secret_key,
|
||||
}
|
||||
if endpoint_url:
|
||||
client_kwargs["endpoint_url"] = endpoint_url
|
||||
|
||||
s3 = boto3.client("s3", **client_kwargs)
|
||||
|
||||
filename = os.path.basename(file_path)
|
||||
s3_key = f"{folder_prefix.rstrip('/')}/{filename}" if folder_prefix else filename
|
||||
|
||||
try:
|
||||
s3.upload_file(file_path, bucket, s3_key, ExtraArgs={"StorageClass": storage_class})
|
||||
except ClientError as exc:
|
||||
raise RuntimeError(f"S3 upload failed: {exc}") from exc
|
||||
|
||||
logger.info("[%s] S3 upload complete: s3://%s/%s", task_id, bucket, s3_key)
|
||||
return {"status": "Completed", "s3_bucket": bucket, "s3_key": s3_key}
|
||||
|
||||
|
||||
def _upload_google_drive(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to Google Drive using per-user OAuth or service-account credentials."""
|
||||
from googleapiclient.discovery import build
|
||||
from googleapiclient.http import MediaFileUpload
|
||||
|
||||
folder_id = cfg.get("folder_id") or ""
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
# Prefer OAuth credentials (client_id + client_secret + refresh_token)
|
||||
client_id = creds.get("client_id") or ""
|
||||
client_secret = creds.get("client_secret") or ""
|
||||
refresh_token = creds.get("refresh_token") or ""
|
||||
credentials_json = creds.get("credentials_json") or ""
|
||||
|
||||
if client_id and client_secret and refresh_token:
|
||||
from google.auth.transport.requests import Request
|
||||
from google.oauth2.credentials import Credentials as OAuthCredentials
|
||||
|
||||
google_creds = OAuthCredentials(
|
||||
None,
|
||||
refresh_token=refresh_token,
|
||||
token_uri="https://oauth2.googleapis.com/token",
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
scopes=["https://www.googleapis.com/auth/drive.file"],
|
||||
)
|
||||
google_creds.refresh(Request())
|
||||
service = build("drive", "v3", credentials=google_creds)
|
||||
elif credentials_json:
|
||||
from google.oauth2.service_account import Credentials as SACredentials
|
||||
|
||||
creds_dict = json.loads(credentials_json)
|
||||
sa_creds = SACredentials.from_service_account_info(creds_dict, scopes=["https://www.googleapis.com/auth/drive"])
|
||||
service = build("drive", "v3", credentials=sa_creds)
|
||||
else:
|
||||
raise ValueError("Google Drive integration requires either OAuth credentials or credentials_json")
|
||||
|
||||
file_metadata: dict[str, Any] = {"name": filename}
|
||||
if folder_id:
|
||||
file_metadata["parents"] = [folder_id]
|
||||
|
||||
media = MediaFileUpload(file_path, mimetype="application/pdf", resumable=True)
|
||||
file_obj = service.files().create(body=file_metadata, media_body=media, fields="id,name,webViewLink").execute()
|
||||
|
||||
gdrive_id = file_obj.get("id")
|
||||
web_link = file_obj.get("webViewLink")
|
||||
logger.info("[%s] Google Drive upload complete: %s (%s)", task_id, gdrive_id, web_link)
|
||||
return {"status": "Completed", "google_drive_file_id": gdrive_id, "web_link": web_link}
|
||||
|
||||
|
||||
def _upload_onedrive(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to OneDrive using per-user MSAL credentials."""
|
||||
import urllib.parse
|
||||
|
||||
import msal
|
||||
import requests as _requests
|
||||
|
||||
client_id = creds.get("client_id") or ""
|
||||
client_secret = creds.get("client_secret") or ""
|
||||
refresh_token = creds.get("refresh_token") or ""
|
||||
tenant = cfg.get("tenant_id") or "common"
|
||||
folder_path = cfg.get("folder_path") or ""
|
||||
|
||||
if not (client_id and client_secret):
|
||||
raise ValueError("OneDrive integration is missing client_id or client_secret in credentials")
|
||||
|
||||
scopes = ["https://graph.microsoft.com/.default"]
|
||||
msal_app = msal.ConfidentialClientApplication(
|
||||
client_id=client_id,
|
||||
client_credential=client_secret,
|
||||
authority=f"https://login.microsoftonline.com/{tenant}",
|
||||
)
|
||||
|
||||
if refresh_token:
|
||||
token_resp = msal_app.acquire_token_by_refresh_token(refresh_token=refresh_token, scopes=scopes)
|
||||
else:
|
||||
token_resp = msal_app.acquire_token_for_client(scopes=scopes)
|
||||
|
||||
if "access_token" not in token_resp:
|
||||
raise ValueError(f"OneDrive token acquisition failed: {token_resp.get('error_description', 'unknown')}")
|
||||
|
||||
access_token = token_resp["access_token"]
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
# Build upload-session URL
|
||||
if folder_path:
|
||||
folder_path = folder_path.strip("/")
|
||||
encoded_path = "/".join(urllib.parse.quote(p) for p in folder_path.split("/"))
|
||||
encoded_file = urllib.parse.quote(filename)
|
||||
item_path = f"/root:/{encoded_path}/{encoded_file}:/createUploadSession"
|
||||
else:
|
||||
encoded_file = urllib.parse.quote(filename)
|
||||
item_path = f"/root:/{encoded_file}:/createUploadSession"
|
||||
|
||||
session_url = f"https://graph.microsoft.com/v1.0/me/drive{item_path}"
|
||||
headers = {"Authorization": f"Bearer {access_token}", "Content-Type": "application/json"}
|
||||
resp = _requests.post(
|
||||
session_url, headers=headers, json={"item": {"@microsoft.graph.conflictBehavior": "replace"}}, timeout=30
|
||||
)
|
||||
resp.raise_for_status()
|
||||
upload_url = resp.json()["uploadUrl"]
|
||||
|
||||
file_size = os.path.getsize(file_path)
|
||||
chunk_size = 10 * 1024 * 1024
|
||||
with open(file_path, "rb") as fh:
|
||||
chunk_num = 0
|
||||
while True:
|
||||
chunk = fh.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
start = chunk_num * chunk_size
|
||||
end = start + len(chunk) - 1
|
||||
upload_headers = {
|
||||
"Content-Length": str(len(chunk)),
|
||||
"Content-Range": f"bytes {start}-{end}/{file_size}",
|
||||
}
|
||||
upload_resp = _requests.put(upload_url, headers=upload_headers, data=chunk, timeout=120)
|
||||
if upload_resp.status_code not in (201, 202):
|
||||
raise RuntimeError(f"OneDrive chunk upload failed: {upload_resp.status_code}")
|
||||
chunk_num += 1
|
||||
|
||||
logger.info("[%s] OneDrive upload complete: %s/%s", task_id, folder_path, filename)
|
||||
return {"status": "Completed", "onedrive_folder": folder_path, "filename": filename}
|
||||
|
||||
|
||||
def _upload_webdav(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to a WebDAV server using per-user credentials."""
|
||||
import requests as _requests
|
||||
|
||||
url = cfg.get("url") or ""
|
||||
folder = cfg.get("folder") or ""
|
||||
username = creds.get("username") or ""
|
||||
password = creds.get("password") or ""
|
||||
verify_ssl = cfg.get("verify_ssl", True)
|
||||
|
||||
if not url:
|
||||
raise ValueError("WebDAV integration is missing url in config")
|
||||
|
||||
filename = os.path.basename(file_path)
|
||||
folder = folder.lstrip("/")
|
||||
target = urljoin(url.rstrip("/") + "/", folder)
|
||||
if not target.endswith("/"):
|
||||
target += "/"
|
||||
dest = urljoin(target, filename)
|
||||
|
||||
with open(file_path, "rb") as fh:
|
||||
resp = _requests.put(dest, auth=(username, password), data=fh, verify=verify_ssl, timeout=120)
|
||||
|
||||
if resp.status_code not in (200, 201, 204):
|
||||
raise RuntimeError(f"WebDAV upload failed: {resp.status_code} {resp.text[:200]}")
|
||||
|
||||
logger.info("[%s] WebDAV upload complete: %s", task_id, dest)
|
||||
return {"status": "Completed", "webdav_url": dest}
|
||||
|
||||
|
||||
def _upload_nextcloud(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to Nextcloud (WebDAV) using per-user credentials."""
|
||||
# Nextcloud uses WebDAV under the hood; reuse the WebDAV helper.
|
||||
return _upload_webdav(file_path, cfg, creds, task_id)
|
||||
|
||||
|
||||
def _upload_ftp(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to an FTP/FTPS server using per-user credentials."""
|
||||
host = cfg.get("host") or ""
|
||||
port = int(cfg.get("port") or 21)
|
||||
username = cfg.get("username") or ""
|
||||
folder = cfg.get("folder") or ""
|
||||
use_tls = cfg.get("use_tls", True)
|
||||
password = creds.get("password") or ""
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not host:
|
||||
raise ValueError("FTP integration is missing host in config")
|
||||
|
||||
ftp: ftplib.FTP
|
||||
if use_tls:
|
||||
ftp = ftplib.FTP_TLS() # nosec B321 # noqa: S321
|
||||
ftp.connect(host=host, port=port)
|
||||
ftp.login(user=username, passwd=password)
|
||||
ftp.prot_p()
|
||||
else:
|
||||
ftp = ftplib.FTP() # nosec B321 # noqa: S321
|
||||
ftp.connect(host=host, port=port)
|
||||
ftp.login(user=username, passwd=password)
|
||||
|
||||
if folder:
|
||||
folder_stripped = folder.lstrip("/")
|
||||
try:
|
||||
ftp.cwd(folder_stripped)
|
||||
except ftplib.error_perm:
|
||||
parts = folder_stripped.split("/")
|
||||
current = ""
|
||||
for part in parts:
|
||||
if not part:
|
||||
continue
|
||||
current += f"/{part}"
|
||||
try:
|
||||
ftp.cwd(current)
|
||||
except ftplib.error_perm:
|
||||
ftp.mkd(current)
|
||||
ftp.cwd(current)
|
||||
|
||||
with open(file_path, "rb") as fh:
|
||||
ftp.storbinary(f"STOR {filename}", fh)
|
||||
ftp.quit()
|
||||
|
||||
logger.info("[%s] FTP upload complete: %s/%s", task_id, host, filename)
|
||||
return {"status": "Completed", "ftp_host": host, "filename": filename}
|
||||
|
||||
|
||||
def _upload_sftp(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to an SFTP server using per-user credentials."""
|
||||
import paramiko
|
||||
|
||||
host = cfg.get("host") or ""
|
||||
port = int(cfg.get("port") or 22)
|
||||
username = cfg.get("username") or ""
|
||||
folder = cfg.get("folder") or ""
|
||||
password = creds.get("password") or ""
|
||||
private_key_text = creds.get("private_key") or ""
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not host:
|
||||
raise ValueError("SFTP integration is missing host in config")
|
||||
|
||||
ssh = paramiko.SSHClient()
|
||||
ssh.load_system_host_keys()
|
||||
ssh.set_missing_host_key_policy(paramiko.RejectPolicy())
|
||||
|
||||
connect_kwargs: dict[str, Any] = {"hostname": host, "port": port, "username": username}
|
||||
if private_key_text:
|
||||
import io
|
||||
|
||||
pkey = paramiko.RSAKey.from_private_key(io.StringIO(private_key_text))
|
||||
connect_kwargs["pkey"] = pkey
|
||||
elif password:
|
||||
connect_kwargs["password"] = password
|
||||
else:
|
||||
raise ValueError("SFTP integration requires password or private_key in credentials")
|
||||
|
||||
ssh.connect(**connect_kwargs)
|
||||
sftp = ssh.open_sftp()
|
||||
|
||||
remote_path = f"{folder.rstrip('/')}/{filename}" if folder else filename
|
||||
if folder and folder.startswith("/") and not remote_path.startswith("/"):
|
||||
remote_path = "/" + remote_path
|
||||
|
||||
sftp.put(file_path, remote_path)
|
||||
sftp.close()
|
||||
ssh.close()
|
||||
|
||||
logger.info("[%s] SFTP upload complete: %s:%s", task_id, host, remote_path)
|
||||
return {"status": "Completed", "sftp_host": host, "sftp_path": remote_path}
|
||||
|
||||
|
||||
def _upload_paperless(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to a Paperless-ngx instance using per-user API token."""
|
||||
import time
|
||||
|
||||
import requests as _requests
|
||||
|
||||
host = (cfg.get("host") or "").rstrip("/")
|
||||
api_token = creds.get("api_token") or ""
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not host:
|
||||
raise ValueError("Paperless integration is missing host in config")
|
||||
if not api_token:
|
||||
raise ValueError("Paperless integration is missing api_token in credentials")
|
||||
|
||||
headers = {"Authorization": f"Token {api_token}"}
|
||||
post_url = f"{host}/api/documents/post_document/"
|
||||
|
||||
with open(file_path, "rb") as fh:
|
||||
resp = _requests.post(
|
||||
post_url,
|
||||
headers=headers,
|
||||
files={"document": (filename, fh, "application/pdf")},
|
||||
data={"title": filename},
|
||||
timeout=120,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
raw_task_id = resp.text.strip().strip('"').strip("'")
|
||||
|
||||
# Poll for completion (up to 30 s)
|
||||
task_url = f"{host}/api/tasks/"
|
||||
doc_id = None
|
||||
for _ in range(10):
|
||||
time.sleep(3)
|
||||
try:
|
||||
poll_resp = _requests.get(task_url, headers=headers, params={"task_id": raw_task_id}, timeout=30)
|
||||
poll_resp.raise_for_status()
|
||||
tasks_data = poll_resp.json()
|
||||
if isinstance(tasks_data, dict) and "results" in tasks_data:
|
||||
tasks_data = tasks_data["results"]
|
||||
if tasks_data:
|
||||
info = tasks_data[0]
|
||||
status = info.get("status")
|
||||
if status == "SUCCESS":
|
||||
doc_id = info.get("related_document")
|
||||
break
|
||||
elif status == "FAILURE":
|
||||
raise RuntimeError(f"Paperless processing failed: {info.get('result')}")
|
||||
except RuntimeError:
|
||||
raise
|
||||
except Exception as poll_exc:
|
||||
logger.warning("[%s] Paperless poll error: %s", task_id, poll_exc)
|
||||
|
||||
logger.info("[%s] Paperless upload complete: doc_id=%s", task_id, doc_id)
|
||||
return {"status": "Completed", "paperless_host": host, "paperless_document_id": doc_id}
|
||||
|
||||
|
||||
def _upload_email(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Send *file_path* as an email attachment using per-user SMTP credentials."""
|
||||
import smtplib
|
||||
from email.mime.application import MIMEApplication
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
|
||||
host = cfg.get("host") or ""
|
||||
port = int(cfg.get("port") or 587)
|
||||
username = cfg.get("username") or ""
|
||||
recipient = cfg.get("recipient") or ""
|
||||
use_tls = cfg.get("use_tls", True)
|
||||
sender_name = cfg.get("sender_name") or "DocuElevate"
|
||||
password = creds.get("password") or ""
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not (host and recipient):
|
||||
raise ValueError("Email integration is missing host or recipient in config")
|
||||
|
||||
msg = MIMEMultipart()
|
||||
msg["From"] = f"{sender_name} <{username}>" if username else sender_name
|
||||
msg["To"] = recipient
|
||||
msg["Subject"] = f"Document: {filename}"
|
||||
msg.attach(MIMEText(f"Please find the attached document: {filename}", "plain"))
|
||||
|
||||
with open(file_path, "rb") as fh:
|
||||
part = MIMEApplication(fh.read(), Name=filename)
|
||||
part["Content-Disposition"] = f'attachment; filename="{filename}"'
|
||||
msg.attach(part)
|
||||
|
||||
if use_tls:
|
||||
import ssl
|
||||
|
||||
tls_context = ssl.create_default_context()
|
||||
with smtplib.SMTP(host, port, timeout=30) as smtp:
|
||||
smtp.starttls(context=tls_context)
|
||||
if username and password:
|
||||
smtp.login(username, password)
|
||||
smtp.sendmail(msg["From"], [recipient], msg.as_string())
|
||||
else:
|
||||
# Plaintext SMTP — only use when explicitly configured and TLS is unavailable.
|
||||
# Credentials and content will be transmitted without encryption.
|
||||
with smtplib.SMTP(host, port, timeout=30) as smtp: # nosec B608
|
||||
if username and password:
|
||||
smtp.login(username, password)
|
||||
smtp.sendmail(msg["From"], [recipient], msg.as_string())
|
||||
|
||||
logger.info("[%s] Email upload complete: sent to %s", task_id, recipient)
|
||||
return {"status": "Completed", "recipient": recipient}
|
||||
|
||||
|
||||
def _upload_rclone(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Copy *file_path* to an rclone remote using per-user rclone config."""
|
||||
import re
|
||||
import tempfile
|
||||
|
||||
remote = cfg.get("remote") or ""
|
||||
folder = cfg.get("folder") or ""
|
||||
rclone_conf_text = creds.get("rclone_conf") or ""
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
if not remote:
|
||||
raise ValueError("Rclone integration is missing remote in config")
|
||||
if not rclone_conf_text:
|
||||
raise ValueError("Rclone integration is missing rclone_conf in credentials")
|
||||
|
||||
# Validate remote and folder to prevent shell metacharacter injection.
|
||||
# rclone remote names are alphanumeric + hyphens/underscores followed by ':'.
|
||||
# folder paths must not contain shell-dangerous characters.
|
||||
_SAFE_REMOTE_RE = re.compile(r"^[A-Za-z0-9_\-]+:(/[A-Za-z0-9_.@\-/ ]*)?$")
|
||||
_SAFE_FOLDER_RE = re.compile(r"^[A-Za-z0-9_.@\-/ ]*$")
|
||||
if not _SAFE_REMOTE_RE.match(remote):
|
||||
raise ValueError(f"Rclone remote contains unsafe characters: {remote!r}")
|
||||
if folder and not _SAFE_FOLDER_RE.match(folder):
|
||||
raise ValueError(f"Rclone folder contains unsafe characters: {folder!r}")
|
||||
|
||||
# Write the user's rclone config to a temp file so we don't touch the system config
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".conf", delete=False) as tmp_conf:
|
||||
tmp_conf.write(rclone_conf_text)
|
||||
conf_path = tmp_conf.name
|
||||
|
||||
dest = f"{remote.rstrip('/')}/{folder.strip('/')}/{filename}" if folder else f"{remote.rstrip('/')}/{filename}"
|
||||
dest = dest.replace("//", "/")
|
||||
|
||||
try:
|
||||
result = subprocess.run( # nosec B603 # noqa: S603 S607
|
||||
["rclone", "copyto", f"--config={conf_path}", file_path, dest], # noqa: S603 S607
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=300,
|
||||
check=False,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(f"rclone exited {result.returncode}: {result.stderr[:300]}")
|
||||
finally:
|
||||
os.unlink(conf_path)
|
||||
|
||||
logger.info("[%s] Rclone upload complete: %s", task_id, dest)
|
||||
return {"status": "Completed", "rclone_dest": dest}
|
||||
|
||||
|
||||
def _upload_icloud(file_path: str, cfg: dict[str, Any], creds: dict[str, Any], task_id: str) -> dict[str, Any]:
|
||||
"""Upload *file_path* to iCloud Drive using per-user credentials.
|
||||
|
||||
Expected *cfg* keys:
|
||||
* ``folder`` – target folder path inside iCloud Drive (e.g. ``Documents/Uploads``).
|
||||
* ``cookie_directory`` – (optional) path for session cookie persistence.
|
||||
|
||||
Expected *creds* keys:
|
||||
* ``username`` – Apple ID email address.
|
||||
* ``password`` – app-specific password.
|
||||
"""
|
||||
from app.tasks.upload_to_icloud import _get_icloud_api, _navigate_to_folder
|
||||
|
||||
username = creds.get("username") or ""
|
||||
password = creds.get("password") or ""
|
||||
folder = cfg.get("folder") or ""
|
||||
cookie_directory = cfg.get("cookie_directory") or None
|
||||
|
||||
if not username or not password:
|
||||
raise ValueError("iCloud integration is missing username or password in credentials")
|
||||
|
||||
api = _get_icloud_api(username, password, cookie_directory)
|
||||
folder_node = _navigate_to_folder(api.drive, folder)
|
||||
|
||||
with open(file_path, "rb") as fh:
|
||||
folder_node.upload(fh)
|
||||
|
||||
logger.info("[%s] iCloud Drive upload complete: folder=%s", task_id, folder or "/")
|
||||
return {"status": "Completed", "icloud_folder": folder or "/"}
|
||||
|
||||
|
||||
# Map IntegrationType → upload helper
|
||||
_UPLOAD_HANDLERS = {
|
||||
IntegrationType.DROPBOX: _upload_dropbox,
|
||||
IntegrationType.S3: _upload_s3,
|
||||
IntegrationType.GOOGLE_DRIVE: _upload_google_drive,
|
||||
IntegrationType.ONEDRIVE: _upload_onedrive,
|
||||
IntegrationType.WEBDAV: _upload_webdav,
|
||||
IntegrationType.NEXTCLOUD: _upload_nextcloud,
|
||||
IntegrationType.FTP: _upload_ftp,
|
||||
IntegrationType.SFTP: _upload_sftp,
|
||||
IntegrationType.PAPERLESS: _upload_paperless,
|
||||
IntegrationType.EMAIL: _upload_email,
|
||||
IntegrationType.RCLONE: _upload_rclone,
|
||||
IntegrationType.ICLOUD: _upload_icloud,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Celery task
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_user_integration(self, file_path: str, integration_id: int, file_id: int | None = None) -> dict[str, Any]:
|
||||
"""Upload *file_path* to the destination described by the given UserIntegration record.
|
||||
|
||||
This task is dispatched once per active DESTINATION UserIntegration that
|
||||
belongs to a document's owner. Credentials are decrypted at runtime so
|
||||
they never travel across the Celery message bus in plaintext.
|
||||
|
||||
Args:
|
||||
file_path: Absolute path to the processed document file.
|
||||
integration_id: Primary key of the ``UserIntegration`` record.
|
||||
file_id: Optional ``FileRecord.id`` used for progress logging.
|
||||
|
||||
Returns:
|
||||
A dict with at least ``{"status": "Completed", ...}`` on success.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: When *file_path* does not exist.
|
||||
ValueError: When the integration record is not found or has missing config.
|
||||
RuntimeError: When the underlying upload operation fails.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
filename = os.path.basename(file_path)
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
f"upload_to_user_integration_{integration_id}",
|
||||
"in_progress",
|
||||
f"Uploading {filename} to integration {integration_id}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
if not os.path.exists(file_path):
|
||||
error_msg = f"File not found: {file_path}"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(
|
||||
task_id, f"upload_to_user_integration_{integration_id}", "failure", error_msg, file_id=file_id
|
||||
)
|
||||
raise FileNotFoundError(error_msg)
|
||||
|
||||
with SessionLocal() as db:
|
||||
integration: UserIntegration | None = (
|
||||
db.query(UserIntegration).filter(UserIntegration.id == integration_id).first()
|
||||
)
|
||||
if integration is None:
|
||||
error_msg = f"UserIntegration {integration_id} not found"
|
||||
logger.error("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(
|
||||
task_id, f"upload_to_user_integration_{integration_id}", "failure", error_msg, file_id=file_id
|
||||
)
|
||||
raise ValueError(error_msg)
|
||||
|
||||
itype = integration.integration_type
|
||||
int_name = integration.name
|
||||
owner_id = integration.owner_id
|
||||
|
||||
# Parse config (non-sensitive) and decrypt credentials (sensitive)
|
||||
try:
|
||||
cfg: dict[str, Any] = json.loads(integration.config) if integration.config else {}
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Integration {integration_id} has invalid JSON in config: {exc}") from exc
|
||||
|
||||
try:
|
||||
raw_creds = decrypt_value(integration.credentials) if integration.credentials else None
|
||||
creds: dict[str, Any] = json.loads(raw_creds) if raw_creds else {}
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError(f"Integration {integration_id} has invalid JSON in credentials: {exc}") from exc
|
||||
|
||||
handler = _UPLOAD_HANDLERS.get(itype)
|
||||
if handler is None:
|
||||
error_msg = f"No upload handler registered for integration type '{itype}' (integration {integration_id})"
|
||||
logger.warning("[%s] %s", task_id, error_msg)
|
||||
log_task_progress(
|
||||
task_id, f"upload_to_user_integration_{integration_id}", "skipped", error_msg, file_id=file_id
|
||||
)
|
||||
return {"status": "Skipped", "reason": error_msg}
|
||||
|
||||
logger.info(
|
||||
"[%s] Uploading %s via %s integration '%s' (id=%d, owner=%s)",
|
||||
task_id,
|
||||
filename,
|
||||
itype,
|
||||
int_name,
|
||||
integration_id,
|
||||
owner_id,
|
||||
)
|
||||
|
||||
try:
|
||||
result = handler(file_path, cfg, creds, task_id)
|
||||
|
||||
# Update last_used_at on success
|
||||
with SessionLocal() as db:
|
||||
integ = db.query(UserIntegration).filter(UserIntegration.id == integration_id).first()
|
||||
if integ:
|
||||
integ.last_used_at = datetime.now(timezone.utc)
|
||||
integ.last_error = None
|
||||
db.commit()
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
f"upload_to_user_integration_{integration_id}",
|
||||
"success",
|
||||
f"Uploaded to {itype} '{int_name}': {filename}",
|
||||
file_id=file_id,
|
||||
)
|
||||
return result
|
||||
|
||||
except Exception as exc:
|
||||
error_msg = str(exc)[:_MAX_ERROR_LENGTH]
|
||||
logger.error(
|
||||
"[%s] Upload to integration %d (%s '%s') failed: %s",
|
||||
task_id,
|
||||
integration_id,
|
||||
itype,
|
||||
int_name,
|
||||
error_msg,
|
||||
)
|
||||
|
||||
# Persist error for operator visibility
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
integ = db.query(UserIntegration).filter(UserIntegration.id == integration_id).first()
|
||||
if integ:
|
||||
integ.last_used_at = datetime.now(timezone.utc)
|
||||
integ.last_error = error_msg
|
||||
db.commit()
|
||||
except Exception as db_exc: # noqa: BLE001
|
||||
logger.warning("[%s] Could not persist last_error for integration %d: %s", task_id, integration_id, db_exc)
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
f"upload_to_user_integration_{integration_id}",
|
||||
"failure",
|
||||
f"Upload to {itype} '{int_name}' failed: {error_msg}",
|
||||
file_id=file_id,
|
||||
)
|
||||
raise
|
||||
@@ -8,14 +8,14 @@ import requests
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def upload_to_webdav(self, file_path: str, file_id: int = None):
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_to_webdav(self, file_path: str, file_id: int = None, folder_override: str = None):
|
||||
"""
|
||||
Uploads a file to a WebDAV server in the configured folder.
|
||||
|
||||
@@ -50,7 +50,7 @@ def upload_to_webdav(self, file_path: str, file_id: int = None):
|
||||
raise ValueError(error_msg)
|
||||
|
||||
# Construct the full upload URL
|
||||
webdav_folder = settings.webdav_folder or ""
|
||||
webdav_folder = folder_override if folder_override is not None else (settings.webdav_folder or "")
|
||||
# Ensure folder doesn't have leading slash if we're joining it to the base URL
|
||||
if webdav_folder and webdav_folder.startswith("/"):
|
||||
webdav_folder = webdav_folder[1:]
|
||||
|
||||
@@ -6,13 +6,13 @@ import subprocess
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.tasks.retry_config import UploadTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def upload_with_rclone(self, file_path: str, destination: str):
|
||||
"""
|
||||
Uploads a file using rclone to the specified destination.
|
||||
@@ -107,7 +107,7 @@ def upload_with_rclone(self, file_path: str, destination: str):
|
||||
raise RuntimeError(error_msg) from e
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
@celery.task(base=UploadTaskWithRetry, bind=True)
|
||||
def send_to_all_rclone_destinations(self, file_path: str):
|
||||
"""
|
||||
Uploads a file to all configured rclone destinations.
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,39 @@
|
||||
"""Celery task for asynchronous webhook delivery with retry and backoff.
|
||||
|
||||
Uses :class:`~app.tasks.retry_config.BaseTaskWithRetry` so failed deliveries
|
||||
are automatically retried with exponential backoff (default: 60 s, 300 s,
|
||||
900 s) and ±20 % jitter.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils.webhook import deliver_webhook
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True, name="webhook.deliver")
|
||||
def deliver_webhook_task(self, url: str, payload: dict[str, Any], secret: str | None = None) -> dict[str, Any]:
|
||||
"""Deliver a webhook payload to *url* with automatic retries.
|
||||
|
||||
Args:
|
||||
url: Target webhook URL.
|
||||
payload: The full webhook payload envelope.
|
||||
secret: Optional shared secret for HMAC-SHA256 signing.
|
||||
|
||||
Returns:
|
||||
A dict with ``status`` and ``url`` on success.
|
||||
|
||||
Raises:
|
||||
RuntimeError: Re-raised to trigger Celery retry on delivery failure.
|
||||
"""
|
||||
logger.info("Delivering webhook to %s (attempt %d/%d)", url, self.request.retries + 1, self.max_retries + 1)
|
||||
|
||||
success = deliver_webhook(url, payload, secret)
|
||||
if success:
|
||||
return {"status": "delivered", "url": url}
|
||||
|
||||
raise RuntimeError(f"Webhook delivery to {url} failed")
|
||||
@@ -0,0 +1,63 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>{{ filename }} – DocuElevate</title>
|
||||
<style>
|
||||
body { margin: 0; padding: 0; background-color: #f4f4f5; font-family: Arial, Helvetica, sans-serif; }
|
||||
.wrapper { max-width: 600px; margin: 32px auto; background: #ffffff; border-radius: 8px; box-shadow: 0 2px 8px rgba(0,0,0,.08); overflow: hidden; }
|
||||
.header { background: linear-gradient(135deg, #4f46e5 0%, #7c3aed 100%); padding: 32px 40px; text-align: center; }
|
||||
.header img { max-height: 48px; }
|
||||
.header h1 { color: #ffffff; font-size: 22px; margin: 16px 0 0; }
|
||||
.body { padding: 32px 40px; }
|
||||
.body p { color: #374151; font-size: 15px; line-height: 1.6; margin: 0 0 16px; }
|
||||
.attachment-box { background: #f9fafb; border: 1px solid #e5e7eb; border-radius: 6px; padding: 16px 20px; margin: 24px 0; }
|
||||
.attachment-box .label { font-size: 11px; font-weight: bold; color: #6b7280; text-transform: uppercase; letter-spacing: .05em; margin-bottom: 6px; }
|
||||
.attachment-box .filename { color: #111827; font-size: 15px; font-weight: bold; word-break: break-all; }
|
||||
.metadata-table { width: 100%; border-collapse: collapse; margin-top: 8px; font-size: 13px; }
|
||||
.metadata-table td { padding: 6px 0; color: #374151; vertical-align: top; }
|
||||
.metadata-table td:first-child { font-weight: bold; color: #6b7280; width: 40%; padding-right: 12px; }
|
||||
.footer { background: #f9fafb; border-top: 1px solid #e5e7eb; padding: 20px 40px; text-align: center; color: #9ca3af; font-size: 12px; }
|
||||
.footer a { color: #4f46e5; text-decoration: none; }
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="wrapper">
|
||||
<div class="header">
|
||||
{% if has_logo %}
|
||||
<img src="cid:logo" alt="{{ app_name }} logo">
|
||||
{% endif %}
|
||||
<h1>Document Delivery</h1>
|
||||
</div>
|
||||
<div class="body">
|
||||
<p>{{ message }}</p>
|
||||
|
||||
<div class="attachment-box">
|
||||
<div class="label">Attached file</div>
|
||||
<div class="filename">📎 {{ filename }}</div>
|
||||
</div>
|
||||
|
||||
{% if has_metadata and metadata %}
|
||||
<p style="font-weight:bold; color:#374151; margin-bottom:8px;">Document metadata</p>
|
||||
<table class="metadata-table">
|
||||
{% for key, value in metadata.items() %}
|
||||
<tr>
|
||||
<td>{{ key }}</td>
|
||||
<td>{{ value }}</td>
|
||||
</tr>
|
||||
{% endfor %}
|
||||
</table>
|
||||
{% endif %}
|
||||
|
||||
<p style="margin-top:24px; color:#6b7280; font-size:13px;">
|
||||
This document was sent automatically by {{ app_name }}.{% if app_url %} Visit <a href="{{ app_url }}" style="color:#4f46e5;">{{ app_url }}</a> to manage your documents.{% endif %}
|
||||
</p>
|
||||
</div>
|
||||
<div class="footer">
|
||||
© {{ current_year }} {{ app_name }} · Intelligent Document Processing
|
||||
{% if app_url %}· <a href="{{ app_url }}">{{ app_url }}</a>{% endif %}
|
||||
</div>
|
||||
</div>
|
||||
</body>
|
||||
</html>
|
||||
@@ -131,3 +131,157 @@ ALLOWED_EXTENSIONS: set[str] = {
|
||||
".md",
|
||||
".markdown",
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fine-grained file-type categories used by IMAP ingestion profiles.
|
||||
# Each category groups related MIME types and extensions so that users can
|
||||
# enable/disable a logical collection of formats (e.g. "images") rather than
|
||||
# having to manage individual MIME strings.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
FILE_TYPE_CATEGORIES: dict[str, dict] = {
|
||||
"pdf": {
|
||||
"label": "PDF",
|
||||
"description": "PDF documents (.pdf)",
|
||||
"mime_types": frozenset({"application/pdf"}),
|
||||
"extensions": frozenset({".pdf"}),
|
||||
},
|
||||
"office": {
|
||||
"label": "Microsoft Office",
|
||||
"description": "Word, Excel and PowerPoint files (.doc, .docx, .xls, .xlsx, .ppt, .pptx, …)",
|
||||
"mime_types": frozenset(
|
||||
{
|
||||
"application/msword",
|
||||
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||||
"application/vnd.openxmlformats-officedocument.wordprocessingml.template",
|
||||
"application/vnd.ms-word.document.macroEnabled.12",
|
||||
"application/vnd.ms-word.template.macroEnabled.12",
|
||||
"application/vnd.ms-excel",
|
||||
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
|
||||
"application/vnd.openxmlformats-officedocument.spreadsheetml.template",
|
||||
"application/vnd.ms-excel.sheet.macroEnabled.12",
|
||||
"application/vnd.ms-excel.sheet.binary.macroEnabled.12",
|
||||
"application/vnd.ms-powerpoint",
|
||||
"application/vnd.openxmlformats-officedocument.presentationml.presentation",
|
||||
"application/vnd.openxmlformats-officedocument.presentationml.template",
|
||||
"application/vnd.openxmlformats-officedocument.presentationml.slideshow",
|
||||
"application/vnd.ms-powerpoint.presentation.macroEnabled.12",
|
||||
}
|
||||
),
|
||||
"extensions": frozenset(
|
||||
{
|
||||
".doc",
|
||||
".docx",
|
||||
".docm",
|
||||
".dot",
|
||||
".dotx",
|
||||
".dotm",
|
||||
".xls",
|
||||
".xlsx",
|
||||
".xlsm",
|
||||
".xlsb",
|
||||
".xlt",
|
||||
".xltx",
|
||||
".xlw",
|
||||
".ppt",
|
||||
".pptx",
|
||||
".pptm",
|
||||
".pps",
|
||||
".ppsx",
|
||||
".pot",
|
||||
".potx",
|
||||
}
|
||||
),
|
||||
},
|
||||
"opendocument": {
|
||||
"label": "OpenDocument (LibreOffice)",
|
||||
"description": "LibreOffice / OpenOffice files (.odt, .ods, .odp, …)",
|
||||
"mime_types": frozenset(
|
||||
{
|
||||
"application/vnd.oasis.opendocument.text",
|
||||
"application/vnd.oasis.opendocument.spreadsheet",
|
||||
"application/vnd.oasis.opendocument.presentation",
|
||||
"application/vnd.oasis.opendocument.graphics",
|
||||
"application/vnd.oasis.opendocument.formula",
|
||||
}
|
||||
),
|
||||
"extensions": frozenset({".odt", ".ods", ".odp", ".odg", ".odf"}),
|
||||
},
|
||||
"text": {
|
||||
"label": "Text & Data",
|
||||
"description": "Plain text, CSV and RTF files (.txt, .csv, .rtf)",
|
||||
"mime_types": frozenset(
|
||||
{
|
||||
"text/plain",
|
||||
"text/csv",
|
||||
"application/rtf",
|
||||
"text/rtf",
|
||||
}
|
||||
),
|
||||
"extensions": frozenset({".txt", ".csv", ".rtf"}),
|
||||
},
|
||||
"web": {
|
||||
"label": "Web & Markup",
|
||||
"description": "HTML and Markdown files (.html, .htm, .md, .markdown)",
|
||||
"mime_types": frozenset(
|
||||
{
|
||||
"text/html",
|
||||
"text/markdown",
|
||||
"text/x-markdown",
|
||||
}
|
||||
),
|
||||
"extensions": frozenset({".html", ".htm", ".md", ".markdown"}),
|
||||
},
|
||||
"images": {
|
||||
"label": "Images",
|
||||
"description": "Image files (.jpg, .png, .gif, .bmp, .tiff, .webp, .svg)",
|
||||
"mime_types": frozenset(
|
||||
{
|
||||
"image/jpeg",
|
||||
"image/jpg",
|
||||
"image/png",
|
||||
"image/gif",
|
||||
"image/bmp",
|
||||
"image/tiff",
|
||||
"image/webp",
|
||||
"image/svg+xml",
|
||||
}
|
||||
),
|
||||
"extensions": frozenset(
|
||||
{
|
||||
".jpg",
|
||||
".jpeg",
|
||||
".png",
|
||||
".gif",
|
||||
".bmp",
|
||||
".tiff",
|
||||
".tif",
|
||||
".webp",
|
||||
".svg",
|
||||
}
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
# Default categories for the "documents only" built-in profile (no images)
|
||||
DEFAULT_CATEGORIES: list[str] = ["pdf", "office", "opendocument", "text", "web"]
|
||||
# All categories including images
|
||||
ALL_CATEGORIES: list[str] = ["pdf", "office", "opendocument", "text", "web", "images"]
|
||||
|
||||
|
||||
def get_allowed_types_for_categories(
|
||||
categories: list[str],
|
||||
) -> tuple[frozenset[str], frozenset[str]]:
|
||||
"""Return ``(mime_types, extensions)`` for the given category list.
|
||||
|
||||
Unknown category names are silently ignored so that future categories
|
||||
don't break existing profiles.
|
||||
"""
|
||||
mime_types: set[str] = set()
|
||||
extensions: set[str] = set()
|
||||
for cat in categories:
|
||||
info = FILE_TYPE_CATEGORIES.get(cat)
|
||||
if info:
|
||||
mime_types |= info["mime_types"]
|
||||
extensions |= info["extensions"]
|
||||
return frozenset(mime_types), frozenset(extensions)
|
||||
|
||||
@@ -0,0 +1,331 @@
|
||||
"""
|
||||
Comprehensive audit-event service for DocuElevate.
|
||||
|
||||
Provides helpers to **record** audit events (append-only database writes)
|
||||
and to optionally **forward** them to external SIEM systems.
|
||||
|
||||
Supported SIEM transports:
|
||||
* **Syslog** – RFC 5424 structured-data messages over UDP or TCP.
|
||||
* **HTTP** – JSON POST payloads compatible with Splunk HEC, Logstash
|
||||
HTTP input, Grafana Loki push API, and any generic webhook endpoint.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import socket
|
||||
import threading
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from fastapi import Request
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.middleware.audit_log import get_client_ip, get_username
|
||||
from app.models import AuditLog
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def record_event(
|
||||
db: Session,
|
||||
*,
|
||||
action: str,
|
||||
user: str = "system",
|
||||
resource_type: str | None = None,
|
||||
resource_id: str | None = None,
|
||||
ip_address: str | None = None,
|
||||
details: dict[str, Any] | None = None,
|
||||
severity: str = "info",
|
||||
) -> AuditLog:
|
||||
"""Persist an audit event and optionally forward it to SIEM.
|
||||
|
||||
Args:
|
||||
db: Active SQLAlchemy session.
|
||||
action: Short action identifier (e.g. ``"login"``, ``"document.create"``).
|
||||
user: Username performing the action.
|
||||
resource_type: Category of the affected resource (``"document"``, ``"user"`` …).
|
||||
resource_id: Identifier of the affected resource.
|
||||
ip_address: Client IP address (``None`` when not applicable).
|
||||
details: Arbitrary key/value context serialised as JSON.
|
||||
severity: One of ``info``, ``warning``, ``error``, ``critical``.
|
||||
|
||||
Returns:
|
||||
The newly created :class:`AuditLog` row.
|
||||
"""
|
||||
details_json = json.dumps(details, default=str) if details else None
|
||||
|
||||
entry = AuditLog(
|
||||
user=user,
|
||||
action=action,
|
||||
resource_type=resource_type,
|
||||
resource_id=str(resource_id) if resource_id is not None else None,
|
||||
ip_address=ip_address,
|
||||
details=details_json,
|
||||
severity=severity,
|
||||
)
|
||||
db.add(entry)
|
||||
db.commit()
|
||||
db.refresh(entry)
|
||||
|
||||
# Fire-and-forget SIEM forwarding in a background thread so we never
|
||||
# block the request path.
|
||||
if settings.audit_siem_enabled:
|
||||
payload = _build_siem_payload(entry)
|
||||
thread = threading.Thread(target=_forward_to_siem, args=(payload,), daemon=True)
|
||||
thread.start()
|
||||
|
||||
return entry
|
||||
|
||||
|
||||
def record_event_from_request(
|
||||
db: Session,
|
||||
request: Request,
|
||||
*,
|
||||
action: str,
|
||||
resource_type: str | None = None,
|
||||
resource_id: str | None = None,
|
||||
details: dict[str, Any] | None = None,
|
||||
severity: str = "info",
|
||||
) -> AuditLog:
|
||||
"""Convenience wrapper that extracts user and IP from a :class:`Request`.
|
||||
|
||||
Args:
|
||||
db: Active SQLAlchemy session.
|
||||
request: The current HTTP request.
|
||||
action: Short action identifier.
|
||||
resource_type: Category of the affected resource.
|
||||
resource_id: Identifier of the affected resource.
|
||||
details: Arbitrary key/value context serialised as JSON.
|
||||
severity: One of ``info``, ``warning``, ``error``, ``critical``.
|
||||
|
||||
Returns:
|
||||
The newly created :class:`AuditLog` row.
|
||||
"""
|
||||
return record_event(
|
||||
db,
|
||||
action=action,
|
||||
user=get_username(request),
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
ip_address=get_client_ip(request),
|
||||
details=details,
|
||||
severity=severity,
|
||||
)
|
||||
|
||||
|
||||
def query_events(
|
||||
db: Session,
|
||||
*,
|
||||
action: str | None = None,
|
||||
user: str | None = None,
|
||||
resource_type: str | None = None,
|
||||
severity: str | None = None,
|
||||
since: datetime | None = None,
|
||||
until: datetime | None = None,
|
||||
limit: int = 200,
|
||||
offset: int = 0,
|
||||
) -> list[AuditLog]:
|
||||
"""Query audit log entries with optional filtering.
|
||||
|
||||
Args:
|
||||
db: Active SQLAlchemy session.
|
||||
action: Filter by action string (exact match).
|
||||
user: Filter by username (exact match).
|
||||
resource_type: Filter by resource type (exact match).
|
||||
severity: Filter by severity level (exact match).
|
||||
since: Only events at or after this timestamp.
|
||||
until: Only events at or before this timestamp.
|
||||
limit: Maximum number of rows to return.
|
||||
offset: Number of rows to skip (for pagination).
|
||||
|
||||
Returns:
|
||||
List of :class:`AuditLog` rows ordered by *timestamp descending*.
|
||||
"""
|
||||
q = db.query(AuditLog)
|
||||
if action:
|
||||
q = q.filter(AuditLog.action == action)
|
||||
if user:
|
||||
q = q.filter(AuditLog.user == user)
|
||||
if resource_type:
|
||||
q = q.filter(AuditLog.resource_type == resource_type)
|
||||
if severity:
|
||||
q = q.filter(AuditLog.severity == severity)
|
||||
if since:
|
||||
q = q.filter(AuditLog.timestamp >= since)
|
||||
if until:
|
||||
q = q.filter(AuditLog.timestamp <= until)
|
||||
return q.order_by(AuditLog.timestamp.desc()).offset(offset).limit(limit).all()
|
||||
|
||||
|
||||
def count_events(
|
||||
db: Session,
|
||||
*,
|
||||
action: str | None = None,
|
||||
user: str | None = None,
|
||||
resource_type: str | None = None,
|
||||
severity: str | None = None,
|
||||
since: datetime | None = None,
|
||||
until: datetime | None = None,
|
||||
) -> int:
|
||||
"""Return the total count of events matching the given filters.
|
||||
|
||||
Args:
|
||||
db: Active SQLAlchemy session.
|
||||
action: Filter by action string.
|
||||
user: Filter by username.
|
||||
resource_type: Filter by resource type.
|
||||
severity: Filter by severity level.
|
||||
since: Only events at or after this timestamp.
|
||||
until: Only events at or before this timestamp.
|
||||
|
||||
Returns:
|
||||
Integer count.
|
||||
"""
|
||||
q = db.query(AuditLog)
|
||||
if action:
|
||||
q = q.filter(AuditLog.action == action)
|
||||
if user:
|
||||
q = q.filter(AuditLog.user == user)
|
||||
if resource_type:
|
||||
q = q.filter(AuditLog.resource_type == resource_type)
|
||||
if severity:
|
||||
q = q.filter(AuditLog.severity == severity)
|
||||
if since:
|
||||
q = q.filter(AuditLog.timestamp >= since)
|
||||
if until:
|
||||
q = q.filter(AuditLog.timestamp <= until)
|
||||
return q.count()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SIEM forwarding internals
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_SYSLOG_FACILITY_LOCAL0 = 16
|
||||
_SYSLOG_SEVERITY_MAP = {
|
||||
"info": 6,
|
||||
"warning": 4,
|
||||
"error": 3,
|
||||
"critical": 2,
|
||||
}
|
||||
|
||||
|
||||
def _build_siem_payload(entry: AuditLog) -> dict[str, Any]:
|
||||
"""Convert an :class:`AuditLog` row into a plain dict for SIEM delivery."""
|
||||
ts = entry.timestamp if entry.timestamp else datetime.now(timezone.utc)
|
||||
return {
|
||||
"id": entry.id,
|
||||
"timestamp": ts.isoformat(),
|
||||
"user": entry.user,
|
||||
"action": entry.action,
|
||||
"resource_type": entry.resource_type,
|
||||
"resource_id": entry.resource_id,
|
||||
"ip_address": entry.ip_address,
|
||||
"details": entry.details,
|
||||
"severity": entry.severity,
|
||||
"source": "docuelevate",
|
||||
}
|
||||
|
||||
|
||||
def _forward_to_siem(payload: dict[str, Any]) -> None:
|
||||
"""Route a SIEM payload to the configured transport."""
|
||||
transport = settings.audit_siem_transport.lower()
|
||||
try:
|
||||
if transport == "syslog":
|
||||
_send_syslog(payload)
|
||||
elif transport == "http":
|
||||
_send_http(payload)
|
||||
else:
|
||||
logger.warning("Unknown SIEM transport %r; skipping forwarding", transport)
|
||||
except Exception:
|
||||
logger.exception("Failed to forward audit event to SIEM (%s)", transport)
|
||||
|
||||
|
||||
def _send_syslog(payload: dict[str, Any]) -> None:
|
||||
"""Send a RFC 5424 syslog message to the configured receiver."""
|
||||
severity_num = _SYSLOG_SEVERITY_MAP.get(payload.get("severity", "info"), 6)
|
||||
priority = _SYSLOG_FACILITY_LOCAL0 * 8 + severity_num
|
||||
ts = payload.get("timestamp", datetime.now(timezone.utc).isoformat())
|
||||
hostname = socket.gethostname()
|
||||
app_name = "docuelevate"
|
||||
msg_id = payload.get("action", "-")
|
||||
|
||||
# Structured data (SD) element with key event fields.
|
||||
sd = (
|
||||
f'[docuelevate@0 user="{payload.get("user", "-")}" '
|
||||
f'action="{payload.get("action", "-")}" '
|
||||
f'resource_type="{payload.get("resource_type", "-")}" '
|
||||
f'resource_id="{payload.get("resource_id", "-")}" '
|
||||
f'ip="{payload.get("ip_address", "-")}"]'
|
||||
)
|
||||
message = json.dumps(payload, default=str)
|
||||
syslog_msg = f"<{priority}>1 {ts} {hostname} {app_name} - {msg_id} {sd} {message}"
|
||||
|
||||
proto = settings.audit_siem_syslog_protocol.lower()
|
||||
host = settings.audit_siem_syslog_host
|
||||
port = settings.audit_siem_syslog_port
|
||||
|
||||
if proto == "tcp":
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
||||
sock.settimeout(5)
|
||||
sock.connect((host, port))
|
||||
sock.sendall(syslog_msg.encode("utf-8"))
|
||||
else:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_DGRAM) as sock:
|
||||
sock.settimeout(5)
|
||||
sock.sendto(syslog_msg.encode("utf-8"), (host, port))
|
||||
|
||||
logger.debug("Syslog audit event sent to %s:%s (%s)", host, port, proto)
|
||||
|
||||
|
||||
def _send_http(payload: dict[str, Any]) -> None:
|
||||
"""POST a JSON audit event to the configured HTTP endpoint."""
|
||||
url = settings.audit_siem_http_url
|
||||
if not url:
|
||||
logger.warning("SIEM HTTP URL not configured; skipping HTTP forwarding")
|
||||
return
|
||||
|
||||
headers: dict[str, str] = {"Content-Type": "application/json"}
|
||||
token = settings.audit_siem_http_token
|
||||
if token:
|
||||
headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
# Parse custom headers (comma-separated "Key:Value" pairs).
|
||||
# Reject headers that could override security-critical ones already set,
|
||||
# and validate that header names contain only RFC 7230 token characters.
|
||||
_PROTECTED_HEADERS = {"authorization", "content-type", "host"}
|
||||
_VALID_HEADER_NAME = re.compile(r"^[A-Za-z0-9!#$%&'*+\-.^_`|~]+$")
|
||||
raw_custom = settings.audit_siem_http_custom_headers
|
||||
if raw_custom:
|
||||
for raw_pair in raw_custom.split(","):
|
||||
pair = raw_pair.strip()
|
||||
if ":" in pair:
|
||||
k, _, v = pair.partition(":")
|
||||
name = k.strip()
|
||||
if not name or not _VALID_HEADER_NAME.match(name):
|
||||
logger.warning("Skipping invalid SIEM custom header name: %r", name)
|
||||
continue
|
||||
if name.lower() in _PROTECTED_HEADERS:
|
||||
logger.warning("Skipping protected SIEM custom header: %r", name)
|
||||
continue
|
||||
headers[name] = v.strip()
|
||||
|
||||
# Wrap in Splunk HEC-style envelope when URL contains ``/services/collector``.
|
||||
body: dict[str, Any]
|
||||
if "/services/collector" in url:
|
||||
body = {"event": payload, "sourcetype": "docuelevate:audit", "source": "docuelevate"}
|
||||
else:
|
||||
body = payload
|
||||
|
||||
with httpx.Client(timeout=10) as client:
|
||||
resp = client.post(url, json=body, headers=headers)
|
||||
resp.raise_for_status()
|
||||
|
||||
logger.debug("HTTP audit event forwarded to %s (status %s)", url, resp.status_code)
|
||||
@@ -0,0 +1,433 @@
|
||||
"""Compliance service for managing GDPR, HIPAA, and SOC2 compliance templates.
|
||||
|
||||
Provides pre-built compliance configurations that can be applied with one click
|
||||
to ensure the DocuElevate instance meets regulatory requirements.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.models import ComplianceTemplate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pre-built compliance template definitions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
COMPLIANCE_TEMPLATES: dict[str, dict[str, Any]] = {
|
||||
"gdpr": {
|
||||
"display_name": "GDPR (General Data Protection Regulation)",
|
||||
"description": (
|
||||
"European Union regulation for data protection and privacy. "
|
||||
"Enforces data minimisation, encryption at rest, audit logging, "
|
||||
"and limits PII exposure in telemetry."
|
||||
),
|
||||
"settings": {
|
||||
"auth_enabled": "True",
|
||||
"sentry_send_default_pii": "False",
|
||||
"security_headers_enabled": "True",
|
||||
"security_header_hsts_enabled": "True",
|
||||
"security_header_csp_enabled": "True",
|
||||
"security_header_x_frame_options_enabled": "True",
|
||||
"enable_deduplication": "True",
|
||||
},
|
||||
"checks": [
|
||||
{
|
||||
"key": "auth_enabled",
|
||||
"expected": "True",
|
||||
"label": "Authentication enabled",
|
||||
"description": "User authentication must be enabled to control access to personal data.",
|
||||
},
|
||||
{
|
||||
"key": "sentry_send_default_pii",
|
||||
"expected": "False",
|
||||
"label": "PII excluded from telemetry",
|
||||
"description": "Personally identifiable information must not be sent to external monitoring services.",
|
||||
},
|
||||
{
|
||||
"key": "security_headers_enabled",
|
||||
"expected": "True",
|
||||
"label": "Security headers enabled",
|
||||
"description": "HTTP security headers protect against common web vulnerabilities.",
|
||||
},
|
||||
{
|
||||
"key": "security_header_hsts_enabled",
|
||||
"expected": "True",
|
||||
"label": "HSTS enabled",
|
||||
"description": "HTTP Strict Transport Security ensures encrypted connections.",
|
||||
},
|
||||
{
|
||||
"key": "security_header_csp_enabled",
|
||||
"expected": "True",
|
||||
"label": "Content Security Policy enabled",
|
||||
"description": "CSP headers prevent cross-site scripting and data injection attacks.",
|
||||
},
|
||||
{
|
||||
"key": "security_header_x_frame_options_enabled",
|
||||
"expected": "True",
|
||||
"label": "Clickjacking protection enabled",
|
||||
"description": "X-Frame-Options header prevents clickjacking attacks.",
|
||||
},
|
||||
{
|
||||
"key": "enable_deduplication",
|
||||
"expected": "True",
|
||||
"label": "Deduplication enabled",
|
||||
"description": "Data minimisation: avoid storing duplicate documents.",
|
||||
},
|
||||
],
|
||||
},
|
||||
"hipaa": {
|
||||
"display_name": "HIPAA (Health Insurance Portability and Accountability Act)",
|
||||
"description": (
|
||||
"United States regulation for protecting health information. "
|
||||
"Requires strong access controls, audit trails, encryption, "
|
||||
"and strict session management."
|
||||
),
|
||||
"settings": {
|
||||
"auth_enabled": "True",
|
||||
"multi_user_enabled": "True",
|
||||
"sentry_send_default_pii": "False",
|
||||
"security_headers_enabled": "True",
|
||||
"security_header_hsts_enabled": "True",
|
||||
"security_header_csp_enabled": "True",
|
||||
"security_header_x_frame_options_enabled": "True",
|
||||
"enable_deduplication": "True",
|
||||
},
|
||||
"checks": [
|
||||
{
|
||||
"key": "auth_enabled",
|
||||
"expected": "True",
|
||||
"label": "Authentication enabled",
|
||||
"description": "Access controls are required to protect electronic Protected Health Information (ePHI).",
|
||||
},
|
||||
{
|
||||
"key": "multi_user_enabled",
|
||||
"expected": "True",
|
||||
"label": "Multi-user mode enabled",
|
||||
"description": "Individual user accounts required for access accountability.",
|
||||
},
|
||||
{
|
||||
"key": "sentry_send_default_pii",
|
||||
"expected": "False",
|
||||
"label": "PII excluded from telemetry",
|
||||
"description": "Protected Health Information must not be sent to external services.",
|
||||
},
|
||||
{
|
||||
"key": "security_headers_enabled",
|
||||
"expected": "True",
|
||||
"label": "Security headers enabled",
|
||||
"description": "Security headers protect ePHI during transmission.",
|
||||
},
|
||||
{
|
||||
"key": "security_header_hsts_enabled",
|
||||
"expected": "True",
|
||||
"label": "HSTS enabled",
|
||||
"description": "Encrypted transport required for all ePHI transmissions.",
|
||||
},
|
||||
{
|
||||
"key": "security_header_csp_enabled",
|
||||
"expected": "True",
|
||||
"label": "Content Security Policy enabled",
|
||||
"description": "CSP prevents injection attacks that could expose ePHI.",
|
||||
},
|
||||
{
|
||||
"key": "security_header_x_frame_options_enabled",
|
||||
"expected": "True",
|
||||
"label": "Clickjacking protection enabled",
|
||||
"description": "Prevents embedding the application in unauthorized frames.",
|
||||
},
|
||||
{
|
||||
"key": "enable_deduplication",
|
||||
"expected": "True",
|
||||
"label": "Deduplication enabled",
|
||||
"description": "Minimise data footprint for ePHI.",
|
||||
},
|
||||
],
|
||||
},
|
||||
"soc2": {
|
||||
"display_name": "SOC 2 (Service Organization Control 2)",
|
||||
"description": (
|
||||
"Trust Service Criteria framework for service organisations. "
|
||||
"Focuses on security, availability, processing integrity, "
|
||||
"confidentiality, and privacy."
|
||||
),
|
||||
"settings": {
|
||||
"auth_enabled": "True",
|
||||
"multi_user_enabled": "True",
|
||||
"sentry_send_default_pii": "False",
|
||||
"security_headers_enabled": "True",
|
||||
"security_header_hsts_enabled": "True",
|
||||
"security_header_csp_enabled": "True",
|
||||
"security_header_x_frame_options_enabled": "True",
|
||||
"enable_deduplication": "True",
|
||||
},
|
||||
"checks": [
|
||||
{
|
||||
"key": "auth_enabled",
|
||||
"expected": "True",
|
||||
"label": "Authentication enabled",
|
||||
"description": "Logical access controls required (CC6.1).",
|
||||
},
|
||||
{
|
||||
"key": "multi_user_enabled",
|
||||
"expected": "True",
|
||||
"label": "Multi-user mode enabled",
|
||||
"description": "Individual user accounts for access management (CC6.2).",
|
||||
},
|
||||
{
|
||||
"key": "sentry_send_default_pii",
|
||||
"expected": "False",
|
||||
"label": "PII excluded from telemetry",
|
||||
"description": "Confidential information must not leak to external services (CC6.7).",
|
||||
},
|
||||
{
|
||||
"key": "security_headers_enabled",
|
||||
"expected": "True",
|
||||
"label": "Security headers enabled",
|
||||
"description": "Protection against common web threats (CC6.6).",
|
||||
},
|
||||
{
|
||||
"key": "security_header_hsts_enabled",
|
||||
"expected": "True",
|
||||
"label": "HSTS enabled",
|
||||
"description": "Encrypted transport in transit (CC6.7).",
|
||||
},
|
||||
{
|
||||
"key": "security_header_csp_enabled",
|
||||
"expected": "True",
|
||||
"label": "Content Security Policy enabled",
|
||||
"description": "Application-level security controls (CC6.6).",
|
||||
},
|
||||
{
|
||||
"key": "security_header_x_frame_options_enabled",
|
||||
"expected": "True",
|
||||
"label": "Clickjacking protection enabled",
|
||||
"description": "UI redress attack prevention (CC6.6).",
|
||||
},
|
||||
{
|
||||
"key": "enable_deduplication",
|
||||
"expected": "True",
|
||||
"label": "Deduplication enabled",
|
||||
"description": "Data integrity through deduplication (PI1.1).",
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def seed_compliance_templates(db: Session) -> None:
|
||||
"""Create or update the built-in compliance template rows.
|
||||
|
||||
Called once at application startup to ensure the ``compliance_templates``
|
||||
table always contains the latest definitions.
|
||||
"""
|
||||
for name, defn in COMPLIANCE_TEMPLATES.items():
|
||||
existing = db.query(ComplianceTemplate).filter(ComplianceTemplate.name == name).first()
|
||||
if existing is None:
|
||||
template = ComplianceTemplate(
|
||||
name=name,
|
||||
display_name=defn["display_name"],
|
||||
description=defn["description"],
|
||||
settings_json=json.dumps(defn["settings"]),
|
||||
enabled=False,
|
||||
status="not_applied",
|
||||
)
|
||||
db.add(template)
|
||||
logger.info(f"Seeded compliance template: {name}")
|
||||
else:
|
||||
# Update display_name and description if changed, but preserve user state
|
||||
existing.display_name = defn["display_name"]
|
||||
existing.description = defn["description"]
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to seed compliance templates")
|
||||
|
||||
|
||||
def get_all_templates(db: Session) -> list[dict[str, Any]]:
|
||||
"""Return all compliance templates with their current status."""
|
||||
templates = db.query(ComplianceTemplate).order_by(ComplianceTemplate.name).all()
|
||||
result = []
|
||||
for t in templates:
|
||||
defn = COMPLIANCE_TEMPLATES.get(t.name, {})
|
||||
checks = defn.get("checks", [])
|
||||
result.append(
|
||||
{
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"display_name": t.display_name,
|
||||
"description": t.description,
|
||||
"enabled": t.enabled,
|
||||
"status": t.status,
|
||||
"applied_at": t.applied_at.isoformat() if t.applied_at else None,
|
||||
"applied_by": t.applied_by,
|
||||
"settings": json.loads(t.settings_json) if t.settings_json else {},
|
||||
"checks": checks,
|
||||
"check_count": len(checks),
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def get_template_by_name(db: Session, name: str) -> ComplianceTemplate | None:
|
||||
"""Retrieve a single compliance template by name."""
|
||||
return db.query(ComplianceTemplate).filter(ComplianceTemplate.name == name).first()
|
||||
|
||||
|
||||
def evaluate_template_status(db: Session, name: str) -> dict[str, Any]:
|
||||
"""Evaluate the compliance status of a template against live settings.
|
||||
|
||||
Returns a dict with ``status``, ``total``, ``passed``, ``failed``, and
|
||||
a list of individual ``check_results``.
|
||||
"""
|
||||
from app.config import settings as app_settings
|
||||
from app.utils.settings_service import get_all_settings_from_db
|
||||
|
||||
defn = COMPLIANCE_TEMPLATES.get(name)
|
||||
if defn is None:
|
||||
return {"status": "unknown", "total": 0, "passed": 0, "failed": 0, "check_results": []}
|
||||
|
||||
db_settings = get_all_settings_from_db(db)
|
||||
checks = defn.get("checks", [])
|
||||
results: list[dict[str, Any]] = []
|
||||
passed = 0
|
||||
|
||||
for check in checks:
|
||||
key = check["key"]
|
||||
expected = check["expected"]
|
||||
|
||||
# Resolve effective value: DB > config object
|
||||
if key in db_settings and db_settings[key] is not None:
|
||||
actual = str(db_settings[key])
|
||||
else:
|
||||
actual = str(getattr(app_settings, key, ""))
|
||||
|
||||
is_passing = actual.lower() == expected.lower()
|
||||
if is_passing:
|
||||
passed += 1
|
||||
|
||||
results.append(
|
||||
{
|
||||
"key": key,
|
||||
"label": check["label"],
|
||||
"description": check["description"],
|
||||
"expected": expected,
|
||||
"actual": actual,
|
||||
"passing": is_passing,
|
||||
}
|
||||
)
|
||||
|
||||
total = len(checks)
|
||||
if passed == total:
|
||||
status = "compliant"
|
||||
elif passed > 0:
|
||||
status = "partial"
|
||||
else:
|
||||
status = "non_compliant"
|
||||
|
||||
return {
|
||||
"status": status,
|
||||
"total": total,
|
||||
"passed": passed,
|
||||
"failed": total - passed,
|
||||
"check_results": results,
|
||||
}
|
||||
|
||||
|
||||
def apply_template(db: Session, name: str, applied_by: str = "admin") -> dict[str, Any]:
|
||||
"""Apply a compliance template by writing its settings to the database.
|
||||
|
||||
Returns a summary of what was applied.
|
||||
"""
|
||||
from app.utils.settings_service import save_setting_to_db
|
||||
|
||||
defn = COMPLIANCE_TEMPLATES.get(name)
|
||||
if defn is None:
|
||||
return {"success": False, "error": f"Unknown template: {name}"}
|
||||
|
||||
template = get_template_by_name(db, name)
|
||||
if template is None:
|
||||
return {"success": False, "error": f"Template not found in database: {name}"}
|
||||
|
||||
applied_settings: dict[str, str] = {}
|
||||
errors: list[str] = []
|
||||
|
||||
for key, value in defn["settings"].items():
|
||||
try:
|
||||
save_setting_to_db(db, key, value, changed_by=f"compliance:{name}")
|
||||
applied_settings[key] = value
|
||||
except Exception as e:
|
||||
errors.append(f"{key}: {e}")
|
||||
logger.error(f"Failed to apply compliance setting {key}={value}: {e}")
|
||||
|
||||
# Update the template record
|
||||
now = datetime.now(timezone.utc)
|
||||
template.enabled = True
|
||||
template.settings_json = json.dumps(applied_settings)
|
||||
template.applied_at = now
|
||||
template.applied_by = applied_by
|
||||
|
||||
# Evaluate and store status
|
||||
eval_result = evaluate_template_status(db, name)
|
||||
template.status = eval_result["status"]
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(f"Failed to update compliance template record: {name}")
|
||||
return {"success": False, "error": "Database commit failed"}
|
||||
|
||||
logger.info(f"Applied compliance template '{name}' by {applied_by}: {len(applied_settings)} settings written")
|
||||
|
||||
return {
|
||||
"success": len(errors) == 0,
|
||||
"template": name,
|
||||
"applied_settings": applied_settings,
|
||||
"errors": errors,
|
||||
"status": eval_result,
|
||||
}
|
||||
|
||||
|
||||
def get_compliance_summary(db: Session) -> dict[str, Any]:
|
||||
"""Return a high-level compliance dashboard summary across all templates."""
|
||||
templates = db.query(ComplianceTemplate).order_by(ComplianceTemplate.name).all()
|
||||
summary: list[dict[str, Any]] = []
|
||||
total_checks = 0
|
||||
total_passed = 0
|
||||
|
||||
for t in templates:
|
||||
eval_result = evaluate_template_status(db, t.name)
|
||||
total_checks += eval_result["total"]
|
||||
total_passed += eval_result["passed"]
|
||||
summary.append(
|
||||
{
|
||||
"name": t.name,
|
||||
"display_name": t.display_name,
|
||||
"enabled": t.enabled,
|
||||
"status": eval_result["status"],
|
||||
"total": eval_result["total"],
|
||||
"passed": eval_result["passed"],
|
||||
"failed": eval_result["failed"],
|
||||
"applied_at": t.applied_at.isoformat() if t.applied_at else None,
|
||||
"applied_by": t.applied_by,
|
||||
}
|
||||
)
|
||||
|
||||
overall = "compliant" if total_checks > 0 and total_passed == total_checks else "non_compliant"
|
||||
if 0 < total_passed < total_checks:
|
||||
overall = "partial"
|
||||
|
||||
return {
|
||||
"overall_status": overall,
|
||||
"total_checks": total_checks,
|
||||
"total_passed": total_passed,
|
||||
"total_failed": total_checks - total_passed,
|
||||
"templates": summary,
|
||||
}
|
||||
@@ -144,7 +144,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "dropbox_app_secret", None)
|
||||
and getattr(settings, "dropbox_refresh_token", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "dropbox_enabled", True),
|
||||
"description": "Upload files to Dropbox cloud storage",
|
||||
"details": {
|
||||
"folder": getattr(settings, "dropbox_folder", "Not set"),
|
||||
@@ -154,23 +154,23 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
},
|
||||
}
|
||||
|
||||
# Add Email configuration
|
||||
# Add Email destination configuration (dedicated settings for document delivery)
|
||||
providers["Email"] = {
|
||||
"name": "Email",
|
||||
"icon": "fa-solid fa-envelope",
|
||||
"configured": bool(
|
||||
getattr(settings, "email_host", None) and getattr(settings, "email_default_recipient", None)
|
||||
getattr(settings, "dest_email_host", None) and getattr(settings, "dest_email_default_recipient", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "dest_email_enabled", True),
|
||||
"description": "Send documents via email",
|
||||
"details": {
|
||||
"host": getattr(settings, "email_host", "Not set"),
|
||||
"port": getattr(settings, "email_port", "Not set"),
|
||||
"username": getattr(settings, "email_username", "Not set"),
|
||||
"password": mask_sensitive_value(getattr(settings, "email_password", None)),
|
||||
"use_tls": getattr(settings, "email_use_tls", "Not set"),
|
||||
"sender": getattr(settings, "email_sender", "Not set"),
|
||||
"default_recipient": getattr(settings, "email_default_recipient", "Not set"),
|
||||
"host": getattr(settings, "dest_email_host", "Not set"),
|
||||
"port": getattr(settings, "dest_email_port", "Not set"),
|
||||
"username": getattr(settings, "dest_email_username", "Not set"),
|
||||
"password": mask_sensitive_value(getattr(settings, "dest_email_password", None)),
|
||||
"use_tls": getattr(settings, "dest_email_use_tls", "Not set"),
|
||||
"sender": getattr(settings, "dest_email_sender", "Not set"),
|
||||
"default_recipient": getattr(settings, "dest_email_default_recipient", "Not set"),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -183,7 +183,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "ftp_username", None)
|
||||
and getattr(settings, "ftp_password", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "ftp_enabled", True),
|
||||
"description": "Upload files to FTP server",
|
||||
"details": {
|
||||
"host": getattr(settings, "ftp_host", "Not set"),
|
||||
@@ -214,7 +214,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
"name": "Google Drive",
|
||||
"icon": "fa-brands fa-google-drive",
|
||||
"configured": is_configured and bool(getattr(settings, "google_drive_folder_id", None)),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "google_drive_enabled", True),
|
||||
"description": "Store documents in Google Drive",
|
||||
"details": {
|
||||
"auth_type": "OAuth" if use_oauth else "Service Account",
|
||||
@@ -250,7 +250,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "nextcloud_username", None)
|
||||
and getattr(settings, "nextcloud_password", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "nextcloud_enabled", True),
|
||||
"description": "Store documents in NextCloud",
|
||||
"details": {
|
||||
"url": getattr(settings, "nextcloud_upload_url", "Not set"),
|
||||
@@ -270,7 +270,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "onedrive_client_secret", None)
|
||||
and getattr(settings, "onedrive_refresh_token", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "onedrive_enabled", True),
|
||||
"description": "Store documents in Microsoft OneDrive",
|
||||
"details": {
|
||||
"client_id": getattr(settings, "onedrive_client_id", "Not set"),
|
||||
@@ -288,7 +288,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
"configured": bool(
|
||||
getattr(settings, "paperless_host", None) and getattr(settings, "paperless_ngx_api_token", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "paperless_enabled", True),
|
||||
"description": "Document management system for digital archives",
|
||||
"details": {
|
||||
"host": getattr(settings, "paperless_host", "Not set"),
|
||||
@@ -305,7 +305,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "aws_access_key_id", None)
|
||||
and getattr(settings, "aws_secret_access_key", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "s3_enabled", True),
|
||||
"description": "Store documents in S3-compatible object storage",
|
||||
"details": {
|
||||
"bucket": getattr(settings, "s3_bucket_name", "Not set"),
|
||||
@@ -327,7 +327,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "sftp_username", None)
|
||||
and (getattr(settings, "sftp_password", None) or getattr(settings, "sftp_private_key", None))
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "sftp_enabled", True),
|
||||
"description": "Upload files to SFTP server",
|
||||
"details": {
|
||||
"host": getattr(settings, "sftp_host", "Not set"),
|
||||
@@ -362,7 +362,7 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
and getattr(settings, "webdav_username", None)
|
||||
and getattr(settings, "webdav_password", None)
|
||||
),
|
||||
"enabled": True,
|
||||
"enabled": getattr(settings, "webdav_enabled", True),
|
||||
"description": "Store documents on WebDAV servers",
|
||||
"details": {
|
||||
"url": getattr(settings, "webdav_url", "Not set"),
|
||||
@@ -373,4 +373,19 @@ def get_provider_status() -> dict[str, dict[str, object]]:
|
||||
},
|
||||
}
|
||||
|
||||
# Check iCloud Drive configuration
|
||||
providers["iCloud Drive"] = {
|
||||
"name": "iCloud Drive",
|
||||
"icon": "fa-brands fa-apple",
|
||||
"configured": bool(getattr(settings, "icloud_username", None) and getattr(settings, "icloud_password", None)),
|
||||
"enabled": getattr(settings, "icloud_enabled", True),
|
||||
"description": "Store documents in Apple iCloud Drive",
|
||||
"details": {
|
||||
"username": getattr(settings, "icloud_username", "Not set"),
|
||||
"password": mask_sensitive_value(getattr(settings, "icloud_password", None)),
|
||||
"folder": getattr(settings, "icloud_folder", "Not set"),
|
||||
"cookie_directory": getattr(settings, "icloud_cookie_directory", "Not set"),
|
||||
},
|
||||
}
|
||||
|
||||
return providers
|
||||
|
||||
@@ -106,6 +106,35 @@ def get_settings_for_display(show_values: bool = False) -> dict[str, list[dict[s
|
||||
"email_sender",
|
||||
"email_default_recipient",
|
||||
],
|
||||
"Watch Folders": [
|
||||
"watch_folders",
|
||||
"watch_folder_poll_interval",
|
||||
"watch_folder_delete_after_process",
|
||||
"ftp_ingest_enabled",
|
||||
"ftp_ingest_folder",
|
||||
"ftp_ingest_delete_after_process",
|
||||
"sftp_ingest_enabled",
|
||||
"sftp_ingest_folder",
|
||||
"sftp_ingest_delete_after_process",
|
||||
"dropbox_ingest_enabled",
|
||||
"dropbox_ingest_folder",
|
||||
"dropbox_ingest_delete_after_process",
|
||||
"google_drive_ingest_enabled",
|
||||
"google_drive_ingest_folder_id",
|
||||
"google_drive_ingest_delete_after_process",
|
||||
"onedrive_ingest_enabled",
|
||||
"onedrive_ingest_folder_path",
|
||||
"onedrive_ingest_delete_after_process",
|
||||
"nextcloud_ingest_enabled",
|
||||
"nextcloud_ingest_folder",
|
||||
"nextcloud_ingest_delete_after_process",
|
||||
"s3_ingest_enabled",
|
||||
"s3_ingest_prefix",
|
||||
"s3_ingest_delete_after_process",
|
||||
"webdav_ingest_enabled",
|
||||
"webdav_ingest_folder",
|
||||
"webdav_ingest_delete_after_process",
|
||||
],
|
||||
"IMAP": [
|
||||
"imap1_host",
|
||||
"imap1_port",
|
||||
|
||||
@@ -59,13 +59,43 @@ def validate_auth_config() -> list[str]:
|
||||
and getattr(settings, "authentik_config_url", None)
|
||||
)
|
||||
|
||||
if not using_simple_auth and not using_oidc:
|
||||
issues.append("Neither simple authentication nor OIDC are properly configured")
|
||||
# Check if any social login provider is enabled
|
||||
using_social_login = any(
|
||||
getattr(settings, f"social_auth_{p}_enabled", False) for p in ("google", "microsoft", "apple", "dropbox")
|
||||
)
|
||||
|
||||
if not using_simple_auth and not using_oidc and not using_social_login:
|
||||
issues.append("Neither simple authentication, OIDC, nor social login are properly configured")
|
||||
|
||||
# If using OIDC, check for provider name
|
||||
if using_oidc and not getattr(settings, "oauth_provider_name", None):
|
||||
issues.append("OAUTH_PROVIDER_NAME is not configured but OIDC is enabled")
|
||||
|
||||
# Validate individual social login provider configs
|
||||
if getattr(settings, "social_auth_google_enabled", False):
|
||||
if not getattr(settings, "social_auth_google_client_id", None):
|
||||
issues.append("SOCIAL_AUTH_GOOGLE_CLIENT_ID is required when Google login is enabled")
|
||||
if not getattr(settings, "social_auth_google_client_secret", None):
|
||||
issues.append("SOCIAL_AUTH_GOOGLE_CLIENT_SECRET is required when Google login is enabled")
|
||||
|
||||
if getattr(settings, "social_auth_microsoft_enabled", False):
|
||||
if not getattr(settings, "social_auth_microsoft_client_id", None):
|
||||
issues.append("SOCIAL_AUTH_MICROSOFT_CLIENT_ID is required when Microsoft login is enabled")
|
||||
if not getattr(settings, "social_auth_microsoft_client_secret", None):
|
||||
issues.append("SOCIAL_AUTH_MICROSOFT_CLIENT_SECRET is required when Microsoft login is enabled")
|
||||
|
||||
if getattr(settings, "social_auth_apple_enabled", False):
|
||||
if not getattr(settings, "social_auth_apple_client_id", None):
|
||||
issues.append("SOCIAL_AUTH_APPLE_CLIENT_ID is required when Apple login is enabled")
|
||||
if not getattr(settings, "social_auth_apple_team_id", None):
|
||||
issues.append("SOCIAL_AUTH_APPLE_TEAM_ID is required when Apple login is enabled")
|
||||
|
||||
if getattr(settings, "social_auth_dropbox_enabled", False):
|
||||
if not getattr(settings, "social_auth_dropbox_client_id", None):
|
||||
issues.append("SOCIAL_AUTH_DROPBOX_CLIENT_ID is required when Dropbox login is enabled")
|
||||
if not getattr(settings, "social_auth_dropbox_client_secret", None):
|
||||
issues.append("SOCIAL_AUTH_DROPBOX_CLIENT_SECRET is required when Dropbox login is enabled")
|
||||
|
||||
return issues
|
||||
|
||||
|
||||
@@ -107,12 +137,12 @@ def validate_storage_configs() -> dict[str, list[str]]:
|
||||
|
||||
issues["sftp"] = sftp_issues
|
||||
|
||||
# Validate Email sending
|
||||
# Validate Email sending (destination-specific settings)
|
||||
email_issues = []
|
||||
if not getattr(settings, "email_host", None):
|
||||
email_issues.append("EMAIL_HOST is not configured")
|
||||
if not getattr(settings, "email_default_recipient", None):
|
||||
email_issues.append("EMAIL_DEFAULT_RECIPIENT is not configured")
|
||||
if not getattr(settings, "dest_email_host", None):
|
||||
email_issues.append("DEST_EMAIL_HOST is not configured")
|
||||
if not getattr(settings, "dest_email_default_recipient", None):
|
||||
email_issues.append("DEST_EMAIL_DEFAULT_RECIPIENT is not configured")
|
||||
issues["email"] = email_issues
|
||||
|
||||
# Validate S3
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
"""
|
||||
Database migration utility for transferring data between databases.
|
||||
|
||||
Copies all table rows from a *source* SQLAlchemy database to a *target*
|
||||
database. This is designed for the common scenario of migrating from the
|
||||
built-in SQLite database to an external PostgreSQL / MySQL instance.
|
||||
|
||||
The utility:
|
||||
1. Creates the schema in the target via ``Base.metadata.create_all``.
|
||||
2. Copies rows table-by-table in dependency order.
|
||||
3. Stamps the Alembic version in the target to ``head``.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import MetaData, create_engine, inspect, text
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.engine.url import make_url
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Tables to skip during migration (Alembic manages its own state).
|
||||
_SKIP_TABLES = {"alembic_version"}
|
||||
|
||||
# Ordered list — parent tables first to respect foreign-key constraints.
|
||||
_TABLE_ORDER = [
|
||||
"documents",
|
||||
"files",
|
||||
"file_processing_steps",
|
||||
"processing_logs",
|
||||
"application_settings",
|
||||
"settings_audit_log",
|
||||
"audit_logs",
|
||||
"saved_searches",
|
||||
"webhook_configs",
|
||||
"shared_links",
|
||||
]
|
||||
|
||||
|
||||
def _make_engine(url: str) -> Engine:
|
||||
"""Create a SQLAlchemy engine from *url* with sensible defaults."""
|
||||
parsed = make_url(url)
|
||||
connect_args: dict[str, Any] = {}
|
||||
if parsed.get_backend_name() == "sqlite":
|
||||
connect_args["check_same_thread"] = False
|
||||
return create_engine(url, connect_args=connect_args)
|
||||
|
||||
|
||||
def _ordered_tables(inspector: Any) -> list[str]:
|
||||
"""Return table names in safe insertion order.
|
||||
|
||||
Tables listed in ``_TABLE_ORDER`` come first (in that order); any
|
||||
remaining tables are appended alphabetically.
|
||||
"""
|
||||
existing = set(inspector.get_table_names())
|
||||
ordered: list[str] = []
|
||||
for name in _TABLE_ORDER:
|
||||
if name in existing and name not in _SKIP_TABLES:
|
||||
ordered.append(name)
|
||||
for name in sorted(existing):
|
||||
if name not in ordered and name not in _SKIP_TABLES:
|
||||
ordered.append(name)
|
||||
return ordered
|
||||
|
||||
|
||||
def preview_migration(source_url: str) -> dict[str, Any]:
|
||||
"""Preview what a migration would do without actually copying data.
|
||||
|
||||
Args:
|
||||
source_url: Connection string for the source database.
|
||||
|
||||
Returns:
|
||||
Dict with ``tables`` (list of dicts with ``name`` and ``row_count``)
|
||||
and ``total_rows``.
|
||||
"""
|
||||
try:
|
||||
src_engine = _make_engine(source_url)
|
||||
src_inspector = inspect(src_engine)
|
||||
tables = _ordered_tables(src_inspector)
|
||||
|
||||
result: list[dict[str, Any]] = []
|
||||
total = 0
|
||||
with src_engine.connect() as conn:
|
||||
for table_name in tables:
|
||||
# table_name is safe — sourced from inspect().get_table_names(), not user input
|
||||
quoted_table = conn.dialect.identifier_preparer.quote(table_name)
|
||||
row = conn.execute(text(f"SELECT COUNT(*) FROM {quoted_table}")).fetchone() # noqa: S608
|
||||
count = row[0] if row else 0
|
||||
result.append({"name": table_name, "row_count": count})
|
||||
total += count
|
||||
|
||||
src_engine.dispose()
|
||||
return {"tables": result, "total_rows": total, "success": True}
|
||||
except Exception as exc:
|
||||
logger.error(f"Migration preview failed: {exc}")
|
||||
return {"success": False, "error": str(exc), "tables": [], "total_rows": 0}
|
||||
|
||||
|
||||
def migrate_data(
|
||||
source_url: str,
|
||||
target_url: str,
|
||||
*,
|
||||
batch_size: int = 500,
|
||||
progress_callback: Any | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Copy all data from *source_url* to *target_url*.
|
||||
|
||||
The target schema is created automatically from the application models.
|
||||
Alembic is stamped to ``head`` in the target after a successful copy.
|
||||
|
||||
Args:
|
||||
source_url: SQLAlchemy connection string for the source DB.
|
||||
target_url: SQLAlchemy connection string for the target DB.
|
||||
batch_size: Number of rows to insert per batch.
|
||||
progress_callback: Optional ``callable(table_name, copied, total)``
|
||||
invoked after each batch.
|
||||
|
||||
Returns:
|
||||
Dict with ``success`` (bool), ``tables_copied`` (int),
|
||||
``rows_copied`` (int), and ``errors`` (list of str).
|
||||
"""
|
||||
errors: list[str] = []
|
||||
tables_copied = 0
|
||||
rows_copied = 0
|
||||
|
||||
try:
|
||||
src_engine = _make_engine(source_url)
|
||||
tgt_engine = _make_engine(target_url)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 1. Create schema in target from application models
|
||||
# ------------------------------------------------------------------
|
||||
from app.database import Base # local import to avoid circular deps
|
||||
|
||||
Base.metadata.create_all(bind=tgt_engine)
|
||||
logger.info("Target schema created from application models.")
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 2. Reflect source schema & determine copy order
|
||||
# ------------------------------------------------------------------
|
||||
src_meta = MetaData()
|
||||
src_meta.reflect(bind=src_engine)
|
||||
|
||||
src_inspector = inspect(src_engine)
|
||||
table_names = _ordered_tables(src_inspector)
|
||||
|
||||
SrcSession = sessionmaker(bind=src_engine)
|
||||
TgtSession = sessionmaker(bind=tgt_engine)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 3. Copy data table-by-table
|
||||
# ------------------------------------------------------------------
|
||||
for table_name in table_names:
|
||||
try:
|
||||
src_session = SrcSession()
|
||||
tgt_session = TgtSession()
|
||||
|
||||
src_table = src_meta.tables.get(table_name)
|
||||
if src_table is None:
|
||||
continue
|
||||
|
||||
# Read all rows from source
|
||||
rows = src_session.execute(src_table.select()).fetchall()
|
||||
column_names = [c.name for c in src_table.columns]
|
||||
|
||||
if not rows:
|
||||
logger.info(f"Skipping empty table: {table_name}")
|
||||
tables_copied += 1
|
||||
src_session.close()
|
||||
tgt_session.close()
|
||||
continue
|
||||
|
||||
# Reflect the target table to insert into
|
||||
tgt_meta = MetaData()
|
||||
tgt_meta.reflect(bind=tgt_engine, only=[table_name])
|
||||
tgt_table = tgt_meta.tables.get(table_name)
|
||||
if tgt_table is None:
|
||||
errors.append(f"Target table {table_name} not found after schema creation")
|
||||
src_session.close()
|
||||
tgt_session.close()
|
||||
continue
|
||||
|
||||
# Batch insert
|
||||
total_for_table = len(rows)
|
||||
for i in range(0, total_for_table, batch_size):
|
||||
batch = rows[i : i + batch_size]
|
||||
# strict=False: column count should always match, but tolerate
|
||||
# minor schema drift (e.g. extra columns) to avoid crashing mid-migration.
|
||||
insert_data = [dict(zip(column_names, row, strict=False)) for row in batch]
|
||||
tgt_session.execute(tgt_table.insert(), insert_data)
|
||||
tgt_session.commit()
|
||||
|
||||
rows_copied += len(batch)
|
||||
if progress_callback:
|
||||
progress_callback(table_name, min(i + batch_size, total_for_table), total_for_table)
|
||||
|
||||
tables_copied += 1
|
||||
logger.info(f"Copied {total_for_table} rows from {table_name}")
|
||||
src_session.close()
|
||||
tgt_session.close()
|
||||
|
||||
except Exception as exc:
|
||||
msg = f"Error copying table {table_name}: {exc}"
|
||||
logger.error(msg)
|
||||
errors.append(msg)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# 4. Stamp Alembic to head in the target
|
||||
# ------------------------------------------------------------------
|
||||
try:
|
||||
_stamp_alembic_head(tgt_engine)
|
||||
logger.info("Alembic version stamped to head in target database.")
|
||||
except Exception as exc:
|
||||
msg = f"Failed to stamp Alembic version: {exc}"
|
||||
logger.error(msg)
|
||||
errors.append(msg)
|
||||
|
||||
src_engine.dispose()
|
||||
tgt_engine.dispose()
|
||||
|
||||
return {
|
||||
"success": len(errors) == 0,
|
||||
"tables_copied": tables_copied,
|
||||
"rows_copied": rows_copied,
|
||||
"errors": errors,
|
||||
}
|
||||
|
||||
except Exception as exc:
|
||||
logger.error(f"Migration failed: {exc}")
|
||||
return {
|
||||
"success": False,
|
||||
"tables_copied": tables_copied,
|
||||
"rows_copied": rows_copied,
|
||||
"errors": errors + [str(exc)],
|
||||
}
|
||||
|
||||
|
||||
def _stamp_alembic_head(engine: Engine) -> None:
|
||||
"""Stamp the Alembic version table to ``head`` in the given engine."""
|
||||
from pathlib import Path
|
||||
|
||||
from alembic import command
|
||||
from alembic.config import Config
|
||||
|
||||
migrations_dir = str(Path(__file__).resolve().parent.parent.parent / "migrations")
|
||||
alembic_cfg = Config()
|
||||
alembic_cfg.set_main_option("script_location", migrations_dir)
|
||||
alembic_cfg.set_main_option("sqlalchemy.url", "")
|
||||
|
||||
with engine.begin() as connection:
|
||||
alembic_cfg.attributes["connection"] = connection
|
||||
command.stamp(alembic_cfg, "head")
|
||||
@@ -0,0 +1,257 @@
|
||||
"""
|
||||
Database configuration wizard utilities.
|
||||
|
||||
Provides helpers for building, validating, and testing database connection
|
||||
strings. Used by both the interactive wizard UI and the REST API.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import create_engine, text
|
||||
from sqlalchemy.engine.url import make_url
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Supported database backends with human-readable labels and defaults.
|
||||
SUPPORTED_BACKENDS: list[dict[str, Any]] = [
|
||||
{
|
||||
"id": "sqlite",
|
||||
"label": "SQLite (Development)",
|
||||
"driver": "",
|
||||
"default_port": None,
|
||||
"description": "File-based database. Best for development and single-user setups.",
|
||||
"requires_host": False,
|
||||
},
|
||||
{
|
||||
"id": "postgresql",
|
||||
"label": "PostgreSQL (Recommended for Production)",
|
||||
"driver": "",
|
||||
"default_port": 5432,
|
||||
"description": "Robust, full-featured database. Recommended for production.",
|
||||
"requires_host": True,
|
||||
},
|
||||
{
|
||||
"id": "mysql",
|
||||
"label": "MySQL / MariaDB",
|
||||
"driver": "pymysql",
|
||||
"default_port": 3306,
|
||||
"description": "Popular open-source database. Requires pymysql driver.",
|
||||
"requires_host": True,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def get_supported_backends() -> list[dict[str, Any]]:
|
||||
"""Return the list of supported database backends with metadata.
|
||||
|
||||
Returns:
|
||||
List of backend descriptor dicts.
|
||||
"""
|
||||
return SUPPORTED_BACKENDS
|
||||
|
||||
|
||||
def build_connection_string(
|
||||
backend: str,
|
||||
host: str = "",
|
||||
port: int | None = None,
|
||||
database: str = "",
|
||||
username: str = "",
|
||||
password: str = "",
|
||||
ssl_mode: str = "",
|
||||
extra_options: str = "",
|
||||
sqlite_path: str = "",
|
||||
) -> str:
|
||||
"""Build a SQLAlchemy connection string from individual components.
|
||||
|
||||
Args:
|
||||
backend: Database backend identifier (``sqlite``, ``postgresql``, ``mysql``).
|
||||
host: Database server hostname or IP.
|
||||
port: Database server port (uses backend default when ``None``).
|
||||
database: Database / schema name.
|
||||
username: Authentication username.
|
||||
password: Authentication password.
|
||||
ssl_mode: SSL mode (e.g. ``require``, ``verify-full``). PostgreSQL only.
|
||||
extra_options: Additional query-string options appended to the URL.
|
||||
sqlite_path: File path for SQLite databases.
|
||||
|
||||
Returns:
|
||||
A SQLAlchemy-compatible connection URL string.
|
||||
|
||||
Raises:
|
||||
ValueError: If required fields are missing for the chosen backend.
|
||||
"""
|
||||
if backend == "sqlite":
|
||||
path = sqlite_path.strip() if sqlite_path else ""
|
||||
if not path:
|
||||
path = "./app/database.db"
|
||||
return f"sqlite:///{path}"
|
||||
|
||||
# Resolve driver prefix
|
||||
backend_info = next((b for b in SUPPORTED_BACKENDS if b["id"] == backend), None)
|
||||
if backend_info is None:
|
||||
raise ValueError(f"Unsupported backend: {backend}")
|
||||
|
||||
if not host:
|
||||
raise ValueError("Host is required for non-SQLite backends")
|
||||
if not database:
|
||||
raise ValueError("Database name is required for non-SQLite backends")
|
||||
if not username:
|
||||
raise ValueError("Username is required for non-SQLite backends")
|
||||
|
||||
driver_suffix = f"+{backend_info['driver']}" if backend_info["driver"] else ""
|
||||
scheme = f"{backend}{driver_suffix}"
|
||||
|
||||
resolved_port = port if port else backend_info["default_port"]
|
||||
|
||||
# Build query parameters
|
||||
params: list[str] = []
|
||||
if ssl_mode:
|
||||
params.append(f"sslmode={ssl_mode}")
|
||||
if extra_options:
|
||||
params.append(extra_options)
|
||||
if backend == "mysql" and "charset=" not in extra_options:
|
||||
params.append("charset=utf8mb4")
|
||||
|
||||
query_string = "&".join(params)
|
||||
|
||||
# Construct URL
|
||||
auth = username
|
||||
if password:
|
||||
auth = f"{username}:{password}"
|
||||
|
||||
url = f"{scheme}://{auth}@{host}:{resolved_port}/{database}"
|
||||
if query_string:
|
||||
url = f"{url}?{query_string}"
|
||||
|
||||
return url
|
||||
|
||||
|
||||
def parse_connection_string(url: str) -> dict[str, Any]:
|
||||
"""Parse a SQLAlchemy connection string into its components.
|
||||
|
||||
Args:
|
||||
url: A SQLAlchemy database URL string.
|
||||
|
||||
Returns:
|
||||
Dict with keys: ``backend``, ``host``, ``port``, ``database``,
|
||||
``username``, ``password``, ``ssl_mode``, ``is_sqlite``.
|
||||
"""
|
||||
try:
|
||||
parsed = make_url(url)
|
||||
backend_name = parsed.get_backend_name()
|
||||
return {
|
||||
"backend": backend_name,
|
||||
"host": parsed.host or "",
|
||||
"port": parsed.port,
|
||||
"database": parsed.database or "",
|
||||
"username": parsed.username or "",
|
||||
"password": parsed.password or "",
|
||||
"ssl_mode": "",
|
||||
"is_sqlite": backend_name == "sqlite",
|
||||
"valid": True,
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.warning(f"Failed to parse connection string: {exc}")
|
||||
return {"valid": False, "error": str(exc)}
|
||||
|
||||
|
||||
def test_connection(url: str, timeout: int = 10) -> dict[str, Any]:
|
||||
"""Attempt to connect to a database and return status information.
|
||||
|
||||
The function creates a short-lived engine, executes a simple ``SELECT 1``
|
||||
query, and disposes the engine. It does **not** modify any global state.
|
||||
|
||||
Args:
|
||||
url: SQLAlchemy database URL to test.
|
||||
timeout: Connection timeout in seconds.
|
||||
|
||||
Returns:
|
||||
Dict with ``success`` (bool), ``message`` (str), and optional
|
||||
``server_version`` (str).
|
||||
"""
|
||||
try:
|
||||
parsed = make_url(url)
|
||||
backend = parsed.get_backend_name()
|
||||
|
||||
connect_args: dict[str, Any] = {}
|
||||
kwargs: dict[str, Any] = {"pool_pre_ping": True}
|
||||
|
||||
if backend == "sqlite":
|
||||
connect_args["check_same_thread"] = False
|
||||
else:
|
||||
kwargs["pool_timeout"] = timeout
|
||||
|
||||
test_engine = create_engine(
|
||||
url,
|
||||
connect_args=connect_args,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
with test_engine.connect() as conn:
|
||||
result = conn.execute(text("SELECT 1"))
|
||||
result.fetchone()
|
||||
|
||||
# Try to fetch server version for informational display
|
||||
server_version = _get_server_version(conn, backend)
|
||||
|
||||
test_engine.dispose()
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": "Connection successful",
|
||||
"backend": backend,
|
||||
"server_version": server_version,
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.warning(f"Connection test failed: {exc}")
|
||||
return {
|
||||
"success": False,
|
||||
"message": str(exc),
|
||||
"backend": "",
|
||||
"server_version": "",
|
||||
}
|
||||
|
||||
|
||||
def _get_server_version(conn: Any, backend: str) -> str:
|
||||
"""Retrieve a human-readable server version string.
|
||||
|
||||
Args:
|
||||
conn: An active SQLAlchemy connection.
|
||||
backend: Backend identifier (``sqlite``, ``postgresql``, ``mysql``).
|
||||
|
||||
Returns:
|
||||
Server version string, or empty string on failure.
|
||||
"""
|
||||
try:
|
||||
if backend == "postgresql":
|
||||
row = conn.execute(text("SELECT version()")).fetchone()
|
||||
return str(row[0]) if row else ""
|
||||
elif backend == "mysql":
|
||||
row = conn.execute(text("SELECT version()")).fetchone()
|
||||
return str(row[0]) if row else ""
|
||||
elif backend == "sqlite":
|
||||
row = conn.execute(text("SELECT sqlite_version()")).fetchone()
|
||||
return f"SQLite {row[0]}" if row else ""
|
||||
except Exception:
|
||||
logger.debug("Could not retrieve server version")
|
||||
return ""
|
||||
|
||||
|
||||
def validate_url_format(url: str) -> dict[str, Any]:
|
||||
"""Validate that a connection string is syntactically correct.
|
||||
|
||||
Args:
|
||||
url: The connection string to validate.
|
||||
|
||||
Returns:
|
||||
Dict with ``valid`` (bool) and optional ``error`` (str).
|
||||
"""
|
||||
try:
|
||||
parsed = make_url(url)
|
||||
backend = parsed.get_backend_name()
|
||||
if backend not in ("sqlite", "postgresql", "mysql"):
|
||||
return {"valid": False, "error": f"Unsupported backend: {backend}"}
|
||||
return {"valid": True, "backend": backend}
|
||||
except Exception as exc:
|
||||
return {"valid": False, "error": str(exc)}
|
||||
@@ -8,6 +8,11 @@ from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Pattern for valid filenames (alphanumeric, dash, underscore, period, and space)
|
||||
# Used for validating GPT-provided filenames and other inputs
|
||||
VALID_FILENAME_PATTERN = r"^[\w\-\. ]+$"
|
||||
VALID_FILENAME_RE = re.compile(VALID_FILENAME_PATTERN)
|
||||
|
||||
|
||||
def get_unique_filename(original_path: str, check_exists_func: Callable[[str], bool] | None = None) -> str:
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,765 @@
|
||||
"""Internationalization (i18n) and localization (l10n) utilities.
|
||||
|
||||
Provides a JSON-based translation system for the DocuElevate UI with:
|
||||
|
||||
* **77 supported languages** covering European, Asian, Middle-Eastern, African, and other languages
|
||||
* Browser ``Accept-Language`` detection with cookie & user-profile persistence
|
||||
* AI-powered fallback translation via the configured LLM provider
|
||||
* Locale-aware date, number, and file-size formatting helpers
|
||||
* Jinja2 integration via a ``_()`` global function
|
||||
|
||||
Language resolution order:
|
||||
1. User profile ``preferred_language`` (persisted in DB)
|
||||
2. ``docuelevate_lang`` cookie
|
||||
3. ``Accept-Language`` HTTP header
|
||||
4. Default (``en``)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import date, datetime
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from starlette.requests import Request
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Supported languages (ordered by priority)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SUPPORTED_LANGUAGES: list[dict[str, str]] = [
|
||||
# --- Tier 1: Primary European languages ---
|
||||
{"code": "en", "name": "English", "native": "English", "flag": "🇬🇧"},
|
||||
{"code": "de", "name": "German", "native": "Deutsch", "flag": "🇩🇪"},
|
||||
{"code": "fr", "name": "French", "native": "Français", "flag": "🇫🇷"},
|
||||
{"code": "es", "name": "Spanish", "native": "Español", "flag": "🇪🇸"},
|
||||
{"code": "it", "name": "Italian", "native": "Italiano", "flag": "🇮🇹"},
|
||||
{"code": "pt", "name": "Portuguese", "native": "Português", "flag": "🇵🇹"},
|
||||
# --- Tier 2: Western & Northern European ---
|
||||
{"code": "nl", "name": "Dutch", "native": "Nederlands", "flag": "🇳🇱"},
|
||||
{"code": "nb", "name": "Norwegian Bokmål", "native": "Norsk bokmål", "flag": "🇳🇴"},
|
||||
{"code": "no", "name": "Norwegian", "native": "Norsk", "flag": "🇳🇴"},
|
||||
{"code": "da", "name": "Danish", "native": "Dansk", "flag": "🇩🇰"},
|
||||
{"code": "sv", "name": "Swedish", "native": "Svenska", "flag": "🇸🇪"},
|
||||
{"code": "fi", "name": "Finnish", "native": "Suomi", "flag": "🇫🇮"},
|
||||
{"code": "is", "name": "Icelandic", "native": "Íslenska", "flag": "🇮🇸"},
|
||||
{"code": "ga", "name": "Irish", "native": "Gaeilge", "flag": "🇮🇪"},
|
||||
{"code": "lb", "name": "Luxembourgish", "native": "Lëtzebuergesch", "flag": "🇱🇺"},
|
||||
{"code": "ca", "name": "Catalan", "native": "Català", "flag": "🏴"},
|
||||
{"code": "cy", "name": "Welsh", "native": "Cymraeg", "flag": "🏴"}, # Wales subdivision flag (U+1F3F4 + tag chars)
|
||||
{"code": "fy", "name": "Western Frisian", "native": "Frysk", "flag": "🇳🇱"},
|
||||
{"code": "gl", "name": "Galician", "native": "Galego", "flag": "🇪🇸"},
|
||||
{"code": "li", "name": "Limburgish", "native": "Limburgs", "flag": "🇳🇱"},
|
||||
{"code": "vls", "name": "Flemish", "native": "West-Vlams", "flag": "🇧🇪"},
|
||||
{"code": "nds", "name": "Low German", "native": "Plattdüütsch", "flag": "🇩🇪"},
|
||||
# --- Tier 3: Central & Eastern European ---
|
||||
{"code": "pl", "name": "Polish", "native": "Polski", "flag": "🇵🇱"},
|
||||
{"code": "cs", "name": "Czech", "native": "Čeština", "flag": "🇨🇿"},
|
||||
{"code": "sk", "name": "Slovak", "native": "Slovenčina", "flag": "🇸🇰"},
|
||||
{"code": "hu", "name": "Hungarian", "native": "Magyar", "flag": "🇭🇺"},
|
||||
{"code": "sl", "name": "Slovenian", "native": "Slovenščina", "flag": "🇸🇮"},
|
||||
{"code": "hr", "name": "Croatian", "native": "Hrvatski", "flag": "🇭🇷"},
|
||||
{"code": "ro", "name": "Romanian", "native": "Română", "flag": "🇷🇴"},
|
||||
{"code": "bg", "name": "Bulgarian", "native": "Български", "flag": "🇧🇬"},
|
||||
{"code": "el", "name": "Greek", "native": "Ελληνικά", "flag": "🇬🇷"},
|
||||
{"code": "et", "name": "Estonian", "native": "Eesti", "flag": "🇪🇪"},
|
||||
{"code": "lv", "name": "Latvian", "native": "Latviešu", "flag": "🇱🇻"},
|
||||
{"code": "lt", "name": "Lithuanian", "native": "Lietuvių", "flag": "🇱🇹"},
|
||||
{"code": "sr", "name": "Serbian", "native": "Српски", "flag": "🇷🇸"},
|
||||
# --- Tier 4: Non-EU European, Middle Eastern & African ---
|
||||
{"code": "tr", "name": "Turkish", "native": "Türkçe", "flag": "🇹🇷"},
|
||||
{"code": "uk", "name": "Ukrainian", "native": "Українська", "flag": "🇺🇦"},
|
||||
{"code": "he", "name": "Hebrew", "native": "עברית", "flag": "🇮🇱"},
|
||||
{"code": "ar", "name": "Arabic", "native": "العربية", "flag": "🇸🇦"},
|
||||
{"code": "fa", "name": "Persian", "native": "فارسی", "flag": "🇮🇷"},
|
||||
{"code": "af", "name": "Afrikaans", "native": "Afrikaans", "flag": "🇿🇦"},
|
||||
# --- Tier 5: Asian languages ---
|
||||
{"code": "zh", "name": "Chinese", "native": "中文", "flag": "🇨🇳"},
|
||||
{"code": "zh-TW", "name": "Traditional Chinese", "native": "繁體中文", "flag": "🇹🇼"},
|
||||
{"code": "ja", "name": "Japanese", "native": "日本語", "flag": "🇯🇵"},
|
||||
{"code": "ko", "name": "Korean", "native": "한국어", "flag": "🇰🇷"},
|
||||
{"code": "vi", "name": "Vietnamese", "native": "Tiếng Việt", "flag": "🇻🇳"},
|
||||
{"code": "pa", "name": "Punjabi", "native": "ਪੰਜਾਬੀ", "flag": "🇮🇳"},
|
||||
{"code": "kn", "name": "Kannada", "native": "ಕನ್ನಡ", "flag": "🇮🇳"},
|
||||
{"code": "hi", "name": "Hindi", "native": "हिन्दी", "flag": "🇮🇳"},
|
||||
{"code": "bn", "name": "Bengali", "native": "বাংলা", "flag": "🇧🇩"},
|
||||
{"code": "gu", "name": "Gujarati", "native": "ગુજરાતી", "flag": "🇮🇳"},
|
||||
{"code": "ml", "name": "Malayalam", "native": "മലയാളം", "flag": "🇮🇳"},
|
||||
{"code": "mr", "name": "Marathi", "native": "मराठी", "flag": "🇮🇳"},
|
||||
{"code": "ta", "name": "Tamil", "native": "தமிழ்", "flag": "🇮🇳"},
|
||||
{"code": "te", "name": "Telugu", "native": "తెలుగు", "flag": "🇮🇳"},
|
||||
{"code": "ur", "name": "Urdu", "native": "اردو", "flag": "🇵🇰"},
|
||||
{"code": "si", "name": "Sinhala", "native": "සිංහල", "flag": "🇱🇰"},
|
||||
{"code": "ne", "name": "Nepali", "native": "नेपाली", "flag": "🇳🇵"},
|
||||
{"code": "th", "name": "Thai", "native": "ไทย", "flag": "🇹🇭"},
|
||||
{"code": "km", "name": "Khmer", "native": "ខ្មែរ", "flag": "🇰🇭"},
|
||||
{"code": "id", "name": "Indonesian", "native": "Bahasa Indonesia", "flag": "🇮🇩"},
|
||||
{"code": "ms", "name": "Malay", "native": "Bahasa Melayu", "flag": "🇲🇾"},
|
||||
{"code": "jv", "name": "Javanese", "native": "Basa Jawa", "flag": "🇮🇩"},
|
||||
{"code": "tl", "name": "Tagalog", "native": "Filipino", "flag": "🇵🇭"},
|
||||
{"code": "mn", "name": "Mongolian", "native": "Монгол", "flag": "🇲🇳"},
|
||||
{"code": "kk", "name": "Kazakh", "native": "Қазақ тілі", "flag": "🇰🇿"},
|
||||
{"code": "uz", "name": "Uzbek", "native": "Oʻzbekcha", "flag": "🇺🇿"},
|
||||
{"code": "az", "name": "Azerbaijani", "native": "Azərbaycan dili", "flag": "🇦🇿"},
|
||||
{"code": "hy", "name": "Armenian", "native": "Հայերեն", "flag": "🇦🇲"},
|
||||
{"code": "ka", "name": "Georgian", "native": "ქართული", "flag": "🇬🇪"},
|
||||
# --- Tier 6: African languages ---
|
||||
{"code": "sw", "name": "Swahili", "native": "Kiswahili", "flag": "🇰🇪"},
|
||||
{"code": "am", "name": "Amharic", "native": "አማርኛ", "flag": "🇪🇹"},
|
||||
{"code": "ha", "name": "Hausa", "native": "Hausa", "flag": "🇳🇬"},
|
||||
{"code": "yo", "name": "Yoruba", "native": "Yorùbá", "flag": "🇳🇬"},
|
||||
{"code": "ig", "name": "Igbo", "native": "Igbo", "flag": "🇳🇬"},
|
||||
{"code": "zu", "name": "Zulu", "native": "isiZulu", "flag": "🇿🇦"},
|
||||
# --- Tier 7: Constructed & other languages ---
|
||||
{"code": "eo", "name": "Esperanto", "native": "Esperanto", "flag": "🌍"},
|
||||
]
|
||||
|
||||
SUPPORTED_LANGUAGE_CODES: set[str] = {lang["code"] for lang in SUPPORTED_LANGUAGES}
|
||||
DEFAULT_LANGUAGE = "en"
|
||||
|
||||
# Lookup map for fast code → language-dict resolution
|
||||
_LANG_CODE_MAP: dict[str, dict[str, str]] = {lang["code"]: lang for lang in SUPPORTED_LANGUAGES}
|
||||
|
||||
# Global-usage order used to fill remaining slots in the smart suggestions list
|
||||
_POPULAR_LANGUAGE_CODES: list[str] = ["en", "zh", "es", "ar", "fr", "de", "ja", "pt", "hi", "ko"]
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Translation file loading
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_TRANSLATIONS_DIR = Path(__file__).resolve().parent.parent.parent / "frontend" / "translations"
|
||||
_translation_cache: dict[str, dict[str, str]] = {}
|
||||
|
||||
|
||||
def _load_translations(locale: str) -> dict[str, str]:
|
||||
"""Load the translation JSON file for *locale*, with caching."""
|
||||
if locale in _translation_cache:
|
||||
return _translation_cache[locale]
|
||||
|
||||
filepath = _TRANSLATIONS_DIR / f"{locale}.json"
|
||||
if not filepath.is_file():
|
||||
logger.warning("Translation file not found for locale '%s'", locale)
|
||||
_translation_cache[locale] = {}
|
||||
return {}
|
||||
|
||||
try:
|
||||
data: dict[str, str] = json.loads(filepath.read_text(encoding="utf-8"))
|
||||
_translation_cache[locale] = data
|
||||
return data
|
||||
except (json.JSONDecodeError, OSError):
|
||||
logger.exception("Failed to load translations for '%s'", locale)
|
||||
_translation_cache[locale] = {}
|
||||
return {}
|
||||
|
||||
|
||||
def reload_translations() -> None:
|
||||
"""Clear the translation cache so files are re-read on next access."""
|
||||
_translation_cache.clear()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Core translation function
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def translate(key: str, locale: str | None = None, **kwargs: Any) -> str:
|
||||
"""Return the translated string for *key* in *locale*.
|
||||
|
||||
Falls back through:
|
||||
1. Requested *locale*
|
||||
2. English (``en``)
|
||||
3. The raw key itself (to keep the UI functional)
|
||||
|
||||
Positional placeholders ``{0}``, ``{1}`` or named placeholders
|
||||
``{name}`` in the translated string are interpolated via *kwargs*.
|
||||
"""
|
||||
locale = locale if locale and locale in SUPPORTED_LANGUAGE_CODES else DEFAULT_LANGUAGE
|
||||
|
||||
translations = _load_translations(locale)
|
||||
value = translations.get(key)
|
||||
|
||||
# Fallback to English
|
||||
if value is None and locale != DEFAULT_LANGUAGE:
|
||||
en_translations = _load_translations(DEFAULT_LANGUAGE)
|
||||
value = en_translations.get(key)
|
||||
|
||||
# Fallback to key itself
|
||||
if value is None:
|
||||
value = key
|
||||
|
||||
if kwargs:
|
||||
try:
|
||||
value = value.format(**kwargs)
|
||||
except (KeyError, IndexError):
|
||||
pass # Return unformatted string rather than crash
|
||||
|
||||
return value
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AI fallback translation (best-effort, non-blocking)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_ai_translation_cache: dict[tuple[str, str], str] = {}
|
||||
|
||||
|
||||
def translate_with_ai_fallback(text: str, target_locale: str) -> str:
|
||||
"""Translate *text* using the configured AI provider as a fallback.
|
||||
|
||||
Returns the original *text* unchanged when:
|
||||
* The target locale is English (source language)
|
||||
* The AI provider is unavailable or returns an error
|
||||
* The translation has already been cached
|
||||
|
||||
Results are cached in-memory for the lifetime of the process.
|
||||
"""
|
||||
if target_locale == DEFAULT_LANGUAGE or target_locale not in SUPPORTED_LANGUAGE_CODES:
|
||||
return text
|
||||
|
||||
cache_key = (text, target_locale)
|
||||
if cache_key in _ai_translation_cache:
|
||||
return _ai_translation_cache[cache_key]
|
||||
|
||||
target_name = next(
|
||||
(lang["name"] for lang in SUPPORTED_LANGUAGES if lang["code"] == target_locale),
|
||||
target_locale,
|
||||
)
|
||||
|
||||
try:
|
||||
from litellm import completion # type: ignore[import-untyped]
|
||||
|
||||
from app.config import settings
|
||||
|
||||
model = getattr(settings, "ai_model", None) or getattr(settings, "openai_model", "gpt-4o-mini")
|
||||
response = completion(
|
||||
model=model,
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
f"You are a professional translator. Translate the following UI text "
|
||||
f"from English to {target_name}. Return ONLY the translated text, "
|
||||
f"nothing else. Keep any HTML tags, placeholders like {{name}}, "
|
||||
f"and special characters intact."
|
||||
),
|
||||
},
|
||||
{"role": "user", "content": text},
|
||||
],
|
||||
max_tokens=256,
|
||||
temperature=0.1,
|
||||
)
|
||||
translated = response.choices[0].message.content.strip()
|
||||
_ai_translation_cache[cache_key] = translated
|
||||
return translated
|
||||
except Exception:
|
||||
logger.debug("AI fallback translation failed for '%s' → %s", text[:50], target_locale)
|
||||
return text
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Language detection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def detect_language(request: Request) -> str:
|
||||
"""Determine the preferred UI language from the request context.
|
||||
|
||||
Resolution order:
|
||||
1. ``preferred_language`` stored in the user session
|
||||
2. ``docuelevate_lang`` cookie
|
||||
3. ``Accept-Language`` HTTP header (best match)
|
||||
4. Default → ``en``
|
||||
"""
|
||||
# 1. User session preference
|
||||
if hasattr(request, "session"):
|
||||
session_lang = request.session.get("preferred_language")
|
||||
if isinstance(session_lang, str) and session_lang in SUPPORTED_LANGUAGE_CODES:
|
||||
return session_lang
|
||||
|
||||
# 2. Cookie
|
||||
if hasattr(request, "cookies"):
|
||||
cookie_lang = request.cookies.get("docuelevate_lang")
|
||||
if isinstance(cookie_lang, str) and cookie_lang in SUPPORTED_LANGUAGE_CODES:
|
||||
return cookie_lang
|
||||
|
||||
# 3. Accept-Language header
|
||||
accept = ""
|
||||
if hasattr(request, "headers"):
|
||||
accept = request.headers.get("accept-language", "")
|
||||
lang = _parse_accept_language(accept)
|
||||
if lang:
|
||||
return lang
|
||||
|
||||
return DEFAULT_LANGUAGE
|
||||
|
||||
|
||||
def _parse_accept_language_entries(header: str) -> list[tuple[float, str]]:
|
||||
"""Parse an ``Accept-Language`` header into quality-sorted ``(q, tag)`` pairs."""
|
||||
if not header:
|
||||
return []
|
||||
|
||||
entries: list[tuple[float, str]] = []
|
||||
for raw_part in header.split(","):
|
||||
part = raw_part.strip()
|
||||
if not part:
|
||||
continue
|
||||
if ";q=" in part:
|
||||
lang_tag, _, q_str = part.partition(";q=")
|
||||
try:
|
||||
quality = float(q_str.strip())
|
||||
except ValueError:
|
||||
quality = 0.0
|
||||
else:
|
||||
lang_tag = part
|
||||
quality = 1.0
|
||||
entries.append((quality, lang_tag.strip().lower()))
|
||||
|
||||
entries.sort(key=lambda e: e[0], reverse=True)
|
||||
return entries
|
||||
|
||||
|
||||
def _parse_accept_language(header: str) -> str | None:
|
||||
"""Extract the best matching language from an ``Accept-Language`` header.
|
||||
|
||||
Parses quality values and returns the highest-priority match among
|
||||
:data:`SUPPORTED_LANGUAGE_CODES`, or ``None`` if nothing matches.
|
||||
"""
|
||||
for _quality, tag in _parse_accept_language_entries(header):
|
||||
code = tag.split("-")[0]
|
||||
if code in SUPPORTED_LANGUAGE_CODES:
|
||||
return code
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# Maximum number of languages shown in the compact nav-bar dropdown
|
||||
_SUGGESTED_LANGUAGES_MAX = 6
|
||||
|
||||
|
||||
def get_suggested_languages(current_locale: str, accept_language_header: str = "") -> list[dict[str, str]]:
|
||||
"""Return up to :data:`_SUGGESTED_LANGUAGES_MAX` suggested languages for the compact picker.
|
||||
|
||||
Selection priority:
|
||||
1. The currently active language (always included first).
|
||||
2. Languages listed in the browser's ``Accept-Language`` header.
|
||||
3. Popular global languages (by estimated speaker count) as fillers.
|
||||
|
||||
The resulting list is de-duplicated and capped at
|
||||
:data:`_SUGGESTED_LANGUAGES_MAX` entries.
|
||||
"""
|
||||
candidates: list[str] = []
|
||||
|
||||
# 1. Active locale first
|
||||
if current_locale in SUPPORTED_LANGUAGE_CODES:
|
||||
candidates.append(current_locale)
|
||||
|
||||
# 2. Browser preferences
|
||||
for _quality, tag in _parse_accept_language_entries(accept_language_header):
|
||||
if len(candidates) >= _SUGGESTED_LANGUAGES_MAX:
|
||||
break
|
||||
code = tag.split("-")[0]
|
||||
if code in SUPPORTED_LANGUAGE_CODES and code not in candidates:
|
||||
candidates.append(code)
|
||||
|
||||
# 3. Popular language fillers
|
||||
for code in _POPULAR_LANGUAGE_CODES:
|
||||
if len(candidates) >= _SUGGESTED_LANGUAGES_MAX:
|
||||
break
|
||||
if code not in candidates and code in SUPPORTED_LANGUAGE_CODES:
|
||||
candidates.append(code)
|
||||
|
||||
return [_LANG_CODE_MAP[c] for c in candidates[:_SUGGESTED_LANGUAGES_MAX] if c in _LANG_CODE_MAP]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Localization helpers (l10n)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Locale-specific formatting rules for date/number display
|
||||
_LOCALE_FORMATS: dict[str, dict[str, Any]] = {
|
||||
"en": {
|
||||
"date": "%B %d, %Y",
|
||||
"date_short": "%m/%d/%Y",
|
||||
"datetime": "%B %d, %Y %I:%M %p",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"de": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"fr": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u202f",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"es": {
|
||||
"date": "%d de %B de %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d de %B de %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"it": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"pt": {
|
||||
"date": "%d de %B de %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d de %B de %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"nl": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d-%m-%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"nb": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"da": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"sv": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%Y-%m-%d",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"fi": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"is": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"ga": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"lb": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"ca": {
|
||||
"date": "%d de %B de %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d de %B de %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"pl": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"cs": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"sk": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"hu": {
|
||||
"date": "%Y. %B %d.",
|
||||
"date_short": "%Y.%m.%d.",
|
||||
"datetime": "%Y. %B %d. %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"sl": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"hr": {
|
||||
"date": "%d. %B %Y.",
|
||||
"date_short": "%d.%m.%Y.",
|
||||
"datetime": "%d. %B %Y. %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"ro": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"bg": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"el": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"et": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"lv": {
|
||||
"date": "%Y. gada %d. %B",
|
||||
"date_short": "%d.%m.%Y.",
|
||||
"datetime": "%Y. gada %d. %B %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"lt": {
|
||||
"date": "%Y m. %B %d d.",
|
||||
"date_short": "%Y-%m-%d",
|
||||
"datetime": "%Y m. %B %d d. %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"tr": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"uk": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"zh": {
|
||||
"date": "%Y年%m月%d日",
|
||||
"date_short": "%Y/%m/%d",
|
||||
"datetime": "%Y年%m月%d日 %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"ru": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
# --- New languages ---
|
||||
"no": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"cy": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"fy": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d-%m-%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"gl": {
|
||||
"date": "%d de %B de %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d de %B de %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"li": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d-%m-%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"vls": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"nds": {
|
||||
"date": "%d. %B %Y",
|
||||
"date_short": "%d.%m.%Y",
|
||||
"datetime": "%d. %B %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"sr": {
|
||||
"date": "%d. %B %Y.",
|
||||
"date_short": "%d.%m.%Y.",
|
||||
"datetime": "%d. %B %Y. %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"he": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"ar": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%Y/%m/%d",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"fa": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%Y/%m/%d",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"af": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%Y/%m/%d",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"ja": {
|
||||
"date": "%Y年%m月%d日",
|
||||
"date_short": "%Y/%m/%d",
|
||||
"datetime": "%Y年%m月%d日 %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"ko": {
|
||||
"date": "%Y년 %m월 %d일",
|
||||
"date_short": "%Y.%m.%d",
|
||||
"datetime": "%Y년 %m월 %d일 %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"vi": {
|
||||
"date": "ngày %d tháng %m năm %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "ngày %d tháng %m năm %Y %H:%M",
|
||||
"thousands_sep": ".",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
"pa": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"kn": {
|
||||
"date": "%d %B %Y",
|
||||
"date_short": "%d/%m/%Y",
|
||||
"datetime": "%d %B %Y %H:%M",
|
||||
"thousands_sep": ",",
|
||||
"decimal_sep": ".",
|
||||
},
|
||||
"eo": {
|
||||
"date": "%d-a de %B %Y",
|
||||
"date_short": "%Y-%m-%d",
|
||||
"datetime": "%d-a de %B %Y %H:%M",
|
||||
"thousands_sep": "\u00a0",
|
||||
"decimal_sep": ",",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def format_date(value: date | datetime | None, locale: str = DEFAULT_LANGUAGE, short: bool = False) -> str:
|
||||
"""Format a date/datetime value according to the locale conventions."""
|
||||
if value is None:
|
||||
return ""
|
||||
fmt_key = "date_short" if short else "date"
|
||||
fmt = _LOCALE_FORMATS.get(locale, _LOCALE_FORMATS[DEFAULT_LANGUAGE])[fmt_key]
|
||||
return value.strftime(fmt)
|
||||
|
||||
|
||||
def format_datetime(value: datetime | None, locale: str = DEFAULT_LANGUAGE) -> str:
|
||||
"""Format a datetime value according to the locale conventions."""
|
||||
if value is None:
|
||||
return ""
|
||||
fmt = _LOCALE_FORMATS.get(locale, _LOCALE_FORMATS[DEFAULT_LANGUAGE])["datetime"]
|
||||
return value.strftime(fmt)
|
||||
|
||||
|
||||
def format_number(value: int | float, locale: str = DEFAULT_LANGUAGE) -> str:
|
||||
"""Format a number with locale-appropriate thousand separators."""
|
||||
lf = _LOCALE_FORMATS.get(locale, _LOCALE_FORMATS[DEFAULT_LANGUAGE])
|
||||
if isinstance(value, float):
|
||||
int_part, _, dec_part = f"{value:,.2f}".partition(".")
|
||||
formatted_int = int_part.replace(",", lf["thousands_sep"])
|
||||
return f"{formatted_int}{lf['decimal_sep']}{dec_part}"
|
||||
return f"{value:,}".replace(",", lf["thousands_sep"])
|
||||
|
||||
|
||||
@lru_cache(maxsize=32)
|
||||
def get_language_info(code: str) -> dict[str, str] | None:
|
||||
"""Return the metadata dict for a supported language code, or ``None``."""
|
||||
for lang in SUPPORTED_LANGUAGES:
|
||||
if lang["code"] == code:
|
||||
return lang
|
||||
return None
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user