Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
+2 |
8dae237a0b | ||
|
|
4e2240aada | ||
|
|
f93d20ac42 | ||
|
|
8ff7ff459b | ||
|
|
f48b8ffbf0 | ||
|
|
3560d2f630 | ||
|
|
3aeb691d98 | ||
|
|
e0d6074cd2 | ||
|
|
b1bd3084f0 | ||
|
|
70b28b629e | ||
|
|
3d1e355df7 | ||
|
|
7102a63c82 | ||
|
|
1cea8ec7d4 | ||
|
|
a7a92d2d9b | ||
|
|
2419899ac6 | ||
|
|
60f67c7c17 | ||
|
|
34a55d4524 | ||
|
|
9771898c58 | ||
|
|
752238247c | ||
|
|
3e14524154 | ||
|
|
d8b55afb00 | ||
|
|
62693938a3 | ||
|
|
465d6fe514 | ||
|
|
d6b73ea2f2 | ||
|
|
5774ab4984 | ||
|
|
db05fdaf83 | ||
|
|
d740b545a4 | ||
|
|
678c44c7cd | ||
|
|
5cc55e2278 | ||
|
|
26711c1bcc | ||
|
|
f2cb63140c | ||
|
|
a76a779c01 | ||
|
|
a766521933 | ||
|
|
7da6b82471 | ||
|
|
6089de0b17 | ||
|
|
b87c755574 | ||
|
|
b73538ece7 | ||
|
|
258e9f917b | ||
|
|
6ecba19447 | ||
|
|
90584ab6f3 | ||
|
|
58bc254809 | ||
|
|
0e311a95a7 | ||
|
|
d0e51bde5d | ||
|
|
4dc5c1eb4f | ||
|
|
89669f3fa1 | ||
|
|
f6bd08c852 | ||
|
|
9b577868c8 | ||
|
|
91d9870266 | ||
|
|
d47993385a | ||
|
|
a4eb10269e | ||
|
|
ef77c3d45b | ||
|
|
83f3a9c543 | ||
|
|
d56d74b387 | ||
|
|
0a8a620fb6 | ||
|
|
f162d4de90 | ||
|
|
9f61a6f13c | ||
|
|
f31768e20e | ||
|
|
493f238431 | ||
|
|
3b821e1f3a | ||
|
|
085d3cb1c9 | ||
|
|
116eb7fc55 | ||
|
|
65f55847a1 | ||
|
|
0542df147a | ||
|
|
6cc799b1bb | ||
|
|
b9fc3f367a | ||
|
|
c7b6de6ca4 | ||
|
|
5f76c250f8 | ||
|
|
0e3135f8dc | ||
|
|
7fd94b0e73 | ||
|
|
c4aac0415c | ||
|
|
a27916d1db | ||
|
|
f485309fd6 | ||
|
|
65834432a3 | ||
|
|
46d73c9dcd | ||
|
|
a2875f13c6 | ||
|
|
81383a7df1 | ||
|
|
4790faba73 | ||
|
|
e88e565ab4 | ||
|
|
7ddb9700ff | ||
|
|
b645b0dc23 | ||
|
|
51627555bf | ||
|
|
1824e69a70 | ||
|
|
5127354b3e | ||
|
|
dc6df52a91 | ||
|
|
f246a66810 | ||
|
|
c3c857a3ec | ||
|
|
47329b5032 | ||
|
|
d5e69f182c | ||
|
|
e29d145a1c | ||
|
|
b3ca943da1 | ||
|
|
56c5bc1d34 | ||
|
|
fd25152076 | ||
|
|
51cd43229c | ||
|
|
0d10b946f5 | ||
|
|
28963815d1 | ||
|
|
3e3f138d93 | ||
|
|
4e31fa4427 | ||
|
|
24dd5b461e | ||
|
|
1d501cfa3f | ||
|
|
f6d1969067 | ||
|
|
eb16ae92a5 | ||
|
|
73e28c9393 | ||
|
|
4198f36c01 | ||
|
|
ec9c066961 | ||
|
|
a05a769938 | ||
|
|
e5b5a17426 | ||
|
|
29ee53aaa5 | ||
|
|
60f3ba6b59 | ||
|
|
b7b7b64d31 | ||
|
|
37eba1c5a6 | ||
|
|
8c9f267ad2 | ||
|
|
5afc258c5b | ||
|
|
42694c7c0c | ||
|
|
98c4f264e4 | ||
|
|
f45d0f130e | ||
|
|
98627e42b4 | ||
|
|
f0ec5ee08f | ||
|
|
4a5401b417 | ||
|
|
8d739e2aba | ||
|
|
5087492e25 | ||
|
|
a4d62253df | ||
|
|
7cfb260b8a | ||
|
|
49430de42d | ||
|
|
1be9627dd2 | ||
|
|
f0e0cfcf02 | ||
|
|
55bfc7cbc2 | ||
|
|
4113b15a60 | ||
|
|
e8e655d0de | ||
|
|
e5f31c2e14 | ||
|
|
c5276f62d0 | ||
|
|
8acce144f9 | ||
|
|
e7e752f8e7 | ||
|
|
f44b7a01f5 | ||
|
|
2c7acb9285 | ||
|
|
50363ba66b | ||
|
|
860b90fd17 | ||
|
|
914ccf07ef | ||
|
|
ba83613ff2 | ||
|
|
499129625b | ||
|
|
3271b013a8 | ||
|
|
638c7ab802 | ||
|
|
4023f6722b | ||
|
|
e695d854f2 | ||
|
|
34d569d564 | ||
|
|
e709d6812f | ||
|
|
3dd8255816 | ||
|
|
32cfb5788a | ||
|
|
2dca850cee | ||
|
|
349ea4ea9e | ||
|
|
3c22afc5a6 | ||
|
|
7e453de4f7 | ||
|
|
f1ef09ddc8 | ||
|
|
7b5880ab9e | ||
|
|
3332878321 | ||
|
|
e396af3cc8 | ||
|
|
7cc7b367dc | ||
|
|
5eae0a5cdd | ||
|
|
0e50d2b741 | ||
|
|
7aed64be3c | ||
|
|
43e5905c13 | ||
|
|
128cf41fce | ||
|
|
398718d505 | ||
|
|
bd35809105 | ||
|
|
2e52ad8ff2 | ||
|
|
4d2f189810 | ||
|
|
70a6a24f14 | ||
|
|
2f9e326dba | ||
|
|
a4251d7e45 | ||
|
|
82755acdfd | ||
|
|
c27aa569c8 | ||
|
|
ee1bc5f5db | ||
|
|
3f40d9da70 | ||
|
|
e6297cf414 | ||
|
|
5944eda0ff | ||
|
|
1860874192 | ||
|
|
5dae600ce7 | ||
|
|
f1be85d997 | ||
|
|
ecd74f220c | ||
|
|
8bd23b9145 | ||
|
|
a4ed16999e | ||
|
|
f102060a6d | ||
|
|
26a645f9e6 | ||
|
|
fd93bd3414 | ||
|
|
a3ea7bf043 | ||
|
|
4866bec0f2 | ||
|
|
804f9f3153 | ||
|
|
ee28032fb9 | ||
|
|
37658fd541 | ||
|
|
cced77b584 | ||
|
|
f685edd161 | ||
|
|
18fe17127a | ||
|
|
a209f7f6e0 | ||
|
|
45e49d33e5 | ||
|
|
9a8c4da67d | ||
|
|
39ea7bf63d | ||
|
|
84ec43105c | ||
|
|
cf4218e688 | ||
|
|
33a4d1b412 | ||
|
|
c8ef7b0289 | ||
|
|
c767bcaa73 | ||
|
|
2943955c52 | ||
|
|
715cf9797a | ||
|
|
8979987eed | ||
|
|
cd55c3e212 | ||
|
|
8dba798cce | ||
|
|
9dccd29c94 | ||
|
|
026903399b | ||
|
|
2991d9f1f0 | ||
|
|
611fe0c8a9 | ||
|
|
31406caa79 | ||
|
|
9c64d84ad9 | ||
|
|
40f5b3d135 | ||
|
|
869cf9e848 | ||
|
|
2ddcb30b9a | ||
|
|
96265cf042 | ||
|
|
050c4b97a9 | ||
|
|
d0188f3fe1 | ||
|
|
45f45f5bba | ||
|
|
8936721414 | ||
|
|
d1a0fbe292 | ||
|
|
22cfb3c673 | ||
|
|
51765b619c | ||
|
|
20544d412e | ||
|
|
cb6e77be3e | ||
|
|
d4b90f93bd | ||
|
|
26b8ca5b5e | ||
|
|
57784706e4 | ||
|
|
fc98000aa8 | ||
|
|
21cc828132 | ||
|
|
8172c7e3d5 | ||
|
|
498ff8cdc3 | ||
|
|
e10a00132e | ||
|
|
facb194a07 | ||
|
|
c3c8c605d7 | ||
|
|
3c2c611ba9 | ||
|
|
a359262616 | ||
|
|
d59b933bf2 | ||
|
|
a7d4c53f3a | ||
|
|
25898116ea | ||
|
|
4292358bd5 | ||
|
|
67023037f8 | ||
|
|
e7ff4768f8 | ||
|
|
c47dd7b771 | ||
|
|
4498e6faf2 | ||
|
|
674c1127e2 | ||
|
|
15b89b9218 | ||
|
|
008f1dfbda | ||
|
|
47d413ce7b | ||
|
|
0753409e7b | ||
|
|
b78dabb442 | ||
|
|
83024d00bb | ||
|
|
4f94d21780 | ||
|
|
5ee791d5d2 | ||
|
|
fb5ef978bf | ||
|
|
977d638afe | ||
|
|
de27a12151 | ||
|
|
f6b85700ea | ||
|
|
d40f31982b | ||
|
|
27169124f2 | ||
|
|
b618d84065 | ||
|
|
15f9a8f3f1 | ||
|
|
a2a9a3a42a | ||
|
|
4498c21f4c | ||
|
|
d3df8f1f37 | ||
|
|
96a0b3239b | ||
|
|
36a81ad43b | ||
|
|
71a39dbac1 | ||
|
|
5eab125f13 | ||
|
|
e790e7be7a | ||
|
|
b0df527224 | ||
|
|
92dfa3f2f2 | ||
|
|
406251c2f3 | ||
|
|
ee9db91df0 | ||
|
|
09f6d7ba57 | ||
|
|
674695918e | ||
|
|
588b81eeda | ||
|
|
db7f122cb0 | ||
|
|
aacf95cf76 | ||
|
|
008cd8e6b9 | ||
|
|
c0ac10d5db | ||
|
|
faf935ef52 | ||
|
|
6acaaea59a | ||
|
|
b6db719758 | ||
|
|
bf49358185 | ||
|
|
69b3ec3511 | ||
|
|
a600f67d6b | ||
|
|
09dccdd42f | ||
|
|
be38ca8e81 | ||
|
|
bd3a3635ee | ||
|
|
1dcbfd47fb | ||
|
|
4b8e331333 | ||
|
|
cb0fd6ed41 | ||
|
|
e51b661af0 | ||
|
|
a775fc9b50 | ||
|
|
f9ceb7fa89 | ||
|
|
2112a99b36 | ||
|
|
736a800c5f | ||
|
|
fcedeb9034 | ||
|
|
99f3c554c8 | ||
|
|
e7e006e781 | ||
|
|
803d833908 | ||
|
|
435efa31ce | ||
|
|
8977789177 | ||
|
|
9f1b279e88 | ||
|
|
8e82f0d239 | ||
|
|
c40ea7f29d | ||
|
|
253f416de3 | ||
|
|
53eadb7df7 | ||
|
|
6c243664bc | ||
|
|
4cee67e2be | ||
|
|
730e52a431 | ||
|
|
8c2afb8157 | ||
|
|
ae0316a30e | ||
|
|
a72e8e0223 | ||
|
|
f66b67c8b8 | ||
|
|
b89019a8e1 | ||
|
|
0dd9f462ff | ||
|
|
4dea4fdf54 | ||
|
|
2863e9f8c4 | ||
|
|
a71d927a0c | ||
|
|
640dbb6a28 | ||
|
|
366454e812 | ||
|
|
4578bf52ee | ||
|
|
9190d4b542 | ||
|
|
6fdd19bf14 | ||
|
|
6d6dfbf02c | ||
|
|
342582676a | ||
|
|
60e4d75174 | ||
|
|
4632f200a9 | ||
|
|
9d3e0637c8 | ||
|
|
65ee771fd0 | ||
|
|
0c5399ca53 | ||
|
|
584a9a0920 | ||
|
|
124b7e9154 | ||
|
|
4764dd5d37 | ||
|
|
86472bb445 | ||
|
|
07262fa62c | ||
|
|
a28ea36657 | ||
|
|
3f7fc1a75a | ||
|
|
60676bfdcf | ||
|
|
2734fcad63 | ||
|
|
3d6e5ff8f9 | ||
|
|
1554be9da6 | ||
|
|
d6a9efca68 | ||
|
|
b3e56e0c92 | ||
|
|
15883e5229 | ||
|
|
0e5696de74 | ||
|
|
70c87a1ed1 | ||
|
|
51b200c67b | ||
|
|
e24683d09c | ||
|
|
61558d59c3 | ||
|
|
0f8d7982e1 | ||
|
|
d7a5a903a6 | ||
|
|
16e6cb458f | ||
|
|
04919757e4 | ||
|
|
7025cc94cc | ||
|
|
2d83d0f950 | ||
|
|
354a179f6d | ||
|
|
b898fc0258 | ||
|
|
eb5c95ef8e | ||
|
|
512d090f6c | ||
|
|
b8d4577292 | ||
|
|
18f6ec68b9 | ||
|
|
7bfd9ac8ac | ||
|
|
7c52382c90 | ||
|
|
ebc9ccbe1c | ||
|
|
c8ef5a4f38 | ||
|
|
53583f8d83 | ||
|
|
bae5ff938a | ||
|
|
0638b9f56c | ||
|
|
acaf9ab50c | ||
|
|
d30a0531d4 | ||
|
|
5a2ff8b2e5 | ||
|
|
0d3d824274 | ||
|
|
8e67bf67e1 | ||
|
|
beb938a9ab | ||
|
|
40a7b65695 | ||
|
|
6de9a2af97 | ||
|
|
fe8a3d9f83 | ||
|
|
c6b1c56e9e | ||
|
|
48288e9ce7 | ||
|
|
90319593d0 | ||
|
|
f984c6e79a | ||
|
|
8a0794958d | ||
|
|
378673408e | ||
|
|
e6f38f52c8 | ||
|
|
36d02aa147 | ||
|
|
98570d3547 | ||
|
|
5fd9db8739 | ||
|
|
b10c70cfcf | ||
|
|
eba2b2cd72 | ||
|
|
1c5e84ddf2 | ||
|
|
10b4b86ada | ||
|
|
9cc3ffb4a9 | ||
|
|
3fda6e9eed | ||
|
|
4b35d70078 | ||
|
|
6512e085c4 | ||
|
|
0ad397c048 | ||
|
|
b1e8c7d2aa | ||
|
|
b794d61626 | ||
|
|
a06685a47b | ||
|
|
4777f4fa32 | ||
|
|
1b1d85fe2e | ||
|
|
edb8971c7d | ||
|
|
64da99a322 | ||
|
|
012ce95f27 | ||
|
|
6c2b2f2c3e | ||
|
|
2040095050 | ||
|
|
2388dd7dc3 | ||
|
|
bcb71bb520 | ||
|
|
9bd84258d0 | ||
|
|
66c9bf57da | ||
|
|
abc3bac44a | ||
|
|
2b0cccd5be | ||
|
|
6d736d3c59 | ||
|
|
9ba97890a5 | ||
|
|
58ca51364c | ||
|
|
308fa924a5 | ||
|
|
4c872a8d12 | ||
|
|
925f77f0f5 | ||
|
|
f3f8f9874f | ||
|
|
52a06bd48a | ||
|
|
4567cdc0d9 | ||
|
|
a641325707 | ||
|
|
11f52921dc | ||
|
|
c6ed0b0788 | ||
|
|
16335f866e | ||
|
|
0842acad53 | ||
|
|
0472017cab | ||
|
|
0bdcb5a337 | ||
|
|
1994d65306 | ||
|
|
19274672f9 | ||
|
|
d1df748de3 | ||
|
|
4d058a125b | ||
|
|
f122525310 | ||
|
|
f5f128620a | ||
|
|
06635898d0 | ||
|
|
0ded0b7069 | ||
|
|
0fa246a1c3 | ||
|
|
8cb47aebae | ||
|
|
bfc606a9e3 | ||
|
|
7a21933d10 | ||
|
|
9364e2fb74 | ||
|
|
350d52f515 | ||
|
|
05252e19b5 | ||
|
|
abe42eaf09 | ||
|
|
f2f4baa89a | ||
|
|
15ae3f588b | ||
|
|
75932be880 | ||
|
|
98c7ed965b | ||
|
|
4d50001c41 | ||
|
|
857d7e6f37 | ||
|
|
08ff3bd30f | ||
|
|
debcc3a652 | ||
|
|
7b78c641fe | ||
|
|
94f877ff32 | ||
|
|
aa2f7fbe52 | ||
|
|
cdc2b3bf85 | ||
|
|
90ca2e9b0f | ||
|
|
76ece4049e | ||
|
|
1cf1b2ca17 | ||
|
|
8b6fa1f4ab | ||
|
|
58dcc1b33f | ||
|
|
cb83831041 | ||
|
|
adf7af34ff | ||
|
|
d933991904 | ||
|
|
440b64088b | ||
|
|
0aebdd5f83 | ||
|
|
261aec8c86 | ||
|
|
61cfaab915 | ||
|
|
968462609f | ||
|
|
601bb78358 | ||
|
|
631bd20c35 | ||
|
|
c0fcbc5b4c | ||
|
|
be21db7069 | ||
|
|
8507e5eb0d | ||
|
|
cffbc3558e | ||
|
|
d738044f47 | ||
|
|
e1cdd7e4fe | ||
|
|
94145c99ae | ||
|
|
7eae377c01 | ||
|
|
2376258f02 | ||
|
|
5c4062c648 | ||
|
|
cb7154cedf | ||
|
|
e4de5c5ad1 | ||
|
|
c24a4da17d | ||
|
|
f949d17db1 | ||
|
|
6d7744c219 | ||
|
|
6f06b3d5ed | ||
|
|
7ce1e9415a | ||
|
|
a9c5c787b9 | ||
|
|
f7e07f3ca1 | ||
|
|
9f946dec60 | ||
|
|
7bcfafa6c5 | ||
|
|
4a70aaa162 | ||
|
|
12f0ad28bf | ||
|
|
cf60b1882f | ||
|
|
c479c22438 | ||
|
|
3841e85abb | ||
|
|
24370d5c40 | ||
|
|
69171a4c8b | ||
|
|
bd8aa3b6a0 | ||
|
|
fe7e002fea | ||
|
|
16ee4ac7a2 | ||
|
|
70285fb6ca | ||
|
|
ade617efa8 | ||
|
|
f0d48a4295 | ||
|
|
139e764b2f | ||
|
|
3a4b862e81 | ||
|
|
637cd136c2 | ||
|
|
1dc647f43b | ||
|
|
6890618221 | ||
|
|
a3238aa79f | ||
|
|
108a019cb8 | ||
|
|
1c25b06dca | ||
|
|
36c3fc58b5 | ||
|
|
0f0ba7dadd | ||
|
|
5d7766e1b6 | ||
|
|
eca51269bb | ||
|
|
d577ff1e4a | ||
|
|
6a9d67b5bb | ||
|
|
ebb7ce2092 | ||
|
|
9a6bf78e14 | ||
|
|
59171daa35 | ||
|
|
945275faae | ||
|
|
2e165926de | ||
|
|
dfc2dc2c0b | ||
|
|
d784eb1e9b | ||
|
|
6a004205d8 | ||
|
|
ee9099cab9 | ||
|
|
ee901fcd2c | ||
|
|
7ffcd3908e | ||
|
|
0dcd6ac983 | ||
|
|
93415a48e8 | ||
|
|
f8b3a32caf | ||
|
|
2ae47cf200 | ||
|
|
adcbba34f8 | ||
|
|
218bd7a402 | ||
|
|
464462b22e | ||
|
|
a1aceb5f87 | ||
|
|
52e227f425 | ||
|
|
ea515fa26e | ||
|
|
cc8b5055f2 | ||
|
|
d8fa0f426a | ||
|
|
aa59b32374 | ||
|
|
877bc23afc | ||
|
|
0ad448b8a4 | ||
|
|
93407ba316 | ||
|
|
78dbad5e1e | ||
|
|
5b026b2a3e | ||
|
|
bb3526f4e4 | ||
|
|
17c819a3c2 | ||
|
|
4c8615f01c | ||
|
|
85411e4867 | ||
|
|
4d67c817ec | ||
|
|
8b4ea5bb78 | ||
|
|
b44eacbc5a | ||
|
|
4f0e574201 | ||
|
|
53b8a1f71b | ||
|
|
9a2c60d595 | ||
|
|
5df4277216 | ||
|
|
6769b1967c | ||
|
|
2c80d95c53 | ||
|
|
124e1ea4a7 | ||
|
|
7674e4f093 | ||
|
|
0afb8f681b | ||
|
|
157ae917eb | ||
|
|
f593f92f18 | ||
|
|
af0b7d4683 | ||
|
|
00cb7f5104 | ||
|
|
068e52f877 | ||
|
|
fe772d95e2 | ||
|
|
ecba37070d | ||
|
|
f23296b22d | ||
|
|
10f06a64fe | ||
|
|
8f3144adb5 | ||
|
|
6089a55da6 | ||
|
|
c81b3ef9ce | ||
|
|
de5e0fbc00 | ||
|
|
c53cd78dcb | ||
|
|
9793315e0b | ||
|
|
694fb3776f | ||
|
|
adcc50d337 | ||
|
|
6b66cb5ef6 | ||
|
|
58e78e8946 | ||
|
|
b8ea267f8e | ||
|
|
de3317e26b | ||
|
|
fcf7208352 | ||
|
|
b171b0216b | ||
|
|
b062235d0c | ||
|
|
3107a5363d | ||
|
|
c0385f60ba | ||
|
|
30068afd78 | ||
|
|
e3f3929198 | ||
|
|
54f7861b2e | ||
|
|
68973f9b39 | ||
|
|
a361df840e | ||
|
|
b99f8dabdd | ||
|
|
be6bf76105 | ||
|
|
38cc6e4762 | ||
|
|
e53152123f | ||
|
|
07e650c787 | ||
|
|
753589e51c | ||
|
|
3dea69f658 | ||
|
|
bef5ec2cea | ||
|
|
9465e2918b | ||
|
|
486c004cbb | ||
|
|
14d876c259 | ||
|
|
5787c969f6 | ||
|
|
a229f9ea42 | ||
|
|
bcd313c363 | ||
|
|
02d9d07900 | ||
|
|
bdb7d48cff | ||
|
|
bc9ee7b4d9 | ||
|
|
b535275b17 | ||
|
|
f9756de693 | ||
|
|
b73010bb36 | ||
|
|
bc5b3ec6b8 | ||
|
|
7611762e04 | ||
|
|
1eef5b4f6a | ||
|
|
a43d98965b | ||
|
|
566e25569e | ||
|
|
47e47e42af | ||
|
|
f9d38a073f | ||
|
|
47ab4c71d5 | ||
|
|
39100eca49 | ||
|
|
6862d618ee | ||
|
|
157ff57c40 | ||
|
|
d85b52bfc2 | ||
|
|
05f314bae4 | ||
|
|
ee0d9b7915 | ||
|
|
f3402d3f1f | ||
|
|
dbd0d7d742 | ||
|
|
afa0609ece | ||
|
|
1b1abdd30c | ||
|
|
4a8f995c3f | ||
|
|
7ea1e9cbd0 | ||
|
|
e34ed72e1e | ||
|
|
865880a0b1 | ||
|
|
3a6b5ebb5f | ||
|
|
cee645017d | ||
|
|
035b981e11 | ||
|
|
8da29566a1 | ||
|
|
f2217da94e | ||
|
|
b312318a99 | ||
|
|
0a87c1ecd0 | ||
|
|
a87f015246 | ||
|
|
f1c1004225 | ||
|
|
06657b8109 | ||
|
|
a407a7f1c0 | ||
|
|
418bd05ae0 | ||
|
|
86cce2cd88 | ||
|
|
8970923940 | ||
|
|
83fad5e9f7 | ||
|
|
e4e69a10ec | ||
|
|
c6a1469fad | ||
|
|
97cc94756e | ||
|
|
cb73257f14 | ||
|
|
61366cbcda | ||
|
|
c3e1d2d894 | ||
|
|
0bfacca0a0 | ||
|
|
1364df0913 | ||
|
|
352391fa76 | ||
|
|
2cb28369b7 | ||
|
|
3f350f8659 | ||
|
|
9d8f590fc5 | ||
|
|
defeddf21b | ||
|
|
c97767424f | ||
|
|
d513eb8c4d | ||
|
|
265d1b2824 | ||
|
|
caf3362be8 | ||
|
|
3e952044bd | ||
|
|
7aa7bbc390 | ||
|
|
a97f5adf95 | ||
|
|
138c4cbfcf | ||
|
|
7a3c5c0f8a | ||
|
|
3e513be963 | ||
|
|
f78b238b40 | ||
|
|
bbbe2b66b4 | ||
|
|
63a0befd3c | ||
|
|
710320601a | ||
|
|
67e26fd3af | ||
|
|
2c35bdbcf5 | ||
|
|
abe865f4f3 | ||
|
|
e1d04be897 | ||
|
|
f419830a52 | ||
|
|
61bbb99d9e | ||
|
|
6c159a97b7 | ||
|
|
124ad948fe | ||
|
|
947dcd34bd | ||
|
|
710b5270a1 | ||
|
|
2bff50f736 | ||
|
|
e24299e66d | ||
|
|
f047b6b3ae | ||
|
|
368912ca62 | ||
|
|
b1048fc9bc | ||
|
|
9bb226dc52 | ||
|
|
0948235c3b | ||
|
|
bd456ed10b | ||
|
|
223c14f48b | ||
|
|
d0c3180376 | ||
|
|
3ceaa107ab | ||
|
|
144d8b1bb7 | ||
|
|
989938856f | ||
|
|
8913f37c3d | ||
|
|
80b5896b70 | ||
|
|
967b1137dc | ||
|
|
8cd3bd7997 | ||
|
|
ce0ca894fe | ||
|
|
d1975b740b | ||
|
|
459a60a242 | ||
|
|
9a269ec8ab | ||
|
|
d7efdcce2b | ||
|
|
885c94bda8 | ||
|
|
2e1ef805ff | ||
|
|
95b65ff751 | ||
|
|
35bc831077 | ||
|
|
5d4505c685 | ||
|
|
b4f340806a | ||
|
|
7b2f597b30 | ||
|
|
e303c3da3b | ||
|
|
bc5d519c4f | ||
|
|
7cdff6b1e2 | ||
|
|
b04de83c20 | ||
|
|
dfa2511199 | ||
|
|
d4faa5a5ea | ||
|
|
2108f420ea | ||
|
|
42ecdb5407 | ||
|
|
626fcff417 | ||
|
|
e6b00a8905 | ||
|
|
03c6caac1f | ||
|
|
29160741a3 | ||
|
|
7820a311ba | ||
|
|
5eb9b58488 | ||
|
|
51a2d2b701 | ||
|
|
044fd1bd15 | ||
|
|
70a31a9a57 | ||
|
|
c7d1d1e390 | ||
|
|
2d0b94794f | ||
|
|
fbf315e624 | ||
|
|
b9c0a9c3bf | ||
|
|
6d9996e599 | ||
|
|
7806cd5aef | ||
|
|
b3622474d7 | ||
|
|
d8bb8c58d0 | ||
|
|
d93cb3658d | ||
|
|
200fb093b1 | ||
|
|
4ab831b259 | ||
|
|
576ee92438 | ||
|
|
af4500e504 | ||
|
|
016928722c | ||
|
|
73b69ae408 | ||
|
|
80376a3fdc | ||
|
|
305e591ec2 | ||
|
|
39deadcab1 | ||
|
|
2153c8ec9f | ||
|
|
a70c718a0d | ||
|
|
c73efab192 | ||
|
|
ce54b1df23 | ||
|
|
16701befe7 | ||
|
|
8a6af40d9f | ||
|
|
9cf6108527 | ||
|
|
1c1c1c3100 | ||
|
|
def954134c | ||
|
|
c85afce702 | ||
|
|
a25ecfa856 | ||
|
|
47b007ef19 | ||
|
|
04fae8b357 | ||
|
|
1850a985b5 | ||
|
|
339ed1d72e | ||
|
|
fa1ebfa4fd | ||
|
|
0820abbc64 | ||
|
|
b94e1c9458 | ||
|
|
fe58ef69d9 | ||
|
|
cc6b51e5ae | ||
|
|
cd2c315495 | ||
|
|
4b3ed3e802 | ||
|
|
8cd2157564 | ||
|
|
aaa49bdd6d | ||
|
|
8da02c669e | ||
|
|
828656b35f | ||
|
|
3b97c8d89b | ||
|
|
f5ea1ce250 | ||
|
|
a181b4a731 | ||
|
|
114f709337 | ||
|
|
a6fb5a0460 | ||
|
|
7ef181bc13 | ||
|
|
49a2e5bf57 | ||
|
|
4403c7b6c2 | ||
|
|
b081e33c0a | ||
|
|
f4c38e6001 | ||
|
|
c40f26946f | ||
|
|
627b063b88 | ||
|
|
f962bae983 | ||
|
|
e08341dab3 | ||
|
|
890949abe6 | ||
|
|
6e43861c0c | ||
|
|
ad275351b6 | ||
|
|
7d45459a47 | ||
|
|
5af24b3ebe | ||
|
|
a36692b4a2 | ||
|
|
ca2aaf0321 | ||
|
|
79f0437980 | ||
|
|
10daa64d5b | ||
|
|
e0d4c3ec92 | ||
|
|
65fbbf5e35 | ||
|
|
10baa6e781 | ||
|
|
3de14a53c2 | ||
|
|
fe5c02331b | ||
|
|
d040953c76 | ||
|
|
b5c3395f79 | ||
|
|
ed9ab65b5e | ||
|
|
1a2b360d3d | ||
|
|
4f6cb771f1 | ||
|
|
75683e5197 | ||
|
|
8ea35e3bb4 | ||
|
|
44349fb62b | ||
|
|
fe1941c13a | ||
|
|
933a3bbbd3 | ||
|
|
3909b62ffc | ||
|
|
bec227da30 | ||
|
|
11487d66fc | ||
|
|
395098c6f1 | ||
|
|
72951324df | ||
|
|
0c42cd2c01 | ||
|
|
c701ebe07b | ||
|
|
b338850cc1 | ||
|
|
64957db7b3 | ||
|
|
d7147d6cdd | ||
|
|
6137f7cb7e | ||
|
|
832d0181b6 | ||
|
|
d1dd449f63 | ||
|
|
2751a0f0b6 | ||
|
|
a9e9fe7899 | ||
|
|
702906aee7 | ||
|
|
fe837d80e7 | ||
|
|
860a0b414e | ||
|
|
9c9a18d6d4 | ||
|
|
67893b9a57 | ||
|
|
2e8c4da17b | ||
|
|
ff9f761d65 | ||
|
|
6863ca482c | ||
|
|
5645d5bccc | ||
|
|
201b93bfcc | ||
|
|
0c2e4270bc | ||
|
|
80ad5fd2d0 | ||
|
|
9904566513 | ||
|
|
2054ee0b73 | ||
|
|
93bab8d822 | ||
|
|
259d5ca596 | ||
|
|
597883a179 | ||
|
|
387225eb8b | ||
|
|
c83a42198d | ||
|
|
2cacc2e649 | ||
|
|
c9a78e5476 | ||
|
|
2cbba2a28a | ||
|
|
62ab30f593 | ||
|
|
0fff2fbcab | ||
|
|
fcff9c3afd | ||
|
|
d415edcfcd | ||
|
|
5f304e57d2 | ||
|
|
0b851cf55a | ||
|
|
ff86283be0 | ||
|
|
e9011113b4 | ||
|
|
1b89bee098 | ||
|
|
c436e0366c | ||
|
|
1db36b5eda | ||
|
|
3569280c0b | ||
|
|
a0d6c209c3 | ||
|
|
73617ec7fa | ||
|
|
391a4878e6 | ||
|
|
6d7f21b57b | ||
|
|
fe604a8a9b | ||
|
|
a9d8348cf9 | ||
|
|
c37c0e3490 | ||
|
|
ddedceb7ad | ||
|
|
18865a9fef | ||
|
|
769ef856bc | ||
|
|
ed1b959bc6 | ||
|
|
d2b38127d0 | ||
|
|
3d535db304 | ||
|
|
234306ff57 | ||
|
|
ae28e7d245 | ||
|
|
39b87d9683 | ||
|
|
e83f668107 | ||
|
|
7dda8025fc | ||
|
|
1357dc6737 | ||
|
|
43c30428a6 | ||
|
|
668bd44485 | ||
|
|
a3de0bcc58 | ||
|
|
aed2f69efe | ||
|
|
30ae519226 | ||
|
|
499ca282e5 | ||
|
|
40d90286b6 | ||
|
|
2d27ef4ece | ||
|
|
e9b5eb6ed3 | ||
|
|
6b462ff121 | ||
|
|
c3bac9aa62 | ||
|
|
f7226333c3 | ||
|
|
fc5f399573 | ||
|
|
ff8cf80fb5 | ||
|
|
54cefedf53 | ||
|
|
9440d09114 | ||
|
|
5bb1c42fa8 | ||
|
|
242b3f0c01 | ||
|
|
144c0f3d76 | ||
|
|
18401de254 | ||
|
|
c71beb0a7d | ||
|
|
f5bf2a2ed7 | ||
|
|
ab3f03bbd5 | ||
|
|
5ac502e93f | ||
|
|
c60b0fa0e3 | ||
|
|
9544a80aa0 | ||
|
|
83b17e2ac8 | ||
|
|
3a6c88ade9 | ||
|
|
3be06132db | ||
|
|
bbbcf27dd5 | ||
|
|
cfa16e1a37 | ||
|
|
f60d386b74 | ||
|
|
0324a1bbdd | ||
|
|
a677b212d9 | ||
|
|
179a4ad9ea | ||
|
|
2d82d260cc | ||
|
|
e7a9988893 | ||
|
|
6b01f96eac | ||
|
|
965f242d16 | ||
|
|
758d8fcf31 | ||
|
|
0f8b339f6d | ||
|
|
5d821d21f3 | ||
|
|
d6d9d1c535 | ||
|
|
44ab77b4f5 | ||
|
|
646b64a318 | ||
|
|
bbab64b53e | ||
|
|
4731ccb73c | ||
|
|
4737e1f118 | ||
|
|
7ea6afdf95 | ||
|
|
4654ecbf1b | ||
|
|
527d36e13a | ||
|
|
d7d05a4717 | ||
|
|
419ea1c346 | ||
|
|
59214538bb | ||
|
|
eca9b405eb | ||
|
|
58d685eea4 | ||
|
|
44ed941a5d | ||
|
|
46229a93ce | ||
|
|
50eff6a672 | ||
|
|
1cb74b0bf7 | ||
|
|
c303388296 | ||
|
|
5a08084899 | ||
|
|
b1f292965c | ||
|
|
819ea0d9be | ||
|
|
1f77691b01 | ||
|
|
50e6a19957 | ||
|
|
cb0165827f | ||
|
|
c5225039ab | ||
|
|
f2c3fff278 | ||
|
|
3271a5277c | ||
|
|
8b2160f2f7 | ||
|
|
bee13f72ad | ||
|
|
64ff15a536 | ||
|
|
345f3e3559 | ||
|
|
636ab99ad8 | ||
|
|
f0c71e5a6d | ||
|
|
87d33f6e18 | ||
|
|
fd91fa433a | ||
|
|
484ba91b07 | ||
|
|
acb2147024 | ||
|
|
ace69bba75 | ||
|
|
5beb37c57c | ||
|
|
50f95a4f1a | ||
|
|
39e3f8fb81 | ||
|
|
b2413f914a | ||
|
|
9dff497abf | ||
|
|
e3f21d6c3b | ||
|
|
184e921930 | ||
|
|
5ee5093259 | ||
|
|
81781e6495 | ||
|
|
82959cec88 | ||
|
|
9478c5e7ac | ||
|
|
62e7e0bc09 | ||
|
|
7a16e495dd | ||
|
|
958fbdd5c0 | ||
|
|
5c403fb829 | ||
|
|
538501c88d | ||
|
|
0b6c92baa7 | ||
|
|
64ec73635b | ||
|
|
b36e55cf1f | ||
|
|
2461121637 | ||
|
|
e6fe3ba8ef | ||
|
|
0b867590a8 | ||
|
|
3c8d658160 | ||
|
|
176f9a7816 | ||
|
|
3d99de6771 | ||
|
|
a52e6c2d57 | ||
|
|
1808d7fd2f | ||
|
|
f4a1d99f00 | ||
|
|
8f49725aa5 | ||
|
|
febc66ef2b | ||
|
|
3761b3ac28 | ||
|
|
140ab270af | ||
|
|
6ab452a452 | ||
|
|
8962afd586 | ||
|
|
22f074cf59 | ||
|
|
ec4fe4f390 | ||
|
|
32c68e000b | ||
|
|
1ac3dd4a89 | ||
|
|
55c489146c | ||
|
|
ffcf97e3e1 | ||
|
|
95bde946ba | ||
|
|
895c805e62 | ||
|
|
2ed3055c42 | ||
|
|
1792f668f2 | ||
|
|
1d3d3b2d94 | ||
|
|
9044abf3bb | ||
|
|
424dba443c | ||
|
|
aa649bec6b | ||
|
|
a8a3098782 | ||
|
|
238e9da209 | ||
|
|
49a1b37e5d | ||
|
|
c035ff7d14 | ||
|
|
2558fe1a3b | ||
|
|
e7848ec712 | ||
|
|
39e5422d93 | ||
|
|
f6bd54fb1f | ||
|
|
d9fd2a3f30 | ||
|
|
824eeba56c | ||
|
|
e61406c825 | ||
|
|
4853ededca | ||
|
|
8c127a4814 | ||
|
|
6eba27ee9c | ||
|
|
8f0658e64f | ||
|
|
4b3543d3c0 | ||
|
|
d1b39da911 | ||
|
|
053a33631f | ||
|
|
becac2b2b7 | ||
|
|
342aa84bbe | ||
|
|
9e81e1dda1 | ||
|
|
3ad2ea6f28 | ||
|
|
bab64c9d52 | ||
|
|
0185f3340d | ||
|
|
1cd26372fb | ||
|
|
0ca2e46ade | ||
|
|
30a13b9b2f | ||
|
|
f651809001 | ||
|
|
c341f97cfe | ||
|
|
32aabe6bae | ||
|
|
3c54863414 | ||
|
|
ad9fbfc1af | ||
|
|
29217cb430 | ||
|
|
c0096b2a53 | ||
|
|
5b9efeef4d | ||
|
|
e0087acfb4 | ||
|
|
1f474187a7 | ||
|
|
2beeeb90c2 | ||
|
|
d016cc5771 | ||
|
|
16e567df57 | ||
|
|
1542dad51a | ||
|
|
2ef55972ff | ||
|
|
75c5d9b179 | ||
|
|
713fe1afa7 | ||
|
|
f95cff0895 | ||
|
|
a0dbd41551 | ||
|
|
bf0fb1c449 | ||
|
|
7043751ca4 | ||
|
|
2d5ebf962a | ||
|
|
b48594a166 | ||
|
|
74e771fec6 | ||
|
|
08f1c823ad | ||
|
|
b559606387 | ||
|
|
914c7ba876 | ||
|
|
96ca47ac9f | ||
|
|
a0c82c8e4c | ||
|
|
631e30e22d | ||
|
|
c114fd6876 | ||
|
|
c2172e43eb | ||
|
|
1ad3656872 | ||
|
|
ff7f38d343 | ||
|
|
bc482b9cce | ||
|
|
4c94f5d434 | ||
|
|
4b9f821b58 | ||
|
|
35598b8017 | ||
|
|
45e23c3ad0 | ||
|
|
5759917f54 | ||
|
|
d247adb60c | ||
|
|
8c713a171d | ||
|
|
7e42d727e8 | ||
|
|
9f7dd31e12 | ||
|
|
5d4547f934 | ||
|
|
5522b91c32 | ||
|
|
3242dad8ae | ||
|
|
6d8a6e6d8b | ||
|
|
8265422ba0 | ||
|
|
10c13b686c | ||
|
|
b1dc58ddb7 | ||
|
|
a9312d2537 | ||
|
|
4228bf71c4 | ||
|
|
ac620118c1 | ||
|
|
d650c987ec | ||
|
|
092a358b3c | ||
|
|
ae05586fda | ||
|
|
f5e5632afc | ||
|
|
2a804541e0 | ||
|
|
8c485b260f | ||
|
|
d664922feb | ||
|
|
9950cc8c28 | ||
|
|
3db6d49e57 | ||
|
|
6d67ac371d | ||
|
|
326599b8db | ||
|
|
c5c31ab769 | ||
|
|
43eb2351d2 | ||
|
|
0a700aafe4 | ||
|
|
91a0301c9e | ||
|
|
6ac593209c | ||
|
|
12bea8cd88 | ||
|
|
1dfe546b6b | ||
|
|
ff837031e4 | ||
|
|
139f02a9d9 | ||
|
|
4bef69cc63 | ||
|
|
723185c22f | ||
|
|
35763a352c | ||
|
|
27c76c677a | ||
|
|
2f1344d619 | ||
|
|
8bfab327ec | ||
|
|
56246324b2 | ||
|
|
af5661c2c8 | ||
|
|
f872a178bc | ||
|
|
3dd44c4f19 | ||
|
|
094ed0b48c | ||
|
|
9b55343509 | ||
|
|
8a7f698e9d | ||
|
|
990c638f6c | ||
|
|
a0195cd5ae | ||
|
|
e9d852545c | ||
|
|
49c36238d0 | ||
|
|
74988189b8 | ||
|
|
b8112d72b9 | ||
|
|
71ccedd2bf | ||
|
|
4225791313 | ||
|
|
e5cd1b479b | ||
|
|
61d44aa773 | ||
|
|
ef036529b5 | ||
|
|
e0bdef85ab | ||
|
|
2ce935bdb1 | ||
|
|
23b1e2cca4 | ||
|
|
05b8768fb9 | ||
|
|
173d5631ca | ||
|
|
34cd3d79e8 | ||
|
|
aede1b7a08 | ||
|
|
15b893e651 | ||
|
|
4896d30281 | ||
|
|
9be45f49e4 | ||
|
|
f1053d94c7 | ||
|
|
10cfddccd7 | ||
|
|
e5e39be90f | ||
|
|
337109e99c | ||
|
|
8adf2b33b4 | ||
|
|
1984ce42aa | ||
|
|
656de56a3e | ||
|
|
0b2abe6cb8 | ||
|
|
f364b2d205 | ||
|
|
519ff40cb6 | ||
|
|
7c7fe44328 | ||
|
|
15b5f97f89 | ||
|
|
ca0983f76b | ||
|
|
ef04a704ce | ||
|
|
5fda814669 | ||
|
|
c26e8110af | ||
|
|
9e85055b8b | ||
|
|
6d17de6c67 | ||
|
|
f4e99c80f6 | ||
|
|
09dc28df1e | ||
|
|
c748c3ede7 | ||
|
|
33308022f0 | ||
|
|
f96e8f04fc | ||
|
|
38ae91ae23 | ||
|
|
88401e91c7 | ||
|
|
8c5cfa530d | ||
|
|
d215e46315 | ||
|
|
e10e7d056e | ||
|
|
7a7d902238 | ||
|
|
24179cde2f | ||
|
|
3ae4c618e1 | ||
|
|
4a0d893995 | ||
|
|
b780d5c556 | ||
|
|
911eecac85 | ||
|
|
319d3e8856 | ||
|
|
e1b3e7252c | ||
|
|
f20cc6d7e6 | ||
|
|
f1a1e64d2e | ||
|
|
58e923fe00 | ||
|
|
f2aca781c8 | ||
|
|
ce51c481b8 | ||
|
|
9a2595f070 | ||
|
|
393c0071dc | ||
|
|
7e224e4a53 | ||
|
|
883f1dda0f | ||
|
|
12bad452fa | ||
|
|
5de60dc922 | ||
|
|
a30b106ea3 | ||
|
|
d33ad462aa | ||
|
|
3b61562c82 | ||
|
|
e5d88be4f3 | ||
|
|
b36f8d9314 | ||
|
|
626d236d13 | ||
|
|
79ecbfc757 | ||
|
|
a9b8677cc0 | ||
|
|
0f3f68b0c4 | ||
|
|
abc9b63093 | ||
|
|
64fa26bd28 | ||
|
|
2487c84f1f | ||
|
|
163211a367 | ||
|
|
f027a01ab2 | ||
|
|
370a677a38 | ||
|
|
d1d1efe212 | ||
|
|
b7549d2f6c | ||
|
|
589c4e64c1 | ||
|
|
20de5a87da | ||
|
|
ca6b18ab5c | ||
|
|
97a3b1528d | ||
|
|
d01b1d4880 | ||
|
|
df6e38039f | ||
|
|
b4c3f54f96 | ||
|
|
73776d54b8 | ||
|
|
7bda6bf767 | ||
|
|
ddcec9842f | ||
|
|
4d5b7b3014 | ||
|
|
9886ebb97f | ||
|
|
0b05b2fc7e | ||
|
|
49e7eade15 | ||
|
|
7a7a25766c | ||
|
|
9fc1658085 | ||
|
|
5297dceb2a | ||
|
|
59afbd6f92 | ||
|
|
f7af3f010e | ||
|
|
bb40724d45 | ||
|
|
0e64b31adb | ||
|
|
850a864b02 | ||
|
|
9468d92553 | ||
|
|
c3dc5d5984 | ||
|
|
5291b3dca2 | ||
|
|
87d0c112fa | ||
|
|
2a11175f22 | ||
|
|
3238d94a0e | ||
|
|
8919d8a82a | ||
|
|
ea4ef28da5 | ||
|
|
2ffd8d9277 | ||
|
|
e8499ccdd1 | ||
|
|
0dcbd05e24 | ||
|
|
c653e4ec54 | ||
|
|
8cf32ae2a7 | ||
|
|
c6af296b60 | ||
|
|
531ac70ce5 | ||
|
|
423d8b1817 | ||
|
|
633505460a | ||
|
|
c8bf390680 | ||
|
|
da46c1bbd2 | ||
|
|
05ae44b98d | ||
|
|
a40808579f | ||
|
|
4d024c91d6 | ||
|
|
ccb71a7322 | ||
|
|
efe5416f83 | ||
|
|
a4281f6a7f | ||
|
|
2372b70031 | ||
|
|
d02e826c9d | ||
|
|
dddac2b0ca | ||
|
|
0da57149ae | ||
|
|
0bebb260bf | ||
|
|
390d7663b0 | ||
|
|
ce3a615442 | ||
|
|
4e0cb88583 | ||
|
|
9b30e8f689 | ||
|
|
97331bf11d | ||
|
|
96c07f44a8 | ||
|
|
60ada21c15 | ||
|
|
ab20745ee5 | ||
|
|
f376d4f378 | ||
|
|
89fddcc741 | ||
|
|
773787c74c | ||
|
|
5c1c9a4dcb | ||
|
|
9b925a115a | ||
|
|
e5035ea31e | ||
|
|
c8cbdc8f7f | ||
|
|
64c37ab968 | ||
|
|
f7c5965a70 | ||
|
|
46aa54b7dc | ||
|
|
27944cf7ca | ||
|
|
1973115678 | ||
|
|
a38ad8fc42 | ||
|
|
c2207887b3 | ||
|
|
3fabc085cc | ||
|
|
2f584c9f88 | ||
|
|
ba18d6250a | ||
|
|
4331029926 | ||
|
|
3e56261c5e | ||
|
|
30f72672fa | ||
|
|
4aedfdc547 | ||
|
|
e3a8257690 | ||
|
|
a73cdf4288 | ||
|
|
c259c87806 | ||
|
|
0044902c08 | ||
|
|
cd31b8301b | ||
|
|
8fd5c06e5b | ||
|
|
97afe3bc58 | ||
|
|
3567054325 | ||
|
|
f236192fe1 | ||
|
|
9b1fd86aa7 | ||
|
|
55169e69c0 | ||
|
|
e3e4e1d9d3 | ||
|
|
c2f5cb542e | ||
|
|
1034b74abd | ||
|
|
0a44d80252 | ||
|
|
48a0abb40f | ||
|
|
68c77295bd | ||
|
|
3c7f9aa6a4 | ||
|
|
f7406ff576 | ||
|
|
65904b867e | ||
|
|
e2d09ac361 | ||
|
|
3ae44d11a5 | ||
|
|
aa8c2959ca | ||
|
|
b147616080 | ||
|
|
9747b07ca5 | ||
|
|
0f78451c2b | ||
|
|
42763cbbd8 | ||
|
|
d193c143a5 | ||
|
|
26460917c4 | ||
|
|
fd97ae9bd7 | ||
|
|
c3fa0f30fe | ||
|
|
4852227158 | ||
|
|
9be85b6d3c | ||
|
|
59b98ab730 | ||
|
|
690686f3c7 | ||
|
|
691a04f0dd | ||
|
|
4fa6123c04 | ||
|
|
494cf8b3ef | ||
|
|
73bb600034 | ||
|
|
9cf4d34832 | ||
|
|
284b97bd84 | ||
|
|
7e79f8d1c6 | ||
|
|
258454276e | ||
|
|
938d1b0743 | ||
|
|
b1737040a7 | ||
|
|
26286625f4 | ||
|
|
a214ec40ea | ||
|
|
2c37daef86 | ||
|
|
8e79b3d0bc | ||
|
|
6c0f886cdf | ||
|
|
225036863f | ||
|
|
cac5dd12e9 | ||
|
|
f751c0b46c | ||
|
|
9ed8f50d40 | ||
|
|
d3f2cf7474 | ||
|
|
62750b8980 | ||
|
|
e62649f940 | ||
|
|
0e60c757ce | ||
|
|
68a1e87b66 | ||
|
|
e8a36f033b | ||
|
|
2cf2565e80 | ||
|
|
5669d1062c | ||
|
|
3ace75820e | ||
|
|
020cb0d4bf | ||
|
|
8b75d34a8a | ||
|
|
6320a9aaa9 | ||
|
|
6cd35b185d | ||
|
|
a1ea854b38 | ||
|
|
405dc26cc6 | ||
|
|
fe681abd33 | ||
|
|
43c68468f7 | ||
|
|
ecf3fa2feb | ||
|
|
4aacaeb9b8 | ||
|
|
afc56b9746 | ||
|
|
0a61666197 | ||
|
|
a9e0462e57 | ||
|
|
a2b9986a75 | ||
|
|
654172d757 | ||
|
|
527d48efa9 | ||
|
|
cda08aaed4 | ||
|
|
cfd30581d5 | ||
|
|
3c0313f41b | ||
|
|
d938eb0e76 | ||
|
|
60f2f8c1c4 | ||
|
|
b0c5f7b668 | ||
|
|
767343dc5b | ||
|
|
c22bb4f853 | ||
|
|
117c091b95 | ||
|
|
6719558150 | ||
|
|
6ffce4bccd | ||
|
|
ea9c58ea80 | ||
|
|
b2c2f1bd49 | ||
|
|
679e56c494 | ||
|
|
3da4323ef3 | ||
|
|
75e5a485d2 | ||
|
|
7bb3a827bb | ||
|
|
96f106319e | ||
|
|
a4ad34841b | ||
|
|
599cd2eeeb | ||
|
|
ee5fd1246c | ||
|
|
1441d0d735 | ||
|
|
e5dbfc420d | ||
|
|
643c661a6f | ||
|
|
ab5dfbda54 | ||
|
|
94302de49b | ||
|
|
45d7486485 | ||
|
|
ee27fd8de1 | ||
|
|
89f154630f | ||
|
|
a6ed0ef9f4 | ||
|
|
aac98120c8 | ||
|
|
44e36e5b0d | ||
|
|
f9ab66f51a | ||
|
|
ce8ed5b5ec | ||
|
|
baef422a28 | ||
|
|
68e257849d | ||
|
|
e686554392 | ||
|
|
96a9696383 | ||
|
|
fa84ff5e12 | ||
|
|
66359d5815 | ||
|
|
c111fa0837 | ||
|
|
6e182940e2 | ||
|
|
bc90463ea6 | ||
|
|
93ed4ae2cd | ||
|
|
a10ac774ab | ||
|
|
26a5d8f75d | ||
|
|
5749f78262 | ||
|
|
1eaae9d934 | ||
|
|
567b0776cd | ||
|
|
72f330133a | ||
|
|
665f95eda3 | ||
|
|
8e2b0b6fd2 | ||
|
|
ce50d9bac4 | ||
|
|
934bebd8cd | ||
|
|
171940869b | ||
|
|
33020d826f | ||
|
|
2c12278444 | ||
|
|
57a2024c58 | ||
|
|
d67fe02263 | ||
|
|
57ec2aa088 | ||
|
|
fa859de460 | ||
|
|
36766f157d | ||
|
|
683438b418 | ||
|
|
4a55167759 | ||
|
|
c5c4aef7b1 | ||
|
|
533c7b27eb | ||
|
|
82af218790 | ||
|
|
b272ca5e88 | ||
|
|
10ba2accf7 | ||
|
|
4c8d4e6dbd | ||
|
|
6359628bc3 | ||
|
|
25fd342261 | ||
|
|
f199c486a2 | ||
|
|
1f205a8441 | ||
|
|
b1d5b3b28e | ||
|
|
32810b4152 | ||
|
|
b7e9992d78 | ||
|
|
6c76983999 | ||
|
|
8bf46dcc5d | ||
|
|
5510fa178e | ||
|
|
5ad593e465 | ||
|
|
dff0141160 | ||
|
|
44da9c6523 | ||
|
|
6ab7d54982 | ||
|
|
0c79a566ac | ||
|
|
34773e795b | ||
|
|
66daa15722 | ||
|
|
db80dd2692 | ||
|
|
0dc74a8a2e | ||
|
|
90a057f400 | ||
|
|
78f856e204 | ||
|
|
d2c695eb11 | ||
|
|
46cf40ec82 | ||
|
|
655420fd25 | ||
|
|
52c73390f8 | ||
|
|
c46ef3b63b | ||
|
|
443908b14c | ||
|
|
3bec320bb9 | ||
|
|
fa6f238777 | ||
|
|
86e6b2b68b | ||
|
|
14e51e0977 | ||
|
|
4c6f100b5f | ||
|
|
5a0488bb18 | ||
|
|
9af40624c5 | ||
|
|
0df561c33c | ||
|
|
c7f996d593 | ||
|
|
907dba4517 | ||
|
|
9e5d6069fe | ||
|
|
14f6747dfc | ||
|
|
68b2872ed6 | ||
|
|
1a4bdd2b30 | ||
|
|
886c12c566 | ||
|
|
a3600e8b21 | ||
|
|
5d48e48e15 | ||
|
|
474427c67e | ||
|
|
00b3583dc2 | ||
|
|
4d9a7cc6c0 | ||
|
|
509bd2bebb | ||
|
|
8c70453b2e | ||
|
|
8eebc2aea6 | ||
|
|
a9a0ce6bea | ||
|
|
ecbdef732b | ||
|
|
4615e8f92b | ||
|
|
91faa9fd5a | ||
|
|
85e92fe3b0 | ||
|
|
38bf0b6eec | ||
|
|
6ae3ddd66b | ||
|
|
98cb2d3411 | ||
|
|
be75bc506a | ||
|
|
7a42efec53 | ||
|
|
e9926694c3 | ||
|
|
5cfb7a08cb | ||
|
|
9d642f6354 | ||
|
|
716f2986b9 | ||
|
|
409f565f09 | ||
|
|
26e95f2a92 | ||
|
|
711a2cd738 | ||
|
|
1c1f72f05c | ||
|
|
1d343aeae4 | ||
|
|
e26f6acc3b | ||
|
|
1555252c4a | ||
|
|
de0cbb9073 | ||
|
|
5a075a2c83 | ||
|
|
7da37b4f66 | ||
|
|
6f80cb6b65 | ||
|
|
9617df04ae | ||
|
|
01d5f42755 | ||
|
|
84d76cccde | ||
|
|
af584b46f4 | ||
|
|
0fb4cceec1 | ||
|
|
1dc353433a | ||
|
|
33e8a09880 | ||
|
|
1cb751d184 | ||
|
|
9e596f8616 | ||
|
|
24044b42ea | ||
|
|
84263fc6a6 | ||
|
|
2faab409d3 | ||
|
|
0b5aa6dd60 | ||
|
|
d0c2bfdbff | ||
|
|
242625782f | ||
|
|
826e9ab317 | ||
|
|
182d5e8591 | ||
|
|
3fc866117d | ||
|
|
b464b48f53 | ||
|
|
2b26355002 | ||
|
|
d81a36310c | ||
|
|
d56bb2c383 | ||
|
|
2dd09223f2 | ||
|
|
0c369d195b | ||
|
|
ab99d3b112 | ||
|
|
3f133fad56 | ||
|
|
41d1ccd39c | ||
|
|
7839d043ff | ||
|
|
9b9e6ce2ab | ||
|
|
81510e9d8f | ||
|
|
c0ff925c2a | ||
|
|
2da661fed1 | ||
|
|
59b128bbda | ||
|
|
f2a360cb87 | ||
|
|
8deef788c1 | ||
|
|
f9b0534e0c | ||
|
|
c4de5ea50c | ||
|
|
8646aebaab |
@@ -1,19 +1,33 @@
|
||||
<!--
|
||||
⚠️ CRITICAL CHECKS FOR CONTRIBUTORS (READ, DON'T DELETE) ⚠️
|
||||
1. Target the `dev` branch. PRs targeting `main` will be automatically closed.
|
||||
2. Do NOT delete the CLA section at the bottom. It is required for the bot to accept your PR.
|
||||
-->
|
||||
|
||||
# Pull Request Checklist
|
||||
|
||||
### Note to first-time contributors: Please open a discussion post in [Discussions](https://github.com/open-webui/open-webui/discussions) to discuss your idea/fix with the community before creating a pull request, and describe your changes before submitting a pull request.
|
||||
|
||||
This is to ensure large feature PRs are discussed with the community first, before starting work on it. If the community does not want this feature or it is not relevant for Open WebUI as a project, it can be identified in the discussion before working on the feature and submitting the PR.
|
||||
|
||||
<!--
|
||||
### ⚠️ Important: Your PR is a contribution, not a guarantee of merge.
|
||||
|
||||
The most impactful way to contribute to Open WebUI is through well-written bug reports, detailed feature discussions, and thoughtful ideas. These directly shape the project. If you do open a pull request, please know that Open WebUI is held to the highest standard of code quality, consistency, and architectural coherence, and every line merged becomes something the core team must own, maintain, and support indefinitely. Submitted code may be refactored, rewritten, or used as inspiration for a different implementation. This is not a reflection of your work's quality. It is how we ensure that a small team can deeply understand and evolve every part of the codebase.
|
||||
-->
|
||||
|
||||
**Before submitting, make sure you've checked the following:**
|
||||
|
||||
- [ ] **Target branch:** Verify that the pull request targets the `dev` branch. **Not targeting the `dev` branch will lead to immediate closure of the PR.**
|
||||
- [ ] **Target branch:** Verify that the pull request targets the `dev` branch. **PRs targeting `main` will be immediately closed.**
|
||||
- [ ] **Description:** Provide a concise description of the changes made in this pull request down below.
|
||||
- [ ] **Changelog:** Ensure a changelog entry following the format of [Keep a Changelog](https://keepachangelog.com/) is added at the bottom of the PR description.
|
||||
- [ ] **Documentation:** If necessary, update relevant documentation [Open WebUI Docs](https://github.com/open-webui/docs) like environment variables, the tutorials, or other documentation sources.
|
||||
- [ ] **Dependencies:** Are there any new dependencies? Have you updated the dependency versions in the documentation?
|
||||
- [ ] **Testing:** Perform manual tests to **verify the implemented fix/feature works as intended AND does not break any other functionality**. Take this as an opportunity to **make screenshots of the feature/fix and include it in the PR description**.
|
||||
- [ ] **Documentation:** Add docs in [Open WebUI Docs Repository](https://github.com/open-webui/docs). Document user-facing behavior, environment variables, public APIs/interfaces, or deployment steps.
|
||||
- [ ] **Dependencies:** Are there any new or upgraded dependencies? If so, explain why, update the changelog/docs, and include any compatibility notes. Actually run the code/function that uses updated library to ensure it doesn't crash.
|
||||
- [ ] **Testing:** Perform manual tests to **verify the implemented fix/feature works as intended AND does not break any other functionality**. Include reproducible steps to demonstrate the issue before the fix. Test edge cases (URL encoding, HTML entities, types). Take this as an opportunity to **make screenshots of the feature/fix and include them in the PR description**.
|
||||
- [ ] **Agentic AI Code:** Confirm this Pull Request is **not written by any AI Agent** or has at least **gone through additional human review AND manual testing**. If any AI Agent is the co-author of this PR, it may lead to immediate closure of the PR.
|
||||
- [ ] **Code review:** Have you performed a self-review of your code, addressing any coding standard issues and ensuring adherence to the project's coding standards?
|
||||
- [ ] **Design & Architecture:** Prefer smart defaults over adding new settings; use local state for ephemeral UI logic. Open a Discussion for major architectural or UX changes.
|
||||
- [ ] **Git Hygiene:** Keep PRs atomic (one logical change). Clean up commits and rebase on `dev` to ensure no unrelated commits (e.g. from `main`) are included. Push updates to the existing PR branch instead of closing and reopening.
|
||||
- [ ] **Title Prefix:** To clearly categorize this pull request, prefix the pull request title using one of the following:
|
||||
- **BREAKING CHANGE**: Significant changes that may affect compatibility
|
||||
- **build**: Changes that affect the build system or external dependencies
|
||||
@@ -76,7 +90,15 @@ This is to ensure large feature PRs are discussed with the community first, befo
|
||||
|
||||
### Contributor License Agreement
|
||||
|
||||
By submitting this pull request, I confirm that I have read and fully agree to the [Contributor License Agreement (CLA)](https://github.com/open-webui/open-webui/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT), and I am providing my contributions under its terms.
|
||||
<!--
|
||||
🚨 DO NOT DELETE THE TEXT BELOW 🚨
|
||||
Keep the "Contributor License Agreement" confirmation text intact.
|
||||
Deleting it will trigger the CLA-Bot to INVALIDATE your PR.
|
||||
|
||||
Your PR will NOT be reviewed or merged until you check the box below confirming that you have read and agree to the terms of the CLA.
|
||||
-->
|
||||
|
||||
- [ ] By submitting this pull request, I confirm that I have read and fully agree to the [Contributor License Agreement (CLA)](https://github.com/open-webui/open-webui/blob/main/CONTRIBUTOR_LICENSE_AGREEMENT), and I am providing my contributions under its terms.
|
||||
|
||||
> [!NOTE]
|
||||
> Deleting the CLA section will lead to immediate closure of your PR and it will not be merged in.
|
||||
|
||||
@@ -27,28 +27,17 @@ jobs:
|
||||
echo "::set-output name=version::$VERSION"
|
||||
|
||||
- name: Extract latest CHANGELOG entry
|
||||
id: changelog
|
||||
run: |
|
||||
CHANGELOG_CONTENT=$(awk 'BEGIN {print_section=0;} /^## \[/ {if (print_section == 0) {print_section=1;} else {exit;}} print_section {print;}' CHANGELOG.md)
|
||||
CHANGELOG_ESCAPED=$(echo "$CHANGELOG_CONTENT" | sed ':a;N;$!ba;s/\n/%0A/g')
|
||||
echo "Extracted latest release notes from CHANGELOG.md:"
|
||||
echo -e "$CHANGELOG_CONTENT"
|
||||
echo "::set-output name=content::$CHANGELOG_ESCAPED"
|
||||
VERSION="${{ steps.get_version.outputs.version }}"
|
||||
awk "/^## \[${VERSION}\]/{found=1; next} /^## \[/{if(found) exit} found{print}" CHANGELOG.md > /tmp/release-notes.md
|
||||
|
||||
- name: Create GitHub release
|
||||
uses: actions/github-script@v8
|
||||
with:
|
||||
github-token: ${{ secrets.GITHUB_TOKEN }}
|
||||
script: |
|
||||
const changelog = `${{ steps.changelog.outputs.content }}`;
|
||||
const release = await github.rest.repos.createRelease({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
tag_name: `v${{ steps.get_version.outputs.version }}`,
|
||||
name: `v${{ steps.get_version.outputs.version }}`,
|
||||
body: changelog,
|
||||
})
|
||||
console.log(`Created release ${release.data.html_url}`)
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
gh release create "v${{ steps.get_version.outputs.version }}" \
|
||||
--title "v${{ steps.get_version.outputs.version }}" \
|
||||
--notes-file /tmp/release-notes.md
|
||||
|
||||
- name: Upload package to GitHub release
|
||||
uses: actions/upload-artifact@v4
|
||||
|
||||
@@ -1,64 +0,0 @@
|
||||
name: Deploy to HuggingFace Spaces
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- dev
|
||||
- main
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-secret:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
token-set: ${{ steps.check-key.outputs.defined }}
|
||||
steps:
|
||||
- id: check-key
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
if: "${{ env.HF_TOKEN != '' }}"
|
||||
run: echo "defined=true" >> $GITHUB_OUTPUT
|
||||
|
||||
deploy:
|
||||
runs-on: ubuntu-latest
|
||||
needs: [check-secret]
|
||||
if: needs.check-secret.outputs.token-set == 'true'
|
||||
env:
|
||||
HF_TOKEN: ${{ secrets.HF_TOKEN }}
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v5
|
||||
with:
|
||||
lfs: true
|
||||
|
||||
- name: Remove git history
|
||||
run: rm -rf .git
|
||||
|
||||
- name: Prepend YAML front matter to README.md
|
||||
run: |
|
||||
echo "---" > temp_readme.md
|
||||
echo "title: Open WebUI" >> temp_readme.md
|
||||
echo "emoji: 🐳" >> temp_readme.md
|
||||
echo "colorFrom: purple" >> temp_readme.md
|
||||
echo "colorTo: gray" >> temp_readme.md
|
||||
echo "sdk: docker" >> temp_readme.md
|
||||
echo "app_port: 8080" >> temp_readme.md
|
||||
echo "---" >> temp_readme.md
|
||||
cat README.md >> temp_readme.md
|
||||
mv temp_readme.md README.md
|
||||
|
||||
- name: Configure git
|
||||
run: |
|
||||
git config --global user.email "41898282+github-actions[bot]@users.noreply.github.com"
|
||||
git config --global user.name "github-actions[bot]"
|
||||
- name: Set up Git and push to Space
|
||||
run: |
|
||||
git init --initial-branch=main
|
||||
git lfs install
|
||||
git lfs track "*.ttf"
|
||||
git lfs track "*.jpg"
|
||||
rm demo.png
|
||||
rm banner.png
|
||||
git add .
|
||||
git commit -m "GitHub deploy: ${{ github.sha }}"
|
||||
git push --force https://open-webui:${HF_TOKEN}@huggingface.co/spaces/open-webui/open-webui main
|
||||
@@ -95,6 +95,7 @@ jobs:
|
||||
outputs: type=image,name=${{ env.FULL_IMAGE_NAME }},push-by-digest=true,name-canonical=true,push=true
|
||||
cache-from: type=registry,ref=${{ steps.cache-meta.outputs.tags }}
|
||||
cache-to: type=registry,ref=${{ steps.cache-meta.outputs.tags }},mode=max
|
||||
sbom: true
|
||||
build-args: |
|
||||
BUILD_HASH=${{ github.sha }}
|
||||
|
||||
@@ -199,6 +200,7 @@ jobs:
|
||||
outputs: type=image,name=${{ env.FULL_IMAGE_NAME }},push-by-digest=true,name-canonical=true,push=true
|
||||
cache-from: type=registry,ref=${{ steps.cache-meta.outputs.tags }}
|
||||
cache-to: type=registry,ref=${{ steps.cache-meta.outputs.tags }},mode=max
|
||||
sbom: true
|
||||
build-args: |
|
||||
BUILD_HASH=${{ github.sha }}
|
||||
USE_CUDA=true
|
||||
@@ -304,6 +306,7 @@ jobs:
|
||||
outputs: type=image,name=${{ env.FULL_IMAGE_NAME }},push-by-digest=true,name-canonical=true,push=true
|
||||
cache-from: type=registry,ref=${{ steps.cache-meta.outputs.tags }}
|
||||
cache-to: type=registry,ref=${{ steps.cache-meta.outputs.tags }},mode=max
|
||||
sbom: true
|
||||
build-args: |
|
||||
BUILD_HASH=${{ github.sha }}
|
||||
USE_CUDA=true
|
||||
@@ -407,6 +410,7 @@ jobs:
|
||||
outputs: type=image,name=${{ env.FULL_IMAGE_NAME }},push-by-digest=true,name-canonical=true,push=true
|
||||
cache-from: type=registry,ref=${{ steps.cache-meta.outputs.tags }}
|
||||
cache-to: type=registry,ref=${{ steps.cache-meta.outputs.tags }},mode=max
|
||||
sbom: true
|
||||
build-args: |
|
||||
BUILD_HASH=${{ github.sha }}
|
||||
USE_OLLAMA=true
|
||||
@@ -509,6 +513,7 @@ jobs:
|
||||
outputs: type=image,name=${{ env.FULL_IMAGE_NAME }},push-by-digest=true,name-canonical=true,push=true
|
||||
cache-from: type=registry,ref=${{ steps.cache-meta.outputs.tags }}
|
||||
cache-to: type=registry,ref=${{ steps.cache-meta.outputs.tags }},mode=max
|
||||
sbom: true
|
||||
build-args: |
|
||||
BUILD_HASH=${{ github.sha }}
|
||||
USE_SLIM=true
|
||||
@@ -804,3 +809,109 @@ jobs:
|
||||
- name: Inspect image
|
||||
run: |
|
||||
docker buildx imagetools inspect ${{ env.FULL_IMAGE_NAME }}:${{ steps.meta.outputs.version }}
|
||||
|
||||
# Copy images from GHCR to Docker Hub (best-effort, won't block GHCR)
|
||||
copy-to-dockerhub:
|
||||
runs-on: ubuntu-latest
|
||||
if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags/v')
|
||||
needs: [merge-main-images, merge-cuda-images, merge-cuda126-images, merge-ollama-images, merge-slim-images]
|
||||
continue-on-error: true
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- variant: main
|
||||
suffix: ""
|
||||
- variant: cuda
|
||||
suffix: "-cuda"
|
||||
- variant: cuda126
|
||||
suffix: "-cuda126"
|
||||
- variant: ollama
|
||||
suffix: "-ollama"
|
||||
- variant: slim
|
||||
suffix: "-slim"
|
||||
steps:
|
||||
- name: Set repository and image name to lowercase
|
||||
run: |
|
||||
echo "IMAGE_NAME=${IMAGE_NAME,,}" >>${GITHUB_ENV}
|
||||
echo "FULL_IMAGE_NAME=ghcr.io/${IMAGE_NAME,,}" >>${GITHUB_ENV}
|
||||
env:
|
||||
IMAGE_NAME: '${{ github.repository }}'
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Log in to the Container registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ${{ env.REGISTRY }}
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Log in to Docker Hub
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
username: ${{ secrets.DOCKERHUB_USERNAME }}
|
||||
password: ${{ secrets.DOCKERHUB_TOKEN }}
|
||||
|
||||
- name: Determine source and destination tags
|
||||
id: tags
|
||||
run: |
|
||||
DOCKERHUB_IMAGE="openwebui/open-webui"
|
||||
SUFFIX="${{ matrix.suffix }}"
|
||||
|
||||
if [[ "${{ github.ref }}" == refs/tags/v* ]]; then
|
||||
# For version tags: copy version tag and major.minor tag
|
||||
VERSION="${{ github.ref_name }}"
|
||||
VERSION="${VERSION#v}"
|
||||
MAJOR_MINOR="${VERSION%.*}"
|
||||
|
||||
echo "tags<<EOF" >> $GITHUB_OUTPUT
|
||||
echo "${VERSION}${SUFFIX}" >> $GITHUB_OUTPUT
|
||||
echo "${MAJOR_MINOR}${SUFFIX}" >> $GITHUB_OUTPUT
|
||||
echo "EOF" >> $GITHUB_OUTPUT
|
||||
else
|
||||
# For main branch
|
||||
if [ -z "$SUFFIX" ]; then
|
||||
echo "tags=latest" >> $GITHUB_OUTPUT
|
||||
else
|
||||
# e.g. latest-cuda -> also tag as just "cuda"
|
||||
VARIANT_NAME="${SUFFIX#-}"
|
||||
echo "tags<<EOF" >> $GITHUB_OUTPUT
|
||||
echo "latest${SUFFIX}" >> $GITHUB_OUTPUT
|
||||
echo "${VARIANT_NAME}" >> $GITHUB_OUTPUT
|
||||
echo "EOF" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
fi
|
||||
|
||||
echo "dockerhub_image=${DOCKERHUB_IMAGE}" >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Copy images from GHCR to Docker Hub
|
||||
run: |
|
||||
DOCKERHUB_IMAGE="${{ steps.tags.outputs.dockerhub_image }}"
|
||||
SUFFIX="${{ matrix.suffix }}"
|
||||
|
||||
# Determine the source tag on GHCR
|
||||
if [[ "${{ github.ref }}" == refs/tags/v* ]]; then
|
||||
VERSION="${{ github.ref_name }}"
|
||||
VERSION="${VERSION#v}"
|
||||
SOURCE_TAG="${VERSION}${SUFFIX}"
|
||||
else
|
||||
if [ -z "$SUFFIX" ]; then
|
||||
SOURCE_TAG="latest"
|
||||
else
|
||||
SOURCE_TAG="latest${SUFFIX}"
|
||||
fi
|
||||
fi
|
||||
|
||||
SOURCE="${{ env.FULL_IMAGE_NAME }}:${SOURCE_TAG}"
|
||||
|
||||
echo "Copying from ${SOURCE} to Docker Hub..."
|
||||
|
||||
# Copy each destination tag
|
||||
while IFS= read -r TAG; do
|
||||
[ -z "$TAG" ] && continue
|
||||
DEST="${DOCKERHUB_IMAGE}:${TAG}"
|
||||
echo " -> ${DEST}"
|
||||
docker buildx imagetools create -t "${DEST}" "${SOURCE}"
|
||||
done <<< "${{ steps.tags.outputs.tags }}"
|
||||
|
||||
@@ -40,10 +40,7 @@ jobs:
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install black
|
||||
pip install "ruff>=0.15.5"
|
||||
|
||||
- name: Format backend
|
||||
run: npm run format:backend
|
||||
|
||||
- name: Check for changes after format
|
||||
run: git diff --exit-code
|
||||
- name: Ruff format check
|
||||
run: ruff format --check . --exclude .venv --exclude venv
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
repos:
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.15.5
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--fix, backend]
|
||||
- id: ruff-format
|
||||
args: [backend]
|
||||
+995
-1
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,7 @@
|
||||
# Open WebUI Contributor License Agreement
|
||||
# Contributor License Agreement
|
||||
|
||||
By submitting my contributions to Open WebUI, I grant Open WebUI full freedom to use my work in any way they choose, under any terms they like, both now and in the future. This approach helps ensure the project remains unified, flexible, and easy to maintain, while empowering Open WebUI to respond quickly to the needs of its users and the wider community.
|
||||
By submitting my contributions to this repository in any form, I grant Open WebUI Inc. a perpetual, worldwide, irrevocable, royalty-free license, under copyright and patent, to use, modify, distribute, sublicense, and commercialize my work under any terms they choose, both now and in the future.
|
||||
|
||||
Taking part in this process means my work can be seamlessly integrated and combined with others, ensuring longevity and adaptability for everyone who benefits from the Open WebUI project. This collaborative approach strengthens the project’s future and helps guarantee that improvements can always be shared and distributed in the most effective way possible.
|
||||
I represent that my contributions are my original work (or that I have sufficient rights to grant this license) and that I have the authority to enter into this agreement.
|
||||
|
||||
**_To the fullest extent permitted by law, my contributions are provided on an “as is” basis, with no warranties or guarantees of any kind, and I disclaim any liability for any issues or damages arising from their use or incorporation into the project, regardless of the type of legal claim._**
|
||||
+19
-11
@@ -127,33 +127,41 @@ RUN chown -R $UID:$GID /app $HOME
|
||||
RUN apt-get update && \
|
||||
apt-get install -y --no-install-recommends \
|
||||
git build-essential pandoc gcc netcat-openbsd curl jq \
|
||||
libmariadb-dev \
|
||||
python3-dev \
|
||||
ffmpeg libsm6 libxext6 \
|
||||
ffmpeg libsm6 libxext6 zstd \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# install python dependencies
|
||||
COPY --chown=$UID:$GID ./backend/requirements.txt ./requirements.txt
|
||||
|
||||
RUN pip3 install --no-cache-dir uv && \
|
||||
# Set UV_LINK_MODE to copy to prevent 0-byte file corruption in QEMU arm64 cross-builds
|
||||
ENV UV_LINK_MODE=copy
|
||||
|
||||
RUN set -e; \
|
||||
pip3 install --no-cache-dir uv; \
|
||||
if [ "$USE_CUDA" = "true" ]; then \
|
||||
# If you use CUDA the whisper and embedding model will be downloaded on first use
|
||||
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/$USE_CUDA_DOCKER_VER --no-cache-dir && \
|
||||
uv pip install --system -r requirements.txt --no-cache-dir && \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ['RAG_EMBEDDING_MODEL'], device='cpu')" && \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')" && \
|
||||
# fix: pin torch<=2.9.1 - torch 2.10.0 aarch64 wheels cause SIGILL on ARM devices (RPi 4 Cortex-A72) #21349
|
||||
pip3 install 'torch<=2.9.1' torchvision torchaudio --index-url https://download.pytorch.org/whl/$USE_CUDA_DOCKER_VER --no-cache-dir; \
|
||||
uv pip install --system -r requirements.txt --no-cache-dir; \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ['RAG_EMBEDDING_MODEL'], device='cpu')"; \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')"; \
|
||||
python -c "import os; from faster_whisper import WhisperModel; WhisperModel(os.environ['WHISPER_MODEL'], device='cpu', compute_type='int8', download_root=os.environ['WHISPER_MODEL_DIR'])"; \
|
||||
python -c "import os; import tiktoken; tiktoken.get_encoding(os.environ['TIKTOKEN_ENCODING_NAME'])"; \
|
||||
python -c "import nltk; nltk.download('punkt_tab')"; \
|
||||
else \
|
||||
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu --no-cache-dir && \
|
||||
uv pip install --system -r requirements.txt --no-cache-dir && \
|
||||
pip3 install 'torch<=2.9.1' torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu --no-cache-dir; \
|
||||
uv pip install --system -r requirements.txt --no-cache-dir; \
|
||||
if [ "$USE_SLIM" != "true" ]; then \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ['RAG_EMBEDDING_MODEL'], device='cpu')" && \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')" && \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ['RAG_EMBEDDING_MODEL'], device='cpu')"; \
|
||||
python -c "import os; from sentence_transformers import SentenceTransformer; SentenceTransformer(os.environ.get('AUXILIARY_EMBEDDING_MODEL', 'TaylorAI/bge-micro-v2'), device='cpu')"; \
|
||||
python -c "import os; from faster_whisper import WhisperModel; WhisperModel(os.environ['WHISPER_MODEL'], device='cpu', compute_type='int8', download_root=os.environ['WHISPER_MODEL_DIR'])"; \
|
||||
python -c "import os; import tiktoken; tiktoken.get_encoding(os.environ['TIKTOKEN_ENCODING_NAME'])"; \
|
||||
python -c "import nltk; nltk.download('punkt_tab')"; \
|
||||
fi; \
|
||||
fi; \
|
||||
mkdir -p /app/backend/data && chown -R $UID:$GID /app/backend/data/ && \
|
||||
mkdir -p /app/backend/data; chown -R $UID:$GID /app/backend/data/; \
|
||||
rm -rf /var/lib/apt/lists/*;
|
||||
|
||||
# Install Ollama if requested
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
Open WebUI License
|
||||
|
||||
Copyright (c) 2023- Open WebUI Inc. [Created by Timothy Jaeryang Baek]
|
||||
All rights reserved.
|
||||
|
||||
@@ -15,11 +17,27 @@ modification, are permitted provided that the following conditions are met:
|
||||
contributors may be used to endorse or promote products derived from
|
||||
this software without specific prior written permission.
|
||||
|
||||
4. Notwithstanding any other provision of this License, and as a material condition of the rights granted herein, licensees are strictly prohibited from altering, removing, obscuring, or replacing any "Open WebUI" branding, including but not limited to the name, logo, or any visual, textual, or symbolic identifiers that distinguish the software and its interfaces, in any deployment or distribution, regardless of the number of users, except as explicitly set forth in Clauses 5 and 6 below.
|
||||
4. Notwithstanding any other provision of this License, and as a material
|
||||
condition of the rights granted herein, licensees are strictly prohibited
|
||||
from altering, removing, obscuring, or replacing any "Open WebUI"
|
||||
branding, including but not limited to the name, logo, or any visual,
|
||||
textual, or symbolic identifiers that distinguish the software and its
|
||||
interfaces, in any deployment or distribution, except in the following
|
||||
circumstances: (i) deployments or distributions where the total number
|
||||
of end users (defined as individual natural persons with direct access
|
||||
to the application) does not exceed fifty (50) within any rolling
|
||||
thirty (30) day period; (ii) the licensee has obtained specific prior
|
||||
written permission from the copyright holder; or (iii) where the
|
||||
licensee has obtained a duly executed enterprise license expressly
|
||||
permitting such modification. For all other cases, any removal or
|
||||
alteration of the "Open WebUI" branding shall constitute a material
|
||||
breach of license.
|
||||
|
||||
5. The branding restriction enumerated in Clause 4 shall not apply in the following limited circumstances: (i) deployments or distributions where the total number of end users (defined as individual natural persons with direct access to the application) does not exceed fifty (50) within any rolling thirty (30) day period; (ii) cases in which the licensee is an official contributor to the codebase—with a substantive code change successfully merged into the main branch of the official codebase maintained by the copyright holder—who has obtained specific prior written permission for branding adjustment from the copyright holder; or (iii) where the licensee has obtained a duly executed enterprise license expressly permitting such modification. For all other cases, any removal or alteration of the "Open WebUI" branding shall constitute a material breach of license.
|
||||
Materials governed by prior licenses retain those original license
|
||||
terms, as specified in LICENSE_HISTORY.
|
||||
|
||||
6. All code, modifications, or derivative works incorporated into this project prior to the incorporation of this branding clause remain licensed under the BSD 3-Clause License, and prior contributors retain all BSD-3 rights therein; if any such contributor requests the removal of their BSD-3-licensed code, the copyright holder will do so, and any replacement code will be licensed under the project's primary license then in effect. By contributing after this clause's adoption, you agree to the project's Contributor License Agreement (CLA) and to these updated terms for all new contributions.
|
||||
By contributing to this project, you agree to the project's Contributor
|
||||
License Agreement (CONTRIBUTOR_LICENSE_AGREEMENT).
|
||||
|
||||
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
|
||||

|
||||
|
||||
**Open WebUI is an [extensible](https://docs.openwebui.com/features/plugin/), feature-rich, and user-friendly self-hosted AI platform designed to operate entirely offline.** It supports various LLM runners like **Ollama** and **OpenAI-compatible APIs**, with **built-in inference engine** for RAG, making it a **powerful AI deployment solution**.
|
||||
**Open WebUI is an [extensible](https://docs.openwebui.com/features/extensibility/plugin), feature-rich, and user-friendly self-hosted AI platform designed to operate entirely offline.** It supports various LLM runners like **Ollama** and **OpenAI-compatible APIs**, with **built-in inference engine** for RAG, making it a **powerful AI deployment solution**.
|
||||
|
||||
Passionate about open-source AI? [Join our team →](https://careers.openwebui.com/)
|
||||
|
||||
@@ -47,7 +47,7 @@ For more information, be sure to check out our [Open WebUI Documentation](https:
|
||||
|
||||
- 💾 **Persistent Artifact Storage**: Built-in key-value storage API for artifacts, enabling features like journals, trackers, leaderboards, and collaborative tools with both personal and shared data scopes across sessions.
|
||||
|
||||
- 📚 **Local RAG Integration**: Dive into the future of chat interactions with groundbreaking Retrieval Augmented Generation (RAG) support using your choice of 9 vector databases and multiple content extraction engines (Tika, Docling, Document Intelligence, Mistral OCR, External loaders). Load documents directly into chat or add files to your document library, effortlessly accessing them using the `#` command before a query.
|
||||
- 📚 **Local RAG Integration**: Dive into the future of chat interactions with groundbreaking Retrieval Augmented Generation (RAG) support using your choice of 9 vector databases and multiple content extraction engines (Tika, Docling, Document Intelligence, Mistral OCR, PaddleOCR-vl, External loaders). Load documents directly into chat or add files to your document library, effortlessly accessing them using the `#` command before a query.
|
||||
|
||||
- 🔍 **Web Search for RAG**: Perform web searches using 15+ providers including `SearXNG`, `Google PSE`, `Brave Search`, `Kagi`, `Mojeek`, `Tavily`, `Perplexity`, `serpstack`, `serper`, `Serply`, `DuckDuckGo`, `SearchApi`, `SerpApi`, `Bing`, `Jina`, `Exa`, `Sougou`, `Azure AI Search`, and `Ollama Cloud`, injecting results directly into your chat experience.
|
||||
|
||||
@@ -172,8 +172,6 @@ After installation, you can access Open WebUI at [http://localhost:3000](http://
|
||||
|
||||
We offer various installation alternatives, including non-Docker native installation methods, Docker Compose, Kustomize, and Helm. Visit our [Open WebUI Documentation](https://docs.openwebui.com/getting-started/) or join our [Discord community](https://discord.gg/5rJgQTnV4s) for comprehensive guidance.
|
||||
|
||||
Look at the [Local Development Guide](https://docs.openwebui.com/getting-started/advanced-topics/development) for instructions on setting up a local development environment.
|
||||
|
||||
### Troubleshooting
|
||||
|
||||
Encountering connection issues? Our [Open WebUI Documentation](https://docs.openwebui.com/troubleshooting/) has got you covered. For further assistance and to join our vibrant community, visit the [Open WebUI Discord](https://discord.gg/5rJgQTnV4s).
|
||||
|
||||
+1
-1
@@ -1,3 +1,3 @@
|
||||
export CORS_ALLOW_ORIGIN="http://localhost:5173;http://localhost:8080"
|
||||
PORT="${PORT:-8080}"
|
||||
uvicorn open_webui.main:app --port $PORT --host 0.0.0.0 --forwarded-allow-ips '*' --reload
|
||||
uvicorn open_webui.main:app --port $PORT --host 0.0.0.0 --forwarded-allow-ips "${FORWARDED_ALLOW_IPS:-*}" --reload
|
||||
|
||||
@@ -2,102 +2,95 @@ import base64
|
||||
import os
|
||||
import random
|
||||
from pathlib import Path
|
||||
from typing import Annotated
|
||||
|
||||
import typer
|
||||
import uvicorn
|
||||
from typing import Optional
|
||||
from typing_extensions import Annotated
|
||||
|
||||
app = typer.Typer()
|
||||
|
||||
KEY_FILE = Path.cwd() / ".webui_secret_key"
|
||||
KEY_FILE = Path.cwd() / '.webui_secret_key'
|
||||
|
||||
|
||||
def version_callback(value: bool):
|
||||
def version_callback(value: bool) -> None:
|
||||
if value:
|
||||
from open_webui.env import VERSION
|
||||
|
||||
typer.echo(f"Open WebUI version: {VERSION}")
|
||||
typer.echo(f'Open WebUI version: {VERSION}')
|
||||
raise typer.Exit()
|
||||
|
||||
|
||||
@app.command()
|
||||
def main(
|
||||
version: Annotated[
|
||||
Optional[bool], typer.Option("--version", callback=version_callback)
|
||||
] = None,
|
||||
version: Annotated[bool | None, typer.Option('--version', callback=version_callback)] = None,
|
||||
):
|
||||
pass
|
||||
|
||||
|
||||
@app.command()
|
||||
def serve(
|
||||
host: str = "0.0.0.0",
|
||||
host: str = '0.0.0.0',
|
||||
port: int = 8080,
|
||||
):
|
||||
os.environ["FROM_INIT_PY"] = "true"
|
||||
if os.getenv("WEBUI_SECRET_KEY") is None:
|
||||
typer.echo(
|
||||
"Loading WEBUI_SECRET_KEY from file, not provided as an environment variable."
|
||||
)
|
||||
os.environ['FROM_INIT_PY'] = 'true'
|
||||
if os.getenv('WEBUI_SECRET_KEY') is None:
|
||||
typer.echo('Loading WEBUI_SECRET_KEY from file, not provided as an environment variable.')
|
||||
if not KEY_FILE.exists():
|
||||
typer.echo(f"Generating a new secret key and saving it to {KEY_FILE}")
|
||||
typer.echo(f'Generating a new secret key and saving it to {KEY_FILE}')
|
||||
KEY_FILE.write_bytes(base64.b64encode(random.randbytes(12)))
|
||||
typer.echo(f"Loading WEBUI_SECRET_KEY from {KEY_FILE}")
|
||||
os.environ["WEBUI_SECRET_KEY"] = KEY_FILE.read_text()
|
||||
typer.echo(f'Loading WEBUI_SECRET_KEY from {KEY_FILE}')
|
||||
os.environ['WEBUI_SECRET_KEY'] = KEY_FILE.read_text()
|
||||
|
||||
if os.getenv("USE_CUDA_DOCKER", "false") == "true":
|
||||
typer.echo(
|
||||
"CUDA is enabled, appending LD_LIBRARY_PATH to include torch/cudnn & cublas libraries."
|
||||
)
|
||||
LD_LIBRARY_PATH = os.getenv("LD_LIBRARY_PATH", "").split(":")
|
||||
os.environ["LD_LIBRARY_PATH"] = ":".join(
|
||||
if os.getenv('USE_CUDA_DOCKER', 'false') == 'true':
|
||||
typer.echo('CUDA is enabled, appending LD_LIBRARY_PATH to include torch/cudnn & cublas libraries.')
|
||||
LD_LIBRARY_PATH = os.getenv('LD_LIBRARY_PATH', '').split(':')
|
||||
os.environ['LD_LIBRARY_PATH'] = ':'.join(
|
||||
LD_LIBRARY_PATH
|
||||
+ [
|
||||
"/usr/local/lib/python3.11/site-packages/torch/lib",
|
||||
"/usr/local/lib/python3.11/site-packages/nvidia/cudnn/lib",
|
||||
'/usr/local/lib/python3.11/site-packages/torch/lib',
|
||||
'/usr/local/lib/python3.11/site-packages/nvidia/cudnn/lib',
|
||||
]
|
||||
)
|
||||
try:
|
||||
import torch
|
||||
|
||||
assert torch.cuda.is_available(), "CUDA not available"
|
||||
typer.echo("CUDA seems to be working")
|
||||
assert torch.cuda.is_available(), 'CUDA not available'
|
||||
typer.echo('CUDA seems to be working')
|
||||
except Exception as e:
|
||||
typer.echo(
|
||||
"Error when testing CUDA but USE_CUDA_DOCKER is true. "
|
||||
"Resetting USE_CUDA_DOCKER to false and removing "
|
||||
f"LD_LIBRARY_PATH modifications: {e}"
|
||||
'Error when testing CUDA but USE_CUDA_DOCKER is true. '
|
||||
'Resetting USE_CUDA_DOCKER to false and removing '
|
||||
f'LD_LIBRARY_PATH modifications: {e}'
|
||||
)
|
||||
os.environ["USE_CUDA_DOCKER"] = "false"
|
||||
os.environ["LD_LIBRARY_PATH"] = ":".join(LD_LIBRARY_PATH)
|
||||
os.environ['USE_CUDA_DOCKER'] = 'false'
|
||||
os.environ['LD_LIBRARY_PATH'] = ':'.join(LD_LIBRARY_PATH)
|
||||
|
||||
import open_webui.main # we need set environment variables before importing main
|
||||
import open_webui.main # noqa: F401
|
||||
from open_webui.env import UVICORN_WORKERS # Import the workers setting
|
||||
|
||||
uvicorn.run(
|
||||
"open_webui.main:app",
|
||||
'open_webui.main:app',
|
||||
host=host,
|
||||
port=port,
|
||||
forwarded_allow_ips="*",
|
||||
forwarded_allow_ips='*',
|
||||
workers=UVICORN_WORKERS,
|
||||
)
|
||||
|
||||
|
||||
@app.command()
|
||||
def dev(
|
||||
host: str = "0.0.0.0",
|
||||
host: str = '0.0.0.0',
|
||||
port: int = 8080,
|
||||
reload: bool = True,
|
||||
):
|
||||
uvicorn.run(
|
||||
"open_webui.main:app",
|
||||
'open_webui.main:app',
|
||||
host=host,
|
||||
port=port,
|
||||
reload=reload,
|
||||
forwarded_allow_ips="*",
|
||||
forwarded_allow_ips='*',
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if __name__ == '__main__':
|
||||
app()
|
||||
|
||||
+1856
-1703
File diff suppressed because it is too large
Load Diff
@@ -2,125 +2,118 @@ from enum import Enum
|
||||
|
||||
|
||||
class MESSAGES(str, Enum):
|
||||
DEFAULT = lambda msg="": f"{msg if msg else ''}"
|
||||
MODEL_ADDED = lambda model="": f"The model '{model}' has been added successfully."
|
||||
MODEL_DELETED = (
|
||||
lambda model="": f"The model '{model}' has been deleted successfully."
|
||||
)
|
||||
DEFAULT = lambda msg='': f'{msg if msg else ""}'
|
||||
MODEL_ADDED = lambda model='': f"The model '{model}' has been added successfully."
|
||||
MODEL_DELETED = lambda model='': f"The model '{model}' has been deleted successfully."
|
||||
|
||||
|
||||
class WEBHOOK_MESSAGES(str, Enum):
|
||||
DEFAULT = lambda msg="": f"{msg if msg else ''}"
|
||||
USER_SIGNUP = lambda username="": (
|
||||
f"New user signed up: {username}" if username else "New user signed up"
|
||||
)
|
||||
DEFAULT = lambda msg='': f'{msg if msg else ""}'
|
||||
USER_SIGNUP = lambda username='': f'New user signed up: {username}' if username else 'New user signed up'
|
||||
|
||||
|
||||
class ERROR_MESSAGES(str, Enum):
|
||||
def __str__(self) -> str:
|
||||
return super().__str__()
|
||||
|
||||
DEFAULT = (
|
||||
lambda err="": f'{"Something went wrong :/" if err == "" else "[ERROR: " + str(err) + "]"}'
|
||||
DEFAULT = lambda err='': f'{"Something went wrong :/" if err == "" else "[ERROR: " + str(err) + "]"}'
|
||||
ENV_VAR_NOT_FOUND = 'Required environment variable not found. Terminating now.'
|
||||
CREATE_USER_ERROR = 'Oops! Something went wrong while creating your account. Please try again later. If the issue persists, contact support for assistance.'
|
||||
DELETE_USER_ERROR = 'Oops! Something went wrong. We encountered an issue while trying to delete the user. Please give it another shot.'
|
||||
EMAIL_MISMATCH = 'Uh-oh! This email does not match the email your provider is registered with. Please check your email and try again.'
|
||||
EMAIL_TAKEN = 'Uh-oh! This email is already registered. Sign in with your existing account or choose another email to start anew.'
|
||||
USERNAME_TAKEN = 'Uh-oh! This username is already registered. Please choose another username.'
|
||||
PASSWORD_TOO_LONG = (
|
||||
'Uh-oh! The password you entered is too long. Please make sure your password is less than 72 bytes long.'
|
||||
)
|
||||
ENV_VAR_NOT_FOUND = "Required environment variable not found. Terminating now."
|
||||
CREATE_USER_ERROR = "Oops! Something went wrong while creating your account. Please try again later. If the issue persists, contact support for assistance."
|
||||
DELETE_USER_ERROR = "Oops! Something went wrong. We encountered an issue while trying to delete the user. Please give it another shot."
|
||||
EMAIL_MISMATCH = "Uh-oh! This email does not match the email your provider is registered with. Please check your email and try again."
|
||||
EMAIL_TAKEN = "Uh-oh! This email is already registered. Sign in with your existing account or choose another email to start anew."
|
||||
USERNAME_TAKEN = (
|
||||
"Uh-oh! This username is already registered. Please choose another username."
|
||||
)
|
||||
PASSWORD_TOO_LONG = "Uh-oh! The password you entered is too long. Please make sure your password is less than 72 bytes long."
|
||||
COMMAND_TAKEN = "Uh-oh! This command is already registered. Please choose another command string."
|
||||
FILE_EXISTS = "Uh-oh! This file is already registered. Please choose another file."
|
||||
COMMAND_TAKEN = 'Uh-oh! This command is already registered. Please choose another command string.'
|
||||
FILE_EXISTS = 'Uh-oh! This file is already registered. Please choose another file.'
|
||||
|
||||
ID_TAKEN = "Uh-oh! This id is already registered. Please choose another id string."
|
||||
MODEL_ID_TAKEN = "Uh-oh! This model id is already registered. Please choose another model id string."
|
||||
NAME_TAG_TAKEN = "Uh-oh! This name tag is already registered. Please choose another name tag string."
|
||||
MODEL_ID_TOO_LONG = "The model id is too long. Please make sure your model id is less than 256 characters long."
|
||||
ID_TAKEN = 'Uh-oh! This id is already registered. Please choose another id string.'
|
||||
MODEL_ID_TAKEN = 'Uh-oh! This model id is already registered. Please choose another model id string.'
|
||||
NAME_TAG_TAKEN = 'Uh-oh! This name tag is already registered. Please choose another name tag string.'
|
||||
MODEL_ID_TOO_LONG = 'The model id is too long. Please make sure your model id is less than 256 characters long.'
|
||||
|
||||
INVALID_TOKEN = (
|
||||
"Your session has expired or the token is invalid. Please sign in again."
|
||||
)
|
||||
INVALID_CRED = "The email or password provided is incorrect. Please check for typos and try logging in again."
|
||||
INVALID_TOKEN = 'Your session has expired or the token is invalid. Please sign in again.'
|
||||
INVALID_CRED = 'The email or password provided is incorrect. Please check for typos and try logging in again.'
|
||||
INVALID_EMAIL_FORMAT = "The email format you entered is invalid. Please double-check and make sure you're using a valid email address (e.g., yourname@example.com)."
|
||||
INCORRECT_PASSWORD = (
|
||||
"The password provided is incorrect. Please check for typos and try again."
|
||||
INCORRECT_PASSWORD = 'The password provided is incorrect. Please check for typos and try again.'
|
||||
INVALID_TRUSTED_HEADER = (
|
||||
'Your provider has not provided a trusted header. Please contact your administrator for assistance.'
|
||||
)
|
||||
INVALID_TRUSTED_HEADER = "Your provider has not provided a trusted header. Please contact your administrator for assistance."
|
||||
|
||||
EXISTING_USERS = "You can't turn off authentication because there are existing users. If you want to disable WEBUI_AUTH, make sure your web interface doesn't have any existing users and is a fresh installation."
|
||||
|
||||
UNAUTHORIZED = "401 Unauthorized"
|
||||
ACCESS_PROHIBITED = "You do not have permission to access this resource. Please contact your administrator for assistance."
|
||||
ACTION_PROHIBITED = (
|
||||
"The requested action has been restricted as a security measure."
|
||||
UNAUTHORIZED = '401 Unauthorized'
|
||||
ACCESS_PROHIBITED = (
|
||||
'You do not have permission to access this resource. Please contact your administrator for assistance.'
|
||||
)
|
||||
ACTION_PROHIBITED = 'The requested action has been restricted as a security measure.'
|
||||
|
||||
FILE_NOT_SENT = "FILE_NOT_SENT"
|
||||
FILE_NOT_SENT = 'FILE_NOT_SENT'
|
||||
FILE_NOT_SUPPORTED = "Oops! It seems like the file format you're trying to upload is not supported. Please upload a file with a supported format and try again."
|
||||
|
||||
NOT_FOUND = "We could not find what you're looking for :/"
|
||||
USER_NOT_FOUND = "We could not find what you're looking for :/"
|
||||
API_KEY_NOT_FOUND = "Oops! It looks like there's a hiccup. The API key is missing. Please make sure to provide a valid API key to access this feature."
|
||||
API_KEY_NOT_ALLOWED = "Use of API key is not enabled in the environment."
|
||||
API_KEY_NOT_ALLOWED = 'Use of API key is not enabled in the environment.'
|
||||
|
||||
MALICIOUS = "Unusual activities detected, please try again in a few minutes."
|
||||
MALICIOUS = 'Unusual activities detected, please try again in a few minutes.'
|
||||
|
||||
PANDOC_NOT_INSTALLED = "Pandoc is not installed on the server. Please contact your administrator for assistance."
|
||||
INCORRECT_FORMAT = (
|
||||
lambda err="": f"Invalid format. Please use the correct format{err}"
|
||||
)
|
||||
RATE_LIMIT_EXCEEDED = "API rate limit exceeded"
|
||||
PANDOC_NOT_INSTALLED = 'Pandoc is not installed on the server. Please contact your administrator for assistance.'
|
||||
INCORRECT_FORMAT = lambda err='': f'Invalid format. Please use the correct format{err}'
|
||||
RATE_LIMIT_EXCEEDED = 'API rate limit exceeded'
|
||||
|
||||
MODEL_NOT_FOUND = lambda name="": f"Model '{name}' was not found"
|
||||
OPENAI_NOT_FOUND = lambda name="": "OpenAI API was not found"
|
||||
OLLAMA_NOT_FOUND = "WebUI could not connect to Ollama"
|
||||
CREATE_API_KEY_ERROR = "Oops! Something went wrong while creating your API key. Please try again later. If the issue persists, contact support for assistance."
|
||||
API_KEY_CREATION_NOT_ALLOWED = "API key creation is not allowed in the environment."
|
||||
MODEL_NOT_FOUND = lambda name='': f"Model '{name}' was not found"
|
||||
OPENAI_NOT_FOUND = lambda name='': 'OpenAI API was not found'
|
||||
OLLAMA_NOT_FOUND = 'WebUI could not connect to Ollama'
|
||||
CREATE_API_KEY_ERROR = 'Oops! Something went wrong while creating your API key. Please try again later. If the issue persists, contact support for assistance.'
|
||||
API_KEY_CREATION_NOT_ALLOWED = 'API key creation is not allowed in the environment.'
|
||||
|
||||
EMPTY_CONTENT = "The content provided is empty. Please ensure that there is text or data present before proceeding."
|
||||
EMPTY_CONTENT = 'The content provided is empty. Please ensure that there is text or data present before proceeding.'
|
||||
|
||||
DB_NOT_SQLITE = "This feature is only available when running with SQLite databases."
|
||||
DB_NOT_SQLITE = 'This feature is only available when running with SQLite databases.'
|
||||
|
||||
INVALID_URL = (
|
||||
"Oops! The URL you provided is invalid. Please double-check and try again."
|
||||
INVALID_URL = 'Oops! The URL you provided is invalid. Please double-check and try again.'
|
||||
|
||||
WEB_SEARCH_ERROR = lambda err='': f'{err if err else "Oops! Something went wrong while searching the web."}'
|
||||
|
||||
OLLAMA_API_DISABLED = 'The Ollama API is disabled. Please enable it to use this feature.'
|
||||
|
||||
FILE_TOO_LARGE = lambda size='': (
|
||||
f"Oops! The file you're trying to upload is too large. Please upload a file that is less than {size}."
|
||||
)
|
||||
|
||||
WEB_SEARCH_ERROR = (
|
||||
lambda err="": f"{err if err else 'Oops! Something went wrong while searching the web.'}"
|
||||
DUPLICATE_CONTENT = 'Duplicate content detected. Please provide unique content to proceed.'
|
||||
FILE_NOT_PROCESSED = (
|
||||
'Extracted content is not available for this file. Please ensure that the file is processed before proceeding.'
|
||||
)
|
||||
|
||||
OLLAMA_API_DISABLED = (
|
||||
"The Ollama API is disabled. Please enable it to use this feature."
|
||||
)
|
||||
INVALID_PASSWORD = lambda err='': err if err else 'The password does not meet the required validation criteria.'
|
||||
|
||||
FILE_TOO_LARGE = (
|
||||
lambda size="": f"Oops! The file you're trying to upload is too large. Please upload a file that is less than {size}."
|
||||
)
|
||||
AUTOMATION_LIMIT_EXCEEDED = lambda size='': f'Automation limit reached ({size})'
|
||||
AUTOMATION_TOO_FREQUENT = lambda interval='': f'Schedule too frequent. Minimum interval is {interval} seconds.'
|
||||
AUTOMATION_INVALID_RRULE = lambda err='': f'Invalid RRULE: {err}'
|
||||
AUTOMATION_NO_FUTURE_RUNS = 'RRULE has no future occurrences'
|
||||
|
||||
DUPLICATE_CONTENT = (
|
||||
"Duplicate content detected. Please provide unique content to proceed."
|
||||
)
|
||||
FILE_NOT_PROCESSED = "Extracted content is not available for this file. Please ensure that the file is processed before proceeding."
|
||||
|
||||
INVALID_PASSWORD = lambda err="": (
|
||||
err if err else "The password does not meet the required validation criteria."
|
||||
)
|
||||
FEATURE_DISABLED = lambda name='': f'{name} is disabled'
|
||||
INPUT_TOO_LONG = lambda size='': f'Input prompt exceeds maximum length of {size}'
|
||||
SERVER_CONNECTION_ERROR = 'Open WebUI: Server Connection Error'
|
||||
REQUIRED_FIELD_EMPTY = lambda name='': f'Required field {name} is empty'
|
||||
OAUTH_NOT_CONFIGURED = lambda name='': f"Provider '{name}' is not configured"
|
||||
|
||||
|
||||
class TASKS(str, Enum):
|
||||
def __str__(self) -> str:
|
||||
return super().__str__()
|
||||
|
||||
DEFAULT = lambda task="": f"{task if task else 'generation'}"
|
||||
TITLE_GENERATION = "title_generation"
|
||||
FOLLOW_UP_GENERATION = "follow_up_generation"
|
||||
TAGS_GENERATION = "tags_generation"
|
||||
EMOJI_GENERATION = "emoji_generation"
|
||||
QUERY_GENERATION = "query_generation"
|
||||
IMAGE_PROMPT_GENERATION = "image_prompt_generation"
|
||||
AUTOCOMPLETE_GENERATION = "autocomplete_generation"
|
||||
FUNCTION_CALLING = "function_calling"
|
||||
MOA_RESPONSE_GENERATION = "moa_response_generation"
|
||||
DEFAULT = lambda task='': f'{task if task else "generation"}'
|
||||
TITLE_GENERATION = 'title_generation'
|
||||
FOLLOW_UP_GENERATION = 'follow_up_generation'
|
||||
TAGS_GENERATION = 'tags_generation'
|
||||
EMOJI_GENERATION = 'emoji_generation'
|
||||
QUERY_GENERATION = 'query_generation'
|
||||
IMAGE_PROMPT_GENERATION = 'image_prompt_generation'
|
||||
AUTOCOMPLETE_GENERATION = 'autocomplete_generation'
|
||||
FUNCTION_CALLING = 'function_calling'
|
||||
MOA_RESPONSE_GENERATION = 'moa_response_generation'
|
||||
|
||||
+503
-366
File diff suppressed because it is too large
Load Diff
+113
-117
@@ -34,10 +34,10 @@ from open_webui.utils.plugin import (
|
||||
load_function_module_by_id,
|
||||
get_function_module_from_cache,
|
||||
)
|
||||
from open_webui.utils.tools import get_tools
|
||||
from open_webui.utils.access_control import has_access
|
||||
from open_webui.utils.access_control import check_model_access
|
||||
|
||||
from open_webui.env import GLOBAL_LOG_LEVEL
|
||||
from open_webui.env import GLOBAL_LOG_LEVEL, BYPASS_MODEL_ACCESS_CONTROL
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
|
||||
from open_webui.utils.misc import (
|
||||
add_or_update_system_message,
|
||||
@@ -51,25 +51,22 @@ from open_webui.utils.payload import (
|
||||
apply_system_prompt_to_body,
|
||||
)
|
||||
|
||||
|
||||
logging.basicConfig(stream=sys.stdout, level=GLOBAL_LOG_LEVEL)
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_function_module_by_id(request: Request, pipe_id: str):
|
||||
function_module, _, _ = get_function_module_from_cache(request, pipe_id)
|
||||
async def get_function_module_by_id(request: Request, pipe_id: str):
|
||||
function_module, _, _ = await get_function_module_from_cache(request, pipe_id)
|
||||
|
||||
if hasattr(function_module, "valves") and hasattr(function_module, "Valves"):
|
||||
if hasattr(function_module, 'valves') and hasattr(function_module, 'Valves'):
|
||||
Valves = function_module.Valves
|
||||
valves = Functions.get_function_valves_by_id(pipe_id)
|
||||
valves = await Functions.get_function_valves_by_id(pipe_id)
|
||||
|
||||
if valves:
|
||||
try:
|
||||
function_module.valves = Valves(
|
||||
**{k: v for k, v in valves.items() if v is not None}
|
||||
)
|
||||
function_module.valves = Valves(**{k: v for k, v in valves.items() if v is not None})
|
||||
except Exception as e:
|
||||
log.exception(f"Error loading valves for function {pipe_id}: {e}")
|
||||
log.exception(f'Error loading valves for function {pipe_id}: {e}')
|
||||
raise e
|
||||
else:
|
||||
function_module.valves = Valves()
|
||||
@@ -78,19 +75,19 @@ def get_function_module_by_id(request: Request, pipe_id: str):
|
||||
|
||||
|
||||
async def get_function_models(request):
|
||||
pipes = Functions.get_functions_by_type("pipe", active_only=True)
|
||||
pipes = await Functions.get_functions_by_type('pipe', active_only=True)
|
||||
pipe_models = []
|
||||
|
||||
for pipe in pipes:
|
||||
try:
|
||||
function_module = get_function_module_by_id(request, pipe.id)
|
||||
function_module = await get_function_module_by_id(request, pipe.id)
|
||||
|
||||
has_user_valves = False
|
||||
if hasattr(function_module, "UserValves"):
|
||||
if hasattr(function_module, 'UserValves'):
|
||||
has_user_valves = True
|
||||
|
||||
# Check if function is a manifold
|
||||
if hasattr(function_module, "pipes"):
|
||||
if hasattr(function_module, 'pipes'):
|
||||
sub_pipes = []
|
||||
|
||||
# Handle pipes being a list, sync function, or async function
|
||||
@@ -106,32 +103,30 @@ async def get_function_models(request):
|
||||
log.exception(e)
|
||||
sub_pipes = []
|
||||
|
||||
log.debug(
|
||||
f"get_function_models: function '{pipe.id}' is a manifold of {sub_pipes}"
|
||||
)
|
||||
log.debug(f"get_function_models: function '{pipe.id}' is a manifold of {sub_pipes}")
|
||||
|
||||
for p in sub_pipes:
|
||||
sub_pipe_id = f'{pipe.id}.{p["id"]}'
|
||||
sub_pipe_name = p["name"]
|
||||
sub_pipe_name = p['name']
|
||||
|
||||
if hasattr(function_module, "name"):
|
||||
sub_pipe_name = f"{function_module.name}{sub_pipe_name}"
|
||||
if hasattr(function_module, 'name'):
|
||||
sub_pipe_name = f'{function_module.name}{sub_pipe_name}'
|
||||
|
||||
pipe_flag = {"type": pipe.type}
|
||||
pipe_flag = {'type': pipe.type}
|
||||
|
||||
pipe_models.append(
|
||||
{
|
||||
"id": sub_pipe_id,
|
||||
"name": sub_pipe_name,
|
||||
"object": "model",
|
||||
"created": pipe.created_at,
|
||||
"owned_by": "openai",
|
||||
"pipe": pipe_flag,
|
||||
"has_user_valves": has_user_valves,
|
||||
'id': sub_pipe_id,
|
||||
'name': sub_pipe_name,
|
||||
'object': 'model',
|
||||
'created': pipe.created_at,
|
||||
'owned_by': 'openai',
|
||||
'pipe': pipe_flag,
|
||||
'has_user_valves': has_user_valves,
|
||||
}
|
||||
)
|
||||
else:
|
||||
pipe_flag = {"type": "pipe"}
|
||||
pipe_flag = {'type': 'pipe'}
|
||||
|
||||
log.debug(
|
||||
f"get_function_models: function '{pipe.id}' is a single pipe {{ 'id': {pipe.id}, 'name': {pipe.name} }}"
|
||||
@@ -139,13 +134,13 @@ async def get_function_models(request):
|
||||
|
||||
pipe_models.append(
|
||||
{
|
||||
"id": pipe.id,
|
||||
"name": pipe.name,
|
||||
"object": "model",
|
||||
"created": pipe.created_at,
|
||||
"owned_by": "openai",
|
||||
"pipe": pipe_flag,
|
||||
"has_user_valves": has_user_valves,
|
||||
'id': pipe.id,
|
||||
'name': pipe.name,
|
||||
'object': 'model',
|
||||
'created': pipe.created_at,
|
||||
'owned_by': 'openai',
|
||||
'pipe': pipe_flag,
|
||||
'has_user_valves': has_user_valves,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
@@ -155,9 +150,7 @@ async def get_function_models(request):
|
||||
return pipe_models
|
||||
|
||||
|
||||
async def generate_function_chat_completion(
|
||||
request, form_data, user, models: dict = {}
|
||||
):
|
||||
async def generate_function_chat_completion(request, form_data, user, models: dict = {}):
|
||||
async def execute_pipe(pipe, params):
|
||||
if inspect.iscoroutinefunction(pipe):
|
||||
return await pipe(**params)
|
||||
@@ -168,35 +161,35 @@ async def generate_function_chat_completion(
|
||||
if isinstance(res, str):
|
||||
return res
|
||||
if isinstance(res, Generator):
|
||||
return "".join(map(str, res))
|
||||
return ''.join(map(str, res))
|
||||
if isinstance(res, AsyncGenerator):
|
||||
return "".join([str(stream) async for stream in res])
|
||||
return ''.join([str(stream) async for stream in res])
|
||||
|
||||
def process_line(form_data: dict, line):
|
||||
if isinstance(line, BaseModel):
|
||||
line = line.model_dump_json()
|
||||
line = f"data: {line}"
|
||||
line = f'data: {line}'
|
||||
if isinstance(line, dict):
|
||||
line = f"data: {json.dumps(line)}"
|
||||
line = f'data: {json.dumps(line)}'
|
||||
|
||||
try:
|
||||
line = line.decode("utf-8")
|
||||
line = line.decode('utf-8')
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if line.startswith("data:"):
|
||||
return f"{line}\n\n"
|
||||
if line.startswith('data:'):
|
||||
return f'{line}\n\n'
|
||||
else:
|
||||
line = openai_chat_chunk_message_template(form_data["model"], line)
|
||||
return f"data: {json.dumps(line)}\n\n"
|
||||
line = openai_chat_chunk_message_template(form_data['model'], line)
|
||||
return f'data: {json.dumps(line)}\n\n'
|
||||
|
||||
def get_pipe_id(form_data: dict) -> str:
|
||||
pipe_id = form_data["model"]
|
||||
if "." in pipe_id:
|
||||
pipe_id, _ = pipe_id.split(".", 1)
|
||||
pipe_id = form_data['model']
|
||||
if '.' in pipe_id:
|
||||
pipe_id, _ = pipe_id.split('.', 1)
|
||||
return pipe_id
|
||||
|
||||
def get_function_params(function_module, form_data, user, extra_params=None):
|
||||
async def get_function_params(function_module, form_data, user, extra_params=None):
|
||||
if extra_params is None:
|
||||
extra_params = {}
|
||||
|
||||
@@ -204,27 +197,25 @@ async def generate_function_chat_completion(
|
||||
|
||||
# Get the signature of the function
|
||||
sig = inspect.signature(function_module.pipe)
|
||||
params = {"body": form_data} | {
|
||||
k: v for k, v in extra_params.items() if k in sig.parameters
|
||||
}
|
||||
params = {'body': form_data} | {k: v for k, v in extra_params.items() if k in sig.parameters}
|
||||
|
||||
if "__user__" in params and hasattr(function_module, "UserValves"):
|
||||
user_valves = Functions.get_user_valves_by_id_and_user_id(pipe_id, user.id)
|
||||
if '__user__' in params and hasattr(function_module, 'UserValves'):
|
||||
user_valves = await Functions.get_user_valves_by_id_and_user_id(pipe_id, user.id)
|
||||
try:
|
||||
params["__user__"]["valves"] = function_module.UserValves(**user_valves)
|
||||
params['__user__']['valves'] = function_module.UserValves(**user_valves)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
params["__user__"]["valves"] = function_module.UserValves()
|
||||
params['__user__']['valves'] = function_module.UserValves()
|
||||
|
||||
return params
|
||||
|
||||
model_id = form_data.get("model")
|
||||
model_info = Models.get_model_by_id(model_id)
|
||||
model_id = form_data.get('model')
|
||||
model_info = await Models.get_model_by_id(model_id)
|
||||
|
||||
metadata = form_data.pop("metadata", {})
|
||||
metadata = form_data.pop('metadata', {})
|
||||
|
||||
files = metadata.get("files", [])
|
||||
tool_ids = metadata.get("tool_ids", [])
|
||||
files = metadata.get('files', [])
|
||||
tool_ids = metadata.get('tool_ids', [])
|
||||
# Check if tool_ids is None
|
||||
if tool_ids is None:
|
||||
tool_ids = []
|
||||
@@ -235,66 +226,73 @@ async def generate_function_chat_completion(
|
||||
__task_body__ = None
|
||||
|
||||
if metadata:
|
||||
if all(k in metadata for k in ("session_id", "chat_id", "message_id")):
|
||||
__event_emitter__ = get_event_emitter(metadata)
|
||||
__event_call__ = get_event_call(metadata)
|
||||
__task__ = metadata.get("task", None)
|
||||
__task_body__ = metadata.get("task_body", None)
|
||||
if all(k in metadata for k in ('session_id', 'chat_id', 'message_id')):
|
||||
__event_emitter__ = await get_event_emitter(metadata)
|
||||
__event_call__ = await get_event_call(metadata)
|
||||
__task__ = metadata.get('task', None)
|
||||
__task_body__ = metadata.get('task_body', None)
|
||||
|
||||
oauth_token = None
|
||||
try:
|
||||
if request.cookies.get("oauth_session_id", None):
|
||||
oauth_session_id = request.cookies.get('oauth_session_id', None)
|
||||
if oauth_session_id:
|
||||
oauth_token = await request.app.state.oauth_manager.get_oauth_token(
|
||||
user.id,
|
||||
request.cookies.get("oauth_session_id", None),
|
||||
oauth_session_id,
|
||||
)
|
||||
|
||||
# Fallback: no cookie (automation, API key, etc.) — use most recent session
|
||||
if oauth_token is None:
|
||||
from open_webui.models.oauth_sessions import OAuthSessions
|
||||
|
||||
sessions = await OAuthSessions.get_sessions_by_user_id(user.id)
|
||||
if sessions:
|
||||
best = max(sessions, key=lambda s: s.updated_at)
|
||||
oauth_token = await request.app.state.oauth_manager.get_oauth_token(
|
||||
user.id,
|
||||
best.id,
|
||||
)
|
||||
except Exception as e:
|
||||
log.error(f"Error getting OAuth token: {e}")
|
||||
log.error(f'Error getting OAuth token: {e}')
|
||||
|
||||
extra_params = {
|
||||
"__event_emitter__": __event_emitter__,
|
||||
"__event_call__": __event_call__,
|
||||
"__chat_id__": metadata.get("chat_id", None),
|
||||
"__session_id__": metadata.get("session_id", None),
|
||||
"__message_id__": metadata.get("message_id", None),
|
||||
"__task__": __task__,
|
||||
"__task_body__": __task_body__,
|
||||
"__files__": files,
|
||||
"__user__": user.model_dump() if isinstance(user, UserModel) else {},
|
||||
"__metadata__": metadata,
|
||||
"__oauth_token__": oauth_token,
|
||||
"__request__": request,
|
||||
'__event_emitter__': __event_emitter__,
|
||||
'__event_call__': __event_call__,
|
||||
'__chat_id__': metadata.get('chat_id', None),
|
||||
'__session_id__': metadata.get('session_id', None),
|
||||
'__message_id__': metadata.get('message_id', None),
|
||||
'__task__': __task__,
|
||||
'__task_body__': __task_body__,
|
||||
'__files__': files,
|
||||
'__user__': user.model_dump() if isinstance(user, UserModel) else {},
|
||||
'__metadata__': metadata,
|
||||
'__oauth_token__': oauth_token,
|
||||
'__request__': request,
|
||||
}
|
||||
extra_params["__tools__"] = await get_tools(
|
||||
request,
|
||||
tool_ids,
|
||||
user,
|
||||
{
|
||||
**extra_params,
|
||||
"__model__": models.get(form_data["model"], None),
|
||||
"__messages__": form_data["messages"],
|
||||
"__files__": files,
|
||||
},
|
||||
)
|
||||
extra_params['__tools__'] = metadata.get('tools', {})
|
||||
|
||||
if model_info:
|
||||
if model_info.base_model_id:
|
||||
form_data["model"] = model_info.base_model_id
|
||||
form_data['model'] = model_info.base_model_id
|
||||
|
||||
if not BYPASS_MODEL_ACCESS_CONTROL:
|
||||
bypass = isinstance(user, UserModel) and user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL
|
||||
await check_model_access(user if isinstance(user, UserModel) else UserModel(**user), model_info, bypass)
|
||||
|
||||
params = model_info.params.model_dump()
|
||||
|
||||
if params:
|
||||
system = params.pop("system", None)
|
||||
system = params.pop('system', None)
|
||||
form_data = apply_model_params_to_body_openai(params, form_data)
|
||||
form_data = apply_system_prompt_to_body(system, form_data, metadata, user)
|
||||
|
||||
pipe_id = get_pipe_id(form_data)
|
||||
function_module = get_function_module_by_id(request, pipe_id)
|
||||
function_module = await get_function_module_by_id(request, pipe_id)
|
||||
|
||||
pipe = function_module.pipe
|
||||
params = get_function_params(function_module, form_data, user, extra_params)
|
||||
params = await get_function_params(function_module, form_data, user, extra_params)
|
||||
|
||||
if form_data.get("stream", False):
|
||||
if form_data.get('stream', False):
|
||||
|
||||
async def stream_content():
|
||||
try:
|
||||
@@ -306,17 +304,17 @@ async def generate_function_chat_completion(
|
||||
yield data
|
||||
return
|
||||
if isinstance(res, dict):
|
||||
yield f"data: {json.dumps(res)}\n\n"
|
||||
yield f'data: {json.dumps(res)}\n\n'
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Error: {e}")
|
||||
yield f"data: {json.dumps({'error': {'detail':str(e)}})}\n\n"
|
||||
log.error(f'Error: {e}')
|
||||
yield f'data: {json.dumps({"error": {"detail": str(e)}})}\n\n'
|
||||
return
|
||||
|
||||
if isinstance(res, str):
|
||||
message = openai_chat_chunk_message_template(form_data["model"], res)
|
||||
yield f"data: {json.dumps(message)}\n\n"
|
||||
message = openai_chat_chunk_message_template(form_data['model'], res)
|
||||
yield f'data: {json.dumps(message)}\n\n'
|
||||
|
||||
if isinstance(res, Iterator):
|
||||
for line in res:
|
||||
@@ -327,21 +325,19 @@ async def generate_function_chat_completion(
|
||||
yield process_line(form_data, line)
|
||||
|
||||
if isinstance(res, str) or isinstance(res, Generator):
|
||||
finish_message = openai_chat_chunk_message_template(
|
||||
form_data["model"], ""
|
||||
)
|
||||
finish_message["choices"][0]["finish_reason"] = "stop"
|
||||
yield f"data: {json.dumps(finish_message)}\n\n"
|
||||
yield "data: [DONE]"
|
||||
finish_message = openai_chat_chunk_message_template(form_data['model'], '')
|
||||
finish_message['choices'][0]['finish_reason'] = 'stop'
|
||||
yield f'data: {json.dumps(finish_message)}\n\n'
|
||||
yield 'data: [DONE]'
|
||||
|
||||
return StreamingResponse(stream_content(), media_type="text/event-stream")
|
||||
return StreamingResponse(stream_content(), media_type='text/event-stream')
|
||||
else:
|
||||
try:
|
||||
res = await execute_pipe(pipe, params)
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Error: {e}")
|
||||
return {"error": {"detail": str(e)}}
|
||||
log.error(f'Error: {e}')
|
||||
return {'error': {'detail': str(e)}}
|
||||
|
||||
if isinstance(res, StreamingResponse) or isinstance(res, dict):
|
||||
return res
|
||||
@@ -349,4 +345,4 @@ async def generate_function_chat_completion(
|
||||
return res.model_dump()
|
||||
|
||||
message = await get_message_content(res)
|
||||
return openai_chat_completion_message_template(form_data["model"], message)
|
||||
return openai_chat_completion_message_template(form_data['model'], message)
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
import os
|
||||
import json
|
||||
import logging
|
||||
from contextlib import contextmanager
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from typing import Any, Optional
|
||||
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
|
||||
|
||||
from open_webui.internal.wrappers import register_connection
|
||||
from open_webui.env import (
|
||||
@@ -14,10 +15,18 @@ from open_webui.env import (
|
||||
DATABASE_POOL_SIZE,
|
||||
DATABASE_POOL_TIMEOUT,
|
||||
DATABASE_ENABLE_SQLITE_WAL,
|
||||
DATABASE_ENABLE_SESSION_SHARING,
|
||||
DATABASE_SQLITE_PRAGMA_SYNCHRONOUS,
|
||||
DATABASE_SQLITE_PRAGMA_BUSY_TIMEOUT,
|
||||
DATABASE_SQLITE_PRAGMA_CACHE_SIZE,
|
||||
DATABASE_SQLITE_PRAGMA_TEMP_STORE,
|
||||
DATABASE_SQLITE_PRAGMA_MMAP_SIZE,
|
||||
DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT,
|
||||
ENABLE_DB_MIGRATIONS,
|
||||
)
|
||||
from peewee_migrate import Router
|
||||
from sqlalchemy import Dialect, create_engine, MetaData, event, types
|
||||
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
|
||||
from sqlalchemy.ext.declarative import declarative_base
|
||||
from sqlalchemy.orm import scoped_session, sessionmaker, Session
|
||||
from sqlalchemy.pool import QueuePool, NullPool
|
||||
@@ -27,6 +36,86 @@ from typing_extensions import Self
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── SSL URL normalization (used by sync engine & Alembic migrations) ─
|
||||
#
|
||||
# psycopg2 (sync) needs ``sslmode=`` in the connection string (it does
|
||||
# not recognise the bare ``ssl=`` key that some ORMs emit). The helpers
|
||||
# below strip all SSL-related query params, normalise them, and
|
||||
# reattach them in the canonical libpq form.
|
||||
#
|
||||
# The **async** engine now uses psycopg (v3), which speaks libpq
|
||||
# natively, so it needs no translation at all — the DATABASE_URL is
|
||||
# passed through as-is.
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _pop_first(params: dict[str, list[str]], key: str) -> str | None:
|
||||
"""Pop a single-valued query param, returning ``None`` if absent."""
|
||||
values = params.pop(key, None)
|
||||
return values[0] if values else None
|
||||
|
||||
|
||||
def _is_postgres_url(url: str) -> bool:
|
||||
"""Return True if *url* looks like a PostgreSQL connection string."""
|
||||
return bool(url) and any(url.startswith(p) for p in ('postgresql://', 'postgresql+', 'postgres://'))
|
||||
|
||||
|
||||
def extract_ssl_params_from_url(url: str) -> tuple[str, dict[str, str]]:
|
||||
"""Strip SSL query-string parameters from a PostgreSQL URL.
|
||||
|
||||
Returns ``(url_without_ssl, ssl_dict)`` where *ssl_dict* maps
|
||||
canonical libpq key names (``sslmode``, ``sslrootcert``, …) to
|
||||
their values. Non-PostgreSQL URLs are returned unchanged with an
|
||||
empty dict.
|
||||
"""
|
||||
if not _is_postgres_url(url):
|
||||
return url, {}
|
||||
|
||||
parsed = urlparse(url)
|
||||
qp = parse_qs(parsed.query, keep_blank_values=True)
|
||||
|
||||
# Prefer sslmode (libpq canonical) over the bare ``ssl`` key.
|
||||
sslmode_val = _pop_first(qp, 'sslmode')
|
||||
ssl_val = _pop_first(qp, 'ssl')
|
||||
ssl_mode = sslmode_val or ssl_val
|
||||
|
||||
ssl_dict: dict[str, str] = {}
|
||||
if ssl_mode:
|
||||
ssl_dict['sslmode'] = ssl_mode
|
||||
for key in ('sslrootcert', 'sslcert', 'sslkey', 'sslcrl'):
|
||||
val = _pop_first(qp, key)
|
||||
if val:
|
||||
ssl_dict[key] = val
|
||||
|
||||
if not ssl_dict:
|
||||
return url, ssl_dict
|
||||
|
||||
cleaned_query = urlencode(qp, doseq=True)
|
||||
return urlunparse(parsed._replace(query=cleaned_query)), ssl_dict
|
||||
|
||||
|
||||
def reattach_ssl_params_to_url(url_without_ssl: str, ssl_dict: dict[str, str]) -> str:
|
||||
"""Re-append SSL query-string parameters to a cleaned PostgreSQL URL.
|
||||
|
||||
Used for psycopg2/libpq consumers that expect ``sslmode`` and the
|
||||
certificate-file keys in the connection string.
|
||||
"""
|
||||
if not ssl_dict:
|
||||
return url_without_ssl
|
||||
|
||||
parts = [f'{k}={v}' for k, v in ssl_dict.items() if v]
|
||||
if not parts:
|
||||
return url_without_ssl
|
||||
|
||||
sep = '&' if '?' in url_without_ssl else '?'
|
||||
return f'{url_without_ssl}{sep}{"&".join(parts)}'
|
||||
|
||||
|
||||
# Backwards-compatible aliases for external callers.
|
||||
extract_ssl_mode_from_url = extract_ssl_params_from_url
|
||||
reattach_ssl_mode_to_url = reattach_ssl_params_to_url
|
||||
|
||||
|
||||
class JSONField(types.TypeDecorator):
|
||||
impl = types.Text
|
||||
cache_ok = True
|
||||
@@ -52,20 +141,23 @@ class JSONField(types.TypeDecorator):
|
||||
# Workaround to handle the peewee migration
|
||||
# This is required to ensure the peewee migration is handled before the alembic migration
|
||||
def handle_peewee_migration(DATABASE_URL):
|
||||
# db = None
|
||||
db = None
|
||||
try:
|
||||
# Normalize SSL params so psycopg2 always sees `sslmode=` (never `ssl=`)
|
||||
# and cert-file params are preserved in the connection string.
|
||||
url_without_ssl, ssl_params = extract_ssl_params_from_url(DATABASE_URL)
|
||||
normalized_url = reattach_ssl_params_to_url(url_without_ssl, ssl_params)
|
||||
|
||||
# Replace the postgresql:// with postgres:// to handle the peewee migration
|
||||
db = register_connection(DATABASE_URL.replace("postgresql://", "postgres://"))
|
||||
migrate_dir = OPEN_WEBUI_DIR / "internal" / "migrations"
|
||||
db = register_connection(normalized_url.replace('postgresql://', 'postgres://'))
|
||||
migrate_dir = OPEN_WEBUI_DIR / 'internal' / 'migrations'
|
||||
router = Router(db, logger=log, migrate_dir=migrate_dir)
|
||||
router.run()
|
||||
db.close()
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Failed to initialize the database connection: {e}")
|
||||
log.warning(
|
||||
"Hint: If your database password contains special characters, you may need to URL-encode it."
|
||||
)
|
||||
log.error(f'Failed to initialize the database connection: {e}')
|
||||
log.warning('Hint: If your database password contains special characters, you may need to URL-encode it.')
|
||||
raise
|
||||
finally:
|
||||
# Properly closing the database connection
|
||||
@@ -73,25 +165,61 @@ def handle_peewee_migration(DATABASE_URL):
|
||||
db.close()
|
||||
|
||||
# Assert if db connection has been closed
|
||||
assert db.is_closed(), "Database connection is still open."
|
||||
if db is not None:
|
||||
assert db.is_closed(), 'Database connection is still open.'
|
||||
|
||||
|
||||
if ENABLE_DB_MIGRATIONS:
|
||||
handle_peewee_migration(DATABASE_URL)
|
||||
|
||||
|
||||
SQLALCHEMY_DATABASE_URL = DATABASE_URL
|
||||
# Normalize SSL params from the URL once; the sync engine needs them
|
||||
# reattached in canonical libpq form for psycopg2.
|
||||
_url_without_ssl, _ssl_dict = extract_ssl_params_from_url(DATABASE_URL)
|
||||
|
||||
# For psycopg2 (sync engine), re-append sslmode + cert-file params.
|
||||
SQLALCHEMY_DATABASE_URL = reattach_ssl_params_to_url(_url_without_ssl, _ssl_dict) if _ssl_dict else DATABASE_URL
|
||||
|
||||
|
||||
def _make_async_url(url: str) -> str:
|
||||
"""Convert a sync database URL to its async driver equivalent.
|
||||
|
||||
The async engine uses psycopg (v3) which speaks libpq natively,
|
||||
so all standard connection-string parameters (``sslmode``,
|
||||
``options``, ``target_session_attrs``, etc.) are passed through
|
||||
without any translation.
|
||||
"""
|
||||
if url.startswith('sqlite+sqlcipher://'):
|
||||
raise ValueError(
|
||||
'sqlite+sqlcipher:// URLs are not supported with async engine. '
|
||||
'Use standard sqlite:// or postgresql:// instead.'
|
||||
)
|
||||
if url.startswith('sqlite:///') or url.startswith('sqlite://'):
|
||||
return url.replace('sqlite://', 'sqlite+aiosqlite://', 1)
|
||||
# psycopg v3 — auto-selects async mode with create_async_engine
|
||||
if url.startswith('postgresql+psycopg2://'):
|
||||
return url.replace('postgresql+psycopg2://', 'postgresql+psycopg://', 1)
|
||||
if url.startswith('postgresql://'):
|
||||
return url.replace('postgresql://', 'postgresql+psycopg://', 1)
|
||||
if url.startswith('postgres://'):
|
||||
return url.replace('postgres://', 'postgresql+psycopg://', 1)
|
||||
# For other dialects, return as-is and let SQLAlchemy handle it
|
||||
return url
|
||||
|
||||
|
||||
# ============================================================
|
||||
# SYNC ENGINE (used only for: startup migrations, config loading,
|
||||
# Alembic, peewee migration, health checks)
|
||||
# ============================================================
|
||||
|
||||
# Handle SQLCipher URLs
|
||||
if SQLALCHEMY_DATABASE_URL.startswith("sqlite+sqlcipher://"):
|
||||
database_password = os.environ.get("DATABASE_PASSWORD")
|
||||
if not database_password or database_password.strip() == "":
|
||||
raise ValueError(
|
||||
"DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs"
|
||||
)
|
||||
if SQLALCHEMY_DATABASE_URL.startswith('sqlite+sqlcipher://'):
|
||||
database_password = os.environ.get('DATABASE_PASSWORD')
|
||||
if not database_password or database_password.strip() == '':
|
||||
raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
|
||||
|
||||
# Extract database path from SQLCipher URL
|
||||
db_path = SQLALCHEMY_DATABASE_URL.replace("sqlite+sqlcipher://", "")
|
||||
db_path = SQLALCHEMY_DATABASE_URL.replace('sqlite+sqlcipher://', '')
|
||||
|
||||
# Create a custom creator function that uses sqlcipher3
|
||||
def create_sqlcipher_connection():
|
||||
@@ -101,28 +229,63 @@ if SQLALCHEMY_DATABASE_URL.startswith("sqlite+sqlcipher://"):
|
||||
conn.execute(f"PRAGMA key = '{database_password}'")
|
||||
return conn
|
||||
|
||||
engine = create_engine(
|
||||
"sqlite://", # Dummy URL since we're using creator
|
||||
creator=create_sqlcipher_connection,
|
||||
echo=False,
|
||||
)
|
||||
# The dummy "sqlite://" URL would cause SQLAlchemy to auto-select
|
||||
# SingletonThreadPool, which non-deterministically closes in-use
|
||||
# connections when thread count exceeds pool_size, leading to segfaults
|
||||
# in the native sqlcipher3 C library. Use NullPool by default for safety,
|
||||
# or QueuePool if DATABASE_POOL_SIZE is explicitly configured.
|
||||
if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0:
|
||||
engine = create_engine(
|
||||
'sqlite://',
|
||||
creator=create_sqlcipher_connection,
|
||||
pool_size=DATABASE_POOL_SIZE,
|
||||
max_overflow=DATABASE_POOL_MAX_OVERFLOW,
|
||||
pool_timeout=DATABASE_POOL_TIMEOUT,
|
||||
pool_recycle=DATABASE_POOL_RECYCLE,
|
||||
pool_pre_ping=True,
|
||||
poolclass=QueuePool,
|
||||
echo=False,
|
||||
)
|
||||
else:
|
||||
engine = create_engine(
|
||||
'sqlite://',
|
||||
creator=create_sqlcipher_connection,
|
||||
poolclass=NullPool,
|
||||
echo=False,
|
||||
)
|
||||
|
||||
log.info("Connected to encrypted SQLite database using SQLCipher")
|
||||
log.info('Connected to encrypted SQLite database using SQLCipher')
|
||||
|
||||
elif "sqlite" in SQLALCHEMY_DATABASE_URL:
|
||||
engine = create_engine(
|
||||
SQLALCHEMY_DATABASE_URL, connect_args={"check_same_thread": False}
|
||||
)
|
||||
elif 'sqlite' in SQLALCHEMY_DATABASE_URL:
|
||||
engine = create_engine(SQLALCHEMY_DATABASE_URL, connect_args={'check_same_thread': False})
|
||||
|
||||
def on_connect(dbapi_connection, connection_record):
|
||||
def _apply_sqlite_pragmas(dbapi_connection):
|
||||
"""Apply all configured SQLite PRAGMAs to a raw DBAPI connection."""
|
||||
cursor = dbapi_connection.cursor()
|
||||
if DATABASE_ENABLE_SQLITE_WAL:
|
||||
cursor.execute("PRAGMA journal_mode=WAL")
|
||||
cursor.execute('PRAGMA journal_mode=WAL')
|
||||
else:
|
||||
cursor.execute("PRAGMA journal_mode=DELETE")
|
||||
cursor.execute('PRAGMA journal_mode=DELETE')
|
||||
|
||||
# Each PRAGMA is skipped when its env var is empty, allowing opt-out.
|
||||
if DATABASE_SQLITE_PRAGMA_SYNCHRONOUS:
|
||||
cursor.execute(f'PRAGMA synchronous={DATABASE_SQLITE_PRAGMA_SYNCHRONOUS}')
|
||||
if DATABASE_SQLITE_PRAGMA_BUSY_TIMEOUT:
|
||||
cursor.execute(f'PRAGMA busy_timeout={DATABASE_SQLITE_PRAGMA_BUSY_TIMEOUT}')
|
||||
if DATABASE_SQLITE_PRAGMA_CACHE_SIZE:
|
||||
cursor.execute(f'PRAGMA cache_size={DATABASE_SQLITE_PRAGMA_CACHE_SIZE}')
|
||||
if DATABASE_SQLITE_PRAGMA_TEMP_STORE:
|
||||
cursor.execute(f'PRAGMA temp_store={DATABASE_SQLITE_PRAGMA_TEMP_STORE}')
|
||||
if DATABASE_SQLITE_PRAGMA_MMAP_SIZE:
|
||||
cursor.execute(f'PRAGMA mmap_size={DATABASE_SQLITE_PRAGMA_MMAP_SIZE}')
|
||||
if DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT:
|
||||
cursor.execute(f'PRAGMA journal_size_limit={DATABASE_SQLITE_PRAGMA_JOURNAL_SIZE_LIMIT}')
|
||||
cursor.close()
|
||||
|
||||
event.listen(engine, "connect", on_connect)
|
||||
def on_connect(dbapi_connection, connection_record):
|
||||
_apply_sqlite_pragmas(dbapi_connection)
|
||||
|
||||
event.listen(engine, 'connect', on_connect)
|
||||
else:
|
||||
if isinstance(DATABASE_POOL_SIZE, int):
|
||||
if DATABASE_POOL_SIZE > 0:
|
||||
@@ -136,22 +299,20 @@ else:
|
||||
poolclass=QueuePool,
|
||||
)
|
||||
else:
|
||||
engine = create_engine(
|
||||
SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool
|
||||
)
|
||||
engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True, poolclass=NullPool)
|
||||
else:
|
||||
engine = create_engine(SQLALCHEMY_DATABASE_URL, pool_pre_ping=True)
|
||||
|
||||
|
||||
SessionLocal = sessionmaker(
|
||||
autocommit=False, autoflush=False, bind=engine, expire_on_commit=False
|
||||
)
|
||||
# Sync session — used ONLY for startup config loading (config.py runs at import time)
|
||||
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine, expire_on_commit=False)
|
||||
metadata_obj = MetaData(schema=DATABASE_SCHEMA)
|
||||
Base = declarative_base(metadata=metadata_obj)
|
||||
ScopedSession = scoped_session(SessionLocal)
|
||||
|
||||
|
||||
def get_session():
|
||||
"""Sync session generator — used ONLY for startup/config operations."""
|
||||
db = SessionLocal()
|
||||
try:
|
||||
yield db
|
||||
@@ -162,10 +323,87 @@ def get_session():
|
||||
get_db = contextmanager(get_session)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def get_db_context(db: Optional[Session] = None):
|
||||
if isinstance(db, Session):
|
||||
# ============================================================
|
||||
# ASYNC ENGINE (used for ALL runtime database operations)
|
||||
# ============================================================
|
||||
|
||||
# psycopg (v3) speaks libpq natively — the full DATABASE_URL is passed
|
||||
# through as-is. SSL params, ``options``, ``target_session_attrs``, etc.
|
||||
# all work without any stripping or translation.
|
||||
ASYNC_SQLALCHEMY_DATABASE_URL = _make_async_url(SQLALCHEMY_DATABASE_URL)
|
||||
|
||||
if 'sqlite' in ASYNC_SQLALCHEMY_DATABASE_URL:
|
||||
# Generous default — async coroutines + no session sharing = high connection demand.
|
||||
_sqlite_pool_size = DATABASE_POOL_SIZE if isinstance(DATABASE_POOL_SIZE, int) and DATABASE_POOL_SIZE > 0 else 512
|
||||
async_engine = create_async_engine(
|
||||
ASYNC_SQLALCHEMY_DATABASE_URL,
|
||||
connect_args={'check_same_thread': False},
|
||||
pool_size=_sqlite_pool_size,
|
||||
pool_timeout=DATABASE_POOL_TIMEOUT,
|
||||
pool_recycle=DATABASE_POOL_RECYCLE,
|
||||
pool_pre_ping=True,
|
||||
)
|
||||
|
||||
@event.listens_for(async_engine.sync_engine, 'connect')
|
||||
def _set_sqlite_pragmas(dbapi_connection, connection_record):
|
||||
_apply_sqlite_pragmas(dbapi_connection)
|
||||
else:
|
||||
if isinstance(DATABASE_POOL_SIZE, int):
|
||||
if DATABASE_POOL_SIZE > 0:
|
||||
async_engine = create_async_engine(
|
||||
ASYNC_SQLALCHEMY_DATABASE_URL,
|
||||
pool_size=DATABASE_POOL_SIZE,
|
||||
max_overflow=DATABASE_POOL_MAX_OVERFLOW,
|
||||
pool_timeout=DATABASE_POOL_TIMEOUT,
|
||||
pool_recycle=DATABASE_POOL_RECYCLE,
|
||||
pool_pre_ping=True,
|
||||
)
|
||||
else:
|
||||
async_engine = create_async_engine(
|
||||
ASYNC_SQLALCHEMY_DATABASE_URL,
|
||||
pool_pre_ping=True,
|
||||
poolclass=NullPool,
|
||||
)
|
||||
else:
|
||||
async_engine = create_async_engine(
|
||||
ASYNC_SQLALCHEMY_DATABASE_URL,
|
||||
pool_pre_ping=True,
|
||||
)
|
||||
|
||||
|
||||
AsyncSessionLocal = async_sessionmaker(
|
||||
bind=async_engine,
|
||||
class_=AsyncSession,
|
||||
autocommit=False,
|
||||
autoflush=False,
|
||||
expire_on_commit=False,
|
||||
)
|
||||
|
||||
|
||||
async def get_async_session():
|
||||
"""Async session generator for FastAPI Depends()."""
|
||||
async with AsyncSessionLocal() as db:
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_async_db():
|
||||
"""Async context manager for use outside of FastAPI dependency injection."""
|
||||
async with AsyncSessionLocal() as db:
|
||||
try:
|
||||
yield db
|
||||
finally:
|
||||
await db.close()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_async_db_context(db: Optional[AsyncSession] = None):
|
||||
"""Async context manager that reuses an existing session if provided and session sharing is enabled."""
|
||||
if isinstance(db, AsyncSession) and DATABASE_ENABLE_SESSION_SHARING:
|
||||
yield db
|
||||
else:
|
||||
with get_db() as session:
|
||||
async with get_async_db() as session:
|
||||
yield session
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -57,7 +56,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
active = pw.BooleanField()
|
||||
|
||||
class Meta:
|
||||
table_name = "auth"
|
||||
table_name = 'auth'
|
||||
|
||||
@migrator.create_model
|
||||
class Chat(pw.Model):
|
||||
@@ -68,7 +67,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "chat"
|
||||
table_name = 'chat'
|
||||
|
||||
@migrator.create_model
|
||||
class ChatIdTag(pw.Model):
|
||||
@@ -79,7 +78,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "chatidtag"
|
||||
table_name = 'chatidtag'
|
||||
|
||||
@migrator.create_model
|
||||
class Document(pw.Model):
|
||||
@@ -93,7 +92,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "document"
|
||||
table_name = 'document'
|
||||
|
||||
@migrator.create_model
|
||||
class Modelfile(pw.Model):
|
||||
@@ -104,7 +103,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "modelfile"
|
||||
table_name = 'modelfile'
|
||||
|
||||
@migrator.create_model
|
||||
class Prompt(pw.Model):
|
||||
@@ -116,7 +115,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "prompt"
|
||||
table_name = 'prompt'
|
||||
|
||||
@migrator.create_model
|
||||
class Tag(pw.Model):
|
||||
@@ -126,7 +125,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
data = pw.TextField(null=True)
|
||||
|
||||
class Meta:
|
||||
table_name = "tag"
|
||||
table_name = 'tag'
|
||||
|
||||
@migrator.create_model
|
||||
class User(pw.Model):
|
||||
@@ -138,7 +137,7 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "user"
|
||||
table_name = 'user'
|
||||
|
||||
|
||||
def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
@@ -150,7 +149,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
active = pw.BooleanField()
|
||||
|
||||
class Meta:
|
||||
table_name = "auth"
|
||||
table_name = 'auth'
|
||||
|
||||
@migrator.create_model
|
||||
class Chat(pw.Model):
|
||||
@@ -161,7 +160,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "chat"
|
||||
table_name = 'chat'
|
||||
|
||||
@migrator.create_model
|
||||
class ChatIdTag(pw.Model):
|
||||
@@ -172,7 +171,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "chatidtag"
|
||||
table_name = 'chatidtag'
|
||||
|
||||
@migrator.create_model
|
||||
class Document(pw.Model):
|
||||
@@ -186,7 +185,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "document"
|
||||
table_name = 'document'
|
||||
|
||||
@migrator.create_model
|
||||
class Modelfile(pw.Model):
|
||||
@@ -197,7 +196,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "modelfile"
|
||||
table_name = 'modelfile'
|
||||
|
||||
@migrator.create_model
|
||||
class Prompt(pw.Model):
|
||||
@@ -209,7 +208,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "prompt"
|
||||
table_name = 'prompt'
|
||||
|
||||
@migrator.create_model
|
||||
class Tag(pw.Model):
|
||||
@@ -219,7 +218,7 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
data = pw.TextField(null=True)
|
||||
|
||||
class Meta:
|
||||
table_name = "tag"
|
||||
table_name = 'tag'
|
||||
|
||||
@migrator.create_model
|
||||
class User(pw.Model):
|
||||
@@ -231,24 +230,24 @@ def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
timestamp = pw.BigIntegerField()
|
||||
|
||||
class Meta:
|
||||
table_name = "user"
|
||||
table_name = 'user'
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_model("user")
|
||||
migrator.remove_model('user')
|
||||
|
||||
migrator.remove_model("tag")
|
||||
migrator.remove_model('tag')
|
||||
|
||||
migrator.remove_model("prompt")
|
||||
migrator.remove_model('prompt')
|
||||
|
||||
migrator.remove_model("modelfile")
|
||||
migrator.remove_model('modelfile')
|
||||
|
||||
migrator.remove_model("document")
|
||||
migrator.remove_model('document')
|
||||
|
||||
migrator.remove_model("chatidtag")
|
||||
migrator.remove_model('chatidtag')
|
||||
|
||||
migrator.remove_model("chat")
|
||||
migrator.remove_model('chat')
|
||||
|
||||
migrator.remove_model("auth")
|
||||
migrator.remove_model('auth')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -37,12 +36,10 @@ with suppress(ImportError):
|
||||
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
migrator.add_fields(
|
||||
"chat", share_id=pw.CharField(max_length=255, null=True, unique=True)
|
||||
)
|
||||
migrator.add_fields('chat', share_id=pw.CharField(max_length=255, null=True, unique=True))
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_fields("chat", "share_id")
|
||||
migrator.remove_fields('chat', 'share_id')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -37,12 +36,10 @@ with suppress(ImportError):
|
||||
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
migrator.add_fields(
|
||||
"user", api_key=pw.CharField(max_length=255, null=True, unique=True)
|
||||
)
|
||||
migrator.add_fields('user', api_key=pw.CharField(max_length=255, null=True, unique=True))
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_fields("user", "api_key")
|
||||
migrator.remove_fields('user', 'api_key')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -37,10 +36,10 @@ with suppress(ImportError):
|
||||
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
migrator.add_fields("chat", archived=pw.BooleanField(default=False))
|
||||
migrator.add_fields('chat', archived=pw.BooleanField(default=False))
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_fields("chat", "archived")
|
||||
migrator.remove_fields('chat', 'archived')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -46,22 +45,20 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
# Adding fields created_at and updated_at to the 'chat' table
|
||||
migrator.add_fields(
|
||||
"chat",
|
||||
'chat',
|
||||
created_at=pw.DateTimeField(null=True), # Allow null for transition
|
||||
updated_at=pw.DateTimeField(null=True), # Allow null for transition
|
||||
)
|
||||
|
||||
# Populate the new fields from an existing 'timestamp' field
|
||||
migrator.sql(
|
||||
"UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL"
|
||||
)
|
||||
migrator.sql('UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL')
|
||||
|
||||
# Now that the data has been copied, remove the original 'timestamp' field
|
||||
migrator.remove_fields("chat", "timestamp")
|
||||
migrator.remove_fields('chat', 'timestamp')
|
||||
|
||||
# Update the fields to be not null now that they are populated
|
||||
migrator.change_fields(
|
||||
"chat",
|
||||
'chat',
|
||||
created_at=pw.DateTimeField(null=False),
|
||||
updated_at=pw.DateTimeField(null=False),
|
||||
)
|
||||
@@ -70,22 +67,20 @@ def migrate_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
def migrate_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
# Adding fields created_at and updated_at to the 'chat' table
|
||||
migrator.add_fields(
|
||||
"chat",
|
||||
'chat',
|
||||
created_at=pw.BigIntegerField(null=True), # Allow null for transition
|
||||
updated_at=pw.BigIntegerField(null=True), # Allow null for transition
|
||||
)
|
||||
|
||||
# Populate the new fields from an existing 'timestamp' field
|
||||
migrator.sql(
|
||||
"UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL"
|
||||
)
|
||||
migrator.sql('UPDATE chat SET created_at = timestamp, updated_at = timestamp WHERE timestamp IS NOT NULL')
|
||||
|
||||
# Now that the data has been copied, remove the original 'timestamp' field
|
||||
migrator.remove_fields("chat", "timestamp")
|
||||
migrator.remove_fields('chat', 'timestamp')
|
||||
|
||||
# Update the fields to be not null now that they are populated
|
||||
migrator.change_fields(
|
||||
"chat",
|
||||
'chat',
|
||||
created_at=pw.BigIntegerField(null=False),
|
||||
updated_at=pw.BigIntegerField(null=False),
|
||||
)
|
||||
@@ -102,29 +97,29 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
|
||||
def rollback_sqlite(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
# Recreate the timestamp field initially allowing null values for safe transition
|
||||
migrator.add_fields("chat", timestamp=pw.DateTimeField(null=True))
|
||||
migrator.add_fields('chat', timestamp=pw.DateTimeField(null=True))
|
||||
|
||||
# Copy the earliest created_at date back into the new timestamp field
|
||||
# This assumes created_at was originally a copy of timestamp
|
||||
migrator.sql("UPDATE chat SET timestamp = created_at")
|
||||
migrator.sql('UPDATE chat SET timestamp = created_at')
|
||||
|
||||
# Remove the created_at and updated_at fields
|
||||
migrator.remove_fields("chat", "created_at", "updated_at")
|
||||
migrator.remove_fields('chat', 'created_at', 'updated_at')
|
||||
|
||||
# Finally, alter the timestamp field to not allow nulls if that was the original setting
|
||||
migrator.change_fields("chat", timestamp=pw.DateTimeField(null=False))
|
||||
migrator.change_fields('chat', timestamp=pw.DateTimeField(null=False))
|
||||
|
||||
|
||||
def rollback_external(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
# Recreate the timestamp field initially allowing null values for safe transition
|
||||
migrator.add_fields("chat", timestamp=pw.BigIntegerField(null=True))
|
||||
migrator.add_fields('chat', timestamp=pw.BigIntegerField(null=True))
|
||||
|
||||
# Copy the earliest created_at date back into the new timestamp field
|
||||
# This assumes created_at was originally a copy of timestamp
|
||||
migrator.sql("UPDATE chat SET timestamp = created_at")
|
||||
migrator.sql('UPDATE chat SET timestamp = created_at')
|
||||
|
||||
# Remove the created_at and updated_at fields
|
||||
migrator.remove_fields("chat", "created_at", "updated_at")
|
||||
migrator.remove_fields('chat', 'created_at', 'updated_at')
|
||||
|
||||
# Finally, alter the timestamp field to not allow nulls if that was the original setting
|
||||
migrator.change_fields("chat", timestamp=pw.BigIntegerField(null=False))
|
||||
migrator.change_fields('chat', timestamp=pw.BigIntegerField(null=False))
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -39,45 +38,45 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
|
||||
# Alter the tables with timestamps
|
||||
migrator.change_fields(
|
||||
"chatidtag",
|
||||
'chatidtag',
|
||||
timestamp=pw.BigIntegerField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"document",
|
||||
'document',
|
||||
timestamp=pw.BigIntegerField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"modelfile",
|
||||
'modelfile',
|
||||
timestamp=pw.BigIntegerField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"prompt",
|
||||
'prompt',
|
||||
timestamp=pw.BigIntegerField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"user",
|
||||
'user',
|
||||
timestamp=pw.BigIntegerField(),
|
||||
)
|
||||
# Alter the tables with varchar to text where necessary
|
||||
migrator.change_fields(
|
||||
"auth",
|
||||
'auth',
|
||||
password=pw.TextField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"chat",
|
||||
'chat',
|
||||
title=pw.TextField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"document",
|
||||
'document',
|
||||
title=pw.TextField(),
|
||||
filename=pw.TextField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"prompt",
|
||||
'prompt',
|
||||
title=pw.TextField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"user",
|
||||
'user',
|
||||
profile_image_url=pw.TextField(),
|
||||
)
|
||||
|
||||
@@ -88,43 +87,43 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
if isinstance(database, pw.SqliteDatabase):
|
||||
# Alter the tables with timestamps
|
||||
migrator.change_fields(
|
||||
"chatidtag",
|
||||
'chatidtag',
|
||||
timestamp=pw.DateField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"document",
|
||||
'document',
|
||||
timestamp=pw.DateField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"modelfile",
|
||||
'modelfile',
|
||||
timestamp=pw.DateField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"prompt",
|
||||
'prompt',
|
||||
timestamp=pw.DateField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"user",
|
||||
'user',
|
||||
timestamp=pw.DateField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"auth",
|
||||
'auth',
|
||||
password=pw.CharField(max_length=255),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"chat",
|
||||
'chat',
|
||||
title=pw.CharField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"document",
|
||||
'document',
|
||||
title=pw.CharField(),
|
||||
filename=pw.CharField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"prompt",
|
||||
'prompt',
|
||||
title=pw.CharField(),
|
||||
)
|
||||
migrator.change_fields(
|
||||
"user",
|
||||
'user',
|
||||
profile_image_url=pw.CharField(),
|
||||
)
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -39,7 +38,7 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
|
||||
# Adding fields created_at and updated_at to the 'user' table
|
||||
migrator.add_fields(
|
||||
"user",
|
||||
'user',
|
||||
created_at=pw.BigIntegerField(null=True), # Allow null for transition
|
||||
updated_at=pw.BigIntegerField(null=True), # Allow null for transition
|
||||
last_active_at=pw.BigIntegerField(null=True), # Allow null for transition
|
||||
@@ -51,11 +50,11 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
)
|
||||
|
||||
# Now that the data has been copied, remove the original 'timestamp' field
|
||||
migrator.remove_fields("user", "timestamp")
|
||||
migrator.remove_fields('user', 'timestamp')
|
||||
|
||||
# Update the fields to be not null now that they are populated
|
||||
migrator.change_fields(
|
||||
"user",
|
||||
'user',
|
||||
created_at=pw.BigIntegerField(null=False),
|
||||
updated_at=pw.BigIntegerField(null=False),
|
||||
last_active_at=pw.BigIntegerField(null=False),
|
||||
@@ -66,14 +65,14 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
# Recreate the timestamp field initially allowing null values for safe transition
|
||||
migrator.add_fields("user", timestamp=pw.BigIntegerField(null=True))
|
||||
migrator.add_fields('user', timestamp=pw.BigIntegerField(null=True))
|
||||
|
||||
# Copy the earliest created_at date back into the new timestamp field
|
||||
# This assumes created_at was originally a copy of timestamp
|
||||
migrator.sql('UPDATE "user" SET timestamp = created_at')
|
||||
|
||||
# Remove the created_at and updated_at fields
|
||||
migrator.remove_fields("user", "created_at", "updated_at", "last_active_at")
|
||||
migrator.remove_fields('user', 'created_at', 'updated_at', 'last_active_at')
|
||||
|
||||
# Finally, alter the timestamp field to not allow nulls if that was the original setting
|
||||
migrator.change_fields("user", timestamp=pw.BigIntegerField(null=False))
|
||||
migrator.change_fields('user', timestamp=pw.BigIntegerField(null=False))
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -44,10 +43,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
created_at = pw.BigIntegerField(null=False)
|
||||
|
||||
class Meta:
|
||||
table_name = "memory"
|
||||
table_name = 'memory'
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_model("memory")
|
||||
migrator.remove_model('memory')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -52,10 +51,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
updated_at = pw.BigIntegerField(null=False)
|
||||
|
||||
class Meta:
|
||||
table_name = "model"
|
||||
table_name = 'model'
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_model("model")
|
||||
migrator.remove_model('model')
|
||||
|
||||
@@ -42,12 +42,12 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
# Fetch data from 'modelfile' table and insert into 'model' table
|
||||
migrate_modelfile_to_model(migrator, database)
|
||||
# Drop the 'modelfile' table
|
||||
migrator.remove_model("modelfile")
|
||||
migrator.remove_model('modelfile')
|
||||
|
||||
|
||||
def migrate_modelfile_to_model(migrator: Migrator, database: pw.Database):
|
||||
ModelFile = migrator.orm["modelfile"]
|
||||
Model = migrator.orm["model"]
|
||||
ModelFile = migrator.orm['modelfile']
|
||||
Model = migrator.orm['model']
|
||||
|
||||
modelfiles = ModelFile.select()
|
||||
|
||||
@@ -57,25 +57,25 @@ def migrate_modelfile_to_model(migrator: Migrator, database: pw.Database):
|
||||
modelfile.modelfile = json.loads(modelfile.modelfile)
|
||||
meta = json.dumps(
|
||||
{
|
||||
"description": modelfile.modelfile.get("desc"),
|
||||
"profile_image_url": modelfile.modelfile.get("imageUrl"),
|
||||
"ollama": {"modelfile": modelfile.modelfile.get("content")},
|
||||
"suggestion_prompts": modelfile.modelfile.get("suggestionPrompts"),
|
||||
"categories": modelfile.modelfile.get("categories"),
|
||||
"user": {**modelfile.modelfile.get("user", {}), "community": True},
|
||||
'description': modelfile.modelfile.get('desc'),
|
||||
'profile_image_url': modelfile.modelfile.get('imageUrl'),
|
||||
'ollama': {'modelfile': modelfile.modelfile.get('content')},
|
||||
'suggestion_prompts': modelfile.modelfile.get('suggestionPrompts'),
|
||||
'categories': modelfile.modelfile.get('categories'),
|
||||
'user': {**modelfile.modelfile.get('user', {}), 'community': True},
|
||||
}
|
||||
)
|
||||
|
||||
info = parse_ollama_modelfile(modelfile.modelfile.get("content"))
|
||||
info = parse_ollama_modelfile(modelfile.modelfile.get('content'))
|
||||
|
||||
# Insert the processed data into the 'model' table
|
||||
Model.create(
|
||||
id=f"ollama-{modelfile.tag_name}",
|
||||
id=f'ollama-{modelfile.tag_name}',
|
||||
user_id=modelfile.user_id,
|
||||
base_model_id=info.get("base_model_id"),
|
||||
name=modelfile.modelfile.get("title"),
|
||||
base_model_id=info.get('base_model_id'),
|
||||
name=modelfile.modelfile.get('title'),
|
||||
meta=meta,
|
||||
params=json.dumps(info.get("params", {})),
|
||||
params=json.dumps(info.get('params', {})),
|
||||
created_at=modelfile.timestamp,
|
||||
updated_at=modelfile.timestamp,
|
||||
)
|
||||
@@ -86,7 +86,7 @@ def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
|
||||
recreate_modelfile_table(migrator, database)
|
||||
move_data_back_to_modelfile(migrator, database)
|
||||
migrator.remove_model("model")
|
||||
migrator.remove_model('model')
|
||||
|
||||
|
||||
def recreate_modelfile_table(migrator: Migrator, database: pw.Database):
|
||||
@@ -102,8 +102,8 @@ def recreate_modelfile_table(migrator: Migrator, database: pw.Database):
|
||||
|
||||
|
||||
def move_data_back_to_modelfile(migrator: Migrator, database: pw.Database):
|
||||
Model = migrator.orm["model"]
|
||||
Modelfile = migrator.orm["modelfile"]
|
||||
Model = migrator.orm['model']
|
||||
Modelfile = migrator.orm['modelfile']
|
||||
|
||||
models = Model.select()
|
||||
|
||||
@@ -112,13 +112,13 @@ def move_data_back_to_modelfile(migrator: Migrator, database: pw.Database):
|
||||
meta = json.loads(model.meta)
|
||||
|
||||
modelfile_data = {
|
||||
"title": model.name,
|
||||
"desc": meta.get("description"),
|
||||
"imageUrl": meta.get("profile_image_url"),
|
||||
"content": meta.get("ollama", {}).get("modelfile"),
|
||||
"suggestionPrompts": meta.get("suggestion_prompts"),
|
||||
"categories": meta.get("categories"),
|
||||
"user": {k: v for k, v in meta.get("user", {}).items() if k != "community"},
|
||||
'title': model.name,
|
||||
'desc': meta.get('description'),
|
||||
'imageUrl': meta.get('profile_image_url'),
|
||||
'content': meta.get('ollama', {}).get('modelfile'),
|
||||
'suggestionPrompts': meta.get('suggestion_prompts'),
|
||||
'categories': meta.get('categories'),
|
||||
'user': {k: v for k, v in meta.get('user', {}).items() if k != 'community'},
|
||||
}
|
||||
|
||||
# Insert the processed data back into the 'modelfile' table
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -38,11 +37,11 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
# Adding fields settings to the 'user' table
|
||||
migrator.add_fields("user", settings=pw.TextField(null=True))
|
||||
migrator.add_fields('user', settings=pw.TextField(null=True))
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
# Remove the settings field
|
||||
migrator.remove_fields("user", "settings")
|
||||
migrator.remove_fields('user', 'settings')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -52,10 +51,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
updated_at = pw.BigIntegerField(null=False)
|
||||
|
||||
class Meta:
|
||||
table_name = "tool"
|
||||
table_name = 'tool'
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_model("tool")
|
||||
migrator.remove_model('tool')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -38,11 +37,11 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
# Adding fields info to the 'user' table
|
||||
migrator.add_fields("user", info=pw.TextField(null=True))
|
||||
migrator.add_fields('user', info=pw.TextField(null=True))
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
# Remove the settings field
|
||||
migrator.remove_fields("user", "info")
|
||||
migrator.remove_fields('user', 'info')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -46,10 +45,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
created_at = pw.BigIntegerField(null=False)
|
||||
|
||||
class Meta:
|
||||
table_name = "file"
|
||||
table_name = 'file'
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_model("file")
|
||||
migrator.remove_model('file')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -52,10 +51,10 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
updated_at = pw.BigIntegerField(null=False)
|
||||
|
||||
class Meta:
|
||||
table_name = "function"
|
||||
table_name = 'function'
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_model("function")
|
||||
migrator.remove_model('function')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -37,14 +36,14 @@ with suppress(ImportError):
|
||||
def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
migrator.add_fields("tool", valves=pw.TextField(null=True))
|
||||
migrator.add_fields("function", valves=pw.TextField(null=True))
|
||||
migrator.add_fields("function", is_active=pw.BooleanField(default=False))
|
||||
migrator.add_fields('tool', valves=pw.TextField(null=True))
|
||||
migrator.add_fields('function', valves=pw.TextField(null=True))
|
||||
migrator.add_fields('function', is_active=pw.BooleanField(default=False))
|
||||
|
||||
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_fields("tool", "valves")
|
||||
migrator.remove_fields("function", "valves")
|
||||
migrator.remove_fields("function", "is_active")
|
||||
migrator.remove_fields('tool', 'valves')
|
||||
migrator.remove_fields('function', 'valves')
|
||||
migrator.remove_fields('function', 'is_active')
|
||||
|
||||
@@ -25,7 +25,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -34,7 +33,7 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
migrator.add_fields(
|
||||
"user",
|
||||
'user',
|
||||
oauth_sub=pw.TextField(null=True, unique=True),
|
||||
)
|
||||
|
||||
@@ -42,4 +41,4 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_fields("user", "oauth_sub")
|
||||
migrator.remove_fields('user', 'oauth_sub')
|
||||
|
||||
@@ -29,7 +29,6 @@ from contextlib import suppress
|
||||
import peewee as pw
|
||||
from peewee_migrate import Migrator
|
||||
|
||||
|
||||
with suppress(ImportError):
|
||||
import playhouse.postgres_ext as pw_pext
|
||||
|
||||
@@ -38,7 +37,7 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your migrations here."""
|
||||
|
||||
migrator.add_fields(
|
||||
"function",
|
||||
'function',
|
||||
is_global=pw.BooleanField(default=False),
|
||||
)
|
||||
|
||||
@@ -46,4 +45,4 @@ def migrate(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
def rollback(migrator: Migrator, database: pw.Database, *, fake=False):
|
||||
"""Write your rollback migrations here."""
|
||||
|
||||
migrator.remove_fields("function", "is_global")
|
||||
migrator.remove_fields('function', 'is_global')
|
||||
|
||||
@@ -10,13 +10,13 @@ from playhouse.shortcuts import ReconnectMixin
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
db_state_default = {"closed": None, "conn": None, "ctx": None, "transactions": None}
|
||||
db_state = ContextVar("db_state", default=db_state_default.copy())
|
||||
db_state_default = {'closed': None, 'conn': None, 'ctx': None, 'transactions': None}
|
||||
db_state = ContextVar('db_state', default=db_state_default.copy())
|
||||
|
||||
|
||||
class PeeweeConnectionState(object):
|
||||
def __init__(self, **kwargs):
|
||||
super().__setattr__("_state", db_state)
|
||||
super().__setattr__('_state', db_state)
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def __setattr__(self, name, value):
|
||||
@@ -30,10 +30,10 @@ class PeeweeConnectionState(object):
|
||||
class CustomReconnectMixin(ReconnectMixin):
|
||||
reconnect_errors = (
|
||||
# psycopg2
|
||||
(OperationalError, "termin"),
|
||||
(InterfaceError, "closed"),
|
||||
(OperationalError, 'termin'),
|
||||
(InterfaceError, 'closed'),
|
||||
# peewee
|
||||
(PeeWeeInterfaceError, "closed"),
|
||||
(PeeWeeInterfaceError, 'closed'),
|
||||
)
|
||||
|
||||
|
||||
@@ -43,23 +43,21 @@ class ReconnectingPostgresqlDatabase(CustomReconnectMixin, PostgresqlDatabase):
|
||||
|
||||
def register_connection(db_url):
|
||||
# Check if using SQLCipher protocol
|
||||
if db_url.startswith("sqlite+sqlcipher://"):
|
||||
database_password = os.environ.get("DATABASE_PASSWORD")
|
||||
if not database_password or database_password.strip() == "":
|
||||
raise ValueError(
|
||||
"DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs"
|
||||
)
|
||||
if db_url.startswith('sqlite+sqlcipher://'):
|
||||
database_password = os.environ.get('DATABASE_PASSWORD')
|
||||
if not database_password or database_password.strip() == '':
|
||||
raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
|
||||
from playhouse.sqlcipher_ext import SqlCipherDatabase
|
||||
|
||||
# Parse the database path from SQLCipher URL
|
||||
# Convert sqlite+sqlcipher:///path/to/db.sqlite to /path/to/db.sqlite
|
||||
db_path = db_url.replace("sqlite+sqlcipher://", "")
|
||||
db_path = db_url.replace('sqlite+sqlcipher://', '')
|
||||
|
||||
# Use Peewee's native SqlCipherDatabase with encryption
|
||||
db = SqlCipherDatabase(db_path, passphrase=database_password)
|
||||
db.autoconnect = True
|
||||
db.reuse_if_open = True
|
||||
log.info("Connected to encrypted SQLite database using SQLCipher")
|
||||
log.info('Connected to encrypted SQLite database using SQLCipher')
|
||||
|
||||
else:
|
||||
# Standard database connection (existing logic)
|
||||
@@ -68,7 +66,7 @@ def register_connection(db_url):
|
||||
# Enable autoconnect for SQLite databases, managed by Peewee
|
||||
db.autoconnect = True
|
||||
db.reuse_if_open = True
|
||||
log.info("Connected to PostgreSQL database")
|
||||
log.info('Connected to PostgreSQL database')
|
||||
|
||||
# Get the connection details
|
||||
connection = parse(db_url, unquote_user=True, unquote_password=True)
|
||||
@@ -80,7 +78,7 @@ def register_connection(db_url):
|
||||
# Enable autoconnect for SQLite databases, managed by Peewee
|
||||
db.autoconnect = True
|
||||
db.reuse_if_open = True
|
||||
log.info("Connected to SQLite database")
|
||||
log.info('Connected to SQLite database')
|
||||
else:
|
||||
raise ValueError("Unsupported database connection")
|
||||
raise ValueError('Unsupported database connection')
|
||||
return db
|
||||
|
||||
+1063
-687
File diff suppressed because it is too large
Load Diff
@@ -1,8 +1,11 @@
|
||||
import logging
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from open_webui.models.auths import Auth
|
||||
from open_webui.env import DATABASE_URL, DATABASE_PASSWORD
|
||||
from open_webui.models.calendar import Calendar, CalendarEvent, CalendarEventAttendee # noqa: F401
|
||||
from open_webui.env import DATABASE_URL, DATABASE_PASSWORD, LOG_FORMAT
|
||||
from open_webui.internal.db import extract_ssl_params_from_url, reattach_ssl_params_to_url
|
||||
from sqlalchemy import engine_from_config, pool, create_engine
|
||||
|
||||
# this is the Alembic Config object, which provides
|
||||
@@ -14,6 +17,13 @@ config = context.config
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name, disable_existing_loggers=False)
|
||||
|
||||
# Re-apply JSON formatter after fileConfig replaces handlers.
|
||||
if LOG_FORMAT == 'json':
|
||||
from open_webui.env import JSONFormatter
|
||||
|
||||
for handler in logging.root.handlers:
|
||||
handler.setFormatter(JSONFormatter())
|
||||
|
||||
# add your model's MetaData object here
|
||||
# for 'autogenerate' support
|
||||
# from myapp import mymodel
|
||||
@@ -27,8 +37,12 @@ target_metadata = Auth.metadata
|
||||
|
||||
DB_URL = DATABASE_URL
|
||||
|
||||
# Normalize SSL query params for psycopg2 (Alembic uses psycopg2 for sync migrations).
|
||||
url_without_ssl, ssl_params = extract_ssl_params_from_url(DB_URL)
|
||||
DB_URL = reattach_ssl_params_to_url(url_without_ssl, ssl_params) if ssl_params else DB_URL
|
||||
|
||||
if DB_URL:
|
||||
config.set_main_option("sqlalchemy.url", DB_URL.replace("%", "%%"))
|
||||
config.set_main_option('sqlalchemy.url', DB_URL.replace('%', '%%'))
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
@@ -43,12 +57,12 @@ def run_migrations_offline() -> None:
|
||||
script output.
|
||||
|
||||
"""
|
||||
url = config.get_main_option("sqlalchemy.url")
|
||||
url = config.get_main_option('sqlalchemy.url')
|
||||
context.configure(
|
||||
url=url,
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
dialect_opts={'paramstyle': 'named'},
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
@@ -63,15 +77,13 @@ def run_migrations_online() -> None:
|
||||
|
||||
"""
|
||||
# Handle SQLCipher URLs
|
||||
if DB_URL and DB_URL.startswith("sqlite+sqlcipher://"):
|
||||
if not DATABASE_PASSWORD or DATABASE_PASSWORD.strip() == "":
|
||||
raise ValueError(
|
||||
"DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs"
|
||||
)
|
||||
if DB_URL and DB_URL.startswith('sqlite+sqlcipher://'):
|
||||
if not DATABASE_PASSWORD or DATABASE_PASSWORD.strip() == '':
|
||||
raise ValueError('DATABASE_PASSWORD is required when using sqlite+sqlcipher:// URLs')
|
||||
|
||||
# Extract database path from SQLCipher URL
|
||||
db_path = DB_URL.replace("sqlite+sqlcipher://", "")
|
||||
if db_path.startswith("/"):
|
||||
db_path = DB_URL.replace('sqlite+sqlcipher://', '')
|
||||
if db_path.startswith('/'):
|
||||
db_path = db_path[1:] # Remove leading slash for relative paths
|
||||
|
||||
# Create a custom creator function that uses sqlcipher3
|
||||
@@ -83,7 +95,7 @@ def run_migrations_online() -> None:
|
||||
return conn
|
||||
|
||||
connectable = create_engine(
|
||||
"sqlite://", # Dummy URL since we're using creator
|
||||
'sqlite://', # Dummy URL since we're using creator
|
||||
creator=create_sqlcipher_connection,
|
||||
echo=False,
|
||||
)
|
||||
@@ -91,7 +103,7 @@ def run_migrations_online() -> None:
|
||||
# Standard database connection (existing logic)
|
||||
connectable = engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
prefix='sqlalchemy.',
|
||||
poolclass=pool.NullPool,
|
||||
)
|
||||
|
||||
|
||||
@@ -12,4 +12,4 @@ def get_existing_tables():
|
||||
def get_revision_id():
|
||||
import uuid
|
||||
|
||||
return str(uuid.uuid4()).replace("-", "")[:12]
|
||||
return str(uuid.uuid4()).replace('-', '')[:12]
|
||||
|
||||
@@ -9,38 +9,38 @@ Create Date: 2025-08-13 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "018012973d35"
|
||||
down_revision = "d31026856c01"
|
||||
revision = '018012973d35'
|
||||
down_revision = 'd31026856c01'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
# Chat table indexes
|
||||
op.create_index("folder_id_idx", "chat", ["folder_id"])
|
||||
op.create_index("user_id_pinned_idx", "chat", ["user_id", "pinned"])
|
||||
op.create_index("user_id_archived_idx", "chat", ["user_id", "archived"])
|
||||
op.create_index("updated_at_user_id_idx", "chat", ["updated_at", "user_id"])
|
||||
op.create_index("folder_id_user_id_idx", "chat", ["folder_id", "user_id"])
|
||||
op.create_index('folder_id_idx', 'chat', ['folder_id'])
|
||||
op.create_index('user_id_pinned_idx', 'chat', ['user_id', 'pinned'])
|
||||
op.create_index('user_id_archived_idx', 'chat', ['user_id', 'archived'])
|
||||
op.create_index('updated_at_user_id_idx', 'chat', ['updated_at', 'user_id'])
|
||||
op.create_index('folder_id_user_id_idx', 'chat', ['folder_id', 'user_id'])
|
||||
|
||||
# Tag table index
|
||||
op.create_index("user_id_idx", "tag", ["user_id"])
|
||||
op.create_index('user_id_idx', 'tag', ['user_id'])
|
||||
|
||||
# Function table index
|
||||
op.create_index("is_global_idx", "function", ["is_global"])
|
||||
op.create_index('is_global_idx', 'function', ['is_global'])
|
||||
|
||||
|
||||
def downgrade():
|
||||
# Chat table indexes
|
||||
op.drop_index("folder_id_idx", table_name="chat")
|
||||
op.drop_index("user_id_pinned_idx", table_name="chat")
|
||||
op.drop_index("user_id_archived_idx", table_name="chat")
|
||||
op.drop_index("updated_at_user_id_idx", table_name="chat")
|
||||
op.drop_index("folder_id_user_id_idx", table_name="chat")
|
||||
op.drop_index('folder_id_idx', table_name='chat')
|
||||
op.drop_index('user_id_pinned_idx', table_name='chat')
|
||||
op.drop_index('user_id_archived_idx', table_name='chat')
|
||||
op.drop_index('updated_at_user_id_idx', table_name='chat')
|
||||
op.drop_index('folder_id_user_id_idx', table_name='chat')
|
||||
|
||||
# Tag table index
|
||||
op.drop_index("user_id_idx", table_name="tag")
|
||||
op.drop_index('user_id_idx', table_name='tag')
|
||||
|
||||
# Function table index
|
||||
|
||||
op.drop_index("is_global_idx", table_name="function")
|
||||
op.drop_index('is_global_idx', table_name='function')
|
||||
|
||||
@@ -13,8 +13,8 @@ from sqlalchemy.engine.reflection import Inspector
|
||||
|
||||
import json
|
||||
|
||||
revision = "1af9b942657b"
|
||||
down_revision = "242a2047eae0"
|
||||
revision = '1af9b942657b'
|
||||
down_revision = '242a2047eae0'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
@@ -25,43 +25,40 @@ def upgrade():
|
||||
inspector = Inspector.from_engine(conn)
|
||||
|
||||
# Clean up potential leftover temp table from previous failures
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS _alembic_tmp_tag"))
|
||||
conn.execute(sa.text('DROP TABLE IF EXISTS _alembic_tmp_tag'))
|
||||
|
||||
# Check if the 'tag' table exists
|
||||
tables = inspector.get_table_names()
|
||||
|
||||
# Step 1: Modify Tag table using batch mode for SQLite support
|
||||
if "tag" in tables:
|
||||
if 'tag' in tables:
|
||||
# Get the current columns in the 'tag' table
|
||||
columns = [col["name"] for col in inspector.get_columns("tag")]
|
||||
columns = [col['name'] for col in inspector.get_columns('tag')]
|
||||
|
||||
# Get any existing unique constraints on the 'tag' table
|
||||
current_constraints = inspector.get_unique_constraints("tag")
|
||||
current_constraints = inspector.get_unique_constraints('tag')
|
||||
|
||||
with op.batch_alter_table("tag", schema=None) as batch_op:
|
||||
with op.batch_alter_table('tag', schema=None) as batch_op:
|
||||
# Check if the unique constraint already exists
|
||||
if not any(
|
||||
constraint["name"] == "uq_id_user_id"
|
||||
for constraint in current_constraints
|
||||
):
|
||||
if not any(constraint['name'] == 'uq_id_user_id' for constraint in current_constraints):
|
||||
# Create unique constraint if it doesn't exist
|
||||
batch_op.create_unique_constraint("uq_id_user_id", ["id", "user_id"])
|
||||
batch_op.create_unique_constraint('uq_id_user_id', ['id', 'user_id'])
|
||||
|
||||
# Check if the 'data' column exists before trying to drop it
|
||||
if "data" in columns:
|
||||
batch_op.drop_column("data")
|
||||
if 'data' in columns:
|
||||
batch_op.drop_column('data')
|
||||
|
||||
# Check if the 'meta' column needs to be created
|
||||
if "meta" not in columns:
|
||||
if 'meta' not in columns:
|
||||
# Add the 'meta' column if it doesn't already exist
|
||||
batch_op.add_column(sa.Column("meta", sa.JSON(), nullable=True))
|
||||
batch_op.add_column(sa.Column('meta', sa.JSON(), nullable=True))
|
||||
|
||||
tag = table(
|
||||
"tag",
|
||||
column("id", sa.String()),
|
||||
column("name", sa.String()),
|
||||
column("user_id", sa.String()),
|
||||
column("meta", sa.JSON()),
|
||||
'tag',
|
||||
column('id', sa.String()),
|
||||
column('name', sa.String()),
|
||||
column('user_id', sa.String()),
|
||||
column('meta', sa.JSON()),
|
||||
)
|
||||
|
||||
# Step 2: Migrate tags
|
||||
@@ -70,12 +67,12 @@ def upgrade():
|
||||
|
||||
tag_updates = {}
|
||||
for row in result:
|
||||
new_id = row.name.replace(" ", "_").lower()
|
||||
new_id = row.name.replace(' ', '_').lower()
|
||||
tag_updates[row.id] = new_id
|
||||
|
||||
for tag_id, new_tag_id in tag_updates.items():
|
||||
print(f"Updating tag {tag_id} to {new_tag_id}")
|
||||
if new_tag_id == "pinned":
|
||||
print(f'Updating tag {tag_id} to {new_tag_id}')
|
||||
if new_tag_id == 'pinned':
|
||||
# delete tag
|
||||
delete_stmt = sa.delete(tag).where(tag.c.id == tag_id)
|
||||
conn.execute(delete_stmt)
|
||||
@@ -86,9 +83,7 @@ def upgrade():
|
||||
|
||||
if existing_tag_result:
|
||||
# Handle duplicate case: the new_tag_id already exists
|
||||
print(
|
||||
f"Tag {new_tag_id} already exists. Removing current tag with ID {tag_id} to avoid duplicates."
|
||||
)
|
||||
print(f'Tag {new_tag_id} already exists. Removing current tag with ID {tag_id} to avoid duplicates.')
|
||||
# Option 1: Delete the current tag if an update to new_tag_id would cause duplication
|
||||
delete_stmt = sa.delete(tag).where(tag.c.id == tag_id)
|
||||
conn.execute(delete_stmt)
|
||||
@@ -98,19 +93,15 @@ def upgrade():
|
||||
conn.execute(update_stmt)
|
||||
|
||||
# Add columns `pinned` and `meta` to 'chat'
|
||||
op.add_column("chat", sa.Column("pinned", sa.Boolean(), nullable=True))
|
||||
op.add_column(
|
||||
"chat", sa.Column("meta", sa.JSON(), nullable=False, server_default="{}")
|
||||
)
|
||||
op.add_column('chat', sa.Column('pinned', sa.Boolean(), nullable=True))
|
||||
op.add_column('chat', sa.Column('meta', sa.JSON(), nullable=False, server_default='{}'))
|
||||
|
||||
chatidtag = table(
|
||||
"chatidtag", column("chat_id", sa.String()), column("tag_name", sa.String())
|
||||
)
|
||||
chatidtag = table('chatidtag', column('chat_id', sa.String()), column('tag_name', sa.String()))
|
||||
chat = table(
|
||||
"chat",
|
||||
column("id", sa.String()),
|
||||
column("pinned", sa.Boolean()),
|
||||
column("meta", sa.JSON()),
|
||||
'chat',
|
||||
column('id', sa.String()),
|
||||
column('pinned', sa.Boolean()),
|
||||
column('meta', sa.JSON()),
|
||||
)
|
||||
|
||||
# Fetch existing tags
|
||||
@@ -120,29 +111,27 @@ def upgrade():
|
||||
chat_updates = {}
|
||||
for row in result:
|
||||
chat_id = row.chat_id
|
||||
tag_name = row.tag_name.replace(" ", "_").lower()
|
||||
tag_name = row.tag_name.replace(' ', '_').lower()
|
||||
|
||||
if tag_name == "pinned":
|
||||
if tag_name == 'pinned':
|
||||
# Specifically handle 'pinned' tag
|
||||
if chat_id not in chat_updates:
|
||||
chat_updates[chat_id] = {"pinned": True, "meta": {}}
|
||||
chat_updates[chat_id] = {'pinned': True, 'meta': {}}
|
||||
else:
|
||||
chat_updates[chat_id]["pinned"] = True
|
||||
chat_updates[chat_id]['pinned'] = True
|
||||
else:
|
||||
if chat_id not in chat_updates:
|
||||
chat_updates[chat_id] = {"pinned": False, "meta": {"tags": [tag_name]}}
|
||||
chat_updates[chat_id] = {'pinned': False, 'meta': {'tags': [tag_name]}}
|
||||
else:
|
||||
tags = chat_updates[chat_id]["meta"].get("tags", [])
|
||||
tags = chat_updates[chat_id]['meta'].get('tags', [])
|
||||
tags.append(tag_name)
|
||||
|
||||
chat_updates[chat_id]["meta"]["tags"] = list(set(tags))
|
||||
chat_updates[chat_id]['meta']['tags'] = list(set(tags))
|
||||
|
||||
# Update chats based on accumulated changes
|
||||
for chat_id, updates in chat_updates.items():
|
||||
update_stmt = sa.update(chat).where(chat.c.id == chat_id)
|
||||
update_stmt = update_stmt.values(
|
||||
meta=updates.get("meta", {}), pinned=updates.get("pinned", False)
|
||||
)
|
||||
update_stmt = update_stmt.values(meta=updates.get('meta', {}), pinned=updates.get('pinned', False))
|
||||
conn.execute(update_stmt)
|
||||
pass
|
||||
|
||||
|
||||
@@ -12,8 +12,8 @@ from sqlalchemy.sql import table, select, update
|
||||
|
||||
import json
|
||||
|
||||
revision = "242a2047eae0"
|
||||
down_revision = "6a39f3d8e55c"
|
||||
revision = '242a2047eae0'
|
||||
down_revision = '6a39f3d8e55c'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
@@ -22,39 +22,37 @@ def upgrade():
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
|
||||
columns = inspector.get_columns("chat")
|
||||
column_dict = {col["name"]: col for col in columns}
|
||||
columns = inspector.get_columns('chat')
|
||||
column_dict = {col['name']: col for col in columns}
|
||||
|
||||
chat_column = column_dict.get("chat")
|
||||
old_chat_exists = "old_chat" in column_dict
|
||||
chat_column = column_dict.get('chat')
|
||||
old_chat_exists = 'old_chat' in column_dict
|
||||
|
||||
if chat_column:
|
||||
if isinstance(chat_column["type"], sa.Text):
|
||||
if isinstance(chat_column['type'], sa.Text):
|
||||
print("Converting 'chat' column to JSON")
|
||||
|
||||
if old_chat_exists:
|
||||
print("Dropping old 'old_chat' column")
|
||||
op.drop_column("chat", "old_chat")
|
||||
op.drop_column('chat', 'old_chat')
|
||||
|
||||
# Step 1: Rename current 'chat' column to 'old_chat'
|
||||
print("Renaming 'chat' column to 'old_chat'")
|
||||
op.alter_column(
|
||||
"chat", "chat", new_column_name="old_chat", existing_type=sa.Text()
|
||||
)
|
||||
op.alter_column('chat', 'chat', new_column_name='old_chat', existing_type=sa.Text())
|
||||
|
||||
# Step 2: Add new 'chat' column of type JSON
|
||||
print("Adding new 'chat' column of type JSON")
|
||||
op.add_column("chat", sa.Column("chat", sa.JSON(), nullable=True))
|
||||
op.add_column('chat', sa.Column('chat', sa.JSON(), nullable=True))
|
||||
else:
|
||||
# If the column is already JSON, no need to do anything
|
||||
pass
|
||||
|
||||
# Step 3: Migrate data from 'old_chat' to 'chat'
|
||||
chat_table = table(
|
||||
"chat",
|
||||
sa.Column("id", sa.String(), primary_key=True),
|
||||
sa.Column("old_chat", sa.Text()),
|
||||
sa.Column("chat", sa.JSON()),
|
||||
'chat',
|
||||
sa.Column('id', sa.String(), primary_key=True),
|
||||
sa.Column('old_chat', sa.Text()),
|
||||
sa.Column('chat', sa.JSON()),
|
||||
)
|
||||
|
||||
# - Selecting all data from the table
|
||||
@@ -67,41 +65,33 @@ def upgrade():
|
||||
except json.JSONDecodeError:
|
||||
json_data = None # Handle cases where the text cannot be converted to JSON
|
||||
|
||||
connection.execute(
|
||||
sa.update(chat_table)
|
||||
.where(chat_table.c.id == row.id)
|
||||
.values(chat=json_data)
|
||||
)
|
||||
connection.execute(sa.update(chat_table).where(chat_table.c.id == row.id).values(chat=json_data))
|
||||
|
||||
# Step 4: Drop 'old_chat' column
|
||||
print("Dropping 'old_chat' column")
|
||||
op.drop_column("chat", "old_chat")
|
||||
op.drop_column('chat', 'old_chat')
|
||||
|
||||
|
||||
def downgrade():
|
||||
# Step 1: Add 'old_chat' column back as Text
|
||||
op.add_column("chat", sa.Column("old_chat", sa.Text(), nullable=True))
|
||||
op.add_column('chat', sa.Column('old_chat', sa.Text(), nullable=True))
|
||||
|
||||
# Step 2: Convert 'chat' JSON data back to text and store in 'old_chat'
|
||||
chat_table = table(
|
||||
"chat",
|
||||
sa.Column("id", sa.String(), primary_key=True),
|
||||
sa.Column("chat", sa.JSON()),
|
||||
sa.Column("old_chat", sa.Text()),
|
||||
'chat',
|
||||
sa.Column('id', sa.String(), primary_key=True),
|
||||
sa.Column('chat', sa.JSON()),
|
||||
sa.Column('old_chat', sa.Text()),
|
||||
)
|
||||
|
||||
connection = op.get_bind()
|
||||
results = connection.execute(select(chat_table.c.id, chat_table.c.chat))
|
||||
for row in results:
|
||||
text_data = json.dumps(row.chat) if row.chat is not None else None
|
||||
connection.execute(
|
||||
sa.update(chat_table)
|
||||
.where(chat_table.c.id == row.id)
|
||||
.values(old_chat=text_data)
|
||||
)
|
||||
connection.execute(sa.update(chat_table).where(chat_table.c.id == row.id).values(old_chat=text_data))
|
||||
|
||||
# Step 3: Remove the new 'chat' JSON column
|
||||
op.drop_column("chat", "chat")
|
||||
op.drop_column('chat', 'chat')
|
||||
|
||||
# Step 4: Rename 'old_chat' back to 'chat'
|
||||
op.alter_column("chat", "old_chat", new_column_name="chat", existing_type=sa.Text())
|
||||
op.alter_column('chat', 'old_chat', new_column_name='chat', existing_type=sa.Text())
|
||||
|
||||
+28
-37
@@ -12,21 +12,20 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
import open_webui.internal.db
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "2f1211949ecc"
|
||||
down_revision: Union[str, None] = "37f288994c47"
|
||||
revision: str = '2f1211949ecc'
|
||||
down_revision: Union[str, None] = '37f288994c47'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# New columns to be added to channel_member table
|
||||
op.add_column("channel_member", sa.Column("status", sa.Text(), nullable=True))
|
||||
op.add_column('channel_member', sa.Column('status', sa.Text(), nullable=True))
|
||||
op.add_column(
|
||||
"channel_member",
|
||||
'channel_member',
|
||||
sa.Column(
|
||||
"is_active",
|
||||
'is_active',
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
default=True,
|
||||
@@ -35,9 +34,9 @@ def upgrade() -> None:
|
||||
)
|
||||
|
||||
op.add_column(
|
||||
"channel_member",
|
||||
'channel_member',
|
||||
sa.Column(
|
||||
"is_channel_muted",
|
||||
'is_channel_muted',
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
default=False,
|
||||
@@ -45,9 +44,9 @@ def upgrade() -> None:
|
||||
),
|
||||
)
|
||||
op.add_column(
|
||||
"channel_member",
|
||||
'channel_member',
|
||||
sa.Column(
|
||||
"is_channel_pinned",
|
||||
'is_channel_pinned',
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
default=False,
|
||||
@@ -55,49 +54,41 @@ def upgrade() -> None:
|
||||
),
|
||||
)
|
||||
|
||||
op.add_column("channel_member", sa.Column("data", sa.JSON(), nullable=True))
|
||||
op.add_column("channel_member", sa.Column("meta", sa.JSON(), nullable=True))
|
||||
op.add_column('channel_member', sa.Column('data', sa.JSON(), nullable=True))
|
||||
op.add_column('channel_member', sa.Column('meta', sa.JSON(), nullable=True))
|
||||
|
||||
op.add_column(
|
||||
"channel_member", sa.Column("joined_at", sa.BigInteger(), nullable=False)
|
||||
)
|
||||
op.add_column(
|
||||
"channel_member", sa.Column("left_at", sa.BigInteger(), nullable=True)
|
||||
)
|
||||
op.add_column('channel_member', sa.Column('joined_at', sa.BigInteger(), nullable=False))
|
||||
op.add_column('channel_member', sa.Column('left_at', sa.BigInteger(), nullable=True))
|
||||
|
||||
op.add_column(
|
||||
"channel_member", sa.Column("last_read_at", sa.BigInteger(), nullable=True)
|
||||
)
|
||||
op.add_column('channel_member', sa.Column('last_read_at', sa.BigInteger(), nullable=True))
|
||||
|
||||
op.add_column(
|
||||
"channel_member", sa.Column("updated_at", sa.BigInteger(), nullable=True)
|
||||
)
|
||||
op.add_column('channel_member', sa.Column('updated_at', sa.BigInteger(), nullable=True))
|
||||
|
||||
# New columns to be added to message table
|
||||
op.add_column(
|
||||
"message",
|
||||
'message',
|
||||
sa.Column(
|
||||
"is_pinned",
|
||||
'is_pinned',
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
default=False,
|
||||
server_default=sa.sql.expression.false(),
|
||||
),
|
||||
)
|
||||
op.add_column("message", sa.Column("pinned_at", sa.BigInteger(), nullable=True))
|
||||
op.add_column("message", sa.Column("pinned_by", sa.Text(), nullable=True))
|
||||
op.add_column('message', sa.Column('pinned_at', sa.BigInteger(), nullable=True))
|
||||
op.add_column('message', sa.Column('pinned_by', sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("channel_member", "updated_at")
|
||||
op.drop_column("channel_member", "last_read_at")
|
||||
op.drop_column('channel_member', 'updated_at')
|
||||
op.drop_column('channel_member', 'last_read_at')
|
||||
|
||||
op.drop_column("channel_member", "meta")
|
||||
op.drop_column("channel_member", "data")
|
||||
op.drop_column('channel_member', 'meta')
|
||||
op.drop_column('channel_member', 'data')
|
||||
|
||||
op.drop_column("channel_member", "is_channel_pinned")
|
||||
op.drop_column("channel_member", "is_channel_muted")
|
||||
op.drop_column('channel_member', 'is_channel_pinned')
|
||||
op.drop_column('channel_member', 'is_channel_muted')
|
||||
|
||||
op.drop_column("message", "pinned_by")
|
||||
op.drop_column("message", "pinned_at")
|
||||
op.drop_column("message", "is_pinned")
|
||||
op.drop_column('message', 'pinned_by')
|
||||
op.drop_column('message', 'pinned_at')
|
||||
op.drop_column('message', 'is_pinned')
|
||||
|
||||
@@ -0,0 +1,245 @@
|
||||
"""Add prompt history table
|
||||
|
||||
Revision ID: 374d2f66af06
|
||||
Revises: c440947495f3
|
||||
Create Date: 2026-01-23 17:15:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
import uuid
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision: str = '374d2f66af06'
|
||||
down_revision: Union[str, None] = 'c440947495f3'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# Step 1: Read existing data from OLD table (schema likely command as PK)
|
||||
# We use batch_alter previously, but we want to move to new table.
|
||||
# We need to assume the OLD structure.
|
||||
|
||||
old_prompt_table = sa.table(
|
||||
'prompt',
|
||||
sa.column('command', sa.Text()),
|
||||
sa.column('user_id', sa.Text()),
|
||||
sa.column('title', sa.Text()),
|
||||
sa.column('content', sa.Text()),
|
||||
sa.column('timestamp', sa.BigInteger()),
|
||||
sa.column('access_control', sa.JSON()),
|
||||
)
|
||||
|
||||
# Check if table exists/read data
|
||||
try:
|
||||
existing_prompts = conn.execute(
|
||||
sa.select(
|
||||
old_prompt_table.c.command,
|
||||
old_prompt_table.c.user_id,
|
||||
old_prompt_table.c.title,
|
||||
old_prompt_table.c.content,
|
||||
old_prompt_table.c.timestamp,
|
||||
old_prompt_table.c.access_control,
|
||||
)
|
||||
).fetchall()
|
||||
except Exception:
|
||||
# Fallback if table doesn't exist (new install)
|
||||
existing_prompts = []
|
||||
|
||||
# Step 2: Create new prompt table with 'id' as PRIMARY KEY
|
||||
op.create_table(
|
||||
'prompt_new',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('command', sa.String(), unique=True, index=True),
|
||||
sa.Column('user_id', sa.String(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False),
|
||||
sa.Column('content', sa.Text(), nullable=False),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=False, server_default='1'),
|
||||
sa.Column('version_id', sa.Text(), nullable=True),
|
||||
sa.Column('tags', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
|
||||
# Step 3: Create prompt_history table
|
||||
op.create_table(
|
||||
'prompt_history',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('prompt_id', sa.Text(), nullable=False, index=True),
|
||||
sa.Column('parent_id', sa.Text(), nullable=True),
|
||||
sa.Column('snapshot', sa.JSON(), nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('commit_message', sa.Text(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
|
||||
# Step 4: Migrate data
|
||||
prompt_new_table = sa.table(
|
||||
'prompt_new',
|
||||
sa.column('id', sa.Text()),
|
||||
sa.column('command', sa.String()),
|
||||
sa.column('user_id', sa.String()),
|
||||
sa.column('name', sa.Text()),
|
||||
sa.column('content', sa.Text()),
|
||||
sa.column('data', sa.JSON()),
|
||||
sa.column('meta', sa.JSON()),
|
||||
sa.column('access_control', sa.JSON()),
|
||||
sa.column('is_active', sa.Boolean()),
|
||||
sa.column('version_id', sa.Text()),
|
||||
sa.column('tags', sa.JSON()),
|
||||
sa.column('created_at', sa.BigInteger()),
|
||||
sa.column('updated_at', sa.BigInteger()),
|
||||
)
|
||||
|
||||
prompt_history_table = sa.table(
|
||||
'prompt_history',
|
||||
sa.column('id', sa.Text()),
|
||||
sa.column('prompt_id', sa.Text()),
|
||||
sa.column('parent_id', sa.Text()),
|
||||
sa.column('snapshot', sa.JSON()),
|
||||
sa.column('user_id', sa.Text()),
|
||||
sa.column('commit_message', sa.Text()),
|
||||
sa.column('created_at', sa.BigInteger()),
|
||||
)
|
||||
|
||||
for row in existing_prompts:
|
||||
command = row[0]
|
||||
user_id = row[1]
|
||||
title = row[2]
|
||||
content = row[3]
|
||||
timestamp = row[4]
|
||||
access_control = row[5]
|
||||
|
||||
new_uuid = str(uuid.uuid4())
|
||||
history_uuid = str(uuid.uuid4())
|
||||
clean_command = command[1:] if command and command.startswith('/') else command
|
||||
|
||||
# Insert into prompt_new
|
||||
conn.execute(
|
||||
sa.insert(prompt_new_table).values(
|
||||
id=new_uuid,
|
||||
command=clean_command,
|
||||
user_id=user_id,
|
||||
name=title,
|
||||
content=content,
|
||||
data={},
|
||||
meta={},
|
||||
access_control=access_control,
|
||||
is_active=True,
|
||||
version_id=history_uuid,
|
||||
tags=[],
|
||||
created_at=timestamp,
|
||||
updated_at=timestamp,
|
||||
)
|
||||
)
|
||||
|
||||
# Create initial history entry
|
||||
conn.execute(
|
||||
sa.insert(prompt_history_table).values(
|
||||
id=history_uuid,
|
||||
prompt_id=new_uuid,
|
||||
parent_id=None,
|
||||
snapshot={
|
||||
'name': title,
|
||||
'content': content,
|
||||
'command': clean_command,
|
||||
'data': {},
|
||||
'meta': {},
|
||||
'access_control': access_control,
|
||||
},
|
||||
user_id=user_id,
|
||||
commit_message=None,
|
||||
created_at=timestamp,
|
||||
)
|
||||
)
|
||||
|
||||
# Step 5: Replace old table with new one
|
||||
op.drop_table('prompt')
|
||||
op.rename_table('prompt_new', 'prompt')
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# Step 1: Read new data
|
||||
prompt_table = sa.table(
|
||||
'prompt',
|
||||
sa.column('command', sa.String()),
|
||||
sa.column('name', sa.Text()),
|
||||
sa.column('created_at', sa.BigInteger()),
|
||||
sa.column('user_id', sa.Text()),
|
||||
sa.column('content', sa.Text()),
|
||||
sa.column('access_control', sa.JSON()),
|
||||
)
|
||||
|
||||
try:
|
||||
current_data = conn.execute(
|
||||
sa.select(
|
||||
prompt_table.c.command,
|
||||
prompt_table.c.name,
|
||||
prompt_table.c.created_at,
|
||||
prompt_table.c.user_id,
|
||||
prompt_table.c.content,
|
||||
prompt_table.c.access_control,
|
||||
)
|
||||
).fetchall()
|
||||
except Exception:
|
||||
current_data = []
|
||||
|
||||
# Step 2: Drop history and table
|
||||
op.drop_table('prompt_history')
|
||||
op.drop_table('prompt')
|
||||
|
||||
# Step 3: Recreate old table (command as PK?)
|
||||
# Assuming old schema:
|
||||
op.create_table(
|
||||
'prompt',
|
||||
sa.Column('command', sa.String(), primary_key=True),
|
||||
sa.Column('user_id', sa.String()),
|
||||
sa.Column('title', sa.Text()),
|
||||
sa.Column('content', sa.Text()),
|
||||
sa.Column('timestamp', sa.BigInteger()),
|
||||
sa.Column('access_control', sa.JSON()),
|
||||
sa.Column('id', sa.Integer(), nullable=True),
|
||||
)
|
||||
|
||||
# Step 4: Restore data
|
||||
old_prompt_table = sa.table(
|
||||
'prompt',
|
||||
sa.column('command', sa.String()),
|
||||
sa.column('user_id', sa.String()),
|
||||
sa.column('title', sa.Text()),
|
||||
sa.column('content', sa.Text()),
|
||||
sa.column('timestamp', sa.BigInteger()),
|
||||
sa.column('access_control', sa.JSON()),
|
||||
)
|
||||
|
||||
for row in current_data:
|
||||
command = row[0]
|
||||
name = row[1]
|
||||
created_at = row[2]
|
||||
user_id = row[3]
|
||||
content = row[4]
|
||||
access_control = row[5]
|
||||
|
||||
# Restore leading /
|
||||
old_command = '/' + command if command and not command.startswith('/') else command
|
||||
|
||||
conn.execute(
|
||||
sa.insert(old_prompt_table).values(
|
||||
command=old_command,
|
||||
user_id=user_id,
|
||||
title=name,
|
||||
content=content,
|
||||
timestamp=created_at,
|
||||
access_control=access_control,
|
||||
)
|
||||
)
|
||||
@@ -9,8 +9,8 @@ Create Date: 2024-12-30 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "3781e22d8b01"
|
||||
down_revision = "7826ab40b532"
|
||||
revision = '3781e22d8b01'
|
||||
down_revision = '7826ab40b532'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
@@ -18,9 +18,9 @@ depends_on = None
|
||||
def upgrade():
|
||||
# Add 'type' column to the 'channel' table
|
||||
op.add_column(
|
||||
"channel",
|
||||
'channel',
|
||||
sa.Column(
|
||||
"type",
|
||||
'type',
|
||||
sa.Text(),
|
||||
nullable=True,
|
||||
),
|
||||
@@ -28,43 +28,31 @@ def upgrade():
|
||||
|
||||
# Add 'parent_id' column to the 'message' table for threads
|
||||
op.add_column(
|
||||
"message",
|
||||
sa.Column("parent_id", sa.Text(), nullable=True),
|
||||
'message',
|
||||
sa.Column('parent_id', sa.Text(), nullable=True),
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"message_reaction",
|
||||
sa.Column(
|
||||
"id", sa.Text(), nullable=False, primary_key=True, unique=True
|
||||
), # Unique reaction ID
|
||||
sa.Column("user_id", sa.Text(), nullable=False), # User who reacted
|
||||
sa.Column(
|
||||
"message_id", sa.Text(), nullable=False
|
||||
), # Message that was reacted to
|
||||
sa.Column(
|
||||
"name", sa.Text(), nullable=False
|
||||
), # Reaction name (e.g. "thumbs_up")
|
||||
sa.Column(
|
||||
"created_at", sa.BigInteger(), nullable=True
|
||||
), # Timestamp of when the reaction was added
|
||||
'message_reaction',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True), # Unique reaction ID
|
||||
sa.Column('user_id', sa.Text(), nullable=False), # User who reacted
|
||||
sa.Column('message_id', sa.Text(), nullable=False), # Message that was reacted to
|
||||
sa.Column('name', sa.Text(), nullable=False), # Reaction name (e.g. "thumbs_up")
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True), # Timestamp of when the reaction was added
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"channel_member",
|
||||
sa.Column(
|
||||
"id", sa.Text(), nullable=False, primary_key=True, unique=True
|
||||
), # Record ID for the membership row
|
||||
sa.Column("channel_id", sa.Text(), nullable=False), # Associated channel
|
||||
sa.Column("user_id", sa.Text(), nullable=False), # Associated user
|
||||
sa.Column(
|
||||
"created_at", sa.BigInteger(), nullable=True
|
||||
), # Timestamp of when the user joined the channel
|
||||
'channel_member',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True), # Record ID for the membership row
|
||||
sa.Column('channel_id', sa.Text(), nullable=False), # Associated channel
|
||||
sa.Column('user_id', sa.Text(), nullable=False), # Associated user
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True), # Timestamp of when the user joined the channel
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
# Revert 'type' column addition to the 'channel' table
|
||||
op.drop_column("channel", "type")
|
||||
op.drop_column("message", "parent_id")
|
||||
op.drop_table("message_reaction")
|
||||
op.drop_table("channel_member")
|
||||
op.drop_column('channel', 'type')
|
||||
op.drop_column('message', 'parent_id')
|
||||
op.drop_table('message_reaction')
|
||||
op.drop_table('channel_member')
|
||||
|
||||
@@ -14,10 +14,9 @@ from typing import Sequence, Union
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "37f288994c47"
|
||||
down_revision: Union[str, None] = "a5c220713937"
|
||||
revision: str = '37f288994c47'
|
||||
down_revision: Union[str, None] = 'a5c220713937'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
@@ -25,50 +24,48 @@ depends_on: Union[str, Sequence[str], None] = None
|
||||
def upgrade() -> None:
|
||||
# 1. Create new table
|
||||
op.create_table(
|
||||
"group_member",
|
||||
sa.Column("id", sa.Text(), primary_key=True, unique=True, nullable=False),
|
||||
'group_member',
|
||||
sa.Column('id', sa.Text(), primary_key=True, unique=True, nullable=False),
|
||||
sa.Column(
|
||||
"group_id",
|
||||
'group_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("group.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('group.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"user_id",
|
||||
'user_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("user.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('user.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.UniqueConstraint("group_id", "user_id", name="uq_group_member_group_user"),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.UniqueConstraint('group_id', 'user_id', name='uq_group_member_group_user'),
|
||||
)
|
||||
|
||||
connection = op.get_bind()
|
||||
|
||||
# 2. Read existing group with user_ids JSON column
|
||||
group_table = sa.Table(
|
||||
"group",
|
||||
'group',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("user_ids", sa.JSON()), # JSON stored as text in SQLite + PG
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('user_ids', sa.JSON()), # JSON stored as text in SQLite + PG
|
||||
)
|
||||
|
||||
results = connection.execute(
|
||||
sa.select(group_table.c.id, group_table.c.user_ids)
|
||||
).fetchall()
|
||||
results = connection.execute(sa.select(group_table.c.id, group_table.c.user_ids)).fetchall()
|
||||
|
||||
print(results)
|
||||
|
||||
# 3. Insert members into group_member table
|
||||
gm_table = sa.Table(
|
||||
"group_member",
|
||||
'group_member',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("group_id", sa.Text()),
|
||||
sa.Column("user_id", sa.Text()),
|
||||
sa.Column("created_at", sa.BigInteger()),
|
||||
sa.Column("updated_at", sa.BigInteger()),
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('group_id', sa.Text()),
|
||||
sa.Column('user_id', sa.Text()),
|
||||
sa.Column('created_at', sa.BigInteger()),
|
||||
sa.Column('updated_at', sa.BigInteger()),
|
||||
)
|
||||
|
||||
now = int(time.time())
|
||||
@@ -87,11 +84,11 @@ def upgrade() -> None:
|
||||
|
||||
rows = [
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"group_id": group_id,
|
||||
"user_id": uid,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
'id': str(uuid.uuid4()),
|
||||
'group_id': group_id,
|
||||
'user_id': uid,
|
||||
'created_at': now,
|
||||
'updated_at': now,
|
||||
}
|
||||
for uid in user_ids
|
||||
]
|
||||
@@ -100,47 +97,41 @@ def upgrade() -> None:
|
||||
connection.execute(gm_table.insert(), rows)
|
||||
|
||||
# 4. Optionally drop the old column
|
||||
with op.batch_alter_table("group") as batch:
|
||||
batch.drop_column("user_ids")
|
||||
with op.batch_alter_table('group') as batch:
|
||||
batch.drop_column('user_ids')
|
||||
|
||||
|
||||
def downgrade():
|
||||
# Reverse: restore user_ids column
|
||||
with op.batch_alter_table("group") as batch:
|
||||
batch.add_column(sa.Column("user_ids", sa.JSON()))
|
||||
with op.batch_alter_table('group') as batch:
|
||||
batch.add_column(sa.Column('user_ids', sa.JSON()))
|
||||
|
||||
connection = op.get_bind()
|
||||
gm_table = sa.Table(
|
||||
"group_member",
|
||||
'group_member',
|
||||
sa.MetaData(),
|
||||
sa.Column("group_id", sa.Text()),
|
||||
sa.Column("user_id", sa.Text()),
|
||||
sa.Column("created_at", sa.BigInteger()),
|
||||
sa.Column("updated_at", sa.BigInteger()),
|
||||
sa.Column('group_id', sa.Text()),
|
||||
sa.Column('user_id', sa.Text()),
|
||||
sa.Column('created_at', sa.BigInteger()),
|
||||
sa.Column('updated_at', sa.BigInteger()),
|
||||
)
|
||||
|
||||
group_table = sa.Table(
|
||||
"group",
|
||||
'group',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("user_ids", sa.JSON()),
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('user_ids', sa.JSON()),
|
||||
)
|
||||
|
||||
# Build JSON arrays again
|
||||
results = connection.execute(sa.select(group_table.c.id)).fetchall()
|
||||
|
||||
for (group_id,) in results:
|
||||
members = connection.execute(
|
||||
sa.select(gm_table.c.user_id).where(gm_table.c.group_id == group_id)
|
||||
).fetchall()
|
||||
members = connection.execute(sa.select(gm_table.c.user_id).where(gm_table.c.group_id == group_id)).fetchall()
|
||||
|
||||
member_ids = [m[0] for m in members]
|
||||
|
||||
connection.execute(
|
||||
group_table.update()
|
||||
.where(group_table.c.id == group_id)
|
||||
.values(user_ids=member_ids)
|
||||
)
|
||||
connection.execute(group_table.update().where(group_table.c.id == group_id).values(user_ids=member_ids))
|
||||
|
||||
# Drop the new table
|
||||
op.drop_table("group_member")
|
||||
op.drop_table('group_member')
|
||||
|
||||
@@ -11,10 +11,9 @@ from typing import Sequence, Union
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "38d63c18f30f"
|
||||
down_revision: Union[str, None] = "3af16a1c9fb6"
|
||||
revision: str = '38d63c18f30f'
|
||||
down_revision: Union[str, None] = '3af16a1c9fb6'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
@@ -22,59 +21,55 @@ depends_on: Union[str, Sequence[str], None] = None
|
||||
def upgrade() -> None:
|
||||
# Ensure 'id' column in 'user' table is unique and primary key (ForeignKey constraint)
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
columns = inspector.get_columns("user")
|
||||
columns = inspector.get_columns('user')
|
||||
|
||||
pk_columns = inspector.get_pk_constraint("user")["constrained_columns"]
|
||||
id_column = next((col for col in columns if col["name"] == "id"), None)
|
||||
pk_columns = inspector.get_pk_constraint('user')['constrained_columns']
|
||||
id_column = next((col for col in columns if col['name'] == 'id'), None)
|
||||
|
||||
if id_column and not id_column.get("unique", False):
|
||||
unique_constraints = inspector.get_unique_constraints("user")
|
||||
unique_columns = {tuple(u["column_names"]) for u in unique_constraints}
|
||||
if id_column and not id_column.get('unique', False):
|
||||
unique_constraints = inspector.get_unique_constraints('user')
|
||||
unique_columns = {tuple(u['column_names']) for u in unique_constraints}
|
||||
|
||||
with op.batch_alter_table("user") as batch_op:
|
||||
with op.batch_alter_table('user') as batch_op:
|
||||
# If primary key is wrong, drop it
|
||||
if pk_columns and pk_columns != ["id"]:
|
||||
batch_op.drop_constraint(
|
||||
inspector.get_pk_constraint("user")["name"], type_="primary"
|
||||
)
|
||||
if pk_columns and pk_columns != ['id']:
|
||||
batch_op.drop_constraint(inspector.get_pk_constraint('user')['name'], type_='primary')
|
||||
|
||||
# Add unique constraint if missing
|
||||
if ("id",) not in unique_columns:
|
||||
batch_op.create_unique_constraint("uq_user_id", ["id"])
|
||||
if ('id',) not in unique_columns:
|
||||
batch_op.create_unique_constraint('uq_user_id', ['id'])
|
||||
|
||||
# Re-create correct primary key
|
||||
batch_op.create_primary_key("pk_user_id", ["id"])
|
||||
batch_op.create_primary_key('pk_user_id', ['id'])
|
||||
|
||||
# Create oauth_session table
|
||||
op.create_table(
|
||||
"oauth_session",
|
||||
sa.Column("id", sa.Text(), primary_key=True, nullable=False, unique=True),
|
||||
'oauth_session',
|
||||
sa.Column('id', sa.Text(), primary_key=True, nullable=False, unique=True),
|
||||
sa.Column(
|
||||
"user_id",
|
||||
'user_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("user.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('user.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("provider", sa.Text(), nullable=False),
|
||||
sa.Column("token", sa.Text(), nullable=False),
|
||||
sa.Column("expires_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column('provider', sa.Text(), nullable=False),
|
||||
sa.Column('token', sa.Text(), nullable=False),
|
||||
sa.Column('expires_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
|
||||
# Create indexes for better performance
|
||||
op.create_index("idx_oauth_session_user_id", "oauth_session", ["user_id"])
|
||||
op.create_index("idx_oauth_session_expires_at", "oauth_session", ["expires_at"])
|
||||
op.create_index(
|
||||
"idx_oauth_session_user_provider", "oauth_session", ["user_id", "provider"]
|
||||
)
|
||||
op.create_index('idx_oauth_session_user_id', 'oauth_session', ['user_id'])
|
||||
op.create_index('idx_oauth_session_expires_at', 'oauth_session', ['expires_at'])
|
||||
op.create_index('idx_oauth_session_user_provider', 'oauth_session', ['user_id', 'provider'])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop indexes first
|
||||
op.drop_index("idx_oauth_session_user_provider", table_name="oauth_session")
|
||||
op.drop_index("idx_oauth_session_expires_at", table_name="oauth_session")
|
||||
op.drop_index("idx_oauth_session_user_id", table_name="oauth_session")
|
||||
op.drop_index('idx_oauth_session_user_provider', table_name='oauth_session')
|
||||
op.drop_index('idx_oauth_session_expires_at', table_name='oauth_session')
|
||||
op.drop_index('idx_oauth_session_user_id', table_name='oauth_session')
|
||||
|
||||
# Drop the table
|
||||
op.drop_table("oauth_session")
|
||||
op.drop_table('oauth_session')
|
||||
|
||||
@@ -13,8 +13,8 @@ from sqlalchemy.engine.reflection import Inspector
|
||||
|
||||
import json
|
||||
|
||||
revision = "3ab32c4b8f59"
|
||||
down_revision = "1af9b942657b"
|
||||
revision = '3ab32c4b8f59'
|
||||
down_revision = '1af9b942657b'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
@@ -24,58 +24,55 @@ def upgrade():
|
||||
inspector = Inspector.from_engine(conn)
|
||||
|
||||
# Inspecting the 'tag' table constraints and structure
|
||||
existing_pk = inspector.get_pk_constraint("tag")
|
||||
unique_constraints = inspector.get_unique_constraints("tag")
|
||||
existing_indexes = inspector.get_indexes("tag")
|
||||
existing_pk = inspector.get_pk_constraint('tag')
|
||||
unique_constraints = inspector.get_unique_constraints('tag')
|
||||
existing_indexes = inspector.get_indexes('tag')
|
||||
|
||||
print(f"Primary Key: {existing_pk}")
|
||||
print(f"Unique Constraints: {unique_constraints}")
|
||||
print(f"Indexes: {existing_indexes}")
|
||||
print(f'Primary Key: {existing_pk}')
|
||||
print(f'Unique Constraints: {unique_constraints}')
|
||||
print(f'Indexes: {existing_indexes}')
|
||||
|
||||
with op.batch_alter_table("tag", schema=None) as batch_op:
|
||||
with op.batch_alter_table('tag', schema=None) as batch_op:
|
||||
# Drop existing primary key constraint if it exists
|
||||
if existing_pk and existing_pk.get("constrained_columns"):
|
||||
pk_name = existing_pk.get("name")
|
||||
if existing_pk and existing_pk.get('constrained_columns'):
|
||||
pk_name = existing_pk.get('name')
|
||||
if pk_name:
|
||||
print(f"Dropping primary key constraint: {pk_name}")
|
||||
batch_op.drop_constraint(pk_name, type_="primary")
|
||||
print(f'Dropping primary key constraint: {pk_name}')
|
||||
batch_op.drop_constraint(pk_name, type_='primary')
|
||||
|
||||
# Now create the new primary key with the combination of 'id' and 'user_id'
|
||||
print("Creating new primary key with 'id' and 'user_id'.")
|
||||
batch_op.create_primary_key("pk_id_user_id", ["id", "user_id"])
|
||||
batch_op.create_primary_key('pk_id_user_id', ['id', 'user_id'])
|
||||
|
||||
# Drop unique constraints that could conflict with the new primary key
|
||||
for constraint in unique_constraints:
|
||||
if (
|
||||
constraint["name"] == "uq_id_user_id"
|
||||
constraint['name'] == 'uq_id_user_id'
|
||||
): # Adjust this name according to what is actually returned by the inspector
|
||||
print(f"Dropping unique constraint: {constraint['name']}")
|
||||
batch_op.drop_constraint(constraint["name"], type_="unique")
|
||||
print(f'Dropping unique constraint: {constraint["name"]}')
|
||||
batch_op.drop_constraint(constraint['name'], type_='unique')
|
||||
|
||||
for index in existing_indexes:
|
||||
if index["unique"]:
|
||||
if not any(
|
||||
constraint["name"] == index["name"]
|
||||
for constraint in unique_constraints
|
||||
):
|
||||
if index['unique']:
|
||||
if not any(constraint['name'] == index['name'] for constraint in unique_constraints):
|
||||
# You are attempting to drop unique indexes
|
||||
print(f"Dropping unique index: {index['name']}")
|
||||
batch_op.drop_index(index["name"])
|
||||
print(f'Dropping unique index: {index["name"]}')
|
||||
batch_op.drop_index(index['name'])
|
||||
|
||||
|
||||
def downgrade():
|
||||
conn = op.get_bind()
|
||||
inspector = Inspector.from_engine(conn)
|
||||
|
||||
current_pk = inspector.get_pk_constraint("tag")
|
||||
current_pk = inspector.get_pk_constraint('tag')
|
||||
|
||||
with op.batch_alter_table("tag", schema=None) as batch_op:
|
||||
with op.batch_alter_table('tag', schema=None) as batch_op:
|
||||
# Drop the current primary key first, if it matches the one we know we added in upgrade
|
||||
if current_pk and "pk_id_user_id" == current_pk.get("name"):
|
||||
batch_op.drop_constraint("pk_id_user_id", type_="primary")
|
||||
if current_pk and 'pk_id_user_id' == current_pk.get('name'):
|
||||
batch_op.drop_constraint('pk_id_user_id', type_='primary')
|
||||
|
||||
# Restore the original primary key
|
||||
batch_op.create_primary_key("pk_id", ["id"])
|
||||
batch_op.create_primary_key('pk_id', ['id'])
|
||||
|
||||
# Since primary key on just 'id' is restored, we now add back any unique constraints if necessary
|
||||
batch_op.create_unique_constraint("uq_id_user_id", ["id", "user_id"])
|
||||
batch_op.create_unique_constraint('uq_id_user_id', ['id', 'user_id'])
|
||||
|
||||
@@ -12,21 +12,21 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "3af16a1c9fb6"
|
||||
down_revision: Union[str, None] = "018012973d35"
|
||||
revision: str = '3af16a1c9fb6'
|
||||
down_revision: Union[str, None] = '018012973d35'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("user", sa.Column("username", sa.String(length=50), nullable=True))
|
||||
op.add_column("user", sa.Column("bio", sa.Text(), nullable=True))
|
||||
op.add_column("user", sa.Column("gender", sa.Text(), nullable=True))
|
||||
op.add_column("user", sa.Column("date_of_birth", sa.Date(), nullable=True))
|
||||
op.add_column('user', sa.Column('username', sa.String(length=50), nullable=True))
|
||||
op.add_column('user', sa.Column('bio', sa.Text(), nullable=True))
|
||||
op.add_column('user', sa.Column('gender', sa.Text(), nullable=True))
|
||||
op.add_column('user', sa.Column('date_of_birth', sa.Date(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("user", "username")
|
||||
op.drop_column("user", "bio")
|
||||
op.drop_column("user", "gender")
|
||||
op.drop_column("user", "date_of_birth")
|
||||
op.drop_column('user', 'username')
|
||||
op.drop_column('user', 'bio')
|
||||
op.drop_column('user', 'gender')
|
||||
op.drop_column('user', 'date_of_birth')
|
||||
|
||||
@@ -18,38 +18,38 @@ import json
|
||||
import uuid
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "3e0e00844bb0"
|
||||
down_revision: Union[str, None] = "90ef40d4714e"
|
||||
revision: str = '3e0e00844bb0'
|
||||
down_revision: Union[str, None] = '90ef40d4714e'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"knowledge_file",
|
||||
sa.Column("id", sa.Text(), primary_key=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=False),
|
||||
'knowledge_file',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column(
|
||||
"knowledge_id",
|
||||
'knowledge_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("knowledge.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('knowledge.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"file_id",
|
||||
'file_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("file.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('file.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
# indexes
|
||||
sa.Index("ix_knowledge_file_knowledge_id", "knowledge_id"),
|
||||
sa.Index("ix_knowledge_file_file_id", "file_id"),
|
||||
sa.Index("ix_knowledge_file_user_id", "user_id"),
|
||||
sa.Index('ix_knowledge_file_knowledge_id', 'knowledge_id'),
|
||||
sa.Index('ix_knowledge_file_file_id', 'file_id'),
|
||||
sa.Index('ix_knowledge_file_user_id', 'user_id'),
|
||||
# unique constraints
|
||||
sa.UniqueConstraint(
|
||||
"knowledge_id", "file_id", name="uq_knowledge_file_knowledge_file"
|
||||
'knowledge_id', 'file_id', name='uq_knowledge_file_knowledge_file'
|
||||
), # prevent duplicate entries
|
||||
)
|
||||
|
||||
@@ -57,35 +57,33 @@ def upgrade() -> None:
|
||||
|
||||
# 2. Read existing group with user_ids JSON column
|
||||
knowledge_table = sa.Table(
|
||||
"knowledge",
|
||||
'knowledge',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("user_id", sa.Text()),
|
||||
sa.Column("data", sa.JSON()), # JSON stored as text in SQLite + PG
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('user_id', sa.Text()),
|
||||
sa.Column('data', sa.JSON()), # JSON stored as text in SQLite + PG
|
||||
)
|
||||
|
||||
results = connection.execute(
|
||||
sa.select(
|
||||
knowledge_table.c.id, knowledge_table.c.user_id, knowledge_table.c.data
|
||||
)
|
||||
sa.select(knowledge_table.c.id, knowledge_table.c.user_id, knowledge_table.c.data)
|
||||
).fetchall()
|
||||
|
||||
# 3. Insert members into group_member table
|
||||
kf_table = sa.Table(
|
||||
"knowledge_file",
|
||||
'knowledge_file',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("user_id", sa.Text()),
|
||||
sa.Column("knowledge_id", sa.Text()),
|
||||
sa.Column("file_id", sa.Text()),
|
||||
sa.Column("created_at", sa.BigInteger()),
|
||||
sa.Column("updated_at", sa.BigInteger()),
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('user_id', sa.Text()),
|
||||
sa.Column('knowledge_id', sa.Text()),
|
||||
sa.Column('file_id', sa.Text()),
|
||||
sa.Column('created_at', sa.BigInteger()),
|
||||
sa.Column('updated_at', sa.BigInteger()),
|
||||
)
|
||||
|
||||
file_table = sa.Table(
|
||||
"file",
|
||||
'file',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column('id', sa.Text()),
|
||||
)
|
||||
|
||||
now = int(time.time())
|
||||
@@ -102,50 +100,48 @@ def upgrade() -> None:
|
||||
if not isinstance(data, dict):
|
||||
continue
|
||||
|
||||
file_ids = data.get("file_ids", [])
|
||||
file_ids = data.get('file_ids', [])
|
||||
|
||||
for file_id in file_ids:
|
||||
file_exists = connection.execute(
|
||||
sa.select(file_table.c.id).where(file_table.c.id == file_id)
|
||||
).fetchone()
|
||||
file_exists = connection.execute(sa.select(file_table.c.id).where(file_table.c.id == file_id)).fetchone()
|
||||
|
||||
if not file_exists:
|
||||
continue # skip non-existing files
|
||||
|
||||
row = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_id": user_id,
|
||||
"knowledge_id": knowledge_id,
|
||||
"file_id": file_id,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
'id': str(uuid.uuid4()),
|
||||
'user_id': user_id,
|
||||
'knowledge_id': knowledge_id,
|
||||
'file_id': file_id,
|
||||
'created_at': now,
|
||||
'updated_at': now,
|
||||
}
|
||||
connection.execute(kf_table.insert().values(**row))
|
||||
|
||||
with op.batch_alter_table("knowledge") as batch:
|
||||
batch.drop_column("data")
|
||||
with op.batch_alter_table('knowledge') as batch:
|
||||
batch.drop_column('data')
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# 1. Add back the old data column
|
||||
op.add_column("knowledge", sa.Column("data", sa.JSON(), nullable=True))
|
||||
op.add_column('knowledge', sa.Column('data', sa.JSON(), nullable=True))
|
||||
|
||||
connection = op.get_bind()
|
||||
|
||||
# 2. Read knowledge_file entries and reconstruct data JSON
|
||||
knowledge_table = sa.Table(
|
||||
"knowledge",
|
||||
'knowledge',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("data", sa.JSON()),
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('data', sa.JSON()),
|
||||
)
|
||||
|
||||
kf_table = sa.Table(
|
||||
"knowledge_file",
|
||||
'knowledge_file',
|
||||
sa.MetaData(),
|
||||
sa.Column("id", sa.Text()),
|
||||
sa.Column("knowledge_id", sa.Text()),
|
||||
sa.Column("file_id", sa.Text()),
|
||||
sa.Column('id', sa.Text()),
|
||||
sa.Column('knowledge_id', sa.Text()),
|
||||
sa.Column('file_id', sa.Text()),
|
||||
)
|
||||
|
||||
results = connection.execute(sa.select(knowledge_table.c.id)).fetchall()
|
||||
@@ -157,13 +153,9 @@ def downgrade() -> None:
|
||||
|
||||
file_ids_list = [fid for (fid,) in file_ids]
|
||||
|
||||
data_json = {"file_ids": file_ids_list}
|
||||
data_json = {'file_ids': file_ids_list}
|
||||
|
||||
connection.execute(
|
||||
knowledge_table.update()
|
||||
.where(knowledge_table.c.id == knowledge_id)
|
||||
.values(data=data_json)
|
||||
)
|
||||
connection.execute(knowledge_table.update().where(knowledge_table.c.id == knowledge_id).values(data=data_json))
|
||||
|
||||
# 3. Drop the knowledge_file table
|
||||
op.drop_table("knowledge_file")
|
||||
op.drop_table('knowledge_file')
|
||||
|
||||
+12
-12
@@ -9,56 +9,56 @@ Create Date: 2024-10-23 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "4ace53fd72c8"
|
||||
down_revision = "af906e964978"
|
||||
revision = '4ace53fd72c8'
|
||||
down_revision = 'af906e964978'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
# Perform safe alterations using batch operation
|
||||
with op.batch_alter_table("folder", schema=None) as batch_op:
|
||||
with op.batch_alter_table('folder', schema=None) as batch_op:
|
||||
# Step 1: Remove server defaults for created_at and updated_at
|
||||
batch_op.alter_column(
|
||||
"created_at",
|
||||
'created_at',
|
||||
server_default=None, # Removing server default
|
||||
)
|
||||
batch_op.alter_column(
|
||||
"updated_at",
|
||||
'updated_at',
|
||||
server_default=None, # Removing server default
|
||||
)
|
||||
|
||||
# Step 2: Change the column types to BigInteger for created_at
|
||||
batch_op.alter_column(
|
||||
"created_at",
|
||||
'created_at',
|
||||
type_=sa.BigInteger(),
|
||||
existing_type=sa.DateTime(),
|
||||
existing_nullable=False,
|
||||
postgresql_using="extract(epoch from created_at)::bigint", # Conversion for PostgreSQL
|
||||
postgresql_using='extract(epoch from created_at)::bigint', # Conversion for PostgreSQL
|
||||
)
|
||||
|
||||
# Change the column types to BigInteger for updated_at
|
||||
batch_op.alter_column(
|
||||
"updated_at",
|
||||
'updated_at',
|
||||
type_=sa.BigInteger(),
|
||||
existing_type=sa.DateTime(),
|
||||
existing_nullable=False,
|
||||
postgresql_using="extract(epoch from updated_at)::bigint", # Conversion for PostgreSQL
|
||||
postgresql_using='extract(epoch from updated_at)::bigint', # Conversion for PostgreSQL
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
# Downgrade: Convert columns back to DateTime and restore defaults
|
||||
with op.batch_alter_table("folder", schema=None) as batch_op:
|
||||
with op.batch_alter_table('folder', schema=None) as batch_op:
|
||||
batch_op.alter_column(
|
||||
"created_at",
|
||||
'created_at',
|
||||
type_=sa.DateTime(),
|
||||
existing_type=sa.BigInteger(),
|
||||
existing_nullable=False,
|
||||
server_default=sa.func.now(), # Restoring server default on downgrade
|
||||
)
|
||||
batch_op.alter_column(
|
||||
"updated_at",
|
||||
'updated_at',
|
||||
type_=sa.DateTime(),
|
||||
existing_type=sa.BigInteger(),
|
||||
existing_nullable=False,
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
"""add calendar tables
|
||||
|
||||
Revision ID: 56359461a091
|
||||
Revises: c1d2e3f4a5b6
|
||||
Create Date: 2026-04-19 16:20:58.162045
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = '56359461a091'
|
||||
down_revision: Union[str, None] = 'c1d2e3f4a5b6'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
'calendar',
|
||||
sa.Column('id', sa.Text(), nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False),
|
||||
sa.Column('color', sa.Text(), nullable=True),
|
||||
sa.Column('is_default', sa.Boolean(), nullable=False),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
op.create_index('ix_calendar_user', 'calendar', ['user_id'], unique=False)
|
||||
|
||||
op.create_table(
|
||||
'calendar_event',
|
||||
sa.Column('id', sa.Text(), nullable=False),
|
||||
sa.Column('calendar_id', sa.Text(), nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('title', sa.Text(), nullable=False),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('start_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('end_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('all_day', sa.Boolean(), nullable=False),
|
||||
sa.Column('rrule', sa.Text(), nullable=True),
|
||||
sa.Column('color', sa.Text(), nullable=True),
|
||||
sa.Column('location', sa.Text(), nullable=True),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('is_cancelled', sa.Boolean(), nullable=False),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
op.create_index('ix_calendar_event_calendar', 'calendar_event', ['calendar_id', 'start_at'], unique=False)
|
||||
op.create_index('ix_calendar_event_user_date', 'calendar_event', ['user_id', 'start_at'], unique=False)
|
||||
|
||||
op.create_table(
|
||||
'calendar_event_attendee',
|
||||
sa.Column('id', sa.Text(), nullable=False),
|
||||
sa.Column('event_id', sa.Text(), nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('status', sa.Text(), nullable=False),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('event_id', 'user_id', name='uq_event_attendee'),
|
||||
)
|
||||
op.create_index('ix_calendar_event_attendee_user', 'calendar_event_attendee', ['user_id', 'status'], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index('ix_calendar_event_attendee_user', table_name='calendar_event_attendee')
|
||||
op.drop_table('calendar_event_attendee')
|
||||
op.drop_index('ix_calendar_event_user_date', table_name='calendar_event')
|
||||
op.drop_index('ix_calendar_event_calendar', table_name='calendar_event')
|
||||
op.drop_table('calendar_event')
|
||||
op.drop_index('ix_calendar_user', table_name='calendar')
|
||||
op.drop_table('calendar')
|
||||
@@ -9,40 +9,40 @@ Create Date: 2024-12-22 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "57c599a3cb57"
|
||||
down_revision = "922e7a387820"
|
||||
revision = '57c599a3cb57'
|
||||
down_revision = '922e7a387820'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
"channel",
|
||||
sa.Column("id", sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column("user_id", sa.Text()),
|
||||
sa.Column("name", sa.Text()),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("data", sa.JSON(), nullable=True),
|
||||
sa.Column("meta", sa.JSON(), nullable=True),
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
'channel',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column('user_id', sa.Text()),
|
||||
sa.Column('name', sa.Text()),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
|
||||
op.create_table(
|
||||
"message",
|
||||
sa.Column("id", sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column("user_id", sa.Text()),
|
||||
sa.Column("channel_id", sa.Text(), nullable=True),
|
||||
sa.Column("content", sa.Text()),
|
||||
sa.Column("data", sa.JSON(), nullable=True),
|
||||
sa.Column("meta", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
'message',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column('user_id', sa.Text()),
|
||||
sa.Column('channel_id', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.Text()),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("channel")
|
||||
op.drop_table('channel')
|
||||
|
||||
op.drop_table("message")
|
||||
op.drop_table('message')
|
||||
|
||||
@@ -12,43 +12,40 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
import open_webui.internal.db
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "6283dc0e4d8d"
|
||||
down_revision: Union[str, None] = "3e0e00844bb0"
|
||||
revision: str = '6283dc0e4d8d'
|
||||
down_revision: Union[str, None] = '3e0e00844bb0'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"channel_file",
|
||||
sa.Column("id", sa.Text(), primary_key=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=False),
|
||||
'channel_file',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column(
|
||||
"channel_id",
|
||||
'channel_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("channel.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('channel.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"file_id",
|
||||
'file_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("file.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('file.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
# indexes
|
||||
sa.Index("ix_channel_file_channel_id", "channel_id"),
|
||||
sa.Index("ix_channel_file_file_id", "file_id"),
|
||||
sa.Index("ix_channel_file_user_id", "user_id"),
|
||||
sa.Index('ix_channel_file_channel_id', 'channel_id'),
|
||||
sa.Index('ix_channel_file_file_id', 'file_id'),
|
||||
sa.Index('ix_channel_file_user_id', 'user_id'),
|
||||
# unique constraints
|
||||
sa.UniqueConstraint(
|
||||
"channel_id", "file_id", name="uq_channel_file_channel_file"
|
||||
), # prevent duplicate entries
|
||||
sa.UniqueConstraint('channel_id', 'file_id', name='uq_channel_file_channel_file'), # prevent duplicate entries
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("channel_file")
|
||||
op.drop_table('channel_file')
|
||||
|
||||
@@ -11,38 +11,37 @@ import sqlalchemy as sa
|
||||
from sqlalchemy.sql import table, column, select
|
||||
import json
|
||||
|
||||
|
||||
revision = "6a39f3d8e55c"
|
||||
down_revision = "c0fbf31ca0db"
|
||||
revision = '6a39f3d8e55c'
|
||||
down_revision = 'c0fbf31ca0db'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
# Creating the 'knowledge' table
|
||||
print("Creating knowledge table")
|
||||
print('Creating knowledge table')
|
||||
knowledge_table = op.create_table(
|
||||
"knowledge",
|
||||
sa.Column("id", sa.Text(), primary_key=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=False),
|
||||
sa.Column("name", sa.Text(), nullable=False),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("data", sa.JSON(), nullable=True),
|
||||
sa.Column("meta", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
'knowledge',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
|
||||
print("Migrating data from document table to knowledge table")
|
||||
print('Migrating data from document table to knowledge table')
|
||||
# Representation of the existing 'document' table
|
||||
document_table = table(
|
||||
"document",
|
||||
column("collection_name", sa.String()),
|
||||
column("user_id", sa.String()),
|
||||
column("name", sa.String()),
|
||||
column("title", sa.Text()),
|
||||
column("content", sa.Text()),
|
||||
column("timestamp", sa.BigInteger()),
|
||||
'document',
|
||||
column('collection_name', sa.String()),
|
||||
column('user_id', sa.String()),
|
||||
column('name', sa.String()),
|
||||
column('title', sa.Text()),
|
||||
column('content', sa.Text()),
|
||||
column('timestamp', sa.BigInteger()),
|
||||
)
|
||||
|
||||
# Select all from existing document table
|
||||
@@ -65,9 +64,9 @@ def upgrade():
|
||||
user_id=doc.user_id,
|
||||
description=doc.name,
|
||||
meta={
|
||||
"legacy": True,
|
||||
"document": True,
|
||||
"tags": json.loads(doc.content or "{}").get("tags", []),
|
||||
'legacy': True,
|
||||
'document': True,
|
||||
'tags': json.loads(doc.content or '{}').get('tags', []),
|
||||
},
|
||||
name=doc.title,
|
||||
created_at=doc.timestamp,
|
||||
@@ -77,4 +76,4 @@ def upgrade():
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("knowledge")
|
||||
op.drop_table('knowledge')
|
||||
|
||||
@@ -9,18 +9,18 @@ Create Date: 2024-12-23 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "7826ab40b532"
|
||||
down_revision = "57c599a3cb57"
|
||||
revision = '7826ab40b532'
|
||||
down_revision = '57c599a3cb57'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.add_column(
|
||||
"file",
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
'file',
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_column("file", "access_control")
|
||||
op.drop_column('file', 'access_control')
|
||||
|
||||
@@ -16,7 +16,7 @@ from open_webui.internal.db import JSONField
|
||||
from open_webui.migrations.util import get_existing_tables
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "7e5b5dc7342b"
|
||||
revision: str = '7e5b5dc7342b'
|
||||
down_revision: Union[str, None] = None
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
@@ -26,179 +26,179 @@ def upgrade() -> None:
|
||||
existing_tables = set(get_existing_tables())
|
||||
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
if "auth" not in existing_tables:
|
||||
if 'auth' not in existing_tables:
|
||||
op.create_table(
|
||||
"auth",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("email", sa.String(), nullable=True),
|
||||
sa.Column("password", sa.Text(), nullable=True),
|
||||
sa.Column("active", sa.Boolean(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'auth',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('email', sa.String(), nullable=True),
|
||||
sa.Column('password', sa.Text(), nullable=True),
|
||||
sa.Column('active', sa.Boolean(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "chat" not in existing_tables:
|
||||
if 'chat' not in existing_tables:
|
||||
op.create_table(
|
||||
"chat",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("title", sa.Text(), nullable=True),
|
||||
sa.Column("chat", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("share_id", sa.Text(), nullable=True),
|
||||
sa.Column("archived", sa.Boolean(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("share_id"),
|
||||
'chat',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('title', sa.Text(), nullable=True),
|
||||
sa.Column('chat', sa.Text(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('share_id', sa.Text(), nullable=True),
|
||||
sa.Column('archived', sa.Boolean(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('share_id'),
|
||||
)
|
||||
|
||||
if "chatidtag" not in existing_tables:
|
||||
if 'chatidtag' not in existing_tables:
|
||||
op.create_table(
|
||||
"chatidtag",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("tag_name", sa.String(), nullable=True),
|
||||
sa.Column("chat_id", sa.String(), nullable=True),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("timestamp", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'chatidtag',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('tag_name', sa.String(), nullable=True),
|
||||
sa.Column('chat_id', sa.String(), nullable=True),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('timestamp', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "document" not in existing_tables:
|
||||
if 'document' not in existing_tables:
|
||||
op.create_table(
|
||||
"document",
|
||||
sa.Column("collection_name", sa.String(), nullable=False),
|
||||
sa.Column("name", sa.String(), nullable=True),
|
||||
sa.Column("title", sa.Text(), nullable=True),
|
||||
sa.Column("filename", sa.Text(), nullable=True),
|
||||
sa.Column("content", sa.Text(), nullable=True),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("timestamp", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("collection_name"),
|
||||
sa.UniqueConstraint("name"),
|
||||
'document',
|
||||
sa.Column('collection_name', sa.String(), nullable=False),
|
||||
sa.Column('name', sa.String(), nullable=True),
|
||||
sa.Column('title', sa.Text(), nullable=True),
|
||||
sa.Column('filename', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.Text(), nullable=True),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('timestamp', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('collection_name'),
|
||||
sa.UniqueConstraint('name'),
|
||||
)
|
||||
|
||||
if "file" not in existing_tables:
|
||||
if 'file' not in existing_tables:
|
||||
op.create_table(
|
||||
"file",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("filename", sa.Text(), nullable=True),
|
||||
sa.Column("meta", JSONField(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'file',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('filename', sa.Text(), nullable=True),
|
||||
sa.Column('meta', JSONField(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "function" not in existing_tables:
|
||||
if 'function' not in existing_tables:
|
||||
op.create_table(
|
||||
"function",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("name", sa.Text(), nullable=True),
|
||||
sa.Column("type", sa.Text(), nullable=True),
|
||||
sa.Column("content", sa.Text(), nullable=True),
|
||||
sa.Column("meta", JSONField(), nullable=True),
|
||||
sa.Column("valves", JSONField(), nullable=True),
|
||||
sa.Column("is_active", sa.Boolean(), nullable=True),
|
||||
sa.Column("is_global", sa.Boolean(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'function',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('name', sa.Text(), nullable=True),
|
||||
sa.Column('type', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.Text(), nullable=True),
|
||||
sa.Column('meta', JSONField(), nullable=True),
|
||||
sa.Column('valves', JSONField(), nullable=True),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=True),
|
||||
sa.Column('is_global', sa.Boolean(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "memory" not in existing_tables:
|
||||
if 'memory' not in existing_tables:
|
||||
op.create_table(
|
||||
"memory",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("content", sa.Text(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'memory',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('content', sa.Text(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "model" not in existing_tables:
|
||||
if 'model' not in existing_tables:
|
||||
op.create_table(
|
||||
"model",
|
||||
sa.Column("id", sa.Text(), nullable=False),
|
||||
sa.Column("user_id", sa.Text(), nullable=True),
|
||||
sa.Column("base_model_id", sa.Text(), nullable=True),
|
||||
sa.Column("name", sa.Text(), nullable=True),
|
||||
sa.Column("params", JSONField(), nullable=True),
|
||||
sa.Column("meta", JSONField(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'model',
|
||||
sa.Column('id', sa.Text(), nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=True),
|
||||
sa.Column('base_model_id', sa.Text(), nullable=True),
|
||||
sa.Column('name', sa.Text(), nullable=True),
|
||||
sa.Column('params', JSONField(), nullable=True),
|
||||
sa.Column('meta', JSONField(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "prompt" not in existing_tables:
|
||||
if 'prompt' not in existing_tables:
|
||||
op.create_table(
|
||||
"prompt",
|
||||
sa.Column("command", sa.String(), nullable=False),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("title", sa.Text(), nullable=True),
|
||||
sa.Column("content", sa.Text(), nullable=True),
|
||||
sa.Column("timestamp", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("command"),
|
||||
'prompt',
|
||||
sa.Column('command', sa.String(), nullable=False),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('title', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.Text(), nullable=True),
|
||||
sa.Column('timestamp', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('command'),
|
||||
)
|
||||
|
||||
if "tag" not in existing_tables:
|
||||
if 'tag' not in existing_tables:
|
||||
op.create_table(
|
||||
"tag",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("name", sa.String(), nullable=True),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("data", sa.Text(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'tag',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('name', sa.String(), nullable=True),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('data', sa.Text(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "tool" not in existing_tables:
|
||||
if 'tool' not in existing_tables:
|
||||
op.create_table(
|
||||
"tool",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("user_id", sa.String(), nullable=True),
|
||||
sa.Column("name", sa.Text(), nullable=True),
|
||||
sa.Column("content", sa.Text(), nullable=True),
|
||||
sa.Column("specs", JSONField(), nullable=True),
|
||||
sa.Column("meta", JSONField(), nullable=True),
|
||||
sa.Column("valves", JSONField(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
'tool',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('user_id', sa.String(), nullable=True),
|
||||
sa.Column('name', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.Text(), nullable=True),
|
||||
sa.Column('specs', JSONField(), nullable=True),
|
||||
sa.Column('meta', JSONField(), nullable=True),
|
||||
sa.Column('valves', JSONField(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
)
|
||||
|
||||
if "user" not in existing_tables:
|
||||
if 'user' not in existing_tables:
|
||||
op.create_table(
|
||||
"user",
|
||||
sa.Column("id", sa.String(), nullable=False),
|
||||
sa.Column("name", sa.String(), nullable=True),
|
||||
sa.Column("email", sa.String(), nullable=True),
|
||||
sa.Column("role", sa.String(), nullable=True),
|
||||
sa.Column("profile_image_url", sa.Text(), nullable=True),
|
||||
sa.Column("last_active_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("api_key", sa.String(), nullable=True),
|
||||
sa.Column("settings", JSONField(), nullable=True),
|
||||
sa.Column("info", JSONField(), nullable=True),
|
||||
sa.Column("oauth_sub", sa.Text(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("api_key"),
|
||||
sa.UniqueConstraint("oauth_sub"),
|
||||
'user',
|
||||
sa.Column('id', sa.String(), nullable=False),
|
||||
sa.Column('name', sa.String(), nullable=True),
|
||||
sa.Column('email', sa.String(), nullable=True),
|
||||
sa.Column('role', sa.String(), nullable=True),
|
||||
sa.Column('profile_image_url', sa.Text(), nullable=True),
|
||||
sa.Column('last_active_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('api_key', sa.String(), nullable=True),
|
||||
sa.Column('settings', JSONField(), nullable=True),
|
||||
sa.Column('info', JSONField(), nullable=True),
|
||||
sa.Column('oauth_sub', sa.Text(), nullable=True),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('api_key'),
|
||||
sa.UniqueConstraint('oauth_sub'),
|
||||
)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_table("user")
|
||||
op.drop_table("tool")
|
||||
op.drop_table("tag")
|
||||
op.drop_table("prompt")
|
||||
op.drop_table("model")
|
||||
op.drop_table("memory")
|
||||
op.drop_table("function")
|
||||
op.drop_table("file")
|
||||
op.drop_table("document")
|
||||
op.drop_table("chatidtag")
|
||||
op.drop_table("chat")
|
||||
op.drop_table("auth")
|
||||
op.drop_table('user')
|
||||
op.drop_table('tool')
|
||||
op.drop_table('tag')
|
||||
op.drop_table('prompt')
|
||||
op.drop_table('model')
|
||||
op.drop_table('memory')
|
||||
op.drop_table('function')
|
||||
op.drop_table('file')
|
||||
op.drop_table('document')
|
||||
op.drop_table('chatidtag')
|
||||
op.drop_table('chat')
|
||||
op.drop_table('auth')
|
||||
# ### end Alembic commands ###
|
||||
|
||||
+11
-14
@@ -12,38 +12,35 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
import open_webui.internal.db
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "81cc2ce44d79"
|
||||
down_revision: Union[str, None] = "6283dc0e4d8d"
|
||||
revision: str = '81cc2ce44d79'
|
||||
down_revision: Union[str, None] = '6283dc0e4d8d'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add message_id column to channel_file table
|
||||
with op.batch_alter_table("channel_file", schema=None) as batch_op:
|
||||
with op.batch_alter_table('channel_file', schema=None) as batch_op:
|
||||
batch_op.add_column(
|
||||
sa.Column(
|
||||
"message_id",
|
||||
'message_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey(
|
||||
"message.id", ondelete="CASCADE", name="fk_channel_file_message_id"
|
||||
),
|
||||
sa.ForeignKey('message.id', ondelete='CASCADE', name='fk_channel_file_message_id'),
|
||||
nullable=True,
|
||||
)
|
||||
)
|
||||
|
||||
# Add data column to knowledge table
|
||||
with op.batch_alter_table("knowledge", schema=None) as batch_op:
|
||||
batch_op.add_column(sa.Column("data", sa.JSON(), nullable=True))
|
||||
with op.batch_alter_table('knowledge', schema=None) as batch_op:
|
||||
batch_op.add_column(sa.Column('data', sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Remove message_id column from channel_file table
|
||||
with op.batch_alter_table("channel_file", schema=None) as batch_op:
|
||||
batch_op.drop_column("message_id")
|
||||
with op.batch_alter_table('channel_file', schema=None) as batch_op:
|
||||
batch_op.drop_column('message_id')
|
||||
|
||||
# Remove data column from knowledge table
|
||||
with op.batch_alter_table("knowledge", schema=None) as batch_op:
|
||||
batch_op.drop_column("data")
|
||||
with op.batch_alter_table('knowledge', schema=None) as batch_op:
|
||||
batch_op.drop_column('data')
|
||||
|
||||
@@ -0,0 +1,221 @@
|
||||
"""Add chat_message table
|
||||
|
||||
Revision ID: 8452d01d26d7
|
||||
Revises: 374d2f66af06
|
||||
Create Date: 2026-02-01 04:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
import time
|
||||
import json
|
||||
import logging
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
revision: str = '8452d01d26d7'
|
||||
down_revision: Union[str, None] = '374d2f66af06'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
BATCH_SIZE = 5000
|
||||
|
||||
|
||||
def _flush_batch(conn, table, batch):
|
||||
"""
|
||||
Insert a batch of messages, falling back to row-by-row on error.
|
||||
|
||||
Tries a single bulk insert first (fast path). If that fails (e.g. due to
|
||||
a duplicate key), falls back to individual inserts wrapped in savepoints
|
||||
so the rest of the batch can still succeed.
|
||||
"""
|
||||
savepoint = conn.begin_nested()
|
||||
try:
|
||||
conn.execute(sa.insert(table), batch)
|
||||
savepoint.commit()
|
||||
return len(batch), 0
|
||||
except Exception:
|
||||
savepoint.rollback()
|
||||
# Batch failed - insert one-by-one to isolate the bad row(s)
|
||||
inserted = 0
|
||||
failed = 0
|
||||
for msg in batch:
|
||||
sp = conn.begin_nested()
|
||||
try:
|
||||
conn.execute(sa.insert(table).values(**msg))
|
||||
sp.commit()
|
||||
inserted += 1
|
||||
except Exception as e:
|
||||
sp.rollback()
|
||||
failed += 1
|
||||
log.warning(f'Failed to insert message {msg["id"]}: {e}')
|
||||
return inserted, failed
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Step 1: Create table
|
||||
op.create_table(
|
||||
'chat_message',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('chat_id', sa.Text(), nullable=False, index=True),
|
||||
sa.Column('user_id', sa.Text(), index=True),
|
||||
sa.Column('role', sa.Text(), nullable=False),
|
||||
sa.Column('parent_id', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.JSON(), nullable=True),
|
||||
sa.Column('output', sa.JSON(), nullable=True),
|
||||
sa.Column('model_id', sa.Text(), nullable=True, index=True),
|
||||
sa.Column('files', sa.JSON(), nullable=True),
|
||||
sa.Column('sources', sa.JSON(), nullable=True),
|
||||
sa.Column('embeds', sa.JSON(), nullable=True),
|
||||
sa.Column('done', sa.Boolean(), default=True),
|
||||
sa.Column('status_history', sa.JSON(), nullable=True),
|
||||
sa.Column('error', sa.JSON(), nullable=True),
|
||||
sa.Column('usage', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), index=True),
|
||||
sa.Column('updated_at', sa.BigInteger()),
|
||||
sa.ForeignKeyConstraint(['chat_id'], ['chat.id'], ondelete='CASCADE'),
|
||||
)
|
||||
|
||||
# Create composite indexes
|
||||
op.create_index('chat_message_chat_parent_idx', 'chat_message', ['chat_id', 'parent_id'])
|
||||
op.create_index('chat_message_model_created_idx', 'chat_message', ['model_id', 'created_at'])
|
||||
op.create_index('chat_message_user_created_idx', 'chat_message', ['user_id', 'created_at'])
|
||||
|
||||
# Step 2: Backfill from existing chats
|
||||
conn = op.get_bind()
|
||||
|
||||
chat_table = sa.table(
|
||||
'chat',
|
||||
sa.column('id', sa.Text()),
|
||||
sa.column('user_id', sa.Text()),
|
||||
sa.column('chat', sa.JSON()),
|
||||
)
|
||||
|
||||
chat_message_table = sa.table(
|
||||
'chat_message',
|
||||
sa.column('id', sa.Text()),
|
||||
sa.column('chat_id', sa.Text()),
|
||||
sa.column('user_id', sa.Text()),
|
||||
sa.column('role', sa.Text()),
|
||||
sa.column('parent_id', sa.Text()),
|
||||
sa.column('content', sa.JSON()),
|
||||
sa.column('output', sa.JSON()),
|
||||
sa.column('model_id', sa.Text()),
|
||||
sa.column('files', sa.JSON()),
|
||||
sa.column('sources', sa.JSON()),
|
||||
sa.column('embeds', sa.JSON()),
|
||||
sa.column('done', sa.Boolean()),
|
||||
sa.column('status_history', sa.JSON()),
|
||||
sa.column('error', sa.JSON()),
|
||||
sa.column('usage', sa.JSON()),
|
||||
sa.column('created_at', sa.BigInteger()),
|
||||
sa.column('updated_at', sa.BigInteger()),
|
||||
)
|
||||
|
||||
# Stream rows instead of loading all into memory:
|
||||
# - yield_per: fetches rows in chunks via cursor.fetchmany() (all backends)
|
||||
# - stream_results: enables server-side cursors on PostgreSQL (no-op on SQLite)
|
||||
result = conn.execute(
|
||||
sa.select(chat_table.c.id, chat_table.c.user_id, chat_table.c.chat)
|
||||
.where(~chat_table.c.user_id.like('shared-%'))
|
||||
.execution_options(yield_per=1000, stream_results=True)
|
||||
)
|
||||
|
||||
now = int(time.time())
|
||||
messages_batch = []
|
||||
total_inserted = 0
|
||||
total_failed = 0
|
||||
|
||||
for chat_row in result:
|
||||
chat_id = chat_row[0]
|
||||
user_id = chat_row[1]
|
||||
chat_data = chat_row[2]
|
||||
|
||||
if not chat_data:
|
||||
continue
|
||||
|
||||
# Handle both string and dict chat data
|
||||
if isinstance(chat_data, str):
|
||||
try:
|
||||
chat_data = json.loads(chat_data)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
history = chat_data.get('history', {})
|
||||
if not isinstance(history, dict):
|
||||
continue
|
||||
|
||||
messages = history.get('messages', {})
|
||||
if not isinstance(messages, dict):
|
||||
continue
|
||||
|
||||
for message_id, message in messages.items():
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
|
||||
role = message.get('role')
|
||||
if not role:
|
||||
continue
|
||||
|
||||
timestamp = message.get('timestamp', now)
|
||||
|
||||
try:
|
||||
timestamp = int(float(timestamp))
|
||||
except Exception as e:
|
||||
timestamp = now
|
||||
|
||||
# Normalize timestamp: convert ms to seconds, validate range
|
||||
if timestamp > 10_000_000_000:
|
||||
timestamp = timestamp // 1000
|
||||
# Must be after 2020 and not too far in the future
|
||||
if timestamp < 1577836800 or timestamp > now + 86400:
|
||||
timestamp = now
|
||||
|
||||
messages_batch.append(
|
||||
{
|
||||
'id': f'{chat_id}-{message_id}',
|
||||
'chat_id': chat_id,
|
||||
'user_id': user_id,
|
||||
'role': role,
|
||||
'parent_id': message.get('parentId'),
|
||||
'content': message.get('content'),
|
||||
'output': message.get('output'),
|
||||
'model_id': message.get('model'),
|
||||
'files': message.get('files'),
|
||||
'sources': message.get('sources'),
|
||||
'embeds': message.get('embeds'),
|
||||
'done': message.get('done', True),
|
||||
'status_history': message.get('statusHistory'),
|
||||
'error': message.get('error'),
|
||||
'usage': message.get('usage'),
|
||||
'created_at': timestamp,
|
||||
'updated_at': timestamp,
|
||||
}
|
||||
)
|
||||
|
||||
# Flush batch when full
|
||||
if len(messages_batch) >= BATCH_SIZE:
|
||||
inserted, failed = _flush_batch(conn, chat_message_table, messages_batch)
|
||||
total_inserted += inserted
|
||||
total_failed += failed
|
||||
if total_inserted % 50000 < BATCH_SIZE:
|
||||
log.info(f'Migration progress: {total_inserted} messages inserted...')
|
||||
messages_batch.clear()
|
||||
|
||||
# Flush remaining messages
|
||||
if messages_batch:
|
||||
inserted, failed = _flush_batch(conn, chat_message_table, messages_batch)
|
||||
total_inserted += inserted
|
||||
total_failed += failed
|
||||
|
||||
log.info(f'Backfilled {total_inserted} messages into chat_message table ({total_failed} failed)')
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index('chat_message_user_created_idx', table_name='chat_message')
|
||||
op.drop_index('chat_message_model_created_idx', table_name='chat_message')
|
||||
op.drop_index('chat_message_chat_parent_idx', table_name='chat_message')
|
||||
op.drop_table('chat_message')
|
||||
+32
-35
@@ -12,50 +12,47 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
import open_webui.internal.db
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "90ef40d4714e"
|
||||
down_revision: Union[str, None] = "b10670c03dd5"
|
||||
revision: str = '90ef40d4714e'
|
||||
down_revision: Union[str, None] = 'b10670c03dd5'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Update 'channel' table
|
||||
op.add_column("channel", sa.Column("is_private", sa.Boolean(), nullable=True))
|
||||
op.add_column('channel', sa.Column('is_private', sa.Boolean(), nullable=True))
|
||||
|
||||
op.add_column("channel", sa.Column("archived_at", sa.BigInteger(), nullable=True))
|
||||
op.add_column("channel", sa.Column("archived_by", sa.Text(), nullable=True))
|
||||
op.add_column('channel', sa.Column('archived_at', sa.BigInteger(), nullable=True))
|
||||
op.add_column('channel', sa.Column('archived_by', sa.Text(), nullable=True))
|
||||
|
||||
op.add_column("channel", sa.Column("deleted_at", sa.BigInteger(), nullable=True))
|
||||
op.add_column("channel", sa.Column("deleted_by", sa.Text(), nullable=True))
|
||||
op.add_column('channel', sa.Column('deleted_at', sa.BigInteger(), nullable=True))
|
||||
op.add_column('channel', sa.Column('deleted_by', sa.Text(), nullable=True))
|
||||
|
||||
op.add_column("channel", sa.Column("updated_by", sa.Text(), nullable=True))
|
||||
op.add_column('channel', sa.Column('updated_by', sa.Text(), nullable=True))
|
||||
|
||||
# Update 'channel_member' table
|
||||
op.add_column("channel_member", sa.Column("role", sa.Text(), nullable=True))
|
||||
op.add_column("channel_member", sa.Column("invited_by", sa.Text(), nullable=True))
|
||||
op.add_column(
|
||||
"channel_member", sa.Column("invited_at", sa.BigInteger(), nullable=True)
|
||||
)
|
||||
op.add_column('channel_member', sa.Column('role', sa.Text(), nullable=True))
|
||||
op.add_column('channel_member', sa.Column('invited_by', sa.Text(), nullable=True))
|
||||
op.add_column('channel_member', sa.Column('invited_at', sa.BigInteger(), nullable=True))
|
||||
|
||||
# Create 'channel_webhook' table
|
||||
op.create_table(
|
||||
"channel_webhook",
|
||||
sa.Column("id", sa.Text(), primary_key=True, unique=True, nullable=False),
|
||||
sa.Column("user_id", sa.Text(), nullable=False),
|
||||
'channel_webhook',
|
||||
sa.Column('id', sa.Text(), primary_key=True, unique=True, nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column(
|
||||
"channel_id",
|
||||
'channel_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("channel.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('channel.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("name", sa.Text(), nullable=False),
|
||||
sa.Column("profile_image_url", sa.Text(), nullable=True),
|
||||
sa.Column("token", sa.Text(), nullable=False),
|
||||
sa.Column("last_used_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False),
|
||||
sa.Column('profile_image_url', sa.Text(), nullable=True),
|
||||
sa.Column('token', sa.Text(), nullable=False),
|
||||
sa.Column('last_used_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
|
||||
pass
|
||||
@@ -63,19 +60,19 @@ def upgrade() -> None:
|
||||
|
||||
def downgrade() -> None:
|
||||
# Downgrade 'channel' table
|
||||
op.drop_column("channel", "is_private")
|
||||
op.drop_column("channel", "archived_at")
|
||||
op.drop_column("channel", "archived_by")
|
||||
op.drop_column("channel", "deleted_at")
|
||||
op.drop_column("channel", "deleted_by")
|
||||
op.drop_column("channel", "updated_by")
|
||||
op.drop_column('channel', 'is_private')
|
||||
op.drop_column('channel', 'archived_at')
|
||||
op.drop_column('channel', 'archived_by')
|
||||
op.drop_column('channel', 'deleted_at')
|
||||
op.drop_column('channel', 'deleted_by')
|
||||
op.drop_column('channel', 'updated_by')
|
||||
|
||||
# Downgrade 'channel_member' table
|
||||
op.drop_column("channel_member", "role")
|
||||
op.drop_column("channel_member", "invited_by")
|
||||
op.drop_column("channel_member", "invited_at")
|
||||
op.drop_column('channel_member', 'role')
|
||||
op.drop_column('channel_member', 'invited_by')
|
||||
op.drop_column('channel_member', 'invited_at')
|
||||
|
||||
# Drop 'channel_webhook' table
|
||||
op.drop_table("channel_webhook")
|
||||
op.drop_table('channel_webhook')
|
||||
|
||||
pass
|
||||
|
||||
@@ -9,38 +9,38 @@ Create Date: 2024-11-14 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "922e7a387820"
|
||||
down_revision = "4ace53fd72c8"
|
||||
revision = '922e7a387820'
|
||||
down_revision = '4ace53fd72c8'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
"group",
|
||||
sa.Column("id", sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=True),
|
||||
sa.Column("name", sa.Text(), nullable=True),
|
||||
sa.Column("description", sa.Text(), nullable=True),
|
||||
sa.Column("data", sa.JSON(), nullable=True),
|
||||
sa.Column("meta", sa.JSON(), nullable=True),
|
||||
sa.Column("permissions", sa.JSON(), nullable=True),
|
||||
sa.Column("user_ids", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
'group',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=True),
|
||||
sa.Column('name', sa.Text(), nullable=True),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('permissions', sa.JSON(), nullable=True),
|
||||
sa.Column('user_ids', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
|
||||
# Add 'access_control' column to 'model' table
|
||||
op.add_column(
|
||||
"model",
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
'model',
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# Add 'is_active' column to 'model' table
|
||||
op.add_column(
|
||||
"model",
|
||||
'model',
|
||||
sa.Column(
|
||||
"is_active",
|
||||
'is_active',
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.sql.expression.true(),
|
||||
@@ -49,37 +49,37 @@ def upgrade():
|
||||
|
||||
# Add 'access_control' column to 'knowledge' table
|
||||
op.add_column(
|
||||
"knowledge",
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
'knowledge',
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# Add 'access_control' column to 'prompt' table
|
||||
op.add_column(
|
||||
"prompt",
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
'prompt',
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# Add 'access_control' column to 'tools' table
|
||||
op.add_column(
|
||||
"tool",
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
'tool',
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("group")
|
||||
op.drop_table('group')
|
||||
|
||||
# Drop 'access_control' column from 'model' table
|
||||
op.drop_column("model", "access_control")
|
||||
op.drop_column('model', 'access_control')
|
||||
|
||||
# Drop 'is_active' column from 'model' table
|
||||
op.drop_column("model", "is_active")
|
||||
op.drop_column('model', 'is_active')
|
||||
|
||||
# Drop 'access_control' column from 'knowledge' table
|
||||
op.drop_column("knowledge", "access_control")
|
||||
op.drop_column('knowledge', 'access_control')
|
||||
|
||||
# Drop 'access_control' column from 'prompt' table
|
||||
op.drop_column("prompt", "access_control")
|
||||
op.drop_column('prompt', 'access_control')
|
||||
|
||||
# Drop 'access_control' column from 'tools' table
|
||||
op.drop_column("tool", "access_control")
|
||||
op.drop_column('tool', 'access_control')
|
||||
|
||||
@@ -9,25 +9,25 @@ Create Date: 2025-05-03 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "9f0c9cd09105"
|
||||
down_revision = "3781e22d8b01"
|
||||
revision = '9f0c9cd09105'
|
||||
down_revision = '3781e22d8b01'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
"note",
|
||||
sa.Column("id", sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=True),
|
||||
sa.Column("title", sa.Text(), nullable=True),
|
||||
sa.Column("data", sa.JSON(), nullable=True),
|
||||
sa.Column("meta", sa.JSON(), nullable=True),
|
||||
sa.Column("access_control", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=True),
|
||||
'note',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True, unique=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=True),
|
||||
sa.Column('title', sa.Text(), nullable=True),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('access_control', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("note")
|
||||
op.drop_table('note')
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
"""Add skill table
|
||||
|
||||
Revision ID: a1b2c3d4e5f6
|
||||
Revises: f1e2d3c4b5a6
|
||||
Create Date: 2026-02-11 09:30:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
from open_webui.migrations.util import get_existing_tables
|
||||
|
||||
revision: str = 'a1b2c3d4e5f6'
|
||||
down_revision: Union[str, None] = 'f1e2d3c4b5a6'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
existing_tables = set(get_existing_tables())
|
||||
|
||||
if 'skill' not in existing_tables:
|
||||
op.create_table(
|
||||
'skill',
|
||||
sa.Column('id', sa.String(), nullable=False, primary_key=True),
|
||||
sa.Column('user_id', sa.String(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False, unique=True),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('content', sa.Text(), nullable=False),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
op.create_index('idx_skill_user_id', 'skill', ['user_id'])
|
||||
op.create_index('idx_skill_updated_at', 'skill', ['updated_at'])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index('idx_skill_updated_at', table_name='skill')
|
||||
op.drop_index('idx_skill_user_id', table_name='skill')
|
||||
op.drop_table('skill')
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Add tasks and summary columns to chat table
|
||||
|
||||
Revision ID: a3dd5bedd151
|
||||
Revises: b2c3d4e5f6a7
|
||||
Create Date: 2026-03-29 22:15:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'a3dd5bedd151'
|
||||
down_revision: Union[str, None] = 'b2c3d4e5f6a7'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column('chat', sa.Column('tasks', sa.JSON(), nullable=True))
|
||||
op.add_column('chat', sa.Column('summary', sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column('chat', 'summary')
|
||||
op.drop_column('chat', 'tasks')
|
||||
+5
-5
@@ -12,8 +12,8 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "a5c220713937"
|
||||
down_revision: Union[str, None] = "38d63c18f30f"
|
||||
revision: str = 'a5c220713937'
|
||||
down_revision: Union[str, None] = '38d63c18f30f'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
@@ -21,14 +21,14 @@ depends_on: Union[str, Sequence[str], None] = None
|
||||
def upgrade() -> None:
|
||||
# Add 'reply_to_id' column to the 'message' table for replying to messages
|
||||
op.add_column(
|
||||
"message",
|
||||
sa.Column("reply_to_id", sa.Text(), nullable=True),
|
||||
'message',
|
||||
sa.Column('reply_to_id', sa.Text(), nullable=True),
|
||||
)
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Remove 'reply_to_id' column from the 'message' table
|
||||
op.drop_column("message", "reply_to_id")
|
||||
op.drop_column('message', 'reply_to_id')
|
||||
|
||||
pass
|
||||
|
||||
@@ -10,8 +10,8 @@ from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# Revision identifiers, used by Alembic.
|
||||
revision = "af906e964978"
|
||||
down_revision = "c29facfe716b"
|
||||
revision = 'af906e964978'
|
||||
down_revision = 'c29facfe716b'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
@@ -19,33 +19,23 @@ depends_on = None
|
||||
def upgrade():
|
||||
# ### Create feedback table ###
|
||||
op.create_table(
|
||||
"feedback",
|
||||
'feedback',
|
||||
sa.Column('id', sa.Text(), primary_key=True), # Unique identifier for each feedback (TEXT type)
|
||||
sa.Column('user_id', sa.Text(), nullable=True), # ID of the user providing the feedback (TEXT type)
|
||||
sa.Column('version', sa.BigInteger(), default=0), # Version of feedback (BIGINT type)
|
||||
sa.Column('type', sa.Text(), nullable=True), # Type of feedback (TEXT type)
|
||||
sa.Column('data', sa.JSON(), nullable=True), # Feedback data (JSON type)
|
||||
sa.Column('meta', sa.JSON(), nullable=True), # Metadata for feedback (JSON type)
|
||||
sa.Column('snapshot', sa.JSON(), nullable=True), # snapshot data for feedback (JSON type)
|
||||
sa.Column(
|
||||
"id", sa.Text(), primary_key=True
|
||||
), # Unique identifier for each feedback (TEXT type)
|
||||
sa.Column(
|
||||
"user_id", sa.Text(), nullable=True
|
||||
), # ID of the user providing the feedback (TEXT type)
|
||||
sa.Column(
|
||||
"version", sa.BigInteger(), default=0
|
||||
), # Version of feedback (BIGINT type)
|
||||
sa.Column("type", sa.Text(), nullable=True), # Type of feedback (TEXT type)
|
||||
sa.Column("data", sa.JSON(), nullable=True), # Feedback data (JSON type)
|
||||
sa.Column(
|
||||
"meta", sa.JSON(), nullable=True
|
||||
), # Metadata for feedback (JSON type)
|
||||
sa.Column(
|
||||
"snapshot", sa.JSON(), nullable=True
|
||||
), # snapshot data for feedback (JSON type)
|
||||
sa.Column(
|
||||
"created_at", sa.BigInteger(), nullable=False
|
||||
'created_at', sa.BigInteger(), nullable=False
|
||||
), # Feedback creation timestamp (BIGINT representing epoch)
|
||||
sa.Column(
|
||||
"updated_at", sa.BigInteger(), nullable=False
|
||||
'updated_at', sa.BigInteger(), nullable=False
|
||||
), # Feedback update timestamp (BIGINT representing epoch)
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
# ### Drop feedback table ###
|
||||
op.drop_table("feedback")
|
||||
op.drop_table('feedback')
|
||||
|
||||
@@ -17,8 +17,8 @@ import json
|
||||
import time
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b10670c03dd5"
|
||||
down_revision: Union[str, None] = "2f1211949ecc"
|
||||
revision: str = 'b10670c03dd5'
|
||||
down_revision: Union[str, None] = '2f1211949ecc'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
@@ -33,13 +33,11 @@ def _drop_sqlite_indexes_for_column(table_name, column_name, conn):
|
||||
for idx in indexes:
|
||||
index_name = idx[1] # index name
|
||||
# Get indexed columns
|
||||
idx_info = conn.execute(
|
||||
sa.text(f"PRAGMA index_info('{index_name}')")
|
||||
).fetchall()
|
||||
idx_info = conn.execute(sa.text(f"PRAGMA index_info('{index_name}')")).fetchall()
|
||||
|
||||
indexed_cols = [row[2] for row in idx_info] # col names
|
||||
if column_name in indexed_cols:
|
||||
conn.execute(sa.text(f"DROP INDEX IF EXISTS {index_name}"))
|
||||
conn.execute(sa.text(f'DROP INDEX IF EXISTS {index_name}'))
|
||||
|
||||
|
||||
def _convert_column_to_json(table: str, column: str):
|
||||
@@ -47,9 +45,9 @@ def _convert_column_to_json(table: str, column: str):
|
||||
dialect = conn.dialect.name
|
||||
|
||||
# SQLite cannot ALTER COLUMN → must recreate column
|
||||
if dialect == "sqlite":
|
||||
if dialect == 'sqlite':
|
||||
# 1. Add temporary column
|
||||
op.add_column(table, sa.Column(f"{column}_json", sa.JSON(), nullable=True))
|
||||
op.add_column(table, sa.Column(f'{column}_json', sa.JSON(), nullable=True))
|
||||
|
||||
# 2. Load old data
|
||||
rows = conn.execute(sa.text(f'SELECT id, {column} FROM "{table}"')).fetchall()
|
||||
@@ -66,14 +64,14 @@ def _convert_column_to_json(table: str, column: str):
|
||||
|
||||
conn.execute(
|
||||
sa.text(f'UPDATE "{table}" SET {column}_json = :val WHERE id = :id'),
|
||||
{"val": json.dumps(parsed) if parsed else None, "id": uid},
|
||||
{'val': json.dumps(parsed) if parsed else None, 'id': uid},
|
||||
)
|
||||
|
||||
# 3. Drop old TEXT column
|
||||
op.drop_column(table, column)
|
||||
|
||||
# 4. Rename new JSON column → original name
|
||||
op.alter_column(table, f"{column}_json", new_column_name=column)
|
||||
op.alter_column(table, f'{column}_json', new_column_name=column)
|
||||
|
||||
else:
|
||||
# PostgreSQL supports direct CAST
|
||||
@@ -81,7 +79,7 @@ def _convert_column_to_json(table: str, column: str):
|
||||
table,
|
||||
column,
|
||||
type_=sa.JSON(),
|
||||
postgresql_using=f"{column}::json",
|
||||
postgresql_using=f'{column}::json',
|
||||
)
|
||||
|
||||
|
||||
@@ -89,163 +87,151 @@ def _convert_column_to_text(table: str, column: str):
|
||||
conn = op.get_bind()
|
||||
dialect = conn.dialect.name
|
||||
|
||||
if dialect == "sqlite":
|
||||
op.add_column(table, sa.Column(f"{column}_text", sa.Text(), nullable=True))
|
||||
if dialect == 'sqlite':
|
||||
op.add_column(table, sa.Column(f'{column}_text', sa.Text(), nullable=True))
|
||||
|
||||
rows = conn.execute(sa.text(f'SELECT id, {column} FROM "{table}"')).fetchall()
|
||||
|
||||
for uid, raw in rows:
|
||||
conn.execute(
|
||||
sa.text(f'UPDATE "{table}" SET {column}_text = :val WHERE id = :id'),
|
||||
{"val": json.dumps(raw) if raw else None, "id": uid},
|
||||
{'val': json.dumps(raw) if raw else None, 'id': uid},
|
||||
)
|
||||
|
||||
op.drop_column(table, column)
|
||||
op.alter_column(table, f"{column}_text", new_column_name=column)
|
||||
op.alter_column(table, f'{column}_text', new_column_name=column)
|
||||
|
||||
else:
|
||||
op.alter_column(
|
||||
table,
|
||||
column,
|
||||
type_=sa.Text(),
|
||||
postgresql_using=f"to_json({column})::text",
|
||||
postgresql_using=f'to_json({column})::text',
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column(
|
||||
"user", sa.Column("profile_banner_image_url", sa.Text(), nullable=True)
|
||||
)
|
||||
op.add_column("user", sa.Column("timezone", sa.String(), nullable=True))
|
||||
op.add_column('user', sa.Column('profile_banner_image_url', sa.Text(), nullable=True))
|
||||
op.add_column('user', sa.Column('timezone', sa.String(), nullable=True))
|
||||
|
||||
op.add_column("user", sa.Column("presence_state", sa.String(), nullable=True))
|
||||
op.add_column("user", sa.Column("status_emoji", sa.String(), nullable=True))
|
||||
op.add_column("user", sa.Column("status_message", sa.Text(), nullable=True))
|
||||
op.add_column(
|
||||
"user", sa.Column("status_expires_at", sa.BigInteger(), nullable=True)
|
||||
)
|
||||
op.add_column('user', sa.Column('presence_state', sa.String(), nullable=True))
|
||||
op.add_column('user', sa.Column('status_emoji', sa.String(), nullable=True))
|
||||
op.add_column('user', sa.Column('status_message', sa.Text(), nullable=True))
|
||||
op.add_column('user', sa.Column('status_expires_at', sa.BigInteger(), nullable=True))
|
||||
|
||||
op.add_column("user", sa.Column("oauth", sa.JSON(), nullable=True))
|
||||
op.add_column('user', sa.Column('oauth', sa.JSON(), nullable=True))
|
||||
|
||||
# Convert info (TEXT/JSONField) → JSON
|
||||
_convert_column_to_json("user", "info")
|
||||
_convert_column_to_json('user', 'info')
|
||||
# Convert settings (TEXT/JSONField) → JSON
|
||||
_convert_column_to_json("user", "settings")
|
||||
_convert_column_to_json('user', 'settings')
|
||||
|
||||
op.create_table(
|
||||
"api_key",
|
||||
sa.Column("id", sa.Text(), primary_key=True, unique=True),
|
||||
sa.Column("user_id", sa.Text(), sa.ForeignKey("user.id", ondelete="CASCADE")),
|
||||
sa.Column("key", sa.Text(), unique=True, nullable=False),
|
||||
sa.Column("data", sa.JSON(), nullable=True),
|
||||
sa.Column("expires_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("last_used_at", sa.BigInteger(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
'api_key',
|
||||
sa.Column('id', sa.Text(), primary_key=True, unique=True),
|
||||
sa.Column('user_id', sa.Text(), sa.ForeignKey('user.id', ondelete='CASCADE')),
|
||||
sa.Column('key', sa.Text(), unique=True, nullable=False),
|
||||
sa.Column('data', sa.JSON(), nullable=True),
|
||||
sa.Column('expires_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('last_used_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
|
||||
conn = op.get_bind()
|
||||
users = conn.execute(
|
||||
sa.text('SELECT id, oauth_sub FROM "user" WHERE oauth_sub IS NOT NULL')
|
||||
).fetchall()
|
||||
users = conn.execute(sa.text('SELECT id, oauth_sub FROM "user" WHERE oauth_sub IS NOT NULL')).fetchall()
|
||||
|
||||
for uid, oauth_sub in users:
|
||||
if oauth_sub:
|
||||
# Example formats supported:
|
||||
# provider@sub
|
||||
# plain sub (stored as {"oidc": {"sub": sub}})
|
||||
if "@" in oauth_sub:
|
||||
provider, sub = oauth_sub.split("@", 1)
|
||||
if '@' in oauth_sub:
|
||||
provider, sub = oauth_sub.split('@', 1)
|
||||
else:
|
||||
provider, sub = "oidc", oauth_sub
|
||||
provider, sub = 'oidc', oauth_sub
|
||||
|
||||
oauth_json = json.dumps({provider: {"sub": sub}})
|
||||
oauth_json = json.dumps({provider: {'sub': sub}})
|
||||
conn.execute(
|
||||
sa.text('UPDATE "user" SET oauth = :oauth WHERE id = :id'),
|
||||
{"oauth": oauth_json, "id": uid},
|
||||
{'oauth': oauth_json, 'id': uid},
|
||||
)
|
||||
|
||||
users_with_keys = conn.execute(
|
||||
sa.text('SELECT id, api_key FROM "user" WHERE api_key IS NOT NULL')
|
||||
).fetchall()
|
||||
users_with_keys = conn.execute(sa.text('SELECT id, api_key FROM "user" WHERE api_key IS NOT NULL')).fetchall()
|
||||
now = int(time.time())
|
||||
|
||||
for uid, api_key in users_with_keys:
|
||||
if api_key:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
sa.text("""
|
||||
INSERT INTO api_key (id, user_id, key, created_at, updated_at)
|
||||
VALUES (:id, :user_id, :key, :created_at, :updated_at)
|
||||
"""
|
||||
),
|
||||
"""),
|
||||
{
|
||||
"id": f"key_{uid}",
|
||||
"user_id": uid,
|
||||
"key": api_key,
|
||||
"created_at": now,
|
||||
"updated_at": now,
|
||||
'id': f'key_{uid}',
|
||||
'user_id': uid,
|
||||
'key': api_key,
|
||||
'created_at': now,
|
||||
'updated_at': now,
|
||||
},
|
||||
)
|
||||
|
||||
if conn.dialect.name == "sqlite":
|
||||
_drop_sqlite_indexes_for_column("user", "api_key", conn)
|
||||
_drop_sqlite_indexes_for_column("user", "oauth_sub", conn)
|
||||
if conn.dialect.name == 'sqlite':
|
||||
_drop_sqlite_indexes_for_column('user', 'api_key', conn)
|
||||
_drop_sqlite_indexes_for_column('user', 'oauth_sub', conn)
|
||||
|
||||
with op.batch_alter_table("user") as batch_op:
|
||||
batch_op.drop_column("api_key")
|
||||
batch_op.drop_column("oauth_sub")
|
||||
with op.batch_alter_table('user') as batch_op:
|
||||
batch_op.drop_column('api_key')
|
||||
batch_op.drop_column('oauth_sub')
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# --- 1. Restore old oauth_sub column ---
|
||||
op.add_column("user", sa.Column("oauth_sub", sa.Text(), nullable=True))
|
||||
op.add_column('user', sa.Column('oauth_sub', sa.Text(), nullable=True))
|
||||
|
||||
conn = op.get_bind()
|
||||
users = conn.execute(
|
||||
sa.text('SELECT id, oauth FROM "user" WHERE oauth IS NOT NULL')
|
||||
).fetchall()
|
||||
users = conn.execute(sa.text('SELECT id, oauth FROM "user" WHERE oauth IS NOT NULL')).fetchall()
|
||||
|
||||
for uid, oauth in users:
|
||||
try:
|
||||
data = json.loads(oauth)
|
||||
provider = list(data.keys())[0]
|
||||
sub = data[provider].get("sub")
|
||||
oauth_sub = f"{provider}@{sub}"
|
||||
sub = data[provider].get('sub')
|
||||
oauth_sub = f'{provider}@{sub}'
|
||||
except Exception:
|
||||
oauth_sub = None
|
||||
|
||||
conn.execute(
|
||||
sa.text('UPDATE "user" SET oauth_sub = :oauth_sub WHERE id = :id'),
|
||||
{"oauth_sub": oauth_sub, "id": uid},
|
||||
{'oauth_sub': oauth_sub, 'id': uid},
|
||||
)
|
||||
|
||||
op.drop_column("user", "oauth")
|
||||
op.drop_column('user', 'oauth')
|
||||
|
||||
# --- 2. Restore api_key field ---
|
||||
op.add_column("user", sa.Column("api_key", sa.String(), nullable=True))
|
||||
op.add_column('user', sa.Column('api_key', sa.String(), nullable=True))
|
||||
|
||||
# Restore values from api_key
|
||||
keys = conn.execute(sa.text("SELECT user_id, key FROM api_key")).fetchall()
|
||||
keys = conn.execute(sa.text('SELECT user_id, key FROM api_key')).fetchall()
|
||||
for uid, key in keys:
|
||||
conn.execute(
|
||||
sa.text('UPDATE "user" SET api_key = :key WHERE id = :id'),
|
||||
{"key": key, "id": uid},
|
||||
{'key': key, 'id': uid},
|
||||
)
|
||||
|
||||
# Drop new table
|
||||
op.drop_table("api_key")
|
||||
op.drop_table('api_key')
|
||||
|
||||
with op.batch_alter_table("user") as batch_op:
|
||||
batch_op.drop_column("profile_banner_image_url")
|
||||
batch_op.drop_column("timezone")
|
||||
with op.batch_alter_table('user') as batch_op:
|
||||
batch_op.drop_column('profile_banner_image_url')
|
||||
batch_op.drop_column('timezone')
|
||||
|
||||
batch_op.drop_column("presence_state")
|
||||
batch_op.drop_column("status_emoji")
|
||||
batch_op.drop_column("status_message")
|
||||
batch_op.drop_column("status_expires_at")
|
||||
batch_op.drop_column('presence_state')
|
||||
batch_op.drop_column('status_emoji')
|
||||
batch_op.drop_column('status_message')
|
||||
batch_op.drop_column('status_expires_at')
|
||||
|
||||
# Convert info (JSON) → TEXT
|
||||
_convert_column_to_text("user", "info")
|
||||
_convert_column_to_text('user', 'info')
|
||||
# Convert settings (JSON) → TEXT
|
||||
_convert_column_to_text("user", "settings")
|
||||
_convert_column_to_text('user', 'settings')
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
"""add scim column to user table
|
||||
|
||||
Revision ID: b2c3d4e5f6a7
|
||||
Revises: a1b2c3d4e5f6
|
||||
Create Date: 2026-02-13 14:19:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = 'b2c3d4e5f6a7'
|
||||
down_revision: Union[str, None] = 'a1b2c3d4e5f6'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column('user', sa.Column('scim', sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column('user', 'scim')
|
||||
@@ -0,0 +1,27 @@
|
||||
"""add last_read_at to chat
|
||||
|
||||
Revision ID: b7c8d9e0f1a2
|
||||
Revises: d4e5f6a7b8c9
|
||||
Create Date: 2026-04-01 04:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'b7c8d9e0f1a2'
|
||||
down_revision = 'd4e5f6a7b8c9'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.add_column('chat', sa.Column('last_read_at', sa.BigInteger(), nullable=True))
|
||||
# Set existing chats to be marked as read
|
||||
op.execute('UPDATE chat SET last_read_at = updated_at')
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_column('chat', 'last_read_at')
|
||||
@@ -12,21 +12,21 @@ import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c0fbf31ca0db"
|
||||
down_revision: Union[str, None] = "ca81bd47c050"
|
||||
revision: str = 'c0fbf31ca0db'
|
||||
down_revision: Union[str, None] = 'ca81bd47c050'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.add_column("file", sa.Column("hash", sa.Text(), nullable=True))
|
||||
op.add_column("file", sa.Column("data", sa.JSON(), nullable=True))
|
||||
op.add_column("file", sa.Column("updated_at", sa.BigInteger(), nullable=True))
|
||||
op.add_column('file', sa.Column('hash', sa.Text(), nullable=True))
|
||||
op.add_column('file', sa.Column('data', sa.JSON(), nullable=True))
|
||||
op.add_column('file', sa.Column('updated_at', sa.BigInteger(), nullable=True))
|
||||
|
||||
|
||||
def downgrade():
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_column("file", "updated_at")
|
||||
op.drop_column("file", "data")
|
||||
op.drop_column("file", "hash")
|
||||
op.drop_column('file', 'updated_at')
|
||||
op.drop_column('file', 'data')
|
||||
op.drop_column('file', 'hash')
|
||||
|
||||
@@ -0,0 +1,164 @@
|
||||
"""Add shared_chat table and migrate existing shares
|
||||
|
||||
Revision ID: c1d2e3f4a5b6
|
||||
Revises: e1f2a3b4c5d6
|
||||
Create Date: 2026-04-16 23:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = 'c1d2e3f4a5b6'
|
||||
down_revision = 'e1f2a3b4c5d6'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
# Lightweight table references for data migration (no ORM models needed)
|
||||
chat_t = sa.table(
|
||||
'chat',
|
||||
sa.column('id', sa.Text),
|
||||
sa.column('user_id', sa.Text),
|
||||
sa.column('title', sa.Text),
|
||||
sa.column('chat', sa.JSON),
|
||||
sa.column('share_id', sa.Text),
|
||||
sa.column('created_at', sa.BigInteger),
|
||||
sa.column('updated_at', sa.BigInteger),
|
||||
sa.column('archived', sa.Boolean),
|
||||
sa.column('meta', sa.JSON),
|
||||
)
|
||||
|
||||
shared_chat_t = sa.table(
|
||||
'shared_chat',
|
||||
sa.column('id', sa.Text),
|
||||
sa.column('chat_id', sa.Text),
|
||||
sa.column('user_id', sa.Text),
|
||||
sa.column('title', sa.Text),
|
||||
sa.column('chat', sa.JSON),
|
||||
sa.column('created_at', sa.BigInteger),
|
||||
sa.column('updated_at', sa.BigInteger),
|
||||
)
|
||||
|
||||
chat_message_t = sa.table(
|
||||
'chat_message',
|
||||
sa.column('chat_id', sa.Text),
|
||||
)
|
||||
|
||||
access_grant_t = sa.table(
|
||||
'access_grant',
|
||||
sa.column('id', sa.Text),
|
||||
sa.column('resource_type', sa.Text),
|
||||
sa.column('resource_id', sa.Text),
|
||||
sa.column('principal_type', sa.Text),
|
||||
sa.column('principal_id', sa.Text),
|
||||
sa.column('permission', sa.Text),
|
||||
sa.column('created_at', sa.BigInteger),
|
||||
)
|
||||
|
||||
|
||||
def upgrade():
|
||||
conn = op.get_bind()
|
||||
|
||||
# 1. Create shared_chat table
|
||||
op.create_table(
|
||||
'shared_chat',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('chat_id', sa.Text(), sa.ForeignKey('chat.id', ondelete='CASCADE'), nullable=False),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('title', sa.Text(), nullable=True),
|
||||
sa.Column('chat', sa.JSON(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=True),
|
||||
)
|
||||
|
||||
# 2. Migrate existing shared-* rows
|
||||
shared_rows = conn.execute(
|
||||
sa.select(
|
||||
chat_t.c.id,
|
||||
chat_t.c.user_id,
|
||||
chat_t.c.title,
|
||||
chat_t.c.chat,
|
||||
chat_t.c.created_at,
|
||||
chat_t.c.updated_at,
|
||||
).where(chat_t.c.user_id.like('shared-%'))
|
||||
).fetchall()
|
||||
|
||||
for row in shared_rows:
|
||||
share_token = row.id
|
||||
original_chat_id = row.user_id.replace('shared-', '', 1)
|
||||
|
||||
# Verify original chat still exists
|
||||
original = conn.execute(sa.select(chat_t.c.user_id).where(chat_t.c.id == original_chat_id)).fetchone()
|
||||
|
||||
if not original:
|
||||
continue
|
||||
|
||||
# Insert snapshot into shared_chat
|
||||
conn.execute(
|
||||
shared_chat_t.insert().values(
|
||||
id=share_token,
|
||||
chat_id=original_chat_id,
|
||||
user_id=original.user_id,
|
||||
title=row.title,
|
||||
chat=row.chat,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
)
|
||||
|
||||
# Create user:*:read grant for backward compat
|
||||
conn.execute(
|
||||
access_grant_t.insert().values(
|
||||
id=str(uuid.uuid4()),
|
||||
resource_type='shared_chat',
|
||||
resource_id=original_chat_id,
|
||||
principal_type='user',
|
||||
principal_id='*',
|
||||
permission='read',
|
||||
created_at=row.created_at or int(time.time()),
|
||||
)
|
||||
)
|
||||
|
||||
# 3. Clean up old phantom rows
|
||||
conn.execute(
|
||||
chat_message_t.delete().where(
|
||||
chat_message_t.c.chat_id.in_(sa.select(chat_t.c.id).where(chat_t.c.user_id.like('shared-%')))
|
||||
)
|
||||
)
|
||||
conn.execute(chat_t.delete().where(chat_t.c.user_id.like('shared-%')))
|
||||
|
||||
|
||||
def downgrade():
|
||||
conn = op.get_bind()
|
||||
|
||||
shared_rows = conn.execute(
|
||||
sa.select(
|
||||
shared_chat_t.c.id,
|
||||
shared_chat_t.c.chat_id,
|
||||
shared_chat_t.c.user_id,
|
||||
shared_chat_t.c.title,
|
||||
shared_chat_t.c.chat,
|
||||
shared_chat_t.c.created_at,
|
||||
shared_chat_t.c.updated_at,
|
||||
)
|
||||
).fetchall()
|
||||
|
||||
for row in shared_rows:
|
||||
conn.execute(
|
||||
chat_t.insert().values(
|
||||
id=row.id,
|
||||
user_id=f'shared-{row.chat_id}',
|
||||
title=row.title,
|
||||
chat=row.chat,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
archived=False,
|
||||
meta={},
|
||||
)
|
||||
)
|
||||
|
||||
conn.execute(access_grant_t.delete().where(access_grant_t.c.resource_type == 'shared_chat'))
|
||||
op.drop_table('shared_chat')
|
||||
@@ -12,36 +12,33 @@ import json
|
||||
from sqlalchemy.sql import table, column
|
||||
from sqlalchemy import String, Text, JSON, and_
|
||||
|
||||
|
||||
revision = "c29facfe716b"
|
||||
down_revision = "c69f45358db4"
|
||||
revision = 'c29facfe716b'
|
||||
down_revision = 'c69f45358db4'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
# 1. Add the `path` column to the "file" table.
|
||||
op.add_column("file", sa.Column("path", sa.Text(), nullable=True))
|
||||
op.add_column('file', sa.Column('path', sa.Text(), nullable=True))
|
||||
|
||||
# 2. Convert the `meta` column from Text/JSONField to `JSON()`
|
||||
# Use Alembic's default batch_op for dialect compatibility.
|
||||
with op.batch_alter_table("file", schema=None) as batch_op:
|
||||
with op.batch_alter_table('file', schema=None) as batch_op:
|
||||
batch_op.alter_column(
|
||||
"meta",
|
||||
'meta',
|
||||
type_=sa.JSON(),
|
||||
existing_type=sa.Text(),
|
||||
existing_nullable=True,
|
||||
nullable=True,
|
||||
postgresql_using="meta::json",
|
||||
postgresql_using='meta::json',
|
||||
)
|
||||
|
||||
# 3. Migrate legacy data from `meta` JSONField
|
||||
# Fetch and process `meta` data from the table, add values to the new `path` column as necessary.
|
||||
# We will use SQLAlchemy core bindings to ensure safety across different databases.
|
||||
|
||||
file_table = table(
|
||||
"file", column("id", String), column("meta", JSON), column("path", Text)
|
||||
)
|
||||
file_table = table('file', column('id', String), column('meta', JSON), column('path', Text))
|
||||
|
||||
# Create connection to the database
|
||||
connection = op.get_bind()
|
||||
@@ -56,24 +53,18 @@ def upgrade():
|
||||
|
||||
# Iterate over each row to extract and update the `path` from `meta` column
|
||||
for row in results:
|
||||
if "path" in row.meta:
|
||||
if 'path' in row.meta:
|
||||
# Extract the `path` field from the `meta` JSON
|
||||
path = row.meta.get("path")
|
||||
path = row.meta.get('path')
|
||||
|
||||
# Update the `file` table with the new `path` value
|
||||
connection.execute(
|
||||
file_table.update()
|
||||
.where(file_table.c.id == row.id)
|
||||
.values({"path": path})
|
||||
)
|
||||
connection.execute(file_table.update().where(file_table.c.id == row.id).values({'path': path}))
|
||||
|
||||
|
||||
def downgrade():
|
||||
# 1. Remove the `path` column
|
||||
op.drop_column("file", "path")
|
||||
op.drop_column('file', 'path')
|
||||
|
||||
# 2. Revert the `meta` column back to Text/JSONField
|
||||
with op.batch_alter_table("file", schema=None) as batch_op:
|
||||
batch_op.alter_column(
|
||||
"meta", type_=sa.Text(), existing_type=sa.JSON(), existing_nullable=True
|
||||
)
|
||||
with op.batch_alter_table('file', schema=None) as batch_op:
|
||||
batch_op.alter_column('meta', type_=sa.Text(), existing_type=sa.JSON(), existing_nullable=True)
|
||||
|
||||
@@ -11,47 +11,44 @@ from typing import Sequence, Union
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c440947495f3"
|
||||
down_revision: Union[str, None] = "81cc2ce44d79"
|
||||
revision: str = 'c440947495f3'
|
||||
down_revision: Union[str, None] = '81cc2ce44d79'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"chat_file",
|
||||
sa.Column("id", sa.Text(), primary_key=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=False),
|
||||
'chat_file',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column(
|
||||
"chat_id",
|
||||
'chat_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("chat.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('chat.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"file_id",
|
||||
'file_id',
|
||||
sa.Text(),
|
||||
sa.ForeignKey("file.id", ondelete="CASCADE"),
|
||||
sa.ForeignKey('file.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("message_id", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column("updated_at", sa.BigInteger(), nullable=False),
|
||||
sa.Column('message_id', sa.Text(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
# indexes
|
||||
sa.Index("ix_chat_file_chat_id", "chat_id"),
|
||||
sa.Index("ix_chat_file_file_id", "file_id"),
|
||||
sa.Index("ix_chat_file_message_id", "message_id"),
|
||||
sa.Index("ix_chat_file_user_id", "user_id"),
|
||||
sa.Index('ix_chat_file_chat_id', 'chat_id'),
|
||||
sa.Index('ix_chat_file_file_id', 'file_id'),
|
||||
sa.Index('ix_chat_file_message_id', 'message_id'),
|
||||
sa.Index('ix_chat_file_user_id', 'user_id'),
|
||||
# unique constraints
|
||||
sa.UniqueConstraint(
|
||||
"chat_id", "file_id", name="uq_chat_file_chat_file"
|
||||
), # prevent duplicate entries
|
||||
sa.UniqueConstraint('chat_id', 'file_id', name='uq_chat_file_chat_file'), # prevent duplicate entries
|
||||
)
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("chat_file")
|
||||
op.drop_table('chat_file')
|
||||
pass
|
||||
|
||||
@@ -9,42 +9,40 @@ Create Date: 2024-10-16 02:02:35.241684
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "c69f45358db4"
|
||||
down_revision = "3ab32c4b8f59"
|
||||
revision = 'c69f45358db4'
|
||||
down_revision = '3ab32c4b8f59'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
"folder",
|
||||
sa.Column("id", sa.Text(), nullable=False),
|
||||
sa.Column("parent_id", sa.Text(), nullable=True),
|
||||
sa.Column("user_id", sa.Text(), nullable=False),
|
||||
sa.Column("name", sa.Text(), nullable=False),
|
||||
sa.Column("items", sa.JSON(), nullable=True),
|
||||
sa.Column("meta", sa.JSON(), nullable=True),
|
||||
sa.Column("is_expanded", sa.Boolean(), default=False, nullable=False),
|
||||
'folder',
|
||||
sa.Column('id', sa.Text(), nullable=False),
|
||||
sa.Column('parent_id', sa.Text(), nullable=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False),
|
||||
sa.Column('items', sa.JSON(), nullable=True),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('is_expanded', sa.Boolean(), default=False, nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column(
|
||||
"created_at", sa.DateTime(), server_default=sa.func.now(), nullable=False
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
'updated_at',
|
||||
sa.DateTime(),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
onupdate=sa.func.now(),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("id", "user_id"),
|
||||
sa.PrimaryKeyConstraint('id', 'user_id'),
|
||||
)
|
||||
|
||||
op.add_column(
|
||||
"chat",
|
||||
sa.Column("folder_id", sa.Text(), nullable=True),
|
||||
'chat',
|
||||
sa.Column('folder_id', sa.Text(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_column("chat", "folder_id")
|
||||
op.drop_column('chat', 'folder_id')
|
||||
|
||||
op.drop_table("folder")
|
||||
op.drop_table('folder')
|
||||
|
||||
@@ -12,23 +12,21 @@ import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "ca81bd47c050"
|
||||
down_revision: Union[str, None] = "7e5b5dc7342b"
|
||||
revision: str = 'ca81bd47c050'
|
||||
down_revision: Union[str, None] = '7e5b5dc7342b'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
"config",
|
||||
sa.Column("id", sa.Integer, primary_key=True),
|
||||
sa.Column("data", sa.JSON(), nullable=False),
|
||||
sa.Column("version", sa.Integer, nullable=False),
|
||||
'config',
|
||||
sa.Column('id', sa.Integer, primary_key=True),
|
||||
sa.Column('data', sa.JSON(), nullable=False),
|
||||
sa.Column('version', sa.Integer, nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
|
||||
sa.Column(
|
||||
"created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
'updated_at',
|
||||
sa.DateTime(),
|
||||
nullable=True,
|
||||
server_default=sa.func.now(),
|
||||
@@ -38,4 +36,4 @@ def upgrade():
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_table("config")
|
||||
op.drop_table('config')
|
||||
|
||||
@@ -9,15 +9,15 @@ Create Date: 2025-07-13 03:00:00.000000
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "d31026856c01"
|
||||
down_revision = "9f0c9cd09105"
|
||||
revision = 'd31026856c01'
|
||||
down_revision = '9f0c9cd09105'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.add_column("folder", sa.Column("data", sa.JSON(), nullable=True))
|
||||
op.add_column('folder', sa.Column('data', sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_column("folder", "data")
|
||||
op.drop_column('folder', 'data')
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
"""add automation tables
|
||||
|
||||
Revision ID: d4e5f6a7b8c9
|
||||
Revises: f1e2d3c4b5a6
|
||||
Create Date: 2026-03-30
|
||||
"""
|
||||
|
||||
from typing import Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision: str = 'd4e5f6a7b8c9'
|
||||
down_revision: Union[str, None] = 'a3dd5bedd151'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.create_table(
|
||||
'automation',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('user_id', sa.Text(), nullable=False),
|
||||
sa.Column('name', sa.Text(), nullable=False),
|
||||
sa.Column('data', sa.JSON(), nullable=False),
|
||||
sa.Column('meta', sa.JSON(), nullable=True),
|
||||
sa.Column('is_active', sa.Boolean(), nullable=False, default=True),
|
||||
sa.Column('last_run_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('next_run_at', sa.BigInteger(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.Column('updated_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
op.create_index('ix_automation_next_run', 'automation', ['next_run_at'])
|
||||
|
||||
op.create_table(
|
||||
'automation_run',
|
||||
sa.Column('id', sa.Text(), primary_key=True),
|
||||
sa.Column('automation_id', sa.Text(), nullable=False),
|
||||
sa.Column('chat_id', sa.Text(), nullable=True),
|
||||
sa.Column('status', sa.Text(), nullable=False),
|
||||
sa.Column('error', sa.Text(), nullable=True),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
)
|
||||
op.create_index(
|
||||
'ix_automation_run_automation_id',
|
||||
'automation_run',
|
||||
['automation_id'],
|
||||
)
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_index('ix_automation_run_automation_id')
|
||||
op.drop_table('automation_run')
|
||||
op.drop_index('ix_automation_next_run')
|
||||
op.drop_table('automation')
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Add is_pinned to note table
|
||||
|
||||
Revision ID: e1f2a3b4c5d6
|
||||
Revises: b7c8d9e0f1a2
|
||||
Create Date: 2026-04-14 22:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = 'e1f2a3b4c5d6'
|
||||
down_revision = 'b7c8d9e0f1a2'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade():
|
||||
op.add_column('note', sa.Column('is_pinned', sa.Boolean(), nullable=True))
|
||||
|
||||
|
||||
def downgrade():
|
||||
op.drop_column('note', 'is_pinned')
|
||||
@@ -0,0 +1,344 @@
|
||||
"""Add access_grant table
|
||||
|
||||
Revision ID: f1e2d3c4b5a6
|
||||
Revises: 8452d01d26d7
|
||||
Create Date: 2026-02-05 10:00:00.000000
|
||||
|
||||
Migrates from JSON access_control columns to normalized access_grant table.
|
||||
Access control semantics:
|
||||
- NULL: Public access (all users can read) -> insert user:* for read
|
||||
- {}: Private/owner-only (no grants) -> insert nothing
|
||||
- {read: {...}, write: {...}}: Custom permissions -> insert specific grants
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
from open_webui.migrations.util import get_existing_tables
|
||||
|
||||
revision: str = 'f1e2d3c4b5a6'
|
||||
down_revision: Union[str, None] = '8452d01d26d7'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
existing_tables = set(get_existing_tables())
|
||||
|
||||
# Create access_grant table
|
||||
if 'access_grant' not in existing_tables:
|
||||
op.create_table(
|
||||
'access_grant',
|
||||
sa.Column('id', sa.Text(), nullable=False, primary_key=True),
|
||||
sa.Column('resource_type', sa.Text(), nullable=False),
|
||||
sa.Column('resource_id', sa.Text(), nullable=False),
|
||||
sa.Column('principal_type', sa.Text(), nullable=False),
|
||||
sa.Column('principal_id', sa.Text(), nullable=False),
|
||||
sa.Column('permission', sa.Text(), nullable=False),
|
||||
sa.Column('created_at', sa.BigInteger(), nullable=False),
|
||||
sa.UniqueConstraint(
|
||||
'resource_type',
|
||||
'resource_id',
|
||||
'principal_type',
|
||||
'principal_id',
|
||||
'permission',
|
||||
name='uq_access_grant_grant',
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
'idx_access_grant_resource',
|
||||
'access_grant',
|
||||
['resource_type', 'resource_id'],
|
||||
)
|
||||
op.create_index(
|
||||
'idx_access_grant_principal',
|
||||
'access_grant',
|
||||
['principal_type', 'principal_id'],
|
||||
)
|
||||
|
||||
# Backfill existing access_control JSON data
|
||||
conn = op.get_bind()
|
||||
|
||||
# Tables with access_control JSON columns: (table_name, resource_type)
|
||||
resource_tables = [
|
||||
('knowledge', 'knowledge'),
|
||||
('prompt', 'prompt'),
|
||||
('tool', 'tool'),
|
||||
('model', 'model'),
|
||||
('note', 'note'),
|
||||
('channel', 'channel'),
|
||||
('file', 'file'),
|
||||
]
|
||||
|
||||
now = int(time.time())
|
||||
inserted = set()
|
||||
|
||||
for table_name, resource_type in resource_tables:
|
||||
if table_name not in existing_tables:
|
||||
continue
|
||||
|
||||
# Query all rows
|
||||
try:
|
||||
result = conn.execute(sa.text(f'SELECT id, access_control FROM "{table_name}"'))
|
||||
rows = result.fetchall()
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
for row in rows:
|
||||
resource_id = row[0]
|
||||
access_control_json = row[1]
|
||||
|
||||
# Handle NULL or JSON "null" = public access (user:* for read)
|
||||
# Could be Python None (SQL NULL) or string "null" (JSON null)
|
||||
# EXCEPTION: files with NULL are PRIVATE (owner-only), not public
|
||||
is_null = (
|
||||
access_control_json is None
|
||||
or access_control_json == 'null'
|
||||
or (isinstance(access_control_json, str) and access_control_json.strip().lower() == 'null')
|
||||
)
|
||||
if is_null:
|
||||
# Files: NULL = private (no entry needed, owner has implicit access)
|
||||
# Other resources: NULL = public (insert user:* for read)
|
||||
if resource_type == 'file':
|
||||
continue # Private - no entry needed
|
||||
|
||||
key = (resource_type, resource_id, 'user', '*', 'read')
|
||||
if key not in inserted:
|
||||
try:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO access_grant (id, resource_type, resource_id, principal_type, principal_id, permission, created_at)
|
||||
VALUES (:id, :resource_type, :resource_id, :principal_type, :principal_id, :permission, :created_at)
|
||||
"""),
|
||||
{
|
||||
'id': str(uuid.uuid4()),
|
||||
'resource_type': resource_type,
|
||||
'resource_id': resource_id,
|
||||
'principal_type': 'user',
|
||||
'principal_id': '*',
|
||||
'permission': 'read',
|
||||
'created_at': now,
|
||||
},
|
||||
)
|
||||
inserted.add(key)
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
|
||||
# Handle JSON parsing
|
||||
if isinstance(access_control_json, str):
|
||||
import json
|
||||
|
||||
try:
|
||||
access_control_json = json.loads(access_control_json)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
# Handle {} = private/owner-only - NO entries needed
|
||||
# Owner access is implicit, no grants to store
|
||||
if not access_control_json or not isinstance(access_control_json, dict):
|
||||
continue
|
||||
|
||||
# Check if it's effectively empty (no read/write keys with content)
|
||||
read_data = access_control_json.get('read', {})
|
||||
write_data = access_control_json.get('write', {})
|
||||
|
||||
has_read_grants = read_data.get('group_ids', []) or read_data.get('user_ids', [])
|
||||
has_write_grants = write_data.get('group_ids', []) or write_data.get('user_ids', [])
|
||||
|
||||
if not has_read_grants and not has_write_grants:
|
||||
# Empty permissions = private, no grants needed
|
||||
continue
|
||||
|
||||
# Extract permissions and insert into access_grant table
|
||||
for permission in ['read', 'write']:
|
||||
perm_data = access_control_json.get(permission, {})
|
||||
if not perm_data:
|
||||
continue
|
||||
|
||||
for group_id in perm_data.get('group_ids', []):
|
||||
key = (resource_type, resource_id, 'group', group_id, permission)
|
||||
if key in inserted:
|
||||
continue
|
||||
try:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO access_grant (id, resource_type, resource_id, principal_type, principal_id, permission, created_at)
|
||||
VALUES (:id, :resource_type, :resource_id, :principal_type, :principal_id, :permission, :created_at)
|
||||
"""),
|
||||
{
|
||||
'id': str(uuid.uuid4()),
|
||||
'resource_type': resource_type,
|
||||
'resource_id': resource_id,
|
||||
'principal_type': 'group',
|
||||
'principal_id': group_id,
|
||||
'permission': permission,
|
||||
'created_at': now,
|
||||
},
|
||||
)
|
||||
inserted.add(key)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for user_id in perm_data.get('user_ids', []):
|
||||
key = (resource_type, resource_id, 'user', user_id, permission)
|
||||
if key in inserted:
|
||||
continue
|
||||
try:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO access_grant (id, resource_type, resource_id, principal_type, principal_id, permission, created_at)
|
||||
VALUES (:id, :resource_type, :resource_id, :principal_type, :principal_id, :permission, :created_at)
|
||||
"""),
|
||||
{
|
||||
'id': str(uuid.uuid4()),
|
||||
'resource_type': resource_type,
|
||||
'resource_id': resource_id,
|
||||
'principal_type': 'user',
|
||||
'principal_id': user_id,
|
||||
'permission': permission,
|
||||
'created_at': now,
|
||||
},
|
||||
)
|
||||
inserted.add(key)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Drop access_control columns from resource tables
|
||||
for table_name, _ in resource_tables:
|
||||
if table_name not in existing_tables:
|
||||
continue
|
||||
try:
|
||||
with op.batch_alter_table(table_name) as batch:
|
||||
batch.drop_column('access_control')
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
import json
|
||||
|
||||
conn = op.get_bind()
|
||||
|
||||
# Resource tables mapping: (table_name, resource_type)
|
||||
resource_tables = [
|
||||
('knowledge', 'knowledge'),
|
||||
('prompt', 'prompt'),
|
||||
('tool', 'tool'),
|
||||
('model', 'model'),
|
||||
('note', 'note'),
|
||||
('channel', 'channel'),
|
||||
('file', 'file'),
|
||||
]
|
||||
|
||||
# Step 1: Re-add access_control columns to resource tables
|
||||
for table_name, _ in resource_tables:
|
||||
try:
|
||||
with op.batch_alter_table(table_name) as batch:
|
||||
batch.add_column(sa.Column('access_control', sa.JSON(), nullable=True))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Step 2: Query access_grant table and reconstruct JSON for each resource
|
||||
for table_name, resource_type in resource_tables:
|
||||
try:
|
||||
# Get all grants for this resource type
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
SELECT resource_id, principal_type, principal_id, permission
|
||||
FROM access_grant
|
||||
WHERE resource_type = :resource_type
|
||||
"""),
|
||||
{'resource_type': resource_type},
|
||||
)
|
||||
rows = result.fetchall()
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
# Group by resource_id and reconstruct JSON structure
|
||||
resource_grants = {}
|
||||
for row in rows:
|
||||
resource_id = row[0]
|
||||
principal_type = row[1]
|
||||
principal_id = row[2]
|
||||
permission = row[3]
|
||||
|
||||
if resource_id not in resource_grants:
|
||||
resource_grants[resource_id] = {
|
||||
'is_public': False,
|
||||
'read': {'group_ids': [], 'user_ids': []},
|
||||
'write': {'group_ids': [], 'user_ids': []},
|
||||
}
|
||||
|
||||
# Handle public access (user:* for read)
|
||||
if principal_type == 'user' and principal_id == '*' and permission == 'read':
|
||||
resource_grants[resource_id]['is_public'] = True
|
||||
continue
|
||||
|
||||
# Add to appropriate list
|
||||
if permission in ['read', 'write']:
|
||||
if principal_type == 'group':
|
||||
if principal_id not in resource_grants[resource_id][permission]['group_ids']:
|
||||
resource_grants[resource_id][permission]['group_ids'].append(principal_id)
|
||||
elif principal_type == 'user':
|
||||
if principal_id not in resource_grants[resource_id][permission]['user_ids']:
|
||||
resource_grants[resource_id][permission]['user_ids'].append(principal_id)
|
||||
|
||||
# Step 3: Update each resource with reconstructed JSON
|
||||
for resource_id, grants in resource_grants.items():
|
||||
if grants['is_public']:
|
||||
# Public = NULL
|
||||
access_control_value = None
|
||||
elif (
|
||||
not grants['read']['group_ids']
|
||||
and not grants['read']['user_ids']
|
||||
and not grants['write']['group_ids']
|
||||
and not grants['write']['user_ids']
|
||||
):
|
||||
# No grants = should not happen (would mean no entries), default to {}
|
||||
access_control_value = json.dumps({})
|
||||
else:
|
||||
# Custom permissions
|
||||
access_control_value = json.dumps(
|
||||
{
|
||||
'read': grants['read'],
|
||||
'write': grants['write'],
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
conn.execute(
|
||||
sa.text(f'UPDATE "{table_name}" SET access_control = :access_control WHERE id = :id'),
|
||||
{'access_control': access_control_value, 'id': resource_id},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Step 4: Set all resources WITHOUT entries to private
|
||||
# For files: NULL means private (owner-only), so leave as NULL
|
||||
# For other resources: {} means private, so update to {}
|
||||
if resource_type != 'file':
|
||||
try:
|
||||
conn.execute(
|
||||
sa.text(f"""
|
||||
UPDATE "{table_name}"
|
||||
SET access_control = :private_value
|
||||
WHERE id NOT IN (
|
||||
SELECT DISTINCT resource_id FROM access_grant WHERE resource_type = :resource_type
|
||||
)
|
||||
AND access_control IS NULL
|
||||
"""),
|
||||
{'private_value': json.dumps({}), 'resource_type': resource_type},
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
# For files, NULL stays NULL - no action needed
|
||||
|
||||
# Step 5: Drop the access_grant table
|
||||
op.drop_index('idx_access_grant_principal', table_name='access_grant')
|
||||
op.drop_index('idx_access_grant_resource', table_name='access_grant')
|
||||
op.drop_table('access_grant')
|
||||
@@ -0,0 +1,882 @@
|
||||
import logging
|
||||
import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import select, delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, Text, UniqueConstraint, or_, and_
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
####################
|
||||
# AccessGrant DB Schema
|
||||
####################
|
||||
|
||||
|
||||
class AccessGrant(Base):
|
||||
__tablename__ = 'access_grant'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
resource_type = Column(Text, nullable=False) # "knowledge", "model", "prompt", "tool", "note", "channel", "file"
|
||||
resource_id = Column(Text, nullable=False)
|
||||
principal_type = Column(Text, nullable=False) # "user" or "group"
|
||||
principal_id = Column(Text, nullable=False) # user_id, group_id, or "*" (wildcard for public)
|
||||
permission = Column(Text, nullable=False) # "read" or "write"
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
'resource_type',
|
||||
'resource_id',
|
||||
'principal_type',
|
||||
'principal_id',
|
||||
'permission',
|
||||
name='uq_access_grant_grant',
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class AccessGrantModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
resource_type: str
|
||||
resource_id: str
|
||||
principal_type: str
|
||||
principal_id: str
|
||||
permission: str
|
||||
created_at: int
|
||||
|
||||
|
||||
class AccessGrantResponse(BaseModel):
|
||||
"""Slim grant model for API responses — resource context is implicit from the parent."""
|
||||
|
||||
id: str
|
||||
principal_type: str
|
||||
principal_id: str
|
||||
permission: str
|
||||
|
||||
@classmethod
|
||||
def from_grant(cls, grant: 'AccessGrantModel') -> 'AccessGrantResponse':
|
||||
return cls(
|
||||
id=grant.id,
|
||||
principal_type=grant.principal_type,
|
||||
principal_id=grant.principal_id,
|
||||
permission=grant.permission,
|
||||
)
|
||||
|
||||
|
||||
####################
|
||||
# Conversion utilities
|
||||
####################
|
||||
|
||||
|
||||
def access_control_to_grants(
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
access_control: Optional[dict],
|
||||
) -> list[dict]:
|
||||
"""
|
||||
Convert an old-style access_control JSON dict to a flat list of grant dicts.
|
||||
|
||||
Semantics:
|
||||
- None → public read (user:* read) — except files which are private
|
||||
- {} → private/owner-only (no grants)
|
||||
- {read: {group_ids, user_ids}, write: {group_ids, user_ids}} → specific grants
|
||||
|
||||
Returns a list of dicts with keys: resource_type, resource_id, principal_type, principal_id, permission
|
||||
"""
|
||||
grants = []
|
||||
|
||||
if access_control is None:
|
||||
# NULL → public read (user:* for read)
|
||||
# Exception: files with NULL are private (owner-only), no grants needed
|
||||
if resource_type != 'file':
|
||||
grants.append(
|
||||
{
|
||||
'resource_type': resource_type,
|
||||
'resource_id': resource_id,
|
||||
'principal_type': 'user',
|
||||
'principal_id': '*',
|
||||
'permission': 'read',
|
||||
}
|
||||
)
|
||||
return grants
|
||||
|
||||
# {} → private/owner-only, no grants
|
||||
if not access_control:
|
||||
return grants
|
||||
|
||||
# Parse structured permissions
|
||||
for permission in ['read', 'write']:
|
||||
perm_data = access_control.get(permission, {})
|
||||
if not perm_data:
|
||||
continue
|
||||
|
||||
for group_id in perm_data.get('group_ids', []):
|
||||
grants.append(
|
||||
{
|
||||
'resource_type': resource_type,
|
||||
'resource_id': resource_id,
|
||||
'principal_type': 'group',
|
||||
'principal_id': group_id,
|
||||
'permission': permission,
|
||||
}
|
||||
)
|
||||
|
||||
for user_id in perm_data.get('user_ids', []):
|
||||
grants.append(
|
||||
{
|
||||
'resource_type': resource_type,
|
||||
'resource_id': resource_id,
|
||||
'principal_type': 'user',
|
||||
'principal_id': user_id,
|
||||
'permission': permission,
|
||||
}
|
||||
)
|
||||
|
||||
return grants
|
||||
|
||||
|
||||
def normalize_access_grants(access_grants: Optional[list]) -> list[dict]:
|
||||
"""
|
||||
Normalize direct access_grants payloads from API forms.
|
||||
|
||||
Keeps only valid grants and removes duplicates by
|
||||
(principal_type, principal_id, permission).
|
||||
"""
|
||||
if not access_grants:
|
||||
return []
|
||||
|
||||
deduped = {}
|
||||
for grant in access_grants:
|
||||
if isinstance(grant, BaseModel):
|
||||
grant = grant.model_dump()
|
||||
if not isinstance(grant, dict):
|
||||
continue
|
||||
|
||||
principal_type = grant.get('principal_type')
|
||||
principal_id = grant.get('principal_id')
|
||||
permission = grant.get('permission')
|
||||
|
||||
if principal_type not in ('user', 'group'):
|
||||
continue
|
||||
if permission not in ('read', 'write'):
|
||||
continue
|
||||
if not isinstance(principal_id, str) or not principal_id:
|
||||
continue
|
||||
|
||||
key = (principal_type, principal_id, permission)
|
||||
deduped[key] = {
|
||||
'id': (grant.get('id') if isinstance(grant.get('id'), str) and grant.get('id') else str(uuid.uuid4())),
|
||||
'principal_type': principal_type,
|
||||
'principal_id': principal_id,
|
||||
'permission': permission,
|
||||
}
|
||||
|
||||
return list(deduped.values())
|
||||
|
||||
|
||||
def has_public_read_access_grant(access_grants: Optional[list]) -> bool:
|
||||
"""
|
||||
Returns True when a direct grant list includes wildcard public-read.
|
||||
"""
|
||||
for grant in normalize_access_grants(access_grants):
|
||||
if grant['principal_type'] == 'user' and grant['principal_id'] == '*' and grant['permission'] == 'read':
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def has_public_write_access_grant(access_grants: Optional[list]) -> bool:
|
||||
"""
|
||||
Returns True when a direct grant list includes wildcard public-write.
|
||||
"""
|
||||
for grant in normalize_access_grants(access_grants):
|
||||
if grant['principal_type'] == 'user' and grant['principal_id'] == '*' and grant['permission'] == 'write':
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def has_user_access_grant(access_grants: Optional[list]) -> bool:
|
||||
"""
|
||||
Returns True when a direct grant list includes any non-wildcard user grant.
|
||||
"""
|
||||
for grant in normalize_access_grants(access_grants):
|
||||
if grant['principal_type'] == 'user' and grant['principal_id'] != '*':
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def strip_user_access_grants(access_grants: Optional[list]) -> list:
|
||||
"""
|
||||
Remove all non-wildcard user grants from the list.
|
||||
Keeps group grants and the public wildcard (user:*) intact.
|
||||
"""
|
||||
if not access_grants:
|
||||
return []
|
||||
return [
|
||||
grant
|
||||
for grant in access_grants
|
||||
if not (
|
||||
(grant.get('principal_type') if isinstance(grant, dict) else getattr(grant, 'principal_type', None))
|
||||
== 'user'
|
||||
and (grant.get('principal_id') if isinstance(grant, dict) else getattr(grant, 'principal_id', None)) != '*'
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def grants_to_access_control(grants: list) -> Optional[dict]:
|
||||
"""
|
||||
Convert a list of grant objects (AccessGrantModel or AccessGrantResponse)
|
||||
back to the old-style access_control JSON dict for backward compatibility.
|
||||
|
||||
Semantics:
|
||||
- [] (empty) → {} (private/owner-only)
|
||||
- Contains user:*:read → None (public), but write grants are preserved
|
||||
- Otherwise → {read: {group_ids, user_ids}, write: {group_ids, user_ids}}
|
||||
|
||||
Note: "public" (user:*:read) still allows additional write permissions
|
||||
to coexist. When the wildcard read is present the function returns None
|
||||
for the legacy dict, so callers that need write info should inspect the
|
||||
grants list directly.
|
||||
"""
|
||||
if not grants:
|
||||
return {} # No grants = private/owner-only
|
||||
|
||||
result = {
|
||||
'read': {'group_ids': [], 'user_ids': []},
|
||||
'write': {'group_ids': [], 'user_ids': []},
|
||||
}
|
||||
|
||||
is_public = False
|
||||
for grant in grants:
|
||||
if grant.principal_type == 'user' and grant.principal_id == '*' and grant.permission == 'read':
|
||||
is_public = True
|
||||
continue # Don't add wildcard to user_ids list
|
||||
|
||||
if grant.permission not in ('read', 'write'):
|
||||
continue
|
||||
|
||||
if grant.principal_type == 'group':
|
||||
if grant.principal_id not in result[grant.permission]['group_ids']:
|
||||
result[grant.permission]['group_ids'].append(grant.principal_id)
|
||||
elif grant.principal_type == 'user':
|
||||
if grant.principal_id not in result[grant.permission]['user_ids']:
|
||||
result[grant.permission]['user_ids'].append(grant.principal_id)
|
||||
|
||||
if is_public:
|
||||
return None # Public read access
|
||||
|
||||
return result
|
||||
|
||||
|
||||
####################
|
||||
# Table Operations
|
||||
####################
|
||||
|
||||
|
||||
class AccessGrantsTable:
|
||||
async def grant_access(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
principal_type: str,
|
||||
principal_id: str,
|
||||
permission: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[AccessGrantModel]:
|
||||
"""Add a single access grant. Idempotent (ignores duplicates)."""
|
||||
async with get_async_db_context(db) as db:
|
||||
# Check for existing grant
|
||||
result = await db.execute(
|
||||
select(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
principal_type=principal_type,
|
||||
principal_id=principal_id,
|
||||
permission=permission,
|
||||
)
|
||||
)
|
||||
existing = result.scalars().first()
|
||||
if existing:
|
||||
return AccessGrantModel.model_validate(existing)
|
||||
|
||||
grant = AccessGrant(
|
||||
id=str(uuid.uuid4()),
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
principal_type=principal_type,
|
||||
principal_id=principal_id,
|
||||
permission=permission,
|
||||
created_at=int(time.time()),
|
||||
)
|
||||
db.add(grant)
|
||||
await db.commit()
|
||||
await db.refresh(grant)
|
||||
return AccessGrantModel.model_validate(grant)
|
||||
|
||||
async def revoke_access(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
principal_type: str,
|
||||
principal_id: str,
|
||||
permission: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> bool:
|
||||
"""Remove a single access grant."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
delete(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
principal_type=principal_type,
|
||||
principal_id=principal_id,
|
||||
permission=permission,
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
return result.rowcount > 0
|
||||
|
||||
async def revoke_all_access(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> int:
|
||||
"""Remove all access grants for a resource."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
delete(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
return result.rowcount
|
||||
|
||||
async def set_access_control(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
access_control: Optional[dict],
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[AccessGrantModel]:
|
||||
"""
|
||||
Replace all grants for a resource from an access_control JSON dict.
|
||||
This is the primary bridge for backward compat with the frontend.
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
# Delete all existing grants for this resource
|
||||
await db.execute(
|
||||
delete(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
)
|
||||
)
|
||||
|
||||
# Convert JSON to grant dicts
|
||||
grant_dicts = access_control_to_grants(resource_type, resource_id, access_control)
|
||||
|
||||
# Insert new grants
|
||||
results = []
|
||||
for grant_dict in grant_dicts:
|
||||
grant = AccessGrant(
|
||||
id=str(uuid.uuid4()),
|
||||
**grant_dict,
|
||||
created_at=int(time.time()),
|
||||
)
|
||||
db.add(grant)
|
||||
results.append(grant)
|
||||
|
||||
await db.commit()
|
||||
|
||||
return [AccessGrantModel.model_validate(g) for g in results]
|
||||
|
||||
async def set_access_grants(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
access_grants: Optional[list],
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[AccessGrantModel]:
|
||||
"""
|
||||
Replace all grants for a resource from a direct access_grants list.
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(
|
||||
delete(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
)
|
||||
)
|
||||
|
||||
normalized_grants = normalize_access_grants(access_grants)
|
||||
|
||||
results = []
|
||||
for grant_dict in normalized_grants:
|
||||
grant = AccessGrant(
|
||||
id=str(uuid.uuid4()),
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
principal_type=grant_dict['principal_type'],
|
||||
principal_id=grant_dict['principal_id'],
|
||||
permission=grant_dict['permission'],
|
||||
created_at=int(time.time()),
|
||||
)
|
||||
db.add(grant)
|
||||
results.append(grant)
|
||||
|
||||
await db.commit()
|
||||
return [AccessGrantModel.model_validate(g) for g in results]
|
||||
|
||||
async def get_access_control(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[dict]:
|
||||
"""
|
||||
Reconstruct the old-style access_control JSON dict from grants.
|
||||
For backward compat with the frontend.
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
)
|
||||
)
|
||||
grants = result.scalars().all()
|
||||
grant_models = [AccessGrantModel.model_validate(g) for g in grants]
|
||||
return grants_to_access_control(grant_models)
|
||||
|
||||
async def get_grants_by_resource(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[AccessGrantModel]:
|
||||
"""Get all grants for a specific resource."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
)
|
||||
)
|
||||
grants = result.scalars().all()
|
||||
return [AccessGrantModel.model_validate(g) for g in grants]
|
||||
|
||||
async def get_grants_by_resources(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_ids: list[str],
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, list[AccessGrantModel]]:
|
||||
"""Batch-fetch grants for multiple resources. Returns {resource_id: [grants]}."""
|
||||
if not resource_ids:
|
||||
return {}
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AccessGrant).filter(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id.in_(resource_ids),
|
||||
)
|
||||
)
|
||||
grants = result.scalars().all()
|
||||
result_dict: dict[str, list[AccessGrantModel]] = {rid: [] for rid in resource_ids}
|
||||
for g in grants:
|
||||
result_dict[g.resource_id].append(AccessGrantModel.model_validate(g))
|
||||
return result_dict
|
||||
|
||||
async def has_access(
|
||||
self,
|
||||
user_id: str,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
permission: str = 'read',
|
||||
user_group_ids: Optional[set[str]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Check if a user has the specified permission on a resource.
|
||||
|
||||
Access is granted if any of the following is true:
|
||||
- There's a grant for user:* (public) with the requested permission
|
||||
- There's a grant for the specific user with the requested permission
|
||||
- There's a grant for any of the user's groups with the requested permission
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
# Build conditions for matching grants
|
||||
conditions = [
|
||||
# Public access
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == '*',
|
||||
),
|
||||
# Direct user access
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == user_id,
|
||||
),
|
||||
]
|
||||
|
||||
# Group access
|
||||
if user_group_ids is None:
|
||||
from open_webui.models.groups import Groups
|
||||
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
|
||||
if user_group_ids:
|
||||
conditions.append(
|
||||
and_(
|
||||
AccessGrant.principal_type == 'group',
|
||||
AccessGrant.principal_id.in_(user_group_ids),
|
||||
)
|
||||
)
|
||||
|
||||
result = await db.execute(
|
||||
select(AccessGrant)
|
||||
.filter(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id == resource_id,
|
||||
AccessGrant.permission == permission,
|
||||
or_(*conditions),
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
grant = result.scalars().first()
|
||||
return grant is not None
|
||||
|
||||
async def get_accessible_resource_ids(
|
||||
self,
|
||||
user_id: str,
|
||||
resource_type: str,
|
||||
resource_ids: list[str],
|
||||
permission: str = 'read',
|
||||
user_group_ids: Optional[set[str]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> set[str]:
|
||||
"""
|
||||
Batch check: return the subset of resource_ids that the user can access.
|
||||
|
||||
This replaces calling has_access() in a loop (N+1) with a single query.
|
||||
"""
|
||||
if not resource_ids:
|
||||
return set()
|
||||
|
||||
async with get_async_db_context(db) as db:
|
||||
conditions = [
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == '*',
|
||||
),
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == user_id,
|
||||
),
|
||||
]
|
||||
|
||||
if user_group_ids is None:
|
||||
from open_webui.models.groups import Groups
|
||||
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
|
||||
if user_group_ids:
|
||||
conditions.append(
|
||||
and_(
|
||||
AccessGrant.principal_type == 'group',
|
||||
AccessGrant.principal_id.in_(user_group_ids),
|
||||
)
|
||||
)
|
||||
|
||||
result = await db.execute(
|
||||
select(AccessGrant.resource_id)
|
||||
.filter(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id.in_(resource_ids),
|
||||
AccessGrant.permission == permission,
|
||||
or_(*conditions),
|
||||
)
|
||||
.distinct()
|
||||
)
|
||||
rows = result.all()
|
||||
return {row[0] for row in rows}
|
||||
|
||||
async def get_users_with_access(
|
||||
self,
|
||||
resource_type: str,
|
||||
resource_id: str,
|
||||
permission: str = 'read',
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list:
|
||||
"""
|
||||
Get all users who have the specified permission on a resource.
|
||||
Returns a list of UserModel instances.
|
||||
"""
|
||||
from open_webui.models.users import Users, UserModel
|
||||
from open_webui.models.groups import Groups
|
||||
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AccessGrant).filter_by(
|
||||
resource_type=resource_type,
|
||||
resource_id=resource_id,
|
||||
permission=permission,
|
||||
)
|
||||
)
|
||||
grants = result.scalars().all()
|
||||
|
||||
# Check for public access
|
||||
for grant in grants:
|
||||
if grant.principal_type == 'user' and grant.principal_id == '*':
|
||||
result = await Users.get_users(filter={'roles': ['!pending']}, db=db)
|
||||
return result.get('users', [])
|
||||
|
||||
user_ids_with_access = set()
|
||||
|
||||
for grant in grants:
|
||||
if grant.principal_type == 'user':
|
||||
user_ids_with_access.add(grant.principal_id)
|
||||
elif grant.principal_type == 'group':
|
||||
group_user_ids = await Groups.get_group_user_ids_by_id(grant.principal_id, db=db)
|
||||
if group_user_ids:
|
||||
user_ids_with_access.update(group_user_ids)
|
||||
|
||||
if not user_ids_with_access:
|
||||
return []
|
||||
|
||||
return await Users.get_users_by_user_ids(list(user_ids_with_access), db=db)
|
||||
|
||||
def has_permission_filter(
|
||||
self,
|
||||
db,
|
||||
query,
|
||||
DocumentModel,
|
||||
filter: dict,
|
||||
resource_type: str,
|
||||
permission: str = 'read',
|
||||
):
|
||||
"""
|
||||
Apply access control filtering to a SQLAlchemy query by JOINing with access_grant.
|
||||
|
||||
This replaces the old JSON-column-based filtering with a proper relational JOIN.
|
||||
|
||||
Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself,
|
||||
so it remains synchronous. The caller is responsible for executing the query
|
||||
asynchronously with `await db.execute(...)`.
|
||||
"""
|
||||
group_ids = filter.get('group_ids', [])
|
||||
user_id = filter.get('user_id')
|
||||
|
||||
if permission == 'read_only':
|
||||
return self._has_read_only_permission_filter(db, query, DocumentModel, filter, resource_type)
|
||||
|
||||
# Build principal conditions
|
||||
principal_conditions = []
|
||||
|
||||
if group_ids or user_id:
|
||||
# Public access: user:* read
|
||||
principal_conditions.append(
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == '*',
|
||||
)
|
||||
)
|
||||
|
||||
if user_id:
|
||||
# Owner always has access
|
||||
principal_conditions.append(DocumentModel.user_id == user_id)
|
||||
|
||||
# Direct user grant
|
||||
principal_conditions.append(
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == user_id,
|
||||
)
|
||||
)
|
||||
|
||||
if group_ids:
|
||||
# Group grants
|
||||
principal_conditions.append(
|
||||
and_(
|
||||
AccessGrant.principal_type == 'group',
|
||||
AccessGrant.principal_id.in_(group_ids),
|
||||
)
|
||||
)
|
||||
|
||||
if not principal_conditions:
|
||||
return query
|
||||
|
||||
# LEFT JOIN access_grant and filter
|
||||
# We use a subquery approach to avoid duplicates from multiple matching grants
|
||||
from sqlalchemy import exists as sa_exists
|
||||
|
||||
grant_exists = (
|
||||
select(AccessGrant.id)
|
||||
.where(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id == DocumentModel.id,
|
||||
AccessGrant.permission == permission,
|
||||
or_(
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == '*',
|
||||
),
|
||||
*(
|
||||
[
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == user_id,
|
||||
)
|
||||
]
|
||||
if user_id
|
||||
else []
|
||||
),
|
||||
*(
|
||||
[
|
||||
and_(
|
||||
AccessGrant.principal_type == 'group',
|
||||
AccessGrant.principal_id.in_(group_ids),
|
||||
)
|
||||
]
|
||||
if group_ids
|
||||
else []
|
||||
),
|
||||
),
|
||||
)
|
||||
.correlate(DocumentModel)
|
||||
.exists()
|
||||
)
|
||||
|
||||
# Owner OR has a matching grant
|
||||
owner_or_grant = [grant_exists]
|
||||
if user_id:
|
||||
owner_or_grant.append(DocumentModel.user_id == user_id)
|
||||
|
||||
query = query.filter(or_(*owner_or_grant))
|
||||
return query
|
||||
|
||||
def _has_read_only_permission_filter(
|
||||
self,
|
||||
db,
|
||||
query,
|
||||
DocumentModel,
|
||||
filter: dict,
|
||||
resource_type: str,
|
||||
):
|
||||
"""
|
||||
Filter for items where user has read BUT NOT write access.
|
||||
Public items are NOT considered read_only.
|
||||
|
||||
Note: This method builds SQLAlchemy expressions and does NOT perform I/O itself,
|
||||
so it remains synchronous. The caller is responsible for executing the query
|
||||
asynchronously with `await db.execute(...)`.
|
||||
"""
|
||||
group_ids = filter.get('group_ids', [])
|
||||
user_id = filter.get('user_id')
|
||||
|
||||
from sqlalchemy import exists as sa_exists
|
||||
|
||||
# Has read grant (not public)
|
||||
read_grant_exists = (
|
||||
select(AccessGrant.id)
|
||||
.where(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id == DocumentModel.id,
|
||||
AccessGrant.permission == 'read',
|
||||
or_(
|
||||
*(
|
||||
[
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == user_id,
|
||||
)
|
||||
]
|
||||
if user_id
|
||||
else []
|
||||
),
|
||||
*(
|
||||
[
|
||||
and_(
|
||||
AccessGrant.principal_type == 'group',
|
||||
AccessGrant.principal_id.in_(group_ids),
|
||||
)
|
||||
]
|
||||
if group_ids
|
||||
else []
|
||||
),
|
||||
),
|
||||
)
|
||||
.correlate(DocumentModel)
|
||||
.exists()
|
||||
)
|
||||
|
||||
# Does NOT have write grant
|
||||
write_grant_exists = (
|
||||
select(AccessGrant.id)
|
||||
.where(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id == DocumentModel.id,
|
||||
AccessGrant.permission == 'write',
|
||||
or_(
|
||||
*(
|
||||
[
|
||||
and_(
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == user_id,
|
||||
)
|
||||
]
|
||||
if user_id
|
||||
else []
|
||||
),
|
||||
*(
|
||||
[
|
||||
and_(
|
||||
AccessGrant.principal_type == 'group',
|
||||
AccessGrant.principal_id.in_(group_ids),
|
||||
)
|
||||
]
|
||||
if group_ids
|
||||
else []
|
||||
),
|
||||
),
|
||||
)
|
||||
.correlate(DocumentModel)
|
||||
.exists()
|
||||
)
|
||||
|
||||
# Is NOT public
|
||||
public_grant_exists = (
|
||||
select(AccessGrant.id)
|
||||
.where(
|
||||
AccessGrant.resource_type == resource_type,
|
||||
AccessGrant.resource_id == DocumentModel.id,
|
||||
AccessGrant.permission == 'read',
|
||||
AccessGrant.principal_type == 'user',
|
||||
AccessGrant.principal_id == '*',
|
||||
)
|
||||
.correlate(DocumentModel)
|
||||
.exists()
|
||||
)
|
||||
|
||||
conditions = [read_grant_exists, ~write_grant_exists, ~public_grant_exists]
|
||||
|
||||
# Not owner
|
||||
if user_id:
|
||||
conditions.append(DocumentModel.user_id != user_id)
|
||||
|
||||
query = query.filter(and_(*conditions))
|
||||
return query
|
||||
|
||||
|
||||
AccessGrants = AccessGrantsTable()
|
||||
@@ -2,10 +2,12 @@ import logging
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from open_webui.models.users import UserModel, UserProfileImageResponse, Users
|
||||
from pydantic import BaseModel
|
||||
from sqlalchemy import select, delete, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.users import User, UserModel, UserProfileImageResponse, Users
|
||||
from open_webui.utils.validate import validate_profile_image_url
|
||||
from pydantic import BaseModel, field_validator
|
||||
from sqlalchemy import Boolean, Column, String, Text
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -16,7 +18,7 @@ log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Auth(Base):
|
||||
__tablename__ = "auth"
|
||||
__tablename__ = 'auth'
|
||||
|
||||
id = Column(String, primary_key=True, unique=True)
|
||||
email = Column(String)
|
||||
@@ -72,59 +74,63 @@ class SignupForm(BaseModel):
|
||||
name: str
|
||||
email: str
|
||||
password: str
|
||||
profile_image_url: Optional[str] = "/user.png"
|
||||
profile_image_url: Optional[str] = '/user.png'
|
||||
|
||||
@field_validator('profile_image_url')
|
||||
@classmethod
|
||||
def check_profile_image_url(cls, v: Optional[str]) -> Optional[str]:
|
||||
if v is not None:
|
||||
return validate_profile_image_url(v)
|
||||
return v
|
||||
|
||||
|
||||
class AddUserForm(SignupForm):
|
||||
role: Optional[str] = "pending"
|
||||
role: Optional[str] = 'pending'
|
||||
|
||||
|
||||
class AuthsTable:
|
||||
def insert_new_auth(
|
||||
async def insert_new_auth(
|
||||
self,
|
||||
email: str,
|
||||
password: str,
|
||||
name: str,
|
||||
profile_image_url: str = "/user.png",
|
||||
role: str = "pending",
|
||||
profile_image_url: str = '/user.png',
|
||||
role: str = 'pending',
|
||||
oauth: Optional[dict] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[UserModel]:
|
||||
with get_db_context(db) as db:
|
||||
log.info("insert_new_auth")
|
||||
async with get_async_db_context(db) as db:
|
||||
log.info('insert_new_auth')
|
||||
|
||||
id = str(uuid.uuid4())
|
||||
|
||||
auth = AuthModel(
|
||||
**{"id": id, "email": email, "password": password, "active": True}
|
||||
)
|
||||
auth = AuthModel(**{'id': id, 'email': email, 'password': password, 'active': True})
|
||||
result = Auth(**auth.model_dump())
|
||||
db.add(result)
|
||||
|
||||
user = Users.insert_new_user(
|
||||
id, name, email, profile_image_url, role, oauth=oauth, db=db
|
||||
)
|
||||
user = await Users.insert_new_user(id, name, email, profile_image_url, role, oauth=oauth, db=db)
|
||||
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
|
||||
if result and user:
|
||||
return user
|
||||
else:
|
||||
return None
|
||||
|
||||
def authenticate_user(
|
||||
self, email: str, verify_password: callable, db: Optional[Session] = None
|
||||
async def authenticate_user(
|
||||
self, email: str, verify_password: callable, db: Optional[AsyncSession] = None
|
||||
) -> Optional[UserModel]:
|
||||
log.info(f"authenticate_user: {email}")
|
||||
log.info(f'authenticate_user: {email}')
|
||||
|
||||
user = Users.get_user_by_email(email, db=db)
|
||||
user = await Users.get_user_by_email(email, db=db)
|
||||
if not user:
|
||||
return None
|
||||
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
auth = db.query(Auth).filter_by(id=user.id, active=True).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Auth).filter_by(id=user.id, active=True))
|
||||
auth = result.scalars().first()
|
||||
if auth:
|
||||
if verify_password(auth.password):
|
||||
return user
|
||||
@@ -135,66 +141,66 @@ class AuthsTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def authenticate_user_by_api_key(
|
||||
self, api_key: str, db: Optional[Session] = None
|
||||
async def authenticate_user_by_api_key(
|
||||
self, api_key: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[UserModel]:
|
||||
log.info(f"authenticate_user_by_api_key: {api_key}")
|
||||
log.info(f'authenticate_user_by_api_key')
|
||||
# if no api_key, return None
|
||||
if not api_key:
|
||||
return None
|
||||
|
||||
try:
|
||||
user = Users.get_user_by_api_key(api_key, db=db)
|
||||
user = await Users.get_user_by_api_key(api_key, db=db)
|
||||
return user if user else None
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def authenticate_user_by_email(
|
||||
self, email: str, db: Optional[Session] = None
|
||||
) -> Optional[UserModel]:
|
||||
log.info(f"authenticate_user_by_email: {email}")
|
||||
async def authenticate_user_by_email(self, email: str, db: Optional[AsyncSession] = None) -> Optional[UserModel]:
|
||||
log.info(f'authenticate_user_by_email: {email}')
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
auth = db.query(Auth).filter_by(email=email, active=True).first()
|
||||
if auth:
|
||||
user = Users.get_user_by_id(auth.id, db=db)
|
||||
return user
|
||||
async with get_async_db_context(db) as db:
|
||||
# Single JOIN query instead of two separate queries
|
||||
result = await db.execute(
|
||||
select(Auth, User).join(User, Auth.id == User.id).filter(Auth.email == email, Auth.active == True)
|
||||
)
|
||||
row = result.first()
|
||||
if row:
|
||||
_, user = row
|
||||
return UserModel.model_validate(user)
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def update_user_password_by_id(
|
||||
self, id: str, new_password: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
async def update_user_password_by_id(self, id: str, new_password: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
result = (
|
||||
db.query(Auth).filter_by(id=id).update({"password": new_password})
|
||||
)
|
||||
db.commit()
|
||||
return True if result == 1 else False
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(update(Auth).filter_by(id=id).values(password=new_password))
|
||||
await db.commit()
|
||||
return True if result.rowcount == 1 else False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def update_email_by_id(
|
||||
self, id: str, email: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
async def update_email_by_id(self, id: str, email: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
result = db.query(Auth).filter_by(id=id).update({"email": email})
|
||||
db.commit()
|
||||
return True if result == 1 else False
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(update(Auth).filter_by(id=id).values(email=email))
|
||||
await db.commit()
|
||||
if result.rowcount == 1:
|
||||
await Users.update_user_by_id(id, {'email': email}, db=db)
|
||||
return True
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_auth_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
async def delete_auth_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Delete User
|
||||
result = Users.delete_user_by_id(id, db=db)
|
||||
result = await Users.delete_user_by_id(id, db=db)
|
||||
|
||||
if result:
|
||||
db.query(Auth).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(Auth).filter_by(id=id))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,421 @@
|
||||
import time
|
||||
import logging
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import Column, Text, JSON, Boolean, BigInteger, Index, select, or_, func, cast, String, delete, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
####################
|
||||
# Automation DB Schema
|
||||
####################
|
||||
|
||||
|
||||
class Automation(Base):
|
||||
__tablename__ = 'automation'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
user_id = Column(Text, nullable=False)
|
||||
name = Column(Text, nullable=False)
|
||||
data = Column(JSON, nullable=False) # {prompt, model_id, rrule}
|
||||
meta = Column(JSON, nullable=True)
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
last_run_at = Column(BigInteger, nullable=True)
|
||||
next_run_at = Column(BigInteger, nullable=True)
|
||||
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (Index('ix_automation_next_run', 'next_run_at'),)
|
||||
|
||||
|
||||
class AutomationRun(Base):
|
||||
__tablename__ = 'automation_run'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
automation_id = Column(Text, nullable=False)
|
||||
chat_id = Column(Text, nullable=True)
|
||||
status = Column(Text, nullable=False) # success | error
|
||||
error = Column(Text, nullable=True)
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
Index('ix_automation_run_automation_id', 'automation_id'),
|
||||
Index('ix_automation_run_aid_created', 'automation_id', 'created_at'),
|
||||
)
|
||||
|
||||
|
||||
####################
|
||||
# Pydantic Models
|
||||
####################
|
||||
|
||||
|
||||
class AutomationTerminalConfig(BaseModel):
|
||||
server_id: str
|
||||
cwd: Optional[str] = None
|
||||
|
||||
|
||||
class AutomationData(BaseModel):
|
||||
prompt: str
|
||||
model_id: str
|
||||
rrule: str
|
||||
terminal: Optional[AutomationTerminalConfig] = None
|
||||
|
||||
|
||||
class AutomationModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
data: dict
|
||||
meta: Optional[dict] = None
|
||||
is_active: bool
|
||||
last_run_at: Optional[int] = None
|
||||
next_run_at: Optional[int] = None
|
||||
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
|
||||
class AutomationRunModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
automation_id: str
|
||||
chat_id: Optional[str] = None
|
||||
status: str
|
||||
error: Optional[str] = None
|
||||
created_at: int
|
||||
|
||||
|
||||
class AutomationForm(BaseModel):
|
||||
name: str
|
||||
data: AutomationData
|
||||
meta: Optional[dict] = None
|
||||
is_active: Optional[bool] = True
|
||||
|
||||
|
||||
class AutomationResponse(AutomationModel):
|
||||
last_run: Optional[AutomationRunModel] = None
|
||||
next_runs: Optional[list[int]] = None
|
||||
|
||||
|
||||
class AutomationListResponse(BaseModel):
|
||||
items: list[AutomationModel]
|
||||
total: int
|
||||
|
||||
|
||||
####################
|
||||
# AutomationTable
|
||||
####################
|
||||
|
||||
|
||||
class AutomationTable:
|
||||
async def insert(
|
||||
self,
|
||||
user_id: str,
|
||||
form: AutomationForm,
|
||||
next_run_at: int,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> AutomationModel:
|
||||
async with get_async_db_context(db) as db:
|
||||
now = int(time.time_ns())
|
||||
row = Automation(
|
||||
id=str(uuid4()),
|
||||
user_id=user_id,
|
||||
name=form.name,
|
||||
data=form.data.model_dump(),
|
||||
meta=form.meta,
|
||||
is_active=form.is_active,
|
||||
next_run_at=next_run_at,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(row)
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
return AutomationModel.model_validate(row)
|
||||
|
||||
async def count_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> int:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(func.count()).select_from(Automation).filter_by(user_id=user_id))
|
||||
return result.scalar()
|
||||
|
||||
async def get_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
row = await db.get(Automation, id)
|
||||
return AutomationModel.model_validate(row) if row else None
|
||||
|
||||
async def get_active_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> list[AutomationModel]:
|
||||
"""Get active automations for a user (for calendar RRULE expansion)."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Automation).filter_by(user_id=user_id, is_active=True).order_by(Automation.created_at.desc())
|
||||
)
|
||||
return [AutomationModel.model_validate(r) for r in result.scalars().all()]
|
||||
|
||||
async def search_automations(
|
||||
self,
|
||||
user_id: str,
|
||||
query: Optional[str] = None,
|
||||
status: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> 'AutomationListResponse':
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Automation).filter_by(user_id=user_id)
|
||||
|
||||
if query:
|
||||
search = f'%{query}%'
|
||||
# Search in name and prompt inside JSON data
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
Automation.name.ilike(search),
|
||||
cast(Automation.data, String).ilike(search),
|
||||
)
|
||||
)
|
||||
|
||||
if status == 'active':
|
||||
stmt = stmt.filter(Automation.is_active == True)
|
||||
elif status == 'paused':
|
||||
stmt = stmt.filter(Automation.is_active == False)
|
||||
|
||||
stmt = stmt.order_by(Automation.created_at.desc())
|
||||
|
||||
# Get total count
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
rows = result.scalars().all()
|
||||
return AutomationListResponse(
|
||||
items=[AutomationModel.model_validate(r) for r in rows],
|
||||
total=total,
|
||||
)
|
||||
|
||||
async def update_by_id(
|
||||
self,
|
||||
id: str,
|
||||
form: AutomationForm,
|
||||
next_run_at: int,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[AutomationModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
row = await db.get(Automation, id)
|
||||
if not row:
|
||||
return None
|
||||
row.name = form.name
|
||||
row.data = form.data.model_dump()
|
||||
row.meta = form.meta
|
||||
if form.is_active is not None:
|
||||
row.is_active = form.is_active
|
||||
row.next_run_at = next_run_at
|
||||
row.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
return AutomationModel.model_validate(row)
|
||||
|
||||
async def toggle(
|
||||
self,
|
||||
id: str,
|
||||
next_run_at: Optional[int],
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[AutomationModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
row = await db.get(Automation, id)
|
||||
if not row:
|
||||
return None
|
||||
row.is_active = not row.is_active
|
||||
row.next_run_at = next_run_at if row.is_active else None
|
||||
row.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
return AutomationModel.model_validate(row)
|
||||
|
||||
async def delete(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
row = await db.get(Automation, id)
|
||||
if not row:
|
||||
return False
|
||||
await db.delete(row)
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
async def claim_due(self, now_ns: int, limit: int = 10, db: Optional[AsyncSession] = None) -> list[AutomationModel]:
|
||||
"""
|
||||
Atomically claim due automations for execution.
|
||||
|
||||
Advances next_run_at immediately so the row can never be
|
||||
double-claimed. On PostgreSQL, uses FOR UPDATE SKIP LOCKED
|
||||
for zero-contention distributed work claiming.
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = (
|
||||
select(Automation)
|
||||
.where(
|
||||
Automation.is_active == True,
|
||||
Automation.next_run_at <= now_ns,
|
||||
)
|
||||
.order_by(Automation.next_run_at)
|
||||
.limit(limit)
|
||||
)
|
||||
|
||||
if db.bind.dialect.name == 'postgresql':
|
||||
stmt = stmt.with_for_update(skip_locked=True)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
rows = result.scalars().all()
|
||||
|
||||
from open_webui.utils.automations import next_run_ns
|
||||
|
||||
# Batch-fetch user timezones so rescheduling respects each
|
||||
# user's local timezone instead of falling back to server time.
|
||||
user_ids = list({row.user_id for row in rows})
|
||||
timezone_by_user_id: dict[str, Optional[str]] = {}
|
||||
if user_ids:
|
||||
from open_webui.models.users import User
|
||||
|
||||
tz_result = await db.execute(select(User.id, User.timezone).where(User.id.in_(user_ids)))
|
||||
timezone_by_user_id = {uid: tz for uid, tz in tz_result.all()}
|
||||
|
||||
for row in rows:
|
||||
row.last_run_at = now_ns
|
||||
row.next_run_at = next_run_ns(row.data.get('rrule', ''), tz=timezone_by_user_id.get(row.user_id))
|
||||
|
||||
await db.commit()
|
||||
|
||||
return [AutomationModel.model_validate(r) for r in rows]
|
||||
|
||||
|
||||
####################
|
||||
# AutomationRunTable
|
||||
####################
|
||||
|
||||
|
||||
class AutomationRunTable:
|
||||
async def insert(
|
||||
self,
|
||||
automation_id: str,
|
||||
status: str,
|
||||
chat_id: Optional[str] = None,
|
||||
error: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> AutomationRunModel:
|
||||
async with get_async_db_context(db) as db:
|
||||
row = AutomationRun(
|
||||
id=str(uuid4()),
|
||||
automation_id=automation_id,
|
||||
chat_id=chat_id,
|
||||
status=status,
|
||||
error=error,
|
||||
created_at=int(time.time_ns()),
|
||||
)
|
||||
db.add(row)
|
||||
await db.commit()
|
||||
await db.refresh(row)
|
||||
return AutomationRunModel.model_validate(row)
|
||||
|
||||
async def get_latest(self, automation_id: str, db: Optional[AsyncSession] = None) -> Optional[AutomationRunModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AutomationRun)
|
||||
.filter_by(automation_id=automation_id)
|
||||
.order_by(AutomationRun.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
row = result.scalars().first()
|
||||
return AutomationRunModel.model_validate(row) if row else None
|
||||
|
||||
async def get_latest_batch(
|
||||
self, automation_ids: list[str], db: Optional[AsyncSession] = None
|
||||
) -> dict[str, AutomationRunModel]:
|
||||
"""Fetch the latest run for each automation in a single query."""
|
||||
if not automation_ids:
|
||||
return {}
|
||||
async with get_async_db_context(db) as db:
|
||||
# Subquery: max created_at per automation_id
|
||||
subq = (
|
||||
select(
|
||||
AutomationRun.automation_id,
|
||||
func.max(AutomationRun.created_at).label('max_created'),
|
||||
)
|
||||
.filter(AutomationRun.automation_id.in_(automation_ids))
|
||||
.group_by(AutomationRun.automation_id)
|
||||
.subquery()
|
||||
)
|
||||
result = await db.execute(
|
||||
select(AutomationRun).join(
|
||||
subq,
|
||||
(AutomationRun.automation_id == subq.c.automation_id)
|
||||
& (AutomationRun.created_at == subq.c.max_created),
|
||||
)
|
||||
)
|
||||
rows = result.scalars().all()
|
||||
return {row.automation_id: AutomationRunModel.model_validate(row) for row in rows}
|
||||
|
||||
async def get_by_automation(
|
||||
self,
|
||||
automation_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[AutomationRunModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AutomationRun)
|
||||
.filter_by(automation_id=automation_id)
|
||||
.order_by(AutomationRun.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
)
|
||||
rows = result.scalars().all()
|
||||
return [AutomationRunModel.model_validate(r) for r in rows]
|
||||
|
||||
async def delete_by_automation(self, automation_id: str, db: Optional[AsyncSession] = None) -> int:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(delete(AutomationRun).filter_by(automation_id=automation_id))
|
||||
await db.commit()
|
||||
return result.rowcount
|
||||
|
||||
async def get_runs_by_user_range(
|
||||
self,
|
||||
user_id: str,
|
||||
start_ns: int,
|
||||
end_ns: int,
|
||||
limit: int = 500,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[tuple['AutomationRunModel', 'AutomationModel']]:
|
||||
"""Get runs within a date range for a user, joined with parent automation."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(AutomationRun, Automation)
|
||||
.join(Automation, Automation.id == AutomationRun.automation_id)
|
||||
.filter(
|
||||
Automation.user_id == user_id,
|
||||
AutomationRun.created_at >= start_ns,
|
||||
AutomationRun.created_at < end_ns,
|
||||
)
|
||||
.order_by(AutomationRun.created_at.desc())
|
||||
.limit(limit)
|
||||
)
|
||||
return [
|
||||
(AutomationRunModel.model_validate(run), AutomationModel.model_validate(auto))
|
||||
for run, auto in result.all()
|
||||
]
|
||||
|
||||
|
||||
Automations = AutomationTable()
|
||||
AutomationRuns = AutomationRunTable()
|
||||
@@ -0,0 +1,822 @@
|
||||
import time
|
||||
import logging
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import (
|
||||
Column,
|
||||
Text,
|
||||
JSON,
|
||||
Boolean,
|
||||
BigInteger,
|
||||
Index,
|
||||
UniqueConstraint,
|
||||
select,
|
||||
or_,
|
||||
exists,
|
||||
func,
|
||||
delete,
|
||||
update,
|
||||
)
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import User, UserModel, UserResponse
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
####################
|
||||
# Calendar DB Schema
|
||||
####################
|
||||
|
||||
|
||||
class Calendar(Base):
|
||||
__tablename__ = 'calendar'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
user_id = Column(Text, nullable=False)
|
||||
name = Column(Text, nullable=False)
|
||||
color = Column(Text, nullable=True)
|
||||
is_default = Column(Boolean, nullable=False, default=False)
|
||||
data = Column(JSON, nullable=True)
|
||||
meta = Column(JSON, nullable=True)
|
||||
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (Index('ix_calendar_user', 'user_id'),)
|
||||
|
||||
|
||||
class CalendarEvent(Base):
|
||||
__tablename__ = 'calendar_event'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
calendar_id = Column(Text, nullable=False)
|
||||
user_id = Column(Text, nullable=False)
|
||||
title = Column(Text, nullable=False)
|
||||
description = Column(Text, nullable=True)
|
||||
start_at = Column(BigInteger, nullable=False)
|
||||
end_at = Column(BigInteger, nullable=True)
|
||||
all_day = Column(Boolean, nullable=False, default=False)
|
||||
rrule = Column(Text, nullable=True)
|
||||
color = Column(Text, nullable=True)
|
||||
location = Column(Text, nullable=True)
|
||||
data = Column(JSON, nullable=True)
|
||||
meta = Column(JSON, nullable=True)
|
||||
is_cancelled = Column(Boolean, nullable=False, default=False)
|
||||
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
Index('ix_calendar_event_calendar', 'calendar_id', 'start_at'),
|
||||
Index('ix_calendar_event_user_date', 'user_id', 'start_at'),
|
||||
)
|
||||
|
||||
|
||||
class CalendarEventAttendee(Base):
|
||||
__tablename__ = 'calendar_event_attendee'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
event_id = Column(Text, nullable=False)
|
||||
user_id = Column(Text, nullable=False)
|
||||
status = Column(Text, nullable=False, default='pending')
|
||||
meta = Column(JSON, nullable=True)
|
||||
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint('event_id', 'user_id', name='uq_event_attendee'),
|
||||
Index('ix_calendar_event_attendee_user', 'user_id', 'status'),
|
||||
)
|
||||
|
||||
|
||||
####################
|
||||
# Pydantic Models
|
||||
####################
|
||||
|
||||
|
||||
class CalendarModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
color: Optional[str] = None
|
||||
is_default: bool = False
|
||||
is_system: bool = False
|
||||
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
|
||||
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
||||
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
|
||||
class CalendarEventModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True, extra='allow')
|
||||
|
||||
id: str
|
||||
calendar_id: str
|
||||
user_id: str
|
||||
title: str
|
||||
description: Optional[str] = None
|
||||
start_at: int
|
||||
end_at: Optional[int] = None
|
||||
all_day: bool = False
|
||||
rrule: Optional[str] = None
|
||||
color: Optional[str] = None
|
||||
location: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
is_cancelled: bool = False
|
||||
|
||||
attendees: list['CalendarEventAttendeeModel'] = Field(default_factory=list)
|
||||
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
|
||||
class CalendarEventAttendeeModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
event_id: str
|
||||
user_id: str
|
||||
status: str = 'pending'
|
||||
meta: Optional[dict] = None
|
||||
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
|
||||
####################
|
||||
# Forms
|
||||
####################
|
||||
|
||||
|
||||
class CalendarForm(BaseModel):
|
||||
name: str
|
||||
color: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class CalendarUpdateForm(BaseModel):
|
||||
name: Optional[str] = None
|
||||
color: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class CalendarEventForm(BaseModel):
|
||||
calendar_id: str
|
||||
title: str
|
||||
description: Optional[str] = None
|
||||
start_at: int
|
||||
end_at: Optional[int] = None
|
||||
all_day: bool = False
|
||||
rrule: Optional[str] = None
|
||||
color: Optional[str] = None
|
||||
location: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
attendees: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class CalendarEventUpdateForm(BaseModel):
|
||||
calendar_id: Optional[str] = None
|
||||
title: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
start_at: Optional[int] = None
|
||||
end_at: Optional[int] = None
|
||||
all_day: Optional[bool] = None
|
||||
rrule: Optional[str] = None
|
||||
color: Optional[str] = None
|
||||
location: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
is_cancelled: Optional[bool] = None
|
||||
attendees: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class RSVPForm(BaseModel):
|
||||
status: str # 'accepted' | 'declined' | 'tentative' | 'pending'
|
||||
|
||||
|
||||
####################
|
||||
# Response Models
|
||||
####################
|
||||
|
||||
|
||||
class CalendarEventUserResponse(CalendarEventModel):
|
||||
user: Optional[UserResponse] = None
|
||||
|
||||
|
||||
class CalendarEventListResponse(BaseModel):
|
||||
items: list[CalendarEventUserResponse]
|
||||
total: int
|
||||
|
||||
|
||||
####################
|
||||
# Table Operations
|
||||
####################
|
||||
|
||||
|
||||
class CalendarTable:
|
||||
async def _get_access_grants(self, calendar_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
|
||||
return await AccessGrants.get_grants_by_resource('calendar', calendar_id, db=db)
|
||||
|
||||
async def _to_calendar_model(
|
||||
self,
|
||||
cal: Calendar,
|
||||
access_grants: Optional[list[AccessGrantModel]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> CalendarModel:
|
||||
cal_data = CalendarModel.model_validate(cal).model_dump(exclude={'access_grants'})
|
||||
cal_data['access_grants'] = (
|
||||
access_grants if access_grants is not None else await self._get_access_grants(cal_data['id'], db=db)
|
||||
)
|
||||
return CalendarModel.model_validate(cal_data)
|
||||
|
||||
async def get_or_create_defaults(self, user_id: str, db: Optional[AsyncSession] = None) -> list[CalendarModel]:
|
||||
"""Return user's calendars, creating 'Personal' default if none exist."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Calendar).filter(Calendar.user_id == user_id).order_by(Calendar.created_at.asc())
|
||||
)
|
||||
calendars = result.scalars().all()
|
||||
|
||||
if calendars:
|
||||
return [CalendarModel.model_validate(c) for c in calendars]
|
||||
|
||||
now = int(time.time_ns())
|
||||
cal = Calendar(
|
||||
id=str(uuid4()),
|
||||
user_id=user_id,
|
||||
name='Personal',
|
||||
color='#3b82f6',
|
||||
is_default=True,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(cal)
|
||||
await db.commit()
|
||||
return [CalendarModel.model_validate(cal)]
|
||||
|
||||
async def get_calendars_by_user(self, user_id: str, db: Optional[AsyncSession] = None) -> list[CalendarModel]:
|
||||
"""Owned + shared calendars."""
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [g.id for g in user_groups]
|
||||
|
||||
stmt = select(Calendar)
|
||||
stmt = AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=stmt,
|
||||
DocumentModel=Calendar,
|
||||
filter={'user_id': user_id, 'group_ids': user_group_ids},
|
||||
resource_type='calendar',
|
||||
permission='read',
|
||||
)
|
||||
stmt = stmt.order_by(Calendar.created_at.asc())
|
||||
|
||||
result = await db.execute(stmt)
|
||||
calendars = result.scalars().all()
|
||||
|
||||
if not calendars:
|
||||
return await self.get_or_create_defaults(user_id, db=db)
|
||||
|
||||
cal_ids = [c.id for c in calendars]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('calendar', cal_ids, db=db)
|
||||
|
||||
return [await self._to_calendar_model(c, access_grants=grants_map.get(c.id, []), db=db) for c in calendars]
|
||||
|
||||
async def get_calendar_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[CalendarModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Calendar).filter(Calendar.id == id))
|
||||
cal = result.scalars().first()
|
||||
return await self._to_calendar_model(cal, db=db) if cal else None
|
||||
|
||||
async def insert_new_calendar(
|
||||
self, user_id: str, form_data: CalendarForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[CalendarModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
now = int(time.time_ns())
|
||||
cal = Calendar(
|
||||
id=str(uuid4()),
|
||||
user_id=user_id,
|
||||
name=form_data.name,
|
||||
color=form_data.color,
|
||||
is_default=False,
|
||||
data=form_data.data,
|
||||
meta=form_data.meta,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(cal)
|
||||
await db.commit()
|
||||
if form_data.access_grants is not None:
|
||||
await AccessGrants.set_access_grants('calendar', cal.id, form_data.access_grants, db=db)
|
||||
return await self._to_calendar_model(cal, db=db)
|
||||
|
||||
async def update_calendar_by_id(
|
||||
self, id: str, form_data: CalendarUpdateForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[CalendarModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Calendar).filter(Calendar.id == id))
|
||||
cal = result.scalars().first()
|
||||
if not cal:
|
||||
return None
|
||||
|
||||
update_data = form_data.model_dump(exclude_unset=True)
|
||||
if 'name' in update_data:
|
||||
cal.name = update_data['name']
|
||||
if 'color' in update_data:
|
||||
cal.color = update_data['color']
|
||||
if 'data' in update_data:
|
||||
cal.data = {**(cal.data or {}), **update_data['data']}
|
||||
if 'meta' in update_data:
|
||||
cal.meta = {**(cal.meta or {}), **update_data['meta']}
|
||||
if 'access_grants' in update_data:
|
||||
await AccessGrants.set_access_grants('calendar', id, update_data['access_grants'], db=db)
|
||||
|
||||
cal.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
return await self._to_calendar_model(cal, db=db)
|
||||
|
||||
async def set_default_calendar(
|
||||
self, user_id: str, calendar_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[CalendarModel]:
|
||||
"""Set a calendar as the user's default, clearing all others."""
|
||||
async with get_async_db_context(db) as db:
|
||||
# Clear all defaults for this user
|
||||
await db.execute(
|
||||
update(Calendar)
|
||||
.where(Calendar.user_id == user_id, Calendar.is_default == True)
|
||||
.values(is_default=False)
|
||||
)
|
||||
# Set the new default
|
||||
result = await db.execute(select(Calendar).filter(Calendar.id == calendar_id, Calendar.user_id == user_id))
|
||||
cal = result.scalars().first()
|
||||
if not cal:
|
||||
return None
|
||||
cal.is_default = True
|
||||
cal.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
return await self._to_calendar_model(cal, db=db)
|
||||
|
||||
async def delete_calendar_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Delete a non-default calendar. Cascades to events, attendees, and grants."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Calendar).filter(Calendar.id == id))
|
||||
cal = result.scalars().first()
|
||||
if not cal or cal.is_default:
|
||||
return False
|
||||
|
||||
# Delete attendees for all events in this calendar
|
||||
event_ids_result = await db.execute(select(CalendarEvent.id).filter(CalendarEvent.calendar_id == id))
|
||||
event_ids = [r[0] for r in event_ids_result.all()]
|
||||
if event_ids:
|
||||
await db.execute(
|
||||
delete(CalendarEventAttendee).filter(CalendarEventAttendee.event_id.in_(event_ids))
|
||||
)
|
||||
|
||||
# Delete events
|
||||
await db.execute(delete(CalendarEvent).filter(CalendarEvent.calendar_id == id))
|
||||
|
||||
# Delete access grants
|
||||
await AccessGrants.revoke_all_access('calendar', id, db=db)
|
||||
|
||||
# Delete calendar
|
||||
await db.execute(delete(Calendar).filter(Calendar.id == id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
class CalendarEventTable:
|
||||
async def _get_attendees(
|
||||
self, event_id: str, db: Optional[AsyncSession] = None
|
||||
) -> list[CalendarEventAttendeeModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == event_id))
|
||||
rows = result.scalars().all()
|
||||
return [CalendarEventAttendeeModel.model_validate(r) for r in rows]
|
||||
|
||||
async def _to_event_model(
|
||||
self,
|
||||
event: CalendarEvent,
|
||||
attendees: Optional[list[CalendarEventAttendeeModel]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> CalendarEventModel:
|
||||
event_data = CalendarEventModel.model_validate(event).model_dump(exclude={'attendees'})
|
||||
event_data['attendees'] = (
|
||||
attendees if attendees is not None else await self._get_attendees(event_data['id'], db=db)
|
||||
)
|
||||
return CalendarEventModel.model_validate(event_data)
|
||||
|
||||
async def insert_new_event(
|
||||
self, user_id: str, form_data: CalendarEventForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[CalendarEventModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
now = int(time.time_ns())
|
||||
event = CalendarEvent(
|
||||
id=str(uuid4()),
|
||||
calendar_id=form_data.calendar_id,
|
||||
user_id=user_id,
|
||||
title=form_data.title,
|
||||
description=form_data.description,
|
||||
start_at=form_data.start_at,
|
||||
end_at=form_data.end_at,
|
||||
all_day=form_data.all_day,
|
||||
rrule=form_data.rrule,
|
||||
color=form_data.color,
|
||||
location=form_data.location,
|
||||
data=form_data.data,
|
||||
meta=form_data.meta,
|
||||
is_cancelled=False,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(event)
|
||||
await db.commit()
|
||||
|
||||
# Add attendees
|
||||
if form_data.attendees:
|
||||
await CalendarEventAttendees.set_attendees(event.id, form_data.attendees, db=db)
|
||||
|
||||
return await self._to_event_model(event, db=db)
|
||||
|
||||
async def get_event_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[CalendarEventModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(CalendarEvent).filter(CalendarEvent.id == id))
|
||||
event = result.scalars().first()
|
||||
return await self._to_event_model(event, db=db) if event else None
|
||||
|
||||
async def get_events_by_range(
|
||||
self,
|
||||
user_id: str,
|
||||
start: int,
|
||||
end: int,
|
||||
calendar_ids: Optional[list[str]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[CalendarEventUserResponse]:
|
||||
"""Fetch events visible to user within a date range.
|
||||
|
||||
Visible events = events in owned/shared calendars + events user attends.
|
||||
Recurring events are fetched if they have any rrule (expansion in Python).
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [g.id for g in user_groups]
|
||||
|
||||
# Get calendar IDs accessible to user
|
||||
cal_stmt = select(Calendar.id)
|
||||
cal_stmt = AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=cal_stmt,
|
||||
DocumentModel=Calendar,
|
||||
filter={'user_id': user_id, 'group_ids': user_group_ids},
|
||||
resource_type='calendar',
|
||||
permission='read',
|
||||
)
|
||||
cal_result = await db.execute(cal_stmt)
|
||||
accessible_cal_ids = [r[0] for r in cal_result.all()]
|
||||
|
||||
if calendar_ids:
|
||||
# Filter to requested calendars only
|
||||
accessible_cal_ids = [c for c in accessible_cal_ids if c in calendar_ids]
|
||||
|
||||
# Also get event IDs where user is an attendee
|
||||
attendee_event_ids_result = await db.execute(
|
||||
select(CalendarEventAttendee.event_id).filter(CalendarEventAttendee.user_id == user_id)
|
||||
)
|
||||
attendee_event_ids = [r[0] for r in attendee_event_ids_result.all()]
|
||||
|
||||
# Build conditions for accessible events
|
||||
conditions = []
|
||||
if accessible_cal_ids:
|
||||
conditions.append(CalendarEvent.calendar_id.in_(accessible_cal_ids))
|
||||
if attendee_event_ids:
|
||||
conditions.append(CalendarEvent.id.in_(attendee_event_ids))
|
||||
|
||||
if not conditions:
|
||||
return []
|
||||
|
||||
# Build event query
|
||||
stmt = (
|
||||
select(CalendarEvent, User)
|
||||
.outerjoin(User, User.id == CalendarEvent.user_id)
|
||||
.filter(
|
||||
CalendarEvent.is_cancelled == False,
|
||||
or_(*conditions),
|
||||
or_(
|
||||
# Non-recurring: overlaps the range
|
||||
(
|
||||
CalendarEvent.rrule.is_(None)
|
||||
& (CalendarEvent.start_at < end)
|
||||
& or_(
|
||||
CalendarEvent.end_at.is_(None) & (CalendarEvent.start_at >= start),
|
||||
CalendarEvent.end_at.isnot(None) & (CalendarEvent.end_at > start),
|
||||
)
|
||||
),
|
||||
# Recurring: fetch all (expansion in Python)
|
||||
CalendarEvent.rrule.isnot(None),
|
||||
),
|
||||
)
|
||||
.order_by(CalendarEvent.start_at.asc())
|
||||
)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
if not items:
|
||||
return []
|
||||
|
||||
# Batch-load attendees for all events in one query (avoid N+1)
|
||||
event_ids = [event.id for event, _user in items]
|
||||
att_result = await db.execute(
|
||||
select(CalendarEventAttendee).filter(CalendarEventAttendee.event_id.in_(event_ids))
|
||||
)
|
||||
att_rows = att_result.scalars().all()
|
||||
att_map: dict[str, list[CalendarEventAttendeeModel]] = {}
|
||||
for a in att_rows:
|
||||
att_map.setdefault(a.event_id, []).append(CalendarEventAttendeeModel.model_validate(a))
|
||||
|
||||
events = []
|
||||
for event, user in items:
|
||||
event_data = CalendarEventModel.model_validate(event).model_dump(exclude={'attendees'})
|
||||
event_data['attendees'] = att_map.get(event.id, [])
|
||||
events.append(
|
||||
CalendarEventUserResponse(
|
||||
**event_data,
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
return events
|
||||
|
||||
async def search_events(
|
||||
self,
|
||||
user_id: str,
|
||||
query: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> CalendarEventListResponse:
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [g.id for g in user_groups]
|
||||
|
||||
# Get accessible calendar IDs
|
||||
cal_stmt = select(Calendar.id)
|
||||
cal_stmt = AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=cal_stmt,
|
||||
DocumentModel=Calendar,
|
||||
filter={'user_id': user_id, 'group_ids': user_group_ids},
|
||||
resource_type='calendar',
|
||||
permission='read',
|
||||
)
|
||||
cal_result = await db.execute(cal_stmt)
|
||||
accessible_cal_ids = [r[0] for r in cal_result.all()]
|
||||
if not accessible_cal_ids:
|
||||
return CalendarEventListResponse(items=[], total=0)
|
||||
|
||||
stmt = (
|
||||
select(CalendarEvent, User)
|
||||
.outerjoin(User, User.id == CalendarEvent.user_id)
|
||||
.filter(
|
||||
CalendarEvent.is_cancelled == False,
|
||||
CalendarEvent.calendar_id.in_(accessible_cal_ids),
|
||||
)
|
||||
)
|
||||
|
||||
if query:
|
||||
search = f'%{query}%'
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
CalendarEvent.title.ilike(search),
|
||||
CalendarEvent.description.ilike(search),
|
||||
CalendarEvent.location.ilike(search),
|
||||
)
|
||||
)
|
||||
|
||||
stmt = stmt.order_by(CalendarEvent.start_at.desc())
|
||||
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
if not items:
|
||||
return CalendarEventListResponse(items=[], total=total)
|
||||
|
||||
# Batch-load attendees
|
||||
event_ids = [event.id for event, _user in items]
|
||||
att_result = await db.execute(
|
||||
select(CalendarEventAttendee).filter(CalendarEventAttendee.event_id.in_(event_ids))
|
||||
)
|
||||
att_rows = att_result.scalars().all()
|
||||
att_map: dict[str, list[CalendarEventAttendeeModel]] = {}
|
||||
for a in att_rows:
|
||||
att_map.setdefault(a.event_id, []).append(CalendarEventAttendeeModel.model_validate(a))
|
||||
|
||||
events = []
|
||||
for event, user in items:
|
||||
event_data = CalendarEventModel.model_validate(event).model_dump(exclude={'attendees'})
|
||||
event_data['attendees'] = att_map.get(event.id, [])
|
||||
events.append(
|
||||
CalendarEventUserResponse(
|
||||
**event_data,
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
return CalendarEventListResponse(items=events, total=total)
|
||||
|
||||
async def update_event_by_id(
|
||||
self, id: str, form_data: CalendarEventUpdateForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[CalendarEventModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(CalendarEvent).filter(CalendarEvent.id == id))
|
||||
event = result.scalars().first()
|
||||
if not event:
|
||||
return None
|
||||
|
||||
update_data = form_data.model_dump(exclude_unset=True)
|
||||
for field in [
|
||||
'calendar_id',
|
||||
'title',
|
||||
'description',
|
||||
'start_at',
|
||||
'end_at',
|
||||
'all_day',
|
||||
'rrule',
|
||||
'color',
|
||||
'location',
|
||||
'is_cancelled',
|
||||
]:
|
||||
if field in update_data:
|
||||
setattr(event, field, update_data[field])
|
||||
|
||||
if 'data' in update_data and update_data['data'] is not None:
|
||||
event.data = {**(event.data or {}), **update_data['data']}
|
||||
if 'meta' in update_data and update_data['meta'] is not None:
|
||||
event.meta = {**(event.meta or {}), **update_data['meta']}
|
||||
|
||||
if 'attendees' in update_data and update_data['attendees'] is not None:
|
||||
await CalendarEventAttendees.set_attendees(id, update_data['attendees'], db=db)
|
||||
|
||||
event.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
return await self._to_event_model(event, db=db)
|
||||
|
||||
async def get_upcoming_events(
|
||||
self,
|
||||
now_ns: int,
|
||||
default_lookahead_ns: int,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[tuple[CalendarEventModel, Optional[str]]]:
|
||||
"""Events starting between now and now + lookahead, for alert processing.
|
||||
|
||||
Per-event lookahead is read from meta.alert_minutes (falls back to
|
||||
default_lookahead_ns). Returns (event, user_timezone) pairs.
|
||||
"""
|
||||
from open_webui.models.users import User as UserRow
|
||||
|
||||
# Use the maximum possible lookahead (60 min) to cast a wide net;
|
||||
# per-event filtering happens in Python after fetching.
|
||||
max_lookahead_ns = max(default_lookahead_ns, 60 * 60 * 1_000_000_000)
|
||||
upper = now_ns + max_lookahead_ns
|
||||
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(CalendarEvent, UserRow.timezone)
|
||||
.outerjoin(UserRow, UserRow.id == CalendarEvent.user_id)
|
||||
.filter(
|
||||
CalendarEvent.is_cancelled == False,
|
||||
CalendarEvent.start_at >= now_ns,
|
||||
CalendarEvent.start_at <= upper,
|
||||
)
|
||||
)
|
||||
rows = result.all()
|
||||
|
||||
events = []
|
||||
for event, tz in rows:
|
||||
model = CalendarEventModel.model_validate(event)
|
||||
# Determine per-event alert window
|
||||
alert_minutes = None
|
||||
if model.meta and 'alert_minutes' in model.meta:
|
||||
alert_minutes = model.meta['alert_minutes']
|
||||
|
||||
if alert_minutes is not None:
|
||||
if alert_minutes < 0:
|
||||
# alert_minutes < 0 means "no alert"
|
||||
continue
|
||||
event_lookahead_ns = alert_minutes * 60 * 1_000_000_000
|
||||
else:
|
||||
event_lookahead_ns = default_lookahead_ns
|
||||
|
||||
if model.start_at <= now_ns + event_lookahead_ns:
|
||||
events.append((model, tz))
|
||||
|
||||
return events
|
||||
|
||||
async def delete_event_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == id))
|
||||
await db.execute(delete(CalendarEvent).filter(CalendarEvent.id == id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
class CalendarEventAttendeeTable:
|
||||
async def set_attendees(
|
||||
self, event_id: str, attendees: list[dict], db: Optional[AsyncSession] = None
|
||||
) -> list[CalendarEventAttendeeModel]:
|
||||
"""Replace all attendees for an event.
|
||||
|
||||
Each dict in attendees: {user_id: str, status?: str, meta?: dict}
|
||||
"""
|
||||
async with get_async_db_context(db) as db:
|
||||
# Remove existing
|
||||
await db.execute(delete(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == event_id))
|
||||
|
||||
now = int(time.time_ns())
|
||||
models = []
|
||||
for att in attendees:
|
||||
row = CalendarEventAttendee(
|
||||
id=str(uuid4()),
|
||||
event_id=event_id,
|
||||
user_id=att['user_id'],
|
||||
status=att.get('status', 'pending'),
|
||||
meta=att.get('meta'),
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(row)
|
||||
models.append(CalendarEventAttendeeModel.model_validate(row))
|
||||
|
||||
await db.commit()
|
||||
return models
|
||||
|
||||
async def update_rsvp(
|
||||
self, event_id: str, user_id: str, status: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[CalendarEventAttendeeModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(CalendarEventAttendee).filter(
|
||||
CalendarEventAttendee.event_id == event_id,
|
||||
CalendarEventAttendee.user_id == user_id,
|
||||
)
|
||||
)
|
||||
att = result.scalars().first()
|
||||
if not att:
|
||||
return None
|
||||
|
||||
att.status = status
|
||||
att.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
return CalendarEventAttendeeModel.model_validate(att)
|
||||
|
||||
async def get_attendees_by_event(
|
||||
self, event_id: str, db: Optional[AsyncSession] = None
|
||||
) -> list[CalendarEventAttendeeModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(CalendarEventAttendee).filter(CalendarEventAttendee.event_id == event_id))
|
||||
return [CalendarEventAttendeeModel.model_validate(r) for r in result.scalars().all()]
|
||||
|
||||
async def get_events_by_attendee(self, user_id: str, db: Optional[AsyncSession] = None) -> list[str]:
|
||||
"""Return event IDs where user is an attendee."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(CalendarEventAttendee.event_id).filter(CalendarEventAttendee.user_id == user_id)
|
||||
)
|
||||
return [r[0] for r in result.all()]
|
||||
|
||||
|
||||
Calendars = CalendarTable()
|
||||
CalendarEvents = CalendarEventTable()
|
||||
CalendarEventAttendees = CalendarEventAttendeeTable()
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,600 @@
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Optional
|
||||
|
||||
from sqlalchemy import select, delete, func, cast, Integer
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.utils.response import normalize_usage
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Boolean,
|
||||
Column,
|
||||
ForeignKey,
|
||||
Text,
|
||||
JSON,
|
||||
Index,
|
||||
)
|
||||
|
||||
####################
|
||||
# Helpers
|
||||
####################
|
||||
|
||||
|
||||
def _normalize_timestamp(timestamp: int) -> float:
|
||||
"""Normalize and validate timestamp. Returns current time if invalid."""
|
||||
now = time.time()
|
||||
|
||||
# Convert milliseconds to seconds if needed
|
||||
if timestamp > 10_000_000_000:
|
||||
timestamp = timestamp / 1000
|
||||
|
||||
# Validate: must be after 2020 and not in the future (with 1 day tolerance)
|
||||
min_valid = 1577836800 # 2020-01-01 00:00:00 UTC
|
||||
max_valid = now + 86400 # 1 day in the future (clock skew tolerance)
|
||||
|
||||
if timestamp < min_valid or timestamp > max_valid:
|
||||
return now
|
||||
|
||||
return timestamp
|
||||
|
||||
|
||||
def get_usage(data: dict) -> Optional[dict]:
|
||||
"""Extract and normalize usage from message data."""
|
||||
usage = data.get('usage') or (data.get('info') or {}).get('usage')
|
||||
return normalize_usage(usage) if usage else None
|
||||
|
||||
|
||||
####################
|
||||
# ChatMessage DB Schema
|
||||
####################
|
||||
|
||||
|
||||
class ChatMessage(Base):
|
||||
__tablename__ = 'chat_message'
|
||||
|
||||
# Identity
|
||||
id = Column(Text, primary_key=True)
|
||||
chat_id = Column(Text, ForeignKey('chat.id', ondelete='CASCADE'), nullable=False, index=True)
|
||||
user_id = Column(Text, index=True)
|
||||
|
||||
# Structure
|
||||
role = Column(Text, nullable=False) # user, assistant, system
|
||||
parent_id = Column(Text, nullable=True)
|
||||
|
||||
# Content
|
||||
content = Column(JSON, nullable=True) # Can be str or list of blocks
|
||||
output = Column(JSON, nullable=True)
|
||||
|
||||
# Model (for assistant messages)
|
||||
model_id = Column(Text, nullable=True, index=True)
|
||||
|
||||
# Attachments
|
||||
files = Column(JSON, nullable=True)
|
||||
sources = Column(JSON, nullable=True)
|
||||
embeds = Column(JSON, nullable=True)
|
||||
|
||||
# Status
|
||||
done = Column(Boolean, default=True)
|
||||
status_history = Column(JSON, nullable=True)
|
||||
error = Column(JSON, nullable=True)
|
||||
|
||||
# Usage (tokens, timing, etc.)
|
||||
usage = Column(JSON, nullable=True)
|
||||
|
||||
# Timestamps
|
||||
created_at = Column(BigInteger, index=True)
|
||||
updated_at = Column(BigInteger)
|
||||
|
||||
__table_args__ = (
|
||||
Index('chat_message_chat_parent_idx', 'chat_id', 'parent_id'),
|
||||
Index('chat_message_model_created_idx', 'model_id', 'created_at'),
|
||||
Index('chat_message_user_created_idx', 'user_id', 'created_at'),
|
||||
)
|
||||
|
||||
|
||||
####################
|
||||
# Pydantic Models
|
||||
####################
|
||||
|
||||
|
||||
class ChatMessageModel(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
id: str
|
||||
chat_id: str
|
||||
user_id: str
|
||||
role: str
|
||||
parent_id: Optional[str] = None
|
||||
content: Optional[Any] = None # str or list of blocks
|
||||
output: Optional[list] = None
|
||||
model_id: Optional[str] = None
|
||||
files: Optional[list] = None
|
||||
sources: Optional[list] = None
|
||||
embeds: Optional[list] = None
|
||||
done: bool = True
|
||||
status_history: Optional[list] = None
|
||||
error: Optional[dict | str] = None
|
||||
usage: Optional[dict] = None
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
|
||||
####################
|
||||
# Table Operations
|
||||
####################
|
||||
|
||||
|
||||
class ChatMessageTable:
|
||||
async def upsert_message(
|
||||
self,
|
||||
message_id: str,
|
||||
chat_id: str,
|
||||
user_id: str,
|
||||
data: dict,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[ChatMessageModel]:
|
||||
"""Insert or update a chat message."""
|
||||
async with get_async_db_context(db) as db:
|
||||
now = int(time.time())
|
||||
timestamp = data.get('timestamp', now)
|
||||
|
||||
# Use composite ID: {chat_id}-{message_id}
|
||||
composite_id = f'{chat_id}-{message_id}'
|
||||
|
||||
existing = await db.get(ChatMessage, composite_id)
|
||||
if existing:
|
||||
# Update existing
|
||||
if 'role' in data:
|
||||
existing.role = data['role']
|
||||
if 'parent_id' in data:
|
||||
existing.parent_id = data.get('parent_id') or data.get('parentId')
|
||||
if 'content' in data:
|
||||
existing.content = data.get('content')
|
||||
if 'output' in data:
|
||||
existing.output = data.get('output')
|
||||
if 'model_id' in data or 'model' in data:
|
||||
existing.model_id = data.get('model_id') or data.get('model')
|
||||
if 'files' in data:
|
||||
existing.files = data.get('files')
|
||||
if 'sources' in data:
|
||||
existing.sources = data.get('sources')
|
||||
if 'embeds' in data:
|
||||
existing.embeds = data.get('embeds')
|
||||
if 'done' in data:
|
||||
existing.done = data.get('done', True)
|
||||
if 'status_history' in data or 'statusHistory' in data:
|
||||
existing.status_history = data.get('status_history') or data.get('statusHistory')
|
||||
if 'error' in data:
|
||||
existing.error = data.get('error')
|
||||
# Extract and normalize usage
|
||||
usage = get_usage(data)
|
||||
if usage:
|
||||
# Deep-merge: preserve existing keys not present in new data
|
||||
# This prevents background tasks (follow-ups, title, tags)
|
||||
# from accidentally clearing the primary response's token counts
|
||||
existing.usage = {**(existing.usage or {}), **usage}
|
||||
existing.updated_at = now
|
||||
await db.commit()
|
||||
await db.refresh(existing)
|
||||
return ChatMessageModel.model_validate(existing)
|
||||
else:
|
||||
# Insert new
|
||||
# Extract and normalize usage
|
||||
usage = get_usage(data)
|
||||
message = ChatMessage(
|
||||
id=composite_id,
|
||||
chat_id=chat_id,
|
||||
user_id=user_id,
|
||||
role=data.get('role', 'user'),
|
||||
parent_id=data.get('parent_id') or data.get('parentId'),
|
||||
content=data.get('content'),
|
||||
output=data.get('output'),
|
||||
model_id=data.get('model_id') or data.get('model'),
|
||||
files=data.get('files'),
|
||||
sources=data.get('sources'),
|
||||
embeds=data.get('embeds'),
|
||||
done=data.get('done', True),
|
||||
status_history=data.get('status_history') or data.get('statusHistory'),
|
||||
error=data.get('error'),
|
||||
usage=usage,
|
||||
created_at=timestamp,
|
||||
updated_at=now,
|
||||
)
|
||||
db.add(message)
|
||||
await db.commit()
|
||||
await db.refresh(message)
|
||||
return ChatMessageModel.model_validate(message)
|
||||
|
||||
async def get_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ChatMessageModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
message = await db.get(ChatMessage, id)
|
||||
return ChatMessageModel.model_validate(message) if message else None
|
||||
|
||||
async def get_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[ChatMessageModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(ChatMessage).filter_by(chat_id=chat_id).order_by(ChatMessage.created_at.asc())
|
||||
)
|
||||
messages = result.scalars().all()
|
||||
return [ChatMessageModel.model_validate(message) for message in messages]
|
||||
|
||||
async def get_messages_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[ChatMessageModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(ChatMessage)
|
||||
.filter_by(user_id=user_id)
|
||||
.order_by(ChatMessage.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
)
|
||||
messages = result.scalars().all()
|
||||
return [ChatMessageModel.model_validate(message) for message in messages]
|
||||
|
||||
async def get_messages_by_model_id(
|
||||
self,
|
||||
model_id: str,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[ChatMessageModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(ChatMessage).filter_by(model_id=model_id)
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
stmt = stmt.order_by(ChatMessage.created_at.desc()).offset(skip).limit(limit)
|
||||
result = await db.execute(stmt)
|
||||
messages = result.scalars().all()
|
||||
return [ChatMessageModel.model_validate(message) for message in messages]
|
||||
|
||||
async def get_chat_ids_by_model_id(
|
||||
self,
|
||||
model_id: str,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[str]:
|
||||
"""Get distinct chat_ids that used a specific model."""
|
||||
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(
|
||||
ChatMessage.chat_id,
|
||||
func.max(ChatMessage.created_at).label('last_message_at'),
|
||||
).filter(ChatMessage.model_id == model_id)
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
|
||||
# Group by chat_id and order by most recent message in each chat
|
||||
# Secondary sort on chat_id ensures deterministic pagination
|
||||
stmt = (
|
||||
stmt.group_by(ChatMessage.chat_id)
|
||||
.order_by(func.max(ChatMessage.created_at).desc(), ChatMessage.chat_id)
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
)
|
||||
result = await db.execute(stmt)
|
||||
chat_ids = result.all()
|
||||
return [chat_id for chat_id, _ in chat_ids]
|
||||
|
||||
async def delete_messages_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(ChatMessage).filter_by(chat_id=chat_id))
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
# Analytics methods
|
||||
async def get_message_count_by_model(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
group_id: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, int]:
|
||||
async with get_async_db_context(db) as db:
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
stmt = select(ChatMessage.model_id, func.count(ChatMessage.id).label('count')).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
ChatMessage.model_id.isnot(None),
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
if group_id:
|
||||
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
|
||||
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
|
||||
|
||||
stmt = stmt.group_by(ChatMessage.model_id)
|
||||
result = await db.execute(stmt)
|
||||
return {row.model_id: row.count for row in result.all()}
|
||||
|
||||
async def get_token_usage_by_model(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
group_id: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, dict]:
|
||||
"""Aggregate token usage by model using database-level aggregation."""
|
||||
async with get_async_db_context(db) as db:
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
# We need the dialect to determine JSON extraction syntax
|
||||
# For async sessions, access via get_bind()
|
||||
bind = await db.connection()
|
||||
dialect = bind.dialect.name
|
||||
|
||||
if dialect == 'sqlite':
|
||||
input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer)
|
||||
output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer)
|
||||
elif dialect == 'postgresql':
|
||||
input_tokens = cast(
|
||||
func.json_extract_path_text(ChatMessage.usage, 'input_tokens'),
|
||||
Integer,
|
||||
)
|
||||
output_tokens = cast(
|
||||
func.json_extract_path_text(ChatMessage.usage, 'output_tokens'),
|
||||
Integer,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f'Unsupported dialect: {dialect}')
|
||||
|
||||
stmt = select(
|
||||
ChatMessage.model_id,
|
||||
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
|
||||
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
|
||||
func.count(ChatMessage.id).label('message_count'),
|
||||
).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
ChatMessage.model_id.isnot(None),
|
||||
ChatMessage.usage.isnot(None),
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
if group_id:
|
||||
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
|
||||
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
|
||||
|
||||
stmt = stmt.group_by(ChatMessage.model_id)
|
||||
result = await db.execute(stmt)
|
||||
|
||||
return {
|
||||
row.model_id: {
|
||||
'input_tokens': row.input_tokens,
|
||||
'output_tokens': row.output_tokens,
|
||||
'total_tokens': row.input_tokens + row.output_tokens,
|
||||
'message_count': row.message_count,
|
||||
}
|
||||
for row in result.all()
|
||||
}
|
||||
|
||||
async def get_token_usage_by_user(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
group_id: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, dict]:
|
||||
"""Aggregate token usage by user using database-level aggregation."""
|
||||
async with get_async_db_context(db) as db:
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
bind = await db.connection()
|
||||
dialect = bind.dialect.name
|
||||
|
||||
if dialect == 'sqlite':
|
||||
input_tokens = cast(func.json_extract(ChatMessage.usage, '$.input_tokens'), Integer)
|
||||
output_tokens = cast(func.json_extract(ChatMessage.usage, '$.output_tokens'), Integer)
|
||||
elif dialect == 'postgresql':
|
||||
input_tokens = cast(
|
||||
func.json_extract_path_text(ChatMessage.usage, 'input_tokens'),
|
||||
Integer,
|
||||
)
|
||||
output_tokens = cast(
|
||||
func.json_extract_path_text(ChatMessage.usage, 'output_tokens'),
|
||||
Integer,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f'Unsupported dialect: {dialect}')
|
||||
|
||||
stmt = select(
|
||||
ChatMessage.user_id,
|
||||
func.coalesce(func.sum(input_tokens), 0).label('input_tokens'),
|
||||
func.coalesce(func.sum(output_tokens), 0).label('output_tokens'),
|
||||
func.count(ChatMessage.id).label('message_count'),
|
||||
).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
ChatMessage.user_id.isnot(None),
|
||||
ChatMessage.usage.isnot(None),
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
if group_id:
|
||||
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
|
||||
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
|
||||
|
||||
stmt = stmt.group_by(ChatMessage.user_id)
|
||||
result = await db.execute(stmt)
|
||||
|
||||
return {
|
||||
row.user_id: {
|
||||
'input_tokens': row.input_tokens,
|
||||
'output_tokens': row.output_tokens,
|
||||
'total_tokens': row.input_tokens + row.output_tokens,
|
||||
'message_count': row.message_count,
|
||||
}
|
||||
for row in result.all()
|
||||
}
|
||||
|
||||
async def get_message_count_by_user(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
group_id: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, int]:
|
||||
async with get_async_db_context(db) as db:
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
stmt = select(ChatMessage.user_id, func.count(ChatMessage.id).label('count')).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
if group_id:
|
||||
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
|
||||
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
|
||||
|
||||
stmt = stmt.group_by(ChatMessage.user_id)
|
||||
result = await db.execute(stmt)
|
||||
return {row.user_id: row.count for row in result.all()}
|
||||
|
||||
async def get_message_count_by_chat(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
group_id: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, int]:
|
||||
async with get_async_db_context(db) as db:
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
stmt = select(ChatMessage.chat_id, func.count(ChatMessage.id).label('count')).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
if group_id:
|
||||
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
|
||||
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
|
||||
|
||||
stmt = stmt.group_by(ChatMessage.chat_id)
|
||||
result = await db.execute(stmt)
|
||||
return {row.chat_id: row.count for row in result.all()}
|
||||
|
||||
async def get_daily_message_counts_by_model(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
group_id: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, dict[str, int]]:
|
||||
"""Get message counts grouped by day and model."""
|
||||
async with get_async_db_context(db) as db:
|
||||
from datetime import datetime, timedelta
|
||||
from open_webui.models.groups import GroupMember
|
||||
|
||||
stmt = select(ChatMessage.created_at, ChatMessage.model_id).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
ChatMessage.model_id.isnot(None),
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
if group_id:
|
||||
group_users = select(GroupMember.user_id).filter(GroupMember.group_id == group_id).scalar_subquery()
|
||||
stmt = stmt.filter(ChatMessage.user_id.in_(group_users))
|
||||
|
||||
result = await db.execute(stmt)
|
||||
results = result.all()
|
||||
|
||||
# Group by date -> model -> count
|
||||
daily_counts: dict[str, dict[str, int]] = {}
|
||||
for timestamp, model_id in results:
|
||||
date_str = datetime.fromtimestamp(_normalize_timestamp(timestamp)).strftime('%Y-%m-%d')
|
||||
if date_str not in daily_counts:
|
||||
daily_counts[date_str] = {}
|
||||
daily_counts[date_str][model_id] = daily_counts[date_str].get(model_id, 0) + 1
|
||||
|
||||
# Fill in missing days
|
||||
if start_date and end_date:
|
||||
current = datetime.fromtimestamp(_normalize_timestamp(start_date))
|
||||
end_dt = datetime.fromtimestamp(_normalize_timestamp(end_date))
|
||||
while current <= end_dt:
|
||||
date_str = current.strftime('%Y-%m-%d')
|
||||
if date_str not in daily_counts:
|
||||
daily_counts[date_str] = {}
|
||||
current += timedelta(days=1)
|
||||
|
||||
return daily_counts
|
||||
|
||||
async def get_hourly_message_counts_by_model(
|
||||
self,
|
||||
start_date: Optional[int] = None,
|
||||
end_date: Optional[int] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> dict[str, dict[str, int]]:
|
||||
"""Get message counts grouped by hour and model."""
|
||||
async with get_async_db_context(db) as db:
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
stmt = select(ChatMessage.created_at, ChatMessage.model_id).filter(
|
||||
ChatMessage.role == 'assistant',
|
||||
ChatMessage.model_id.isnot(None),
|
||||
)
|
||||
|
||||
if start_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at >= start_date)
|
||||
if end_date:
|
||||
stmt = stmt.filter(ChatMessage.created_at <= end_date)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
results = result.all()
|
||||
|
||||
# Group by hour -> model -> count
|
||||
hourly_counts: dict[str, dict[str, int]] = {}
|
||||
for timestamp, model_id in results:
|
||||
hour_str = datetime.fromtimestamp(_normalize_timestamp(timestamp)).strftime('%Y-%m-%d %H:00')
|
||||
if hour_str not in hourly_counts:
|
||||
hourly_counts[hour_str] = {}
|
||||
hourly_counts[hour_str][model_id] = hourly_counts[hour_str].get(model_id, 0) + 1
|
||||
|
||||
# Fill in missing hours
|
||||
if start_date and end_date:
|
||||
current = datetime.fromtimestamp(_normalize_timestamp(start_date)).replace(
|
||||
minute=0, second=0, microsecond=0
|
||||
)
|
||||
end_dt = datetime.fromtimestamp(_normalize_timestamp(end_date))
|
||||
while current <= end_dt:
|
||||
hour_str = current.strftime('%Y-%m-%d %H:00')
|
||||
if hour_str not in hourly_counts:
|
||||
hourly_counts[hour_str] = {}
|
||||
current += timedelta(hours=1)
|
||||
|
||||
return hourly_counts
|
||||
|
||||
|
||||
ChatMessages = ChatMessageTable()
|
||||
+835
-708
File diff suppressed because it is too large
Load Diff
@@ -3,9 +3,10 @@ import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from open_webui.models.users import User
|
||||
from sqlalchemy import select, delete, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.users import User, UserModel
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean
|
||||
@@ -19,7 +20,7 @@ log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Feedback(Base):
|
||||
__tablename__ = "feedback"
|
||||
__tablename__ = 'feedback'
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
user_id = Column(Text)
|
||||
version = Column(BigInteger, default=0)
|
||||
@@ -81,7 +82,7 @@ class RatingData(BaseModel):
|
||||
sibling_model_ids: Optional[list[str]] = None
|
||||
reason: Optional[str] = None
|
||||
comment: Optional[str] = None
|
||||
model_config = ConfigDict(extra="allow", protected_namespaces=())
|
||||
model_config = ConfigDict(extra='allow', protected_namespaces=())
|
||||
|
||||
|
||||
class MetaData(BaseModel):
|
||||
@@ -89,12 +90,12 @@ class MetaData(BaseModel):
|
||||
chat_id: Optional[str] = None
|
||||
message_id: Optional[str] = None
|
||||
tags: Optional[list[str]] = None
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
|
||||
class SnapshotData(BaseModel):
|
||||
chat: Optional[dict] = None
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
|
||||
class FeedbackForm(BaseModel):
|
||||
@@ -102,14 +103,14 @@ class FeedbackForm(BaseModel):
|
||||
data: Optional[RatingData] = None
|
||||
meta: Optional[dict] = None
|
||||
snapshot: Optional[SnapshotData] = None
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
|
||||
class UserResponse(BaseModel):
|
||||
id: str
|
||||
name: str
|
||||
email: str
|
||||
role: str = "pending"
|
||||
role: str = 'pending'
|
||||
|
||||
last_active_at: int # timestamp in epoch
|
||||
updated_at: int # timestamp in epoch
|
||||
@@ -139,139 +140,148 @@ class ModelHistoryResponse(BaseModel):
|
||||
|
||||
|
||||
class FeedbackTable:
|
||||
def insert_new_feedback(
|
||||
self, user_id: str, form_data: FeedbackForm, db: Optional[Session] = None
|
||||
async def insert_new_feedback(
|
||||
self, user_id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FeedbackModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
id = str(uuid.uuid4())
|
||||
feedback = FeedbackModel(
|
||||
**{
|
||||
"id": id,
|
||||
"user_id": user_id,
|
||||
"version": 0,
|
||||
'id': id,
|
||||
'user_id': user_id,
|
||||
'version': 0,
|
||||
**form_data.model_dump(),
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
try:
|
||||
result = Feedback(**feedback.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return FeedbackModel.model_validate(result)
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Error creating a new feedback: {e}")
|
||||
log.exception(f'Error creating a new feedback: {e}')
|
||||
return None
|
||||
|
||||
def get_feedback_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[FeedbackModel]:
|
||||
async def get_feedback_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FeedbackModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
feedback = db.query(Feedback).filter_by(id=id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(id=id))
|
||||
feedback = result.scalars().first()
|
||||
if not feedback:
|
||||
return None
|
||||
return FeedbackModel.model_validate(feedback)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_feedback_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_feedback_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FeedbackModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id))
|
||||
feedback = result.scalars().first()
|
||||
if not feedback:
|
||||
return None
|
||||
return FeedbackModel.model_validate(feedback)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_feedback_items(
|
||||
async def get_feedbacks_by_chat_id(self, chat_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
|
||||
"""Get all feedbacks for a specific chat."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
# meta.chat_id stores the chat reference
|
||||
result = await db.execute(
|
||||
select(Feedback)
|
||||
.filter(Feedback.meta['chat_id'].as_string() == chat_id)
|
||||
.order_by(Feedback.created_at.desc())
|
||||
)
|
||||
feedbacks = result.scalars().all()
|
||||
return [FeedbackModel.model_validate(fb) for fb in feedbacks]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
async def get_feedback_items(
|
||||
self,
|
||||
filter: dict = {},
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> FeedbackListResponse:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Feedback, User).join(User, Feedback.user_id == User.id)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Feedback, User).join(User, Feedback.user_id == User.id)
|
||||
|
||||
if filter:
|
||||
order_by = filter.get("order_by")
|
||||
direction = filter.get("direction")
|
||||
# Apply model_id filter (exact match)
|
||||
model_id = filter.get('model_id')
|
||||
if model_id:
|
||||
stmt = stmt.filter(Feedback.data['model_id'].as_string() == model_id)
|
||||
|
||||
if order_by == "username":
|
||||
if direction == "asc":
|
||||
query = query.order_by(User.name.asc())
|
||||
order_by = filter.get('order_by')
|
||||
direction = filter.get('direction')
|
||||
|
||||
if order_by == 'username':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(User.name.asc())
|
||||
else:
|
||||
query = query.order_by(User.name.desc())
|
||||
elif order_by == "model_id":
|
||||
# it's stored in feedback.data['model_id']
|
||||
if direction == "asc":
|
||||
query = query.order_by(
|
||||
Feedback.data["model_id"].as_string().asc()
|
||||
)
|
||||
stmt = stmt.order_by(User.name.desc())
|
||||
elif order_by == 'model_id':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Feedback.data['model_id'].as_string().asc())
|
||||
else:
|
||||
query = query.order_by(
|
||||
Feedback.data["model_id"].as_string().desc()
|
||||
)
|
||||
elif order_by == "rating":
|
||||
# it's stored in feedback.data['rating']
|
||||
if direction == "asc":
|
||||
query = query.order_by(
|
||||
Feedback.data["rating"].as_string().asc()
|
||||
)
|
||||
stmt = stmt.order_by(Feedback.data['model_id'].as_string().desc())
|
||||
elif order_by == 'rating':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Feedback.data['rating'].as_string().asc())
|
||||
else:
|
||||
query = query.order_by(
|
||||
Feedback.data["rating"].as_string().desc()
|
||||
)
|
||||
elif order_by == "updated_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Feedback.updated_at.asc())
|
||||
stmt = stmt.order_by(Feedback.data['rating'].as_string().desc())
|
||||
elif order_by == 'updated_at':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Feedback.updated_at.asc())
|
||||
else:
|
||||
query = query.order_by(Feedback.updated_at.desc())
|
||||
stmt = stmt.order_by(Feedback.updated_at.desc())
|
||||
|
||||
else:
|
||||
query = query.order_by(Feedback.created_at.desc())
|
||||
stmt = stmt.order_by(Feedback.created_at.desc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
total = query.count()
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
items = query.all()
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
feedbacks = []
|
||||
for feedback, user in items:
|
||||
feedback_model = FeedbackModel.model_validate(feedback)
|
||||
user_model = UserResponse.model_validate(user)
|
||||
feedbacks.append(
|
||||
FeedbackUserResponse(**feedback_model.model_dump(), user=user_model)
|
||||
)
|
||||
feedbacks.append(FeedbackUserResponse(**feedback_model.model_dump(), user=user_model))
|
||||
|
||||
return FeedbackListResponse(items=feedbacks, total=total)
|
||||
|
||||
def get_all_feedbacks(self, db: Optional[Session] = None) -> list[FeedbackModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FeedbackModel.model_validate(feedback)
|
||||
for feedback in db.query(Feedback)
|
||||
.order_by(Feedback.updated_at.desc())
|
||||
.all()
|
||||
]
|
||||
async def get_all_feedbacks(self, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).order_by(Feedback.updated_at.desc()))
|
||||
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
|
||||
|
||||
def get_all_feedback_ids(
|
||||
self, db: Optional[Session] = None
|
||||
) -> list[FeedbackIdResponse]:
|
||||
with get_db_context(db) as db:
|
||||
async def get_all_feedback_ids(self, db: Optional[AsyncSession] = None) -> list[FeedbackIdResponse]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Feedback.id, Feedback.user_id, Feedback.created_at, Feedback.updated_at).order_by(
|
||||
Feedback.updated_at.desc()
|
||||
)
|
||||
)
|
||||
return [
|
||||
FeedbackIdResponse(
|
||||
id=row.id,
|
||||
@@ -279,28 +289,28 @@ class FeedbackTable:
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
for row in db.query(
|
||||
Feedback.id,
|
||||
Feedback.user_id,
|
||||
Feedback.created_at,
|
||||
Feedback.updated_at,
|
||||
)
|
||||
.order_by(Feedback.updated_at.desc())
|
||||
.all()
|
||||
for row in result.all()
|
||||
]
|
||||
|
||||
def get_feedbacks_for_leaderboard(
|
||||
self, db: Optional[Session] = None
|
||||
) -> list[LeaderboardFeedbackData]:
|
||||
async def get_distinct_model_ids(self, db: Optional[AsyncSession] = None) -> list[str]:
|
||||
"""Get distinct model_ids from feedback data for filter dropdowns."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Feedback.data['model_id'].as_string())
|
||||
.filter(Feedback.data['model_id'].as_string().isnot(None))
|
||||
.distinct()
|
||||
)
|
||||
rows = result.all()
|
||||
return sorted([row[0] for row in rows if row[0]])
|
||||
|
||||
async def get_feedbacks_for_leaderboard(self, db: Optional[AsyncSession] = None) -> list[LeaderboardFeedbackData]:
|
||||
"""Fetch only id and data for leaderboard computation (excludes snapshot/meta)."""
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
LeaderboardFeedbackData(id=row.id, data=row.data)
|
||||
for row in db.query(Feedback.id, Feedback.data).all()
|
||||
]
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback.id, Feedback.data))
|
||||
return [LeaderboardFeedbackData(id=row.id, data=row.data) for row in result.all()]
|
||||
|
||||
def get_model_evaluation_history(
|
||||
self, model_id: str, days: int = 30, db: Optional[Session] = None
|
||||
async def get_model_evaluation_history(
|
||||
self, model_id: str, days: int = 30, db: Optional[AsyncSession] = None
|
||||
) -> list[ModelHistoryEntry]:
|
||||
"""
|
||||
Get daily wins/losses for a specific model over the past N days.
|
||||
@@ -310,36 +320,35 @@ class FeedbackTable:
|
||||
from datetime import datetime, timedelta
|
||||
from collections import defaultdict
|
||||
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
if days == 0:
|
||||
# All time - no cutoff
|
||||
rows = db.query(Feedback.created_at, Feedback.data).all()
|
||||
result = await db.execute(select(Feedback.created_at, Feedback.data))
|
||||
else:
|
||||
cutoff = int(time.time()) - (days * 86400)
|
||||
rows = (
|
||||
db.query(Feedback.created_at, Feedback.data)
|
||||
.filter(Feedback.created_at >= cutoff)
|
||||
.all()
|
||||
result = await db.execute(
|
||||
select(Feedback.created_at, Feedback.data).filter(Feedback.created_at >= cutoff)
|
||||
)
|
||||
rows = result.all()
|
||||
|
||||
daily_counts = defaultdict(lambda: {"won": 0, "lost": 0})
|
||||
daily_counts = defaultdict(lambda: {'won': 0, 'lost': 0})
|
||||
first_date = None
|
||||
|
||||
for created_at, data in rows:
|
||||
if not data:
|
||||
continue
|
||||
if data.get("model_id") != model_id:
|
||||
if data.get('model_id') != model_id:
|
||||
continue
|
||||
|
||||
rating_str = str(data.get("rating", ""))
|
||||
if rating_str not in ("1", "-1"):
|
||||
rating_str = str(data.get('rating', ''))
|
||||
if rating_str not in ('1', '-1'):
|
||||
continue
|
||||
|
||||
date_str = datetime.fromtimestamp(created_at).strftime("%Y-%m-%d")
|
||||
if rating_str == "1":
|
||||
daily_counts[date_str]["won"] += 1
|
||||
date_str = datetime.fromtimestamp(created_at).strftime('%Y-%m-%d')
|
||||
if rating_str == '1':
|
||||
daily_counts[date_str]['won'] += 1
|
||||
else:
|
||||
daily_counts[date_str]["lost"] += 1
|
||||
daily_counts[date_str]['lost'] += 1
|
||||
|
||||
# Track first date for this model
|
||||
if first_date is None or date_str < first_date:
|
||||
@@ -351,7 +360,7 @@ class FeedbackTable:
|
||||
|
||||
if days == 0 and first_date:
|
||||
# All time: start from first feedback date
|
||||
start_date = datetime.strptime(first_date, "%Y-%m-%d").date()
|
||||
start_date = datetime.strptime(first_date, '%Y-%m-%d').date()
|
||||
num_days = (today - start_date).days + 1
|
||||
else:
|
||||
# Fixed range
|
||||
@@ -360,43 +369,28 @@ class FeedbackTable:
|
||||
|
||||
for i in range(num_days):
|
||||
d = start_date + timedelta(days=i)
|
||||
date_str = d.strftime("%Y-%m-%d")
|
||||
counts = daily_counts.get(date_str, {"won": 0, "lost": 0})
|
||||
result.append(
|
||||
ModelHistoryEntry(date=date_str, won=counts["won"], lost=counts["lost"])
|
||||
)
|
||||
date_str = d.strftime('%Y-%m-%d')
|
||||
counts = daily_counts.get(date_str, {'won': 0, 'lost': 0})
|
||||
result.append(ModelHistoryEntry(date=date_str, won=counts['won'], lost=counts['lost']))
|
||||
|
||||
return result
|
||||
|
||||
def get_feedbacks_by_type(
|
||||
self, type: str, db: Optional[Session] = None
|
||||
) -> list[FeedbackModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FeedbackModel.model_validate(feedback)
|
||||
for feedback in db.query(Feedback)
|
||||
.filter_by(type=type)
|
||||
.order_by(Feedback.updated_at.desc())
|
||||
.all()
|
||||
]
|
||||
async def get_feedbacks_by_type(self, type: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(type=type).order_by(Feedback.updated_at.desc()))
|
||||
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
|
||||
|
||||
def get_feedbacks_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> list[FeedbackModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FeedbackModel.model_validate(feedback)
|
||||
for feedback in db.query(Feedback)
|
||||
.filter_by(user_id=user_id)
|
||||
.order_by(Feedback.updated_at.desc())
|
||||
.all()
|
||||
]
|
||||
async def get_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FeedbackModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(user_id=user_id).order_by(Feedback.updated_at.desc()))
|
||||
return [FeedbackModel.model_validate(feedback) for feedback in result.scalars().all()]
|
||||
|
||||
def update_feedback_by_id(
|
||||
self, id: str, form_data: FeedbackForm, db: Optional[Session] = None
|
||||
async def update_feedback_by_id(
|
||||
self, id: str, form_data: FeedbackForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FeedbackModel]:
|
||||
with get_db_context(db) as db:
|
||||
feedback = db.query(Feedback).filter_by(id=id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(id=id))
|
||||
feedback = result.scalars().first()
|
||||
if not feedback:
|
||||
return None
|
||||
|
||||
@@ -409,18 +403,19 @@ class FeedbackTable:
|
||||
|
||||
feedback.updated_at = int(time.time())
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return FeedbackModel.model_validate(feedback)
|
||||
|
||||
def update_feedback_by_id_and_user_id(
|
||||
async def update_feedback_by_id_and_user_id(
|
||||
self,
|
||||
id: str,
|
||||
user_id: str,
|
||||
form_data: FeedbackForm,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[FeedbackModel]:
|
||||
with get_db_context(db) as db:
|
||||
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id))
|
||||
feedback = result.scalars().first()
|
||||
if not feedback:
|
||||
return None
|
||||
|
||||
@@ -433,50 +428,40 @@ class FeedbackTable:
|
||||
|
||||
feedback.updated_at = int(time.time())
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return FeedbackModel.model_validate(feedback)
|
||||
|
||||
def delete_feedback_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
feedback = db.query(Feedback).filter_by(id=id).first()
|
||||
async def delete_feedback_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(id=id))
|
||||
feedback = result.scalars().first()
|
||||
if not feedback:
|
||||
return False
|
||||
db.delete(feedback)
|
||||
db.commit()
|
||||
await db.delete(feedback)
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
def delete_feedback_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
feedback = db.query(Feedback).filter_by(id=id, user_id=user_id).first()
|
||||
async def delete_feedback_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Feedback).filter_by(id=id, user_id=user_id))
|
||||
feedback = result.scalars().first()
|
||||
if not feedback:
|
||||
return False
|
||||
db.delete(feedback)
|
||||
db.commit()
|
||||
await db.delete(feedback)
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
def delete_feedbacks_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
feedbacks = db.query(Feedback).filter_by(user_id=user_id).all()
|
||||
if not feedbacks:
|
||||
return False
|
||||
for feedback in feedbacks:
|
||||
db.delete(feedback)
|
||||
db.commit()
|
||||
return True
|
||||
async def delete_feedbacks_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(delete(Feedback).filter_by(user_id=user_id))
|
||||
await db.commit()
|
||||
return result.rowcount > 0
|
||||
|
||||
def delete_all_feedbacks(self, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
feedbacks = db.query(Feedback).all()
|
||||
if not feedbacks:
|
||||
return False
|
||||
for feedback in feedbacks:
|
||||
db.delete(feedback)
|
||||
db.commit()
|
||||
return True
|
||||
async def delete_all_feedbacks(self, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(delete(Feedback))
|
||||
await db.commit()
|
||||
return result.rowcount > 0
|
||||
|
||||
|
||||
Feedbacks = FeedbackTable()
|
||||
|
||||
+163
-135
@@ -2,20 +2,24 @@ import logging
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import select, delete, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.utils.misc import sanitize_metadata
|
||||
from pydantic import BaseModel, ConfigDict, model_validator
|
||||
from sqlalchemy import BigInteger, Column, String, Text, JSON
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
####################
|
||||
# Files DB Schema
|
||||
# What is written here bears witness. Let the testimony
|
||||
# remain as it was given, and let none tamper with it.
|
||||
####################
|
||||
|
||||
|
||||
class File(Base):
|
||||
__tablename__ = "file"
|
||||
__tablename__ = 'file'
|
||||
id = Column(String, primary_key=True, unique=True)
|
||||
user_id = Column(String)
|
||||
hash = Column(Text, nullable=True)
|
||||
@@ -26,8 +30,6 @@ class File(Base):
|
||||
data = Column(JSON, nullable=True)
|
||||
meta = Column(JSON, nullable=True)
|
||||
|
||||
access_control = Column(JSON, nullable=True)
|
||||
|
||||
created_at = Column(BigInteger)
|
||||
updated_at = Column(BigInteger)
|
||||
|
||||
@@ -45,8 +47,6 @@ class FileModel(BaseModel):
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
|
||||
access_control: Optional[dict] = None
|
||||
|
||||
created_at: Optional[int] # timestamp in epoch
|
||||
updated_at: Optional[int] # timestamp in epoch
|
||||
|
||||
@@ -61,7 +61,24 @@ class FileMeta(BaseModel):
|
||||
content_type: Optional[str] = None
|
||||
size: Optional[int] = None
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
@model_validator(mode='before')
|
||||
@classmethod
|
||||
def sanitize_meta(cls, data):
|
||||
"""Sanitize metadata fields to handle malformed legacy data."""
|
||||
if not isinstance(data, dict):
|
||||
return data
|
||||
|
||||
# Handle content_type that may be a list like ['application/pdf', None]
|
||||
content_type = data.get('content_type')
|
||||
if isinstance(content_type, list):
|
||||
# Extract first non-None string value
|
||||
data['content_type'] = next((item for item in content_type if isinstance(item, str)), None)
|
||||
elif content_type is not None and not isinstance(content_type, str):
|
||||
data['content_type'] = None
|
||||
|
||||
return data
|
||||
|
||||
|
||||
class FileModelResponse(BaseModel):
|
||||
@@ -71,12 +88,12 @@ class FileModelResponse(BaseModel):
|
||||
|
||||
filename: str
|
||||
data: Optional[dict] = None
|
||||
meta: FileMeta
|
||||
meta: Optional[FileMeta] = None
|
||||
|
||||
created_at: int # timestamp in epoch
|
||||
updated_at: int # timestamp in epoch
|
||||
updated_at: Optional[int] = None # timestamp in epoch, optional for legacy files
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
|
||||
class FileMetadataResponse(BaseModel):
|
||||
@@ -87,6 +104,11 @@ class FileMetadataResponse(BaseModel):
|
||||
updated_at: int # timestamp in epoch
|
||||
|
||||
|
||||
class FileListResponse(BaseModel):
|
||||
items: list[FileModelResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class FileForm(BaseModel):
|
||||
id: str
|
||||
hash: Optional[str] = None
|
||||
@@ -94,7 +116,6 @@ class FileForm(BaseModel):
|
||||
path: str
|
||||
data: dict = {}
|
||||
meta: dict = {}
|
||||
access_control: Optional[dict] = None
|
||||
|
||||
|
||||
class FileUpdateForm(BaseModel):
|
||||
@@ -103,57 +124,58 @@ class FileUpdateForm(BaseModel):
|
||||
meta: Optional[dict] = None
|
||||
|
||||
|
||||
class FileListResponse(BaseModel):
|
||||
items: list[FileModel]
|
||||
total: int
|
||||
|
||||
|
||||
class FilesTable:
|
||||
def insert_new_file(
|
||||
self, user_id: str, form_data: FileForm, db: Optional[Session] = None
|
||||
async def insert_new_file(
|
||||
self, user_id: str, form_data: FileForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
file_data = form_data.model_dump()
|
||||
|
||||
# Sanitize meta to remove non-JSON-serializable objects
|
||||
# (e.g. callable tool functions, MCP client instances from middleware)
|
||||
if file_data.get('meta'):
|
||||
file_data['meta'] = sanitize_metadata(file_data['meta'])
|
||||
|
||||
file = FileModel(
|
||||
**{
|
||||
**form_data.model_dump(),
|
||||
"user_id": user_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
**file_data,
|
||||
'user_id': user_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
result = File(**file.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return FileModel.model_validate(result)
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Error inserting a new file: {e}")
|
||||
log.exception(f'Error inserting a new file: {e}')
|
||||
return None
|
||||
|
||||
def get_file_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[FileModel]:
|
||||
async def get_file_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FileModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.get(File, id)
|
||||
return FileModel.model_validate(file)
|
||||
file = await db.get(File, id)
|
||||
return FileModel.model_validate(file) if file else None
|
||||
except Exception:
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_file_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_file_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id, user_id=user_id).first()
|
||||
result = await db.execute(select(File).filter_by(id=id, user_id=user_id))
|
||||
file = result.scalars().first()
|
||||
if file:
|
||||
return FileModel.model_validate(file)
|
||||
else:
|
||||
@@ -161,12 +183,14 @@ class FilesTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_file_metadata_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
async def get_file_metadata_by_id(
|
||||
self, id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileMetadataResponse]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.get(File, id)
|
||||
file = await db.get(File, id)
|
||||
if not file:
|
||||
return None
|
||||
return FileMetadataResponse(
|
||||
id=file.id,
|
||||
hash=file.hash,
|
||||
@@ -177,14 +201,13 @@ class FilesTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_files(self, db: Optional[Session] = None) -> list[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [FileModel.model_validate(file) for file in db.query(File).all()]
|
||||
async def get_files(self, db: Optional[AsyncSession] = None) -> list[FileModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(File))
|
||||
return [FileModel.model_validate(file) for file in result.scalars().all()]
|
||||
|
||||
def check_access_by_user_id(
|
||||
self, id, user_id, permission="write", db: Optional[Session] = None
|
||||
) -> bool:
|
||||
file = self.get_file_by_id(id, db=db)
|
||||
async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool:
|
||||
file = await self.get_file_by_id(id, db=db)
|
||||
if not file:
|
||||
return False
|
||||
if file.user_id == user_id:
|
||||
@@ -192,46 +215,55 @@ class FilesTable:
|
||||
# Implement additional access control logic here as needed
|
||||
return False
|
||||
|
||||
def get_files_by_ids(
|
||||
self, ids: list[str], db: Optional[Session] = None
|
||||
) -> list[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FileModel.model_validate(file)
|
||||
for file in db.query(File)
|
||||
async def get_files_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FileModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(File).filter(File.id.in_(ids)).order_by(File.updated_at.desc()))
|
||||
return [FileModel.model_validate(file) for file in result.scalars().all()]
|
||||
|
||||
async def get_file_metadatas_by_ids(
|
||||
self, ids: list[str], db: Optional[AsyncSession] = None
|
||||
) -> list[FileMetadataResponse]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(File.id, File.hash, File.meta, File.created_at, File.updated_at)
|
||||
.filter(File.id.in_(ids))
|
||||
.order_by(File.updated_at.desc())
|
||||
.all()
|
||||
]
|
||||
|
||||
def get_file_metadatas_by_ids(
|
||||
self, ids: list[str], db: Optional[Session] = None
|
||||
) -> list[FileMetadataResponse]:
|
||||
with get_db_context(db) as db:
|
||||
)
|
||||
return [
|
||||
FileMetadataResponse(
|
||||
id=file.id,
|
||||
hash=file.hash,
|
||||
meta=file.meta,
|
||||
created_at=file.created_at,
|
||||
updated_at=file.updated_at,
|
||||
id=row.id,
|
||||
hash=row.hash,
|
||||
meta=row.meta,
|
||||
created_at=row.created_at,
|
||||
updated_at=row.updated_at,
|
||||
)
|
||||
for file in db.query(
|
||||
File.id, File.hash, File.meta, File.created_at, File.updated_at
|
||||
)
|
||||
.filter(File.id.in_(ids))
|
||||
.order_by(File.updated_at.desc())
|
||||
.all()
|
||||
for row in result.all()
|
||||
]
|
||||
|
||||
def get_files_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> list[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FileModel.model_validate(file)
|
||||
for file in db.query(File).filter_by(user_id=user_id).all()
|
||||
]
|
||||
async def get_files_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FileModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(File).filter_by(user_id=user_id))
|
||||
return [FileModel.model_validate(file) for file in result.scalars().all()]
|
||||
|
||||
async def get_file_list(
|
||||
self,
|
||||
user_id: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> 'FileListResponse':
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(File)
|
||||
if user_id:
|
||||
stmt = stmt.filter_by(user_id=user_id)
|
||||
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
result = await db.execute(stmt.order_by(File.updated_at.desc(), File.id.desc()).offset(skip).limit(limit))
|
||||
items = [FileModelResponse.model_validate(file, from_attributes=True) for file in result.scalars().all()]
|
||||
|
||||
return FileListResponse(items=items, total=total)
|
||||
|
||||
@staticmethod
|
||||
def _glob_to_like_pattern(glob: str) -> str:
|
||||
@@ -249,20 +281,20 @@ class FilesTable:
|
||||
A SQL LIKE compatible pattern with proper escaping.
|
||||
"""
|
||||
# Escape SQL special characters first, then convert glob wildcards
|
||||
pattern = glob.replace("\\", "\\\\")
|
||||
pattern = pattern.replace("%", "\\%")
|
||||
pattern = pattern.replace("_", "\\_")
|
||||
pattern = pattern.replace("*", "%")
|
||||
pattern = pattern.replace("?", "_")
|
||||
pattern = glob.replace('\\', '\\\\')
|
||||
pattern = pattern.replace('%', '\\%')
|
||||
pattern = pattern.replace('_', '\\_')
|
||||
pattern = pattern.replace('*', '%')
|
||||
pattern = pattern.replace('?', '_')
|
||||
return pattern
|
||||
|
||||
def search_files(
|
||||
async def search_files(
|
||||
self,
|
||||
user_id: Optional[str] = None,
|
||||
filename: str = "*",
|
||||
filename: str = '*',
|
||||
skip: int = 0,
|
||||
limit: int = 100,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[FileModel]:
|
||||
"""
|
||||
Search files with glob pattern matching, optional user filter, and pagination.
|
||||
@@ -275,32 +307,28 @@ class FilesTable:
|
||||
db: Optional database session.
|
||||
|
||||
Returns:
|
||||
List of matching FileModel objects, ordered by updated_at descending.
|
||||
List of matching FileModel objects, ordered by created_at descending.
|
||||
"""
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(File)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(File)
|
||||
|
||||
if user_id:
|
||||
query = query.filter_by(user_id=user_id)
|
||||
stmt = stmt.filter_by(user_id=user_id)
|
||||
|
||||
pattern = self._glob_to_like_pattern(filename)
|
||||
if pattern != "%":
|
||||
query = query.filter(File.filename.ilike(pattern, escape="\\"))
|
||||
if pattern != '%':
|
||||
stmt = stmt.filter(File.filename.ilike(pattern, escape='\\'))
|
||||
|
||||
return [
|
||||
FileModel.model_validate(file)
|
||||
for file in query.order_by(File.updated_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
]
|
||||
result = await db.execute(stmt.order_by(File.created_at.desc(), File.id.desc()).offset(skip).limit(limit))
|
||||
return [FileModel.model_validate(file) for file in result.scalars().all()]
|
||||
|
||||
def update_file_by_id(
|
||||
self, id: str, form_data: FileUpdateForm, db: Optional[Session] = None
|
||||
async def update_file_by_id(
|
||||
self, id: str, form_data: FileUpdateForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
result = await db.execute(select(File).filter_by(id=id))
|
||||
file = result.scalars().first()
|
||||
|
||||
if form_data.hash is not None:
|
||||
file.hash = form_data.hash
|
||||
@@ -312,70 +340,70 @@ class FilesTable:
|
||||
file.meta = {**(file.meta if file.meta else {}), **form_data.meta}
|
||||
|
||||
file.updated_at = int(time.time())
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return FileModel.model_validate(file)
|
||||
except Exception as e:
|
||||
log.exception(f"Error updating file completely by id: {e}")
|
||||
log.exception(f'Error updating file completely by id: {e}')
|
||||
return None
|
||||
|
||||
def update_file_hash_by_id(
|
||||
self, id: str, hash: Optional[str], db: Optional[Session] = None
|
||||
async def update_file_hash_by_id(
|
||||
self, id: str, hash: Optional[str], db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
result = await db.execute(select(File).filter_by(id=id))
|
||||
file = result.scalars().first()
|
||||
file.hash = hash
|
||||
file.updated_at = int(time.time())
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
return FileModel.model_validate(file)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def update_file_data_by_id(
|
||||
self, id: str, data: dict, db: Optional[Session] = None
|
||||
async def update_file_data_by_id(
|
||||
self, id: str, data: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
result = await db.execute(select(File).filter_by(id=id))
|
||||
file = result.scalars().first()
|
||||
file.data = {**(file.data if file.data else {}), **data}
|
||||
file.updated_at = int(time.time())
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return FileModel.model_validate(file)
|
||||
except Exception as e:
|
||||
|
||||
return None
|
||||
|
||||
def update_file_metadata_by_id(
|
||||
self, id: str, meta: dict, db: Optional[Session] = None
|
||||
async def update_file_metadata_by_id(
|
||||
self, id: str, meta: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
file = db.query(File).filter_by(id=id).first()
|
||||
result = await db.execute(select(File).filter_by(id=id))
|
||||
file = result.scalars().first()
|
||||
file.meta = {**(file.meta if file.meta else {}), **meta}
|
||||
file.updated_at = int(time.time())
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return FileModel.model_validate(file)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
return False
|
||||
|
||||
def delete_file_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_file_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(File).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(File).filter_by(id=id))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_all_files(self, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_all_files(self, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(File).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(File))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
|
||||
@@ -6,22 +6,23 @@ import re
|
||||
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean, func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean, func, select, delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
####################
|
||||
# Folder DB Schema
|
||||
# Let every room in this house shelter someone who needs it,
|
||||
# and let no chamber stand empty while there is want.
|
||||
####################
|
||||
|
||||
|
||||
class Folder(Base):
|
||||
__tablename__ = "folder"
|
||||
__tablename__ = 'folder'
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
parent_id = Column(Text, nullable=True)
|
||||
user_id = Column(Text)
|
||||
@@ -72,55 +73,57 @@ class FolderForm(BaseModel):
|
||||
name: str
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
model_config = ConfigDict(extra="allow")
|
||||
parent_id: Optional[str] = None
|
||||
model_config = ConfigDict(extra='forbid')
|
||||
|
||||
|
||||
class FolderUpdateForm(BaseModel):
|
||||
name: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='forbid')
|
||||
|
||||
|
||||
class FolderTable:
|
||||
def insert_new_folder(
|
||||
async def insert_new_folder(
|
||||
self,
|
||||
user_id: str,
|
||||
form_data: FolderForm,
|
||||
parent_id: Optional[str] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[FolderModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
id = str(uuid.uuid4())
|
||||
folder = FolderModel(
|
||||
**{
|
||||
"id": id,
|
||||
"user_id": user_id,
|
||||
'id': id,
|
||||
'user_id': user_id,
|
||||
**(form_data.model_dump(exclude_unset=True) or {}),
|
||||
"parent_id": parent_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
'parent_id': parent_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
try:
|
||||
result = Folder(**folder.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return FolderModel.model_validate(result)
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Error inserting a new folder: {e}")
|
||||
log.exception(f'Error inserting a new folder: {e}')
|
||||
return None
|
||||
|
||||
def get_folder_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_folder_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FolderModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
|
||||
folder = result.scalars().first()
|
||||
|
||||
if not folder:
|
||||
return None
|
||||
@@ -129,85 +132,75 @@ class FolderTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_children_folders_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_children_folders_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[list[FolderModel]]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
folders = []
|
||||
|
||||
def get_children(folder):
|
||||
children = self.get_folders_by_parent_id_and_user_id(
|
||||
folder.id, user_id, db=db
|
||||
)
|
||||
async def get_children(folder):
|
||||
children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
|
||||
for child in children:
|
||||
get_children(child)
|
||||
await get_children(child)
|
||||
folders.append(child)
|
||||
|
||||
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
|
||||
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
|
||||
folder = result.scalars().first()
|
||||
if not folder:
|
||||
return None
|
||||
|
||||
get_children(folder)
|
||||
await get_children(folder)
|
||||
return folders
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_folders_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> list[FolderModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FolderModel.model_validate(folder)
|
||||
for folder in db.query(Folder).filter_by(user_id=user_id).all()
|
||||
]
|
||||
async def get_folders_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[FolderModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(user_id=user_id))
|
||||
return [FolderModel.model_validate(folder) for folder in result.scalars().all()]
|
||||
|
||||
def get_folder_by_parent_id_and_user_id_and_name(
|
||||
async def get_folder_by_parent_id_and_user_id_and_name(
|
||||
self,
|
||||
parent_id: Optional[str],
|
||||
user_id: str,
|
||||
name: str,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[FolderModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Check if folder exists
|
||||
folder = (
|
||||
db.query(Folder)
|
||||
.filter_by(parent_id=parent_id, user_id=user_id)
|
||||
.filter(Folder.name.ilike(name))
|
||||
.first()
|
||||
result = await db.execute(
|
||||
select(Folder).filter_by(parent_id=parent_id, user_id=user_id).filter(Folder.name.ilike(name))
|
||||
)
|
||||
folder = result.scalars().first()
|
||||
|
||||
if not folder:
|
||||
return None
|
||||
|
||||
return FolderModel.model_validate(folder)
|
||||
except Exception as e:
|
||||
log.error(f"get_folder_by_parent_id_and_user_id_and_name: {e}")
|
||||
log.error(f'get_folder_by_parent_id_and_user_id_and_name: {e}')
|
||||
return None
|
||||
|
||||
def get_folders_by_parent_id_and_user_id(
|
||||
self, parent_id: Optional[str], user_id: str, db: Optional[Session] = None
|
||||
async def get_folders_by_parent_id_and_user_id(
|
||||
self, parent_id: Optional[str], user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> list[FolderModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FolderModel.model_validate(folder)
|
||||
for folder in db.query(Folder)
|
||||
.filter_by(parent_id=parent_id, user_id=user_id)
|
||||
.all()
|
||||
]
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(parent_id=parent_id, user_id=user_id))
|
||||
return [FolderModel.model_validate(folder) for folder in result.scalars().all()]
|
||||
|
||||
def update_folder_parent_id_by_id_and_user_id(
|
||||
async def update_folder_parent_id_by_id_and_user_id(
|
||||
self,
|
||||
id: str,
|
||||
user_id: str,
|
||||
parent_id: str,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[FolderModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
|
||||
folder = result.scalars().first()
|
||||
|
||||
if not folder:
|
||||
return None
|
||||
@@ -215,69 +208,70 @@ class FolderTable:
|
||||
folder.parent_id = parent_id
|
||||
folder.updated_at = int(time.time())
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
return FolderModel.model_validate(folder)
|
||||
except Exception as e:
|
||||
log.error(f"update_folder: {e}")
|
||||
log.error(f'update_folder: {e}')
|
||||
return
|
||||
|
||||
def update_folder_by_id_and_user_id(
|
||||
async def update_folder_by_id_and_user_id(
|
||||
self,
|
||||
id: str,
|
||||
user_id: str,
|
||||
form_data: FolderUpdateForm,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[FolderModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
|
||||
folder = result.scalars().first()
|
||||
|
||||
if not folder:
|
||||
return None
|
||||
|
||||
form_data = form_data.model_dump(exclude_unset=True)
|
||||
|
||||
existing_folder = (
|
||||
db.query(Folder)
|
||||
.filter_by(
|
||||
name=form_data.get("name"),
|
||||
existing_result = await db.execute(
|
||||
select(Folder).filter_by(
|
||||
name=form_data.get('name'),
|
||||
parent_id=folder.parent_id,
|
||||
user_id=user_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
existing_folder = existing_result.scalars().first()
|
||||
|
||||
if existing_folder and existing_folder.id != id:
|
||||
return None
|
||||
|
||||
folder.name = form_data.get("name", folder.name)
|
||||
if "data" in form_data:
|
||||
folder.name = form_data.get('name', folder.name)
|
||||
if 'data' in form_data:
|
||||
folder.data = {
|
||||
**(folder.data or {}),
|
||||
**form_data["data"],
|
||||
**form_data['data'],
|
||||
}
|
||||
|
||||
if "meta" in form_data:
|
||||
if 'meta' in form_data:
|
||||
folder.meta = {
|
||||
**(folder.meta or {}),
|
||||
**form_data["meta"],
|
||||
**form_data['meta'],
|
||||
}
|
||||
|
||||
folder.updated_at = int(time.time())
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
return FolderModel.model_validate(folder)
|
||||
except Exception as e:
|
||||
log.error(f"update_folder: {e}")
|
||||
log.error(f'update_folder: {e}')
|
||||
return
|
||||
|
||||
def update_folder_is_expanded_by_id_and_user_id(
|
||||
self, id: str, user_id: str, is_expanded: bool, db: Optional[Session] = None
|
||||
async def update_folder_is_expanded_by_id_and_user_id(
|
||||
self, id: str, user_id: str, is_expanded: bool, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FolderModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
|
||||
folder = result.scalars().first()
|
||||
|
||||
if not folder:
|
||||
return None
|
||||
@@ -285,54 +279,53 @@ class FolderTable:
|
||||
folder.is_expanded = is_expanded
|
||||
folder.updated_at = int(time.time())
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
return FolderModel.model_validate(folder)
|
||||
except Exception as e:
|
||||
log.error(f"update_folder: {e}")
|
||||
log.error(f'update_folder: {e}')
|
||||
return
|
||||
|
||||
def delete_folder_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def delete_folder_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> list[str]:
|
||||
try:
|
||||
folder_ids = []
|
||||
with get_db_context(db) as db:
|
||||
folder = db.query(Folder).filter_by(id=id, user_id=user_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(id=id, user_id=user_id))
|
||||
folder = result.scalars().first()
|
||||
if not folder:
|
||||
return folder_ids
|
||||
|
||||
folder_ids.append(folder.id)
|
||||
|
||||
# Delete all children folders
|
||||
def delete_children(folder):
|
||||
folder_children = self.get_folders_by_parent_id_and_user_id(
|
||||
folder.id, user_id, db=db
|
||||
)
|
||||
async def delete_children(folder):
|
||||
folder_children = await self.get_folders_by_parent_id_and_user_id(folder.id, user_id, db=db)
|
||||
for folder_child in folder_children:
|
||||
|
||||
delete_children(folder_child)
|
||||
await delete_children(folder_child)
|
||||
folder_ids.append(folder_child.id)
|
||||
|
||||
folder = db.query(Folder).filter_by(id=folder_child.id).first()
|
||||
db.delete(folder)
|
||||
db.commit()
|
||||
child_result = await db.execute(select(Folder).filter_by(id=folder_child.id))
|
||||
child_folder = child_result.scalars().first()
|
||||
await db.delete(child_folder)
|
||||
await db.commit()
|
||||
|
||||
delete_children(folder)
|
||||
db.delete(folder)
|
||||
db.commit()
|
||||
await delete_children(folder)
|
||||
await db.delete(folder)
|
||||
await db.commit()
|
||||
return folder_ids
|
||||
except Exception as e:
|
||||
log.error(f"delete_folder: {e}")
|
||||
log.error(f'delete_folder: {e}')
|
||||
return []
|
||||
|
||||
def normalize_folder_name(self, name: str) -> str:
|
||||
# Replace _ and space with a single space, lower case, collapse multiple spaces
|
||||
name = re.sub(r"[\s_]+", " ", name)
|
||||
name = re.sub(r'[\s_]+', ' ', name)
|
||||
return name.strip().lower()
|
||||
|
||||
def search_folders_by_names(
|
||||
self, user_id: str, queries: list[str], db: Optional[Session] = None
|
||||
async def search_folders_by_names(
|
||||
self, user_id: str, queries: list[str], db: Optional[AsyncSession] = None
|
||||
) -> list[FolderModel]:
|
||||
"""
|
||||
Search for folders for a user where the name matches any of the queries, treating _ and space as equivalent, case-insensitive.
|
||||
@@ -342,18 +335,18 @@ class FolderTable:
|
||||
return []
|
||||
|
||||
results = {}
|
||||
with get_db_context(db) as db:
|
||||
folders = db.query(Folder).filter_by(user_id=user_id).all()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(user_id=user_id))
|
||||
folders = result.scalars().all()
|
||||
for folder in folders:
|
||||
if self.normalize_folder_name(folder.name) in normalized_queries:
|
||||
results[folder.id] = FolderModel.model_validate(folder)
|
||||
|
||||
# get children folders
|
||||
children = self.get_children_folders_by_id_and_user_id(
|
||||
folder.id, user_id, db=db
|
||||
)
|
||||
for child in children:
|
||||
results[child.id] = child
|
||||
children = await self.get_children_folders_by_id_and_user_id(folder.id, user_id, db=db)
|
||||
if children:
|
||||
for child in children:
|
||||
results[child.id] = child
|
||||
|
||||
# Return the results as a list
|
||||
if not results:
|
||||
@@ -362,16 +355,17 @@ class FolderTable:
|
||||
results = list(results.values())
|
||||
return results
|
||||
|
||||
def search_folders_by_name_contains(
|
||||
self, user_id: str, query: str, db: Optional[Session] = None
|
||||
async def search_folders_by_name_contains(
|
||||
self, user_id: str, query: str, db: Optional[AsyncSession] = None
|
||||
) -> list[FolderModel]:
|
||||
"""
|
||||
Partial match: normalized name contains (as substring) the normalized query.
|
||||
"""
|
||||
normalized_query = self.normalize_folder_name(query)
|
||||
results = []
|
||||
with get_db_context(db) as db:
|
||||
folders = db.query(Folder).filter_by(user_id=user_id).all()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Folder).filter_by(user_id=user_id))
|
||||
folders = result.scalars().all()
|
||||
for folder in folders:
|
||||
norm_name = self.normalize_folder_name(folder.name)
|
||||
if normalized_query in norm_name:
|
||||
|
||||
@@ -2,9 +2,10 @@ import logging
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from open_webui.models.users import Users, UserModel
|
||||
from sqlalchemy import select, delete, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.users import Users, UserModel, UserResponse
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Boolean, Column, String, Text, Index
|
||||
|
||||
@@ -12,11 +13,13 @@ log = logging.getLogger(__name__)
|
||||
|
||||
####################
|
||||
# Functions DB Schema
|
||||
# Each function here is a promise made. Let no promise
|
||||
# go unkept, and let none be called who cannot answer.
|
||||
####################
|
||||
|
||||
|
||||
class Function(Base):
|
||||
__tablename__ = "function"
|
||||
__tablename__ = 'function'
|
||||
|
||||
id = Column(String, primary_key=True, unique=True)
|
||||
user_id = Column(String)
|
||||
@@ -30,13 +33,13 @@ class Function(Base):
|
||||
updated_at = Column(BigInteger)
|
||||
created_at = Column(BigInteger)
|
||||
|
||||
__table_args__ = (Index("is_global_idx", "is_global"),)
|
||||
__table_args__ = (Index('is_global_idx', 'is_global'),)
|
||||
|
||||
|
||||
class FunctionMeta(BaseModel):
|
||||
description: Optional[str] = None
|
||||
manifest: Optional[dict] = {}
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
|
||||
class FunctionModel(BaseModel):
|
||||
@@ -75,10 +78,6 @@ class FunctionWithValvesModel(BaseModel):
|
||||
####################
|
||||
|
||||
|
||||
class FunctionUserResponse(FunctionModel):
|
||||
user: Optional[UserModel] = None
|
||||
|
||||
|
||||
class FunctionResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
@@ -90,6 +89,12 @@ class FunctionResponse(BaseModel):
|
||||
updated_at: int # timestamp in epoch
|
||||
created_at: int # timestamp in epoch
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class FunctionUserResponse(FunctionResponse):
|
||||
user: Optional[UserResponse] = None
|
||||
|
||||
|
||||
class FunctionForm(BaseModel):
|
||||
id: str
|
||||
@@ -103,48 +108,49 @@ class FunctionValves(BaseModel):
|
||||
|
||||
|
||||
class FunctionsTable:
|
||||
def insert_new_function(
|
||||
async def insert_new_function(
|
||||
self,
|
||||
user_id: str,
|
||||
type: str,
|
||||
form_data: FunctionForm,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[FunctionModel]:
|
||||
function = FunctionModel(
|
||||
**{
|
||||
**form_data.model_dump(),
|
||||
"user_id": user_id,
|
||||
"type": type,
|
||||
"updated_at": int(time.time()),
|
||||
"created_at": int(time.time()),
|
||||
'user_id': user_id,
|
||||
'type': type,
|
||||
'updated_at': int(time.time()),
|
||||
'created_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = Function(**function.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return FunctionModel.model_validate(result)
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Error creating a new function: {e}")
|
||||
log.exception(f'Error creating a new function: {e}')
|
||||
return None
|
||||
|
||||
def sync_functions(
|
||||
async def sync_functions(
|
||||
self,
|
||||
user_id: str,
|
||||
functions: list[FunctionWithValvesModel],
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[FunctionWithValvesModel]:
|
||||
# Synchronize functions for a user by updating existing ones, inserting new ones, and removing those that are no longer present.
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Get existing functions
|
||||
existing_functions = db.query(Function).all()
|
||||
result = await db.execute(select(Function))
|
||||
existing_functions = result.scalars().all()
|
||||
existing_ids = {func.id for func in existing_functions}
|
||||
|
||||
# Prepare a set of new function IDs
|
||||
@@ -153,19 +159,21 @@ class FunctionsTable:
|
||||
# Update or insert functions
|
||||
for func in functions:
|
||||
if func.id in existing_ids:
|
||||
db.query(Function).filter_by(id=func.id).update(
|
||||
{
|
||||
await db.execute(
|
||||
update(Function)
|
||||
.filter_by(id=func.id)
|
||||
.values(
|
||||
**func.model_dump(),
|
||||
"user_id": user_id,
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
user_id=user_id,
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
else:
|
||||
new_func = Function(
|
||||
**{
|
||||
**func.model_dump(),
|
||||
"user_id": user_id,
|
||||
"updated_at": int(time.time()),
|
||||
'user_id': user_id,
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
db.add(new_func)
|
||||
@@ -173,64 +181,78 @@ class FunctionsTable:
|
||||
# Remove functions that are no longer present
|
||||
for func in existing_functions:
|
||||
if func.id not in new_function_ids:
|
||||
db.delete(func)
|
||||
await db.delete(func)
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
return [
|
||||
FunctionModel.model_validate(func)
|
||||
for func in db.query(Function).all()
|
||||
]
|
||||
result = await db.execute(select(Function))
|
||||
return [FunctionModel.model_validate(func) for func in result.scalars().all()]
|
||||
except Exception as e:
|
||||
log.exception(f"Error syncing functions for user {user_id}: {e}")
|
||||
log.exception(f'Error syncing functions for user {user_id}: {e}')
|
||||
return []
|
||||
|
||||
def get_function_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[FunctionModel]:
|
||||
async def get_function_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[FunctionModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
function = db.get(Function, id)
|
||||
return FunctionModel.model_validate(function)
|
||||
async with get_async_db_context(db) as db:
|
||||
function = await db.get(Function, id)
|
||||
return FunctionModel.model_validate(function) if function else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_functions(
|
||||
self, active_only=False, include_valves=False, db: Optional[Session] = None
|
||||
) -> list[FunctionModel | FunctionWithValvesModel]:
|
||||
with get_db_context(db) as db:
|
||||
if active_only:
|
||||
functions = db.query(Function).filter_by(is_active=True).all()
|
||||
async def get_functions_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[FunctionModel]:
|
||||
"""
|
||||
Batch fetch multiple functions by their IDs in a single query.
|
||||
Returns functions in the same order as the input IDs (None entries filtered out).
|
||||
"""
|
||||
if not ids:
|
||||
return []
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Function).filter(Function.id.in_(ids)))
|
||||
functions = result.scalars().all()
|
||||
# Create a dict for O(1) lookup
|
||||
func_dict = {f.id: FunctionModel.model_validate(f) for f in functions}
|
||||
# Return in original order, filtering out any not found
|
||||
return [func_dict[id] for id in ids if id in func_dict]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
async def get_functions(
|
||||
self, active_only=False, include_valves=False, db: Optional[AsyncSession] = None
|
||||
) -> list[FunctionModel | FunctionWithValvesModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
if active_only:
|
||||
result = await db.execute(select(Function).filter_by(is_active=True))
|
||||
else:
|
||||
functions = db.query(Function).all()
|
||||
result = await db.execute(select(Function))
|
||||
|
||||
functions = result.scalars().all()
|
||||
|
||||
if include_valves:
|
||||
return [
|
||||
FunctionWithValvesModel.model_validate(function)
|
||||
for function in functions
|
||||
]
|
||||
return [FunctionWithValvesModel.model_validate(function) for function in functions]
|
||||
else:
|
||||
return [
|
||||
FunctionModel.model_validate(function) for function in functions
|
||||
]
|
||||
return [FunctionModel.model_validate(function) for function in functions]
|
||||
|
||||
def get_function_list(
|
||||
self, db: Optional[Session] = None
|
||||
) -> list[FunctionUserResponse]:
|
||||
with get_db_context(db) as db:
|
||||
functions = db.query(Function).order_by(Function.updated_at.desc()).all()
|
||||
async def get_function_list(self, db: Optional[AsyncSession] = None) -> list[FunctionUserResponse]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Function).order_by(Function.updated_at.desc()))
|
||||
functions = result.scalars().all()
|
||||
user_ids = list(set(func.user_id for func in functions))
|
||||
|
||||
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users_dict = {user.id: user for user in users}
|
||||
|
||||
return [
|
||||
FunctionUserResponse.model_validate(
|
||||
{
|
||||
**FunctionModel.model_validate(func).model_dump(),
|
||||
"user": (
|
||||
users_dict.get(func.user_id).model_dump()
|
||||
**FunctionResponse.model_validate(func).model_dump(),
|
||||
'user': (
|
||||
UserResponse(
|
||||
id=users_dict[func.user_id].id,
|
||||
name=users_dict[func.user_id].name,
|
||||
role=users_dict[func.user_id].role,
|
||||
email=users_dict[func.user_id].email,
|
||||
).model_dump()
|
||||
if func.user_id in users_dict
|
||||
else None
|
||||
),
|
||||
@@ -239,76 +261,72 @@ class FunctionsTable:
|
||||
for func in functions
|
||||
]
|
||||
|
||||
def get_functions_by_type(
|
||||
self, type: str, active_only=False, db: Optional[Session] = None
|
||||
async def get_functions_by_type(
|
||||
self, type: str, active_only=False, db: Optional[AsyncSession] = None
|
||||
) -> list[FunctionModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
if active_only:
|
||||
return [
|
||||
FunctionModel.model_validate(function)
|
||||
for function in db.query(Function)
|
||||
.filter_by(type=type, is_active=True)
|
||||
.all()
|
||||
]
|
||||
result = await db.execute(select(Function).filter_by(type=type, is_active=True))
|
||||
else:
|
||||
return [
|
||||
FunctionModel.model_validate(function)
|
||||
for function in db.query(Function).filter_by(type=type).all()
|
||||
]
|
||||
result = await db.execute(select(Function).filter_by(type=type))
|
||||
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
|
||||
|
||||
def get_global_filter_functions(
|
||||
self, db: Optional[Session] = None
|
||||
) -> list[FunctionModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FunctionModel.model_validate(function)
|
||||
for function in db.query(Function)
|
||||
.filter_by(type="filter", is_active=True, is_global=True)
|
||||
.all()
|
||||
]
|
||||
async def get_global_filter_functions(self, db: Optional[AsyncSession] = None) -> list[FunctionModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Function).filter_by(type='filter', is_active=True, is_global=True))
|
||||
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
|
||||
|
||||
def get_global_action_functions(
|
||||
self, db: Optional[Session] = None
|
||||
) -> list[FunctionModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
FunctionModel.model_validate(function)
|
||||
for function in db.query(Function)
|
||||
.filter_by(type="action", is_active=True, is_global=True)
|
||||
.all()
|
||||
]
|
||||
async def get_global_action_functions(self, db: Optional[AsyncSession] = None) -> list[FunctionModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Function).filter_by(type='action', is_active=True, is_global=True))
|
||||
return [FunctionModel.model_validate(function) for function in result.scalars().all()]
|
||||
|
||||
def get_function_valves_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[dict]:
|
||||
with get_db_context(db) as db:
|
||||
async def get_function_valves_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[dict]:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
function = db.get(Function, id)
|
||||
function = await db.get(Function, id)
|
||||
return function.valves if function.valves else {}
|
||||
except Exception as e:
|
||||
log.exception(f"Error getting function valves by id {id}: {e}")
|
||||
log.exception(f'Error getting function valves by id {id}: {e}')
|
||||
return None
|
||||
|
||||
def update_function_valves_by_id(
|
||||
self, id: str, valves: dict, db: Optional[Session] = None
|
||||
async def get_function_valves_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, dict]:
|
||||
"""
|
||||
Batch fetch valves for multiple functions in a single query.
|
||||
Returns a dict mapping function_id -> valves dict.
|
||||
Functions without valves are mapped to {}.
|
||||
"""
|
||||
if not ids:
|
||||
return {}
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Function.id, Function.valves).filter(Function.id.in_(ids)))
|
||||
functions = result.all()
|
||||
return {f.id: (f.valves if f.valves else {}) for f in functions}
|
||||
except Exception as e:
|
||||
log.exception(f'Error batch-fetching function valves: {e}')
|
||||
return {}
|
||||
|
||||
async def update_function_valves_by_id(
|
||||
self, id: str, valves: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FunctionValves]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
function = db.get(Function, id)
|
||||
function = await db.get(Function, id)
|
||||
function.valves = valves
|
||||
function.updated_at = int(time.time())
|
||||
db.commit()
|
||||
db.refresh(function)
|
||||
return self.get_function_by_id(id, db=db)
|
||||
await db.commit()
|
||||
await db.refresh(function)
|
||||
return FunctionModel.model_validate(function)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def update_function_metadata_by_id(
|
||||
self, id: str, metadata: dict, db: Optional[Session] = None
|
||||
async def update_function_metadata_by_id(
|
||||
self, id: str, metadata: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FunctionModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
function = db.get(Function, id)
|
||||
function = await db.get(Function, id)
|
||||
|
||||
if function:
|
||||
if function.meta:
|
||||
@@ -317,93 +335,94 @@ class FunctionsTable:
|
||||
function.meta = metadata
|
||||
|
||||
function.updated_at = int(time.time())
|
||||
db.commit()
|
||||
db.refresh(function)
|
||||
return self.get_function_by_id(id, db=db)
|
||||
await db.commit()
|
||||
await db.refresh(function)
|
||||
return FunctionModel.model_validate(function)
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Error updating function metadata by id {id}: {e}")
|
||||
log.exception(f'Error updating function metadata by id {id}: {e}')
|
||||
return None
|
||||
|
||||
def get_user_valves_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_user_valves_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[dict]:
|
||||
try:
|
||||
user = Users.get_user_by_id(user_id, db=db)
|
||||
user = await Users.get_user_by_id(user_id, db=db)
|
||||
user_settings = user.settings.model_dump() if user.settings else {}
|
||||
|
||||
# Check if user has "functions" and "valves" settings
|
||||
if "functions" not in user_settings:
|
||||
user_settings["functions"] = {}
|
||||
if "valves" not in user_settings["functions"]:
|
||||
user_settings["functions"]["valves"] = {}
|
||||
if 'functions' not in user_settings:
|
||||
user_settings['functions'] = {}
|
||||
if 'valves' not in user_settings['functions']:
|
||||
user_settings['functions']['valves'] = {}
|
||||
|
||||
return user_settings["functions"]["valves"].get(id, {})
|
||||
return user_settings['functions']['valves'].get(id, {})
|
||||
except Exception as e:
|
||||
log.exception(f"Error getting user values by id {id} and user id {user_id}")
|
||||
log.exception(f'Error getting user values by id {id} and user id {user_id}')
|
||||
return None
|
||||
|
||||
def update_user_valves_by_id_and_user_id(
|
||||
self, id: str, user_id: str, valves: dict, db: Optional[Session] = None
|
||||
async def update_user_valves_by_id_and_user_id(
|
||||
self, id: str, user_id: str, valves: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[dict]:
|
||||
try:
|
||||
user = Users.get_user_by_id(user_id, db=db)
|
||||
user = await Users.get_user_by_id(user_id, db=db)
|
||||
user_settings = user.settings.model_dump() if user.settings else {}
|
||||
|
||||
# Check if user has "functions" and "valves" settings
|
||||
if "functions" not in user_settings:
|
||||
user_settings["functions"] = {}
|
||||
if "valves" not in user_settings["functions"]:
|
||||
user_settings["functions"]["valves"] = {}
|
||||
if 'functions' not in user_settings:
|
||||
user_settings['functions'] = {}
|
||||
if 'valves' not in user_settings['functions']:
|
||||
user_settings['functions']['valves'] = {}
|
||||
|
||||
user_settings["functions"]["valves"][id] = valves
|
||||
user_settings['functions']['valves'][id] = valves
|
||||
|
||||
# Update the user settings in the database
|
||||
Users.update_user_by_id(user_id, {"settings": user_settings}, db=db)
|
||||
await Users.update_user_by_id(user_id, {'settings': user_settings}, db=db)
|
||||
|
||||
return user_settings["functions"]["valves"][id]
|
||||
return user_settings['functions']['valves'][id]
|
||||
except Exception as e:
|
||||
log.exception(
|
||||
f"Error updating user valves by id {id} and user_id {user_id}: {e}"
|
||||
)
|
||||
log.exception(f'Error updating user valves by id {id} and user_id {user_id}: {e}')
|
||||
return None
|
||||
|
||||
def update_function_by_id(
|
||||
self, id: str, updated: dict, db: Optional[Session] = None
|
||||
async def update_function_by_id(
|
||||
self, id: str, updated: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[FunctionModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Function).filter_by(id=id).update(
|
||||
{
|
||||
await db.execute(
|
||||
update(Function)
|
||||
.filter_by(id=id)
|
||||
.values(
|
||||
**updated,
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
return self.get_function_by_id(id, db=db)
|
||||
await db.commit()
|
||||
function = await db.get(Function, id)
|
||||
return FunctionModel.model_validate(function) if function else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def deactivate_all_functions(self, db: Optional[Session] = None) -> Optional[bool]:
|
||||
with get_db_context(db) as db:
|
||||
async def deactivate_all_functions(self, db: Optional[AsyncSession] = None) -> Optional[bool]:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Function).update(
|
||||
{
|
||||
"is_active": False,
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
await db.execute(
|
||||
update(Function).values(
|
||||
is_active=False,
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def delete_function_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_function_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Function).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(Function).filter_by(id=id))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
|
||||
+255
-222
@@ -4,8 +4,10 @@ import time
|
||||
from typing import Optional
|
||||
import uuid
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, update, func, and_, or_, cast, String
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.env import DEFAULT_GROUP_SHARE_PERMISSION
|
||||
|
||||
from open_webui.models.files import FileMetadataResponse
|
||||
|
||||
@@ -14,26 +16,22 @@ from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Column,
|
||||
String,
|
||||
Text,
|
||||
JSON,
|
||||
and_,
|
||||
func,
|
||||
ForeignKey,
|
||||
cast,
|
||||
or_,
|
||||
)
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
####################
|
||||
# UserGroup DB Schema
|
||||
# Let none who belong to this house be turned away,
|
||||
# and let the covenant hold for every member.
|
||||
####################
|
||||
|
||||
|
||||
class Group(Base):
|
||||
__tablename__ = "group"
|
||||
__tablename__ = 'group'
|
||||
|
||||
id = Column(Text, unique=True, primary_key=True)
|
||||
user_id = Column(Text)
|
||||
@@ -69,12 +67,12 @@ class GroupModel(BaseModel):
|
||||
|
||||
|
||||
class GroupMember(Base):
|
||||
__tablename__ = "group_member"
|
||||
__tablename__ = 'group_member'
|
||||
|
||||
id = Column(Text, unique=True, primary_key=True)
|
||||
group_id = Column(
|
||||
Text,
|
||||
ForeignKey("group.id", ondelete="CASCADE"),
|
||||
ForeignKey('group.id', ondelete='CASCADE'),
|
||||
nullable=False,
|
||||
)
|
||||
user_id = Column(Text, nullable=False)
|
||||
@@ -99,6 +97,16 @@ class GroupResponse(GroupModel):
|
||||
member_count: Optional[int] = None
|
||||
|
||||
|
||||
class GroupInfoResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
description: str
|
||||
member_count: Optional[int] = None
|
||||
created_at: int
|
||||
updated_at: int
|
||||
|
||||
|
||||
class GroupForm(BaseModel):
|
||||
name: str
|
||||
description: str
|
||||
@@ -120,25 +128,36 @@ class GroupListResponse(BaseModel):
|
||||
|
||||
|
||||
class GroupTable:
|
||||
def insert_new_group(
|
||||
self, user_id: str, form_data: GroupForm, db: Optional[Session] = None
|
||||
def _ensure_default_share_config(self, group_data: dict) -> dict:
|
||||
"""Ensure the group data dict has a default share config if not already set."""
|
||||
if 'data' not in group_data or group_data['data'] is None:
|
||||
group_data['data'] = {}
|
||||
if 'config' not in group_data['data']:
|
||||
group_data['data']['config'] = {}
|
||||
if 'share' not in group_data['data']['config']:
|
||||
group_data['data']['config']['share'] = DEFAULT_GROUP_SHARE_PERMISSION
|
||||
return group_data
|
||||
|
||||
async def insert_new_group(
|
||||
self, user_id: str, form_data: GroupForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[GroupModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
group_data = self._ensure_default_share_config(form_data.model_dump(exclude_none=True))
|
||||
group = GroupModel(
|
||||
**{
|
||||
**form_data.model_dump(exclude_none=True),
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_id": user_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
**group_data,
|
||||
'id': str(uuid.uuid4()),
|
||||
'user_id': user_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
result = Group(**group.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return GroupModel.model_validate(result)
|
||||
else:
|
||||
@@ -147,199 +166,208 @@ class GroupTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_all_groups(self, db: Optional[Session] = None) -> list[GroupModel]:
|
||||
with get_db_context(db) as db:
|
||||
groups = db.query(Group).order_by(Group.updated_at.desc()).all()
|
||||
async def get_all_groups(self, db: Optional[AsyncSession] = None) -> list[GroupModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Group).order_by(Group.updated_at.desc()))
|
||||
groups = result.scalars().all()
|
||||
return [GroupModel.model_validate(group) for group in groups]
|
||||
|
||||
def get_groups(self, filter, db: Optional[Session] = None) -> list[GroupResponse]:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Group)
|
||||
async def get_group_by_name(self, name: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Group).filter(Group.name == name))
|
||||
group = result.scalars().first()
|
||||
return GroupModel.model_validate(group) if group else None
|
||||
|
||||
async def get_groups(self, filter, db: Optional[AsyncSession] = None) -> list[GroupResponse]:
|
||||
async with get_async_db_context(db) as db:
|
||||
member_count = (
|
||||
select(func.count(GroupMember.user_id))
|
||||
.where(GroupMember.group_id == Group.id)
|
||||
.correlate(Group)
|
||||
.scalar_subquery()
|
||||
.label('member_count')
|
||||
)
|
||||
stmt = select(Group, member_count)
|
||||
|
||||
if filter:
|
||||
if "query" in filter:
|
||||
query = query.filter(Group.name.ilike(f"%{filter['query']}%"))
|
||||
if 'query' in filter:
|
||||
stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%'))
|
||||
|
||||
# When share filter is present, member check is handled in the share logic
|
||||
if "share" in filter:
|
||||
share_value = filter["share"]
|
||||
member_id = filter.get("member_id")
|
||||
json_share = Group.data["config"]["share"]
|
||||
json_share_bool = json_share.as_boolean()
|
||||
if 'share' in filter:
|
||||
share_value = filter['share']
|
||||
member_id = filter.get('member_id')
|
||||
json_share = Group.data['config']['share']
|
||||
json_share_str = json_share.as_string()
|
||||
json_share_lower = func.lower(json_share_str)
|
||||
|
||||
if share_value:
|
||||
# Groups open to anyone: data is null, share is null, or share is true
|
||||
anyone_can_share = or_(
|
||||
Group.data.is_(None),
|
||||
json_share_bool.is_(None),
|
||||
json_share_bool == True,
|
||||
json_share_str.is_(None),
|
||||
json_share_lower == 'true',
|
||||
json_share_lower == '1', # Handle SQLite boolean true
|
||||
)
|
||||
|
||||
if member_id:
|
||||
# Also include member-only groups where user is a member
|
||||
member_groups_subq = (
|
||||
db.query(GroupMember.group_id)
|
||||
.filter(GroupMember.user_id == member_id)
|
||||
.subquery()
|
||||
)
|
||||
member_groups_select = select(GroupMember.group_id).where(GroupMember.user_id == member_id)
|
||||
members_only_and_is_member = and_(
|
||||
json_share_str == "members",
|
||||
Group.id.in_(member_groups_subq),
|
||||
)
|
||||
query = query.filter(
|
||||
or_(anyone_can_share, members_only_and_is_member)
|
||||
json_share_lower == 'members',
|
||||
Group.id.in_(member_groups_select),
|
||||
)
|
||||
stmt = stmt.filter(or_(anyone_can_share, members_only_and_is_member))
|
||||
else:
|
||||
query = query.filter(anyone_can_share)
|
||||
stmt = stmt.filter(anyone_can_share)
|
||||
else:
|
||||
query = query.filter(
|
||||
and_(Group.data.isnot(None), json_share_bool == False)
|
||||
)
|
||||
stmt = stmt.filter(and_(Group.data.isnot(None), json_share_lower == 'false'))
|
||||
|
||||
else:
|
||||
# Only apply member_id filter when share filter is NOT present
|
||||
if "member_id" in filter:
|
||||
query = query.join(
|
||||
GroupMember, GroupMember.group_id == Group.id
|
||||
).filter(GroupMember.user_id == filter["member_id"])
|
||||
if 'member_id' in filter:
|
||||
stmt = stmt.filter(
|
||||
Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id']))
|
||||
)
|
||||
|
||||
result = await db.execute(stmt.order_by(Group.updated_at.desc()))
|
||||
rows = result.all()
|
||||
|
||||
groups = query.order_by(Group.updated_at.desc()).all()
|
||||
return [
|
||||
GroupResponse.model_validate(
|
||||
{
|
||||
**GroupModel.model_validate(group).model_dump(),
|
||||
"member_count": self.get_group_member_count_by_id(
|
||||
group.id, db=db
|
||||
),
|
||||
'member_count': count or 0,
|
||||
}
|
||||
)
|
||||
for group in groups
|
||||
for group, count in rows
|
||||
]
|
||||
|
||||
def search_groups(
|
||||
async def search_groups(
|
||||
self,
|
||||
filter: Optional[dict] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> GroupListResponse:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Group)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Group)
|
||||
|
||||
if filter:
|
||||
if "query" in filter:
|
||||
query = query.filter(Group.name.ilike(f"%{filter['query']}%"))
|
||||
if "member_id" in filter:
|
||||
query = query.join(
|
||||
GroupMember, GroupMember.group_id == Group.id
|
||||
).filter(GroupMember.user_id == filter["member_id"])
|
||||
|
||||
if "share" in filter:
|
||||
# 'share' is stored in data JSON, support both sqlite and postgres
|
||||
share_value = filter["share"]
|
||||
print("Filtering by share:", share_value)
|
||||
query = query.filter(
|
||||
Group.data.op("->>")("share") == str(share_value)
|
||||
if 'query' in filter:
|
||||
stmt = stmt.filter(Group.name.ilike(f'%{filter["query"]}%'))
|
||||
if 'member_id' in filter:
|
||||
stmt = stmt.filter(
|
||||
Group.id.in_(select(GroupMember.group_id).where(GroupMember.user_id == filter['member_id']))
|
||||
)
|
||||
|
||||
total = query.count()
|
||||
query = query.order_by(Group.updated_at.desc())
|
||||
groups = query.offset(skip).limit(limit).all()
|
||||
if 'share' in filter:
|
||||
share_value = filter['share']
|
||||
stmt = stmt.filter(Group.data.op('->>')('share') == str(share_value))
|
||||
|
||||
# Get total count
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
member_count = (
|
||||
select(func.count(GroupMember.user_id))
|
||||
.where(GroupMember.group_id == Group.id)
|
||||
.correlate(Group)
|
||||
.scalar_subquery()
|
||||
.label('member_count')
|
||||
)
|
||||
result = await db.execute(
|
||||
select(Group, member_count)
|
||||
.where(Group.id.in_(select(stmt.subquery().c.id)))
|
||||
.order_by(Group.updated_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
)
|
||||
rows = result.all()
|
||||
|
||||
return {
|
||||
"items": [
|
||||
'items': [
|
||||
GroupResponse.model_validate(
|
||||
**GroupModel.model_validate(group).model_dump(),
|
||||
member_count=self.get_group_member_count_by_id(group.id, db=db),
|
||||
{
|
||||
**GroupModel.model_validate(group).model_dump(),
|
||||
'member_count': count or 0,
|
||||
}
|
||||
)
|
||||
for group in groups
|
||||
for group, count in rows
|
||||
],
|
||||
"total": total,
|
||||
'total': total,
|
||||
}
|
||||
|
||||
def get_groups_by_member_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> list[GroupModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
GroupModel.model_validate(group)
|
||||
for group in db.query(Group)
|
||||
async def get_groups_by_member_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[GroupModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Group)
|
||||
.join(GroupMember, GroupMember.group_id == Group.id)
|
||||
.filter(GroupMember.user_id == user_id)
|
||||
.order_by(Group.updated_at.desc())
|
||||
.all()
|
||||
]
|
||||
)
|
||||
return [GroupModel.model_validate(group) for group in result.scalars().all()]
|
||||
|
||||
def get_groups_by_member_ids(
|
||||
self, user_ids: list[str], db: Optional[Session] = None
|
||||
async def get_groups_by_member_ids(
|
||||
self, user_ids: list[str], db: Optional[AsyncSession] = None
|
||||
) -> dict[str, list[GroupModel]]:
|
||||
"""Fetch groups for multiple users in a single query to avoid N+1."""
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Query GroupMember joined with Group, filtering by user_ids
|
||||
results = (
|
||||
db.query(GroupMember.user_id, Group)
|
||||
result = await db.execute(
|
||||
select(GroupMember.user_id, Group)
|
||||
.join(Group, Group.id == GroupMember.group_id)
|
||||
.filter(GroupMember.user_id.in_(user_ids))
|
||||
.order_by(Group.updated_at.desc())
|
||||
.all()
|
||||
)
|
||||
rows = result.all()
|
||||
|
||||
# Group groups by user_id
|
||||
user_groups: dict[str, list[GroupModel]] = {uid: [] for uid in user_ids}
|
||||
for user_id, group in results:
|
||||
for user_id, group in rows:
|
||||
user_groups[user_id].append(GroupModel.model_validate(group))
|
||||
|
||||
return user_groups
|
||||
|
||||
def get_group_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[GroupModel]:
|
||||
async def get_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[GroupModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
group = db.query(Group).filter_by(id=id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Group).filter_by(id=id))
|
||||
group = result.scalars().first()
|
||||
return GroupModel.model_validate(group) if group else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_group_user_ids_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[list[str]]:
|
||||
with get_db_context(db) as db:
|
||||
members = (
|
||||
db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all()
|
||||
)
|
||||
async def get_group_user_ids_by_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(GroupMember.user_id).filter(GroupMember.group_id == id))
|
||||
members = result.all()
|
||||
|
||||
if not members:
|
||||
return None
|
||||
return []
|
||||
|
||||
return [m[0] for m in members]
|
||||
|
||||
def get_group_user_ids_by_ids(
|
||||
self, group_ids: list[str], db: Optional[Session] = None
|
||||
async def get_group_user_ids_by_ids(
|
||||
self, group_ids: list[str], db: Optional[AsyncSession] = None
|
||||
) -> dict[str, list[str]]:
|
||||
with get_db_context(db) as db:
|
||||
members = (
|
||||
db.query(GroupMember.group_id, GroupMember.user_id)
|
||||
.filter(GroupMember.group_id.in_(group_ids))
|
||||
.all()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(GroupMember.group_id, GroupMember.user_id).filter(GroupMember.group_id.in_(group_ids))
|
||||
)
|
||||
members = result.all()
|
||||
|
||||
group_user_ids: dict[str, list[str]] = {
|
||||
group_id: [] for group_id in group_ids
|
||||
}
|
||||
group_user_ids: dict[str, list[str]] = {group_id: [] for group_id in group_ids}
|
||||
|
||||
for group_id, user_id in members:
|
||||
group_user_ids[group_id].append(user_id)
|
||||
|
||||
return group_user_ids
|
||||
|
||||
def set_group_user_ids_by_id(
|
||||
self, group_id: str, user_ids: list[str], db: Optional[Session] = None
|
||||
async def set_group_user_ids_by_id(
|
||||
self, group_id: str, user_ids: list[str], db: Optional[AsyncSession] = None
|
||||
) -> None:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Delete existing members
|
||||
db.query(GroupMember).filter(GroupMember.group_id == group_id).delete()
|
||||
await db.execute(delete(GroupMember).filter(GroupMember.group_id == group_id))
|
||||
|
||||
# Insert new members
|
||||
now = int(time.time())
|
||||
@@ -355,142 +383,149 @@ class GroupTable:
|
||||
]
|
||||
|
||||
db.add_all(new_members)
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
def get_group_member_count_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> int:
|
||||
with get_db_context(db) as db:
|
||||
count = (
|
||||
db.query(func.count(GroupMember.user_id))
|
||||
.filter(GroupMember.group_id == id)
|
||||
.scalar()
|
||||
)
|
||||
async def get_group_member_count_by_id(self, id: str, db: Optional[AsyncSession] = None) -> int:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(func.count(GroupMember.user_id)).filter(GroupMember.group_id == id))
|
||||
count = result.scalar()
|
||||
return count if count else 0
|
||||
|
||||
def update_group_by_id(
|
||||
async def get_group_member_counts_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> dict[str, int]:
|
||||
if not ids:
|
||||
return {}
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(GroupMember.group_id, func.count(GroupMember.user_id))
|
||||
.filter(GroupMember.group_id.in_(ids))
|
||||
.group_by(GroupMember.group_id)
|
||||
)
|
||||
rows = result.all()
|
||||
return {group_id: count for group_id, count in rows}
|
||||
|
||||
async def update_group_by_id(
|
||||
self,
|
||||
id: str,
|
||||
form_data: GroupUpdateForm,
|
||||
overwrite: bool = False,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[GroupModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Group).filter_by(id=id).update(
|
||||
{
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(
|
||||
update(Group)
|
||||
.filter_by(id=id)
|
||||
.values(
|
||||
**form_data.model_dump(exclude_none=True),
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
return self.get_group_by_id(id=id, db=db)
|
||||
await db.commit()
|
||||
return await self.get_group_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
def delete_group_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
async def delete_group_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Group).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(Group).filter_by(id=id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_all_groups(self, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_all_groups(self, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Group).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(Group))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def remove_user_from_all_groups(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def remove_user_from_all_groups(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
# Find all groups the user belongs to
|
||||
groups = (
|
||||
db.query(Group)
|
||||
result = await db.execute(
|
||||
select(Group)
|
||||
.join(GroupMember, GroupMember.group_id == Group.id)
|
||||
.filter(GroupMember.user_id == user_id)
|
||||
.all()
|
||||
)
|
||||
groups = result.scalars().all()
|
||||
|
||||
# Remove the user from each group
|
||||
for group in groups:
|
||||
db.query(GroupMember).filter(
|
||||
GroupMember.group_id == group.id, GroupMember.user_id == user_id
|
||||
).delete()
|
||||
|
||||
db.query(Group).filter_by(id=group.id).update(
|
||||
{"updated_at": int(time.time())}
|
||||
await db.execute(
|
||||
delete(GroupMember).filter(GroupMember.group_id == group.id, GroupMember.user_id == user_id)
|
||||
)
|
||||
|
||||
db.commit()
|
||||
await db.execute(update(Group).filter_by(id=group.id).values(updated_at=int(time.time())))
|
||||
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
except Exception:
|
||||
db.rollback()
|
||||
await db.rollback()
|
||||
return False
|
||||
|
||||
def create_groups_by_group_names(
|
||||
self, user_id: str, group_names: list[str], db: Optional[Session] = None
|
||||
async def create_groups_by_group_names(
|
||||
self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None
|
||||
) -> list[GroupModel]:
|
||||
|
||||
# check for existing groups
|
||||
existing_groups = self.get_all_groups(db=db)
|
||||
existing_groups = await self.get_all_groups(db=db)
|
||||
existing_group_names = {group.name for group in existing_groups}
|
||||
|
||||
new_groups = []
|
||||
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
for group_name in group_names:
|
||||
if group_name not in existing_group_names:
|
||||
new_group = GroupModel(
|
||||
id=str(uuid.uuid4()),
|
||||
user_id=user_id,
|
||||
name=group_name,
|
||||
description="",
|
||||
description='',
|
||||
data={
|
||||
'config': {
|
||||
'share': DEFAULT_GROUP_SHARE_PERMISSION,
|
||||
}
|
||||
},
|
||||
created_at=int(time.time()),
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
try:
|
||||
result = Group(**new_group.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
new_groups.append(GroupModel.model_validate(result))
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
continue
|
||||
return new_groups
|
||||
|
||||
def sync_groups_by_group_names(
|
||||
self, user_id: str, group_names: list[str], db: Optional[Session] = None
|
||||
async def sync_groups_by_group_names(
|
||||
self, user_id: str, group_names: list[str], db: Optional[AsyncSession] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
now = int(time.time())
|
||||
|
||||
# 1. Groups that SHOULD contain the user
|
||||
target_groups = (
|
||||
db.query(Group).filter(Group.name.in_(group_names)).all()
|
||||
)
|
||||
result = await db.execute(select(Group).filter(Group.name.in_(group_names)))
|
||||
target_groups = result.scalars().all()
|
||||
target_group_ids = {g.id for g in target_groups}
|
||||
|
||||
# 2. Groups the user is CURRENTLY in
|
||||
existing_group_ids = {
|
||||
g.id
|
||||
for g in db.query(Group)
|
||||
result = await db.execute(
|
||||
select(Group)
|
||||
.join(GroupMember, GroupMember.group_id == Group.id)
|
||||
.filter(GroupMember.user_id == user_id)
|
||||
.all()
|
||||
}
|
||||
)
|
||||
existing_group_ids = {g.id for g in result.scalars().all()}
|
||||
|
||||
# 3. Determine adds + removals
|
||||
groups_to_add = target_group_ids - existing_group_ids
|
||||
@@ -498,15 +533,15 @@ class GroupTable:
|
||||
|
||||
# 4. Remove in one bulk delete
|
||||
if groups_to_remove:
|
||||
db.query(GroupMember).filter(
|
||||
GroupMember.user_id == user_id,
|
||||
GroupMember.group_id.in_(groups_to_remove),
|
||||
).delete(synchronize_session=False)
|
||||
|
||||
db.query(Group).filter(Group.id.in_(groups_to_remove)).update(
|
||||
{"updated_at": now}, synchronize_session=False
|
||||
await db.execute(
|
||||
delete(GroupMember).filter(
|
||||
GroupMember.user_id == user_id,
|
||||
GroupMember.group_id.in_(groups_to_remove),
|
||||
)
|
||||
)
|
||||
|
||||
await db.execute(update(Group).filter(Group.id.in_(groups_to_remove)).values(updated_at=now))
|
||||
|
||||
# 5. Bulk insert missing memberships
|
||||
for group_id in groups_to_add:
|
||||
db.add(
|
||||
@@ -520,27 +555,26 @@ class GroupTable:
|
||||
)
|
||||
|
||||
if groups_to_add:
|
||||
db.query(Group).filter(Group.id.in_(groups_to_add)).update(
|
||||
{"updated_at": now}, synchronize_session=False
|
||||
)
|
||||
await db.execute(update(Group).filter(Group.id.in_(groups_to_add)).values(updated_at=now))
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
db.rollback()
|
||||
await db.rollback()
|
||||
return False
|
||||
|
||||
def add_users_to_group(
|
||||
async def add_users_to_group(
|
||||
self,
|
||||
id: str,
|
||||
user_ids: Optional[list[str]] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[GroupModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
group = db.query(Group).filter_by(id=id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Group).filter_by(id=id))
|
||||
group = result.scalars().first()
|
||||
if not group:
|
||||
return None
|
||||
|
||||
@@ -557,15 +591,14 @@ class GroupTable:
|
||||
updated_at=now,
|
||||
)
|
||||
)
|
||||
db.flush() # Detect unique constraint violation early
|
||||
await db.flush() # Detect unique constraint violation early
|
||||
except Exception:
|
||||
db.rollback() # Clear failed INSERT
|
||||
db.begin() # Start a new transaction
|
||||
await db.rollback() # Clear failed INSERT
|
||||
continue # Duplicate → ignore
|
||||
|
||||
group.updated_at = now
|
||||
db.commit()
|
||||
db.refresh(group)
|
||||
await db.commit()
|
||||
await db.refresh(group)
|
||||
|
||||
return GroupModel.model_validate(group)
|
||||
|
||||
@@ -573,32 +606,32 @@ class GroupTable:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
def remove_users_from_group(
|
||||
async def remove_users_from_group(
|
||||
self,
|
||||
id: str,
|
||||
user_ids: Optional[list[str]] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[GroupModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
group = db.query(Group).filter_by(id=id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Group).filter_by(id=id))
|
||||
group = result.scalars().first()
|
||||
if not group:
|
||||
return None
|
||||
|
||||
if not user_ids:
|
||||
return GroupModel.model_validate(group)
|
||||
|
||||
# Remove each user from group_member
|
||||
for user_id in user_ids:
|
||||
db.query(GroupMember).filter(
|
||||
GroupMember.group_id == id, GroupMember.user_id == user_id
|
||||
).delete()
|
||||
# Remove users from group_member in batch
|
||||
await db.execute(
|
||||
delete(GroupMember).filter(GroupMember.group_id == id, GroupMember.user_id.in_(user_ids))
|
||||
)
|
||||
|
||||
# Update group timestamp
|
||||
group.updated_at = int(time.time())
|
||||
|
||||
db.commit()
|
||||
db.refresh(group)
|
||||
await db.commit()
|
||||
await db.refresh(group)
|
||||
return GroupModel.model_validate(group)
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -4,8 +4,9 @@ import time
|
||||
from typing import Optional
|
||||
import uuid
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, update, or_, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
|
||||
from open_webui.models.files import (
|
||||
File,
|
||||
@@ -15,9 +16,10 @@ from open_webui.models.files import (
|
||||
)
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import User, UserModel, Users, UserResponse
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import (
|
||||
BigInteger,
|
||||
Column,
|
||||
@@ -26,22 +28,19 @@ from sqlalchemy import (
|
||||
Text,
|
||||
JSON,
|
||||
UniqueConstraint,
|
||||
or_,
|
||||
)
|
||||
|
||||
from open_webui.utils.access_control import has_access
|
||||
from open_webui.utils.db.access_control import has_permission
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
####################
|
||||
# Knowledge DB Schema
|
||||
# Let what was gathered here outlast the one who gathered it,
|
||||
# and still teach when the builder is gone.
|
||||
####################
|
||||
|
||||
|
||||
class Knowledge(Base):
|
||||
__tablename__ = "knowledge"
|
||||
__tablename__ = 'knowledge'
|
||||
|
||||
id = Column(Text, unique=True, primary_key=True)
|
||||
user_id = Column(Text)
|
||||
@@ -50,22 +49,6 @@ class Knowledge(Base):
|
||||
description = Column(Text)
|
||||
|
||||
meta = Column(JSON, nullable=True)
|
||||
access_control = Column(JSON, nullable=True) # Controls data access levels.
|
||||
# Defines access control rules for this entry.
|
||||
# - `None`: Public access, available to all users with the "user" role.
|
||||
# - `{}`: Private access, restricted exclusively to the owner.
|
||||
# - Custom permissions: Specific access control for reading and writing;
|
||||
# Can specify group or user-level restrictions:
|
||||
# {
|
||||
# "read": {
|
||||
# "group_ids": ["group_id1", "group_id2"],
|
||||
# "user_ids": ["user_id1", "user_id2"]
|
||||
# },
|
||||
# "write": {
|
||||
# "group_ids": ["group_id1", "group_id2"],
|
||||
# "user_ids": ["user_id1", "user_id2"]
|
||||
# }
|
||||
# }
|
||||
|
||||
created_at = Column(BigInteger)
|
||||
updated_at = Column(BigInteger)
|
||||
@@ -82,31 +65,25 @@ class KnowledgeModel(BaseModel):
|
||||
|
||||
meta: Optional[dict] = None
|
||||
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
||||
|
||||
created_at: int # timestamp in epoch
|
||||
updated_at: int # timestamp in epoch
|
||||
|
||||
|
||||
class KnowledgeFile(Base):
|
||||
__tablename__ = "knowledge_file"
|
||||
__tablename__ = 'knowledge_file'
|
||||
|
||||
id = Column(Text, unique=True, primary_key=True)
|
||||
|
||||
knowledge_id = Column(
|
||||
Text, ForeignKey("knowledge.id", ondelete="CASCADE"), nullable=False
|
||||
)
|
||||
file_id = Column(Text, ForeignKey("file.id", ondelete="CASCADE"), nullable=False)
|
||||
knowledge_id = Column(Text, ForeignKey('knowledge.id', ondelete='CASCADE'), nullable=False)
|
||||
file_id = Column(Text, ForeignKey('file.id', ondelete='CASCADE'), nullable=False)
|
||||
user_id = Column(Text, nullable=False)
|
||||
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=False)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"knowledge_id", "file_id", name="uq_knowledge_file_knowledge_file"
|
||||
),
|
||||
)
|
||||
__table_args__ = (UniqueConstraint('knowledge_id', 'file_id', name='uq_knowledge_file_knowledge_file'),)
|
||||
|
||||
|
||||
class KnowledgeFileModel(BaseModel):
|
||||
@@ -139,7 +116,7 @@ class KnowledgeUserResponse(KnowledgeUserModel):
|
||||
class KnowledgeForm(BaseModel):
|
||||
name: str
|
||||
description: str
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class FileUserResponse(FileModelResponse):
|
||||
@@ -157,43 +134,61 @@ class KnowledgeFileListResponse(BaseModel):
|
||||
|
||||
|
||||
class KnowledgeTable:
|
||||
def insert_new_knowledge(
|
||||
self, user_id: str, form_data: KnowledgeForm, db: Optional[Session] = None
|
||||
async def _get_access_grants(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
|
||||
return await AccessGrants.get_grants_by_resource('knowledge', knowledge_id, db=db)
|
||||
|
||||
async def _to_knowledge_model(
|
||||
self,
|
||||
knowledge: Knowledge,
|
||||
access_grants: Optional[list[AccessGrantModel]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> KnowledgeModel:
|
||||
knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump(exclude={'access_grants'})
|
||||
knowledge_data['access_grants'] = (
|
||||
access_grants if access_grants is not None else await self._get_access_grants(knowledge_data['id'], db=db)
|
||||
)
|
||||
return KnowledgeModel.model_validate(knowledge_data)
|
||||
|
||||
async def insert_new_knowledge(
|
||||
self, user_id: str, form_data: KnowledgeForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[KnowledgeModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
knowledge = KnowledgeModel(
|
||||
**{
|
||||
**form_data.model_dump(),
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_id": user_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
**form_data.model_dump(exclude={'access_grants'}),
|
||||
'id': str(uuid.uuid4()),
|
||||
'user_id': user_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
'access_grants': [],
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
result = Knowledge(**knowledge.model_dump())
|
||||
result = Knowledge(**knowledge.model_dump(exclude={'access_grants'}))
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
await AccessGrants.set_access_grants('knowledge', result.id, form_data.access_grants, db=db)
|
||||
if result:
|
||||
return KnowledgeModel.model_validate(result)
|
||||
return await self._to_knowledge_model(result, db=db)
|
||||
else:
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_knowledge_bases(
|
||||
self, skip: int = 0, limit: int = 30, db: Optional[Session] = None
|
||||
async def get_knowledge_bases(
|
||||
self, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None
|
||||
) -> list[KnowledgeUserModel]:
|
||||
with get_db_context(db) as db:
|
||||
all_knowledge = (
|
||||
db.query(Knowledge).order_by(Knowledge.updated_at.desc()).all()
|
||||
)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Knowledge).order_by(Knowledge.updated_at.desc()))
|
||||
all_knowledge = result.scalars().all()
|
||||
user_ids = list(set(knowledge.user_id for knowledge in all_knowledge))
|
||||
knowledge_ids = [knowledge.id for knowledge in all_knowledge]
|
||||
|
||||
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users_dict = {user.id: user for user in users}
|
||||
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
|
||||
|
||||
knowledge_bases = []
|
||||
for knowledge in all_knowledge:
|
||||
@@ -201,68 +196,87 @@ class KnowledgeTable:
|
||||
knowledge_bases.append(
|
||||
KnowledgeUserModel.model_validate(
|
||||
{
|
||||
**KnowledgeModel.model_validate(knowledge).model_dump(),
|
||||
"user": user.model_dump() if user else None,
|
||||
**(
|
||||
await self._to_knowledge_model(
|
||||
knowledge,
|
||||
access_grants=grants_map.get(knowledge.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
'user': user.model_dump() if user else None,
|
||||
}
|
||||
)
|
||||
)
|
||||
return knowledge_bases
|
||||
|
||||
def search_knowledge_bases(
|
||||
async def search_knowledge_bases(
|
||||
self,
|
||||
user_id: str,
|
||||
filter: dict,
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> KnowledgeListResponse:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Knowledge, User).outerjoin(
|
||||
User, User.id == Knowledge.user_id
|
||||
)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Knowledge, User).outerjoin(User, User.id == Knowledge.user_id)
|
||||
|
||||
if filter:
|
||||
query_key = filter.get("query")
|
||||
query_key = filter.get('query')
|
||||
if query_key:
|
||||
query = query.filter(
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
Knowledge.name.ilike(f"%{query_key}%"),
|
||||
Knowledge.description.ilike(f"%{query_key}%"),
|
||||
Knowledge.name.ilike(f'%{query_key}%'),
|
||||
Knowledge.description.ilike(f'%{query_key}%'),
|
||||
User.name.ilike(f'%{query_key}%'),
|
||||
User.email.ilike(f'%{query_key}%'),
|
||||
User.username.ilike(f'%{query_key}%'),
|
||||
)
|
||||
)
|
||||
|
||||
view_option = filter.get("view_option")
|
||||
if view_option == "created":
|
||||
query = query.filter(Knowledge.user_id == user_id)
|
||||
elif view_option == "shared":
|
||||
query = query.filter(Knowledge.user_id != user_id)
|
||||
view_option = filter.get('view_option')
|
||||
if view_option == 'created':
|
||||
stmt = stmt.filter(Knowledge.user_id == user_id)
|
||||
elif view_option == 'shared':
|
||||
stmt = stmt.filter(Knowledge.user_id != user_id)
|
||||
|
||||
query = has_permission(db, Knowledge, query, filter)
|
||||
stmt = AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=stmt,
|
||||
DocumentModel=Knowledge,
|
||||
filter=filter,
|
||||
resource_type='knowledge',
|
||||
permission='read',
|
||||
)
|
||||
|
||||
query = query.order_by(Knowledge.updated_at.desc())
|
||||
stmt = stmt.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc())
|
||||
|
||||
total = query.count()
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
items = query.all()
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
knowledge_ids = [kb.id for kb, _ in items]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
|
||||
|
||||
knowledge_bases = []
|
||||
for knowledge_base, user in items:
|
||||
knowledge_bases.append(
|
||||
KnowledgeUserModel.model_validate(
|
||||
{
|
||||
**KnowledgeModel.model_validate(
|
||||
knowledge_base
|
||||
**(
|
||||
await self._to_knowledge_model(
|
||||
knowledge_base,
|
||||
access_grants=grants_map.get(knowledge_base.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
"user": (
|
||||
UserModel.model_validate(user).model_dump()
|
||||
if user
|
||||
else None
|
||||
),
|
||||
'user': (UserModel.model_validate(user).model_dump() if user else None),
|
||||
}
|
||||
)
|
||||
)
|
||||
@@ -272,218 +286,227 @@ class KnowledgeTable:
|
||||
print(e)
|
||||
return KnowledgeListResponse(items=[], total=0)
|
||||
|
||||
def search_knowledge_files(
|
||||
self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[Session] = None
|
||||
async def search_knowledge_files(
|
||||
self, filter: dict, skip: int = 0, limit: int = 30, db: Optional[AsyncSession] = None
|
||||
) -> KnowledgeFileListResponse:
|
||||
"""
|
||||
Scalable version: search files across all knowledge bases the user has
|
||||
READ access to, without loading all KBs or using large IN() lists.
|
||||
"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Base query: join Knowledge → KnowledgeFile → File
|
||||
query = (
|
||||
db.query(File, User, Knowledge)
|
||||
stmt = (
|
||||
select(File, User, Knowledge)
|
||||
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
|
||||
.join(Knowledge, KnowledgeFile.knowledge_id == Knowledge.id)
|
||||
.outerjoin(User, User.id == KnowledgeFile.user_id)
|
||||
)
|
||||
|
||||
# Apply access-control directly to the joined query
|
||||
# This makes the database handle filtering, even with 10k+ KBs
|
||||
query = has_permission(db, Knowledge, query, filter)
|
||||
stmt = AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=stmt,
|
||||
DocumentModel=Knowledge,
|
||||
filter=filter,
|
||||
resource_type='knowledge',
|
||||
permission='read',
|
||||
)
|
||||
|
||||
# Apply filename search
|
||||
if filter:
|
||||
q = filter.get("query")
|
||||
q = filter.get('query')
|
||||
if q:
|
||||
query = query.filter(File.filename.ilike(f"%{q}%"))
|
||||
stmt = stmt.filter(File.filename.ilike(f'%{q}%'))
|
||||
|
||||
# Order by file changes
|
||||
query = query.order_by(File.updated_at.desc())
|
||||
stmt = stmt.order_by(File.updated_at.desc(), File.id.asc())
|
||||
|
||||
# Count before pagination
|
||||
total = query.count()
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
rows = query.all()
|
||||
result = await db.execute(stmt)
|
||||
rows = result.all()
|
||||
|
||||
items = []
|
||||
for file, user, knowledge in rows:
|
||||
items.append(
|
||||
FileUserResponse(
|
||||
**FileModel.model_validate(file).model_dump(),
|
||||
user=(
|
||||
UserResponse(
|
||||
**UserModel.model_validate(user).model_dump()
|
||||
)
|
||||
if user
|
||||
else None
|
||||
),
|
||||
collection=KnowledgeModel.model_validate(
|
||||
knowledge
|
||||
).model_dump(),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
collection=(await self._to_knowledge_model(knowledge, db=db)).model_dump(),
|
||||
)
|
||||
)
|
||||
|
||||
return KnowledgeFileListResponse(items=items, total=total)
|
||||
|
||||
except Exception as e:
|
||||
print("search_knowledge_files error:", e)
|
||||
print('search_knowledge_files error:', e)
|
||||
return KnowledgeFileListResponse(items=[], total=0)
|
||||
|
||||
def check_access_by_user_id(
|
||||
self, id, user_id, permission="write", db: Optional[Session] = None
|
||||
) -> bool:
|
||||
knowledge = self.get_knowledge_by_id(id, db=db)
|
||||
async def check_access_by_user_id(self, id, user_id, permission='write', db: Optional[AsyncSession] = None) -> bool:
|
||||
knowledge = await self.get_knowledge_by_id(id, db=db)
|
||||
if not knowledge:
|
||||
return False
|
||||
if knowledge.user_id == user_id:
|
||||
return True
|
||||
user_group_ids = {
|
||||
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
|
||||
}
|
||||
return has_access(user_id, permission, knowledge.access_control, user_group_ids)
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
return await AccessGrants.has_access(
|
||||
user_id=user_id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge.id,
|
||||
permission=permission,
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
)
|
||||
|
||||
def get_knowledge_bases_by_user_id(
|
||||
self, user_id: str, permission: str = "write", db: Optional[Session] = None
|
||||
async def get_knowledge_bases_by_user_id(
|
||||
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
|
||||
) -> list[KnowledgeUserModel]:
|
||||
knowledge_bases = self.get_knowledge_bases(db=db)
|
||||
user_group_ids = {
|
||||
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
|
||||
}
|
||||
return [
|
||||
knowledge_base
|
||||
for knowledge_base in knowledge_bases
|
||||
if knowledge_base.user_id == user_id
|
||||
or has_access(
|
||||
user_id, permission, knowledge_base.access_control, user_group_ids
|
||||
)
|
||||
]
|
||||
knowledge_bases = await self.get_knowledge_bases(db=db)
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
|
||||
def get_knowledge_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[KnowledgeModel]:
|
||||
result = []
|
||||
for knowledge_base in knowledge_bases:
|
||||
if knowledge_base.user_id == user_id:
|
||||
result.append(knowledge_base)
|
||||
elif await AccessGrants.has_access(
|
||||
user_id=user_id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge_base.id,
|
||||
permission=permission,
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
):
|
||||
result.append(knowledge_base)
|
||||
return result
|
||||
|
||||
async def get_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
knowledge = db.query(Knowledge).filter_by(id=id).first()
|
||||
return KnowledgeModel.model_validate(knowledge) if knowledge else None
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Knowledge).filter_by(id=id))
|
||||
knowledge = result.scalars().first()
|
||||
return await self._to_knowledge_model(knowledge, db=db) if knowledge else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_knowledge_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_knowledge_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[KnowledgeModel]:
|
||||
knowledge = self.get_knowledge_by_id(id, db=db)
|
||||
knowledge = await self.get_knowledge_by_id(id, db=db)
|
||||
if not knowledge:
|
||||
return None
|
||||
|
||||
if knowledge.user_id == user_id:
|
||||
return knowledge
|
||||
|
||||
user_group_ids = {
|
||||
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
|
||||
}
|
||||
if has_access(user_id, "write", knowledge.access_control, user_group_ids):
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
if await AccessGrants.has_access(
|
||||
user_id=user_id,
|
||||
resource_type='knowledge',
|
||||
resource_id=knowledge.id,
|
||||
permission='write',
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
):
|
||||
return knowledge
|
||||
return None
|
||||
|
||||
def get_knowledges_by_file_id(
|
||||
self, file_id: str, db: Optional[Session] = None
|
||||
) -> list[KnowledgeModel]:
|
||||
async def get_knowledges_by_file_id(self, file_id: str, db: Optional[AsyncSession] = None) -> list[KnowledgeModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
knowledges = (
|
||||
db.query(Knowledge)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Knowledge)
|
||||
.join(KnowledgeFile, Knowledge.id == KnowledgeFile.knowledge_id)
|
||||
.filter(KnowledgeFile.file_id == file_id)
|
||||
.all()
|
||||
)
|
||||
knowledges = result.scalars().all()
|
||||
knowledge_ids = [k.id for k in knowledges]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('knowledge', knowledge_ids, db=db)
|
||||
return [
|
||||
KnowledgeModel.model_validate(knowledge) for knowledge in knowledges
|
||||
await self._to_knowledge_model(
|
||||
knowledge,
|
||||
access_grants=grants_map.get(knowledge.id, []),
|
||||
db=db,
|
||||
)
|
||||
for knowledge in knowledges
|
||||
]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def search_files_by_id(
|
||||
async def search_files_by_id(
|
||||
self,
|
||||
knowledge_id: str,
|
||||
user_id: str,
|
||||
filter: dict,
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> KnowledgeFileListResponse:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
query = (
|
||||
db.query(File, User)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = (
|
||||
select(File, User)
|
||||
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
|
||||
.outerjoin(User, User.id == KnowledgeFile.user_id)
|
||||
.filter(KnowledgeFile.knowledge_id == knowledge_id)
|
||||
)
|
||||
|
||||
# Default sort: updated_at descending
|
||||
primary_sort = File.updated_at.desc()
|
||||
|
||||
if filter:
|
||||
query_key = filter.get("query")
|
||||
query_key = filter.get('query')
|
||||
if query_key:
|
||||
query = query.filter(or_(File.filename.ilike(f"%{query_key}%")))
|
||||
stmt = stmt.filter(or_(File.filename.ilike(f'%{query_key}%')))
|
||||
|
||||
view_option = filter.get("view_option")
|
||||
if view_option == "created":
|
||||
query = query.filter(KnowledgeFile.user_id == user_id)
|
||||
elif view_option == "shared":
|
||||
query = query.filter(KnowledgeFile.user_id != user_id)
|
||||
view_option = filter.get('view_option')
|
||||
if view_option == 'created':
|
||||
stmt = stmt.filter(KnowledgeFile.user_id == user_id)
|
||||
elif view_option == 'shared':
|
||||
stmt = stmt.filter(KnowledgeFile.user_id != user_id)
|
||||
|
||||
order_by = filter.get("order_by")
|
||||
direction = filter.get("direction")
|
||||
order_by = filter.get('order_by')
|
||||
direction = filter.get('direction')
|
||||
is_asc = direction == 'asc'
|
||||
|
||||
if order_by == "name":
|
||||
if direction == "asc":
|
||||
query = query.order_by(File.filename.asc())
|
||||
else:
|
||||
query = query.order_by(File.filename.desc())
|
||||
elif order_by == "created_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(File.created_at.asc())
|
||||
else:
|
||||
query = query.order_by(File.created_at.desc())
|
||||
elif order_by == "updated_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(File.updated_at.asc())
|
||||
else:
|
||||
query = query.order_by(File.updated_at.desc())
|
||||
else:
|
||||
query = query.order_by(File.updated_at.desc())
|
||||
if order_by == 'name':
|
||||
primary_sort = File.filename.asc() if is_asc else File.filename.desc()
|
||||
elif order_by == 'created_at':
|
||||
primary_sort = File.created_at.asc() if is_asc else File.created_at.desc()
|
||||
elif order_by == 'updated_at':
|
||||
primary_sort = File.updated_at.asc() if is_asc else File.updated_at.desc()
|
||||
|
||||
else:
|
||||
query = query.order_by(File.updated_at.desc())
|
||||
# Apply sort with secondary key for deterministic pagination
|
||||
stmt = stmt.order_by(primary_sort, File.id.asc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
total = query.count()
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
items = query.all()
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
files = []
|
||||
for file, user in items:
|
||||
files.append(
|
||||
FileUserResponse(
|
||||
**FileModel.model_validate(file).model_dump(),
|
||||
user=(
|
||||
UserResponse(
|
||||
**UserModel.model_validate(user).model_dump()
|
||||
)
|
||||
if user
|
||||
else None
|
||||
),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -492,55 +515,52 @@ class KnowledgeTable:
|
||||
print(e)
|
||||
return KnowledgeFileListResponse(items=[], total=0)
|
||||
|
||||
def get_files_by_id(
|
||||
self, knowledge_id: str, db: Optional[Session] = None
|
||||
) -> list[FileModel]:
|
||||
async def get_files_by_id(self, knowledge_id: str, db: Optional[AsyncSession] = None) -> list[FileModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
files = (
|
||||
db.query(File)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(File)
|
||||
.join(KnowledgeFile, File.id == KnowledgeFile.file_id)
|
||||
.filter(KnowledgeFile.knowledge_id == knowledge_id)
|
||||
.all()
|
||||
)
|
||||
files = result.scalars().all()
|
||||
return [FileModel.model_validate(file) for file in files]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def get_file_metadatas_by_id(
|
||||
self, knowledge_id: str, db: Optional[Session] = None
|
||||
async def get_file_metadatas_by_id(
|
||||
self, knowledge_id: str, db: Optional[AsyncSession] = None
|
||||
) -> list[FileMetadataResponse]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
files = self.get_files_by_id(knowledge_id, db=db)
|
||||
return [FileMetadataResponse(**file.model_dump()) for file in files]
|
||||
files = await self.get_files_by_id(knowledge_id, db=db)
|
||||
return [FileMetadataResponse(**file.model_dump()) for file in files]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def add_file_to_knowledge_by_id(
|
||||
async def add_file_to_knowledge_by_id(
|
||||
self,
|
||||
knowledge_id: str,
|
||||
file_id: str,
|
||||
user_id: str,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[KnowledgeFileModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
knowledge_file = KnowledgeFileModel(
|
||||
**{
|
||||
"id": str(uuid.uuid4()),
|
||||
"knowledge_id": knowledge_id,
|
||||
"file_id": file_id,
|
||||
"user_id": user_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
'id': str(uuid.uuid4()),
|
||||
'knowledge_id': knowledge_id,
|
||||
'file_id': file_id,
|
||||
'user_id': user_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
result = KnowledgeFile(**knowledge_file.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return KnowledgeFileModel.model_validate(result)
|
||||
else:
|
||||
@@ -548,95 +568,107 @@ class KnowledgeTable:
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def remove_file_from_knowledge_by_id(
|
||||
self, knowledge_id: str, file_id: str, db: Optional[Session] = None
|
||||
async def has_file(self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Check whether a file belongs to a knowledge base."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id).limit(1)
|
||||
)
|
||||
return result.scalars().first() is not None
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
async def remove_file_from_knowledge_by_id(
|
||||
self, knowledge_id: str, file_id: str, db: Optional[AsyncSession] = None
|
||||
) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(KnowledgeFile).filter_by(
|
||||
knowledge_id=knowledge_id, file_id=file_id
|
||||
).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=knowledge_id, file_id=file_id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def reset_knowledge_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[KnowledgeModel]:
|
||||
async def reset_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[KnowledgeModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Delete all knowledge_file entries for this knowledge_id
|
||||
db.query(KnowledgeFile).filter_by(knowledge_id=id).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(KnowledgeFile).filter_by(knowledge_id=id))
|
||||
await db.commit()
|
||||
|
||||
# Update the knowledge entry's updated_at timestamp
|
||||
db.query(Knowledge).filter_by(id=id).update(
|
||||
{
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
)
|
||||
db.commit()
|
||||
await db.execute(update(Knowledge).filter_by(id=id).values(updated_at=int(time.time())))
|
||||
await db.commit()
|
||||
|
||||
return self.get_knowledge_by_id(id=id, db=db)
|
||||
return await self.get_knowledge_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
def update_knowledge_by_id(
|
||||
async def update_knowledge_by_id(
|
||||
self,
|
||||
id: str,
|
||||
form_data: KnowledgeForm,
|
||||
overwrite: bool = False,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[KnowledgeModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
knowledge = self.get_knowledge_by_id(id=id, db=db)
|
||||
db.query(Knowledge).filter_by(id=id).update(
|
||||
{
|
||||
**form_data.model_dump(),
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(
|
||||
update(Knowledge)
|
||||
.filter_by(id=id)
|
||||
.values(
|
||||
**form_data.model_dump(exclude={'access_grants'}),
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
return self.get_knowledge_by_id(id=id, db=db)
|
||||
await db.commit()
|
||||
if form_data.access_grants is not None:
|
||||
await AccessGrants.set_access_grants('knowledge', id, form_data.access_grants, db=db)
|
||||
return await self.get_knowledge_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
def update_knowledge_data_by_id(
|
||||
self, id: str, data: dict, db: Optional[Session] = None
|
||||
async def update_knowledge_data_by_id(
|
||||
self, id: str, data: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[KnowledgeModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
knowledge = self.get_knowledge_by_id(id=id, db=db)
|
||||
db.query(Knowledge).filter_by(id=id).update(
|
||||
{
|
||||
"data": data,
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(
|
||||
update(Knowledge)
|
||||
.filter_by(id=id)
|
||||
.values(
|
||||
data=data,
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
return self.get_knowledge_by_id(id=id, db=db)
|
||||
await db.commit()
|
||||
return await self.get_knowledge_by_id(id=id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(e)
|
||||
return None
|
||||
|
||||
def delete_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
async def delete_knowledge_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Knowledge).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await AccessGrants.revoke_all_access('knowledge', id, db=db)
|
||||
await db.execute(delete(Knowledge).filter_by(id=id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_all_knowledge(self, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_all_knowledge(self, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Knowledge).delete()
|
||||
db.commit()
|
||||
result = await db.execute(select(Knowledge.id))
|
||||
knowledge_ids = [row[0] for row in result.all()]
|
||||
for knowledge_id in knowledge_ids:
|
||||
await AccessGrants.revoke_all_access('knowledge', knowledge_id, db=db)
|
||||
await db.execute(delete(Knowledge))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
|
||||
@@ -2,18 +2,21 @@ import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, get_db, get_db_context
|
||||
from sqlalchemy import select, delete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, String, Text
|
||||
|
||||
####################
|
||||
# Memory DB Schema
|
||||
# What was learned at cost should not need to be paid
|
||||
# for again. Let the memory hold.
|
||||
####################
|
||||
|
||||
|
||||
class Memory(Base):
|
||||
__tablename__ = "memory"
|
||||
__tablename__ = 'memory'
|
||||
|
||||
id = Column(String, primary_key=True, unique=True)
|
||||
user_id = Column(String)
|
||||
@@ -38,117 +41,112 @@ class MemoryModel(BaseModel):
|
||||
|
||||
|
||||
class MemoriesTable:
|
||||
def insert_new_memory(
|
||||
async def insert_new_memory(
|
||||
self,
|
||||
user_id: str,
|
||||
content: str,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[MemoryModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
id = str(uuid.uuid4())
|
||||
|
||||
memory = MemoryModel(
|
||||
**{
|
||||
"id": id,
|
||||
"user_id": user_id,
|
||||
"content": content,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
'id': id,
|
||||
'user_id': user_id,
|
||||
'content': content,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
result = Memory(**memory.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
if result:
|
||||
return MemoryModel.model_validate(result)
|
||||
else:
|
||||
return None
|
||||
|
||||
def update_memory_by_id_and_user_id(
|
||||
async def update_memory_by_id_and_user_id(
|
||||
self,
|
||||
id: str,
|
||||
user_id: str,
|
||||
content: str,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[MemoryModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
memory = db.get(Memory, id)
|
||||
memory = await db.get(Memory, id)
|
||||
if not memory or memory.user_id != user_id:
|
||||
return None
|
||||
|
||||
memory.content = content
|
||||
memory.updated_at = int(time.time())
|
||||
|
||||
db.commit()
|
||||
return self.get_memory_by_id(id)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_memories(self, db: Optional[Session] = None) -> list[MemoryModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
memories = db.query(Memory).all()
|
||||
return [MemoryModel.model_validate(memory) for memory in memories]
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_memories_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> list[MemoryModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
memories = db.query(Memory).filter_by(user_id=user_id).all()
|
||||
return [MemoryModel.model_validate(memory) for memory in memories]
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_memory_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[MemoryModel]:
|
||||
with get_db_context(db) as db:
|
||||
try:
|
||||
memory = db.get(Memory, id)
|
||||
await db.commit()
|
||||
await db.refresh(memory)
|
||||
return MemoryModel.model_validate(memory)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def delete_memory_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def get_memories(self, db: Optional[AsyncSession] = None) -> list[MemoryModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Memory).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
result = await db.execute(select(Memory))
|
||||
memories = result.scalars().all()
|
||||
return [MemoryModel.model_validate(memory) for memory in memories]
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def get_memories_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> list[MemoryModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
result = await db.execute(select(Memory).filter_by(user_id=user_id))
|
||||
memories = result.scalars().all()
|
||||
return [MemoryModel.model_validate(memory) for memory in memories]
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def get_memory_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[MemoryModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
memory = await db.get(Memory, id)
|
||||
return MemoryModel.model_validate(memory) if memory else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def delete_memory_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
await db.execute(delete(Memory).filter_by(id=id))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_memories_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_memories_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
db.query(Memory).filter_by(user_id=user_id).delete()
|
||||
db.commit()
|
||||
await db.execute(delete(Memory).filter_by(user_id=user_id))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_memory_by_id_and_user_id(
|
||||
self, id: str, user_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
async def delete_memory_by_id_and_user_id(self, id: str, user_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
memory = db.get(Memory, id)
|
||||
memory = await db.get(Memory, id)
|
||||
if not memory or memory.user_id != user_id:
|
||||
return None
|
||||
|
||||
# Delete the memory
|
||||
db.delete(memory)
|
||||
db.commit()
|
||||
await db.delete(memory)
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
|
||||
@@ -3,8 +3,9 @@ import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.tags import TagModel, Tag, Tags
|
||||
from open_webui.models.users import Users, User, UserNameResponse
|
||||
from open_webui.models.channels import Channels, ChannelMember
|
||||
@@ -12,7 +13,7 @@ from open_webui.models.channels import Channels, ChannelMember
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, field_validator
|
||||
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON
|
||||
from sqlalchemy import or_, func, select, and_, text
|
||||
from sqlalchemy import or_, func, and_, text
|
||||
from sqlalchemy.sql import exists
|
||||
|
||||
####################
|
||||
@@ -21,7 +22,7 @@ from sqlalchemy.sql import exists
|
||||
|
||||
|
||||
class MessageReaction(Base):
|
||||
__tablename__ = "message_reaction"
|
||||
__tablename__ = 'message_reaction'
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
user_id = Column(Text)
|
||||
message_id = Column(Text)
|
||||
@@ -40,7 +41,7 @@ class MessageReactionModel(BaseModel):
|
||||
|
||||
|
||||
class Message(Base):
|
||||
__tablename__ = "message"
|
||||
__tablename__ = 'message'
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
|
||||
user_id = Column(Text)
|
||||
@@ -112,7 +113,7 @@ class MessageUserResponse(MessageModel):
|
||||
class MessageUserSlimResponse(MessageUserResponse):
|
||||
data: bool | None = None
|
||||
|
||||
@field_validator("data", mode="before")
|
||||
@field_validator('data', mode='before')
|
||||
def convert_data_to_bool(cls, v):
|
||||
# No data or not a dict → False
|
||||
if not isinstance(v, dict):
|
||||
@@ -137,248 +138,211 @@ class MessageResponse(MessageReplyToResponse):
|
||||
|
||||
|
||||
class MessageTable:
|
||||
def insert_new_message(
|
||||
async def insert_new_message(
|
||||
self,
|
||||
form_data: MessageForm,
|
||||
channel_id: str,
|
||||
user_id: str,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[MessageModel]:
|
||||
with get_db_context(db) as db:
|
||||
channel_member = Channels.join_channel(channel_id, user_id)
|
||||
async with get_async_db_context(db) as db:
|
||||
channel_member = await Channels.join_channel(channel_id, user_id)
|
||||
|
||||
id = str(uuid.uuid4())
|
||||
ts = int(time.time_ns())
|
||||
|
||||
message = MessageModel(
|
||||
**{
|
||||
"id": id,
|
||||
"user_id": user_id,
|
||||
"channel_id": channel_id,
|
||||
"reply_to_id": form_data.reply_to_id,
|
||||
"parent_id": form_data.parent_id,
|
||||
"is_pinned": False,
|
||||
"pinned_at": None,
|
||||
"pinned_by": None,
|
||||
"content": form_data.content,
|
||||
"data": form_data.data,
|
||||
"meta": form_data.meta,
|
||||
"created_at": ts,
|
||||
"updated_at": ts,
|
||||
'id': id,
|
||||
'user_id': user_id,
|
||||
'channel_id': channel_id,
|
||||
'reply_to_id': form_data.reply_to_id,
|
||||
'parent_id': form_data.parent_id,
|
||||
'is_pinned': False,
|
||||
'pinned_at': None,
|
||||
'pinned_by': None,
|
||||
'content': form_data.content,
|
||||
'data': form_data.data,
|
||||
'meta': form_data.meta,
|
||||
'created_at': ts,
|
||||
'updated_at': ts,
|
||||
}
|
||||
)
|
||||
result = Message(**message.model_dump())
|
||||
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
return MessageModel.model_validate(result) if result else None
|
||||
|
||||
def get_message_by_id(
|
||||
async def get_message_by_id(
|
||||
self,
|
||||
id: str,
|
||||
include_thread_replies: Optional[bool] = True,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[MessageResponse]:
|
||||
with get_db_context(db) as db:
|
||||
message = db.get(Message, id)
|
||||
async with get_async_db_context(db) as db:
|
||||
message = await db.get(Message, id)
|
||||
if not message:
|
||||
return None
|
||||
|
||||
reply_to_message = (
|
||||
self.get_message_by_id(
|
||||
message.reply_to_id, include_thread_replies=False, db=db
|
||||
)
|
||||
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
|
||||
if message.reply_to_id
|
||||
else None
|
||||
)
|
||||
|
||||
reactions = self.get_reactions_by_message_id(id, db=db)
|
||||
reactions = await self.get_reactions_by_message_id(id, db=db)
|
||||
|
||||
thread_replies = []
|
||||
if include_thread_replies:
|
||||
thread_replies = self.get_thread_replies_by_message_id(id, db=db)
|
||||
thread_replies = await self.get_thread_replies_by_message_id(id, db=db)
|
||||
|
||||
# Check if message was sent by webhook (webhook info in meta takes precedence)
|
||||
webhook_info = message.meta.get("webhook") if message.meta else None
|
||||
if webhook_info and webhook_info.get("id"):
|
||||
webhook_info = message.meta.get('webhook') if message.meta else None
|
||||
if webhook_info and webhook_info.get('id'):
|
||||
# Look up webhook by ID to get current name
|
||||
webhook = Channels.get_webhook_by_id(webhook_info.get("id"), db=db)
|
||||
webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
|
||||
if webhook:
|
||||
user_info = {
|
||||
"id": webhook.id,
|
||||
"name": webhook.name,
|
||||
"role": "webhook",
|
||||
'id': webhook.id,
|
||||
'name': webhook.name,
|
||||
'role': 'webhook',
|
||||
}
|
||||
else:
|
||||
# Webhook was deleted, use placeholder
|
||||
user_info = {
|
||||
"id": webhook_info.get("id"),
|
||||
"name": "Deleted Webhook",
|
||||
"role": "webhook",
|
||||
'id': webhook_info.get('id'),
|
||||
'name': 'Deleted Webhook',
|
||||
'role': 'webhook',
|
||||
}
|
||||
else:
|
||||
user = Users.get_user_by_id(message.user_id, db=db)
|
||||
user = await Users.get_user_by_id(message.user_id, db=db)
|
||||
user_info = user.model_dump() if user else None
|
||||
|
||||
return MessageResponse.model_validate(
|
||||
{
|
||||
**MessageModel.model_validate(message).model_dump(),
|
||||
"user": user_info,
|
||||
"reply_to_message": (
|
||||
reply_to_message.model_dump() if reply_to_message else None
|
||||
),
|
||||
"latest_reply_at": (
|
||||
thread_replies[0].created_at if thread_replies else None
|
||||
),
|
||||
"reply_count": len(thread_replies),
|
||||
"reactions": reactions,
|
||||
'user': user_info,
|
||||
'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
|
||||
'latest_reply_at': (thread_replies[0].created_at if thread_replies else None),
|
||||
'reply_count': len(thread_replies),
|
||||
'reactions': reactions,
|
||||
}
|
||||
)
|
||||
|
||||
def get_thread_replies_by_message_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
async def _resolve_user_info(self, message: Message, db: AsyncSession) -> Optional[dict]:
|
||||
"""Resolve user info from message, handling webhook messages."""
|
||||
webhook_info = message.meta.get('webhook') if message.meta else None
|
||||
if webhook_info and webhook_info.get('id'):
|
||||
webhook = await Channels.get_webhook_by_id(webhook_info.get('id'), db=db)
|
||||
if webhook:
|
||||
return {
|
||||
'id': webhook.id,
|
||||
'name': webhook.name,
|
||||
'role': 'webhook',
|
||||
}
|
||||
else:
|
||||
return {
|
||||
'id': webhook_info.get('id'),
|
||||
'name': 'Deleted Webhook',
|
||||
'role': 'webhook',
|
||||
}
|
||||
return None
|
||||
|
||||
async def get_thread_replies_by_message_id(
|
||||
self, id: str, db: Optional[AsyncSession] = None
|
||||
) -> list[MessageReplyToResponse]:
|
||||
with get_db_context(db) as db:
|
||||
all_messages = (
|
||||
db.query(Message)
|
||||
.filter_by(parent_id=id)
|
||||
.order_by(Message.created_at.desc())
|
||||
.all()
|
||||
)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Message).filter_by(parent_id=id).order_by(Message.created_at.desc()))
|
||||
all_messages = result.scalars().all()
|
||||
|
||||
messages = []
|
||||
for message in all_messages:
|
||||
reply_to_message = (
|
||||
self.get_message_by_id(
|
||||
message.reply_to_id, include_thread_replies=False, db=db
|
||||
)
|
||||
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
|
||||
if message.reply_to_id
|
||||
else None
|
||||
)
|
||||
|
||||
webhook_info = message.meta.get("webhook") if message.meta else None
|
||||
user_info = None
|
||||
if webhook_info and webhook_info.get("id"):
|
||||
webhook = Channels.get_webhook_by_id(webhook_info.get("id"), db=db)
|
||||
if webhook:
|
||||
user_info = {
|
||||
"id": webhook.id,
|
||||
"name": webhook.name,
|
||||
"role": "webhook",
|
||||
}
|
||||
else:
|
||||
user_info = {
|
||||
"id": webhook_info.get("id"),
|
||||
"name": "Deleted Webhook",
|
||||
"role": "webhook",
|
||||
}
|
||||
user_info = await self._resolve_user_info(message, db)
|
||||
|
||||
messages.append(
|
||||
MessageReplyToResponse.model_validate(
|
||||
{
|
||||
**MessageModel.model_validate(message).model_dump(),
|
||||
"user": user_info,
|
||||
"reply_to_message": (
|
||||
reply_to_message.model_dump()
|
||||
if reply_to_message
|
||||
else None
|
||||
),
|
||||
'user': user_info,
|
||||
'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
|
||||
}
|
||||
)
|
||||
)
|
||||
return messages
|
||||
|
||||
def get_reply_user_ids_by_message_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> list[str]:
|
||||
with get_db_context(db) as db:
|
||||
return [
|
||||
message.user_id
|
||||
for message in db.query(Message).filter_by(parent_id=id).all()
|
||||
]
|
||||
async def get_reply_user_ids_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[str]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Message.user_id).filter_by(parent_id=id))
|
||||
return [row[0] for row in result.all()]
|
||||
|
||||
def get_messages_by_channel_id(
|
||||
async def get_messages_by_channel_id(
|
||||
self,
|
||||
channel_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[MessageReplyToResponse]:
|
||||
with get_db_context(db) as db:
|
||||
all_messages = (
|
||||
db.query(Message)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Message)
|
||||
.filter_by(channel_id=channel_id, parent_id=None)
|
||||
.order_by(Message.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
all_messages = result.scalars().all()
|
||||
|
||||
messages = []
|
||||
for message in all_messages:
|
||||
reply_to_message = (
|
||||
self.get_message_by_id(
|
||||
message.reply_to_id, include_thread_replies=False, db=db
|
||||
)
|
||||
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
|
||||
if message.reply_to_id
|
||||
else None
|
||||
)
|
||||
|
||||
webhook_info = message.meta.get("webhook") if message.meta else None
|
||||
user_info = None
|
||||
if webhook_info and webhook_info.get("id"):
|
||||
webhook = Channels.get_webhook_by_id(webhook_info.get("id"), db=db)
|
||||
if webhook:
|
||||
user_info = {
|
||||
"id": webhook.id,
|
||||
"name": webhook.name,
|
||||
"role": "webhook",
|
||||
}
|
||||
else:
|
||||
user_info = {
|
||||
"id": webhook_info.get("id"),
|
||||
"name": "Deleted Webhook",
|
||||
"role": "webhook",
|
||||
}
|
||||
user_info = await self._resolve_user_info(message, db)
|
||||
|
||||
messages.append(
|
||||
MessageReplyToResponse.model_validate(
|
||||
{
|
||||
**MessageModel.model_validate(message).model_dump(),
|
||||
"user": user_info,
|
||||
"reply_to_message": (
|
||||
reply_to_message.model_dump()
|
||||
if reply_to_message
|
||||
else None
|
||||
),
|
||||
'user': user_info,
|
||||
'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
|
||||
}
|
||||
)
|
||||
)
|
||||
return messages
|
||||
|
||||
def get_messages_by_parent_id(
|
||||
async def get_messages_by_parent_id(
|
||||
self,
|
||||
channel_id: str,
|
||||
parent_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[MessageReplyToResponse]:
|
||||
with get_db_context(db) as db:
|
||||
message = db.get(Message, parent_id)
|
||||
async with get_async_db_context(db) as db:
|
||||
message = await db.get(Message, parent_id)
|
||||
|
||||
if not message:
|
||||
return []
|
||||
|
||||
all_messages = (
|
||||
db.query(Message)
|
||||
result = await db.execute(
|
||||
select(Message)
|
||||
.filter_by(channel_id=channel_id, parent_id=parent_id)
|
||||
.order_by(Message.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
all_messages = list(result.scalars().all())
|
||||
|
||||
# If length of all_messages is less than limit, then add the parent message
|
||||
if len(all_messages) < limit:
|
||||
@@ -387,80 +351,57 @@ class MessageTable:
|
||||
messages = []
|
||||
for message in all_messages:
|
||||
reply_to_message = (
|
||||
self.get_message_by_id(
|
||||
message.reply_to_id, include_thread_replies=False, db=db
|
||||
)
|
||||
await self.get_message_by_id(message.reply_to_id, include_thread_replies=False, db=db)
|
||||
if message.reply_to_id
|
||||
else None
|
||||
)
|
||||
|
||||
webhook_info = message.meta.get("webhook") if message.meta else None
|
||||
user_info = None
|
||||
if webhook_info and webhook_info.get("id"):
|
||||
webhook = Channels.get_webhook_by_id(webhook_info.get("id"), db=db)
|
||||
if webhook:
|
||||
user_info = {
|
||||
"id": webhook.id,
|
||||
"name": webhook.name,
|
||||
"role": "webhook",
|
||||
}
|
||||
else:
|
||||
user_info = {
|
||||
"id": webhook_info.get("id"),
|
||||
"name": "Deleted Webhook",
|
||||
"role": "webhook",
|
||||
}
|
||||
user_info = await self._resolve_user_info(message, db)
|
||||
|
||||
messages.append(
|
||||
MessageReplyToResponse.model_validate(
|
||||
{
|
||||
**MessageModel.model_validate(message).model_dump(),
|
||||
"user": user_info,
|
||||
"reply_to_message": (
|
||||
reply_to_message.model_dump()
|
||||
if reply_to_message
|
||||
else None
|
||||
),
|
||||
'user': user_info,
|
||||
'reply_to_message': (reply_to_message.model_dump() if reply_to_message else None),
|
||||
}
|
||||
)
|
||||
)
|
||||
return messages
|
||||
|
||||
def get_last_message_by_channel_id(
|
||||
self, channel_id: str, db: Optional[Session] = None
|
||||
async def get_last_message_by_channel_id(
|
||||
self, channel_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[MessageModel]:
|
||||
with get_db_context(db) as db:
|
||||
message = (
|
||||
db.query(Message)
|
||||
.filter_by(channel_id=channel_id)
|
||||
.order_by(Message.created_at.desc())
|
||||
.first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Message).filter_by(channel_id=channel_id).order_by(Message.created_at.desc()).limit(1)
|
||||
)
|
||||
message = result.scalars().first()
|
||||
return MessageModel.model_validate(message) if message else None
|
||||
|
||||
def get_pinned_messages_by_channel_id(
|
||||
async def get_pinned_messages_by_channel_id(
|
||||
self,
|
||||
channel_id: str,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[MessageModel]:
|
||||
with get_db_context(db) as db:
|
||||
all_messages = (
|
||||
db.query(Message)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Message)
|
||||
.filter_by(channel_id=channel_id, is_pinned=True)
|
||||
.order_by(Message.pinned_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
)
|
||||
all_messages = result.scalars().all()
|
||||
return [MessageModel.model_validate(message) for message in all_messages]
|
||||
|
||||
def update_message_by_id(
|
||||
self, id: str, form_data: MessageForm, db: Optional[Session] = None
|
||||
async def update_message_by_id(
|
||||
self, id: str, form_data: MessageForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[MessageModel]:
|
||||
with get_db_context(db) as db:
|
||||
message = db.get(Message, id)
|
||||
async with get_async_db_context(db) as db:
|
||||
message = await db.get(Message, id)
|
||||
message.content = form_data.content
|
||||
message.data = {
|
||||
**(message.data if message.data else {}),
|
||||
@@ -471,53 +412,51 @@ class MessageTable:
|
||||
**(form_data.meta if form_data.meta else {}),
|
||||
}
|
||||
message.updated_at = int(time.time_ns())
|
||||
db.commit()
|
||||
db.refresh(message)
|
||||
await db.commit()
|
||||
await db.refresh(message)
|
||||
return MessageModel.model_validate(message) if message else None
|
||||
|
||||
def update_is_pinned_by_id(
|
||||
async def update_is_pinned_by_id(
|
||||
self,
|
||||
id: str,
|
||||
is_pinned: bool,
|
||||
pinned_by: Optional[str] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[MessageModel]:
|
||||
with get_db_context(db) as db:
|
||||
message = db.get(Message, id)
|
||||
async with get_async_db_context(db) as db:
|
||||
message = await db.get(Message, id)
|
||||
message.is_pinned = is_pinned
|
||||
message.pinned_at = int(time.time_ns()) if is_pinned else None
|
||||
message.pinned_by = pinned_by if is_pinned else None
|
||||
db.commit()
|
||||
db.refresh(message)
|
||||
await db.commit()
|
||||
await db.refresh(message)
|
||||
return MessageModel.model_validate(message) if message else None
|
||||
|
||||
def get_unread_message_count(
|
||||
async def get_unread_message_count(
|
||||
self,
|
||||
channel_id: str,
|
||||
user_id: str,
|
||||
last_read_at: Optional[int] = None,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> int:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Message).filter(
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(func.count(Message.id)).filter(
|
||||
Message.channel_id == channel_id,
|
||||
Message.parent_id == None, # only count top-level messages
|
||||
Message.created_at > (last_read_at if last_read_at else 0),
|
||||
)
|
||||
if user_id:
|
||||
query = query.filter(Message.user_id != user_id)
|
||||
return query.count()
|
||||
stmt = stmt.filter(Message.user_id != user_id)
|
||||
result = await db.execute(stmt)
|
||||
return result.scalar()
|
||||
|
||||
def add_reaction_to_message(
|
||||
self, id: str, user_id: str, name: str, db: Optional[Session] = None
|
||||
async def add_reaction_to_message(
|
||||
self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[MessageReactionModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# check for existing reaction
|
||||
existing_reaction = (
|
||||
db.query(MessageReaction)
|
||||
.filter_by(message_id=id, user_id=user_id, name=name)
|
||||
.first()
|
||||
)
|
||||
result = await db.execute(select(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name))
|
||||
existing_reaction = result.scalars().first()
|
||||
if existing_reaction:
|
||||
return MessageReactionModel.model_validate(existing_reaction)
|
||||
|
||||
@@ -531,102 +470,94 @@ class MessageTable:
|
||||
)
|
||||
result = MessageReaction(**reaction.model_dump())
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
return MessageReactionModel.model_validate(result) if result else None
|
||||
|
||||
def get_reactions_by_message_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> list[Reactions]:
|
||||
with get_db_context(db) as db:
|
||||
async def get_reactions_by_message_id(self, id: str, db: Optional[AsyncSession] = None) -> list[Reactions]:
|
||||
async with get_async_db_context(db) as db:
|
||||
# JOIN User so all user info is fetched in one query
|
||||
results = (
|
||||
db.query(MessageReaction, User)
|
||||
result = await db.execute(
|
||||
select(MessageReaction, User)
|
||||
.join(User, MessageReaction.user_id == User.id)
|
||||
.filter(MessageReaction.message_id == id)
|
||||
.all()
|
||||
)
|
||||
results = result.all()
|
||||
|
||||
reactions = {}
|
||||
|
||||
for reaction, user in results:
|
||||
if reaction.name not in reactions:
|
||||
reactions[reaction.name] = {
|
||||
"name": reaction.name,
|
||||
"users": [],
|
||||
"count": 0,
|
||||
'name': reaction.name,
|
||||
'users': [],
|
||||
'count': 0,
|
||||
}
|
||||
|
||||
reactions[reaction.name]["users"].append(
|
||||
reactions[reaction.name]['users'].append(
|
||||
{
|
||||
"id": user.id,
|
||||
"name": user.name,
|
||||
'id': user.id,
|
||||
'name': user.name,
|
||||
}
|
||||
)
|
||||
reactions[reaction.name]["count"] += 1
|
||||
reactions[reaction.name]['count'] += 1
|
||||
|
||||
return [Reactions(**reaction) for reaction in reactions.values()]
|
||||
|
||||
def remove_reaction_by_id_and_user_id_and_name(
|
||||
self, id: str, user_id: str, name: str, db: Optional[Session] = None
|
||||
async def remove_reaction_by_id_and_user_id_and_name(
|
||||
self, id: str, user_id: str, name: str, db: Optional[AsyncSession] = None
|
||||
) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
db.query(MessageReaction).filter_by(
|
||||
message_id=id, user_id=user_id, name=name
|
||||
).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(MessageReaction).filter_by(message_id=id, user_id=user_id, name=name))
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
def delete_reactions_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
db.query(MessageReaction).filter_by(message_id=id).delete()
|
||||
db.commit()
|
||||
async def delete_reactions_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(MessageReaction).filter_by(message_id=id))
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
def delete_replies_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Message).filter_by(parent_id=id).delete()
|
||||
db.commit()
|
||||
async def delete_replies_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(Message).filter_by(parent_id=id))
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
def delete_message_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Message).filter_by(id=id).delete()
|
||||
async def delete_message_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(Message).filter_by(id=id))
|
||||
|
||||
# Delete all reactions to this message
|
||||
db.query(MessageReaction).filter_by(message_id=id).delete()
|
||||
await db.execute(delete(MessageReaction).filter_by(message_id=id))
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
def search_messages_by_channel_ids(
|
||||
async def search_messages_by_channel_ids(
|
||||
self,
|
||||
channel_ids: list[str],
|
||||
query: str,
|
||||
start_timestamp: Optional[int] = None,
|
||||
end_timestamp: Optional[int] = None,
|
||||
limit: int = 10,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[MessageModel]:
|
||||
"""Search messages in specified channels by content."""
|
||||
with get_db_context(db) as db:
|
||||
query_builder = db.query(Message).filter(
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Message).filter(
|
||||
Message.channel_id.in_(channel_ids),
|
||||
Message.content.ilike(f"%{query}%"),
|
||||
Message.content.ilike(f'%{query}%'),
|
||||
)
|
||||
|
||||
if start_timestamp:
|
||||
query_builder = query_builder.filter(
|
||||
Message.created_at >= start_timestamp
|
||||
)
|
||||
stmt = stmt.filter(Message.created_at >= start_timestamp)
|
||||
if end_timestamp:
|
||||
query_builder = query_builder.filter(
|
||||
Message.created_at <= end_timestamp
|
||||
)
|
||||
stmt = stmt.filter(Message.created_at <= end_timestamp)
|
||||
|
||||
messages = (
|
||||
query_builder.order_by(Message.created_at.desc()).limit(limit).all()
|
||||
)
|
||||
stmt = stmt.order_by(Message.created_at.desc()).limit(limit)
|
||||
result = await db.execute(stmt)
|
||||
messages = result.scalars().all()
|
||||
return [MessageModel.model_validate(msg) for msg in messages]
|
||||
|
||||
|
||||
|
||||
+337
-227
@@ -1,43 +1,41 @@
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, update, or_, func, String, cast
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import User, UserModel, Users, UserResponse
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from sqlalchemy import String, cast, or_, and_, func
|
||||
from sqlalchemy.dialects import postgresql, sqlite
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON, Boolean
|
||||
|
||||
|
||||
from open_webui.utils.access_control import has_access
|
||||
|
||||
from sqlalchemy import BigInteger, Column, Text, Boolean
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
####################
|
||||
# Models DB Schema
|
||||
# A misconfigured model wastes the time of everyone
|
||||
# who trusts it. Let what is set here be set with care.
|
||||
####################
|
||||
|
||||
|
||||
# ModelParams is a model for the data stored in the params field of the Model table
|
||||
class ModelParams(BaseModel):
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
pass
|
||||
|
||||
|
||||
# ModelMeta is a model for the data stored in the meta field of the Model table
|
||||
class ModelMeta(BaseModel):
|
||||
profile_image_url: Optional[str] = "/static/favicon.png"
|
||||
profile_image_url: Optional[str] = '/static/favicon.png'
|
||||
|
||||
description: Optional[str] = None
|
||||
"""
|
||||
@@ -46,13 +44,26 @@ class ModelMeta(BaseModel):
|
||||
|
||||
capabilities: Optional[dict] = None
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
model_config = ConfigDict(extra='allow')
|
||||
|
||||
pass
|
||||
@model_validator(mode='before')
|
||||
@classmethod
|
||||
def normalize_tags(cls, data):
|
||||
if isinstance(data, dict) and 'tags' in data:
|
||||
raw_tags = data['tags']
|
||||
if isinstance(raw_tags, list):
|
||||
normalized = []
|
||||
for tag in raw_tags:
|
||||
if isinstance(tag, str):
|
||||
normalized.append({'name': tag})
|
||||
elif isinstance(tag, dict) and 'name' in tag:
|
||||
normalized.append(tag)
|
||||
data['tags'] = normalized
|
||||
return data
|
||||
|
||||
|
||||
class Model(Base):
|
||||
__tablename__ = "model"
|
||||
__tablename__ = 'model'
|
||||
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
"""
|
||||
@@ -80,23 +91,6 @@ class Model(Base):
|
||||
Holds a JSON encoded blob of metadata, see `ModelMeta`.
|
||||
"""
|
||||
|
||||
access_control = Column(JSON, nullable=True) # Controls data access levels.
|
||||
# Defines access control rules for this entry.
|
||||
# - `None`: Public access, available to all users with the "user" role.
|
||||
# - `{}`: Private access, restricted exclusively to the owner.
|
||||
# - Custom permissions: Specific access control for reading and writing;
|
||||
# Can specify group or user-level restrictions:
|
||||
# {
|
||||
# "read": {
|
||||
# "group_ids": ["group_id1", "group_id2"],
|
||||
# "user_ids": ["user_id1", "user_id2"]
|
||||
# },
|
||||
# "write": {
|
||||
# "group_ids": ["group_id1", "group_id2"],
|
||||
# "user_ids": ["user_id1", "user_id2"]
|
||||
# }
|
||||
# }
|
||||
|
||||
is_active = Column(Boolean, default=True)
|
||||
|
||||
updated_at = Column(BigInteger)
|
||||
@@ -112,7 +106,7 @@ class ModelModel(BaseModel):
|
||||
params: ModelParams
|
||||
meta: ModelMeta
|
||||
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
||||
|
||||
is_active: bool
|
||||
updated_at: int # timestamp in epoch
|
||||
@@ -149,54 +143,81 @@ class ModelAccessListResponse(BaseModel):
|
||||
|
||||
|
||||
class ModelForm(BaseModel):
|
||||
model_config = ConfigDict(extra='ignore')
|
||||
|
||||
id: str
|
||||
base_model_id: Optional[str] = None
|
||||
name: str
|
||||
meta: ModelMeta
|
||||
params: ModelParams
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
class ModelsTable:
|
||||
def insert_new_model(
|
||||
self, form_data: ModelForm, user_id: str, db: Optional[Session] = None
|
||||
) -> Optional[ModelModel]:
|
||||
model = ModelModel(
|
||||
**{
|
||||
**form_data.model_dump(),
|
||||
"user_id": user_id,
|
||||
"created_at": int(time.time()),
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
async def _get_access_grants(self, model_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
|
||||
return await AccessGrants.get_grants_by_resource('model', model_id, db=db)
|
||||
|
||||
async def _to_model_model(
|
||||
self,
|
||||
model: Model,
|
||||
access_grants: Optional[list[AccessGrantModel]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> ModelModel:
|
||||
model_data = ModelModel.model_validate(model).model_dump(exclude={'access_grants'})
|
||||
model_data['access_grants'] = (
|
||||
access_grants if access_grants is not None else await self._get_access_grants(model_data['id'], db=db)
|
||||
)
|
||||
return ModelModel.model_validate(model_data)
|
||||
|
||||
async def insert_new_model(
|
||||
self, form_data: ModelForm, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[ModelModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
result = Model(**model.model_dump())
|
||||
async with get_async_db_context(db) as db:
|
||||
result = Model(
|
||||
**{
|
||||
**form_data.model_dump(exclude={'access_grants'}),
|
||||
'user_id': user_id,
|
||||
'created_at': int(time.time()),
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
await AccessGrants.set_access_grants('model', result.id, form_data.access_grants, db=db)
|
||||
|
||||
if result:
|
||||
return ModelModel.model_validate(result)
|
||||
return await self._to_model_model(result, db=db)
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.exception(f"Failed to insert a new model: {e}")
|
||||
log.exception(f'Failed to insert a new model: {e}')
|
||||
return None
|
||||
|
||||
def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]:
|
||||
with get_db_context(db) as db:
|
||||
return [ModelModel.model_validate(model) for model in db.query(Model).all()]
|
||||
async def get_all_models(self, db: Optional[AsyncSession] = None) -> list[ModelModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model))
|
||||
all_models = result.scalars().all()
|
||||
model_ids = [model.id for model in all_models]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
return [
|
||||
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)
|
||||
for model in all_models
|
||||
]
|
||||
|
||||
def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]:
|
||||
with get_db_context(db) as db:
|
||||
all_models = db.query(Model).filter(Model.base_model_id != None).all()
|
||||
async def get_models(self, db: Optional[AsyncSession] = None) -> list[ModelUserResponse]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model).filter(Model.base_model_id != None))
|
||||
all_models = result.scalars().all()
|
||||
|
||||
user_ids = list(set(model.user_id for model in all_models))
|
||||
model_ids = [model.id for model in all_models]
|
||||
|
||||
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users_dict = {user.id: user for user in users}
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
|
||||
models = []
|
||||
for model in all_models:
|
||||
@@ -204,252 +225,328 @@ class ModelsTable:
|
||||
models.append(
|
||||
ModelUserResponse.model_validate(
|
||||
{
|
||||
**ModelModel.model_validate(model).model_dump(),
|
||||
"user": user.model_dump() if user else None,
|
||||
**(
|
||||
await self._to_model_model(
|
||||
model,
|
||||
access_grants=grants_map.get(model.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
'user': user.model_dump() if user else None,
|
||||
}
|
||||
)
|
||||
)
|
||||
return models
|
||||
|
||||
def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]:
|
||||
with get_db_context(db) as db:
|
||||
async def get_base_models(self, db: Optional[AsyncSession] = None) -> list[ModelModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model).filter(Model.base_model_id == None))
|
||||
all_models = result.scalars().all()
|
||||
model_ids = [model.id for model in all_models]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
return [
|
||||
ModelModel.model_validate(model)
|
||||
for model in db.query(Model).filter(Model.base_model_id == None).all()
|
||||
await self._to_model_model(model, access_grants=grants_map.get(model.id, []), db=db)
|
||||
for model in all_models
|
||||
]
|
||||
|
||||
def get_models_by_user_id(
|
||||
self, user_id: str, permission: str = "write", db: Optional[Session] = None
|
||||
async def get_models_by_user_id(
|
||||
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
|
||||
) -> list[ModelUserResponse]:
|
||||
models = self.get_models(db=db)
|
||||
user_group_ids = {
|
||||
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
|
||||
}
|
||||
return [
|
||||
model
|
||||
for model in models
|
||||
if model.user_id == user_id
|
||||
or has_access(user_id, permission, model.access_control, user_group_ids)
|
||||
]
|
||||
models = await self.get_models(db=db)
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
|
||||
def _has_permission(self, db, query, filter: dict, permission: str = "read"):
|
||||
group_ids = filter.get("group_ids", [])
|
||||
user_id = filter.get("user_id")
|
||||
result = []
|
||||
for model in models:
|
||||
if model.user_id == user_id:
|
||||
result.append(model)
|
||||
elif await AccessGrants.has_access(
|
||||
user_id=user_id,
|
||||
resource_type='model',
|
||||
resource_id=model.id,
|
||||
permission=permission,
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
):
|
||||
result.append(model)
|
||||
return result
|
||||
|
||||
dialect_name = db.bind.dialect.name
|
||||
def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
|
||||
return AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=query,
|
||||
DocumentModel=Model,
|
||||
filter=filter,
|
||||
resource_type='model',
|
||||
permission=permission,
|
||||
)
|
||||
|
||||
# Public access
|
||||
conditions = []
|
||||
if group_ids or user_id:
|
||||
conditions.extend(
|
||||
[
|
||||
Model.access_control.is_(None),
|
||||
cast(Model.access_control, String) == "null",
|
||||
]
|
||||
)
|
||||
|
||||
# User-level permission
|
||||
if user_id:
|
||||
conditions.append(Model.user_id == user_id)
|
||||
|
||||
# Group-level permission
|
||||
if group_ids:
|
||||
group_conditions = []
|
||||
for gid in group_ids:
|
||||
if dialect_name == "sqlite":
|
||||
group_conditions.append(
|
||||
Model.access_control[permission]["group_ids"].contains([gid])
|
||||
)
|
||||
elif dialect_name == "postgresql":
|
||||
group_conditions.append(
|
||||
cast(
|
||||
Model.access_control[permission]["group_ids"],
|
||||
JSONB,
|
||||
).contains([gid])
|
||||
)
|
||||
conditions.append(or_(*group_conditions))
|
||||
|
||||
if conditions:
|
||||
query = query.filter(or_(*conditions))
|
||||
|
||||
return query
|
||||
|
||||
def search_models(
|
||||
async def search_models(
|
||||
self,
|
||||
user_id: str,
|
||||
filter: dict = {},
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> ModelListResponse:
|
||||
with get_db_context(db) as db:
|
||||
# Join GroupMember so we can order by group_id when requested
|
||||
query = db.query(Model, User).outerjoin(User, User.id == Model.user_id)
|
||||
query = query.filter(Model.base_model_id != None)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Model, User).outerjoin(User, User.id == Model.user_id)
|
||||
stmt = stmt.filter(Model.base_model_id != None)
|
||||
|
||||
if filter:
|
||||
query_key = filter.get("query")
|
||||
query_key = filter.get('query')
|
||||
if query_key:
|
||||
query = query.filter(
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
Model.name.ilike(f"%{query_key}%"),
|
||||
Model.base_model_id.ilike(f"%{query_key}%"),
|
||||
Model.name.ilike(f'%{query_key}%'),
|
||||
Model.base_model_id.ilike(f'%{query_key}%'),
|
||||
User.name.ilike(f'%{query_key}%'),
|
||||
User.email.ilike(f'%{query_key}%'),
|
||||
User.username.ilike(f'%{query_key}%'),
|
||||
)
|
||||
)
|
||||
|
||||
view_option = filter.get("view_option")
|
||||
if view_option == "created":
|
||||
query = query.filter(Model.user_id == user_id)
|
||||
elif view_option == "shared":
|
||||
query = query.filter(Model.user_id != user_id)
|
||||
view_option = filter.get('view_option')
|
||||
if view_option == 'created':
|
||||
stmt = stmt.filter(Model.user_id == user_id)
|
||||
elif view_option == 'shared':
|
||||
stmt = stmt.filter(Model.user_id != user_id)
|
||||
|
||||
# Apply access control filtering
|
||||
query = self._has_permission(
|
||||
stmt = self._has_permission(
|
||||
db,
|
||||
query,
|
||||
stmt,
|
||||
filter,
|
||||
permission="read",
|
||||
permission='read',
|
||||
)
|
||||
|
||||
tag = filter.get("tag")
|
||||
tag = filter.get('tag')
|
||||
if tag:
|
||||
# TODO: This is a simple implementation and should be improved for performance
|
||||
like_pattern = f'%"{tag.lower()}"%' # `"tag"` inside JSON array
|
||||
meta_text = func.lower(cast(Model.meta, String))
|
||||
|
||||
query = query.filter(meta_text.like(like_pattern))
|
||||
|
||||
order_by = filter.get("order_by")
|
||||
direction = filter.get("direction")
|
||||
|
||||
if order_by == "name":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Model.name.asc())
|
||||
# SQLite stores JSON text via json.dumps(ensure_ascii=True),
|
||||
# so non-ASCII chars are \uXXXX-escaped. PostgreSQL native JSONB
|
||||
# stores literal Unicode. Use the right pattern for each.
|
||||
if db.bind.dialect.name == 'sqlite':
|
||||
if tag.isascii():
|
||||
meta_text = func.lower(cast(Model.meta, String))
|
||||
pattern = f'%{json.dumps(tag.lower())}%'
|
||||
else:
|
||||
meta_text = cast(Model.meta, String)
|
||||
pattern = f'%{json.dumps(tag)}%'
|
||||
else:
|
||||
query = query.order_by(Model.name.desc())
|
||||
elif order_by == "created_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Model.created_at.asc())
|
||||
meta_text = func.lower(cast(Model.meta, String))
|
||||
pattern = f'%{json.dumps(tag.lower(), ensure_ascii=False)}%'
|
||||
stmt = stmt.filter(meta_text.like(pattern))
|
||||
|
||||
order_by = filter.get('order_by')
|
||||
direction = filter.get('direction')
|
||||
|
||||
if order_by == 'name':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Model.name.asc())
|
||||
else:
|
||||
query = query.order_by(Model.created_at.desc())
|
||||
elif order_by == "updated_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Model.updated_at.asc())
|
||||
stmt = stmt.order_by(Model.name.desc())
|
||||
elif order_by == 'created_at':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Model.created_at.asc())
|
||||
else:
|
||||
query = query.order_by(Model.updated_at.desc())
|
||||
stmt = stmt.order_by(Model.created_at.desc())
|
||||
elif order_by == 'updated_at':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Model.updated_at.asc())
|
||||
else:
|
||||
stmt = stmt.order_by(Model.updated_at.desc())
|
||||
|
||||
else:
|
||||
query = query.order_by(Model.created_at.desc())
|
||||
stmt = stmt.order_by(Model.created_at.desc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
total = query.count()
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
items = query.all()
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
model_ids = [model.id for model, _ in items]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
|
||||
models = []
|
||||
for model, user in items:
|
||||
models.append(
|
||||
ModelUserResponse(
|
||||
**ModelModel.model_validate(model).model_dump(),
|
||||
user=(
|
||||
UserResponse(**UserModel.model_validate(user).model_dump())
|
||||
if user
|
||||
else None
|
||||
),
|
||||
**(
|
||||
await self._to_model_model(
|
||||
model,
|
||||
access_grants=grants_map.get(model.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
return ModelListResponse(items=models, total=total)
|
||||
|
||||
def get_model_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[ModelModel]:
|
||||
async def get_model_meta_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[tuple[dict, int]]:
|
||||
"""Return (meta, updated_at) for a model, skipping access grant resolution."""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
model = db.get(Model, id)
|
||||
return ModelModel.model_validate(model)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model.meta, Model.updated_at).filter_by(id=id))
|
||||
return result.first()
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_models_by_ids(
|
||||
self, ids: list[str], db: Optional[Session] = None
|
||||
) -> list[ModelModel]:
|
||||
async def get_all_tags(
|
||||
self,
|
||||
user_id: str,
|
||||
is_admin: bool = False,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> set[str]:
|
||||
"""Extract unique tag names from model meta, querying only the meta column."""
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Model.meta).filter(Model.base_model_id != None)
|
||||
|
||||
if not is_admin:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [group.id for group in user_groups]
|
||||
|
||||
filter_dict = {'user_id': user_id}
|
||||
if user_group_ids:
|
||||
filter_dict['group_ids'] = user_group_ids
|
||||
|
||||
stmt = self._has_permission(db, stmt, filter_dict, permission='read')
|
||||
|
||||
result = await db.execute(stmt)
|
||||
rows = result.scalars().all()
|
||||
|
||||
tags_set: set[str] = set()
|
||||
for meta in rows:
|
||||
if not meta:
|
||||
continue
|
||||
for tag in meta.get('tags', []):
|
||||
try:
|
||||
name = tag.get('name') if isinstance(tag, dict) else str(tag)
|
||||
if name:
|
||||
tags_set.add(name)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
return tags_set
|
||||
|
||||
async def get_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
models = db.query(Model).filter(Model.id.in_(ids)).all()
|
||||
return [ModelModel.model_validate(model) for model in models]
|
||||
async with get_async_db_context(db) as db:
|
||||
model = await db.get(Model, id)
|
||||
return await self._to_model_model(model, db=db) if model else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def get_models_by_ids(self, ids: list[str], db: Optional[AsyncSession] = None) -> list[ModelModel]:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model).filter(Model.id.in_(ids)))
|
||||
models = result.scalars().all()
|
||||
model_ids = [model.id for model in models]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
return [
|
||||
await self._to_model_model(
|
||||
model,
|
||||
access_grants=grants_map.get(model.id, []),
|
||||
db=db,
|
||||
)
|
||||
for model in models
|
||||
]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def toggle_model_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
) -> Optional[ModelModel]:
|
||||
with get_db_context(db) as db:
|
||||
async def toggle_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
try:
|
||||
is_active = db.query(Model).filter_by(id=id).first().is_active
|
||||
result = await db.execute(select(Model).filter_by(id=id))
|
||||
model = result.scalars().first()
|
||||
if not model:
|
||||
return None
|
||||
|
||||
db.query(Model).filter_by(id=id).update(
|
||||
{
|
||||
"is_active": not is_active,
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
)
|
||||
db.commit()
|
||||
model.is_active = not model.is_active
|
||||
model.updated_at = int(time.time())
|
||||
await db.commit()
|
||||
await db.refresh(model)
|
||||
|
||||
return self.get_model_by_id(id, db=db)
|
||||
return await self._to_model_model(model, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def update_model_by_id(
|
||||
self, id: str, model: ModelForm, db: Optional[Session] = None
|
||||
async def update_model_by_id(
|
||||
self, id: str, model: ModelForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[ModelModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# update only the fields that are present in the model
|
||||
data = model.model_dump(exclude={"id"})
|
||||
result = db.query(Model).filter_by(id=id).update(data)
|
||||
data = model.model_dump(exclude={'id', 'access_grants'})
|
||||
data['updated_at'] = int(time.time())
|
||||
await db.execute(update(Model).filter_by(id=id).values(**data))
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
if model.access_grants is not None:
|
||||
await AccessGrants.set_access_grants('model', id, model.access_grants, db=db)
|
||||
|
||||
model = db.get(Model, id)
|
||||
db.refresh(model)
|
||||
return ModelModel.model_validate(model)
|
||||
return await self.get_model_by_id(id, db=db)
|
||||
except Exception as e:
|
||||
log.exception(f"Failed to update the model by id {id}: {e}")
|
||||
log.exception(f'Failed to update the model by id {id}: {e}')
|
||||
return None
|
||||
|
||||
def delete_model_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
async def update_model_updated_at_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[ModelModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Model).filter_by(id=id).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model).filter_by(id=id))
|
||||
model_obj = result.scalars().first()
|
||||
if not model_obj:
|
||||
return None
|
||||
model_obj.updated_at = int(time.time())
|
||||
await db.commit()
|
||||
await db.refresh(model_obj)
|
||||
return await self._to_model_model(model_obj, db=db)
|
||||
except Exception as e:
|
||||
log.exception(f'Failed to update the model updated_at by id {id}: {e}')
|
||||
return None
|
||||
|
||||
async def delete_model_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
await AccessGrants.revoke_all_access('model', id, db=db)
|
||||
await db.execute(delete(Model).filter_by(id=id))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def delete_all_models(self, db: Optional[Session] = None) -> bool:
|
||||
async def delete_all_models(self, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Model).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Model.id))
|
||||
model_ids = [row[0] for row in result.all()]
|
||||
for model_id in model_ids:
|
||||
await AccessGrants.revoke_all_access('model', model_id, db=db)
|
||||
await db.execute(delete(Model))
|
||||
await db.commit()
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def sync_models(
|
||||
self, user_id: str, models: list[ModelModel], db: Optional[Session] = None
|
||||
async def sync_models(
|
||||
self, user_id: str, models: list[ModelModel], db: Optional[AsyncSession] = None
|
||||
) -> list[ModelModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Get existing models
|
||||
existing_models = db.query(Model).all()
|
||||
result = await db.execute(select(Model))
|
||||
existing_models = result.scalars().all()
|
||||
existing_ids = {model.id for model in existing_models}
|
||||
|
||||
# Prepare a set of new model IDs
|
||||
@@ -458,35 +555,48 @@ class ModelsTable:
|
||||
# Update or insert models
|
||||
for model in models:
|
||||
if model.id in existing_ids:
|
||||
db.query(Model).filter_by(id=model.id).update(
|
||||
{
|
||||
**model.model_dump(),
|
||||
"user_id": user_id,
|
||||
"updated_at": int(time.time()),
|
||||
}
|
||||
await db.execute(
|
||||
update(Model)
|
||||
.filter_by(id=model.id)
|
||||
.values(
|
||||
**model.model_dump(exclude={'access_grants'}),
|
||||
user_id=user_id,
|
||||
updated_at=int(time.time()),
|
||||
)
|
||||
)
|
||||
else:
|
||||
new_model = Model(
|
||||
**{
|
||||
**model.model_dump(),
|
||||
"user_id": user_id,
|
||||
"updated_at": int(time.time()),
|
||||
**model.model_dump(exclude={'access_grants'}),
|
||||
'user_id': user_id,
|
||||
'updated_at': int(time.time()),
|
||||
}
|
||||
)
|
||||
db.add(new_model)
|
||||
await AccessGrants.set_access_grants('model', model.id, model.access_grants, db=db)
|
||||
|
||||
# Remove models that are no longer present
|
||||
for model in existing_models:
|
||||
if model.id not in new_model_ids:
|
||||
db.delete(model)
|
||||
await AccessGrants.revoke_all_access('model', model.id, db=db)
|
||||
await db.delete(model)
|
||||
|
||||
db.commit()
|
||||
await db.commit()
|
||||
|
||||
result = await db.execute(select(Model))
|
||||
all_models = result.scalars().all()
|
||||
model_ids = [model.id for model in all_models]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('model', model_ids, db=db)
|
||||
return [
|
||||
ModelModel.model_validate(model) for model in db.query(Model).all()
|
||||
await self._to_model_model(
|
||||
model,
|
||||
access_grants=grants_map.get(model.id, []),
|
||||
db=db,
|
||||
)
|
||||
for model in all_models
|
||||
]
|
||||
except Exception as e:
|
||||
log.exception(f"Error syncing models for user {user_id}: {e}")
|
||||
log.exception(f'Error syncing models for user {user_id}: {e}')
|
||||
return []
|
||||
|
||||
|
||||
|
||||
+194
-245
@@ -4,20 +4,16 @@ import uuid
|
||||
from typing import Optional
|
||||
from functools import lru_cache
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, get_db, get_db_context
|
||||
from sqlalchemy import Boolean, select, delete, update, or_, func, cast
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.utils.access_control import has_access
|
||||
from open_webui.models.users import User, UserModel, Users, UserResponse
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
|
||||
from sqlalchemy import or_, func, select, and_, text, cast, or_, and_, func
|
||||
from sqlalchemy.sql import exists
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON
|
||||
|
||||
####################
|
||||
# Note DB Schema
|
||||
@@ -25,7 +21,7 @@ from sqlalchemy.sql import exists
|
||||
|
||||
|
||||
class Note(Base):
|
||||
__tablename__ = "note"
|
||||
__tablename__ = 'note'
|
||||
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
user_id = Column(Text)
|
||||
@@ -33,8 +29,7 @@ class Note(Base):
|
||||
title = Column(Text)
|
||||
data = Column(JSON, nullable=True)
|
||||
meta = Column(JSON, nullable=True)
|
||||
|
||||
access_control = Column(JSON, nullable=True)
|
||||
is_pinned = Column(Boolean, default=False, nullable=True)
|
||||
|
||||
created_at = Column(BigInteger)
|
||||
updated_at = Column(BigInteger)
|
||||
@@ -49,8 +44,9 @@ class NoteModel(BaseModel):
|
||||
title: str
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
is_pinned: Optional[bool] = False
|
||||
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
||||
|
||||
created_at: int # timestamp in epoch
|
||||
updated_at: int # timestamp in epoch
|
||||
@@ -65,14 +61,14 @@ class NoteForm(BaseModel):
|
||||
title: str
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class NoteUpdateForm(BaseModel):
|
||||
title: Optional[str] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
access_control: Optional[dict] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
|
||||
|
||||
class NoteUserResponse(NoteModel):
|
||||
@@ -83,6 +79,7 @@ class NoteItemResponse(BaseModel):
|
||||
id: str
|
||||
title: str
|
||||
data: Optional[dict]
|
||||
is_pinned: Optional[bool] = False
|
||||
updated_at: int
|
||||
created_at: int
|
||||
user: Optional[UserResponse] = None
|
||||
@@ -94,316 +91,268 @@ class NoteListResponse(BaseModel):
|
||||
|
||||
|
||||
class NoteTable:
|
||||
def _has_permission(self, db, query, filter: dict, permission: str = "read"):
|
||||
group_ids = filter.get("group_ids", [])
|
||||
user_id = filter.get("user_id")
|
||||
dialect_name = db.bind.dialect.name
|
||||
async def _get_access_grants(self, note_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
|
||||
return await AccessGrants.get_grants_by_resource('note', note_id, db=db)
|
||||
|
||||
conditions = []
|
||||
async def _to_note_model(
|
||||
self,
|
||||
note: Note,
|
||||
access_grants: Optional[list[AccessGrantModel]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> NoteModel:
|
||||
note_data = NoteModel.model_validate(note).model_dump(exclude={'access_grants'})
|
||||
note_data['access_grants'] = (
|
||||
access_grants if access_grants is not None else await self._get_access_grants(note_data['id'], db=db)
|
||||
)
|
||||
return NoteModel.model_validate(note_data)
|
||||
|
||||
# Handle read_only permission separately
|
||||
if permission == "read_only":
|
||||
# For read_only, we want items where:
|
||||
# 1. User has explicit read permission (via groups or user-level)
|
||||
# 2. BUT does NOT have write permission
|
||||
# 3. Public items are NOT considered read_only
|
||||
def _has_permission(self, db, query, filter: dict, permission: str = 'read'):
|
||||
return AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=query,
|
||||
DocumentModel=Note,
|
||||
filter=filter,
|
||||
resource_type='note',
|
||||
permission=permission,
|
||||
)
|
||||
|
||||
read_conditions = []
|
||||
|
||||
# Group-level read permission
|
||||
if group_ids:
|
||||
group_read_conditions = []
|
||||
for gid in group_ids:
|
||||
if dialect_name == "sqlite":
|
||||
group_read_conditions.append(
|
||||
Note.access_control["read"]["group_ids"].contains([gid])
|
||||
)
|
||||
elif dialect_name == "postgresql":
|
||||
group_read_conditions.append(
|
||||
cast(
|
||||
Note.access_control["read"]["group_ids"],
|
||||
JSONB,
|
||||
).contains([gid])
|
||||
)
|
||||
|
||||
if group_read_conditions:
|
||||
read_conditions.append(or_(*group_read_conditions))
|
||||
|
||||
# Combine read conditions
|
||||
if read_conditions:
|
||||
has_read = or_(*read_conditions)
|
||||
else:
|
||||
# If no read conditions, return empty result
|
||||
return query.filter(False)
|
||||
|
||||
# Now exclude items where user has write permission
|
||||
write_exclusions = []
|
||||
|
||||
# Exclude items owned by user (they have implicit write)
|
||||
if user_id:
|
||||
write_exclusions.append(Note.user_id != user_id)
|
||||
|
||||
# Exclude items where user has explicit write permission via groups
|
||||
if group_ids:
|
||||
group_write_conditions = []
|
||||
for gid in group_ids:
|
||||
if dialect_name == "sqlite":
|
||||
group_write_conditions.append(
|
||||
Note.access_control["write"]["group_ids"].contains([gid])
|
||||
)
|
||||
elif dialect_name == "postgresql":
|
||||
group_write_conditions.append(
|
||||
cast(
|
||||
Note.access_control["write"]["group_ids"],
|
||||
JSONB,
|
||||
).contains([gid])
|
||||
)
|
||||
|
||||
if group_write_conditions:
|
||||
# User should NOT have write permission
|
||||
write_exclusions.append(~or_(*group_write_conditions))
|
||||
|
||||
# Exclude public items (items without access_control)
|
||||
write_exclusions.append(Note.access_control.isnot(None))
|
||||
write_exclusions.append(cast(Note.access_control, String) != "null")
|
||||
|
||||
# Combine: has read AND does not have write AND not public
|
||||
if write_exclusions:
|
||||
query = query.filter(and_(has_read, *write_exclusions))
|
||||
else:
|
||||
query = query.filter(has_read)
|
||||
|
||||
return query
|
||||
|
||||
# Original logic for other permissions (read, write, etc.)
|
||||
# Public access conditions
|
||||
if group_ids or user_id:
|
||||
conditions.extend(
|
||||
[
|
||||
Note.access_control.is_(None),
|
||||
cast(Note.access_control, String) == "null",
|
||||
]
|
||||
)
|
||||
|
||||
# User-level permission (owner has all permissions)
|
||||
if user_id:
|
||||
conditions.append(Note.user_id == user_id)
|
||||
|
||||
# Group-level permission
|
||||
if group_ids:
|
||||
group_conditions = []
|
||||
for gid in group_ids:
|
||||
if dialect_name == "sqlite":
|
||||
group_conditions.append(
|
||||
Note.access_control[permission]["group_ids"].contains([gid])
|
||||
)
|
||||
elif dialect_name == "postgresql":
|
||||
group_conditions.append(
|
||||
cast(
|
||||
Note.access_control[permission]["group_ids"],
|
||||
JSONB,
|
||||
).contains([gid])
|
||||
)
|
||||
conditions.append(or_(*group_conditions))
|
||||
|
||||
if conditions:
|
||||
query = query.filter(or_(*conditions))
|
||||
|
||||
return query
|
||||
|
||||
def insert_new_note(
|
||||
self, user_id: str, form_data: NoteForm, db: Optional[Session] = None
|
||||
async def insert_new_note(
|
||||
self, user_id: str, form_data: NoteForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
note = NoteModel(
|
||||
**{
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_id": user_id,
|
||||
**form_data.model_dump(),
|
||||
"created_at": int(time.time_ns()),
|
||||
"updated_at": int(time.time_ns()),
|
||||
'id': str(uuid.uuid4()),
|
||||
'user_id': user_id,
|
||||
**form_data.model_dump(exclude={'access_grants'}),
|
||||
'created_at': int(time.time_ns()),
|
||||
'updated_at': int(time.time_ns()),
|
||||
'access_grants': [],
|
||||
}
|
||||
)
|
||||
|
||||
new_note = Note(**note.model_dump())
|
||||
new_note = Note(**note.model_dump(exclude={'access_grants'}))
|
||||
|
||||
db.add(new_note)
|
||||
db.commit()
|
||||
return note
|
||||
await db.commit()
|
||||
await AccessGrants.set_access_grants('note', note.id, form_data.access_grants, db=db)
|
||||
return await self._to_note_model(new_note, db=db)
|
||||
|
||||
def get_notes(
|
||||
self, skip: int = 0, limit: int = 50, db: Optional[Session] = None
|
||||
) -> list[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Note).order_by(Note.updated_at.desc())
|
||||
async def get_notes(self, skip: int = 0, limit: int = 50, db: Optional[AsyncSession] = None) -> list[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Note).order_by(Note.updated_at.desc())
|
||||
if skip is not None:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit is not None:
|
||||
query = query.limit(limit)
|
||||
notes = query.all()
|
||||
return [NoteModel.model_validate(note) for note in notes]
|
||||
stmt = stmt.limit(limit)
|
||||
result = await db.execute(stmt)
|
||||
notes = result.scalars().all()
|
||||
note_ids = [note.id for note in notes]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
|
||||
def search_notes(
|
||||
async def search_notes(
|
||||
self,
|
||||
user_id: str,
|
||||
filter: dict = {},
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> NoteListResponse:
|
||||
with get_db_context(db) as db:
|
||||
query = db.query(Note, User).outerjoin(User, User.id == Note.user_id)
|
||||
async with get_async_db_context(db) as db:
|
||||
stmt = select(Note, User).outerjoin(User, User.id == Note.user_id)
|
||||
if filter:
|
||||
query_key = filter.get("query")
|
||||
query_key = filter.get('query')
|
||||
if query_key:
|
||||
# Normalize search by removing hyphens and spaces (e.g., "todo" matches "to-do" and "to do")
|
||||
normalized_query = query_key.replace("-", "").replace(" ", "")
|
||||
query = query.filter(
|
||||
or_(
|
||||
func.replace(
|
||||
func.replace(Note.title, "-", ""), " ", ""
|
||||
).ilike(f"%{normalized_query}%"),
|
||||
func.replace(
|
||||
# Split query into individual words and normalize each
|
||||
# (strip hyphens so "todo" matches "to-do").
|
||||
# All words must match somewhere in title OR content (AND semantics).
|
||||
search_words = query_key.split()
|
||||
normalized_words = [w.replace('-', '') for w in search_words if w.replace('-', '')]
|
||||
for word in normalized_words:
|
||||
stmt = stmt.filter(
|
||||
or_(
|
||||
func.replace(func.replace(Note.title, '-', ''), ' ', '').ilike(f'%{word}%'),
|
||||
func.replace(
|
||||
cast(Note.data["content"]["md"], Text), "-", ""
|
||||
),
|
||||
" ",
|
||||
"",
|
||||
).ilike(f"%{normalized_query}%"),
|
||||
func.replace(cast(Note.data['content']['md'], Text), '-', ''),
|
||||
' ',
|
||||
'',
|
||||
).ilike(f'%{word}%'),
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
view_option = filter.get("view_option")
|
||||
if view_option == "created":
|
||||
query = query.filter(Note.user_id == user_id)
|
||||
elif view_option == "shared":
|
||||
query = query.filter(Note.user_id != user_id)
|
||||
view_option = filter.get('view_option')
|
||||
if view_option == 'created':
|
||||
stmt = stmt.filter(Note.user_id == user_id)
|
||||
elif view_option == 'shared':
|
||||
stmt = stmt.filter(Note.user_id != user_id)
|
||||
|
||||
# Apply access control filtering
|
||||
if "permission" in filter:
|
||||
permission = filter["permission"]
|
||||
if 'permission' in filter:
|
||||
permission = filter['permission']
|
||||
else:
|
||||
permission = "write"
|
||||
permission = 'write'
|
||||
|
||||
query = self._has_permission(
|
||||
stmt = self._has_permission(
|
||||
db,
|
||||
query,
|
||||
stmt,
|
||||
filter,
|
||||
permission=permission,
|
||||
)
|
||||
|
||||
order_by = filter.get("order_by")
|
||||
direction = filter.get("direction")
|
||||
order_by = filter.get('order_by')
|
||||
direction = filter.get('direction')
|
||||
|
||||
if order_by == "name":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Note.title.asc())
|
||||
if order_by == 'name':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Note.title.asc())
|
||||
else:
|
||||
query = query.order_by(Note.title.desc())
|
||||
elif order_by == "created_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Note.created_at.asc())
|
||||
stmt = stmt.order_by(Note.title.desc())
|
||||
elif order_by == 'created_at':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Note.created_at.asc())
|
||||
else:
|
||||
query = query.order_by(Note.created_at.desc())
|
||||
elif order_by == "updated_at":
|
||||
if direction == "asc":
|
||||
query = query.order_by(Note.updated_at.asc())
|
||||
stmt = stmt.order_by(Note.created_at.desc())
|
||||
elif order_by == 'updated_at':
|
||||
if direction == 'asc':
|
||||
stmt = stmt.order_by(Note.updated_at.asc())
|
||||
else:
|
||||
query = query.order_by(Note.updated_at.desc())
|
||||
stmt = stmt.order_by(Note.updated_at.desc())
|
||||
else:
|
||||
query = query.order_by(Note.updated_at.desc())
|
||||
stmt = stmt.order_by(Note.updated_at.desc())
|
||||
|
||||
else:
|
||||
query = query.order_by(Note.updated_at.desc())
|
||||
stmt = stmt.order_by(Note.updated_at.desc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
total = query.count()
|
||||
count_result = await db.execute(select(func.count()).select_from(stmt.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
items = query.all()
|
||||
result = await db.execute(stmt)
|
||||
items = result.all()
|
||||
|
||||
note_ids = [note.id for note, _ in items]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
|
||||
notes = []
|
||||
for note, user in items:
|
||||
notes.append(
|
||||
NoteUserResponse(
|
||||
**NoteModel.model_validate(note).model_dump(),
|
||||
user=(
|
||||
UserResponse(**UserModel.model_validate(user).model_dump())
|
||||
if user
|
||||
else None
|
||||
),
|
||||
**(
|
||||
await self._to_note_model(
|
||||
note,
|
||||
access_grants=grants_map.get(note.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
return NoteListResponse(items=notes, total=total)
|
||||
|
||||
def get_notes_by_user_id(
|
||||
async def get_notes_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
permission: str = "read",
|
||||
permission: str = 'read',
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
user_group_ids = [
|
||||
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
|
||||
]
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [group.id for group in user_groups]
|
||||
|
||||
query = db.query(Note).order_by(Note.updated_at.desc())
|
||||
query = self._has_permission(
|
||||
db, query, {"user_id": user_id, "group_ids": user_group_ids}, permission
|
||||
)
|
||||
stmt = select(Note).order_by(Note.updated_at.desc())
|
||||
stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
|
||||
|
||||
if skip is not None:
|
||||
query = query.offset(skip)
|
||||
stmt = stmt.offset(skip)
|
||||
if limit is not None:
|
||||
query = query.limit(limit)
|
||||
stmt = stmt.limit(limit)
|
||||
|
||||
notes = query.all()
|
||||
return [NoteModel.model_validate(note) for note in notes]
|
||||
result = await db.execute(stmt)
|
||||
notes = result.scalars().all()
|
||||
note_ids = [note.id for note in notes]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
|
||||
def get_note_by_id(
|
||||
self, id: str, db: Optional[Session] = None
|
||||
async def get_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Note).filter(Note.id == id))
|
||||
note = result.scalars().first()
|
||||
return await self._to_note_model(note, db=db) if note else None
|
||||
|
||||
async def update_note_by_id(
|
||||
self, id: str, form_data: NoteUpdateForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
note = db.query(Note).filter(Note.id == id).first()
|
||||
return NoteModel.model_validate(note) if note else None
|
||||
|
||||
def update_note_by_id(
|
||||
self, id: str, form_data: NoteUpdateForm, db: Optional[Session] = None
|
||||
) -> Optional[NoteModel]:
|
||||
with get_db_context(db) as db:
|
||||
note = db.query(Note).filter(Note.id == id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Note).filter(Note.id == id))
|
||||
note = result.scalars().first()
|
||||
if not note:
|
||||
return None
|
||||
|
||||
form_data = form_data.model_dump(exclude_unset=True)
|
||||
|
||||
if "title" in form_data:
|
||||
note.title = form_data["title"]
|
||||
if "data" in form_data:
|
||||
note.data = {**note.data, **form_data["data"]}
|
||||
if "meta" in form_data:
|
||||
note.meta = {**note.meta, **form_data["meta"]}
|
||||
if 'title' in form_data:
|
||||
note.title = form_data['title']
|
||||
if 'data' in form_data:
|
||||
note.data = {**note.data, **form_data['data']}
|
||||
if 'meta' in form_data:
|
||||
note.meta = {**note.meta, **form_data['meta']}
|
||||
|
||||
if "access_control" in form_data:
|
||||
note.access_control = form_data["access_control"]
|
||||
if 'access_grants' in form_data:
|
||||
await AccessGrants.set_access_grants('note', id, form_data['access_grants'], db=db)
|
||||
|
||||
note.updated_at = int(time.time_ns())
|
||||
|
||||
db.commit()
|
||||
return NoteModel.model_validate(note) if note else None
|
||||
await db.commit()
|
||||
return await self._to_note_model(note, db=db) if note else None
|
||||
|
||||
def delete_note_by_id(self, id: str, db: Optional[Session] = None) -> bool:
|
||||
async def toggle_note_pinned_by_id(self, id: str, db: Optional[AsyncSession] = None) -> Optional[NoteModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Note).filter(Note.id == id).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Note).filter(Note.id == id))
|
||||
note = result.scalars().first()
|
||||
if not note:
|
||||
return None
|
||||
note.is_pinned = not note.is_pinned
|
||||
note.updated_at = int(time.time_ns())
|
||||
await db.commit()
|
||||
return await self._to_note_model(note, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def get_pinned_notes_by_user_id(
|
||||
self,
|
||||
user_id: str,
|
||||
permission: str = 'read',
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[NoteModel]:
|
||||
async with get_async_db_context(db) as db:
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = [group.id for group in user_groups]
|
||||
|
||||
stmt = select(Note).filter(Note.is_pinned == True).order_by(Note.updated_at.desc())
|
||||
stmt = self._has_permission(db, stmt, {'user_id': user_id, 'group_ids': user_group_ids}, permission)
|
||||
|
||||
result = await db.execute(stmt)
|
||||
notes = result.scalars().all()
|
||||
note_ids = [note.id for note in notes]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('note', note_ids, db=db)
|
||||
return [await self._to_note_model(note, access_grants=grants_map.get(note.id, []), db=db) for note in notes]
|
||||
|
||||
async def delete_note_by_id(self, id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
await AccessGrants.revoke_all_access('note', id, db=db)
|
||||
await db.execute(delete(Note).filter(Note.id == id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@@ -8,8 +8,9 @@ import json
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.env import OAUTH_SESSION_TOKEN_ENCRYPTION_KEY
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
@@ -23,23 +24,21 @@ log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OAuthSession(Base):
|
||||
__tablename__ = "oauth_session"
|
||||
__tablename__ = 'oauth_session'
|
||||
|
||||
id = Column(Text, primary_key=True, unique=True)
|
||||
user_id = Column(Text, nullable=False)
|
||||
provider = Column(Text, nullable=False)
|
||||
token = Column(
|
||||
Text, nullable=False
|
||||
) # JSON with access_token, id_token, refresh_token
|
||||
token = Column(Text, nullable=False) # JSON with access_token, id_token, refresh_token
|
||||
expires_at = Column(BigInteger, nullable=False)
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
updated_at = Column(BigInteger, nullable=False)
|
||||
|
||||
# Add indexes for better performance
|
||||
__table_args__ = (
|
||||
Index("idx_oauth_session_user_id", "user_id"),
|
||||
Index("idx_oauth_session_expires_at", "expires_at"),
|
||||
Index("idx_oauth_session_user_provider", "user_id", "provider"),
|
||||
Index('idx_oauth_session_user_id', 'user_id'),
|
||||
Index('idx_oauth_session_expires_at', 'expires_at'),
|
||||
Index('idx_oauth_session_user_provider', 'user_id', 'provider'),
|
||||
)
|
||||
|
||||
|
||||
@@ -71,7 +70,7 @@ class OAuthSessionTable:
|
||||
def __init__(self):
|
||||
self.encryption_key = OAUTH_SESSION_TOKEN_ENCRYPTION_KEY
|
||||
if not self.encryption_key:
|
||||
raise Exception("OAUTH_SESSION_TOKEN_ENCRYPTION_KEY is not set")
|
||||
raise Exception('OAUTH_SESSION_TOKEN_ENCRYPTION_KEY is not set')
|
||||
|
||||
# check if encryption key is in the right format for Fernet (32 url-safe base64-encoded bytes)
|
||||
if len(self.encryption_key) != 44:
|
||||
@@ -83,7 +82,7 @@ class OAuthSessionTable:
|
||||
try:
|
||||
self.fernet = Fernet(self.encryption_key)
|
||||
except Exception as e:
|
||||
log.error(f"Error initializing Fernet with provided key: {e}")
|
||||
log.error(f'Error initializing Fernet with provided key: {e}')
|
||||
raise
|
||||
|
||||
def _encrypt_token(self, token) -> str:
|
||||
@@ -93,7 +92,7 @@ class OAuthSessionTable:
|
||||
encrypted = self.fernet.encrypt(token_json.encode()).decode()
|
||||
return encrypted
|
||||
except Exception as e:
|
||||
log.error(f"Error encrypting tokens: {e}")
|
||||
log.error(f'Error encrypting tokens: {e}')
|
||||
raise
|
||||
|
||||
def _decrypt_token(self, token: str):
|
||||
@@ -102,186 +101,247 @@ class OAuthSessionTable:
|
||||
decrypted = self.fernet.decrypt(token.encode()).decode()
|
||||
return json.loads(decrypted)
|
||||
except Exception as e:
|
||||
log.error(f"Error decrypting tokens: {e}")
|
||||
log.error(f'Error decrypting tokens: {type(e).__name__}: {e}')
|
||||
raise
|
||||
|
||||
def create_session(
|
||||
async def create_session(
|
||||
self,
|
||||
user_id: str,
|
||||
provider: str,
|
||||
token: dict,
|
||||
db: Optional[Session] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[OAuthSessionModel]:
|
||||
"""Create a new OAuth session"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
current_time = int(time.time())
|
||||
id = str(uuid.uuid4())
|
||||
|
||||
result = OAuthSession(
|
||||
**{
|
||||
"id": id,
|
||||
"user_id": user_id,
|
||||
"provider": provider,
|
||||
"token": self._encrypt_token(token),
|
||||
"expires_at": token.get("expires_at"),
|
||||
"created_at": current_time,
|
||||
"updated_at": current_time,
|
||||
'id': id,
|
||||
'user_id': user_id,
|
||||
'provider': provider,
|
||||
'token': self._encrypt_token(token),
|
||||
'expires_at': token.get('expires_at') or int(time.time() + 3600),
|
||||
'created_at': current_time,
|
||||
'updated_at': current_time,
|
||||
}
|
||||
)
|
||||
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
|
||||
if result:
|
||||
result.token = token # Return decrypted token
|
||||
return OAuthSessionModel.model_validate(result)
|
||||
# Make a copy of the model data before closing session
|
||||
model = OAuthSessionModel(
|
||||
id=result.id,
|
||||
user_id=result.user_id,
|
||||
provider=result.provider,
|
||||
token=token, # Return decrypted token
|
||||
expires_at=result.expires_at,
|
||||
created_at=result.created_at,
|
||||
updated_at=result.updated_at,
|
||||
)
|
||||
return model
|
||||
else:
|
||||
return None
|
||||
except Exception as e:
|
||||
log.error(f"Error creating OAuth session: {e}")
|
||||
log.error(f'Error creating OAuth session: {e}')
|
||||
return None
|
||||
|
||||
def get_session_by_id(
|
||||
self, session_id: str, db: Optional[Session] = None
|
||||
async def get_session_by_id(
|
||||
self, session_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[OAuthSessionModel]:
|
||||
"""Get OAuth session by ID"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
session = db.query(OAuthSession).filter_by(id=session_id).first()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(OAuthSession).filter_by(id=session_id))
|
||||
session = result.scalars().first()
|
||||
if session:
|
||||
session.token = self._decrypt_token(session.token)
|
||||
return OAuthSessionModel.model_validate(session)
|
||||
return OAuthSessionModel(
|
||||
id=session.id,
|
||||
user_id=session.user_id,
|
||||
provider=session.provider,
|
||||
token=self._decrypt_token(session.token),
|
||||
expires_at=session.expires_at,
|
||||
created_at=session.created_at,
|
||||
updated_at=session.updated_at,
|
||||
)
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
log.error(f"Error getting OAuth session by ID: {e}")
|
||||
log.error(f'Error getting OAuth session by ID: {e}')
|
||||
return None
|
||||
|
||||
def get_session_by_id_and_user_id(
|
||||
self, session_id: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_session_by_id_and_user_id(
|
||||
self, session_id: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[OAuthSessionModel]:
|
||||
"""Get OAuth session by ID and user ID"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
session = (
|
||||
db.query(OAuthSession)
|
||||
.filter_by(id=session_id, user_id=user_id)
|
||||
.first()
|
||||
)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(OAuthSession).filter_by(id=session_id, user_id=user_id))
|
||||
session = result.scalars().first()
|
||||
if session:
|
||||
session.token = self._decrypt_token(session.token)
|
||||
return OAuthSessionModel.model_validate(session)
|
||||
return OAuthSessionModel(
|
||||
id=session.id,
|
||||
user_id=session.user_id,
|
||||
provider=session.provider,
|
||||
token=self._decrypt_token(session.token),
|
||||
expires_at=session.expires_at,
|
||||
created_at=session.created_at,
|
||||
updated_at=session.updated_at,
|
||||
)
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
log.error(f"Error getting OAuth session by ID: {e}")
|
||||
log.error(f'Error getting OAuth session by ID: {e}')
|
||||
return None
|
||||
|
||||
def get_session_by_provider_and_user_id(
|
||||
self, provider: str, user_id: str, db: Optional[Session] = None
|
||||
async def get_session_by_provider_and_user_id(
|
||||
self, provider: str, user_id: str, db: Optional[AsyncSession] = None
|
||||
) -> Optional[OAuthSessionModel]:
|
||||
"""Get OAuth session by provider and user ID"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
session = (
|
||||
db.query(OAuthSession)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(OAuthSession)
|
||||
.filter_by(provider=provider, user_id=user_id)
|
||||
.first()
|
||||
.order_by(OAuthSession.created_at.desc())
|
||||
)
|
||||
session = result.scalars().first()
|
||||
if session:
|
||||
session.token = self._decrypt_token(session.token)
|
||||
return OAuthSessionModel.model_validate(session)
|
||||
return OAuthSessionModel(
|
||||
id=session.id,
|
||||
user_id=session.user_id,
|
||||
provider=session.provider,
|
||||
token=self._decrypt_token(session.token),
|
||||
expires_at=session.expires_at,
|
||||
created_at=session.created_at,
|
||||
updated_at=session.updated_at,
|
||||
)
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
log.error(f"Error getting OAuth session by provider and user ID: {e}")
|
||||
log.error(f'Error getting OAuth session by provider and user ID: {e}')
|
||||
return None
|
||||
|
||||
def get_sessions_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> List[OAuthSessionModel]:
|
||||
async def get_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> List[OAuthSessionModel]:
|
||||
"""Get all OAuth sessions for a user"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
sessions = db.query(OAuthSession).filter_by(user_id=user_id).all()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(OAuthSession).filter_by(user_id=user_id))
|
||||
sessions = result.scalars().all()
|
||||
|
||||
results = []
|
||||
for session in sessions:
|
||||
session.token = self._decrypt_token(session.token)
|
||||
results.append(OAuthSessionModel.model_validate(session))
|
||||
try:
|
||||
results.append(
|
||||
OAuthSessionModel(
|
||||
id=session.id,
|
||||
user_id=session.user_id,
|
||||
provider=session.provider,
|
||||
token=self._decrypt_token(session.token),
|
||||
expires_at=session.expires_at,
|
||||
created_at=session.created_at,
|
||||
updated_at=session.updated_at,
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
log.warning(
|
||||
f'Skipping OAuth session {session.id} due to decryption failure, deleting corrupted session: {type(e).__name__}: {e}'
|
||||
)
|
||||
await db.execute(delete(OAuthSession).filter_by(id=session.id))
|
||||
await db.commit()
|
||||
|
||||
return results
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Error getting OAuth sessions by user ID: {e}")
|
||||
log.error(f'Error getting OAuth sessions by user ID: {e}')
|
||||
return []
|
||||
|
||||
def update_session_by_id(
|
||||
self, session_id: str, token: dict, db: Optional[Session] = None
|
||||
async def update_session_by_id(
|
||||
self, session_id: str, token: dict, db: Optional[AsyncSession] = None
|
||||
) -> Optional[OAuthSessionModel]:
|
||||
"""Update OAuth session tokens"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
async with get_async_db_context(db) as db:
|
||||
current_time = int(time.time())
|
||||
|
||||
db.query(OAuthSession).filter_by(id=session_id).update(
|
||||
{
|
||||
"token": self._encrypt_token(token),
|
||||
"expires_at": token.get("expires_at"),
|
||||
"updated_at": current_time,
|
||||
}
|
||||
await db.execute(
|
||||
update(OAuthSession)
|
||||
.filter_by(id=session_id)
|
||||
.values(
|
||||
token=self._encrypt_token(token),
|
||||
expires_at=token.get('expires_at') or int(time.time() + 3600),
|
||||
updated_at=current_time,
|
||||
)
|
||||
)
|
||||
db.commit()
|
||||
session = db.query(OAuthSession).filter_by(id=session_id).first()
|
||||
await db.commit()
|
||||
result = await db.execute(select(OAuthSession).filter_by(id=session_id))
|
||||
session = result.scalars().first()
|
||||
|
||||
if session:
|
||||
session.token = self._decrypt_token(session.token)
|
||||
return OAuthSessionModel.model_validate(session)
|
||||
return OAuthSessionModel(
|
||||
id=session.id,
|
||||
user_id=session.user_id,
|
||||
provider=session.provider,
|
||||
token=self._decrypt_token(session.token),
|
||||
expires_at=session.expires_at,
|
||||
created_at=session.created_at,
|
||||
updated_at=session.updated_at,
|
||||
)
|
||||
|
||||
return None
|
||||
except Exception as e:
|
||||
log.error(f"Error updating OAuth session tokens: {e}")
|
||||
log.error(f'Error updating OAuth session tokens: {e}')
|
||||
return None
|
||||
|
||||
def delete_session_by_id(
|
||||
self, session_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
async def delete_session_by_id(self, session_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Delete an OAuth session"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
result = db.query(OAuthSession).filter_by(id=session_id).delete()
|
||||
db.commit()
|
||||
return result > 0
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(delete(OAuthSession).filter_by(id=session_id))
|
||||
await db.commit()
|
||||
return result.rowcount > 0
|
||||
except Exception as e:
|
||||
log.error(f"Error deleting OAuth session: {e}")
|
||||
log.error(f'Error deleting OAuth session: {e}')
|
||||
return False
|
||||
|
||||
def delete_sessions_by_user_id(
|
||||
self, user_id: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
async def delete_sessions_by_user_id(self, user_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Delete all OAuth sessions for a user"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
result = db.query(OAuthSession).filter_by(user_id=user_id).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(OAuthSession).filter_by(user_id=user_id))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception as e:
|
||||
log.error(f"Error deleting OAuth sessions by user ID: {e}")
|
||||
log.error(f'Error deleting OAuth sessions by user ID: {e}')
|
||||
return False
|
||||
|
||||
def delete_sessions_by_provider(
|
||||
self, provider: str, db: Optional[Session] = None
|
||||
async def delete_sessions_by_user_id_and_provider(
|
||||
self, user_id: str, provider: str, db: Optional[AsyncSession] = None
|
||||
) -> bool:
|
||||
"""Delete all OAuth sessions for a specific user and provider"""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(delete(OAuthSession).filter_by(user_id=user_id, provider=provider))
|
||||
await db.commit()
|
||||
return result.rowcount > 0
|
||||
except Exception as e:
|
||||
log.error(f'Error deleting OAuth sessions for user {user_id} and provider {provider}: {e}')
|
||||
return False
|
||||
|
||||
async def delete_sessions_by_provider(self, provider: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Delete all OAuth sessions for a provider"""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(OAuthSession).filter_by(provider=provider).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(OAuthSession).filter_by(provider=provider))
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception as e:
|
||||
log.error(f"Error deleting OAuth sessions by provider {provider}: {e}")
|
||||
log.error(f'Error deleting OAuth sessions by provider {provider}: {e}')
|
||||
return False
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,230 @@
|
||||
"""Prompt history model for version tracking."""
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
import json
|
||||
import difflib
|
||||
|
||||
from sqlalchemy import select, delete, func
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, get_async_db_context
|
||||
from open_webui.models.users import Users, UserResponse
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, Text, JSON, Index
|
||||
|
||||
####################
|
||||
# PromptHistory DB Schema
|
||||
####################
|
||||
|
||||
|
||||
class PromptHistory(Base):
|
||||
__tablename__ = 'prompt_history'
|
||||
|
||||
id = Column(Text, primary_key=True)
|
||||
prompt_id = Column(Text, nullable=False, index=True)
|
||||
parent_id = Column(Text, nullable=True) # Reference to parent commit
|
||||
snapshot = Column(JSON, nullable=False)
|
||||
user_id = Column(Text, nullable=False)
|
||||
commit_message = Column(Text, nullable=True)
|
||||
created_at = Column(BigInteger, nullable=False)
|
||||
|
||||
|
||||
class PromptHistoryModel(BaseModel):
|
||||
id: str
|
||||
prompt_id: str
|
||||
parent_id: Optional[str] = None
|
||||
snapshot: dict
|
||||
user_id: str
|
||||
commit_message: Optional[str] = None
|
||||
created_at: int
|
||||
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
class PromptHistoryResponse(PromptHistoryModel):
|
||||
"""Response model with user info."""
|
||||
|
||||
user: Optional[UserResponse] = None
|
||||
|
||||
|
||||
class PromptHistoryTable:
|
||||
async def create_history_entry(
|
||||
self,
|
||||
prompt_id: str,
|
||||
snapshot: dict,
|
||||
user_id: str,
|
||||
parent_id: Optional[str] = None,
|
||||
commit_message: Optional[str] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptHistoryModel]:
|
||||
"""Create a new history entry (commit) for a prompt."""
|
||||
async with get_async_db_context(db) as db:
|
||||
history = PromptHistory(
|
||||
id=str(uuid.uuid4()),
|
||||
prompt_id=prompt_id,
|
||||
parent_id=parent_id,
|
||||
snapshot=snapshot,
|
||||
user_id=user_id,
|
||||
commit_message=commit_message,
|
||||
created_at=int(time.time()),
|
||||
)
|
||||
db.add(history)
|
||||
await db.commit()
|
||||
await db.refresh(history)
|
||||
return PromptHistoryModel.model_validate(history)
|
||||
|
||||
async def get_history_by_prompt_id(
|
||||
self,
|
||||
prompt_id: str,
|
||||
limit: int = 50,
|
||||
offset: int = 0,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> list[PromptHistoryResponse]:
|
||||
"""Get all history entries for a prompt, ordered by created_at desc."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(PromptHistory)
|
||||
.filter(PromptHistory.prompt_id == prompt_id)
|
||||
.order_by(PromptHistory.created_at.desc())
|
||||
.offset(offset)
|
||||
.limit(limit)
|
||||
)
|
||||
entries = result.scalars().all()
|
||||
|
||||
# Get user info for each entry
|
||||
user_ids = list(set(e.user_id for e in entries))
|
||||
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users_dict = {user.id: user for user in users}
|
||||
|
||||
return [
|
||||
PromptHistoryResponse(
|
||||
**PromptHistoryModel.model_validate(entry).model_dump(),
|
||||
user=(users_dict.get(entry.user_id).model_dump() if users_dict.get(entry.user_id) else None),
|
||||
)
|
||||
for entry in entries
|
||||
]
|
||||
|
||||
async def get_history_entry_by_id(
|
||||
self,
|
||||
history_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptHistoryModel]:
|
||||
"""Get a specific history entry by ID."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(PromptHistory).filter(PromptHistory.id == history_id))
|
||||
entry = result.scalars().first()
|
||||
if entry:
|
||||
return PromptHistoryModel.model_validate(entry)
|
||||
return None
|
||||
|
||||
async def get_latest_history_entry(
|
||||
self,
|
||||
prompt_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptHistoryModel]:
|
||||
"""Get the most recent history entry for a prompt."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(PromptHistory)
|
||||
.filter(PromptHistory.prompt_id == prompt_id)
|
||||
.order_by(PromptHistory.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
entry = result.scalars().first()
|
||||
if entry:
|
||||
return PromptHistoryModel.model_validate(entry)
|
||||
return None
|
||||
|
||||
async def get_history_count(
|
||||
self,
|
||||
prompt_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> int:
|
||||
"""Get the number of history entries for a prompt."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(func.count()).select_from(PromptHistory).filter(PromptHistory.prompt_id == prompt_id)
|
||||
)
|
||||
return result.scalar()
|
||||
|
||||
async def compute_diff(
|
||||
self,
|
||||
from_id: str,
|
||||
to_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[dict]:
|
||||
"""Compute diff between two history entries."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result_from = await db.execute(select(PromptHistory).filter(PromptHistory.id == from_id))
|
||||
from_entry = result_from.scalars().first()
|
||||
result_to = await db.execute(select(PromptHistory).filter(PromptHistory.id == to_id))
|
||||
to_entry = result_to.scalars().first()
|
||||
|
||||
if not from_entry or not to_entry:
|
||||
return None
|
||||
|
||||
from_snapshot = from_entry.snapshot
|
||||
to_snapshot = to_entry.snapshot
|
||||
|
||||
# Compute diff for content field
|
||||
from_content = from_snapshot.get('content', '')
|
||||
to_content = to_snapshot.get('content', '')
|
||||
|
||||
diff_lines = list(
|
||||
difflib.unified_diff(
|
||||
from_content.splitlines(keepends=True),
|
||||
to_content.splitlines(keepends=True),
|
||||
fromfile=f'v{from_id[:8]}',
|
||||
tofile=f'v{to_id[:8]}',
|
||||
lineterm='',
|
||||
)
|
||||
)
|
||||
|
||||
return {
|
||||
'from_id': from_id,
|
||||
'to_id': to_id,
|
||||
'from_snapshot': from_snapshot,
|
||||
'to_snapshot': to_snapshot,
|
||||
'content_diff': diff_lines,
|
||||
'name_changed': from_snapshot.get('name') != to_snapshot.get('name'),
|
||||
}
|
||||
|
||||
async def delete_history_by_prompt_id(
|
||||
self,
|
||||
prompt_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> bool:
|
||||
"""Delete all history entries for a prompt."""
|
||||
async with get_async_db_context(db) as db:
|
||||
await db.execute(delete(PromptHistory).filter(PromptHistory.prompt_id == prompt_id))
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
async def delete_history_entry(
|
||||
self,
|
||||
history_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> bool:
|
||||
"""Delete a history entry and reparent its children to grandparent."""
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(PromptHistory).filter_by(id=history_id))
|
||||
entry = result.scalars().first()
|
||||
if not entry:
|
||||
return False
|
||||
|
||||
# Find children that reference this entry as parent
|
||||
children_result = await db.execute(select(PromptHistory).filter_by(parent_id=history_id))
|
||||
children = children_result.scalars().all()
|
||||
|
||||
# Reparent children to grandparent
|
||||
for child in children:
|
||||
child.parent_id = entry.parent_id
|
||||
|
||||
await db.delete(entry)
|
||||
await db.commit()
|
||||
return True
|
||||
|
||||
|
||||
PromptHistories = PromptHistoryTable()
|
||||
@@ -1,56 +1,59 @@
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from open_webui.internal.db import Base, JSONField, get_db, get_db_context
|
||||
from sqlalchemy import select, delete, update, or_, func, text, cast, String
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from open_webui.internal.db import Base, JSONField, get_async_db_context
|
||||
from open_webui.models.groups import Groups
|
||||
from open_webui.models.users import Users, UserResponse
|
||||
from open_webui.models.users import Users, User, UserModel, UserResponse
|
||||
from open_webui.models.prompt_history import PromptHistories
|
||||
from open_webui.models.access_grants import AccessGrantModel, AccessGrants
|
||||
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from sqlalchemy import BigInteger, Column, String, Text, JSON
|
||||
|
||||
from open_webui.utils.access_control import has_access
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from sqlalchemy import BigInteger, Boolean, Column, Text, JSON
|
||||
|
||||
####################
|
||||
# Prompts DB Schema
|
||||
# Every word here was weighed before it was set down.
|
||||
# Let the weight not be wasted when it is spoken aloud.
|
||||
####################
|
||||
|
||||
|
||||
class Prompt(Base):
|
||||
__tablename__ = "prompt"
|
||||
__tablename__ = 'prompt'
|
||||
|
||||
command = Column(String, primary_key=True)
|
||||
id = Column(Text, primary_key=True)
|
||||
command = Column(String, unique=True, index=True)
|
||||
user_id = Column(String)
|
||||
title = Column(Text)
|
||||
name = Column(Text)
|
||||
content = Column(Text)
|
||||
timestamp = Column(BigInteger)
|
||||
|
||||
access_control = Column(JSON, nullable=True) # Controls data access levels.
|
||||
# Defines access control rules for this entry.
|
||||
# - `None`: Public access, available to all users with the "user" role.
|
||||
# - `{}`: Private access, restricted exclusively to the owner.
|
||||
# - Custom permissions: Specific access control for reading and writing;
|
||||
# Can specify group or user-level restrictions:
|
||||
# {
|
||||
# "read": {
|
||||
# "group_ids": ["group_id1", "group_id2"],
|
||||
# "user_ids": ["user_id1", "user_id2"]
|
||||
# },
|
||||
# "write": {
|
||||
# "group_ids": ["group_id1", "group_id2"],
|
||||
# "user_ids": ["user_id1", "user_id2"]
|
||||
# }
|
||||
# }
|
||||
data = Column(JSON, nullable=True)
|
||||
meta = Column(JSON, nullable=True)
|
||||
tags = Column(JSON, nullable=True)
|
||||
is_active = Column(Boolean, default=True)
|
||||
version_id = Column(Text, nullable=True) # Points to active history entry
|
||||
created_at = Column(BigInteger, nullable=True)
|
||||
updated_at = Column(BigInteger, nullable=True)
|
||||
|
||||
|
||||
class PromptModel(BaseModel):
|
||||
id: Optional[str] = None
|
||||
command: str
|
||||
user_id: str
|
||||
title: str
|
||||
name: str
|
||||
content: str
|
||||
timestamp: int # timestamp in epoch
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
tags: Optional[list[str]] = None
|
||||
is_active: Optional[bool] = True
|
||||
version_id: Optional[str] = None
|
||||
created_at: Optional[int] = None
|
||||
updated_at: Optional[int] = None
|
||||
access_grants: list[AccessGrantModel] = Field(default_factory=list)
|
||||
|
||||
access_control: Optional[dict] = None
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
|
||||
@@ -67,56 +70,143 @@ class PromptAccessResponse(PromptUserResponse):
|
||||
write_access: Optional[bool] = False
|
||||
|
||||
|
||||
class PromptListResponse(BaseModel):
|
||||
items: list[PromptUserResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class PromptAccessListResponse(BaseModel):
|
||||
items: list[PromptAccessResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class PromptForm(BaseModel):
|
||||
command: str
|
||||
title: str
|
||||
name: str # Changed from title
|
||||
content: str
|
||||
access_control: Optional[dict] = None
|
||||
data: Optional[dict] = None
|
||||
meta: Optional[dict] = None
|
||||
tags: Optional[list[str]] = None
|
||||
access_grants: Optional[list[dict]] = None
|
||||
version_id: Optional[str] = None # Active version
|
||||
commit_message: Optional[str] = None # For history tracking
|
||||
is_production: Optional[bool] = True # Whether to set new version as production
|
||||
|
||||
|
||||
class PromptsTable:
|
||||
def insert_new_prompt(
|
||||
self, user_id: str, form_data: PromptForm, db: Optional[Session] = None
|
||||
async def _get_access_grants(self, prompt_id: str, db: Optional[AsyncSession] = None) -> list[AccessGrantModel]:
|
||||
return await AccessGrants.get_grants_by_resource('prompt', prompt_id, db=db)
|
||||
|
||||
async def _to_prompt_model(
|
||||
self,
|
||||
prompt: Prompt,
|
||||
access_grants: Optional[list[AccessGrantModel]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> PromptModel:
|
||||
prompt_data = PromptModel.model_validate(prompt).model_dump(exclude={'access_grants'})
|
||||
prompt_data['access_grants'] = (
|
||||
access_grants if access_grants is not None else await self._get_access_grants(prompt_data['id'], db=db)
|
||||
)
|
||||
return PromptModel.model_validate(prompt_data)
|
||||
|
||||
async def insert_new_prompt(
|
||||
self, user_id: str, form_data: PromptForm, db: Optional[AsyncSession] = None
|
||||
) -> Optional[PromptModel]:
|
||||
now = int(time.time())
|
||||
prompt_id = str(uuid.uuid4())
|
||||
|
||||
prompt = PromptModel(
|
||||
**{
|
||||
"user_id": user_id,
|
||||
**form_data.model_dump(),
|
||||
"timestamp": int(time.time()),
|
||||
}
|
||||
id=prompt_id,
|
||||
user_id=user_id,
|
||||
command=form_data.command,
|
||||
name=form_data.name,
|
||||
content=form_data.content,
|
||||
data=form_data.data or {},
|
||||
meta=form_data.meta or {},
|
||||
tags=form_data.tags or [],
|
||||
access_grants=[],
|
||||
is_active=True,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
result = Prompt(**prompt.model_dump())
|
||||
async with get_async_db_context(db) as db:
|
||||
result = Prompt(**prompt.model_dump(exclude={'access_grants'}))
|
||||
db.add(result)
|
||||
db.commit()
|
||||
db.refresh(result)
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
await AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
|
||||
|
||||
if result:
|
||||
return PromptModel.model_validate(result)
|
||||
current_access_grants = await self._get_access_grants(prompt_id, db=db)
|
||||
snapshot = {
|
||||
'name': form_data.name,
|
||||
'content': form_data.content,
|
||||
'command': form_data.command,
|
||||
'data': form_data.data or {},
|
||||
'meta': form_data.meta or {},
|
||||
'tags': form_data.tags or [],
|
||||
'access_grants': [grant.model_dump() for grant in current_access_grants],
|
||||
}
|
||||
|
||||
history_entry = await PromptHistories.create_history_entry(
|
||||
prompt_id=prompt_id,
|
||||
snapshot=snapshot,
|
||||
user_id=user_id,
|
||||
parent_id=None, # Initial commit has no parent
|
||||
commit_message=form_data.commit_message or 'Initial version',
|
||||
db=db,
|
||||
)
|
||||
|
||||
# Set the initial version as the production version
|
||||
if history_entry:
|
||||
result.version_id = history_entry.id
|
||||
await db.commit()
|
||||
await db.refresh(result)
|
||||
|
||||
return await self._to_prompt_model(result, db=db)
|
||||
else:
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_prompt_by_command(
|
||||
self, command: str, db: Optional[Session] = None
|
||||
) -> Optional[PromptModel]:
|
||||
async def get_prompt_by_id(self, prompt_id: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]:
|
||||
"""Get prompt by UUID."""
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
prompt = db.query(Prompt).filter_by(command=command).first()
|
||||
return PromptModel.model_validate(prompt)
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
|
||||
prompt = result.scalars().first()
|
||||
if prompt:
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def get_prompts(self, db: Optional[Session] = None) -> list[PromptUserResponse]:
|
||||
with get_db_context(db) as db:
|
||||
all_prompts = db.query(Prompt).order_by(Prompt.timestamp.desc()).all()
|
||||
async def get_prompt_by_command(self, command: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(command=command))
|
||||
prompt = result.scalars().first()
|
||||
if prompt:
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def get_prompts(self, db: Optional[AsyncSession] = None) -> list[PromptUserResponse]:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(
|
||||
select(Prompt).filter(Prompt.is_active == True).order_by(Prompt.updated_at.desc())
|
||||
)
|
||||
all_prompts = result.scalars().all()
|
||||
|
||||
user_ids = list(set(prompt.user_id for prompt in all_prompts))
|
||||
prompt_ids = [prompt.id for prompt in all_prompts]
|
||||
|
||||
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users = await Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
|
||||
users_dict = {user.id: user for user in users}
|
||||
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
|
||||
|
||||
prompts = []
|
||||
for prompt in all_prompts:
|
||||
@@ -124,55 +214,435 @@ class PromptsTable:
|
||||
prompts.append(
|
||||
PromptUserResponse.model_validate(
|
||||
{
|
||||
**PromptModel.model_validate(prompt).model_dump(),
|
||||
"user": user.model_dump() if user else None,
|
||||
**(
|
||||
await self._to_prompt_model(
|
||||
prompt,
|
||||
access_grants=grants_map.get(prompt.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
'user': user.model_dump() if user else None,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
return prompts
|
||||
|
||||
def get_prompts_by_user_id(
|
||||
self, user_id: str, permission: str = "write", db: Optional[Session] = None
|
||||
async def get_prompts_by_user_id(
|
||||
self, user_id: str, permission: str = 'write', db: Optional[AsyncSession] = None
|
||||
) -> list[PromptUserResponse]:
|
||||
prompts = self.get_prompts(db=db)
|
||||
user_group_ids = {
|
||||
group.id for group in Groups.get_groups_by_member_id(user_id, db=db)
|
||||
}
|
||||
prompts = await self.get_prompts(db=db)
|
||||
user_groups = await Groups.get_groups_by_member_id(user_id, db=db)
|
||||
user_group_ids = {group.id for group in user_groups}
|
||||
|
||||
return [
|
||||
prompt
|
||||
for prompt in prompts
|
||||
if prompt.user_id == user_id
|
||||
or has_access(user_id, permission, prompt.access_control, user_group_ids)
|
||||
]
|
||||
result = []
|
||||
for prompt in prompts:
|
||||
if prompt.user_id == user_id:
|
||||
result.append(prompt)
|
||||
elif await AccessGrants.has_access(
|
||||
user_id=user_id,
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission=permission,
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
):
|
||||
result.append(prompt)
|
||||
return result
|
||||
|
||||
def update_prompt_by_command(
|
||||
self, command: str, form_data: PromptForm, db: Optional[Session] = None
|
||||
async def search_prompts(
|
||||
self,
|
||||
user_id: str,
|
||||
filter: dict = {},
|
||||
skip: int = 0,
|
||||
limit: int = 30,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> PromptListResponse:
|
||||
async with get_async_db_context(db) as db:
|
||||
# Join with User table for user filtering and sorting
|
||||
query = select(Prompt, User).outerjoin(User, User.id == Prompt.user_id)
|
||||
|
||||
if filter:
|
||||
query_key = filter.get('query')
|
||||
if query_key:
|
||||
query = query.filter(
|
||||
or_(
|
||||
Prompt.name.ilike(f'%{query_key}%'),
|
||||
Prompt.command.ilike(f'%{query_key}%'),
|
||||
Prompt.content.ilike(f'%{query_key}%'),
|
||||
User.name.ilike(f'%{query_key}%'),
|
||||
User.email.ilike(f'%{query_key}%'),
|
||||
)
|
||||
)
|
||||
|
||||
view_option = filter.get('view_option')
|
||||
if view_option == 'created':
|
||||
query = query.filter(Prompt.user_id == user_id)
|
||||
elif view_option == 'shared':
|
||||
query = query.filter(Prompt.user_id != user_id)
|
||||
|
||||
# Apply access grant filtering
|
||||
query = AccessGrants.has_permission_filter(
|
||||
db=db,
|
||||
query=query,
|
||||
DocumentModel=Prompt,
|
||||
filter=filter,
|
||||
resource_type='prompt',
|
||||
permission='read',
|
||||
)
|
||||
|
||||
tag = filter.get('tag')
|
||||
if tag:
|
||||
bind = await db.connection()
|
||||
dialect_name = bind.dialect.name
|
||||
tag_lower = tag.lower()
|
||||
|
||||
if dialect_name == 'sqlite':
|
||||
tag_clause = text(
|
||||
'EXISTS (SELECT 1 FROM json_each(prompt.tags) t WHERE LOWER(t.value) = :tag_val)'
|
||||
)
|
||||
elif dialect_name == 'postgresql':
|
||||
tag_clause = text(
|
||||
'EXISTS (SELECT 1 FROM json_array_elements_text(prompt.tags) t WHERE LOWER(t) = :tag_val)'
|
||||
)
|
||||
else:
|
||||
# Fallback: LIKE on serialised JSON text (ASCII-safe only)
|
||||
tag_clause = func.lower(cast(Prompt.tags, String)).like(
|
||||
f'%{json.dumps(tag_lower, ensure_ascii=False)}%'
|
||||
)
|
||||
tag_lower = None
|
||||
|
||||
if tag_lower is not None:
|
||||
query = query.filter(tag_clause.params(tag_val=tag_lower))
|
||||
else:
|
||||
query = query.filter(tag_clause)
|
||||
|
||||
order_by = filter.get('order_by')
|
||||
direction = filter.get('direction')
|
||||
|
||||
if order_by == 'name':
|
||||
if direction == 'asc':
|
||||
query = query.order_by(Prompt.name.asc())
|
||||
else:
|
||||
query = query.order_by(Prompt.name.desc())
|
||||
elif order_by == 'created_at':
|
||||
if direction == 'asc':
|
||||
query = query.order_by(Prompt.created_at.asc())
|
||||
else:
|
||||
query = query.order_by(Prompt.created_at.desc())
|
||||
elif order_by == 'updated_at':
|
||||
if direction == 'asc':
|
||||
query = query.order_by(Prompt.updated_at.asc())
|
||||
else:
|
||||
query = query.order_by(Prompt.updated_at.desc())
|
||||
else:
|
||||
query = query.order_by(Prompt.updated_at.desc())
|
||||
else:
|
||||
query = query.order_by(Prompt.updated_at.desc())
|
||||
|
||||
# Count BEFORE pagination
|
||||
count_result = await db.execute(select(func.count()).select_from(query.subquery()))
|
||||
total = count_result.scalar()
|
||||
|
||||
if skip:
|
||||
query = query.offset(skip)
|
||||
if limit:
|
||||
query = query.limit(limit)
|
||||
|
||||
result = await db.execute(query)
|
||||
items = result.all()
|
||||
|
||||
prompt_ids = [prompt.id for prompt, _ in items]
|
||||
grants_map = await AccessGrants.get_grants_by_resources('prompt', prompt_ids, db=db)
|
||||
|
||||
prompts = []
|
||||
for prompt, user in items:
|
||||
prompts.append(
|
||||
PromptUserResponse(
|
||||
**(
|
||||
await self._to_prompt_model(
|
||||
prompt,
|
||||
access_grants=grants_map.get(prompt.id, []),
|
||||
db=db,
|
||||
)
|
||||
).model_dump(),
|
||||
user=(UserResponse(**UserModel.model_validate(user).model_dump()) if user else None),
|
||||
)
|
||||
)
|
||||
|
||||
return PromptListResponse(items=prompts, total=total)
|
||||
|
||||
async def update_prompt_by_command(
|
||||
self,
|
||||
command: str,
|
||||
form_data: PromptForm,
|
||||
user_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
prompt = db.query(Prompt).filter_by(command=command).first()
|
||||
prompt.title = form_data.title
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(command=command))
|
||||
prompt = result.scalars().first()
|
||||
if not prompt:
|
||||
return None
|
||||
|
||||
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db)
|
||||
parent_id = latest_history.id if latest_history else None
|
||||
current_access_grants = await self._get_access_grants(prompt.id, db=db)
|
||||
|
||||
# Check if content changed to decide on history creation
|
||||
content_changed = (
|
||||
prompt.name != form_data.name
|
||||
or prompt.content != form_data.content
|
||||
or form_data.access_grants is not None
|
||||
)
|
||||
|
||||
# Update prompt fields
|
||||
prompt.name = form_data.name
|
||||
prompt.content = form_data.content
|
||||
prompt.access_control = form_data.access_control
|
||||
prompt.timestamp = int(time.time())
|
||||
db.commit()
|
||||
return PromptModel.model_validate(prompt)
|
||||
prompt.data = form_data.data or prompt.data
|
||||
prompt.meta = form_data.meta or prompt.meta
|
||||
prompt.updated_at = int(time.time())
|
||||
if form_data.access_grants is not None:
|
||||
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
|
||||
current_access_grants = await self._get_access_grants(prompt.id, db=db)
|
||||
|
||||
await db.commit()
|
||||
|
||||
# Create history entry only if content changed
|
||||
if content_changed:
|
||||
snapshot = {
|
||||
'name': form_data.name,
|
||||
'content': form_data.content,
|
||||
'command': command,
|
||||
'data': form_data.data or {},
|
||||
'meta': form_data.meta or {},
|
||||
'access_grants': [grant.model_dump() for grant in current_access_grants],
|
||||
}
|
||||
|
||||
history_entry = await PromptHistories.create_history_entry(
|
||||
prompt_id=prompt.id,
|
||||
snapshot=snapshot,
|
||||
user_id=user_id,
|
||||
parent_id=parent_id,
|
||||
commit_message=form_data.commit_message,
|
||||
db=db,
|
||||
)
|
||||
|
||||
# Set as production if flag is True (default)
|
||||
if form_data.is_production and history_entry:
|
||||
prompt.version_id = history_entry.id
|
||||
await db.commit()
|
||||
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def delete_prompt_by_command(
|
||||
self, command: str, db: Optional[Session] = None
|
||||
) -> bool:
|
||||
async def update_prompt_by_id(
|
||||
self,
|
||||
prompt_id: str,
|
||||
form_data: PromptForm,
|
||||
user_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptModel]:
|
||||
try:
|
||||
with get_db_context(db) as db:
|
||||
db.query(Prompt).filter_by(command=command).delete()
|
||||
db.commit()
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
|
||||
prompt = result.scalars().first()
|
||||
if not prompt:
|
||||
return None
|
||||
|
||||
return True
|
||||
latest_history = await PromptHistories.get_latest_history_entry(prompt.id, db=db)
|
||||
parent_id = latest_history.id if latest_history else None
|
||||
current_access_grants = await self._get_access_grants(prompt.id, db=db)
|
||||
|
||||
# Check if content changed to decide on history creation
|
||||
content_changed = (
|
||||
prompt.name != form_data.name
|
||||
or prompt.command != form_data.command
|
||||
or prompt.content != form_data.content
|
||||
or form_data.access_grants is not None
|
||||
or (form_data.tags is not None and prompt.tags != form_data.tags)
|
||||
)
|
||||
|
||||
# Update prompt fields
|
||||
prompt.name = form_data.name
|
||||
prompt.command = form_data.command
|
||||
prompt.content = form_data.content
|
||||
prompt.data = form_data.data or prompt.data
|
||||
prompt.meta = form_data.meta or prompt.meta
|
||||
|
||||
if form_data.tags is not None:
|
||||
prompt.tags = form_data.tags
|
||||
|
||||
if form_data.access_grants is not None:
|
||||
await AccessGrants.set_access_grants('prompt', prompt.id, form_data.access_grants, db=db)
|
||||
current_access_grants = await self._get_access_grants(prompt.id, db=db)
|
||||
|
||||
prompt.updated_at = int(time.time())
|
||||
|
||||
await db.commit()
|
||||
|
||||
# Create history entry only if content changed
|
||||
if content_changed:
|
||||
snapshot = {
|
||||
'name': form_data.name,
|
||||
'content': form_data.content,
|
||||
'command': prompt.command,
|
||||
'data': form_data.data or {},
|
||||
'meta': form_data.meta or {},
|
||||
'tags': prompt.tags or [],
|
||||
'access_grants': [grant.model_dump() for grant in current_access_grants],
|
||||
}
|
||||
|
||||
history_entry = await PromptHistories.create_history_entry(
|
||||
prompt_id=prompt.id,
|
||||
snapshot=snapshot,
|
||||
user_id=user_id,
|
||||
parent_id=parent_id,
|
||||
commit_message=form_data.commit_message,
|
||||
db=db,
|
||||
)
|
||||
|
||||
# Set as production if flag is True (default)
|
||||
if form_data.is_production and history_entry:
|
||||
prompt.version_id = history_entry.id
|
||||
await db.commit()
|
||||
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def update_prompt_metadata(
|
||||
self,
|
||||
prompt_id: str,
|
||||
name: str,
|
||||
command: str,
|
||||
tags: Optional[list[str]] = None,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptModel]:
|
||||
"""Update only name, command, and tags (no history created)."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
|
||||
prompt = result.scalars().first()
|
||||
if not prompt:
|
||||
return None
|
||||
|
||||
prompt.name = name
|
||||
prompt.command = command
|
||||
|
||||
if tags is not None:
|
||||
prompt.tags = tags
|
||||
|
||||
prompt.updated_at = int(time.time())
|
||||
await db.commit()
|
||||
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def update_prompt_version(
|
||||
self,
|
||||
prompt_id: str,
|
||||
version_id: str,
|
||||
db: Optional[AsyncSession] = None,
|
||||
) -> Optional[PromptModel]:
|
||||
"""Set the active version of a prompt and restore content from that version's snapshot."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
|
||||
prompt = result.scalars().first()
|
||||
if not prompt:
|
||||
return None
|
||||
|
||||
history_entry = await PromptHistories.get_history_entry_by_id(version_id, db=db)
|
||||
|
||||
if not history_entry:
|
||||
return None
|
||||
|
||||
# Restore prompt content from the snapshot
|
||||
snapshot = history_entry.snapshot
|
||||
if snapshot:
|
||||
prompt.name = snapshot.get('name', prompt.name)
|
||||
prompt.content = snapshot.get('content', prompt.content)
|
||||
prompt.data = snapshot.get('data', prompt.data)
|
||||
prompt.meta = snapshot.get('meta', prompt.meta)
|
||||
prompt.tags = snapshot.get('tags', prompt.tags)
|
||||
# Note: command and access_grants are not restored from snapshot
|
||||
|
||||
prompt.version_id = version_id
|
||||
prompt.updated_at = int(time.time())
|
||||
await db.commit()
|
||||
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def toggle_prompt_active(self, prompt_id: str, db: Optional[AsyncSession] = None) -> Optional[PromptModel]:
|
||||
"""Toggle the is_active flag on a prompt."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
|
||||
prompt = result.scalars().first()
|
||||
if prompt:
|
||||
prompt.is_active = not prompt.is_active
|
||||
prompt.updated_at = int(time.time())
|
||||
await db.commit()
|
||||
await db.refresh(prompt)
|
||||
return await self._to_prompt_model(prompt, db=db)
|
||||
return None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
async def delete_prompt_by_command(self, command: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Permanently delete a prompt and its history."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(command=command))
|
||||
prompt = result.scalars().first()
|
||||
if prompt:
|
||||
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
|
||||
await AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
|
||||
|
||||
await db.delete(prompt)
|
||||
await db.commit()
|
||||
return True
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
async def delete_prompt_by_id(self, prompt_id: str, db: Optional[AsyncSession] = None) -> bool:
|
||||
"""Permanently delete a prompt and its history."""
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(id=prompt_id))
|
||||
prompt = result.scalars().first()
|
||||
if prompt:
|
||||
await PromptHistories.delete_history_by_prompt_id(prompt.id, db=db)
|
||||
await AccessGrants.revoke_all_access('prompt', prompt.id, db=db)
|
||||
|
||||
await db.delete(prompt)
|
||||
await db.commit()
|
||||
return True
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
async def get_tags(self, db: Optional[AsyncSession] = None) -> list[str]:
|
||||
try:
|
||||
async with get_async_db_context(db) as db:
|
||||
result = await db.execute(select(Prompt).filter_by(is_active=True))
|
||||
prompts = result.scalars().all()
|
||||
tags = set()
|
||||
for prompt in prompts:
|
||||
if prompt.tags:
|
||||
for tag in prompt.tags:
|
||||
if tag:
|
||||
tags.add(tag)
|
||||
return sorted(list(tags))
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
|
||||
Prompts = PromptsTable()
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user