Compare commits
1755 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e9595c7868 | |||
| 58b14ae769 | |||
| 23c5bac666 | |||
| d925dc5cd3 | |||
| b8ddd2f8d2 | |||
| 301ca9d186 | |||
| 3bd8a52ea2 | |||
| 789e8c6236 | |||
| a3ea215a1c | |||
| c6e0b80bec | |||
| e86e1b9f13 | |||
| 46a9a30af0 | |||
| bdfa3ba1e0 | |||
| 8295279ec9 | |||
| 152ee15b06 | |||
| a75e8b9297 | |||
| ee664f83fb | |||
| 91ef089aa7 | |||
| 925864ddca | |||
| 57db4c7c82 | |||
| 35752c9092 | |||
| 9b9882c4d6 | |||
| 6a77533795 | |||
| 69053bfb08 | |||
| f1cf5d0e76 | |||
| c547ad1acc | |||
| 1625896e30 | |||
| 470f08d893 | |||
| a57766ed7e | |||
| 76f202f7f1 | |||
| 33484b236a | |||
| 6927e7643f | |||
| aeb50c21d2 | |||
| 45e41338dc | |||
| f0d3563029 | |||
| 12a35f9b30 | |||
| 4136033bf0 | |||
| 48331f6e91 | |||
| cafc0e4523 | |||
| a2c9915908 | |||
| c3124b08bd | |||
| 4faba2ec08 | |||
| 2f3c22000c | |||
| aca12858c1 | |||
| ae524bb94e | |||
| 5af4dbcb25 | |||
| 94a090da77 | |||
| 3fd8b32724 | |||
| 7f20c903ef | |||
| 8fcc223ef1 | |||
| ce050b542a | |||
| 114b69a8c2 | |||
| f6591d10fc | |||
| c26c376e2e | |||
| 627a8579de | |||
| f041f28d9f | |||
| 7dec570ce6 | |||
| c7d3ec57c3 | |||
| 11a49eb7fd | |||
| 527fb666d0 | |||
| 07bdee94b0 | |||
| b7e33af669 | |||
| 965647940b | |||
| 2f5e2a0fcd | |||
| c9bb2b6807 | |||
| c03ce8cdb2 | |||
| 0b8f967eb5 | |||
| 315d85c443 | |||
| f8f95085fc | |||
| 7bd9d20091 | |||
| c61afb2c33 | |||
| dd5603bdd0 | |||
| 0c3ee6f484 | |||
| 6f510d5a2d | |||
| 2014a93c1b | |||
| 3be93be35a | |||
| 4cac9fbe9b | |||
| bcdbf9d178 | |||
| fc1365dfec | |||
| 52e8e535ff | |||
| ef5528dcef | |||
| ea2dad0c08 | |||
| 8e26e3aaa7 | |||
| 2a5296d7e7 | |||
| a052b2fbe0 | |||
| a384b222f1 | |||
| 326adb1858 | |||
| e330a611d0 | |||
| 15dd1a8471 | |||
| 2ee6bfc7ea | |||
| 248619d91e | |||
| 8984d4da70 | |||
| 0596206e17 | |||
| 26963a8464 | |||
| 1e7f2275d3 | |||
| 88368f7f76 | |||
| 7fbcf5c593 | |||
| 01c04c20ce | |||
| cfcce57e35 | |||
| 10297ede37 | |||
| 78bd5b5904 | |||
| 9153b1f7f0 | |||
| f9b4975093 | |||
| 47595818b4 | |||
| bad369548b | |||
| 7ea8b17fd2 | |||
| cc5e879ea9 | |||
| 7490462c67 | |||
| c25e1b0e21 | |||
| a10f8e628e | |||
| 1018ea17d9 | |||
| 7c1967b728 | |||
| 06b0fced38 | |||
| 341839fe5e | |||
| 28d4bced0c | |||
| d22175310a | |||
| 7755f5a1ed | |||
| cee6d6d4e1 | |||
| 0497fbbbad | |||
| 57795ee487 | |||
| a4bd1d7178 | |||
| d94e9ca4bc | |||
| 82c6915c42 | |||
| d71945b7b9 | |||
| 91f36e0d5a | |||
| 1e69c55947 | |||
| 9b748db4d4 | |||
| eeae47ddec | |||
| be500e1a2b | |||
| 45d3ac8cf0 | |||
| 4df4673628 | |||
| 9642020887 | |||
| 89dec45062 | |||
| 34457f9775 | |||
| b0fe1a014a | |||
| 1d9bd15a70 | |||
| b50a534454 | |||
| 80de3b6743 | |||
| 8b4280d5dd | |||
| 93629ff440 | |||
| c4e10bee5e | |||
| 084171395d | |||
| 958b195e79 | |||
| c5ef1ec50c | |||
| b4e0067a27 | |||
| 6188003897 | |||
| ef897f660d | |||
| 6cb9feacab | |||
| 76c0e91500 | |||
| 0c7ea6748d | |||
| 78077fa8c7 | |||
| 242846aa9c | |||
| 868613ac49 | |||
| 33a0e49acd | |||
| 14b3031e63 | |||
| 1d7df13c94 | |||
| ce4bca0186 | |||
| 4b07e996ad | |||
| 720c9c11b0 | |||
| 425472c839 | |||
| 55afa4981b | |||
| 63f7b62fc0 | |||
| 48a303d498 | |||
| 8f1fe79411 | |||
| 3e1b352930 | |||
| 46772fc746 | |||
| 25d32a9006 | |||
| 8c6a02885d | |||
| 899cc56638 | |||
| 9822ba583d | |||
| 2288b89cd7 | |||
| 3d0bdf7836 | |||
| 41844c4b60 | |||
| 94aa2ebe57 | |||
| be97a757a3 | |||
| a5df6dc9cb | |||
| d4cc44a72f | |||
| 61dee5ba52 | |||
| 5f94e64734 | |||
| 9be03d8690 | |||
| 5c5b3ac054 | |||
| 9458055661 | |||
| bb116dcdd3 | |||
| 725bf98352 | |||
| 962495ba8c | |||
| 3843bce596 | |||
| 2df92ce469 | |||
| b25aaf879f | |||
| 4120a502df | |||
| 7d6128d78f | |||
| 9c98a8438a | |||
| 28cd5e565b | |||
| 49b816c878 | |||
| e7be6ff907 | |||
| 029bbb2c85 | |||
| 4a35aabdaa | |||
| cd66c5eb4d | |||
| 961cbaea50 | |||
| c5d52fc797 | |||
| b2912da4dc | |||
| 5f5e18d261 | |||
| f64d04fffd | |||
| c084cfabe6 | |||
| f2d69a6e27 | |||
| e84e26ea84 | |||
| 124b802c8f | |||
| 6f2752bdf8 | |||
| b202f10e1a | |||
| bd2fd7a241 | |||
| 83afc6c8f6 | |||
| 096e224b6a | |||
| 44855f03d0 | |||
| 2ea35d419c | |||
| e1643f20e2 | |||
| 5f1911f0b5 | |||
| 5d5622bd36 | |||
| c17afe8c11 | |||
| 0287a165cf | |||
| 840a5bcd5b | |||
| 54a0ba1023 | |||
| 2ca015b38c | |||
| d71add1484 | |||
| f852ba9783 | |||
| 5a3ddcc1f0 | |||
| c70b607939 | |||
| b737b83811 | |||
| 12fa6fefe9 | |||
| f3abe87d85 | |||
| b1aa09c28d | |||
| d89f18edd8 | |||
| 8f0905033c | |||
| cb0fe93812 | |||
| b5ed16c1c8 | |||
| a5a8cd94c9 | |||
| a7a88218c3 | |||
| 4b6412734b | |||
| ad795e200a | |||
| c22bb66c4b | |||
| 30a124d85b | |||
| 4d85b8d03a | |||
| de7b9ec7cf | |||
| bb92b592dd | |||
| bab963ecff | |||
| aee2292ab0 | |||
| 9f5d045648 | |||
| e5feee5aae | |||
| 49cd41e2e4 | |||
| 9dc1000d63 | |||
| 00e0d5fa45 | |||
| 9241b9df5c | |||
| 918005d26b | |||
| 094542e5b1 | |||
| ee007885dd | |||
| 3e6fbb49c4 | |||
| ba5aedcc7b | |||
| 0a192eeeca | |||
| bd26f31778 | |||
| 600ab5fbf2 | |||
| 350c0d14db | |||
| 11ad9c22cf | |||
| 74a8c22478 | |||
| 8cc292c6ee | |||
| 982c222717 | |||
| 2339765866 | |||
| 0e09351c01 | |||
| 059e092510 | |||
| 1a53218a53 | |||
| ae8be68df9 | |||
| 9a94ff62ed | |||
| 1e026c2fe9 | |||
| 051d763a7e | |||
| d487a60484 | |||
| f9b6d93213 | |||
| 690d9d96ca | |||
| a27a4ce130 | |||
| 7da93b5b15 | |||
| 64ee68b1aa | |||
| 71e1a6fe1c | |||
| 167145579f | |||
| 9f8a9b349d | |||
| f1e5ab6ce4 | |||
| 81484ad770 | |||
| 5b04504b0e | |||
| 9229be88ff | |||
| 35caf24e3c | |||
| bc122e351d | |||
| 2f6dbea1ce | |||
| 902f109551 | |||
| db88cde66e | |||
| ee1810692a | |||
| e8de3b3761 | |||
| 0a44b06b6d | |||
| 34ff7f8de8 | |||
| e518bce922 | |||
| b70341a062 | |||
| 5b2d51f647 | |||
| 6f5a73f98a | |||
| 65cf33ce89 | |||
| f9f36df38d | |||
| 831e1c602d | |||
| 95263f132c | |||
| ae675485af | |||
| b3a238744d | |||
| d6c21b8026 | |||
| 1e1e6e6280 | |||
| a8423064ec | |||
| 23c8c76b39 | |||
| ed01952610 | |||
| 1afd42bc57 | |||
| 28aa72ae4d | |||
| 72b0b49a7f | |||
| 91eecd9396 | |||
| 74ed8b9bd9 | |||
| 8f9abac014 | |||
| b5fce418fb | |||
| 3b5ca04ebc | |||
| 906d76f08b | |||
| 7c28cdda07 | |||
| b13713dab8 | |||
| 5e3e2b1999 | |||
| 0e6a4c5084 | |||
| 445d629949 | |||
| b1723b4c5f | |||
| fa4d09c5b6 | |||
| 250cce4f4d | |||
| f2ba74a483 | |||
| 78c3717661 | |||
| d6de835aed | |||
| c81e29cd46 | |||
| 3d286df8af | |||
| 60e3ea030a | |||
| 84fe8543c1 | |||
| 1465040864 | |||
| 57f9e90e45 | |||
| 579bd261ce | |||
| d8fc75d5a5 | |||
| 910fb297ba | |||
| 867b269322 | |||
| d1f9819f4e | |||
| 6541529250 | |||
| 67c17e7baa | |||
| 689c616e44 | |||
| 933fb940f9 | |||
| f7e4f81773 | |||
| 0252f11cc0 | |||
| f020a3e292 | |||
| 044ae72c50 | |||
| 9076394440 | |||
| ec882214e2 | |||
| d5c18ccf07 | |||
| cde966012c | |||
| d53390cada | |||
| aec6c3944d | |||
| 547ce4abc4 | |||
| cfe83d7efa | |||
| f549505bfd | |||
| 1559686f90 | |||
| 136631762b | |||
| 1572f322d7 | |||
| 5c15a2395a | |||
| 0aab5bcbf7 | |||
| 6c699d1904 | |||
| 24bbe3889b | |||
| 3e6ff61117 | |||
| 441a2b5c2e | |||
| ad730c1e71 | |||
| 6cc4599d1e | |||
| 1d636d866a | |||
| 5869dd6fa6 | |||
| 31a72026a8 | |||
| 71a7a57adc | |||
| d34b8bceb9 | |||
| 34ea9333fb | |||
| b12e891682 | |||
| 9c5bd73794 | |||
| 2941f6e177 | |||
| be65703875 | |||
| 9ac8f29448 | |||
| 786c909765 | |||
| ba4ebd83b8 | |||
| 1773c12cb1 | |||
| e1bd976697 | |||
| e4749b4e7c | |||
| f2b7db88ba | |||
| 5ef82050b8 | |||
| 98327edfc2 | |||
| 6cc185a507 | |||
| 55127fee68 | |||
| f91c57eacb | |||
| a4aaebfe66 | |||
| 11c8d80d59 | |||
| d8d2016f85 | |||
| 2192783737 | |||
| 30608a7eb2 | |||
| ec459c54ff | |||
| 057933ff2f | |||
| ba8c88bc17 | |||
| a8eb6504ac | |||
| 6727253958 | |||
| e7eda8af5e | |||
| 020bc6a9c7 | |||
| bf4be55779 | |||
| 7e786712f4 | |||
| dc0a19bd11 | |||
| 2db65647ee | |||
| f6a2bae05e | |||
| a08b103271 | |||
| 1d7286c4c6 | |||
| 342e2d1614 | |||
| c9f9001244 | |||
| 6bf121f02f | |||
| bf6b9177af | |||
| 4a07c49bc7 | |||
| a3c657b947 | |||
| a77a29444d | |||
| 3ac49965a5 | |||
| ad329d0aa7 | |||
| faa68adaa1 | |||
| 571cc81789 | |||
| 70b193e07d | |||
| 723b14e660 | |||
| 56bf665397 | |||
| 0f6a1ee1ec | |||
| 9fe50a87d1 | |||
| 27ebbddd5f | |||
| bcf2d00c33 | |||
| c79dc4abef | |||
| 264089cce4 | |||
| 30c2e9afef | |||
| 8905031d16 | |||
| 7e1ed14ee6 | |||
| c7d67031d4 | |||
| 581adf0e26 | |||
| 03712cfb08 | |||
| bd03d1971b | |||
| a022dec9c2 | |||
| b8075b6821 | |||
| a0dace75a3 | |||
| 7b7554decf | |||
| a007b4fd98 | |||
| aa49fa3ae6 | |||
| 4d019d53d9 | |||
| f75b125992 | |||
| 935e8a626e | |||
| b3dce16838 | |||
| 0b0f43fa99 | |||
| 0e472d3515 | |||
| 15fc90f240 | |||
| 1350aa6a5e | |||
| d5e1d92d6f | |||
| d1ebac74a1 | |||
| 475c41d375 | |||
| fef74450c7 | |||
| 8d366f3b1e | |||
| 204000aabc | |||
| 7dffdc0554 | |||
| 51821092f4 | |||
| bdc26846bd | |||
| 651b48658c | |||
| 96bfba8057 | |||
| 421744865f | |||
| a4588a57cb | |||
| a88d790445 | |||
| 8a3ae8652e | |||
| 20e61db050 | |||
| 4d302b495c | |||
| 6d03d107fd | |||
| b48ed0b2ac | |||
| 64403f17f1 | |||
| 793b6d3e25 | |||
| d322ec6dc7 | |||
| 84b9461cfe | |||
| 75a6379aad | |||
| 1bd9840c81 | |||
| 1caebdded0 | |||
| a7b993ad42 | |||
| 5b09b7b9ae | |||
| 0f449abe17 | |||
| 480591fcee | |||
| 242ef587e1 | |||
| afb168d3e8 | |||
| 45c2cd7c6e | |||
| 9b21872007 | |||
| da6b337a8b | |||
| b2492b497c | |||
| 89b3a6a3c0 | |||
| 3cd895d57c | |||
| 872ff6c53b | |||
| dca5ee2ef3 | |||
| a764e52f77 | |||
| 134e05b5a3 | |||
| b6af39c567 | |||
| acfba4b58c | |||
| daf236b270 | |||
| b5438f4265 | |||
| a579233710 | |||
| 5054ec8a93 | |||
| 4c8926a1fa | |||
| 84ce65f532 | |||
| c6e4cc22b2 | |||
| fee5a09954 | |||
| 4db4df8621 | |||
| 0e144710b5 | |||
| 95a1a8eaea | |||
| 1ff9b44a54 | |||
| 5131635d44 | |||
| 0e271b7162 | |||
| f6d69fc826 | |||
| fca90e5bf1 | |||
| b906b7bf9d | |||
| 803331d61f | |||
| 04c6823079 | |||
| 24328c584e | |||
| 195c3c3446 | |||
| 1c81678c9f | |||
| 6d48976bc5 | |||
| 66696d2b07 | |||
| fca6b0fb81 | |||
| 419714e013 | |||
| de623b69f2 | |||
| bebeab3191 | |||
| 4de6b439ce | |||
| 10c11815f6 | |||
| 1bec4e02e2 | |||
| 21e2f89827 | |||
| 905d128dd8 | |||
| b8df5ba253 | |||
| 16f89433da | |||
| 036210dd45 | |||
| df89bbf79b | |||
| b13e2a6049 | |||
| 6c532627f1 | |||
| af788a8ea1 | |||
| 718d82c815 | |||
| 825f3cc3a8 | |||
| a5ef7e4903 | |||
| 98a4d7a72c | |||
| bb39bc1ccd | |||
| f5a5d1f2f4 | |||
| 82f9ca35cd | |||
| 9bc505d2f2 | |||
| 77df9628c3 | |||
| 5a7095495b | |||
| 280d508425 | |||
| 438b79ae04 | |||
| 3e541ea655 | |||
| cd65ae4343 | |||
| b78a328626 | |||
| ebbbd3d62e | |||
| 4619e6b8a6 | |||
| 62f246df9c | |||
| 0444b14d87 | |||
| e8d60006c2 | |||
| 97fe83fba2 | |||
| c81914f78d | |||
| 0c10fcbafc | |||
| 7bd667e271 | |||
| 075a505085 | |||
| b7a3b301a3 | |||
| 8cc817930d | |||
| 3a221a62cd | |||
| 0e72a965d5 | |||
| 491aface29 | |||
| 1b1cbfce39 | |||
| f2255f9a1c | |||
| 9d6bfde288 | |||
| 9a256741d5 | |||
| 15a9ed9435 | |||
| 8125a01f11 | |||
| 3beb243b31 | |||
| ddea132b68 | |||
| 3447a408db | |||
| cf2ffc97f6 | |||
| 40cc10c6e2 | |||
| 1f1157e86f | |||
| 2a8a4b7471 | |||
| cd7322d989 | |||
| 9dc2ec3a7a | |||
| 6bb695b2ae | |||
| b8a1ac52b3 | |||
| 66fdb11e39 | |||
| 6837198409 | |||
| e10f0bff42 | |||
| 0f312160bc | |||
| cb1e81355d | |||
| 7ff91af2cb | |||
| 5734df2d50 | |||
| 1bba4899ba | |||
| 275f706c87 | |||
| d30ac49c85 | |||
| fd15c36665 | |||
| ff1fde5e48 | |||
| 4f7f33cf1e | |||
| c7ff177e17 | |||
| 46c4031276 | |||
| 012be0dffb | |||
| 7798ac3b57 | |||
| dd7c8f0342 | |||
| 70c46c8ec0 | |||
| 6e497aadca | |||
| 4d4706e078 | |||
| b64c8b9d34 | |||
| 8ce41d723e | |||
| f68f8d8e31 | |||
| 683af42fe8 | |||
| b5ac98889c | |||
| 8ad90d7da9 | |||
| ff369a2ac1 | |||
| 92996bc2f5 | |||
| 740d18555b | |||
| dcfa1ab70c | |||
| 5e986cb6d6 | |||
| e679b71356 | |||
| bcd49aa793 | |||
| 8f8e11cffe | |||
| 1c8cebb718 | |||
| 3be78708cd | |||
| 992adad978 | |||
| df1fa51800 | |||
| be6023464a | |||
| c9940965a6 | |||
| e4e2521aec | |||
| fb9fc01780 | |||
| 7b21a69ceb | |||
| 827979598e | |||
| b9b8796153 | |||
| b3d1824d66 | |||
| 00be8d7b8e | |||
| e29e222833 | |||
| 6e2e54aac6 | |||
| 0d606d480f | |||
| a9af06ed17 | |||
| 6da1e6fd81 | |||
| 056292dbe6 | |||
| 50846360b3 | |||
| 13a156f4e3 | |||
| c00a35bbac | |||
| 635966b099 | |||
| 30718218cc | |||
| 5b3c7da644 | |||
| 21f9998706 | |||
| 68398522e7 | |||
| e5cce5d184 | |||
| 1d51266208 | |||
| 3088459c70 | |||
| 1102a495e5 | |||
| b8db664c2e | |||
| fffb7cf357 | |||
| fbd4f83730 | |||
| 320a2acedd | |||
| 2471921204 | |||
| d58c43c7b5 | |||
| b4118f6162 | |||
| 71d2d4100e | |||
| 6f5f4d9d49 | |||
| 80a0ddcbfc | |||
| 9a85615811 | |||
| b290cffb98 | |||
| c941738644 | |||
| 18c49c6b2d | |||
| d1f64ebfba | |||
| ac35c5e6fa | |||
| 94dc6f967d | |||
| fb8aef3e2a | |||
| bb59233d33 | |||
| 0425d46c44 | |||
| f24c39a027 | |||
| ca2d023d81 | |||
| 7242f3c168 | |||
| fa9b037d5a | |||
| ff4093c9a1 | |||
| 8bb6457c65 | |||
| c1657a01a7 | |||
| dab881b9b6 | |||
| 705c801158 | |||
| cc2a07b090 | |||
| fe20e02f78 | |||
| 2c68b3c197 | |||
| 84c6e1c5dd | |||
| 275a5ad6fa | |||
| d8372c6fb8 | |||
| 7fadbfa992 | |||
| b4e28046fc | |||
| 6a044753eb | |||
| d18c05c36d | |||
| e4e3ac4077 | |||
| d8906aece0 | |||
| 040f4dcdd4 | |||
| fa36ec6987 | |||
| 5529b32d54 | |||
| 2cfbea29a9 | |||
| 433d1eb639 | |||
| 6e3e6238a1 | |||
| 726e4dfdc4 | |||
| c76e51391b | |||
| 522cefad93 | |||
| df64aece2c | |||
| 8eb2e97113 | |||
| 9c1be9ec10 | |||
| 7a004f782e | |||
| df4b4ae18c | |||
| c3d06d1876 | |||
| ae9ed6e9a7 | |||
| 5e911ed268 | |||
| 829e95d674 | |||
| 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 | |||
| 2b698cc694 | |||
| 13aa14b8e4 | |||
| b85fc1d277 | |||
| ac6e052788 | |||
| c8987d966b | |||
| 542fb46ee7 | |||
| 491c424580 | |||
| 6a83d51d88 | |||
| 289dcc375c | |||
| ce2a76fb77 | |||
| 666f739f4e | |||
| d167be8274 | |||
| 23498e0a98 | |||
| bd898605de | |||
| 1e5e35a26a | |||
| df051e8b81 | |||
| e95693d684 | |||
| c286e1b394 | |||
| 40d56f0396 | |||
| 653c137222 | |||
| 3b491ea84c | |||
| ff76855f29 | |||
| 35db9f88de | |||
| 73295cad33 | |||
| b4131e0d19 | |||
| 3afd406c59 | |||
| 954e67640c | |||
| 91e50a4441 | |||
| 76353349e7 | |||
| 341dad643f | |||
| 9ee3249146 | |||
| e6dfa079cc | |||
| 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 |
@@ -0,0 +1,91 @@
|
||||
# =============================================================================
|
||||
# Docker build context exclusions
|
||||
# Reducing the build context speeds up builds and prevents unnecessary cache
|
||||
# invalidation when unrelated files change.
|
||||
# =============================================================================
|
||||
|
||||
# ── Version control ──────────────────────────────────────────────────────────
|
||||
.git
|
||||
|
||||
# ── GitHub / CI tooling ──────────────────────────────────────────────────────
|
||||
.github
|
||||
|
||||
# ── IDE / local dev ──────────────────────────────────────────────────────────
|
||||
.vscode
|
||||
.jules
|
||||
|
||||
# ── Pre-commit / linting config (not needed at runtime) ──────────────────────
|
||||
.pre-commit-config.yaml
|
||||
pyproject.toml
|
||||
codecov.yml
|
||||
crowdin.yml
|
||||
|
||||
# ── Test suite ───────────────────────────────────────────────────────────────
|
||||
tests/
|
||||
requirements-dev.txt
|
||||
coverage.json
|
||||
COVERAGE_REPORT.md
|
||||
.coverage
|
||||
htmlcov/
|
||||
.pytest_cache/
|
||||
junit.xml
|
||||
coverage.xml
|
||||
|
||||
# ── Mobile app / browser extension / legacy placeholder ─────────────────────
|
||||
# backend/ is an empty placeholder directory not part of the Python application
|
||||
mobile/
|
||||
browser-extension/
|
||||
backend/
|
||||
|
||||
# ── Helm charts ──────────────────────────────────────────────────────────────
|
||||
helm/
|
||||
|
||||
# ── Scripts (run before Docker build, output files are COPYd separately) ─────
|
||||
scripts/
|
||||
|
||||
# ── Benchmark and one-off utility scripts ────────────────────────────────────
|
||||
benchmark_*.py
|
||||
fix_test*.py
|
||||
run_fast_tests.sh
|
||||
|
||||
# ── Root-level Markdown files (docs/ is kept for docs-builder stage) ─────────
|
||||
# Note: *.md only matches files at the root level, not inside subdirectories
|
||||
*.md
|
||||
|
||||
# ── Python bytecode / compiled artifacts ─────────────────────────────────────
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.pyo
|
||||
*.pyd
|
||||
*.so
|
||||
*.egg
|
||||
*.egg-info/
|
||||
|
||||
# ── Virtual environments ──────────────────────────────────────────────────────
|
||||
.venv/
|
||||
venv/
|
||||
env/
|
||||
|
||||
# ── Environment / secret files ───────────────────────────────────────────────
|
||||
.env
|
||||
.env.local
|
||||
.env.*.local
|
||||
|
||||
# ── Runtime state files ───────────────────────────────────────────────────────
|
||||
*.log
|
||||
celerybeat-schedule
|
||||
celerybeat.pid
|
||||
|
||||
# ── Build artifacts ───────────────────────────────────────────────────────────
|
||||
build/
|
||||
dist/
|
||||
.cache/
|
||||
.mypy_cache/
|
||||
.ruff_cache/
|
||||
site/
|
||||
docs_build/
|
||||
|
||||
# ── Editor temp files ─────────────────────────────────────────────────────────
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
@@ -5,6 +5,29 @@ 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)
|
||||
|
||||
# **System Reset / Factory Reset**
|
||||
# FACTORY_RESET_ON_STARTUP=false # Wipe all user data on every startup (demo/testing only)
|
||||
# ENABLE_FACTORY_RESET=false # Show the System Reset page in admin UI
|
||||
|
||||
# **Logging**
|
||||
# LOG_LEVEL controls the Python root-logger level.
|
||||
# Accepted values: DEBUG, INFO, WARNING, ERROR, CRITICAL (default: INFO).
|
||||
# When DEBUG=true and LOG_LEVEL is not set, the level is automatically lowered to DEBUG.
|
||||
# LOG_LEVEL=INFO
|
||||
# DEBUG=false
|
||||
|
||||
# Log output format: "text" (human-readable, default) or "json" (structured JSON lines).
|
||||
# Use "json" when shipping logs to Grafana Loki, Splunk, ELK, Datadog, or any SIEM.
|
||||
# LOG_FORMAT=text
|
||||
|
||||
# Forward application logs to a syslog receiver (in addition to stdout).
|
||||
# Useful for traditional (non-container) deployments and centralised SIEM ingestion.
|
||||
# LOG_SYSLOG_ENABLED=false
|
||||
# LOG_SYSLOG_HOST=localhost
|
||||
# LOG_SYSLOG_PORT=514
|
||||
# LOG_SYSLOG_PROTOCOL=udp # udp | tcp
|
||||
|
||||
# **UI / Appearance**
|
||||
# Default colour scheme: system (follow OS), light, or dark
|
||||
@@ -16,6 +39,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 +119,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
|
||||
@@ -119,16 +164,76 @@ AUTH_ENABLED=true
|
||||
# Generate a secure random string, for example:
|
||||
# python -c "import secrets; print(secrets.token_hex(32))"
|
||||
SESSION_SECRET=b39fd43f68d0491ca942f28a16e484b1e763fe9accf4445ca2669a5f3b179eb4
|
||||
|
||||
# Session lifetime in days (default: 30). Common values: 30, 60, 90.
|
||||
# Determines how long a user stays logged in before needing to re-authenticate.
|
||||
# SESSION_LIFETIME_DAYS=30
|
||||
# Override with a custom value (takes precedence over SESSION_LIFETIME_DAYS):
|
||||
# SESSION_LIFETIME_CUSTOM_DAYS=
|
||||
|
||||
# Time-to-live in seconds for QR code login challenges (default: 120 = 2 minutes).
|
||||
# QR_LOGIN_CHALLENGE_TTL_SECONDS=120
|
||||
|
||||
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
|
||||
@@ -167,17 +272,90 @@ OPENAI_MODEL=gpt-4o-mini
|
||||
# AZURE_OPENAI_API_VERSION=2024-02-01
|
||||
# AI_MODEL=gpt-4o # deployment name in Azure
|
||||
|
||||
# **Document Translation**
|
||||
# After processing, documents whose detected language differs from the default
|
||||
# target language are automatically translated. Only the original and this
|
||||
# default-language version are persisted; other translations are on-the-fly.
|
||||
# Users can override this in their profile settings.
|
||||
# DEFAULT_DOCUMENT_LANGUAGE=en
|
||||
|
||||
# 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 +378,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 +396,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 +420,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,13 +440,24 @@ 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
|
||||
ONEDRIVE_REFRESH_TOKEN=your-refresh-token
|
||||
ONEDRIVE_FOLDER_PATH=Documents/Uploads
|
||||
|
||||
# SharePoint
|
||||
SHAREPOINT_CLIENT_ID=your-client-id
|
||||
SHAREPOINT_CLIENT_SECRET=your-client-secret
|
||||
SHAREPOINT_TENANT_ID=common
|
||||
SHAREPOINT_REFRESH_TOKEN=your-refresh-token
|
||||
SHAREPOINT_SITE_URL=https://tenant.sharepoint.com/sites/sitename
|
||||
SHAREPOINT_DOCUMENT_LIBRARY=Documents
|
||||
SHAREPOINT_FOLDER_PATH=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 +465,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 +477,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 +490,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 +525,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 +556,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.
|
||||
|
||||
+72
-197
@@ -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,161 @@ 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
|
||||
- run: ruff check app/ tests/
|
||||
- run: ruff format --check app/ tests/
|
||||
|
||||
- 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)
|
||||
# ══════════════════════════════════════════════════════════════════════════
|
||||
migration-chain:
|
||||
name: Alembic Migration Chain Check
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.11"
|
||||
- name: Validate migration chain
|
||||
run: python scripts/check_alembic_migrations.py
|
||||
|
||||
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 --ignore-vuln CVE-2026-4539
|
||||
|
||||
- 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, migration-chain]
|
||||
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 +183,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 +225,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 +232,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,30 @@
|
||||
## 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.
|
||||
## 2026-03-22 - B310: urllib.request.urlopen replaced with httpx
|
||||
**Vulnerability:** The `_test_webdav_connection` function used `urllib.request.urlopen`, which natively supports dangerous schemes like `file://` or `ftp://` and follows redirects by default, potentially allowing SSRF bypasses or Local File Inclusion.
|
||||
**Learning:** `urllib.request` should be avoided for user-supplied URLs. Even when URL schemes are manually validated, `urllib`'s default redirect following behavior can bypass SSRF protections (e.g. redirecting to `127.0.0.1`).
|
||||
**Prevention:** Use a modern, safer HTTP client like `httpx` with `follow_redirects=False` when testing user-provided URLs.
|
||||
|
||||
## 2026-03-20 - Safe Path Traversal Prevention in Low-Level Utilities
|
||||
**Vulnerability:** The generic file utility `hash_file` in `app/utils/file_operations.py` accepted any file path and was vulnerable to reading arbitrary files via path traversal (e.g., `../../../etc/passwd`) or absolute paths if an attacker could control the `filepath` argument.
|
||||
**Learning:** Naively checking for `".." in path` breaks legitimate relative paths used internally by the application. Blocking absolute paths entirely also breaks functionality. Input validation should occur at the API boundary, but for defense-in-depth, low-level utilities must enforce expected boundaries (e.g., the application's `workdir`).
|
||||
**Prevention:** Use `pathlib.Path.resolve()` on both the target path and the allowed base directory (`settings.workdir`). Ensure the resolved target path is strictly within the allowed boundary using `filepath_obj.relative_to(workdir_obj)`, catching the `ValueError` that is raised when the path is out of bounds. This safely blocks both relative traversal attacks and arbitrary absolute paths.
|
||||
## 2025-05-18 - [SSRF Bypass via DNS Resolution Failure]
|
||||
**Vulnerability:** The `is_private_ip` function in `app/utils/network.py` failed open (returned `False`) when a hostname could not be resolved (`socket.gaierror`).
|
||||
**Learning:** This fail-open pattern was originally added to allow external domains in tests, but in production, it created a severe SSRF risk. An attacker could bypass SSRF protections by providing a URL that fails to resolve during the security check but resolves later (DNS rebinding), or by exploiting internal routing behaviors via unresolvable addresses.
|
||||
**Prevention:** Always fail securely in network authorization functions. If a domain cannot be resolved to verify its safety, the request must be blocked (`return True` / default-deny). Tests should mock DNS resolution correctly instead of compromising production security logic.
|
||||
## 2026-03-26 - SSRF in Integration Connection Tests
|
||||
**Vulnerability:** The `_test_imap_connection` and `_test_s3_connection` functions in `app/api/integrations.py` did not validate user-provided `host` and `endpoint_url` variables against `is_private_ip()`. This allowed an attacker to test the presence of internal IMAP servers or direct S3 SDK API calls to internal infrastructure via SSRF.
|
||||
**Learning:** Any time a new generic connection or integration test is added, SSRF validation may be forgotten if the core network utility (`is_private_ip`) is not systematically applied to all outbound network operations, regardless of the protocol (e.g., IMAP, S3).
|
||||
**Prevention:** Establish a pattern where any user-configurable host or endpoint URL is immediately passed through the centralized `is_private_ip` validation function before any network call or third-party client initialization.
|
||||
|
||||
## 2024-05-27 - SSRF Bypass via HTTP Redirects
|
||||
**Vulnerability:** In `app/api/url_upload.py`, the `validate_url_safety` function was correctly verifying the initially requested URL to prevent fetching internal IPs or cloud metadata endpoints. However, the subsequent `httpx.AsyncClient` was configured with `follow_redirects=True` without validating the destination of those redirects. An attacker could bypass SSRF protections by providing a URL to an attacker-controlled server that responds with a 301/302 redirect pointing to an internal target (e.g., `http://127.0.0.1` or `http://169.254.169.254`).
|
||||
**Learning:** Checking the URL before sending the request is insufficient if the HTTP client automatically follows redirects. The target of every single redirect must be subject to the same strict validation as the initial request.
|
||||
**Prevention:** Avoid `follow_redirects=True` for user-provided URLs when possible. If redirects must be followed, attach an event hook (e.g., `event_hooks={"response": [hook_function]}`) to the `httpx` client to intercept the response, calculate the redirect destination from the `Location` header, and run the URL safety validation logic before the redirect is actually followed.
|
||||
## 2026-03-27 - SSRF Bypass via HTTP Redirects in httpx
|
||||
**Vulnerability:** The `/process-url` endpoint used `httpx.AsyncClient(follow_redirects=True)` after validating the initial user-provided URL against SSRF protections. However, it did not validate the target URLs of any subsequent HTTP redirects, allowing an attacker to provide a safe URL that redirects to an internal/private IP, bypassing the security check.
|
||||
**Learning:** Initial URL validation is insufficient when the HTTP client is configured to follow redirects automatically. The client must be explicitly configured to validate every redirect target.
|
||||
**Prevention:** When using `httpx.AsyncClient(follow_redirects=True)` for user-provided URLs, always implement a redirect validator hook function (e.g., using `event_hooks={'response': [validate_redirect]}`) that resolves the `Location` header and passes it through the same SSRF validation logic before the redirect is followed.
|
||||
@@ -48,6 +48,16 @@ repos:
|
||||
.env.demo
|
||||
)$
|
||||
|
||||
# Alembic migration chain validation
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: check-alembic-migrations
|
||||
name: Check Alembic migration chain
|
||||
entry: python scripts/check_alembic_migrations.py
|
||||
language: python
|
||||
pass_filenames: false
|
||||
files: ^migrations/versions/.*\.py$
|
||||
|
||||
# Conventional commits validation
|
||||
- repo: https://github.com/compilerla/conventional-pre-commit
|
||||
rev: v3.0.0
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
2026-03-01T17:23:25Z
|
||||
2026-04-07T09:34:53Z
|
||||
|
||||
+3939
File diff suppressed because it is too large
Load Diff
+71
-15
@@ -1,20 +1,69 @@
|
||||
# Use multi-stage build for a smaller final image
|
||||
FROM python:3.14.1 AS builder
|
||||
# syntax=docker/dockerfile:1
|
||||
|
||||
WORKDIR /app
|
||||
# ── Stage 1: Python dependency builder ──────────────────────────────────────
|
||||
# Use the same slim variant as the runtime to keep Python versions in sync.
|
||||
# build-essential + libffi-dev cover the few packages (e.g. cryptography) that
|
||||
# need a C compiler; they are discarded after this stage.
|
||||
FROM python:3.14.3-slim AS builder
|
||||
|
||||
# Copy requirements first for better layer caching
|
||||
COPY requirements.txt /app/
|
||||
WORKDIR /build
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
libffi-dev \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Create an isolated virtual environment so only installed packages are copied
|
||||
# to the runtime image (no pip, setuptools, or other builder artefacts).
|
||||
RUN python -m venv /opt/venv
|
||||
|
||||
ENV PATH="/opt/venv/bin:$PATH" \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PIP_NO_CACHE_DIR=1
|
||||
|
||||
COPY requirements.txt /build/
|
||||
RUN pip install --no-cache-dir -r requirements.txt \
|
||||
# Remove bytecode and cache to keep the venv lean
|
||||
&& find /opt/venv -type f -name "*.pyc" -delete \
|
||||
&& find /opt/venv -type d -name "__pycache__" -exec rm -rf {} + 2>/dev/null || true
|
||||
|
||||
# ── Stage 2: Frontend asset builder ─────────────────────────────────────────
|
||||
# Compiles Tailwind CSS (a devDependency) into the minified styles.css.
|
||||
# npm ci installs ALL deps (including devDependencies) so the tailwindcss CLI
|
||||
# is available; using --omit=dev would cause 'tailwindcss: not found'.
|
||||
FROM node:20-slim AS frontend-builder
|
||||
|
||||
WORKDIR /frontend
|
||||
|
||||
COPY frontend/package.json frontend/package-lock.json ./
|
||||
RUN npm ci
|
||||
|
||||
COPY frontend/ ./
|
||||
RUN npm run build
|
||||
|
||||
# ── Stage 4: Documentation builder ──────────────────────────────────────────
|
||||
FROM python:3.14.3-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
|
||||
|
||||
# Second stage for the actual runtime
|
||||
# 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
|
||||
|
||||
# ── Stage 5: Runtime image ───────────────────────────────────────────────────
|
||||
FROM python:3.14.3-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Copy installed packages from builder stage
|
||||
COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
|
||||
COPY --from=builder /usr/local/bin /usr/local/bin
|
||||
# Copy only the pre-built virtual environment from the builder
|
||||
COPY --from=builder /opt/venv /opt/venv
|
||||
|
||||
# Install system-level OCR tools required for local OCR workflows:
|
||||
# tesseract-ocr – OCR engine used by pytesseract and ocrmypdf
|
||||
@@ -33,6 +82,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,15 +92,20 @@ COPY ./BUILD_DATE /app/BUILD_DATE
|
||||
COPY ./GIT_SHA /app/GIT_SHA
|
||||
COPY ./RUNTIME_INFO /app/RUNTIME_INFO
|
||||
|
||||
# Create runtime_info directory
|
||||
RUN mkdir -p /app/runtime_info
|
||||
# Copy the pre-built MkDocs documentation site (served at /help)
|
||||
COPY --from=docs-builder /docs/docs_build /app/docs_build
|
||||
|
||||
# Create necessary directories
|
||||
RUN mkdir -p /workdir
|
||||
# Copy the compiled Tailwind CSS (built in the frontend-builder stage)
|
||||
COPY --from=frontend-builder /frontend/static/styles.css /app/frontend/static/styles.css
|
||||
|
||||
# Create necessary runtime directories in a single layer
|
||||
RUN mkdir -p /app/runtime_info /workdir
|
||||
|
||||
# Set environment variables
|
||||
ENV PYTHONPATH=/app
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV PATH="/opt/venv/bin:$PATH" \
|
||||
PYTHONPATH=/app \
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1
|
||||
|
||||
# Expose the port the app runs on
|
||||
EXPOSE 8000
|
||||
|
||||
+50
-10
@@ -1,45 +1,85 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
|
||||
# Local development Dockerfile (avoids CI-only build metadata files)
|
||||
FROM python:3.14.1 AS builder
|
||||
|
||||
WORKDIR /app
|
||||
# ── Stage 1: Python dependency builder ──────────────────────────────────────
|
||||
FROM python:3.14.3-slim AS builder
|
||||
|
||||
COPY requirements.txt /app/
|
||||
WORKDIR /build
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
build-essential \
|
||||
libffi-dev \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Create an isolated virtual environment
|
||||
RUN python -m venv /opt/venv
|
||||
|
||||
ENV PATH="/opt/venv/bin:$PATH" \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PIP_NO_CACHE_DIR=1
|
||||
|
||||
COPY requirements.txt /build/
|
||||
RUN pip install --no-cache-dir -r requirements.txt \
|
||||
&& find /opt/venv -type f -name "*.pyc" -delete \
|
||||
&& find /opt/venv -type d -name "__pycache__" -exec rm -rf {} + 2>/dev/null || true
|
||||
|
||||
# ── Stage 2: Documentation builder ──────────────────────────────────────────
|
||||
FROM python:3.14.3-slim AS docs-builder
|
||||
|
||||
WORKDIR /docs
|
||||
|
||||
COPY docs/requirements.txt /docs/requirements.txt
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
FROM python:3.14.1-slim
|
||||
COPY docs /docs/docs
|
||||
COPY mkdocs.yml /docs/mkdocs.yml
|
||||
|
||||
RUN mkdocs build --config-file /docs/mkdocs.yml --site-dir /docs/docs_build
|
||||
|
||||
# ── Stage 3: Runtime image ───────────────────────────────────────────────────
|
||||
FROM python:3.14.3-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
COPY --from=builder /usr/local/lib/python3.14/site-packages /usr/local/lib/python3.14/site-packages
|
||||
COPY --from=builder /usr/local/bin /usr/local/bin
|
||||
COPY --from=builder /opt/venv /opt/venv
|
||||
|
||||
# Install system-level OCR tools required for local OCR workflows:
|
||||
# tesseract-ocr – OCR engine used by pytesseract and ocrmypdf
|
||||
# ghostscript – required by ocrmypdf for PDF/PS operations
|
||||
# poppler-utils – provides pdfinfo/pdftoppm used by pdf2image
|
||||
# unpaper – optional deskewing pre-processor used by ocrmypdf
|
||||
# wget – used by ocr_language_manager to download tessdata files
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
tesseract-ocr \
|
||||
ghostscript \
|
||||
poppler-utils \
|
||||
unpaper \
|
||||
wget \
|
||||
&& apt-get clean && rm -rf /var/lib/apt/lists/*
|
||||
|
||||
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
|
||||
|
||||
RUN mkdir -p /app/runtime_info
|
||||
RUN mkdir -p /workdir
|
||||
# Create necessary runtime directories in a single layer
|
||||
RUN mkdir -p /app/runtime_info /workdir
|
||||
|
||||
ENV PYTHONPATH=/app
|
||||
ENV PYTHONUNBUFFERED=1
|
||||
ENV PATH="/opt/venv/bin:$PATH" \
|
||||
PYTHONPATH=/app \
|
||||
PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1
|
||||
|
||||
EXPOSE 8000
|
||||
|
||||
|
||||
@@ -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.172.9
|
||||
Build Date: 2026-04-07T09:34:53Z
|
||||
Git Commit: 3bd8a52ea201b33d6071c9b3a7fdace582e65fd5
|
||||
Git Short SHA: 3bd8a52
|
||||
Git Branch: main
|
||||
Commit Date: 2026-03-01T18:23:06+01:00
|
||||
Build Timestamp: 2026-03-01T17:23:25Z
|
||||
Commit Date: 2026-04-07T11:34:28+02:00
|
||||
Build Timestamp: 2026-04-07T09:34:53Z
|
||||
==============================
|
||||
|
||||
@@ -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,56 @@ 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.automation import router as automation_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.classification_rules import router as classification_rules_router
|
||||
from app.api.comments import router as comments_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.qr_auth import router as qr_auth_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.sessions import router as sessions_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.sharing import router as sharing_router
|
||||
from app.api.similarity import router as similarity_router
|
||||
from app.api.subscriptions import router as subscriptions_router
|
||||
from app.api.system_reset import router as system_reset_router
|
||||
from app.api.translation import router as translation_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 +64,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 +82,33 @@ 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(sessions_router)
|
||||
router.include_router(qr_auth_router)
|
||||
router.include_router(compliance_router)
|
||||
router.include_router(system_reset_router)
|
||||
router.include_router(translation_router)
|
||||
router.include_router(classification_rules_router)
|
||||
router.include_router(automation_router)
|
||||
router.include_router(comments_router)
|
||||
router.include_router(sharing_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,321 @@
|
||||
"""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, timedelta, 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"
|
||||
|
||||
#: Name prefix used for tokens created by the mobile app flow.
|
||||
MOBILE_TOKEN_PREFIX = "Mobile App"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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()
|
||||
|
||||
|
||||
def _token_to_dict(t: ApiToken) -> dict[str, Any]:
|
||||
"""Convert an ``ApiToken`` ORM instance to a serialisable dict."""
|
||||
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,
|
||||
"expires_at": t.expires_at,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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")
|
||||
expires_in_days: int | None = Field(
|
||||
default=None,
|
||||
ge=1,
|
||||
le=3650, # Maximum 10 years; keeps tokens from being effectively permanent while allowing long-lived CI/CD tokens.
|
||||
description="Optional lifetime in days. If omitted the token never expires.",
|
||||
)
|
||||
|
||||
|
||||
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
|
||||
expires_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
|
||||
|
||||
expires_at = None
|
||||
if body.expires_in_days is not None:
|
||||
expires_at = datetime.now(timezone.utc) + timedelta(days=body.expires_in_days)
|
||||
|
||||
db_token = ApiToken(
|
||||
owner_id=owner_id,
|
||||
name=body.name,
|
||||
token_hash=token_hash_value,
|
||||
token_prefix=prefix,
|
||||
expires_at=expires_at,
|
||||
)
|
||||
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,
|
||||
"expires_at": db_token.expires_at,
|
||||
"token": plaintext,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/", response_model=list[TokenResponse])
|
||||
async def list_tokens(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List non-mobile API tokens for the authenticated user.
|
||||
|
||||
Mobile tokens (whose names start with ``"Mobile App"``) are excluded
|
||||
from this list; they are managed on the dedicated Devices page via
|
||||
``GET /api/api-tokens/mobile``.
|
||||
"""
|
||||
tokens = (
|
||||
db.query(ApiToken)
|
||||
.filter(
|
||||
ApiToken.owner_id == owner_id,
|
||||
~ApiToken.name.startswith(MOBILE_TOKEN_PREFIX),
|
||||
)
|
||||
.order_by(ApiToken.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_token_to_dict(t) for t in tokens]
|
||||
|
||||
|
||||
@router.get("/mobile", response_model=list[TokenResponse])
|
||||
async def list_mobile_tokens(
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List mobile API tokens for the authenticated user.
|
||||
|
||||
Returns tokens whose names start with ``"Mobile App"`` — these are
|
||||
created via the mobile SSO flow or QR code login.
|
||||
"""
|
||||
tokens = (
|
||||
db.query(ApiToken)
|
||||
.filter(
|
||||
ApiToken.owner_id == owner_id,
|
||||
ApiToken.name.startswith(MOBILE_TOKEN_PREFIX),
|
||||
)
|
||||
.order_by(ApiToken.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
return [_token_to_dict(t) for t in tokens]
|
||||
|
||||
|
||||
@router.delete("/{token_id}", status_code=status.HTTP_200_OK)
|
||||
async def revoke_or_delete_token(
|
||||
token_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Revoke or permanently delete an API token.
|
||||
|
||||
* **Active token** – soft-revoked: the row is kept for audit purposes
|
||||
but marked inactive with a ``revoked_at`` timestamp.
|
||||
* **Already-revoked token** – hard-deleted: the row is permanently
|
||||
removed from the database.
|
||||
"""
|
||||
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 db_token.is_active:
|
||||
# Soft-revoke the active token.
|
||||
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"}
|
||||
|
||||
# Hard-delete an already-revoked token.
|
||||
try:
|
||||
db.delete(db_token)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("API token permanently deleted: id=%s owner=%s", token_id, owner_id)
|
||||
return {"detail": "Token deleted"}
|
||||
|
||||
|
||||
@router.post("/{token_id}/reactivate", status_code=status.HTTP_200_OK, response_model=TokenResponse)
|
||||
async def reactivate_token(
|
||||
token_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Reactivate a previously revoked API token.
|
||||
|
||||
Clears the ``revoked_at`` timestamp and sets ``is_active`` back to
|
||||
``True``. The token can be used for authentication again immediately.
|
||||
If the token had an ``expires_at`` in the past the caller should
|
||||
consider re-creating a new token instead.
|
||||
"""
|
||||
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 db_token.is_active:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Token is already active")
|
||||
|
||||
try:
|
||||
db_token.is_active = True
|
||||
db_token.revoked_at = None
|
||||
db.commit()
|
||||
db.refresh(db_token)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("API token reactivated: id=%s owner=%s", token_id, owner_id)
|
||||
return _token_to_dict(db_token)
|
||||
@@ -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,311 @@
|
||||
"""API endpoints for Zapier / Make.com automation integration.
|
||||
|
||||
Provides a REST hooks subscription interface for outgoing triggers and
|
||||
incoming action endpoints that external automation platforms can call.
|
||||
|
||||
Outgoing triggers:
|
||||
External platforms subscribe to DocuElevate events via
|
||||
``POST /api/automation/hooks/subscribe``. When a subscribed event
|
||||
fires, DocuElevate POSTs a flat Zapier-compatible JSON payload to the
|
||||
registered ``target_url``.
|
||||
|
||||
Incoming actions:
|
||||
``POST /api/automation/actions/upload`` allows automation platforms to
|
||||
push documents into DocuElevate for processing.
|
||||
|
||||
Authentication:
|
||||
All endpoints require a valid API token via ``Authorization: Bearer``
|
||||
header.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.models import AutomationHook
|
||||
from app.utils.automation_hooks import SAMPLE_PAYLOADS
|
||||
from app.utils.webhook import VALID_EVENTS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/automation", tags=["automation"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Auth helper – require a valid API token (Bearer)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _require_api_user(request: Request) -> dict:
|
||||
"""Ensure the caller is authenticated via session or API token.
|
||||
|
||||
Raises:
|
||||
HTTPException: 401 if not authenticated, 403 if automation hooks are disabled.
|
||||
"""
|
||||
if not settings.automation_hooks_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Automation hooks are disabled",
|
||||
)
|
||||
|
||||
# Check for API-token user first (set by auth middleware)
|
||||
user = getattr(request.state, "api_token_user", None)
|
||||
if user:
|
||||
return user
|
||||
|
||||
# Fall back to session user
|
||||
user = request.session.get("user")
|
||||
if user:
|
||||
return user
|
||||
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Authentication required (Bearer token or session)",
|
||||
)
|
||||
|
||||
|
||||
AuthUser = Annotated[dict, Depends(_require_api_user)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class HookSubscribe(BaseModel):
|
||||
"""Schema for subscribing to automation hook events."""
|
||||
|
||||
target_url: str = Field(..., min_length=1, max_length=2048, description="URL to POST event payloads to")
|
||||
events: list[str] = Field(..., min_length=1, description="Event types to subscribe to")
|
||||
secret: str | None = Field(default=None, max_length=512, description="Optional HMAC-SHA256 signing secret")
|
||||
hook_type: str = Field(
|
||||
default="generic",
|
||||
max_length=50,
|
||||
description="Platform identifier (zapier, make, generic)",
|
||||
)
|
||||
description: str | None = Field(default=None, max_length=500, description="Optional human-readable label")
|
||||
|
||||
|
||||
class HookResponse(BaseModel):
|
||||
"""Schema returned when listing or creating hooks."""
|
||||
|
||||
id: int
|
||||
target_url: str
|
||||
events: list[str]
|
||||
is_active: bool
|
||||
hook_type: str
|
||||
description: str | None
|
||||
has_secret: bool
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class ActionUploadResponse(BaseModel):
|
||||
"""Response after an automation action uploads a document."""
|
||||
|
||||
status: str
|
||||
filename: str
|
||||
task_id: str | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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: {', '.join(sorted(VALID_EVENTS))}",
|
||||
)
|
||||
|
||||
|
||||
def _hook_to_response(hook: AutomationHook) -> dict[str, Any]:
|
||||
"""Convert a DB model instance to a response dict."""
|
||||
try:
|
||||
events = json.loads(hook.events)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
events = []
|
||||
return {
|
||||
"id": hook.id,
|
||||
"target_url": hook.target_url,
|
||||
"events": events,
|
||||
"is_active": hook.is_active,
|
||||
"hook_type": hook.hook_type,
|
||||
"description": hook.description,
|
||||
"has_secret": hook.secret is not None and len(hook.secret) > 0,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Outgoing triggers – REST hooks subscription endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/hooks/subscribe",
|
||||
status_code=status.HTTP_201_CREATED,
|
||||
summary="Subscribe to automation events (REST hooks)",
|
||||
)
|
||||
def subscribe_hook(body: HookSubscribe, db: DbSession, user: AuthUser) -> dict[str, Any]:
|
||||
"""Register a new automation hook subscription.
|
||||
|
||||
Zapier and Make.com call this endpoint to subscribe to DocuElevate
|
||||
events. When an event fires, a flat JSON payload is POSTed to
|
||||
``target_url``.
|
||||
"""
|
||||
_validate_events(body.events)
|
||||
|
||||
hook = AutomationHook(
|
||||
target_url=body.target_url,
|
||||
secret=body.secret,
|
||||
events=json.dumps(sorted(body.events)),
|
||||
is_active=True,
|
||||
hook_type=body.hook_type or "generic",
|
||||
description=body.description,
|
||||
)
|
||||
try:
|
||||
db.add(hook)
|
||||
db.commit()
|
||||
db.refresh(hook)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Automation hook %d created (type=%s) for events %s", hook.id, hook.hook_type, body.events)
|
||||
return _hook_to_response(hook)
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/hooks/{hook_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
summary="Unsubscribe an automation hook",
|
||||
)
|
||||
def unsubscribe_hook(hook_id: int, db: DbSession, user: AuthUser) -> None:
|
||||
"""Remove an automation hook subscription.
|
||||
|
||||
Zapier calls this endpoint when a Zap is turned off or deleted.
|
||||
"""
|
||||
hook = db.query(AutomationHook).filter(AutomationHook.id == hook_id).first()
|
||||
if not hook:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Hook not found")
|
||||
|
||||
try:
|
||||
db.delete(hook)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Automation hook %d deleted", hook_id)
|
||||
|
||||
|
||||
@router.get("/hooks", summary="List automation hook subscriptions")
|
||||
def list_hooks(db: DbSession, user: AuthUser) -> list[dict[str, Any]]:
|
||||
"""Return all active automation hook subscriptions."""
|
||||
hooks = db.query(AutomationHook).order_by(AutomationHook.id).all()
|
||||
return [_hook_to_response(h) for h in hooks]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Outgoing triggers – sample data for Zapier field mapping
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/triggers/sample/{event}", summary="Get sample trigger data")
|
||||
def get_trigger_sample(event: str, user: AuthUser) -> list[dict[str, Any]]:
|
||||
"""Return sample payload data for the given event type.
|
||||
|
||||
Zapier uses this during Zap setup to discover available fields and
|
||||
provide a mapping interface. The response is wrapped in an array
|
||||
as Zapier expects.
|
||||
"""
|
||||
if event not in VALID_EVENTS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail=f"Unknown event: {event}. Valid: {', '.join(sorted(VALID_EVENTS))}",
|
||||
)
|
||||
|
||||
sample = SAMPLE_PAYLOADS.get(event, {"id": "evt_sample", "event": event, "timestamp": 0})
|
||||
return [sample]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Outgoing triggers – list valid events
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/events", summary="List valid automation event types")
|
||||
def list_events(user: AuthUser) -> list[str]:
|
||||
"""Return the list of valid event types that automation hooks can subscribe to."""
|
||||
return sorted(VALID_EVENTS)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Incoming actions – endpoints that Zapier / Make.com can call
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/actions/upload", summary="Upload a document (incoming action)")
|
||||
def action_upload(
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
user: AuthUser,
|
||||
file: UploadFile = File(...),
|
||||
) -> dict[str, Any]:
|
||||
"""Accept a document upload from an automation platform.
|
||||
|
||||
This endpoint allows Zapier or Make.com to push a document into
|
||||
DocuElevate for processing. The file is saved to the work directory
|
||||
and a background processing task is queued.
|
||||
"""
|
||||
if not file.filename:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
|
||||
|
||||
# Sanitise filename to prevent path traversal attacks
|
||||
safe_filename = os.path.basename(file.filename)
|
||||
if not safe_filename:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Filename is required")
|
||||
|
||||
owner_id = user.get("preferred_username") or user.get("email") or user.get("id", "automation")
|
||||
workdir = settings.workdir or tempfile.gettempdir()
|
||||
upload_dir = os.path.join(workdir, "uploads")
|
||||
os.makedirs(upload_dir, exist_ok=True)
|
||||
|
||||
dest_path = os.path.join(upload_dir, safe_filename)
|
||||
try:
|
||||
contents = file.file.read()
|
||||
with open(dest_path, "wb") as f:
|
||||
f.write(contents)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to save uploaded file: %s", exc)
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to save file")
|
||||
|
||||
# Queue background processing
|
||||
task_id = None
|
||||
try:
|
||||
from app.tasks.process_document import process_document
|
||||
|
||||
result = process_document.delay(dest_path, owner_id)
|
||||
task_id = result.id
|
||||
logger.info("Automation upload queued: file=%s, task=%s, owner=%s", safe_filename, task_id, owner_id)
|
||||
except Exception as exc:
|
||||
logger.warning("Could not queue processing task (Celery may be unavailable): %s", exc)
|
||||
|
||||
return {
|
||||
"status": "accepted",
|
||||
"filename": safe_filename,
|
||||
"task_id": task_id,
|
||||
}
|
||||
@@ -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,618 @@
|
||||
"""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 plan %s",
|
||||
checkout_session.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")
|
||||
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(request, "billing_success.html")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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,325 @@
|
||||
"""Classification Rules API endpoints.
|
||||
|
||||
Provides CRUD operations for managing custom document classification rules.
|
||||
System-wide rules (``owner_id IS NULL``) can only be managed by admins.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
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.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import ClassificationRuleModel
|
||||
from app.utils.classification_rules import (
|
||||
BUILTIN_CATEGORIES,
|
||||
RULE_TYPE_CONTENT,
|
||||
RULE_TYPE_FILENAME,
|
||||
RULE_TYPE_METADATA,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/classification-rules", tags=["classification"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
_VALID_RULE_TYPES = {RULE_TYPE_FILENAME, RULE_TYPE_CONTENT, RULE_TYPE_METADATA}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_user_id(request: Request) -> str:
|
||||
"""Extract the user identifier from the request session."""
|
||||
user = getattr(request.state, "user", None)
|
||||
if user and hasattr(user, "get"):
|
||||
return user.get("sub") or user.get("email") or "anonymous"
|
||||
return "anonymous"
|
||||
|
||||
|
||||
def _is_admin(request: Request) -> bool:
|
||||
"""Check whether the current user is an admin."""
|
||||
user = getattr(request.state, "user", None)
|
||||
if user and hasattr(user, "get"):
|
||||
groups = user.get("groups", [])
|
||||
return "admin" in groups or "Admin" in groups
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pydantic schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RuleCreate(BaseModel):
|
||||
"""Schema for creating a classification rule."""
|
||||
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
category: str = Field(..., min_length=1, max_length=100)
|
||||
rule_type: str = Field(..., description="One of: filename_pattern, content_keyword, metadata_match")
|
||||
pattern: str = Field(..., min_length=1, max_length=1000)
|
||||
priority: int = Field(default=0, ge=0, le=1000)
|
||||
case_sensitive: bool = False
|
||||
enabled: bool = True
|
||||
|
||||
|
||||
class RuleUpdate(BaseModel):
|
||||
"""Schema for updating a classification rule."""
|
||||
|
||||
name: str | None = Field(default=None, min_length=1, max_length=255)
|
||||
category: str | None = Field(default=None, min_length=1, max_length=100)
|
||||
rule_type: str | None = Field(default=None)
|
||||
pattern: str | None = Field(default=None, min_length=1, max_length=1000)
|
||||
priority: int | None = Field(default=None, ge=0, le=1000)
|
||||
case_sensitive: bool | None = None
|
||||
enabled: bool | None = None
|
||||
|
||||
|
||||
class RuleResponse(BaseModel):
|
||||
"""Schema for a classification rule response."""
|
||||
|
||||
id: int
|
||||
owner_id: str | None
|
||||
name: str
|
||||
category: str
|
||||
rule_type: str
|
||||
pattern: str
|
||||
priority: int
|
||||
case_sensitive: bool
|
||||
enabled: bool
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/categories")
|
||||
@require_login
|
||||
async def list_categories(request: Request) -> dict[str, str]:
|
||||
"""Return all built-in classification categories.
|
||||
|
||||
Custom categories created via rules are not included here; they are
|
||||
discovered dynamically when rules are evaluated.
|
||||
"""
|
||||
return BUILTIN_CATEGORIES
|
||||
|
||||
|
||||
@router.get("/rule-types")
|
||||
@require_login
|
||||
async def list_rule_types(request: Request) -> list[dict[str, str]]:
|
||||
"""Return the supported rule types with descriptions."""
|
||||
return [
|
||||
{
|
||||
"type": RULE_TYPE_FILENAME,
|
||||
"label": "Filename Pattern",
|
||||
"description": "Regex pattern matched against the original filename.",
|
||||
},
|
||||
{
|
||||
"type": RULE_TYPE_CONTENT,
|
||||
"label": "Content Keyword",
|
||||
"description": "Pipe-separated keywords matched against the OCR text.",
|
||||
},
|
||||
{
|
||||
"type": RULE_TYPE_METADATA,
|
||||
"label": "Metadata Match",
|
||||
"description": "field=value pattern matched against existing AI metadata.",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
@router.get("/")
|
||||
@require_login
|
||||
async def list_rules(request: Request, db: DbSession) -> list[dict[str, Any]]:
|
||||
"""List classification rules visible to the current user.
|
||||
|
||||
Returns both system rules (``owner_id IS NULL``) and the user's own rules.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
rules = (
|
||||
db.query(ClassificationRuleModel)
|
||||
.filter((ClassificationRuleModel.owner_id.is_(None)) | (ClassificationRuleModel.owner_id == user_id))
|
||||
.order_by(ClassificationRuleModel.priority.desc(), ClassificationRuleModel.id)
|
||||
.all()
|
||||
)
|
||||
return [
|
||||
{
|
||||
"id": r.id,
|
||||
"owner_id": r.owner_id,
|
||||
"name": r.name,
|
||||
"category": r.category,
|
||||
"rule_type": r.rule_type,
|
||||
"pattern": r.pattern,
|
||||
"priority": r.priority,
|
||||
"case_sensitive": r.case_sensitive,
|
||||
"enabled": r.enabled,
|
||||
}
|
||||
for r in rules
|
||||
]
|
||||
|
||||
|
||||
@router.post("/", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
async def create_rule(request: Request, body: RuleCreate, db: DbSession) -> dict[str, Any]:
|
||||
"""Create a new custom classification rule.
|
||||
|
||||
The rule is owned by the current user. Admins may create system-wide
|
||||
rules by setting ``owner_id`` to ``null`` (not yet exposed).
|
||||
"""
|
||||
if body.rule_type not in _VALID_RULE_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid rule_type. Must be one of: {', '.join(sorted(_VALID_RULE_TYPES))}",
|
||||
)
|
||||
|
||||
user_id = _get_user_id(request)
|
||||
|
||||
# Check for duplicate name within the user's scope
|
||||
existing = (
|
||||
db.query(ClassificationRuleModel)
|
||||
.filter(ClassificationRuleModel.owner_id == user_id, ClassificationRuleModel.name == body.name)
|
||||
.first()
|
||||
)
|
||||
if existing:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"A rule named '{body.name}' already exists.",
|
||||
)
|
||||
|
||||
rule = ClassificationRuleModel(
|
||||
owner_id=user_id,
|
||||
name=body.name,
|
||||
category=body.category,
|
||||
rule_type=body.rule_type,
|
||||
pattern=body.pattern,
|
||||
priority=body.priority,
|
||||
case_sensitive=body.case_sensitive,
|
||||
enabled=body.enabled,
|
||||
)
|
||||
try:
|
||||
db.add(rule)
|
||||
db.commit()
|
||||
db.refresh(rule)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Classification rule created: id=%s, user=%s", rule.id, user_id)
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"category": rule.category,
|
||||
"rule_type": rule.rule_type,
|
||||
"pattern": rule.pattern,
|
||||
"priority": rule.priority,
|
||||
"case_sensitive": rule.case_sensitive,
|
||||
"enabled": rule.enabled,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/{rule_id}")
|
||||
@require_login
|
||||
async def get_rule(request: Request, rule_id: int, db: DbSession) -> dict[str, Any]:
|
||||
"""Get a single classification rule by ID."""
|
||||
user_id = _get_user_id(request)
|
||||
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
|
||||
if rule is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
# Users can see system rules and their own rules
|
||||
if rule.owner_id is not None and rule.owner_id != user_id and not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"category": rule.category,
|
||||
"rule_type": rule.rule_type,
|
||||
"pattern": rule.pattern,
|
||||
"priority": rule.priority,
|
||||
"case_sensitive": rule.case_sensitive,
|
||||
"enabled": rule.enabled,
|
||||
}
|
||||
|
||||
|
||||
@router.put("/{rule_id}")
|
||||
@require_login
|
||||
async def update_rule(request: Request, rule_id: int, body: RuleUpdate, db: DbSession) -> dict[str, Any]:
|
||||
"""Update an existing classification rule.
|
||||
|
||||
Users can only update their own rules. Admins can update any rule.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
|
||||
if rule is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
if rule.owner_id != user_id and not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot modify this rule")
|
||||
|
||||
if body.rule_type is not None and body.rule_type not in _VALID_RULE_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Invalid rule_type. Must be one of: {', '.join(sorted(_VALID_RULE_TYPES))}",
|
||||
)
|
||||
|
||||
update_data = body.model_dump(exclude_unset=True)
|
||||
for field_name, value in update_data.items():
|
||||
setattr(rule, field_name, value)
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(rule)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Classification rule updated: id=%s, user=%s", rule.id, user_id)
|
||||
return {
|
||||
"id": rule.id,
|
||||
"owner_id": rule.owner_id,
|
||||
"name": rule.name,
|
||||
"category": rule.category,
|
||||
"rule_type": rule.rule_type,
|
||||
"pattern": rule.pattern,
|
||||
"priority": rule.priority,
|
||||
"case_sensitive": rule.case_sensitive,
|
||||
"enabled": rule.enabled,
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{rule_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
async def delete_rule(request: Request, rule_id: int, db: DbSession) -> None:
|
||||
"""Delete a classification rule.
|
||||
|
||||
Users can only delete their own rules. Admins can delete any rule.
|
||||
"""
|
||||
user_id = _get_user_id(request)
|
||||
rule = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.id == rule_id).first()
|
||||
if rule is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Rule not found")
|
||||
|
||||
if rule.owner_id != user_id and not _is_admin(request):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Cannot delete this rule")
|
||||
|
||||
try:
|
||||
db.delete(rule)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
|
||||
logger.info("Classification rule deleted: id=%s, user=%s", rule_id, user_id)
|
||||
@@ -0,0 +1,751 @@
|
||||
"""Document comments and annotations API endpoints.
|
||||
|
||||
Provides CRUD operations for threaded comments on documents,
|
||||
text annotations on PDF pages, and a list of mentionable users
|
||||
for the @mention feature.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import get_current_user_id, require_login
|
||||
from app.database import get_db
|
||||
from app.models import (
|
||||
FILE_SHARE_ROLE_VIEWER,
|
||||
DocumentAnnotation,
|
||||
DocumentComment,
|
||||
FileRecord,
|
||||
FileShare,
|
||||
UserProfile,
|
||||
)
|
||||
from app.utils.user_scope import get_current_owner_id, has_file_role
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(tags=["comments"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# Constraints
|
||||
MAX_COMMENT_BODY_LENGTH = 10_000
|
||||
MAX_ANNOTATION_CONTENT_LENGTH = 5_000
|
||||
|
||||
# Allowed annotation types
|
||||
ALLOWED_ANNOTATION_TYPES = frozenset({"note", "highlight", "underline", "strikethrough"})
|
||||
|
||||
# Simple pattern for @mentions – matches @username tokens inside comment body
|
||||
_MENTION_PATTERN = re.compile(r"@([\w.\-]+)")
|
||||
|
||||
|
||||
def _extract_mentions(body: str) -> list[str]:
|
||||
"""Extract unique @mentioned usernames from a comment body.
|
||||
|
||||
Args:
|
||||
body: The raw comment text.
|
||||
|
||||
Returns:
|
||||
A deduplicated list of mentioned usernames (without the ``@`` prefix).
|
||||
"""
|
||||
return list(dict.fromkeys(_MENTION_PATTERN.findall(body)))
|
||||
|
||||
|
||||
def _serialize_comment(c: DocumentComment) -> dict[str, Any]:
|
||||
"""Serialize a DocumentComment to a JSON-friendly dict.
|
||||
|
||||
Args:
|
||||
c: The comment model instance.
|
||||
|
||||
Returns:
|
||||
A dictionary representation of the comment.
|
||||
"""
|
||||
mentions: list[str] = []
|
||||
if c.mentions:
|
||||
try:
|
||||
mentions = json.loads(c.mentions)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
return {
|
||||
"id": c.id,
|
||||
"file_id": c.file_id,
|
||||
"user_id": c.user_id,
|
||||
"parent_id": c.parent_id,
|
||||
"body": c.body,
|
||||
"mentions": mentions,
|
||||
"is_resolved": c.is_resolved,
|
||||
"created_at": c.created_at.isoformat() if c.created_at else None,
|
||||
"updated_at": c.updated_at.isoformat() if c.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _serialize_annotation(a: DocumentAnnotation) -> dict[str, Any]:
|
||||
"""Serialize a DocumentAnnotation to a JSON-friendly dict.
|
||||
|
||||
Args:
|
||||
a: The annotation model instance.
|
||||
|
||||
Returns:
|
||||
A dictionary representation of the annotation.
|
||||
"""
|
||||
return {
|
||||
"id": a.id,
|
||||
"file_id": a.file_id,
|
||||
"user_id": a.user_id,
|
||||
"page": a.page,
|
||||
"x": a.x,
|
||||
"y": a.y,
|
||||
"width": a.width,
|
||||
"height": a.height,
|
||||
"content": a.content,
|
||||
"annotation_type": a.annotation_type,
|
||||
"color": a.color,
|
||||
"created_at": a.created_at.isoformat() if a.created_at else None,
|
||||
"updated_at": a.updated_at.isoformat() if a.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _build_thread_tree(comments: list[DocumentComment]) -> list[dict[str, Any]]:
|
||||
"""Organize a flat list of comments into a threaded tree structure.
|
||||
|
||||
Top-level comments (``parent_id is None``) appear as root nodes.
|
||||
Replies are nested inside their parent's ``replies`` list.
|
||||
|
||||
Args:
|
||||
comments: All comments for a given document, ordered by ``created_at``.
|
||||
|
||||
Returns:
|
||||
A list of root-level comment dicts, each with a ``replies`` key.
|
||||
"""
|
||||
by_id: dict[int, dict[str, Any]] = {}
|
||||
roots: list[dict[str, Any]] = []
|
||||
|
||||
for c in comments:
|
||||
node = _serialize_comment(c)
|
||||
node["replies"] = []
|
||||
by_id[c.id] = node
|
||||
|
||||
for c in comments:
|
||||
node = by_id[c.id]
|
||||
if c.parent_id and c.parent_id in by_id:
|
||||
by_id[c.parent_id]["replies"].append(node)
|
||||
else:
|
||||
roots.append(node)
|
||||
|
||||
return roots
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Comments endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/comments")
|
||||
@require_login
|
||||
def list_comments(request: Request, file_id: int, db: DbSession):
|
||||
"""List all comments for a document, organized into threads.
|
||||
|
||||
Returns a threaded tree where top-level comments contain nested
|
||||
``replies``. Requires at least viewer access.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
|
||||
Returns:
|
||||
A dict with ``file_id``, ``comments`` (threaded), and ``total``.
|
||||
"""
|
||||
user_id = get_current_owner_id(request)
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(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")
|
||||
|
||||
if not is_admin and not has_file_role(file_record, user_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
comments = (
|
||||
db.query(DocumentComment).filter(DocumentComment.file_id == file_id).order_by(DocumentComment.created_at).all()
|
||||
)
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"comments": _build_thread_tree(comments),
|
||||
"total": len(comments),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/files/{file_id}/comments", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
def create_comment(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
body: str = Body(..., embed=True),
|
||||
parent_id: int | None = Body(None, embed=True),
|
||||
):
|
||||
"""Create a new comment on a document.
|
||||
|
||||
Automatically extracts @mentions from the comment body and stores
|
||||
them for later notification or UI highlighting. When multi-user
|
||||
mode is enabled, any mentioned user that does not already have
|
||||
access to the document is automatically granted ``viewer`` access by
|
||||
the file owner so they can read the file and continue the discussion.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document to comment on.
|
||||
|
||||
Request body (JSON):
|
||||
body: Comment text (required, max 10 000 characters).
|
||||
parent_id: ID of the parent comment for threaded replies (optional).
|
||||
|
||||
Returns:
|
||||
The created comment object.
|
||||
"""
|
||||
user_id = get_current_user_id(request)
|
||||
owner_id = get_current_owner_id(request)
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(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")
|
||||
|
||||
if not is_admin and not has_file_role(file_record, owner_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
if not isinstance(body, str) or not body.strip():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="body is required and must be non-empty",
|
||||
)
|
||||
body = body.strip()
|
||||
if len(body) > MAX_COMMENT_BODY_LENGTH:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"body must be at most {MAX_COMMENT_BODY_LENGTH} characters",
|
||||
)
|
||||
|
||||
if parent_id is not None:
|
||||
parent = (
|
||||
db.query(DocumentComment)
|
||||
.filter(DocumentComment.id == parent_id, DocumentComment.file_id == file_id)
|
||||
.first()
|
||||
)
|
||||
if not parent:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="Parent comment not found",
|
||||
)
|
||||
|
||||
mentions = _extract_mentions(body)
|
||||
|
||||
comment = DocumentComment(
|
||||
file_id=file_id,
|
||||
user_id=user_id,
|
||||
parent_id=parent_id,
|
||||
body=body,
|
||||
mentions=json.dumps(mentions) if mentions else None,
|
||||
)
|
||||
|
||||
try:
|
||||
db.add(comment)
|
||||
db.flush() # write comment so we can get its id before committing
|
||||
|
||||
# Auto-share the file with mentioned users that don't have access yet.
|
||||
# Only do this in multi-user mode and only when the file has an owner
|
||||
# (unowned files are already visible to all authenticated users).
|
||||
if mentions and file_record.owner_id is not None:
|
||||
from app.config import settings as _settings
|
||||
|
||||
if _settings.multi_user_enabled:
|
||||
for mentioned_user in mentions:
|
||||
# Skip the file owner (already has full access) and the commenter
|
||||
# themselves (they already have access to be posting a comment).
|
||||
if mentioned_user in {file_record.owner_id, owner_id}:
|
||||
continue
|
||||
existing_share = (
|
||||
db.query(FileShare)
|
||||
.filter(
|
||||
FileShare.file_id == file_id,
|
||||
FileShare.shared_with_user_id == mentioned_user,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if not existing_share:
|
||||
auto_share = FileShare(
|
||||
file_id=file_id,
|
||||
owner_id=file_record.owner_id,
|
||||
shared_with_user_id=mentioned_user,
|
||||
role=FILE_SHARE_ROLE_VIEWER,
|
||||
)
|
||||
db.add(auto_share)
|
||||
logger.info(
|
||||
"Auto-shared file_id=%s with mentioned user=%s as viewer",
|
||||
file_id,
|
||||
mentioned_user,
|
||||
)
|
||||
|
||||
db.commit()
|
||||
db.refresh(comment)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create comment on file_id=%s", file_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to create comment",
|
||||
)
|
||||
|
||||
logger.info("Comment created: id=%s, file_id=%s, user=%s", comment.id, file_id, user_id)
|
||||
return _serialize_comment(comment)
|
||||
|
||||
|
||||
@router.put("/files/{file_id}/comments/{comment_id}")
|
||||
@require_login
|
||||
def update_comment(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
comment_id: int,
|
||||
db: DbSession,
|
||||
body: str = Body(..., embed=True),
|
||||
):
|
||||
"""Update the body of an existing comment.
|
||||
|
||||
Only the comment author may update the comment. Mentions are
|
||||
re-extracted from the updated body.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
comment_id: The ID of the comment to update.
|
||||
|
||||
Request body (JSON):
|
||||
body: New comment text (required).
|
||||
|
||||
Returns:
|
||||
The updated comment object.
|
||||
"""
|
||||
user_id = get_current_user_id(request)
|
||||
|
||||
comment = (
|
||||
db.query(DocumentComment).filter(DocumentComment.id == comment_id, DocumentComment.file_id == file_id).first()
|
||||
)
|
||||
if not comment:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Comment not found")
|
||||
|
||||
if comment.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only edit your own comments")
|
||||
|
||||
if not isinstance(body, str) or not body.strip():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="body is required and must be non-empty",
|
||||
)
|
||||
body = body.strip()
|
||||
if len(body) > MAX_COMMENT_BODY_LENGTH:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"body must be at most {MAX_COMMENT_BODY_LENGTH} characters",
|
||||
)
|
||||
|
||||
mentions = _extract_mentions(body)
|
||||
comment.body = body
|
||||
comment.mentions = json.dumps(mentions) if mentions else None
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(comment)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to update comment id=%s", comment_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update comment",
|
||||
)
|
||||
|
||||
logger.info("Comment updated: id=%s, user=%s", comment_id, user_id)
|
||||
return _serialize_comment(comment)
|
||||
|
||||
|
||||
@router.delete("/files/{file_id}/comments/{comment_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
def delete_comment(request: Request, file_id: int, comment_id: int, db: DbSession):
|
||||
"""Delete a comment.
|
||||
|
||||
Only the comment author may delete the comment. Replies to the
|
||||
deleted comment are **not** removed — they become orphaned root
|
||||
comments so that conversation context is preserved.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
comment_id: The ID of the comment to delete.
|
||||
"""
|
||||
user_id = get_current_user_id(request)
|
||||
|
||||
comment = (
|
||||
db.query(DocumentComment).filter(DocumentComment.id == comment_id, DocumentComment.file_id == file_id).first()
|
||||
)
|
||||
if not comment:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Comment not found")
|
||||
|
||||
if comment.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only delete your own comments")
|
||||
|
||||
try:
|
||||
db.delete(comment)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to delete comment id=%s", comment_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to delete comment",
|
||||
)
|
||||
|
||||
logger.info("Comment deleted: id=%s, user=%s", comment_id, user_id)
|
||||
|
||||
|
||||
@router.patch("/files/{file_id}/comments/{comment_id}/resolve")
|
||||
@require_login
|
||||
def resolve_comment(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
comment_id: int,
|
||||
db: DbSession,
|
||||
is_resolved: bool = Body(..., embed=True),
|
||||
):
|
||||
"""Mark a top-level comment thread as resolved or unresolved.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
comment_id: The ID of the comment to resolve / unresolve.
|
||||
|
||||
Request body (JSON):
|
||||
is_resolved: ``true`` to resolve, ``false`` to unresolve.
|
||||
|
||||
Returns:
|
||||
The updated comment object.
|
||||
"""
|
||||
comment = (
|
||||
db.query(DocumentComment).filter(DocumentComment.id == comment_id, DocumentComment.file_id == file_id).first()
|
||||
)
|
||||
if not comment:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Comment not found")
|
||||
|
||||
comment.is_resolved = is_resolved
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(comment)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to resolve comment id=%s", comment_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update comment",
|
||||
)
|
||||
|
||||
logger.info("Comment %s: id=%s", "resolved" if is_resolved else "unresolved", comment_id)
|
||||
return _serialize_comment(comment)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Annotations endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/annotations")
|
||||
@require_login
|
||||
def list_annotations(request: Request, file_id: int, db: DbSession):
|
||||
"""List all annotations for a document.
|
||||
|
||||
Requires at least viewer access.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
|
||||
Returns:
|
||||
A dict with ``file_id``, ``annotations``, and ``total``.
|
||||
"""
|
||||
user_id = get_current_owner_id(request)
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(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")
|
||||
|
||||
if not is_admin and not has_file_role(file_record, user_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
annotations = (
|
||||
db.query(DocumentAnnotation)
|
||||
.filter(DocumentAnnotation.file_id == file_id)
|
||||
.order_by(DocumentAnnotation.page, DocumentAnnotation.created_at)
|
||||
.all()
|
||||
)
|
||||
|
||||
return {
|
||||
"file_id": file_id,
|
||||
"annotations": [_serialize_annotation(a) for a in annotations],
|
||||
"total": len(annotations),
|
||||
}
|
||||
|
||||
|
||||
@router.post("/files/{file_id}/annotations", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
def create_annotation(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
page: int = Body(..., embed=True),
|
||||
x: float = Body(..., embed=True),
|
||||
y: float = Body(..., embed=True),
|
||||
content: str = Body(..., embed=True),
|
||||
width: float = Body(0, embed=True),
|
||||
height: float = Body(0, embed=True),
|
||||
annotation_type: str = Body("note", embed=True),
|
||||
color: str | None = Body(None, embed=True),
|
||||
):
|
||||
"""Create a new annotation on a PDF page.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
|
||||
Request body (JSON):
|
||||
page: Page number (1-based, required).
|
||||
x: Horizontal position on the page (required).
|
||||
y: Vertical position on the page (required).
|
||||
content: Annotation text (required, max 5 000 characters).
|
||||
width: Width of the annotation bounding box (default 0).
|
||||
height: Height of the annotation bounding box (default 0).
|
||||
annotation_type: One of ``note``, ``highlight``, ``underline``,
|
||||
``strikethrough`` (default ``note``).
|
||||
color: Optional CSS colour string (e.g. ``#ff0000``).
|
||||
|
||||
Returns:
|
||||
The created annotation object.
|
||||
"""
|
||||
user_id = get_current_user_id(request)
|
||||
owner_id = get_current_owner_id(request)
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(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")
|
||||
|
||||
if not is_admin and not has_file_role(file_record, owner_id, db, minimum_role=FILE_SHARE_ROLE_VIEWER):
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
if not isinstance(content, str) or not content.strip():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="content is required and must be non-empty",
|
||||
)
|
||||
content = content.strip()
|
||||
if len(content) > MAX_ANNOTATION_CONTENT_LENGTH:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"content must be at most {MAX_ANNOTATION_CONTENT_LENGTH} characters",
|
||||
)
|
||||
|
||||
if page < 1:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="page must be >= 1",
|
||||
)
|
||||
|
||||
if annotation_type not in ALLOWED_ANNOTATION_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"annotation_type must be one of: {', '.join(sorted(ALLOWED_ANNOTATION_TYPES))}",
|
||||
)
|
||||
|
||||
annotation = DocumentAnnotation(
|
||||
file_id=file_id,
|
||||
user_id=user_id,
|
||||
page=page,
|
||||
x=x,
|
||||
y=y,
|
||||
width=width,
|
||||
height=height,
|
||||
content=content,
|
||||
annotation_type=annotation_type,
|
||||
color=color,
|
||||
)
|
||||
|
||||
try:
|
||||
db.add(annotation)
|
||||
db.commit()
|
||||
db.refresh(annotation)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create annotation on file_id=%s", file_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to create annotation",
|
||||
)
|
||||
|
||||
logger.info("Annotation created: id=%s, file_id=%s, user=%s", annotation.id, file_id, user_id)
|
||||
return _serialize_annotation(annotation)
|
||||
|
||||
|
||||
@router.put("/files/{file_id}/annotations/{annotation_id}")
|
||||
@require_login
|
||||
def update_annotation(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
annotation_id: int,
|
||||
db: DbSession,
|
||||
content: str | None = Body(None, embed=True),
|
||||
x: float | None = Body(None, embed=True),
|
||||
y: float | None = Body(None, embed=True),
|
||||
width: float | None = Body(None, embed=True),
|
||||
height: float | None = Body(None, embed=True),
|
||||
annotation_type: str | None = Body(None, embed=True),
|
||||
color: str | None = Body(None, embed=True),
|
||||
):
|
||||
"""Update an existing annotation.
|
||||
|
||||
Only the annotation author may update the annotation.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
annotation_id: The ID of the annotation to update.
|
||||
|
||||
Request body (JSON):
|
||||
Any subset of ``content``, ``x``, ``y``, ``width``, ``height``,
|
||||
``annotation_type``, and ``color``.
|
||||
|
||||
Returns:
|
||||
The updated annotation object.
|
||||
"""
|
||||
user_id = get_current_user_id(request)
|
||||
|
||||
annotation = (
|
||||
db.query(DocumentAnnotation)
|
||||
.filter(DocumentAnnotation.id == annotation_id, DocumentAnnotation.file_id == file_id)
|
||||
.first()
|
||||
)
|
||||
if not annotation:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Annotation not found")
|
||||
|
||||
if annotation.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only edit your own annotations")
|
||||
|
||||
if content is not None:
|
||||
content = content.strip() if isinstance(content, str) else ""
|
||||
if not content:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="content must be non-empty",
|
||||
)
|
||||
if len(content) > MAX_ANNOTATION_CONTENT_LENGTH:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"content must be at most {MAX_ANNOTATION_CONTENT_LENGTH} characters",
|
||||
)
|
||||
annotation.content = content
|
||||
|
||||
if x is not None:
|
||||
annotation.x = x
|
||||
if y is not None:
|
||||
annotation.y = y
|
||||
if width is not None:
|
||||
annotation.width = width
|
||||
if height is not None:
|
||||
annotation.height = height
|
||||
if annotation_type is not None:
|
||||
if annotation_type not in ALLOWED_ANNOTATION_TYPES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"annotation_type must be one of: {', '.join(sorted(ALLOWED_ANNOTATION_TYPES))}",
|
||||
)
|
||||
annotation.annotation_type = annotation_type
|
||||
if color is not None:
|
||||
annotation.color = color
|
||||
|
||||
try:
|
||||
db.commit()
|
||||
db.refresh(annotation)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to update annotation id=%s", annotation_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update annotation",
|
||||
)
|
||||
|
||||
logger.info("Annotation updated: id=%s, user=%s", annotation_id, user_id)
|
||||
return _serialize_annotation(annotation)
|
||||
|
||||
|
||||
@router.delete("/files/{file_id}/annotations/{annotation_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
def delete_annotation(request: Request, file_id: int, annotation_id: int, db: DbSession):
|
||||
"""Delete an annotation.
|
||||
|
||||
Only the annotation author may delete the annotation.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
annotation_id: The ID of the annotation to delete.
|
||||
"""
|
||||
user_id = get_current_user_id(request)
|
||||
|
||||
annotation = (
|
||||
db.query(DocumentAnnotation)
|
||||
.filter(DocumentAnnotation.id == annotation_id, DocumentAnnotation.file_id == file_id)
|
||||
.first()
|
||||
)
|
||||
if not annotation:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Annotation not found")
|
||||
|
||||
if annotation.user_id != user_id:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="You can only delete your own annotations")
|
||||
|
||||
try:
|
||||
db.delete(annotation)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to delete annotation id=%s", annotation_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to delete annotation",
|
||||
)
|
||||
|
||||
logger.info("Annotation deleted: id=%s, user=%s", annotation_id, user_id)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Mentionable users endpoint
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/users/mentionable")
|
||||
@require_login
|
||||
def list_mentionable_users(request: Request, db: DbSession):
|
||||
"""List users that can be @mentioned in comments.
|
||||
|
||||
Returns all user profiles that are not blocked, sorted by
|
||||
``display_name``.
|
||||
|
||||
Returns:
|
||||
A list of ``{user_id, display_name}`` objects.
|
||||
"""
|
||||
profiles = db.query(UserProfile).filter(UserProfile.is_blocked.is_(False)).order_by(UserProfile.display_name).all()
|
||||
|
||||
return [
|
||||
{
|
||||
"user_id": p.user_id,
|
||||
"display_name": p.display_name or p.user_id,
|
||||
}
|
||||
for p in profiles
|
||||
]
|
||||
@@ -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,18 +2,172 @@
|
||||
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()
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unauthenticated probe endpoints for Kubernetes liveness / readiness checks.
|
||||
# These intentionally skip authentication so that kubelet can reach them
|
||||
# without credentials. They live under /diagnostic/healthz/* so that the
|
||||
# existing authenticated /diagnostic/health endpoint is unaffected.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/diagnostic/healthz/live")
|
||||
async def liveness_probe() -> JSONResponse:
|
||||
"""Lightweight liveness probe for Kubernetes.
|
||||
|
||||
Returns **200 OK** as long as the process is running. Kubernetes uses
|
||||
this to decide whether to *restart* the container — it should therefore
|
||||
be as cheap as possible and **never** check external dependencies.
|
||||
|
||||
**Authentication:** None (designed for kubelet probes).
|
||||
"""
|
||||
return JSONResponse(content={"status": "ok"}, status_code=200)
|
||||
|
||||
|
||||
@router.get("/diagnostic/healthz/ready")
|
||||
async def readiness_probe() -> JSONResponse:
|
||||
"""Readiness probe for Kubernetes.
|
||||
|
||||
Verifies that the application can serve traffic by checking the database
|
||||
and Redis. Kubernetes uses this to decide whether to *route traffic* to
|
||||
the pod.
|
||||
|
||||
Returns **200 OK** when all critical subsystems are reachable, or
|
||||
**503 Service Unavailable** when the database is down.
|
||||
|
||||
**Authentication:** None (designed for kubelet probes).
|
||||
"""
|
||||
checks: dict[str, dict[str, str]] = {}
|
||||
db_ok = False
|
||||
|
||||
# ── Database check ─────────────────────────────────────────────────
|
||||
try:
|
||||
with engine.connect() as conn:
|
||||
conn.execute(text("SELECT 1"))
|
||||
checks["database"] = {"status": "ok"}
|
||||
db_ok = True
|
||||
except Exception as exc:
|
||||
logger.warning("Readiness probe: database check 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("Readiness probe: Redis check failed: %s", exc)
|
||||
checks["redis"] = {"status": "error", "detail": str(exc)}
|
||||
|
||||
http_status = 503 if not db_ok else 200
|
||||
overall = "ready" if db_ok else "not_ready"
|
||||
return JSONResponse(content={"status": overall, "checks": checks}, status_code=http_status)
|
||||
|
||||
|
||||
@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
|
||||
|
||||
+235
-47
@@ -5,7 +5,9 @@ Dropbox API endpoints
|
||||
import logging
|
||||
import os
|
||||
from typing import Annotated, Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
import requests
|
||||
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -23,6 +25,104 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
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)]
|
||||
|
||||
|
||||
def _build_dropbox_redirect_uri(request: Request) -> str:
|
||||
"""Build the Dropbox OAuth callback redirect URI.
|
||||
|
||||
Uses ``PUBLIC_BASE_URL`` when configured (recommended for deployments behind
|
||||
a reverse proxy that doesn't forward ``X-Forwarded-Proto``). Falls back to
|
||||
deriving the URI from the incoming request's scheme and host headers.
|
||||
"""
|
||||
if settings.public_base_url:
|
||||
return settings.public_base_url.rstrip("/") + "/dropbox-callback"
|
||||
return f"{request.url.scheme}://{request.url.netloc}/dropbox-callback"
|
||||
|
||||
|
||||
@router.get("/dropbox/global-authorize-url")
|
||||
@require_login
|
||||
async def dropbox_global_authorize_url(request: Request):
|
||||
"""Return the Dropbox OAuth authorization URL using the global app credentials.
|
||||
|
||||
This endpoint is used when ``DROPBOX_ALLOW_GLOBAL_CREDENTIALS_FOR_INTEGRATIONS``
|
||||
is enabled so that users can authorize their personal Dropbox integration without
|
||||
needing to supply their own app key/secret. Only the public ``app_key`` is
|
||||
embedded in the URL; the ``app_secret`` is never sent to the browser.
|
||||
"""
|
||||
if not settings.dropbox_allow_global_credentials_for_integrations:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Global credentials for integrations are not enabled",
|
||||
)
|
||||
if not settings.dropbox_app_key or not settings.dropbox_app_secret:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Global Dropbox credentials are not configured",
|
||||
)
|
||||
redirect_uri = _build_dropbox_redirect_uri(request)
|
||||
authorize_url = (
|
||||
"https://www.dropbox.com/oauth2/authorize"
|
||||
f"?client_id={settings.dropbox_app_key}"
|
||||
"&response_type=code"
|
||||
"&token_access_type=offline"
|
||||
f"&redirect_uri={quote(redirect_uri, safe='')}"
|
||||
)
|
||||
return {"authorize_url": authorize_url}
|
||||
|
||||
|
||||
@router.post("/dropbox/exchange-token-global")
|
||||
@require_login
|
||||
async def exchange_dropbox_token_global(
|
||||
request: Request,
|
||||
code: Annotated[str, Form(...)],
|
||||
redirect_uri: Annotated[str, Form(...)],
|
||||
):
|
||||
"""Exchange an authorization code using the global Dropbox app credentials.
|
||||
|
||||
Used when ``DROPBOX_ALLOW_GLOBAL_CREDENTIALS_FOR_INTEGRATIONS`` is enabled so
|
||||
that the ``app_secret`` is never exposed to the browser. Only the OAuth code
|
||||
and redirect URI need to be supplied by the client.
|
||||
"""
|
||||
if not settings.dropbox_allow_global_credentials_for_integrations:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Global credentials for integrations are not enabled",
|
||||
)
|
||||
if not settings.dropbox_app_key or not settings.dropbox_app_secret:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="Global Dropbox credentials are not configured",
|
||||
)
|
||||
|
||||
token_url = "https://api.dropboxapi.com/oauth2/token"
|
||||
payload = {
|
||||
"client_id": settings.dropbox_app_key,
|
||||
"client_secret": settings.dropbox_app_secret,
|
||||
"code": code,
|
||||
"redirect_uri": redirect_uri,
|
||||
"grant_type": "authorization_code",
|
||||
}
|
||||
|
||||
token_data = exchange_oauth_token(provider_name="Dropbox", token_url=token_url, payload=payload)
|
||||
|
||||
return {
|
||||
"refresh_token": token_data["refresh_token"],
|
||||
"access_token": token_data["access_token"],
|
||||
"expires_in": token_data.get("expires_in", 14400),
|
||||
# Return the public app_key so the callback can store it in the integration
|
||||
"app_key": settings.dropbox_app_key,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/dropbox/exchange-token")
|
||||
@require_login
|
||||
async def exchange_dropbox_token(
|
||||
@@ -132,57 +232,60 @@ async def test_dropbox_token(request: Request):
|
||||
"message": "Dropbox credentials are not fully configured",
|
||||
}
|
||||
|
||||
# Check token validity by getting current account info
|
||||
headers = {"Authorization": f"Bearer {settings.dropbox_refresh_token}"}
|
||||
response = requests.post(
|
||||
"https://api.dropboxapi.com/2/users/get_current_account",
|
||||
headers=headers,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
# If token is invalid, try refreshing it
|
||||
if response.status_code == 401:
|
||||
logger.info("Dropbox access token invalid or expired, trying to refresh")
|
||||
|
||||
# Get a new access token using the refresh token
|
||||
refresh_url = "https://api.dropbox.com/oauth2/token"
|
||||
refresh_data = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": settings.dropbox_refresh_token,
|
||||
"client_id": settings.dropbox_app_key,
|
||||
"client_secret": settings.dropbox_app_secret,
|
||||
}
|
||||
|
||||
refresh_response = requests.post(refresh_url, data=refresh_data, timeout=settings.http_request_timeout)
|
||||
|
||||
if refresh_response.status_code != 200:
|
||||
logger.error(f"Failed to refresh Dropbox token: {refresh_response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "Refresh token has expired or is invalid",
|
||||
"needs_reauth": True,
|
||||
}
|
||||
|
||||
token_info = refresh_response.json()
|
||||
access_token = token_info.get("access_token")
|
||||
|
||||
# Try again with the new access token
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
response = requests.post(
|
||||
async with httpx.AsyncClient() as client:
|
||||
# Check token validity by getting current account info
|
||||
headers = {"Authorization": f"Bearer {settings.dropbox_refresh_token}"}
|
||||
response = await client.post(
|
||||
"https://api.dropboxapi.com/2/users/get_current_account",
|
||||
headers=headers,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Dropbox token test failed: {response.status_code} {response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": f"Token validation failed with status {response.status_code}: {response.text}",
|
||||
}
|
||||
# If token is invalid, try refreshing it
|
||||
if response.status_code == 401:
|
||||
logger.info("Dropbox access token invalid or expired, trying to refresh")
|
||||
|
||||
# Get account info
|
||||
account_info = response.json()
|
||||
# Get a new access token using the refresh token
|
||||
refresh_url = "https://api.dropbox.com/oauth2/token"
|
||||
refresh_data = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": settings.dropbox_refresh_token,
|
||||
"client_id": settings.dropbox_app_key,
|
||||
"client_secret": settings.dropbox_app_secret,
|
||||
}
|
||||
|
||||
refresh_response = await client.post(
|
||||
refresh_url, data=refresh_data, timeout=settings.http_request_timeout
|
||||
)
|
||||
|
||||
if refresh_response.status_code != 200:
|
||||
logger.error(f"Failed to refresh Dropbox token: {refresh_response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "Refresh token has expired or is invalid",
|
||||
"needs_reauth": True,
|
||||
}
|
||||
|
||||
token_info = refresh_response.json()
|
||||
access_token = token_info.get("access_token")
|
||||
|
||||
# Try again with the new access token
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
response = await client.post(
|
||||
"https://api.dropboxapi.com/2/users/get_current_account",
|
||||
headers=headers,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Dropbox token test failed: {response.status_code} {response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": f"Token validation failed with status {response.status_code}: {response.text}",
|
||||
}
|
||||
|
||||
# Get account info
|
||||
account_info = response.json()
|
||||
account_email = account_info.get("email", "Unknown account")
|
||||
account_name = account_info.get("name", {}).get("display_name", "Unknown user")
|
||||
|
||||
@@ -207,15 +310,100 @@ async def test_dropbox_token(request: Request):
|
||||
return {"status": "error", "message": f"Connection error: {str(e)}"}
|
||||
|
||||
|
||||
@router.post("/dropbox/save-settings")
|
||||
@router.post("/dropbox/list-folders")
|
||||
@require_login
|
||||
async def list_dropbox_folders(
|
||||
request: Request,
|
||||
access_token: Annotated[str, Form(...)],
|
||||
path: Annotated[str, Form()] = "",
|
||||
):
|
||||
"""
|
||||
List folders in a Dropbox account for the directory selector.
|
||||
|
||||
Accepts an OAuth access token (short-lived) and a path to list.
|
||||
Returns a flat list of folder entries under the given path.
|
||||
"""
|
||||
try:
|
||||
# Normalize path: Dropbox API uses "" for root, otherwise "/path"
|
||||
folder_path = path.strip()
|
||||
if folder_path == "/":
|
||||
folder_path = ""
|
||||
elif folder_path and not folder_path.startswith("/"):
|
||||
folder_path = f"/{folder_path}"
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
|
||||
payload = {
|
||||
"path": folder_path,
|
||||
"recursive": False,
|
||||
"include_deleted": False,
|
||||
"include_has_explicit_shared_members": False,
|
||||
"include_mounted_folders": True,
|
||||
}
|
||||
|
||||
response = requests.post(
|
||||
"https://api.dropboxapi.com/2/files/list_folder",
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Access token is invalid or expired. Please re-authorize.",
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Dropbox list_folder failed: {response.status_code} {response.text}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"Failed to list Dropbox folders: {response.text}",
|
||||
)
|
||||
|
||||
data = response.json()
|
||||
folders = []
|
||||
for entry in data.get("entries", []):
|
||||
if entry.get(".tag") == "folder":
|
||||
folders.append(
|
||||
{
|
||||
"name": entry["name"],
|
||||
"path": entry["path_display"],
|
||||
"id": entry.get("id", ""),
|
||||
}
|
||||
)
|
||||
|
||||
# Sort folders alphabetically
|
||||
folders.sort(key=lambda f: f["name"].lower())
|
||||
|
||||
return {
|
||||
"folders": folders,
|
||||
"path": folder_path or "/",
|
||||
"has_more": data.get("has_more", False),
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.exception(f"Error listing Dropbox folders: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to list folders: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
@router.post("/dropbox/save-settings")
|
||||
async def save_dropbox_settings(
|
||||
request: Request,
|
||||
refresh_token: Annotated[str, Form(...)],
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
app_key: Annotated[Optional[str], Form()] = None,
|
||||
app_secret: Annotated[Optional[str], Form()] = None,
|
||||
folder_path: Annotated[Optional[str], Form()] = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
Save Dropbox settings to database (primary) and .env file (best-effort).
|
||||
|
||||
@@ -0,0 +1,235 @@
|
||||
"""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
|
||||
|
||||
if dup_hashes:
|
||||
# Fetch all matching files (both original and duplicates) in a single batch query
|
||||
all_records = (
|
||||
db.query(FileRecord).filter(FileRecord.filehash.in_(dup_hashes)).order_by(FileRecord.id.asc()).all()
|
||||
)
|
||||
|
||||
# Group records by hash
|
||||
originals_by_hash = {}
|
||||
duplicates_by_hash = {h: [] for h in dup_hashes}
|
||||
|
||||
for record in all_records:
|
||||
h = record.filehash
|
||||
if not record.is_duplicate:
|
||||
# Store only the first original record per hash, matching the old .first() behaviour
|
||||
if h not in originals_by_hash:
|
||||
originals_by_hash[h] = record
|
||||
else:
|
||||
duplicates_by_hash[h].append(record)
|
||||
total_duplicate_files += 1
|
||||
|
||||
for filehash in dup_hashes:
|
||||
original = originals_by_hash.get(filehash)
|
||||
duplicates = duplicates_by_hash.get(filehash, [])
|
||||
|
||||
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,
|
||||
}
|
||||
+409
-31
@@ -11,7 +11,8 @@ import zipfile
|
||||
from datetime import datetime, timezone
|
||||
from typing import Annotated, List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile
|
||||
import aiofiles
|
||||
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
|
||||
@@ -19,14 +20,17 @@ from sqlalchemy.orm import Session
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.middleware.upload_rate_limit import require_upload_rate_limit
|
||||
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, get_file_role
|
||||
|
||||
# Set up logging
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -49,7 +53,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 +73,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 +89,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 +102,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 +212,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 +245,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")
|
||||
@@ -285,18 +300,33 @@ def delete_file_record(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
Delete a file record from the database.
|
||||
This only removes the database entry, not the actual file.
|
||||
Only the file owner (or an admin) may delete a document.
|
||||
"""
|
||||
# Check if file deletion is allowed
|
||||
if not settings.allow_file_delete:
|
||||
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")
|
||||
|
||||
# Enforce owner-only deletion in multi-user mode
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
|
||||
if not is_admin:
|
||||
owner_id = get_current_owner_id(request)
|
||||
role = get_file_role(file_record, owner_id, db)
|
||||
if role != "owner":
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Only the file owner can delete this document",
|
||||
)
|
||||
|
||||
# Log the deletion
|
||||
logger.info(f"Deleting file record: ID={file_id}, Filename={file_record.original_filename}")
|
||||
|
||||
@@ -323,6 +353,7 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
"""
|
||||
Delete multiple file records from the database.
|
||||
This only removes the database entries, not the actual files.
|
||||
Only the file owner (or an admin) may delete each document.
|
||||
"""
|
||||
# Check if file deletion is allowed
|
||||
if not settings.allow_file_delete:
|
||||
@@ -330,11 +361,25 @@ def bulk_delete_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
|
||||
try:
|
||||
# Find all file records
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
|
||||
# Enforce owner-only deletion in multi-user mode
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(user.get("is_admin"))
|
||||
if not is_admin:
|
||||
owner_id = get_current_owner_id(request)
|
||||
non_owner_ids = [f.id for f in file_records if get_file_role(f, owner_id, db) != "owner"]
|
||||
if non_owner_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"You can only delete files you own. Not owner of file IDs: {non_owner_ids}",
|
||||
)
|
||||
|
||||
deleted_count = len(file_records)
|
||||
deleted_ids = [f.id for f in file_records]
|
||||
|
||||
@@ -369,7 +414,9 @@ def bulk_reprocess_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
"""
|
||||
try:
|
||||
# Find all file records
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -441,7 +488,9 @@ def bulk_reprocess_files_cloud_ocr(request: Request, file_ids: List[int], db: Db
|
||||
Useful for re-running OCR on files with poor text quality or missing OCR text.
|
||||
"""
|
||||
try:
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -521,7 +570,9 @@ def bulk_download_files(request: Request, file_ids: List[int], db: DbSession):
|
||||
Files not found on disk are silently skipped.
|
||||
"""
|
||||
try:
|
||||
file_records = db.query(FileRecord).filter(FileRecord.id.in_(file_ids)).all()
|
||||
query = db.query(FileRecord).filter(FileRecord.id.in_(file_ids))
|
||||
query = apply_owner_filter(query, request)
|
||||
file_records = query.all()
|
||||
|
||||
if not file_records:
|
||||
raise HTTPException(status_code=404, detail="No files found with the provided IDs")
|
||||
@@ -603,7 +654,9 @@ def reprocess_single_file(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
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 with ID {file_id} not found")
|
||||
@@ -659,7 +712,9 @@ def reprocess_with_cloud_ocr(request: Request, file_id: int, db: DbSession):
|
||||
"""
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
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 with ID {file_id} not found")
|
||||
@@ -922,7 +977,9 @@ def retry_subtask(
|
||||
"""
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
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 with ID {file_id} not found")
|
||||
@@ -1064,7 +1121,9 @@ def get_file_preview(
|
||||
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
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 with ID {file_id} not found")
|
||||
@@ -1144,7 +1203,9 @@ def download_file(
|
||||
|
||||
try:
|
||||
# Find the file record
|
||||
file_record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
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 with ID {file_id} not found")
|
||||
@@ -1202,9 +1263,79 @@ def download_file(
|
||||
raise HTTPException(status_code=500, detail=f"Error downloading file: {str(e)}")
|
||||
|
||||
|
||||
async def _save_upload_file_chunks(file: UploadFile, target_path: str, max_size: int) -> int:
|
||||
"""Save an uploaded file in chunks and enforce the maximum size limit."""
|
||||
try:
|
||||
written_size = 0
|
||||
with open(target_path, "wb") as f:
|
||||
chunk_size = 65536 # 64 KB chunks
|
||||
while True:
|
||||
chunk = await file.read(chunk_size)
|
||||
if not chunk:
|
||||
break
|
||||
written_size += len(chunk)
|
||||
if written_size > max_size:
|
||||
# Exceeded limit mid-stream; clean up and reject
|
||||
f.close()
|
||||
os.remove(target_path)
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"File too large: exceeded {max_size} bytes during upload. "
|
||||
f"See SECURITY_AUDIT.md for configuration details.",
|
||||
)
|
||||
f.write(chunk)
|
||||
return written_size
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
if os.path.exists(target_path):
|
||||
os.remove(target_path)
|
||||
raise HTTPException(status_code=500, detail=f"Failed to save file: {e}")
|
||||
|
||||
|
||||
def _check_for_exact_duplicate(db: DbSession, target_path: str, safe_filename: str) -> dict | None:
|
||||
"""Check for an exact duplicate of the uploaded file.
|
||||
|
||||
Returns a dict with duplicate info when the file's SHA-256 hash matches an
|
||||
already-processed document, or ``None`` when no duplicate is found (or
|
||||
deduplication is disabled).
|
||||
"""
|
||||
if not settings.enable_deduplication:
|
||||
return None
|
||||
|
||||
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:
|
||||
logger.info(f"Exact duplicate detected on upload: '{safe_filename}' matches file ID {existing.id}")
|
||||
return {
|
||||
"duplicate_type": "exact",
|
||||
"original_file_id": existing.id,
|
||||
"original_filename": existing.original_filename,
|
||||
"message": (
|
||||
"This file is an exact duplicate of an already-processed document. "
|
||||
"It has not been queued for processing again."
|
||||
),
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"Duplicate check failed for uploaded file '{safe_filename}': {e}")
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@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(...),
|
||||
_rate_ok: None = Depends(require_upload_rate_limit),
|
||||
):
|
||||
"""Endpoint to accept a user-uploaded file and enqueue it for processing."""
|
||||
workdir = settings.workdir
|
||||
|
||||
@@ -1241,11 +1372,28 @@ 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:
|
||||
written_size = 0
|
||||
with open(target_path, "wb") as f:
|
||||
async with aiofiles.open(target_path, "wb") as f:
|
||||
chunk_size = 65536 # 64 KB chunks
|
||||
while True:
|
||||
chunk = await file.read(chunk_size)
|
||||
@@ -1254,14 +1402,14 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
written_size += len(chunk)
|
||||
if written_size > max_size:
|
||||
# Exceeded limit mid-stream; clean up and reject
|
||||
f.close()
|
||||
await f.close()
|
||||
os.remove(target_path)
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"File too large: exceeded {max_size} bytes during upload. "
|
||||
f"See SECURITY_AUDIT.md for configuration details.",
|
||||
)
|
||||
f.write(chunk)
|
||||
await f.write(chunk)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@@ -1273,6 +1421,25 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
logger.info(f"Saved uploaded file '{safe_filename}' as '{target_filename}'")
|
||||
file_size = written_size
|
||||
|
||||
# ── Early duplicate rejection ──────────────────────────────────────────
|
||||
# Check for exact duplicates (same SHA-256 hash) BEFORE enqueuing a
|
||||
# processing task. When deduplication is enabled and the file already
|
||||
# exists, we skip processing entirely, clean up the temp file, and
|
||||
# return the existing file's information to the caller.
|
||||
exact_duplicate = _check_for_exact_duplicate(db, target_path, safe_filename)
|
||||
if exact_duplicate:
|
||||
# Remove the just-saved temp file — it's a duplicate.
|
||||
try:
|
||||
os.remove(target_path)
|
||||
except OSError:
|
||||
pass
|
||||
return {
|
||||
"status": "duplicate",
|
||||
"original_filename": safe_filename,
|
||||
"stored_filename": target_filename,
|
||||
"duplicate_of": exact_duplicate,
|
||||
}
|
||||
|
||||
# Determine if the file is a PDF or needs conversion
|
||||
mime_type, _ = mimetypes.guess_type(target_path)
|
||||
file_ext = os.path.splitext(target_path)[1].lower()
|
||||
@@ -1301,7 +1468,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 +1491,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",
|
||||
@@ -1336,18 +1503,20 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
".tif",
|
||||
".webp",
|
||||
".svg",
|
||||
".heic",
|
||||
".heif",
|
||||
}:
|
||||
# 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 {
|
||||
"task_id": task.id,
|
||||
@@ -1355,3 +1524,212 @@ async def ui_upload(request: Request, file: UploadFile = File(...)):
|
||||
"original_filename": safe_filename,
|
||||
"stored_filename": target_filename,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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("File %d claimed by user", file_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("Bulk claim: claimed=%s, skipped=%s", claimed, [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")
|
||||
|
||||
logger.info("Admin assigned owner to %d file(s)", updated)
|
||||
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}
|
||||
|
||||
+26
-13
@@ -23,6 +23,17 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
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)]
|
||||
|
||||
|
||||
@router.post("/google-drive/exchange-token")
|
||||
@require_login
|
||||
async def exchange_google_drive_token(
|
||||
@@ -362,15 +373,15 @@ def format_time_remaining(time_delta):
|
||||
|
||||
|
||||
@router.post("/google-drive/save-settings")
|
||||
@require_login
|
||||
async def save_dropbox_settings(
|
||||
async def save_google_drive_settings(
|
||||
request: Request,
|
||||
refresh_token: Annotated[str, Form(...)],
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
client_id: Annotated[Optional[str], Form()] = None,
|
||||
client_secret: Annotated[Optional[str], Form()] = None,
|
||||
folder_id: Annotated[Optional[str], Form()] = None,
|
||||
use_oauth: Annotated[str, Form()] = "true",
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
Save Google Drive settings to the .env file (best-effort) and persist to database.
|
||||
@@ -403,9 +414,10 @@ async def save_dropbox_settings(
|
||||
if folder_id:
|
||||
drive_settings["GOOGLE_DRIVE_FOLDER_ID"] = folder_id
|
||||
|
||||
# Try to update the .env file, but don't fail if it doesn't exist (for Docker containers)
|
||||
if os.path.exists(env_path):
|
||||
try:
|
||||
# Best-effort .env file write — failures here are non-fatal
|
||||
env_file_written = False
|
||||
try:
|
||||
if os.path.exists(env_path):
|
||||
logger.info(f"Updating Google Drive settings in {env_path}")
|
||||
|
||||
# Read the current .env file
|
||||
@@ -438,12 +450,13 @@ async def save_dropbox_settings(
|
||||
f.write("\n".join(new_env_lines) + "\n")
|
||||
|
||||
logger.info("Successfully updated Google Drive settings in .env file")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to update .env file: {str(e)}, but will continue with in-memory update")
|
||||
else:
|
||||
logger.warning(
|
||||
f".env file not found at {env_path}, skipping file update but continuing with in-memory update"
|
||||
)
|
||||
env_file_written = True
|
||||
else:
|
||||
logger.warning(
|
||||
f".env file not found at {env_path}, skipping file update but continuing with in-memory update"
|
||||
)
|
||||
except Exception as env_err:
|
||||
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
|
||||
|
||||
# Update the settings in memory (this always happens)
|
||||
if refresh_token:
|
||||
@@ -481,7 +494,7 @@ async def save_dropbox_settings(
|
||||
return {
|
||||
"status": "success",
|
||||
"message": "Google Drive settings have been saved",
|
||||
"in_memory_only": not os.path.exists(env_path),
|
||||
"in_memory_only": not env_file_written,
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -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,377 @@
|
||||
"""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.network import is_private_ip
|
||||
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}``.
|
||||
"""
|
||||
|
||||
# Security: Prevent SSRF by blocking connections to internal IPs
|
||||
if is_private_ip(host):
|
||||
logger.warning("SSRF blocked: Attempt to connect to private IP %s", host)
|
||||
return {"success": False, "message": "Connection error: Invalid hostname or IP address"}
|
||||
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,754 @@
|
||||
"""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
|
||||
|
||||
# Optional Dropbox SDK — imported at module level so tests can patch it cleanly.
|
||||
try:
|
||||
import dropbox as dbx_lib
|
||||
from dropbox.exceptions import AuthError as _DropboxAuthError
|
||||
from dropbox.exceptions import BadInputError as _DropboxBadInputError
|
||||
except ImportError: # pragma: no cover
|
||||
dbx_lib = None # type: ignore[assignment]
|
||||
|
||||
class _DropboxAuthError(Exception): # type: ignore[no-redef]
|
||||
"""Stub — only used when the dropbox package is missing."""
|
||||
|
||||
class _DropboxBadInputError(Exception): # type: ignore[no-redef]
|
||||
"""Stub — only used when the dropbox package is missing."""
|
||||
|
||||
|
||||
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"}
|
||||
|
||||
from app.utils.network import is_private_ip
|
||||
|
||||
if is_private_ip(host):
|
||||
logger.warning("SSRF blocked: Attempt to connect to private IP %s", host)
|
||||
return {"success": False, "message": "Connection error: Invalid hostname or IP address"}
|
||||
|
||||
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")
|
||||
endpoint_url = cfg.get("endpoint_url")
|
||||
|
||||
if not bucket:
|
||||
return {"success": False, "message": "Missing required field: bucket"}
|
||||
|
||||
if endpoint_url:
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from app.utils.network import is_private_ip
|
||||
|
||||
parsed_url = urlparse(endpoint_url)
|
||||
if parsed_url.hostname and is_private_ip(parsed_url.hostname):
|
||||
logger.warning("SSRF blocked: Attempt to connect to private IP via S3 endpoint %s", endpoint_url)
|
||||
return {"success": False, "message": "Connection error: Invalid endpoint URL or private IP"}
|
||||
|
||||
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=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_dropbox_connection(config: dict[str, Any] | None, credentials: dict[str, Any] | None) -> dict[str, Any]:
|
||||
"""Test a Dropbox connection by verifying OAuth credentials via the Dropbox API."""
|
||||
if dbx_lib is None:
|
||||
return {"success": False, "message": "dropbox package is not installed"} # pragma: no cover
|
||||
|
||||
creds = credentials or {}
|
||||
app_key = creds.get("app_key", "")
|
||||
app_secret = creds.get("app_secret", "")
|
||||
refresh_token = creds.get("refresh_token", "")
|
||||
|
||||
if not refresh_token:
|
||||
return {"success": False, "message": "Missing required credential: refresh_token"}
|
||||
if not app_key or not app_secret:
|
||||
return {"success": False, "message": "Missing required credentials: app_key and app_secret"}
|
||||
|
||||
try:
|
||||
dbx = dbx_lib.Dropbox(
|
||||
app_key=app_key,
|
||||
app_secret=app_secret,
|
||||
oauth2_refresh_token=refresh_token,
|
||||
)
|
||||
account = dbx.users_get_current_account()
|
||||
display_name = getattr(account, "name", None)
|
||||
name_str = ""
|
||||
if display_name:
|
||||
name_str = f" ({getattr(display_name, 'display_name', '') or ''})"
|
||||
return {"success": True, "message": f"Dropbox connection successful{name_str}"}
|
||||
except _DropboxAuthError as exc:
|
||||
logger.warning("Dropbox auth error: %s", exc)
|
||||
return {
|
||||
"success": False,
|
||||
"message": "Dropbox authentication failed — check app_key, app_secret, and refresh_token",
|
||||
}
|
||||
except _DropboxBadInputError as exc:
|
||||
logger.warning("Dropbox bad input error: %s", exc)
|
||||
return {"success": False, "message": "Dropbox connection failed — invalid credentials format"}
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("Dropbox connection error: %s", exc)
|
||||
return {"success": False, "message": "Dropbox connection failed — check credentials and network connectivity"}
|
||||
|
||||
|
||||
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 httpx
|
||||
|
||||
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:
|
||||
auth = (username, password) if username and password else None
|
||||
headers = {"Depth": "0"}
|
||||
|
||||
# Use httpx for secure connection testing, avoiding urllib vulnerabilities
|
||||
resp = httpx.request("PROPFIND", url, auth=auth, headers=headers, timeout=10.0, follow_redirects=False)
|
||||
if resp.status_code < 400:
|
||||
return {"success": True, "message": "WebDAV connection successful"}
|
||||
return {"success": False, "message": f"WebDAV returned HTTP {resp.status_code}"}
|
||||
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.DROPBOX: _test_dropbox_connection,
|
||||
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(
|
||||
request,
|
||||
"signup.html",
|
||||
context={
|
||||
"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(request, "verify_email_sent.html")
|
||||
|
||||
|
||||
@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(
|
||||
request,
|
||||
"forgot_username.html",
|
||||
context={
|
||||
"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(
|
||||
request,
|
||||
"forgot_password.html",
|
||||
context={
|
||||
"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(
|
||||
request,
|
||||
"password_reset_form.html",
|
||||
context={
|
||||
"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,362 @@
|
||||
"""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
|
||||
preferred_language: str | None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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_200_OK)
|
||||
@require_login
|
||||
async def deactivate_device(
|
||||
request: Request,
|
||||
device_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, str]:
|
||||
"""Deactivate or permanently delete a push-notification device registration.
|
||||
|
||||
* **Active device** – soft-deactivated: the record is kept for audit
|
||||
purposes but will no longer receive push notifications.
|
||||
* **Already-inactive device** – hard-deleted: the record is permanently
|
||||
removed from the database.
|
||||
"""
|
||||
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")
|
||||
|
||||
if device.is_active:
|
||||
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)
|
||||
return {"detail": "Device deactivated"}
|
||||
|
||||
# Hard-delete an already-inactive device.
|
||||
try:
|
||||
db.delete(device)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
logger.info("Mobile device permanently deleted: id=%s owner=%s", device_id, owner_id)
|
||||
return {"detail": "Device deleted"}
|
||||
|
||||
|
||||
@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,
|
||||
"preferred_language": profile.preferred_language if profile else None,
|
||||
}
|
||||
@@ -0,0 +1,483 @@
|
||||
"""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:
|
||||
# Pre-fetch existing preferences for this user to avoid N+1 queries
|
||||
existing_prefs = (
|
||||
db.query(UserNotificationPreference).filter(UserNotificationPreference.owner_id == owner_id).all()
|
||||
)
|
||||
|
||||
# Build a fast lookup dictionary keyed by (event_type, channel_type, target_id)
|
||||
prefs_dict = {(pref.event_type, pref.channel_type, pref.target_id): pref for pref in existing_prefs}
|
||||
|
||||
for item in body.preferences:
|
||||
existing = prefs_dict.get((item.event_type, item.channel_type, item.target_id))
|
||||
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}
|
||||
+150
-113
@@ -3,10 +3,10 @@ OneDrive API endpoints
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Annotated, Optional
|
||||
|
||||
import httpx
|
||||
import requests
|
||||
from fastapi import APIRouter, Depends, Form, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
@@ -14,6 +14,7 @@ from sqlalchemy.orm import Session
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.utils.env_utils import update_env_file
|
||||
from app.utils.oauth_helper import exchange_oauth_token
|
||||
from app.utils.settings_service import save_setting_to_db
|
||||
from app.utils.settings_sync import notify_settings_updated
|
||||
@@ -24,6 +25,17 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
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)]
|
||||
|
||||
|
||||
@router.post("/onedrive/exchange-token")
|
||||
@require_login
|
||||
async def exchange_onedrive_token(
|
||||
@@ -56,6 +68,7 @@ async def exchange_onedrive_token(
|
||||
# Return just what's needed by the frontend
|
||||
return {
|
||||
"refresh_token": token_data["refresh_token"],
|
||||
"access_token": token_data.get("access_token", ""),
|
||||
"expires_in": token_data.get("expires_in", 3600),
|
||||
}
|
||||
|
||||
@@ -92,17 +105,18 @@ async def test_onedrive_token(request: Request):
|
||||
"scope": "offline_access Files.ReadWrite",
|
||||
}
|
||||
|
||||
response = requests.post(token_url, data=refresh_data, timeout=settings.http_request_timeout)
|
||||
async with httpx.AsyncClient(timeout=settings.http_request_timeout) as client:
|
||||
response = await client.post(token_url, data=refresh_data)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Failed to refresh OneDrive token: {response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "Refresh token has expired or is invalid",
|
||||
"needs_reauth": True,
|
||||
}
|
||||
if response.status_code != 200:
|
||||
logger.error(f"Failed to refresh OneDrive token: {response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": "Refresh token has expired or is invalid",
|
||||
"needs_reauth": True,
|
||||
}
|
||||
|
||||
token_data = response.json()
|
||||
token_data = response.json()
|
||||
access_token = token_data.get("access_token")
|
||||
expires_in = token_data.get("expires_in", 3600) # Default to 1 hour if not specified
|
||||
|
||||
@@ -115,32 +129,7 @@ async def test_onedrive_token(request: Request):
|
||||
settings.onedrive_refresh_token = new_refresh_token
|
||||
|
||||
# Also try to update .env file if it exists
|
||||
try:
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
|
||||
if os.path.exists(env_path):
|
||||
with open(env_path, "r") as f:
|
||||
env_lines = f.readlines()
|
||||
|
||||
updated_lines = []
|
||||
updated = False
|
||||
|
||||
for line in env_lines:
|
||||
if line.startswith("ONEDRIVE_REFRESH_TOKEN="):
|
||||
updated_lines.append(f"ONEDRIVE_REFRESH_TOKEN={new_refresh_token}\n")
|
||||
updated = True
|
||||
else:
|
||||
updated_lines.append(line)
|
||||
|
||||
if not updated:
|
||||
updated_lines.append(f"ONEDRIVE_REFRESH_TOKEN={new_refresh_token}\n")
|
||||
|
||||
with open(env_path, "w") as f:
|
||||
f.writelines(updated_lines)
|
||||
|
||||
logger.info("Updated refresh token in .env file")
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to update refresh token in .env file: {e}")
|
||||
update_env_file({"ONEDRIVE_REFRESH_TOKEN": new_refresh_token})
|
||||
|
||||
# Persist the rotated refresh token to the database
|
||||
try:
|
||||
@@ -164,17 +153,18 @@ async def test_onedrive_token(request: Request):
|
||||
user_info_url = "https://graph.microsoft.com/v1.0/me"
|
||||
headers = {"Authorization": f"Bearer {access_token}"}
|
||||
|
||||
user_response = requests.get(user_info_url, headers=headers, timeout=settings.http_request_timeout)
|
||||
async with httpx.AsyncClient(timeout=settings.http_request_timeout) as client:
|
||||
user_response = await client.get(user_info_url, headers=headers)
|
||||
|
||||
if user_response.status_code != 200:
|
||||
logger.error(f"OneDrive token test failed: {user_response.status_code} {user_response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": f"Token validation failed with status {user_response.status_code}: {user_response.text}",
|
||||
}
|
||||
if user_response.status_code != 200:
|
||||
logger.error(f"OneDrive token test failed: {user_response.status_code} {user_response.text}")
|
||||
return {
|
||||
"status": "error",
|
||||
"message": f"Token validation failed with status {user_response.status_code}: {user_response.text}",
|
||||
}
|
||||
|
||||
# Get user info
|
||||
user_info = user_response.json()
|
||||
# Get user info
|
||||
user_info = user_response.json()
|
||||
display_name = user_info.get("displayName", "Unknown user")
|
||||
email = user_info.get("userPrincipalName", "Unknown email")
|
||||
|
||||
@@ -206,6 +196,102 @@ async def test_onedrive_token(request: Request):
|
||||
return {"status": "error", "message": f"Connection error: {str(e)}"}
|
||||
|
||||
|
||||
@router.post("/onedrive/list-folders")
|
||||
@require_login
|
||||
async def list_onedrive_folders(
|
||||
request: Request,
|
||||
access_token: Annotated[str, Form(...)],
|
||||
path: Annotated[str, Form()] = "",
|
||||
):
|
||||
"""
|
||||
List folders in a OneDrive account for the directory selector.
|
||||
|
||||
Accepts an OAuth access token (short-lived) and a path to list.
|
||||
Returns a flat list of folder entries under the given path.
|
||||
"""
|
||||
try:
|
||||
folder_path = path.strip().strip("/")
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
}
|
||||
|
||||
# Build the Graph API URL for listing children
|
||||
if not folder_path or folder_path == "root":
|
||||
url = "https://graph.microsoft.com/v1.0/me/drive/root/children"
|
||||
else:
|
||||
url = f"https://graph.microsoft.com/v1.0/me/drive/root:/{folder_path}:/children"
|
||||
|
||||
# Only request folders and minimal fields
|
||||
params = {
|
||||
"$filter": "folder ne null",
|
||||
"$select": "name,id,parentReference,folder",
|
||||
"$top": "200",
|
||||
}
|
||||
|
||||
response = requests.get(
|
||||
url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
timeout=settings.http_request_timeout,
|
||||
)
|
||||
|
||||
if response.status_code == 401:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="Access token is invalid or expired. Please re-authorize.",
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
logger.error(f"OneDrive list children failed: {response.status_code} {response.text}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"Failed to list OneDrive folders: {response.text}",
|
||||
)
|
||||
|
||||
data = response.json()
|
||||
folders = []
|
||||
for item in data.get("value", []):
|
||||
if "folder" in item:
|
||||
parent_path = ""
|
||||
if item.get("parentReference", {}).get("path"):
|
||||
# parentReference.path looks like /drive/root:/some/path
|
||||
raw_parent = item["parentReference"]["path"]
|
||||
prefix = "/drive/root:"
|
||||
if raw_parent.startswith(prefix):
|
||||
parent_path = raw_parent[len(prefix) :]
|
||||
elif raw_parent == "/drive/root":
|
||||
parent_path = ""
|
||||
|
||||
item_path = f"{parent_path}/{item['name']}" if parent_path else f"/{item['name']}"
|
||||
|
||||
folders.append(
|
||||
{
|
||||
"name": item["name"],
|
||||
"path": item_path,
|
||||
"id": item.get("id", ""),
|
||||
"child_count": item.get("folder", {}).get("childCount", 0),
|
||||
}
|
||||
)
|
||||
|
||||
# Sort folders alphabetically
|
||||
folders.sort(key=lambda f: f["name"].lower())
|
||||
|
||||
return {
|
||||
"folders": folders,
|
||||
"path": f"/{folder_path}" if folder_path else "/",
|
||||
}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.exception(f"Error listing OneDrive folders: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to list folders: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
def format_time_remaining(time_delta):
|
||||
"""Format a timedelta into a human-readable string."""
|
||||
if time_delta.total_seconds() <= 0:
|
||||
@@ -227,15 +313,15 @@ def format_time_remaining(time_delta):
|
||||
|
||||
|
||||
@router.post("/onedrive/save-settings")
|
||||
@require_login
|
||||
async def save_onedrive_settings(
|
||||
request: Request,
|
||||
refresh_token: Annotated[str, Form(...)],
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
client_id: Annotated[Optional[str], Form()] = None,
|
||||
client_secret: Annotated[Optional[str], Form()] = None,
|
||||
tenant_id: Annotated[str, Form()] = "common",
|
||||
folder_path: Annotated[Optional[str], Form()] = None,
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""
|
||||
Saves to database (primary) and .env file (best-effort).
|
||||
@@ -246,75 +332,26 @@ async def save_onedrive_settings(
|
||||
user.get("preferred_username") or user.get("username") or user.get("email") or user.get("id") or "wizard"
|
||||
)
|
||||
|
||||
# Best-effort .env file write
|
||||
try:
|
||||
env_path = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), ".env")
|
||||
if not os.path.exists(env_path):
|
||||
logger.warning(f".env file not found at {env_path}, skipping file write")
|
||||
else:
|
||||
logger.info(f"Updating OneDrive settings in {env_path}")
|
||||
# Build settings dictionary mapped to database/memory keys
|
||||
onedrive_settings = {
|
||||
"onedrive_refresh_token": refresh_token,
|
||||
"onedrive_client_id": client_id,
|
||||
"onedrive_client_secret": client_secret,
|
||||
"onedrive_tenant_id": tenant_id,
|
||||
"onedrive_folder_path": folder_path,
|
||||
}
|
||||
|
||||
with open(env_path, "r") as f:
|
||||
env_lines = f.readlines()
|
||||
# Filter out None values
|
||||
onedrive_settings = {k: v for k, v in onedrive_settings.items() if v is not None}
|
||||
|
||||
onedrive_settings = {"ONEDRIVE_REFRESH_TOKEN": refresh_token}
|
||||
if client_id:
|
||||
onedrive_settings["ONEDRIVE_CLIENT_ID"] = client_id
|
||||
if client_secret:
|
||||
onedrive_settings["ONEDRIVE_CLIENT_SECRET"] = client_secret
|
||||
if tenant_id:
|
||||
onedrive_settings["ONEDRIVE_TENANT_ID"] = tenant_id
|
||||
if folder_path:
|
||||
onedrive_settings["ONEDRIVE_FOLDER_PATH"] = folder_path
|
||||
# Best-effort .env file write using the new utility
|
||||
env_settings = {k.upper(): v for k, v in onedrive_settings.items()}
|
||||
update_env_file(env_settings)
|
||||
|
||||
updated = set()
|
||||
new_env_lines = []
|
||||
for line in env_lines:
|
||||
stripped_line = line.rstrip()
|
||||
is_updated = False
|
||||
for key, value in onedrive_settings.items():
|
||||
if stripped_line.startswith(f"{key}=") or stripped_line.startswith(f"# {key}="):
|
||||
new_env_lines.append(f"{key}={value}")
|
||||
updated.add(key)
|
||||
is_updated = True
|
||||
break
|
||||
if not is_updated:
|
||||
new_env_lines.append(stripped_line)
|
||||
|
||||
for key, value in onedrive_settings.items():
|
||||
if key not in updated:
|
||||
new_env_lines.append(f"{key}={value}")
|
||||
|
||||
with open(env_path, "w") as f:
|
||||
f.write("\n".join(new_env_lines) + "\n")
|
||||
|
||||
logger.info("Successfully updated OneDrive settings in .env file")
|
||||
except Exception as env_err:
|
||||
logger.warning(f"Failed to write .env file (non-fatal): {env_err}")
|
||||
|
||||
# Update the settings in memory
|
||||
if refresh_token:
|
||||
settings.onedrive_refresh_token = refresh_token
|
||||
if client_id:
|
||||
settings.onedrive_client_id = client_id
|
||||
if client_secret:
|
||||
settings.onedrive_client_secret = client_secret
|
||||
if tenant_id:
|
||||
settings.onedrive_tenant_id = tenant_id
|
||||
if folder_path:
|
||||
settings.onedrive_folder_path = folder_path
|
||||
|
||||
# Persist to database (primary)
|
||||
if refresh_token:
|
||||
save_setting_to_db(db, "onedrive_refresh_token", refresh_token, changed_by=changed_by)
|
||||
if client_id:
|
||||
save_setting_to_db(db, "onedrive_client_id", client_id, changed_by=changed_by)
|
||||
if client_secret:
|
||||
save_setting_to_db(db, "onedrive_client_secret", client_secret, changed_by=changed_by)
|
||||
if tenant_id:
|
||||
save_setting_to_db(db, "onedrive_tenant_id", tenant_id, changed_by=changed_by)
|
||||
if folder_path:
|
||||
save_setting_to_db(db, "onedrive_folder_path", folder_path, changed_by=changed_by)
|
||||
# Update in-memory settings and persist to database dynamically
|
||||
for key, value in onedrive_settings.items():
|
||||
setattr(settings, key, value)
|
||||
save_setting_to_db(db, key, value, changed_by=changed_by)
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
|
||||
@@ -0,0 +1,933 @@
|
||||
"""
|
||||
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 built-in and custom rules (filename patterns, content keywords, metadata matching).",
|
||||
"config_schema": {
|
||||
"use_builtin_rules": {
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "Include the pre-built classification rules (invoice, contract, receipt, etc.).",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 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,288 @@
|
||||
"""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
|
||||
|
||||
# Fetch all requested plans in a single query to avoid N+1
|
||||
plan_ids = body.order
|
||||
plans = db.query(SubscriptionPlan).filter(SubscriptionPlan.plan_id.in_(plan_ids)).all()
|
||||
|
||||
# Build a map for fast O(1) lookup
|
||||
plan_map = {p.plan_id: p for p in plans}
|
||||
|
||||
for sort_order, plan_id in enumerate(plan_ids):
|
||||
plan = plan_map.get(plan_id)
|
||||
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,358 @@
|
||||
"""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
|
||||
default_document_language: str | None
|
||||
"""ISO 639-1 code for the user's preferred document translation target language."""
|
||||
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'")
|
||||
default_document_language: str | None = Field(
|
||||
default=None,
|
||||
description="ISO 639-1 code for the default document translation target language, e.g. 'en', 'de'",
|
||||
)
|
||||
|
||||
|
||||
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]
|
||||
default_document_language=profile.default_document_language, # 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]
|
||||
|
||||
# Validate default document language
|
||||
if body.default_document_language is not None:
|
||||
doc_lang = body.default_document_language.lower().strip()
|
||||
if doc_lang and doc_lang not in SUPPORTED_LANGUAGE_CODES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"Unsupported language code: {doc_lang}",
|
||||
)
|
||||
profile.default_document_language = doc_lang 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]
|
||||
default_document_language=profile.default_document_language, # 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,249 @@
|
||||
"""QR code login API endpoints for mobile app authentication.
|
||||
|
||||
Provides a secure challenge-response flow for logging into the mobile app
|
||||
by scanning a QR code displayed in the web interface:
|
||||
|
||||
1. **Web user** calls ``POST /qr-auth/challenge`` → receives a time-limited
|
||||
challenge token (encoded in the QR code).
|
||||
2. **Web UI** polls ``GET /qr-auth/challenge/{id}/status`` to detect when
|
||||
the mobile app has claimed the challenge.
|
||||
3. **Mobile app** scans the QR code and calls ``POST /qr-auth/claim`` with
|
||||
the challenge token + device name → receives an API token.
|
||||
|
||||
Security properties:
|
||||
* Challenges expire after a configurable TTL (default 2 minutes).
|
||||
* Single-use: once claimed, a challenge cannot be reused (replay-safe).
|
||||
* Cryptographically random 64-byte tokens.
|
||||
* IP addresses are logged for audit.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import io
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any
|
||||
|
||||
import segno
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
from app.middleware.audit_log import get_client_ip
|
||||
from app.utils.session_manager import (
|
||||
claim_qr_challenge,
|
||||
create_qr_challenge,
|
||||
get_challenge_status,
|
||||
)
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/qr-auth", tags=["qr-auth"])
|
||||
|
||||
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 CreateChallengeResponse(BaseModel):
|
||||
"""Response after creating a QR login challenge."""
|
||||
|
||||
challenge_id: int
|
||||
challenge_token: str
|
||||
expires_at: datetime
|
||||
ttl_seconds: int = Field(description="Seconds until the challenge expires (use for client-side countdown).")
|
||||
qr_payload: str = Field(description="The string to encode in the QR code.")
|
||||
qr_code_svg: str = Field(description="Base64-encoded SVG data URI of the QR code, ready for use in an <img> src.")
|
||||
|
||||
|
||||
class ChallengeStatusResponse(BaseModel):
|
||||
"""Response for polling the status of a QR challenge."""
|
||||
|
||||
id: int
|
||||
status: str # "pending", "claimed", "expired", "cancelled"
|
||||
device_name: str | None = None
|
||||
claimed_at: datetime | None = None
|
||||
expires_at: datetime
|
||||
|
||||
|
||||
class ClaimChallengeRequest(BaseModel):
|
||||
"""Request body for claiming a QR login challenge."""
|
||||
|
||||
challenge_token: str = Field(min_length=1, max_length=256)
|
||||
device_name: str = Field(
|
||||
default="Mobile App",
|
||||
min_length=1,
|
||||
max_length=120,
|
||||
description="Human-readable device name.",
|
||||
)
|
||||
|
||||
|
||||
class ClaimChallengeResponse(BaseModel):
|
||||
"""Response after successfully claiming a QR challenge."""
|
||||
|
||||
token: str
|
||||
token_id: int
|
||||
name: str
|
||||
owner_id: str
|
||||
created_at: datetime
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# QR code rendering parameters
|
||||
_QR_ERROR_LEVEL = "M" # Medium error correction (~15% recovery); sufficient for on-screen display
|
||||
_QR_SCALE = 4 # Each QR module is rendered as 4×4 SVG pixels
|
||||
|
||||
|
||||
def _generate_qr_svg(payload: str) -> str:
|
||||
"""Generate a QR code for *payload* and return it as a base64 SVG data URI.
|
||||
|
||||
Using ``segno`` (pure-Python, no Pillow dependency) and SVG output so the
|
||||
QR code scales crisply at any resolution without requiring a canvas or any
|
||||
client-side JavaScript library.
|
||||
"""
|
||||
qr = segno.make(payload, error=_QR_ERROR_LEVEL)
|
||||
buf = io.BytesIO()
|
||||
qr.save(buf, kind="svg", scale=_QR_SCALE, xmldecl=False, svgclass=None, lineclass=None, omitsize=True)
|
||||
svg_bytes = buf.getvalue()
|
||||
return "data:image/svg+xml;base64," + base64.b64encode(svg_bytes).decode("ascii")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/challenge", status_code=status.HTTP_201_CREATED, response_model=CreateChallengeResponse)
|
||||
@require_login
|
||||
async def create_challenge(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Create a new QR login challenge.
|
||||
|
||||
The returned ``qr_payload`` should be encoded into a QR code and
|
||||
displayed to the user. The mobile app scans this QR code and
|
||||
calls the ``/claim`` endpoint.
|
||||
"""
|
||||
if not settings.qr_login_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="QR login feature is currently disabled. Please contact your administrator to enable it.",
|
||||
)
|
||||
ip = get_client_ip(request)
|
||||
challenge = create_qr_challenge(db, owner_id, ip_address=ip)
|
||||
|
||||
# The QR payload is a JSON-like string with enough info for the mobile
|
||||
# app to know the server URL and challenge token.
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
qr_payload = f"docuelevate://qr-login?token={challenge.challenge_token}&server={base_url}"
|
||||
|
||||
# Compute the TTL in seconds so the client can run a countdown timer
|
||||
# without comparing absolute timestamps (which breaks when client and
|
||||
# server clocks are out of sync).
|
||||
ttl_seconds = max(0, int((challenge.expires_at - challenge.created_at).total_seconds()))
|
||||
|
||||
return {
|
||||
"challenge_id": challenge.id,
|
||||
"challenge_token": challenge.challenge_token,
|
||||
"expires_at": challenge.expires_at,
|
||||
"ttl_seconds": ttl_seconds,
|
||||
"qr_payload": qr_payload,
|
||||
"qr_code_svg": _generate_qr_svg(qr_payload),
|
||||
}
|
||||
|
||||
|
||||
@router.get("/challenge/{challenge_id}/status", response_model=ChallengeStatusResponse)
|
||||
@require_login
|
||||
async def poll_challenge_status(
|
||||
request: Request,
|
||||
challenge_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Poll the status of a QR login challenge.
|
||||
|
||||
The web UI calls this endpoint every few seconds to check if the
|
||||
mobile app has scanned the QR code and claimed the challenge.
|
||||
"""
|
||||
if not settings.qr_login_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="QR login feature is currently disabled. Please contact your administrator to enable it.",
|
||||
)
|
||||
result = get_challenge_status(db, challenge_id, owner_id)
|
||||
if not result:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Challenge not found")
|
||||
return result
|
||||
|
||||
|
||||
@router.post("/claim", response_model=ClaimChallengeResponse)
|
||||
async def claim_challenge(
|
||||
request: Request,
|
||||
body: ClaimChallengeRequest,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Claim a QR login challenge and receive an API token.
|
||||
|
||||
This endpoint is called by the mobile app after scanning a QR code.
|
||||
It does **not** require authentication — the challenge token itself
|
||||
serves as proof that the user authorized this login from their web
|
||||
session.
|
||||
"""
|
||||
if not settings.qr_login_enabled:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="QR login feature is currently disabled. Please contact your administrator to enable it.",
|
||||
)
|
||||
ip = get_client_ip(request)
|
||||
result = claim_qr_challenge(db, body.challenge_token, device_name=body.device_name, ip_address=ip)
|
||||
|
||||
if not result:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Invalid, expired, or already claimed challenge.",
|
||||
)
|
||||
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="qr_login_claimed",
|
||||
user=result["owner_id"],
|
||||
resource_type="session",
|
||||
ip_address=ip,
|
||||
details={"device_name": body.device_name, "token_id": result["token_id"]},
|
||||
severity="info",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write QR login audit event", exc_info=True)
|
||||
|
||||
return result
|
||||
@@ -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,
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
"""API endpoints for managing user sessions.
|
||||
|
||||
Provides endpoints for listing active sessions, revoking individual sessions,
|
||||
and the "log off everywhere" feature that invalidates all sessions and API
|
||||
tokens across all devices.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.middleware.audit_log import get_client_ip
|
||||
from app.utils.session_manager import (
|
||||
get_session_lifetime_days,
|
||||
list_user_sessions,
|
||||
revoke_all_sessions,
|
||||
revoke_session,
|
||||
)
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
|
||||
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)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Response schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class SessionResponse(BaseModel):
|
||||
"""Serialised user session for the management UI."""
|
||||
|
||||
id: int
|
||||
device_info: str | None
|
||||
ip_address: str | None
|
||||
created_at: datetime
|
||||
last_active_at: datetime
|
||||
expires_at: datetime
|
||||
is_current: bool = False
|
||||
|
||||
|
||||
class SessionListResponse(BaseModel):
|
||||
"""Response for listing active sessions."""
|
||||
|
||||
sessions: list[SessionResponse]
|
||||
session_lifetime_days: int
|
||||
|
||||
|
||||
class RevokeAllResponse(BaseModel):
|
||||
"""Response after revoking all sessions."""
|
||||
|
||||
revoked_count: int
|
||||
message: str
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/", response_model=SessionListResponse)
|
||||
@require_login
|
||||
async def list_sessions(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""List all active sessions for the current user."""
|
||||
sessions = list_user_sessions(db, owner_id)
|
||||
|
||||
# Determine which session is the current one
|
||||
current_token = request.session.get("_session_token")
|
||||
|
||||
session_list = []
|
||||
for s in sessions:
|
||||
session_list.append(
|
||||
{
|
||||
"id": s.id,
|
||||
"device_info": s.device_info,
|
||||
"ip_address": s.ip_address,
|
||||
"created_at": s.created_at,
|
||||
"last_active_at": s.last_active_at,
|
||||
"expires_at": s.expires_at,
|
||||
"is_current": s.session_token == current_token if current_token else False,
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"sessions": session_list,
|
||||
"session_lifetime_days": get_session_lifetime_days(),
|
||||
}
|
||||
|
||||
|
||||
@router.delete("/{session_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
@require_login
|
||||
async def revoke_single_session(
|
||||
request: Request,
|
||||
session_id: int,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> None:
|
||||
"""Revoke a specific session by ID."""
|
||||
success = revoke_session(db, session_id, owner_id)
|
||||
if not success:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
|
||||
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="session_revoked",
|
||||
user=owner_id,
|
||||
resource_type="session",
|
||||
resource_id=str(session_id),
|
||||
ip_address=get_client_ip(request),
|
||||
severity="info",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write session revocation audit event", exc_info=True)
|
||||
|
||||
|
||||
@router.post("/revoke-all", response_model=RevokeAllResponse)
|
||||
@require_login
|
||||
async def revoke_all(
|
||||
request: Request,
|
||||
owner_id: CurrentOwner,
|
||||
db: DbSession,
|
||||
) -> dict[str, Any]:
|
||||
"""Revoke all sessions except the current one ("log off everywhere").
|
||||
|
||||
Also revokes all active API tokens for the user, which invalidates
|
||||
mobile app sessions and any programmatic access.
|
||||
"""
|
||||
# Find current session to preserve it
|
||||
current_token = request.session.get("_session_token")
|
||||
current_session_id = None
|
||||
if current_token:
|
||||
from app.models import UserSession
|
||||
|
||||
current = db.query(UserSession).filter(UserSession.session_token == current_token).first()
|
||||
if current:
|
||||
current_session_id = current.id
|
||||
|
||||
count = revoke_all_sessions(
|
||||
db,
|
||||
owner_id,
|
||||
except_session_id=current_session_id,
|
||||
revoke_api_tokens=True,
|
||||
)
|
||||
|
||||
try:
|
||||
from app.utils.audit_service import record_event
|
||||
|
||||
record_event(
|
||||
db,
|
||||
action="revoke_all_sessions",
|
||||
user=owner_id,
|
||||
resource_type="session",
|
||||
ip_address=get_client_ip(request),
|
||||
details={"revoked_count": count},
|
||||
severity="warning",
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Failed to write revoke-all audit event", exc_info=True)
|
||||
|
||||
return {
|
||||
"revoked_count": count,
|
||||
"message": f"Successfully revoked {count} session(s) and all API tokens.",
|
||||
}
|
||||
@@ -55,6 +55,12 @@ class SettingUpdate(BaseModel):
|
||||
value: Optional[str] = Field(None, description="Setting value (None to delete)")
|
||||
|
||||
|
||||
class SettingValueUpdate(BaseModel):
|
||||
"""Model for updating a setting value by key (key is provided in the URL path)."""
|
||||
|
||||
value: Optional[str] = Field(None, description="Setting value (None to delete)")
|
||||
|
||||
|
||||
class SettingResponse(BaseModel):
|
||||
"""Model for setting response"""
|
||||
|
||||
@@ -323,6 +329,62 @@ async def update_setting(
|
||||
)
|
||||
|
||||
|
||||
@router.put("/{key}")
|
||||
async def put_setting(
|
||||
key: str,
|
||||
body: SettingValueUpdate,
|
||||
request: Request,
|
||||
db: DbSession,
|
||||
admin: AdminUser,
|
||||
):
|
||||
"""
|
||||
Update a specific setting by key (RESTful PUT).
|
||||
|
||||
Accepts a body with only ``value``; the key is taken from the URL path.
|
||||
This is the endpoint used by the admin Connections wizard.
|
||||
Admin only.
|
||||
"""
|
||||
validate_setting_key(key)
|
||||
try:
|
||||
if body.value is not None:
|
||||
is_valid, error_message = validate_setting_value(key, body.value)
|
||||
if not is_valid:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=error_message)
|
||||
|
||||
user = request.session.get("user", {}) if hasattr(request, "session") else {}
|
||||
changed_by = (
|
||||
user.get("preferred_username") or user.get("username") or user.get("email") or user.get("id") or "admin"
|
||||
)
|
||||
|
||||
success = save_setting_to_db(db, key, body.value, changed_by=changed_by)
|
||||
if not success:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to save setting to database",
|
||||
)
|
||||
|
||||
notify_settings_updated()
|
||||
|
||||
metadata = get_setting_metadata(key)
|
||||
restart_required = metadata.get("restart_required", False)
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"message": f"Setting '{key}' updated successfully",
|
||||
"restart_required": restart_required,
|
||||
"key": key,
|
||||
"value": body.value,
|
||||
}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"Error updating setting {key}: {e}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Failed to update setting: {key}",
|
||||
)
|
||||
|
||||
|
||||
@router.delete("/{key}")
|
||||
async def delete_setting(key: str, request: Request, db: DbSession, admin: AdminUser):
|
||||
"""
|
||||
@@ -446,6 +508,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,502 @@
|
||||
"""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, FileRecord.original_filename)
|
||||
.outerjoin(FileRecord, SharedLink.file_id == FileRecord.id)
|
||||
.filter(SharedLink.owner_id == owner_id)
|
||||
)
|
||||
if active_only:
|
||||
q = q.filter(SharedLink.is_active.is_(True))
|
||||
links_with_filenames = q.order_by(SharedLink.created_at.desc()).all()
|
||||
|
||||
base_url = str(request.base_url).rstrip("/")
|
||||
result = []
|
||||
for link, filename in links_with_filenames:
|
||||
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,355 @@
|
||||
"""File-sharing API endpoints.
|
||||
|
||||
Provides CRUD operations for ``FileShare`` records, which grant named
|
||||
users ``viewer`` or ``editor`` access to a document owned by someone
|
||||
else. Only the file owner may create, update, or revoke shares.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Request, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.auth import require_login
|
||||
from app.database import get_db
|
||||
from app.models import FILE_SHARE_ROLE_VIEWER, FILE_SHARE_ROLES, FileRecord, FileShare, UserProfile
|
||||
from app.utils.user_scope import get_current_owner_id, get_file_role
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(tags=["sharing"])
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _serialize_share(share: FileShare) -> dict[str, Any]:
|
||||
"""Serialize a ``FileShare`` to a JSON-friendly dict."""
|
||||
return {
|
||||
"id": share.id,
|
||||
"file_id": share.file_id,
|
||||
"owner_id": share.owner_id,
|
||||
"shared_with_user_id": share.shared_with_user_id,
|
||||
"role": share.role,
|
||||
"created_at": share.created_at.isoformat() if share.created_at else None,
|
||||
"updated_at": share.updated_at.isoformat() if share.updated_at else None,
|
||||
}
|
||||
|
||||
|
||||
def _require_owner(file_record: FileRecord, user_id: str | None, db: Session) -> None:
|
||||
"""Raise 403 unless the calling user is the file owner."""
|
||||
if get_file_role(file_record, user_id, db) != "owner":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Only the file owner can manage shares",
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# List shares
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/shares")
|
||||
@require_login
|
||||
def list_shares(request: Request, file_id: int, db: DbSession):
|
||||
"""List all shares for a document.
|
||||
|
||||
Only the file owner (or an admin) may call this endpoint.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
|
||||
Returns:
|
||||
A list of share objects.
|
||||
"""
|
||||
user_id = get_current_owner_id(request)
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(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")
|
||||
|
||||
role = get_file_role(file_record, user_id, db)
|
||||
if role is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
if role != "owner" and not is_admin:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="Only the file owner can view shares",
|
||||
)
|
||||
|
||||
shares = db.query(FileShare).filter(FileShare.file_id == file_id).all()
|
||||
return [_serialize_share(s) for s in shares]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Create share
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post("/files/{file_id}/shares", status_code=status.HTTP_201_CREATED)
|
||||
@require_login
|
||||
def create_share(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
shared_with_user_id: str = Body(..., embed=True),
|
||||
role: str = Body(FILE_SHARE_ROLE_VIEWER, embed=True),
|
||||
):
|
||||
"""Share a document with another user.
|
||||
|
||||
Only the file owner may share the document. Sharing with a user
|
||||
that already has access updates their role instead of creating a
|
||||
duplicate record.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document to share.
|
||||
|
||||
Request body (JSON):
|
||||
shared_with_user_id: The stable user identifier of the recipient.
|
||||
role: ``"viewer"`` (default) or ``"editor"``.
|
||||
|
||||
Returns:
|
||||
The created or updated share object.
|
||||
"""
|
||||
owner_id = get_current_owner_id(request)
|
||||
|
||||
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")
|
||||
|
||||
_require_owner(file_record, owner_id, db)
|
||||
|
||||
if role not in FILE_SHARE_ROLES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"role must be one of: {', '.join(FILE_SHARE_ROLES)}",
|
||||
)
|
||||
|
||||
if not shared_with_user_id or not shared_with_user_id.strip():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="shared_with_user_id must be a non-empty string",
|
||||
)
|
||||
shared_with_user_id = shared_with_user_id.strip()
|
||||
|
||||
# Cannot share with yourself
|
||||
if shared_with_user_id == owner_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="You cannot share a file with yourself",
|
||||
)
|
||||
|
||||
try:
|
||||
existing = (
|
||||
db.query(FileShare)
|
||||
.filter(FileShare.file_id == file_id, FileShare.shared_with_user_id == shared_with_user_id)
|
||||
.first()
|
||||
)
|
||||
|
||||
if existing:
|
||||
# Update role if different
|
||||
if existing.role != role:
|
||||
existing.role = role
|
||||
db.commit()
|
||||
db.refresh(existing)
|
||||
logger.info(
|
||||
"Share updated: file_id=%s, shared_with=%s, role=%s, by owner=%s",
|
||||
file_id,
|
||||
shared_with_user_id,
|
||||
role,
|
||||
owner_id,
|
||||
)
|
||||
return _serialize_share(existing)
|
||||
|
||||
share = FileShare(
|
||||
file_id=file_id,
|
||||
owner_id=owner_id,
|
||||
shared_with_user_id=shared_with_user_id,
|
||||
role=role,
|
||||
)
|
||||
db.add(share)
|
||||
db.commit()
|
||||
db.refresh(share)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to create share: file_id=%s, shared_with=%s", file_id, shared_with_user_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to create share",
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Share created: id=%s, file_id=%s, shared_with=%s, role=%s, by owner=%s",
|
||||
share.id,
|
||||
file_id,
|
||||
shared_with_user_id,
|
||||
role,
|
||||
owner_id,
|
||||
)
|
||||
return _serialize_share(share)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Update share role
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.put("/files/{file_id}/shares/{share_id}")
|
||||
@require_login
|
||||
def update_share(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
share_id: int,
|
||||
db: DbSession,
|
||||
role: str = Body(..., embed=True),
|
||||
):
|
||||
"""Update the role of an existing share.
|
||||
|
||||
Only the file owner may change the role of a share.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
share_id: The ID of the share record to update.
|
||||
|
||||
Request body (JSON):
|
||||
role: New role — ``"viewer"`` or ``"editor"``.
|
||||
|
||||
Returns:
|
||||
The updated share object.
|
||||
"""
|
||||
owner_id = get_current_owner_id(request)
|
||||
|
||||
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")
|
||||
|
||||
_require_owner(file_record, owner_id, db)
|
||||
|
||||
if role not in FILE_SHARE_ROLES:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail=f"role must be one of: {', '.join(FILE_SHARE_ROLES)}",
|
||||
)
|
||||
|
||||
share = db.query(FileShare).filter(FileShare.id == share_id, FileShare.file_id == file_id).first()
|
||||
if not share:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Share not found")
|
||||
|
||||
try:
|
||||
share.role = role
|
||||
db.commit()
|
||||
db.refresh(share)
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to update share: share_id=%s", share_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to update share",
|
||||
)
|
||||
|
||||
logger.info("Share updated: id=%s, file_id=%s, new_role=%s, by owner=%s", share_id, file_id, role, owner_id)
|
||||
return _serialize_share(share)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Revoke share
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.delete("/files/{file_id}/shares/{share_id}", status_code=status.HTTP_200_OK)
|
||||
@require_login
|
||||
def revoke_share(request: Request, file_id: int, share_id: int, db: DbSession):
|
||||
"""Revoke a share, removing the user's access.
|
||||
|
||||
Only the file owner may revoke shares.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
share_id: The ID of the share record to delete.
|
||||
|
||||
Returns:
|
||||
A success message.
|
||||
"""
|
||||
owner_id = get_current_owner_id(request)
|
||||
|
||||
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")
|
||||
|
||||
_require_owner(file_record, owner_id, db)
|
||||
|
||||
share = db.query(FileShare).filter(FileShare.id == share_id, FileShare.file_id == file_id).first()
|
||||
if not share:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Share not found")
|
||||
|
||||
try:
|
||||
db.delete(share)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception("Failed to revoke share: share_id=%s", share_id)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="Failed to revoke share",
|
||||
)
|
||||
|
||||
logger.info("Share revoked: id=%s, file_id=%s, by owner=%s", share_id, file_id, owner_id)
|
||||
return {"status": "success", "message": "Share revoked successfully"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# List users that the file is already shared with (for the share-picker UI)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/shared-with")
|
||||
@require_login
|
||||
def list_shared_with(request: Request, file_id: int, db: DbSession):
|
||||
"""Return the list of users a document is shared with and their roles.
|
||||
|
||||
Accessible to any user that has at least viewer access to the file,
|
||||
so that editors/viewers can see who else has access.
|
||||
|
||||
Path Parameters:
|
||||
file_id: The ID of the document.
|
||||
|
||||
Returns:
|
||||
A list of ``{share_id, user_id, display_name, role}`` objects.
|
||||
"""
|
||||
user_id = get_current_owner_id(request)
|
||||
user = request.session.get("user")
|
||||
is_admin = isinstance(user, dict) and bool(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")
|
||||
|
||||
role = get_file_role(file_record, user_id, db)
|
||||
if role is None and not is_admin:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
|
||||
shares = db.query(FileShare).filter(FileShare.file_id == file_id).all()
|
||||
|
||||
results = []
|
||||
for s in shares:
|
||||
profile = db.query(UserProfile).filter(UserProfile.user_id == s.shared_with_user_id).first()
|
||||
results.append(
|
||||
{
|
||||
"share_id": s.id,
|
||||
"user_id": s.shared_with_user_id,
|
||||
"display_name": (profile.display_name if profile and profile.display_name else s.shared_with_user_id),
|
||||
"role": s.role,
|
||||
}
|
||||
)
|
||||
return results
|
||||
@@ -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(),
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
"""
|
||||
System reset API endpoints for DocuElevate.
|
||||
|
||||
Provides admin-only REST endpoints for:
|
||||
- Full system reset (wipe all user data)
|
||||
- Reset with re-import (move originals → reimport folder, wipe, re-ingest)
|
||||
|
||||
Both operations require the ``ENABLE_FACTORY_RESET=True`` feature flag and
|
||||
admin privileges.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.config import settings
|
||||
from app.database import get_db
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/admin/system-reset", tags=["system-reset"])
|
||||
|
||||
|
||||
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)]
|
||||
|
||||
|
||||
def _require_feature_enabled() -> None:
|
||||
"""Raise 404 when the factory-reset feature flag is off."""
|
||||
if not settings.enable_factory_reset:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="System reset is not enabled. Set ENABLE_FACTORY_RESET=True to activate.",
|
||||
)
|
||||
|
||||
|
||||
class ResetRequest(BaseModel):
|
||||
"""Body for system reset endpoints. Requires explicit confirmation."""
|
||||
|
||||
confirmation: str
|
||||
|
||||
|
||||
@router.post("/full")
|
||||
async def full_reset(
|
||||
body: ResetRequest,
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""Wipe all user data (database + work-files).
|
||||
|
||||
The caller must send ``{"confirmation": "DELETE"}`` to proceed.
|
||||
"""
|
||||
_require_feature_enabled()
|
||||
|
||||
if body.confirmation != "DELETE":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail='Confirmation required: send {"confirmation": "DELETE"} to proceed.',
|
||||
)
|
||||
|
||||
from app.utils.system_reset import perform_full_reset
|
||||
|
||||
try:
|
||||
result = perform_full_reset(db)
|
||||
except Exception as exc:
|
||||
logger.exception("Full system reset failed")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"System reset failed: {exc}",
|
||||
) from exc
|
||||
|
||||
return {"status": "ok", "result": result}
|
||||
|
||||
|
||||
@router.post("/reimport")
|
||||
async def reset_and_reimport(
|
||||
body: ResetRequest,
|
||||
_admin: AdminUser,
|
||||
db: Session = Depends(get_db),
|
||||
) -> dict:
|
||||
"""Move original files to a reimport folder, wipe everything, and
|
||||
configure the reimport folder as a watch folder for automatic
|
||||
re-ingestion.
|
||||
|
||||
The caller must send ``{"confirmation": "REIMPORT"}`` to proceed.
|
||||
"""
|
||||
_require_feature_enabled()
|
||||
|
||||
if body.confirmation != "REIMPORT":
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail='Confirmation required: send {"confirmation": "REIMPORT"} to proceed.',
|
||||
)
|
||||
|
||||
from app.utils.system_reset import perform_reset_and_reimport
|
||||
|
||||
try:
|
||||
result = perform_reset_and_reimport(db)
|
||||
except Exception as exc:
|
||||
logger.exception("Reset-and-reimport failed")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"Reset and reimport failed: {exc}",
|
||||
) from exc
|
||||
|
||||
return {"status": "ok", "result": result}
|
||||
|
||||
|
||||
@router.get("/status")
|
||||
async def reset_status(_admin: AdminUser) -> dict:
|
||||
"""Return whether the system reset feature is enabled."""
|
||||
return {
|
||||
"enabled": settings.enable_factory_reset,
|
||||
"factory_reset_on_startup": settings.factory_reset_on_startup,
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
"""
|
||||
API endpoints for document translation.
|
||||
|
||||
Provides on-the-fly translation via the AI provider and access to the
|
||||
persisted default-language translation.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
|
||||
from fastapi.responses import JSONResponse
|
||||
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
|
||||
from app.utils.ai_provider import get_ai_provider
|
||||
from app.utils.user_scope import apply_owner_filter
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
DbSession = Annotated[Session, Depends(get_db)]
|
||||
|
||||
# Maximum characters sent to the AI provider for a single translation request.
|
||||
_MAX_TRANSLATION_INPUT = 50_000
|
||||
|
||||
|
||||
def _get_file_or_404(db: Session, file_id: int, request: Request) -> FileRecord:
|
||||
"""Fetch a FileRecord visible to the current user or raise 404."""
|
||||
query = db.query(FileRecord).filter(FileRecord.id == file_id)
|
||||
query = apply_owner_filter(query, request)
|
||||
record = query.first()
|
||||
if not record:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="File not found")
|
||||
return record
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/translation/default")
|
||||
@require_login
|
||||
def get_default_translation(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
) -> JSONResponse:
|
||||
"""Return the persisted default-language translation for a document.
|
||||
|
||||
Returns 404 if no default-language translation has been generated yet
|
||||
(e.g. because the document is already in the default language).
|
||||
"""
|
||||
record = _get_file_or_404(db, file_id, request)
|
||||
|
||||
if not record.default_language_text:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="No default-language translation available for this file",
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
content={
|
||||
"file_id": record.id,
|
||||
"detected_language": record.detected_language,
|
||||
"default_language_code": record.default_language_code,
|
||||
"text": record.default_language_text,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.get("/files/{file_id}/translate")
|
||||
@require_login
|
||||
def translate_on_the_fly(
|
||||
request: Request,
|
||||
file_id: int,
|
||||
db: DbSession,
|
||||
lang: str = Query(..., min_length=2, max_length=10, description="Target language ISO 639-1 code"),
|
||||
) -> JSONResponse:
|
||||
"""Translate a document's extracted text into an arbitrary language on the fly.
|
||||
|
||||
The translation is generated via the configured AI provider and is **not**
|
||||
persisted. For the default-language translation, use the
|
||||
``/files/{file_id}/translation/default`` endpoint instead.
|
||||
"""
|
||||
record = _get_file_or_404(db, file_id, request)
|
||||
|
||||
source_text = record.ocr_text
|
||||
if not source_text:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="No extracted text available for this file — translation requires OCR text",
|
||||
)
|
||||
|
||||
# If the requested language matches what is already stored, return it directly.
|
||||
if record.default_language_code and lang == record.default_language_code and record.default_language_text:
|
||||
return JSONResponse(
|
||||
content={
|
||||
"file_id": record.id,
|
||||
"source_language": record.detected_language,
|
||||
"target_language": lang,
|
||||
"text": record.default_language_text,
|
||||
"cached": True,
|
||||
}
|
||||
)
|
||||
|
||||
# If the detected language already matches, return the original text.
|
||||
detected = record.detected_language
|
||||
if detected and detected == lang:
|
||||
return JSONResponse(
|
||||
content={
|
||||
"file_id": record.id,
|
||||
"source_language": detected,
|
||||
"target_language": lang,
|
||||
"text": source_text,
|
||||
"cached": True,
|
||||
}
|
||||
)
|
||||
|
||||
# Truncate to keep AI costs bounded.
|
||||
text_to_translate = source_text[:_MAX_TRANSLATION_INPUT]
|
||||
|
||||
try:
|
||||
provider = get_ai_provider()
|
||||
model = settings.ai_model or settings.openai_model
|
||||
translated = provider.chat_completion(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
f"You are a professional translator. Translate the following text "
|
||||
f"into {lang}. Preserve the original formatting, paragraph structure, "
|
||||
f"and meaning. Do not add any commentary — output ONLY the translated text."
|
||||
),
|
||||
},
|
||||
{"role": "user", "content": text_to_translate},
|
||||
],
|
||||
model=model,
|
||||
temperature=0.3,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception(f"On-the-fly translation failed for file {file_id}: {exc}")
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail="Translation failed — the AI provider returned an error",
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
content={
|
||||
"file_id": record.id,
|
||||
"source_language": detected or "unknown",
|
||||
"target_language": lang,
|
||||
"text": translated,
|
||||
"cached": False,
|
||||
}
|
||||
)
|
||||
+104
-90
@@ -2,7 +2,6 @@
|
||||
API endpoint for processing files from URLs
|
||||
"""
|
||||
|
||||
import ipaddress
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
@@ -10,15 +9,18 @@ import urllib.parse
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
import aiofiles
|
||||
import httpx
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel, HttpUrl, field_validator
|
||||
|
||||
from app.auth import require_login
|
||||
from app.config import settings
|
||||
from app.middleware.upload_rate_limit import require_upload_rate_limit
|
||||
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 +44,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).
|
||||
@@ -135,9 +106,32 @@ def validate_file_type(content_type: str, filename: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
async def verify_redirect(response: httpx.Response) -> None:
|
||||
"""
|
||||
Event hook to intercept redirects and validate the new destination URL.
|
||||
Prevents SSRF bypasses via redirects to internal networks or metadata endpoints.
|
||||
"""
|
||||
if response.status_code in (301, 302, 303, 307, 308):
|
||||
location = response.headers.get("Location")
|
||||
if location:
|
||||
# Resolve relative redirects
|
||||
new_url = str(response.url.join(location))
|
||||
# Validate the new URL
|
||||
try:
|
||||
validate_url_safety(new_url)
|
||||
except HTTPException as e:
|
||||
# Map the validation error to an httpx exception so it can be handled
|
||||
# properly by the caller, avoiding raw HTTPExceptions escaping the client scope
|
||||
raise httpx.RequestError(f"Redirect to unsafe URL blocked: {e.detail}", request=response.request) from e
|
||||
|
||||
|
||||
@router.post("/process-url")
|
||||
@require_login
|
||||
async def process_url(request: Request, url_request: URLUploadRequest):
|
||||
async def process_url(
|
||||
request: Request,
|
||||
url_request: URLUploadRequest,
|
||||
_rate_ok: None = Depends(require_upload_rate_limit),
|
||||
):
|
||||
"""
|
||||
Download a file from a URL and enqueue it for processing.
|
||||
|
||||
@@ -176,6 +170,19 @@ async def process_url(request: Request, url_request: URLUploadRequest):
|
||||
if not safe_filename:
|
||||
safe_filename = "download"
|
||||
|
||||
# Hook to validate redirects and prevent SSRF
|
||||
async def validate_redirect(response: httpx.Response):
|
||||
if response.is_redirect:
|
||||
location = response.headers.get("Location")
|
||||
if location:
|
||||
# Resolve relative URLs
|
||||
next_url = urllib.parse.urljoin(str(response.url), location)
|
||||
try:
|
||||
validate_url_safety(next_url)
|
||||
except HTTPException as e:
|
||||
# Reraise as a RequestError so httpx aborts the request
|
||||
raise httpx.RequestError(f"Unsafe redirect target: {e.detail}", request=response.request)
|
||||
|
||||
# Download file with security measures
|
||||
# Initialize target_path to None to prevent UnboundLocalError in exception handlers
|
||||
# that may execute before target_path is assigned during error cases
|
||||
@@ -184,67 +191,74 @@ async def process_url(request: Request, url_request: URLUploadRequest):
|
||||
logger.info(f"Downloading file from URL: {url}")
|
||||
|
||||
# Use configured timeout to prevent hanging
|
||||
response = requests.get(
|
||||
url,
|
||||
async with httpx.AsyncClient(
|
||||
timeout=settings.http_request_timeout,
|
||||
stream=True, # Stream to handle large files
|
||||
allow_redirects=True, # Follow redirects
|
||||
follow_redirects=True,
|
||||
headers={
|
||||
"User-Agent": "DocuElevate/1.0", # Identify ourselves
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
event_hooks={"response": [validate_redirect, verify_redirect]},
|
||||
) as client:
|
||||
async with client.stream("GET", url) as response:
|
||||
response.raise_for_status()
|
||||
|
||||
# Validate content type
|
||||
content_type = response.headers.get("Content-Type", "")
|
||||
if not validate_file_type(content_type, safe_filename):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Unsupported file type: {content_type}. "
|
||||
"Supported types: PDF, Office documents, images, plain text",
|
||||
)
|
||||
# Validate content type
|
||||
content_type = response.headers.get("Content-Type", "")
|
||||
if not validate_file_type(content_type, safe_filename):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Unsupported file type: {content_type}. "
|
||||
"Supported types: PDF, Office documents, images, plain text",
|
||||
)
|
||||
|
||||
# Check content length before downloading
|
||||
content_length = response.headers.get("Content-Length")
|
||||
if content_length:
|
||||
file_size = int(content_length)
|
||||
max_size = settings.max_upload_size
|
||||
if file_size > max_size:
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"File too large: {file_size} bytes (max {max_size} bytes)",
|
||||
)
|
||||
|
||||
# Generate unique filename
|
||||
unique_id = str(uuid.uuid4())
|
||||
if "." in safe_filename:
|
||||
file_extension = safe_filename.rsplit(".", 1)[1]
|
||||
target_filename = f"{unique_id}.{file_extension}"
|
||||
else:
|
||||
target_filename = unique_id
|
||||
|
||||
target_path = os.path.join(settings.workdir, target_filename)
|
||||
|
||||
# Download file in chunks to handle large files
|
||||
downloaded_size = 0
|
||||
max_size = settings.max_upload_size
|
||||
|
||||
with open(target_path, "wb") as f:
|
||||
for chunk in response.iter_content(chunk_size=8192):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
downloaded_size += len(chunk)
|
||||
|
||||
# Check size during download
|
||||
if downloaded_size > max_size:
|
||||
# Remove partial file
|
||||
f.close()
|
||||
os.remove(target_path)
|
||||
# Check content length before downloading
|
||||
content_length = response.headers.get("Content-Length")
|
||||
if content_length:
|
||||
file_size = int(content_length)
|
||||
max_size = settings.max_upload_size
|
||||
if file_size > max_size:
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"File too large: exceeded {max_size} bytes during download",
|
||||
detail=f"File too large: {file_size} bytes (max {max_size} bytes)",
|
||||
)
|
||||
|
||||
# Generate unique filename
|
||||
unique_id = str(uuid.uuid4())
|
||||
|
||||
# Check for extension using original_filename to avoid any CodeQL issues
|
||||
# with safe_filename which is derived from the URL directly.
|
||||
if "." in original_filename:
|
||||
_, ext = os.path.splitext(original_filename)
|
||||
# Strip out the leading dot and any non-alphanumeric chars
|
||||
clean_ext = "".join(c for c in ext if c.isalnum())
|
||||
if not clean_ext:
|
||||
clean_ext = "bin"
|
||||
target_filename = f"{unique_id}.{clean_ext}"
|
||||
else:
|
||||
target_filename = unique_id
|
||||
|
||||
target_path = os.path.join(settings.workdir, target_filename)
|
||||
|
||||
# Download file in chunks to handle large files
|
||||
downloaded_size = 0
|
||||
max_size = settings.max_upload_size
|
||||
|
||||
async with aiofiles.open(target_path, "wb") as f:
|
||||
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||
if chunk:
|
||||
await f.write(chunk)
|
||||
downloaded_size += len(chunk)
|
||||
|
||||
# Check size during download
|
||||
if downloaded_size > max_size:
|
||||
# Remove partial file
|
||||
await f.close()
|
||||
os.remove(target_path)
|
||||
raise HTTPException(
|
||||
status_code=413,
|
||||
detail=f"File too large: exceeded {max_size} bytes during download",
|
||||
)
|
||||
|
||||
logger.info(f"Downloaded file from URL '{url}' as '{target_filename}' ({downloaded_size} bytes)")
|
||||
|
||||
# Enqueue for processing
|
||||
@@ -258,19 +272,19 @@ async def process_url(request: Request, url_request: URLUploadRequest):
|
||||
"size": downloaded_size,
|
||||
}
|
||||
|
||||
except requests.exceptions.Timeout:
|
||||
except httpx.TimeoutException:
|
||||
logger.error(f"Timeout while downloading file from URL: {url}")
|
||||
raise HTTPException(status_code=408, detail="Request timeout: server took too long to respond")
|
||||
|
||||
except requests.exceptions.ConnectionError as e:
|
||||
except httpx.ConnectError as e:
|
||||
logger.error(f"Connection error while downloading file from URL: {url} - {str(e)}")
|
||||
raise HTTPException(status_code=502, detail=f"Failed to connect to URL: {str(e)}")
|
||||
|
||||
except requests.exceptions.HTTPError as e:
|
||||
except httpx.HTTPStatusError as e:
|
||||
logger.error(f"HTTP error while downloading file from URL: {url} - {str(e)}")
|
||||
raise HTTPException(status_code=e.response.status_code, detail=f"HTTP error: {str(e)}")
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
except httpx.RequestError as e:
|
||||
logger.error(f"Error downloading file from URL: {url} - {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=f"Failed to download file: {str(e)}")
|
||||
|
||||
|
||||
+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)
|
||||
+1214
-40
File diff suppressed because it is too large
Load Diff
+78
-3
@@ -1,10 +1,15 @@
|
||||
# app/celery_app.py
|
||||
|
||||
import logging
|
||||
import os
|
||||
|
||||
from celery import Celery
|
||||
from celery.signals import task_failure
|
||||
from celery.signals import task_failure, worker_ready
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
celery = Celery(
|
||||
"document_processor",
|
||||
broker=settings.redis_url,
|
||||
@@ -21,6 +26,72 @@ celery.conf.task_routes = {
|
||||
"app.tasks.*": {"queue": "document_processor"},
|
||||
}
|
||||
|
||||
# Mapping of document pipeline task names to the positional index of ``file_id``
|
||||
# in their ``args`` tuple. These indices correspond to the task signatures:
|
||||
# process_with_ocr(filename, file_id, ...) → index 1
|
||||
# extract_metadata_with_gpt(filename, text, file_id) → index 2
|
||||
# embed_metadata_into_pdf(path, text, metadata, file_id) → index 3
|
||||
# Tasks that always pass ``file_id`` as a keyword argument
|
||||
# (e.g. ``process_document``, ``finalize_document_storage``) are not listed
|
||||
# here — their ``file_id`` is found via ``kwargs`` instead.
|
||||
_FILE_ID_ARG_INDEX: dict[str, int] = {
|
||||
"app.tasks.process_with_ocr.process_with_ocr": 1,
|
||||
"app.tasks.extract_metadata_with_gpt.extract_metadata_with_gpt": 2,
|
||||
"app.tasks.embed_metadata_into_pdf.embed_metadata_into_pdf": 3,
|
||||
}
|
||||
|
||||
|
||||
def _dispatch_user_failure_notification(sender, exception, args: list | None, kwargs: dict | None) -> None:
|
||||
"""Best-effort per-user failure notification for document pipeline tasks.
|
||||
|
||||
Extracts ``file_id`` from the failed task's arguments, looks up the owning
|
||||
user from the database, and dispatches a ``document.failed`` notification.
|
||||
"""
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord
|
||||
from app.utils.user_notification import notify_user_document_failed
|
||||
|
||||
task_name = sender.name if sender else ""
|
||||
if not task_name.startswith("app.tasks."):
|
||||
return
|
||||
|
||||
# 1. Resolve file_id from kwargs or positional args
|
||||
file_id = (kwargs or {}).get("file_id")
|
||||
if file_id is None:
|
||||
idx = _FILE_ID_ARG_INDEX.get(task_name)
|
||||
if idx is not None and args and len(args) > idx:
|
||||
val = args[idx]
|
||||
if isinstance(val, int):
|
||||
file_id = val
|
||||
|
||||
if file_id is None:
|
||||
return
|
||||
|
||||
# 2. Look up owner from the database
|
||||
with SessionLocal() as db:
|
||||
record = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if not record or not record.owner_id:
|
||||
return
|
||||
owner_id = record.owner_id
|
||||
filename = record.original_filename or record.local_filename or "unknown"
|
||||
|
||||
# 3. Dispatch per-user notification
|
||||
error_msg = f"{type(exception).__name__}: {exception}" if exception else "Unknown error"
|
||||
notify_user_document_failed(
|
||||
owner_id=owner_id,
|
||||
filename=os.path.basename(filename),
|
||||
error=error_msg,
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
|
||||
@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(
|
||||
@@ -40,6 +111,10 @@ def task_failure_handler(
|
||||
kwargs=kwargs or {},
|
||||
)
|
||||
except Exception as e:
|
||||
import logging
|
||||
logger.exception(f"Failed to send task failure notification: {e}")
|
||||
|
||||
logging.exception(f"Failed to send task failure notification: {e}")
|
||||
# Also dispatch a per-user failure notification for document pipeline tasks
|
||||
try:
|
||||
_dispatch_user_failure_notification(sender, exception, args, kwargs)
|
||||
except Exception:
|
||||
logger.warning("Could not dispatch per-user failure notification", exc_info=True)
|
||||
|
||||
+155
-9
@@ -1,5 +1,7 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
import logging
|
||||
|
||||
from celery.schedules import crontab
|
||||
|
||||
# Ensure tasks are loaded
|
||||
@@ -8,36 +10,63 @@ 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.automation_tasks import deliver_automation_hook_task # noqa: F401
|
||||
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.classify_document import classify_document_task # noqa: F401
|
||||
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
|
||||
from app.tasks.translate_to_default_language import translate_to_default_language # 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_sharepoint import upload_to_sharepoint # 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 +83,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 +118,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()
|
||||
+996
-2
File diff suppressed because it is too large
Load Diff
+59
-10
@@ -10,6 +10,7 @@ from typing import Any
|
||||
from sqlalchemy import create_engine, exc
|
||||
from sqlalchemy.engine.url import make_url
|
||||
from sqlalchemy.orm import Session, declarative_base, sessionmaker
|
||||
from sqlalchemy.pool import NullPool, QueuePool
|
||||
|
||||
from app.config import settings
|
||||
|
||||
@@ -17,17 +18,51 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
Base = declarative_base()
|
||||
|
||||
# Parse the DATABASE_URL
|
||||
# ---------------------------------------------------------------------------
|
||||
# Engine construction
|
||||
# ---------------------------------------------------------------------------
|
||||
DB_URL = settings.database_url
|
||||
engine = create_engine(DB_URL, connect_args={"check_same_thread": False})
|
||||
_parsed_url = make_url(DB_URL)
|
||||
|
||||
_connect_args: dict[str, Any] = {}
|
||||
_engine_kwargs: dict[str, Any] = {
|
||||
"pool_pre_ping": True, # detect stale / dropped connections before use
|
||||
}
|
||||
|
||||
if _parsed_url.get_backend_name() == "sqlite":
|
||||
# SQLite does not benefit from connection pooling and is prone to
|
||||
# QueuePool exhaustion under concurrent access. NullPool opens a fresh
|
||||
# connection for each request and closes it immediately afterwards,
|
||||
# completely avoiding the "QueuePool limit reached" TimeoutError.
|
||||
_connect_args["check_same_thread"] = False
|
||||
_engine_kwargs["poolclass"] = NullPool
|
||||
else:
|
||||
# PostgreSQL / MySQL — use a bounded QueuePool with configurable limits.
|
||||
_engine_kwargs["poolclass"] = QueuePool
|
||||
_engine_kwargs.update(
|
||||
{
|
||||
"pool_size": settings.db_pool_size,
|
||||
"max_overflow": settings.db_max_overflow,
|
||||
"pool_timeout": settings.db_pool_timeout,
|
||||
"pool_recycle": settings.db_pool_recycle,
|
||||
}
|
||||
)
|
||||
|
||||
engine = create_engine(DB_URL, connect_args=_connect_args, **_engine_kwargs)
|
||||
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 +82,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 +239,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 +294,17 @@ 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})"))
|
||||
# SECURITY: Quoted identifiers to prevent SQL injection during index creation
|
||||
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")
|
||||
|
||||
|
||||
+253
-12
@@ -1,8 +1,11 @@
|
||||
#!/usr/bin/env python3
|
||||
import json as _json_mod
|
||||
import logging
|
||||
import os
|
||||
import pathlib
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime as _dt
|
||||
from datetime import timezone as _tz
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request, status
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
@@ -16,6 +19,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 +31,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
|
||||
@@ -33,6 +39,114 @@ from app.views import router as frontend_router
|
||||
# Explicitly include the files router
|
||||
from app.views.files import router as files_router
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Configure Python root logging level early so that *all* loggers (including
|
||||
# those already created via ``logging.getLogger(__name__)`` in other modules)
|
||||
# respect the configured level.
|
||||
#
|
||||
# Standard behaviour (matches Django, Flask, 12-factor conventions):
|
||||
# • ``LOG_LEVEL`` env var takes precedence when explicitly set.
|
||||
# • When ``DEBUG=True`` and ``LOG_LEVEL`` is **not** set, the effective
|
||||
# level is automatically lowered to ``DEBUG``.
|
||||
# • Default (neither flag set): ``INFO``.
|
||||
#
|
||||
# ``LOG_FORMAT=json`` enables structured JSON lines on stdout, suitable for
|
||||
# Promtail, Fluentd, Filebeat, Datadog, Splunk UF, or any log collector.
|
||||
#
|
||||
# ``LOG_SYSLOG_ENABLED=true`` adds a Python SysLogHandler so that every log
|
||||
# message is also forwarded to the configured syslog receiver — useful for
|
||||
# traditional (non-container) deployments and centralised SIEM ingestion.
|
||||
#
|
||||
# Noisy third-party loggers (httpx, httpcore, authlib, etc.) are pinned to
|
||||
# WARNING when the app-level is DEBUG to keep output useful.
|
||||
# ---------------------------------------------------------------------------
|
||||
_explicit_log_level = os.environ.get("LOG_LEVEL")
|
||||
if settings.debug and _explicit_log_level is None:
|
||||
_effective_level = "DEBUG"
|
||||
else:
|
||||
_effective_level = settings.log_level.upper()
|
||||
|
||||
_effective_level_int = getattr(logging, _effective_level, logging.INFO)
|
||||
|
||||
|
||||
class _JsonFormatter(logging.Formatter):
|
||||
"""Emit one JSON object per log line for machine consumption.
|
||||
|
||||
Fields emitted: ``timestamp``, ``level``, ``logger``, ``message``,
|
||||
``module``, ``funcName``, ``lineno``, and — when present — ``exc_info``.
|
||||
Compatible with Grafana Loki, Splunk, ELK, Datadog, and most SIEM tools.
|
||||
"""
|
||||
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry: dict = {
|
||||
"timestamp": _dt.fromtimestamp(record.created, tz=_tz.utc).isoformat(),
|
||||
"level": record.levelname,
|
||||
"logger": record.name,
|
||||
"message": record.getMessage(),
|
||||
"module": record.module,
|
||||
"funcName": record.funcName,
|
||||
"lineno": record.lineno,
|
||||
}
|
||||
if record.exc_info and record.exc_info[1] is not None:
|
||||
log_entry["exc_info"] = self.formatException(record.exc_info)
|
||||
return _json_mod.dumps(log_entry, default=str)
|
||||
|
||||
|
||||
# Choose formatter based on LOG_FORMAT setting
|
||||
if settings.log_format.lower() == "json":
|
||||
_handler = logging.StreamHandler()
|
||||
_handler.setFormatter(_JsonFormatter())
|
||||
logging.root.handlers = [_handler]
|
||||
logging.root.setLevel(_effective_level_int)
|
||||
else:
|
||||
logging.basicConfig(
|
||||
level=_effective_level_int,
|
||||
format="%(asctime)s - %(name)s - %(levelname)s - %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
force=True,
|
||||
)
|
||||
|
||||
# Optional: forward application logs to a syslog receiver
|
||||
if settings.log_syslog_enabled:
|
||||
import logging.handlers as _lh
|
||||
import socket as _socket
|
||||
|
||||
_proto = settings.log_syslog_protocol.lower()
|
||||
_socktype = _socket.SOCK_STREAM if _proto == "tcp" else _socket.SOCK_DGRAM
|
||||
_syslog_handler = _lh.SysLogHandler(
|
||||
address=(settings.log_syslog_host, settings.log_syslog_port),
|
||||
socktype=_socktype,
|
||||
)
|
||||
_syslog_handler.setLevel(_effective_level_int)
|
||||
# Use the same formatter as stdout (text or JSON)
|
||||
if settings.log_format.lower() == "json":
|
||||
_syslog_handler.setFormatter(_JsonFormatter())
|
||||
else:
|
||||
_syslog_handler.setFormatter(logging.Formatter("%(name)s - %(levelname)s - %(message)s"))
|
||||
logging.root.addHandler(_syslog_handler)
|
||||
|
||||
# Keep noisy third-party loggers quiet at DEBUG level
|
||||
if _effective_level_int <= logging.DEBUG:
|
||||
for _noisy in (
|
||||
"httpx",
|
||||
"httpcore",
|
||||
"authlib",
|
||||
"urllib3",
|
||||
"hpack",
|
||||
"multipart",
|
||||
"watchfiles",
|
||||
):
|
||||
logging.getLogger(_noisy).setLevel(logging.WARNING)
|
||||
|
||||
_startup_logger = logging.getLogger(__name__)
|
||||
_startup_logger.info(
|
||||
"Root logging level set to %s (debug=%s, format=%s, syslog=%s)",
|
||||
_effective_level,
|
||||
settings.debug,
|
||||
settings.log_format,
|
||||
settings.log_syslog_enabled,
|
||||
)
|
||||
|
||||
# Load configuration from .env for the session key
|
||||
config = Config(".env")
|
||||
# Use settings.session_secret which has proper validation
|
||||
@@ -56,6 +170,12 @@ async def lifespan(app: FastAPI):
|
||||
# Startup: Initialize database
|
||||
init_db() # Create tables if they don't exist
|
||||
|
||||
# Factory reset on startup — wipe all user data before anything else
|
||||
if settings.factory_reset_on_startup:
|
||||
from app.utils.system_reset import perform_startup_reset
|
||||
|
||||
perform_startup_reset()
|
||||
|
||||
# Load settings from database after DB initialization
|
||||
from app.database import SessionLocal
|
||||
from app.utils.config_loader import load_settings_from_db
|
||||
@@ -69,6 +189,22 @@ async def lifespan(app: FastAPI):
|
||||
finally:
|
||||
db.close()
|
||||
|
||||
# Re-register OAuth / social-login providers now that DB settings are
|
||||
# loaded. auth.py runs its initial registration at import time (before
|
||||
# the lifespan runs), so providers that are only configured in the
|
||||
# database would not be registered yet. Calling refresh here ensures
|
||||
# they are active immediately on startup without any manual restart.
|
||||
try:
|
||||
from app.auth import refresh_social_providers
|
||||
|
||||
refresh_social_providers()
|
||||
except Exception as e:
|
||||
logging.warning(f"Could not refresh social login providers on startup: {e}")
|
||||
|
||||
# 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,17 +235,83 @@ 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
|
||||
|
||||
# Shutdown: Cleanup tasks
|
||||
logging.info("Application shutting down")
|
||||
try:
|
||||
logging.info("Application shutting down")
|
||||
except Exception:
|
||||
_startup_logger.exception("Error during shutdown logging")
|
||||
|
||||
# Send shutdown notification
|
||||
notify_shutdown()
|
||||
try:
|
||||
notify_shutdown()
|
||||
except Exception:
|
||||
_startup_logger.exception("Error sending shutdown notification")
|
||||
|
||||
|
||||
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)
|
||||
@@ -140,8 +342,19 @@ app.add_middleware(CSRFMiddleware, config=settings)
|
||||
# See SECURITY_AUDIT.md – Infrastructure Security section
|
||||
app.add_middleware(AuditLogMiddleware, config=settings)
|
||||
|
||||
|
||||
# 3) Session Middleware (for request.session to work)
|
||||
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET)
|
||||
def _get_session_max_age() -> int:
|
||||
"""Compute session max-age at startup time."""
|
||||
try:
|
||||
from app.utils.session_manager import get_session_max_age_seconds
|
||||
|
||||
return get_session_max_age_seconds()
|
||||
except Exception:
|
||||
return 30 * 86400 # 30 days default fallback
|
||||
|
||||
|
||||
app.add_middleware(SessionMiddleware, secret_key=SESSION_SECRET, max_age=_get_session_max_age())
|
||||
|
||||
# 3a) CORS Middleware - handles cross-origin requests and preflight (OPTIONS) responses.
|
||||
# Disabled by default: set CORS_ENABLED=True only when NOT using a reverse proxy
|
||||
@@ -174,8 +387,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,17 +428,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(request, "404.html", 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(
|
||||
request,
|
||||
"404.html", # Reuse 404 template for other errors, or create a generic error template
|
||||
{"request": request},
|
||||
status_code=exc.status_code,
|
||||
)
|
||||
|
||||
@@ -216,10 +455,10 @@ 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(
|
||||
request,
|
||||
"500.html",
|
||||
{"request": request, "exc": exc},
|
||||
context={"exc": exc},
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
@@ -233,4 +472,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")
|
||||
|
||||
@@ -20,6 +20,9 @@ How it works:
|
||||
|
||||
Exempt paths (CSRF is not checked even for state-changing methods):
|
||||
- ``/oauth-callback`` – OAuth 2.0 callback; protected by the ``state`` parameter.
|
||||
- ``/api/qr-auth/claim`` – Called by the unauthenticated mobile app; the
|
||||
cryptographically-random, single-use challenge token provides equivalent
|
||||
protection.
|
||||
"""
|
||||
|
||||
import logging
|
||||
@@ -39,6 +42,10 @@ CSRF_PROTECTED_METHODS = {"POST", "PUT", "DELETE", "PATCH"}
|
||||
# their own replay-protection mechanism).
|
||||
CSRF_EXEMPT_PATHS = {
|
||||
"/oauth-callback",
|
||||
# The mobile app calls this endpoint without a browser session/CSRF token.
|
||||
# The cryptographically-random, single-use challenge token already provides
|
||||
# equivalent protection against cross-site request forgery.
|
||||
"/api/qr-auth/claim",
|
||||
}
|
||||
|
||||
|
||||
@@ -109,6 +116,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 +162,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:
|
||||
|
||||
@@ -0,0 +1,290 @@
|
||||
"""Per-user, health-aware upload rate limiter for DocuElevate.
|
||||
|
||||
This module provides a FastAPI dependency that enforces per-user upload rate
|
||||
limits using a Redis-backed sliding window counter. The effective limit is
|
||||
dynamically reduced when the system is under heavy load (high Celery queue
|
||||
depth or elevated CPU load average), ensuring the server remains responsive
|
||||
to all users even during bulk-upload scenarios.
|
||||
|
||||
Usage in an endpoint::
|
||||
|
||||
from app.middleware.upload_rate_limit import require_upload_rate_limit
|
||||
|
||||
@router.post("/ui-upload")
|
||||
@require_login
|
||||
async def ui_upload(
|
||||
request: Request,
|
||||
_rate_ok: None = Depends(require_upload_rate_limit),
|
||||
...
|
||||
):
|
||||
...
|
||||
|
||||
See ``docs/ConfigurationGuide.md`` for the configuration options
|
||||
(``UPLOAD_RATE_LIMIT_PER_USER``, ``UPLOAD_RATE_LIMIT_WINDOW``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import redis
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
||||
from app.config import settings
|
||||
from app.utils.user_scope import get_current_owner_id
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Redis key prefix
|
||||
# ---------------------------------------------------------------------------
|
||||
_KEY_PREFIX = "docuelevate:upload_rate"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Health-check queue names (Celery defaults used by DocuElevate)
|
||||
# ---------------------------------------------------------------------------
|
||||
_CELERY_QUEUES = ("document_processor", "default", "celery")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Singleton Redis client (lazy-initialised; fail-open when unavailable)
|
||||
# ---------------------------------------------------------------------------
|
||||
_redis_client: redis.Redis | None = None
|
||||
|
||||
|
||||
def _get_redis() -> redis.Redis | None:
|
||||
"""Return a shared Redis client, or *None* when Redis is unavailable."""
|
||||
global _redis_client
|
||||
if _redis_client is not None:
|
||||
return _redis_client
|
||||
try:
|
||||
_redis_client = redis.Redis.from_url(
|
||||
settings.redis_url,
|
||||
decode_responses=True,
|
||||
socket_connect_timeout=2,
|
||||
socket_timeout=2,
|
||||
)
|
||||
# Quick connectivity check – raises on failure.
|
||||
_redis_client.ping()
|
||||
return _redis_client
|
||||
except Exception: # noqa: BLE001
|
||||
logger.debug("Redis unavailable for upload rate limiter – falling back to allow-all", exc_info=True)
|
||||
_redis_client = None
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Health metrics helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_queue_depth(r: redis.Redis) -> int:
|
||||
"""Return the total number of pending tasks across all Celery queues."""
|
||||
total = 0
|
||||
for queue_name in _CELERY_QUEUES:
|
||||
try:
|
||||
total += r.llen(queue_name)
|
||||
except Exception: # noqa: BLE001, S110
|
||||
logger.debug("Could not read queue length for %r", queue_name, exc_info=True)
|
||||
return total
|
||||
|
||||
|
||||
def _get_cpu_load_ratio() -> float:
|
||||
"""Return the 1-minute load average divided by the number of CPU cores.
|
||||
|
||||
Returns ``0.0`` on platforms that do not support :func:`os.getloadavg`
|
||||
(e.g. Windows) so that the limiter never penalises on those systems.
|
||||
"""
|
||||
try:
|
||||
load_1m = os.getloadavg()[0]
|
||||
cpu_count = os.cpu_count() or 1
|
||||
return load_1m / cpu_count
|
||||
except (OSError, AttributeError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def compute_effective_limit(
|
||||
base_limit: int,
|
||||
queue_depth: int = 0,
|
||||
cpu_load_ratio: float = 0.0,
|
||||
) -> tuple[int, float, str]:
|
||||
"""Compute the effective upload rate limit based on system health.
|
||||
|
||||
The function applies a *reduction factor* (``0.0 < factor ≤ 1.0``) to the
|
||||
configured base limit. Both queue depth and CPU load contribute
|
||||
independently; the lowest factor wins.
|
||||
|
||||
Args:
|
||||
base_limit: The configured maximum uploads per window.
|
||||
queue_depth: Total pending tasks in Celery queues.
|
||||
cpu_load_ratio: 1-minute load average divided by CPU count.
|
||||
|
||||
Returns:
|
||||
A 3-tuple of ``(effective_limit, factor, reason)`` where *reason*
|
||||
is a human-readable tag for logging.
|
||||
"""
|
||||
factor = 1.0
|
||||
reason = "normal"
|
||||
|
||||
# --- Queue-depth thresholds ---
|
||||
if queue_depth > 200:
|
||||
factor, reason = min(factor, 0.10), f"critical_queue({queue_depth})"
|
||||
elif queue_depth > 100:
|
||||
factor, reason = min(factor, 0.25), f"high_queue({queue_depth})"
|
||||
elif queue_depth > 50:
|
||||
factor, reason = min(factor, 0.50), f"moderate_queue({queue_depth})"
|
||||
|
||||
# --- CPU-load thresholds ---
|
||||
if cpu_load_ratio > 3.0:
|
||||
new_factor = 0.10
|
||||
if new_factor < factor:
|
||||
factor, reason = new_factor, f"critical_cpu({cpu_load_ratio:.1f})"
|
||||
elif cpu_load_ratio > 2.0:
|
||||
new_factor = 0.25
|
||||
if new_factor < factor:
|
||||
factor, reason = new_factor, f"high_cpu({cpu_load_ratio:.1f})"
|
||||
elif cpu_load_ratio > 1.5:
|
||||
new_factor = 0.50
|
||||
if new_factor < factor:
|
||||
factor, reason = new_factor, f"moderate_cpu({cpu_load_ratio:.1f})"
|
||||
|
||||
effective = max(1, int(base_limit * factor))
|
||||
return effective, factor, reason
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Core sliding-window check (Redis sorted set)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _check_and_record(
|
||||
r: redis.Redis,
|
||||
user_id: str,
|
||||
window: int,
|
||||
effective_limit: int,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Atomically check the user's upload count and record the new upload.
|
||||
|
||||
Uses a Redis sorted set where each member is a unique timestamp-based ID
|
||||
and the score is the Unix timestamp. Entries older than *window* seconds
|
||||
are pruned on every call so the set never grows unbounded.
|
||||
|
||||
Returns:
|
||||
``None`` if the request is allowed, or a ``dict`` with ``count``,
|
||||
``limit``, and ``retry_after`` if the limit is exceeded.
|
||||
"""
|
||||
key = f"{_KEY_PREFIX}:{user_id}"
|
||||
now = time.time()
|
||||
window_start = now - window
|
||||
|
||||
pipe = r.pipeline(transaction=True)
|
||||
# 1. Remove entries outside the window
|
||||
pipe.zremrangebyscore(key, "-inf", window_start)
|
||||
# 2. Count current entries
|
||||
pipe.zcard(key)
|
||||
# 3. Retrieve the oldest entry's score (to compute retry_after)
|
||||
pipe.zrange(key, 0, 0, withscores=True)
|
||||
results = pipe.execute()
|
||||
|
||||
current_count: int = results[1]
|
||||
oldest_entries: list = results[2]
|
||||
|
||||
if current_count >= effective_limit:
|
||||
# Compute how long until the oldest entry expires from the window.
|
||||
if oldest_entries:
|
||||
oldest_score = oldest_entries[0][1]
|
||||
retry_after = max(1, int((oldest_score + window) - now))
|
||||
else:
|
||||
retry_after = max(1, window // 2)
|
||||
return {
|
||||
"count": current_count,
|
||||
"limit": effective_limit,
|
||||
"retry_after": retry_after,
|
||||
}
|
||||
|
||||
# 4. Record this upload (unique member = timestamp with random suffix)
|
||||
member = f"{now}:{os.urandom(4).hex()}"
|
||||
pipe2 = r.pipeline(transaction=True)
|
||||
pipe2.zadd(key, {member: now})
|
||||
pipe2.expire(key, window + 60) # TTL slightly longer than window
|
||||
pipe2.execute()
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FastAPI dependency
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def require_upload_rate_limit(request: Request) -> None:
|
||||
"""FastAPI dependency that enforces per-user upload rate limits.
|
||||
|
||||
The dependency is designed to **fail open**: if Redis is unavailable the
|
||||
request is allowed through so that uploads are never blocked by a
|
||||
monitoring outage.
|
||||
|
||||
Raises:
|
||||
HTTPException: 429 Too Many Requests when the per-user upload limit
|
||||
is exceeded. The ``Retry-After`` header indicates how many
|
||||
seconds the client should wait before retrying.
|
||||
"""
|
||||
r = _get_redis()
|
||||
if r is None:
|
||||
# Redis unavailable – fail open.
|
||||
return
|
||||
|
||||
# Identify the user (owner_id for multi-user, IP fallback).
|
||||
user_id = get_current_owner_id(request)
|
||||
if not user_id:
|
||||
user_id = f"ip:{request.client.host}" if request.client else "ip:unknown"
|
||||
|
||||
base_limit: int = settings.upload_rate_limit_per_user
|
||||
window: int = settings.upload_rate_limit_window
|
||||
|
||||
# Gather health metrics and compute effective limit.
|
||||
try:
|
||||
queue_depth = _get_queue_depth(r)
|
||||
except Exception: # noqa: BLE001
|
||||
queue_depth = 0
|
||||
|
||||
cpu_load_ratio = _get_cpu_load_ratio()
|
||||
effective_limit, factor, health_reason = compute_effective_limit(base_limit, queue_depth, cpu_load_ratio)
|
||||
|
||||
# Sliding-window check.
|
||||
try:
|
||||
rejection = _check_and_record(r, user_id, window, effective_limit)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("Upload rate-limit check failed (allowing request): %s", exc)
|
||||
return
|
||||
|
||||
if rejection is not None:
|
||||
retry_after = rejection["retry_after"]
|
||||
logger.warning(
|
||||
"Upload rate limit exceeded: user=%s count=%d/%d window=%ds health=%s retry_after=%ds",
|
||||
user_id,
|
||||
rejection["count"],
|
||||
rejection["limit"],
|
||||
window,
|
||||
health_reason,
|
||||
retry_after,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=(
|
||||
f"Upload rate limit exceeded ({rejection['count']}/{rejection['limit']} "
|
||||
f"in {window}s). Retry after {retry_after}s."
|
||||
),
|
||||
headers={"Retry-After": str(retry_after)},
|
||||
)
|
||||
|
||||
if factor < 1.0:
|
||||
logger.info(
|
||||
"Upload allowed with reduced limit: user=%s effective=%d/%d health=%s",
|
||||
user_id,
|
||||
effective_limit,
|
||||
base_limit,
|
||||
health_reason,
|
||||
)
|
||||
+1144
-1
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,44 @@
|
||||
"""Celery task for asynchronous automation hook 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="automation.deliver_hook")
|
||||
def deliver_automation_hook_task(self, url: str, payload: dict[str, Any], secret: str | None = None) -> dict[str, Any]:
|
||||
"""Deliver an automation hook payload to *url* with automatic retries.
|
||||
|
||||
Args:
|
||||
url: Target webhook URL (provided by Zapier / Make.com).
|
||||
payload: The flat Zapier-compatible payload.
|
||||
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 automation hook 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"Automation hook delivery to {url} failed")
|
||||
@@ -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 rule-based document classification.
|
||||
|
||||
This task is executed as a pipeline step (``step_type="classify"``). It
|
||||
applies built-in and user-defined classification rules against the document's
|
||||
filename, OCR text, and existing AI metadata to assign a ``document_type``
|
||||
category.
|
||||
|
||||
The result is stored in the ``ai_metadata`` JSON blob on the
|
||||
:class:`~app.models.FileRecord` (field ``classification``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.database import SessionLocal
|
||||
from app.models import ClassificationRuleModel, FileRecord
|
||||
from app.tasks.retry_config import BaseTaskWithRetry
|
||||
from app.utils import log_task_progress
|
||||
from app.utils.classification_rules import (
|
||||
ClassificationResult,
|
||||
classify_document,
|
||||
db_rule_to_engine_rule,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
STEP_NAME = "classify_document"
|
||||
|
||||
|
||||
def _load_custom_rules(owner_id: str | None) -> list[Any]:
|
||||
"""Load enabled custom classification rules from the database.
|
||||
|
||||
Returns engine-level :class:`ClassificationRule` dataclass instances.
|
||||
Rules are loaded in priority-descending order. System rules
|
||||
(``owner_id IS NULL``) and the user's own rules are both included.
|
||||
"""
|
||||
with SessionLocal() as db:
|
||||
query = db.query(ClassificationRuleModel).filter(ClassificationRuleModel.enabled.is_(True))
|
||||
if owner_id:
|
||||
query = query.filter(
|
||||
(ClassificationRuleModel.owner_id.is_(None)) | (ClassificationRuleModel.owner_id == owner_id)
|
||||
)
|
||||
else:
|
||||
query = query.filter(ClassificationRuleModel.owner_id.is_(None))
|
||||
rules = query.order_by(ClassificationRuleModel.priority.desc()).all()
|
||||
return [db_rule_to_engine_rule(r) for r in rules]
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def classify_document_task(
|
||||
self: Any,
|
||||
file_id: int,
|
||||
owner_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Classify a document using rule-based matching.
|
||||
|
||||
This task:
|
||||
1. Loads the :class:`FileRecord` from the database.
|
||||
2. Gathers filename, OCR text, and existing AI metadata.
|
||||
3. Loads built-in + user-defined classification rules.
|
||||
4. Runs the classification engine.
|
||||
5. Persists the result into ``ai_metadata.classification``.
|
||||
|
||||
Args:
|
||||
file_id: Primary key of the :class:`FileRecord` to classify.
|
||||
owner_id: Owner identifier for loading user-specific rules.
|
||||
|
||||
Returns:
|
||||
Dict with ``category``, ``confidence``, and ``matched_rules``.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"in_progress",
|
||||
f"Starting classification for file {file_id}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
try:
|
||||
with SessionLocal() as db:
|
||||
file_record: FileRecord | None = db.query(FileRecord).filter(FileRecord.id == file_id).first()
|
||||
if file_record is None:
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"failure",
|
||||
f"FileRecord {file_id} not found",
|
||||
file_id=file_id,
|
||||
)
|
||||
return {"status": "error", "detail": "File not found"}
|
||||
|
||||
# Gather inputs
|
||||
filename = file_record.original_filename or ""
|
||||
text = file_record.ocr_text or ""
|
||||
existing_metadata: dict[str, Any] = {}
|
||||
if file_record.ai_metadata:
|
||||
try:
|
||||
existing_metadata = json.loads(file_record.ai_metadata)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
logger.warning("Failed to parse ai_metadata for file %s, starting fresh", file_id)
|
||||
existing_metadata = {}
|
||||
|
||||
# Load custom rules
|
||||
effective_owner = owner_id or file_record.owner_id
|
||||
custom_rules = _load_custom_rules(effective_owner)
|
||||
|
||||
# Run classification engine
|
||||
result: ClassificationResult = classify_document(
|
||||
filename=filename,
|
||||
text=text,
|
||||
metadata=existing_metadata,
|
||||
custom_rules=custom_rules,
|
||||
)
|
||||
|
||||
# Persist result into ai_metadata
|
||||
classification_data = {
|
||||
"category": result.category,
|
||||
"confidence": result.confidence,
|
||||
"matched_rules": [
|
||||
{
|
||||
"rule_name": m.rule_name,
|
||||
"rule_type": m.rule_type,
|
||||
"category": m.category,
|
||||
"confidence": m.confidence,
|
||||
}
|
||||
for m in result.matched_rules
|
||||
],
|
||||
}
|
||||
|
||||
existing_metadata["classification"] = classification_data
|
||||
|
||||
# If no document_type was set yet, populate it from the classification
|
||||
if not existing_metadata.get("document_type"):
|
||||
from app.utils.classification_rules import BUILTIN_CATEGORIES
|
||||
|
||||
existing_metadata["document_type"] = BUILTIN_CATEGORIES.get(
|
||||
result.category, result.category.replace("_", " ").title()
|
||||
)
|
||||
|
||||
file_record.ai_metadata = json.dumps(existing_metadata, ensure_ascii=False)
|
||||
db.commit()
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"success",
|
||||
f"Classified as '{result.category}' with confidence {result.confidence}",
|
||||
file_id=file_id,
|
||||
detail=f"Matched {len(result.matched_rules)} rule(s)",
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"category": result.category,
|
||||
"confidence": result.confidence,
|
||||
"matched_rules": len(result.matched_rules),
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.exception("Classification failed for file %s: %s", file_id, e)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
STEP_NAME,
|
||||
"failure",
|
||||
f"Classification failed: {e}",
|
||||
file_id=file_id,
|
||||
)
|
||||
raise
|
||||
@@ -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,454 @@
|
||||
"""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",
|
||||
"--", # end-of-options separator: prevents file paths from being interpreted as options
|
||||
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}
|
||||
@@ -216,6 +216,30 @@ def embed_metadata_into_pdf(self, local_file_path: str, extracted_text: str, met
|
||||
except Exception as search_exc:
|
||||
logger.warning(f"[{task_id}] Meilisearch indexing failed (non-fatal): {search_exc}")
|
||||
|
||||
# Cache the detected language on the FileRecord and trigger
|
||||
# default-language translation when the document is in a
|
||||
# different language.
|
||||
detected_lang = metadata.get("language") if metadata else None
|
||||
if detected_lang and extracted_text:
|
||||
try:
|
||||
file_record.detected_language = detected_lang
|
||||
db.commit()
|
||||
|
||||
from app.tasks.translate_to_default_language import translate_to_default_language
|
||||
|
||||
translate_to_default_language.delay(
|
||||
file_id,
|
||||
extracted_text,
|
||||
detected_lang,
|
||||
owner_id=file_record.owner_id,
|
||||
)
|
||||
logger.info(
|
||||
f"[{task_id}] Queued default-language translation for file {file_id} "
|
||||
f"(detected: {detected_lang})"
|
||||
)
|
||||
except Exception as trans_exc:
|
||||
logger.warning(f"[{task_id}] Could not queue translation task (non-fatal): {trans_exc}")
|
||||
|
||||
# Persist the metadata into a JSON file with the same base name.
|
||||
# Include file path references for traceability
|
||||
logger.info(f"[{task_id}] Persisting metadata to JSON")
|
||||
|
||||
@@ -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, 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,14 +10,20 @@ 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
|
||||
|
||||
# Import notification utility
|
||||
# Import notification utilities
|
||||
from app.utils.notification import notify_file_processed
|
||||
from app.utils.user_notification import notify_user_document_processed
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -26,13 +32,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 +55,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 +64,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)
|
||||
|
||||
@@ -90,4 +140,15 @@ def finalize_document_storage(self, original_file: str, processed_file: str, met
|
||||
except Exception as e:
|
||||
logger.warning(f"[WARNING] Failed to send file processed notification: {e}")
|
||||
|
||||
# 6. Send per-user notification
|
||||
if owner_id:
|
||||
try:
|
||||
notify_user_document_processed(
|
||||
owner_id=owner_id,
|
||||
filename=os.path.basename(processed_file),
|
||||
file_id=file_id,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"[WARNING] Failed to send per-user processed notification: {e}")
|
||||
|
||||
return {"status": "Completed", "file": processed_file}
|
||||
|
||||
+676
-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.
|
||||
|
||||
+213
-15
@@ -6,17 +6,19 @@ 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
|
||||
from app.tasks.upload_to_s3 import upload_to_s3
|
||||
from app.tasks.upload_to_sftp import upload_to_sftp
|
||||
from app.tasks.upload_to_sharepoint import upload_to_sharepoint
|
||||
from app.tasks.upload_to_webdav import upload_to_webdav
|
||||
from app.utils.config_validator import get_provider_status
|
||||
from app.utils.logging import log_task_progress
|
||||
@@ -25,18 +27,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 +67,78 @@ 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 _should_upload_to_sharepoint():
|
||||
return bool(
|
||||
settings.sharepoint_client_id
|
||||
and settings.sharepoint_client_secret
|
||||
and settings.sharepoint_site_url
|
||||
and (
|
||||
settings.sharepoint_refresh_token
|
||||
or (settings.sharepoint_tenant_id and settings.sharepoint_tenant_id != "common")
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
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 +153,21 @@ def get_configured_services_from_validator():
|
||||
"Email": "email",
|
||||
"OneDrive": "onedrive",
|
||||
"S3 Storage": "s3",
|
||||
"SharePoint": "sharepoint",
|
||||
"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 +176,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 +264,16 @@ 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": "sharepoint",
|
||||
"should_upload": _should_upload_to_sharepoint,
|
||||
"upload_func": upload_to_sharepoint,
|
||||
},
|
||||
{
|
||||
"name": "icloud",
|
||||
"should_upload": _should_upload_to_icloud,
|
||||
"upload_func": upload_to_icloud,
|
||||
},
|
||||
]
|
||||
|
||||
# Optionally get configuration status from validator
|
||||
@@ -236,7 +311,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 +329,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}
|
||||
@@ -0,0 +1,141 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Celery task to translate extracted document text into the default target language.
|
||||
|
||||
This task is triggered after metadata extraction when the detected document
|
||||
language differs from the user's (or system) default document language. The
|
||||
translated text is persisted in ``FileRecord.default_language_text`` so that
|
||||
users can always read a reference copy in their preferred language.
|
||||
|
||||
Other ad-hoc translations are generated on the fly via the ``/api/files/{id}/translate``
|
||||
endpoint and are NOT persisted.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from app.celery_app import celery
|
||||
from app.config import settings
|
||||
from app.database import SessionLocal
|
||||
from app.models import FileRecord, UserProfile
|
||||
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 _resolve_default_language(owner_id: str | None) -> str:
|
||||
"""Return the default document language for the given owner.
|
||||
|
||||
Resolution order:
|
||||
1. ``UserProfile.default_document_language`` (per-user override)
|
||||
2. ``settings.default_document_language`` (global setting)
|
||||
"""
|
||||
if owner_id:
|
||||
with SessionLocal() as db:
|
||||
profile = db.query(UserProfile).filter_by(user_id=owner_id).first()
|
||||
if profile and profile.default_document_language:
|
||||
return profile.default_document_language
|
||||
return settings.default_document_language
|
||||
|
||||
|
||||
@celery.task(base=BaseTaskWithRetry, bind=True)
|
||||
def translate_to_default_language(
|
||||
self,
|
||||
file_id: int,
|
||||
extracted_text: str,
|
||||
detected_language: str,
|
||||
owner_id: str | None = None,
|
||||
) -> dict:
|
||||
"""Translate *extracted_text* into the default document language and persist the result.
|
||||
|
||||
Args:
|
||||
file_id: Primary key of the :class:`FileRecord`.
|
||||
extracted_text: The OCR / refined text in the document's original language.
|
||||
detected_language: ISO 639-1 code of the document's detected language.
|
||||
owner_id: Owner identifier used to resolve per-user language preference.
|
||||
|
||||
Returns:
|
||||
A dict with ``status``, ``target_language``, and the translated text length.
|
||||
"""
|
||||
task_id = self.request.id
|
||||
target_language = _resolve_default_language(owner_id)
|
||||
|
||||
# Nothing to do when the document is already in the target language.
|
||||
if detected_language == target_language:
|
||||
logger.info(
|
||||
f"[{task_id}] Document {file_id} already in target language '{target_language}', skipping translation"
|
||||
)
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"translate_to_default_language",
|
||||
"skipped",
|
||||
f"Document already in {target_language}",
|
||||
file_id=file_id,
|
||||
)
|
||||
return {"status": "skipped", "reason": "already_in_target_language"}
|
||||
|
||||
logger.info(f"[{task_id}] Translating document {file_id} from '{detected_language}' to '{target_language}'")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"translate_to_default_language",
|
||||
"in_progress",
|
||||
f"Translating from {detected_language} to {target_language}",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
try:
|
||||
provider = get_ai_provider()
|
||||
model = settings.ai_model or settings.openai_model
|
||||
translated_text = provider.chat_completion(
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": (
|
||||
f"You are a professional translator. Translate the following text "
|
||||
f"from {detected_language} to {target_language}. "
|
||||
f"Preserve the original formatting, paragraph structure, and meaning. "
|
||||
f"Do not add any commentary or explanation — output ONLY the translated text."
|
||||
),
|
||||
},
|
||||
{"role": "user", "content": extracted_text},
|
||||
],
|
||||
model=model,
|
||||
temperature=0.3,
|
||||
)
|
||||
|
||||
# Persist the translation.
|
||||
with SessionLocal() as db:
|
||||
record = db.query(FileRecord).filter_by(id=file_id).first()
|
||||
if record:
|
||||
record.default_language_text = translated_text
|
||||
record.default_language_code = target_language
|
||||
record.detected_language = detected_language
|
||||
db.commit()
|
||||
logger.info(
|
||||
f"[{task_id}] Stored default-language translation ({len(translated_text)} chars) for file {file_id}"
|
||||
)
|
||||
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"translate_to_default_language",
|
||||
"success",
|
||||
f"Translated {len(extracted_text)} → {len(translated_text)} chars ({detected_language} → {target_language})",
|
||||
file_id=file_id,
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"target_language": target_language,
|
||||
"translated_length": len(translated_text),
|
||||
}
|
||||
|
||||
except Exception as exc:
|
||||
logger.exception(f"[{task_id}] Translation failed for file {file_id}: {exc}")
|
||||
log_task_progress(
|
||||
task_id,
|
||||
"translate_to_default_language",
|
||||
"failure",
|
||||
f"Exception: {exc}",
|
||||
file_id=file_id,
|
||||
)
|
||||
raise
|
||||
@@ -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
|
||||
|
||||
@@ -11,11 +11,12 @@ from email.mime.image import MIMEImage
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
|
||||
import pypdf
|
||||
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__)
|
||||
@@ -23,6 +24,15 @@ logger = logging.getLogger(__name__)
|
||||
# Constants
|
||||
_LOGO_FILENAME = "logo.png"
|
||||
|
||||
# Mapping from PDF metadata keys (with leading slash stripped) to application-specific names.
|
||||
# This mirrors the inverse of the mapping used in app/tasks/embed_metadata_into_pdf.py.
|
||||
_PDF_METADATA_KEY_MAP = {
|
||||
"Title": "filename",
|
||||
"Author": "absender",
|
||||
"Subject": "document_type",
|
||||
"Keywords": "tags",
|
||||
}
|
||||
|
||||
|
||||
def get_email_template(template_name="default.html"):
|
||||
"""
|
||||
@@ -63,9 +73,12 @@ def extract_metadata_from_file(file_path):
|
||||
"""
|
||||
Try to extract metadata from a file using several methods:
|
||||
1. Check for a .json metadata file with the same name
|
||||
2. Extract metadata from PDF if it's embedded
|
||||
2. Extract embedded metadata from PDF using pypdf
|
||||
|
||||
Returns a dictionary of metadata or None if not found
|
||||
JSON metadata takes precedence; embedded PDF metadata fills in any missing
|
||||
fields using the application's standard key mapping (e.g., /Title → filename).
|
||||
|
||||
Returns a dictionary of metadata (may be empty if none found).
|
||||
"""
|
||||
metadata = {}
|
||||
|
||||
@@ -76,12 +89,28 @@ def extract_metadata_from_file(file_path):
|
||||
with open(metadata_path, "r", encoding="utf-8") as f:
|
||||
metadata = json.load(f)
|
||||
logger.info(f"Loaded metadata from external JSON file: {metadata_path}")
|
||||
return metadata
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to load metadata from JSON file: {str(e)}")
|
||||
|
||||
# TODO: For PDF files, try to extract embedded metadata using PyPDF2
|
||||
# This would require additional dependencies, so for now we'll just check for external JSON
|
||||
# Try to extract embedded metadata from PDF
|
||||
if file_path.lower().endswith(".pdf") and os.path.exists(file_path):
|
||||
try:
|
||||
with open(file_path, "rb") as f:
|
||||
pdf_reader = pypdf.PdfReader(f)
|
||||
pdf_metadata = pdf_reader.metadata
|
||||
if pdf_metadata:
|
||||
for key, value in pdf_metadata.items():
|
||||
# Remove the leading slash from PDF metadata keys (e.g., '/Title' -> 'Title')
|
||||
clean_key = key[1:] if key.startswith("/") else key
|
||||
# Map to application-specific key names where possible
|
||||
mapped_key = _PDF_METADATA_KEY_MAP.get(clean_key, clean_key)
|
||||
# Only set if not already present (JSON metadata takes precedence)
|
||||
if mapped_key not in metadata:
|
||||
metadata[mapped_key] = str(value)
|
||||
|
||||
logger.info(f"Extracted embedded metadata from PDF: {file_path}")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to extract metadata from PDF {file_path}: {str(e)}")
|
||||
|
||||
return metadata
|
||||
|
||||
@@ -125,11 +154,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 +168,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 +186,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 +234,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 +265,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.
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user