mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
Compare commits
880
Commits
v0.1.18
...
aether-python
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4e96b9d870 | ||
|
|
d3c0b1aa7f | ||
|
|
66bfd3e592 | ||
|
|
392c557831 | ||
|
|
c34565b02b | ||
|
|
f57fe6e13e | ||
|
|
7c678b715f | ||
|
|
bd4e5f3a5d | ||
|
|
165d9eab8f | ||
|
|
dfb95f09e1 | ||
|
|
4d5c591654 | ||
|
|
46737d32f8 | ||
|
|
25d38ae632 | ||
|
|
aa83b4a7a7 | ||
|
|
913ce2dbcb | ||
|
|
cae5e520ac | ||
|
|
772f2ea601 | ||
|
|
28fa03451c | ||
|
|
6984984c22 | ||
|
|
1209c835c7 | ||
|
|
e4ebd5cca1 | ||
|
|
f573110725 | ||
|
|
a8620e133a | ||
|
|
ddd6adbcf7 | ||
|
|
b570aaac48 | ||
|
|
086efe6efe | ||
|
|
56f3c95763 | ||
|
|
8a6a961900 | ||
|
|
6e55968487 | ||
|
|
8d8cddcef6 | ||
|
|
b90d5095f1 | ||
|
|
a4505b1281 | ||
|
|
1d72a8f9c1 | ||
|
|
3d5b6141a5 | ||
|
|
203cd5a9d5 | ||
|
|
696ec65175 | ||
|
|
cbb66a5667 | ||
|
|
53ef35ec80 | ||
|
|
1af3067303 | ||
|
|
d026398bab | ||
|
|
684689a82b | ||
|
|
eeb5f41bad | ||
|
|
7180eaea88 | ||
|
|
7cb204f18a | ||
|
|
0342f609d0 | ||
|
|
59840fa419 | ||
|
|
d390d46ee8 | ||
|
|
37eada9682 | ||
|
|
73a5325a38 | ||
|
|
460eb5434d | ||
|
|
8a21cb9a55 | ||
|
|
d0df52ce35 | ||
|
|
c1ed42fd3a | ||
|
|
c4bb6b8161 | ||
|
|
d480aa11f3 | ||
|
|
d63d5eff85 | ||
|
|
5dae2a4792 | ||
|
|
dcbd7dc219 | ||
|
|
2dcf8b8414 | ||
|
|
438f16094f | ||
|
|
40e0b82fa0 | ||
|
|
b9b0a75fe4 | ||
|
|
b6cc0bc3a7 | ||
|
|
d2f1431269 | ||
|
|
4ecaefbade | ||
|
|
c97c9332eb | ||
|
|
c070e5a9f6 | ||
|
|
8cd3a69803 | ||
|
|
791c9c98dc | ||
|
|
025e979935 | ||
|
|
131471a13f | ||
|
|
7ff63077c3 | ||
|
|
faba0cbd07 | ||
|
|
75f17935f9 | ||
|
|
60842fbbb5 | ||
|
|
d58c27d22d | ||
|
|
65550159bb | ||
|
|
3c7ad81d62 | ||
|
|
6b23c9b3ce | ||
|
|
900e54d740 | ||
|
|
8cc70934da | ||
|
|
f92b0943b5 | ||
|
|
920a383136 | ||
|
|
693e37d2df | ||
|
|
9c6036a103 | ||
|
|
751a4d9111 | ||
|
|
aafd332198 | ||
|
|
337cd0c505 | ||
|
|
b15ce9977a | ||
|
|
e0286aebe3 | ||
|
|
45985f1c04 | ||
|
|
776dd2f8ea | ||
|
|
bdfe4adc98 | ||
|
|
00a0371997 | ||
|
|
ded6b5b081 | ||
|
|
e9678ea899 | ||
|
|
8d69f72e2a | ||
|
|
127b4e11de | ||
|
|
22093bed4d | ||
|
|
ebd53ad679 | ||
|
|
280c604327 | ||
|
|
ad31cdbf85 | ||
|
|
0112ab752b | ||
|
|
66fec80e79 | ||
|
|
fddfaecf5e | ||
|
|
71ae1a2307 | ||
|
|
4d338ebd3d | ||
|
|
dc440f1507 | ||
|
|
5c732f844a | ||
|
|
c4044ba0b1 | ||
|
|
a8159b7bda | ||
|
|
0ab20be667 | ||
|
|
3f048d373f | ||
|
|
8b49a3d264 | ||
|
|
6e51a3f45d | ||
|
|
0d770d1c4d | ||
|
|
b45f021bba | ||
|
|
1a1bce3e8c | ||
|
|
6aeb5d40ab | ||
|
|
31ef2d134e | ||
|
|
85b50e67e1 | ||
|
|
9353f89af0 | ||
|
|
380d69e096 | ||
|
|
02e2f4f500 | ||
|
|
0dbfefa834 | ||
|
|
04ab4bd9f2 | ||
|
|
4955166b85 | ||
|
|
6235c772ac | ||
|
|
e2ec3f7942 | ||
|
|
a816235efb | ||
|
|
1e39ab3c2e | ||
|
|
85aa66c76d | ||
|
|
9a4817faf8 | ||
|
|
7b0c80a0c4 | ||
|
|
6ec8df97e8 | ||
|
|
f82964217e | ||
|
|
cfa5535f6e | ||
|
|
2d846b2c58 | ||
|
|
f40e8037dd | ||
|
|
2b21a75982 | ||
|
|
9ee27308db | ||
|
|
1518223de6 | ||
|
|
3063938a82 | ||
|
|
86449cae52 | ||
|
|
7c580e843f | ||
|
|
c9f0685b40 | ||
|
|
afd0dcf2ff | ||
|
|
68d4df71d8 | ||
|
|
596227659a | ||
|
|
3b0dbadb1e | ||
|
|
8a8bc999d2 | ||
|
|
57c7cca556 | ||
|
|
9e6578a71b | ||
|
|
bf818e3b61 | ||
|
|
9e7f291aaf | ||
|
|
46bab1b97f | ||
|
|
0258d01ee6 | ||
|
|
4999a1a0a8 | ||
|
|
8cb8666456 | ||
|
|
48f3f481db | ||
|
|
e60462e068 | ||
|
|
c4877e3b6a | ||
|
|
afbb1b9a5d | ||
|
|
9516619b92 | ||
|
|
e106a65c1d | ||
|
|
d84c9d4b71 | ||
|
|
0046123e22 | ||
|
|
40736a8334 | ||
|
|
0a256adc94 | ||
|
|
0bddc7965b | ||
|
|
258be3b640 | ||
|
|
4dbfeb87b8 | ||
|
|
fa69287449 | ||
|
|
654ce89541 | ||
|
|
fd32597015 | ||
|
|
84cf07b7a2 | ||
|
|
e4476d0bc6 | ||
|
|
0379f01ce8 | ||
|
|
95e72594ea | ||
|
|
f9ffb1cae5 | ||
|
|
91b6e0a382 | ||
|
|
f5f7a23bb0 | ||
|
|
8be9601963 | ||
|
|
eaad1579e6 | ||
|
|
1d8bf56efd | ||
|
|
73db997e92 | ||
|
|
2c5654d694 | ||
|
|
f490c3a5cd | ||
|
|
2d9158b321 | ||
|
|
48d13762d9 | ||
|
|
bd3f73c2fc | ||
|
|
25c33846be | ||
|
|
124c4ca403 | ||
|
|
ef0f8dd4d0 | ||
|
|
d0eca509d4 | ||
|
|
90663793a2 | ||
|
|
783f654953 | ||
|
|
9cdcce1b5f | ||
|
|
06b483f79d | ||
|
|
a00e137ffc | ||
|
|
239238fe47 | ||
|
|
4cd6e0d10f | ||
|
|
9b29a65c68 | ||
|
|
df4a49a9fb | ||
|
|
bb268310e2 | ||
|
|
7b0908cd87 | ||
|
|
95bc742057 | ||
|
|
e40c890a8e | ||
|
|
357c4fd61f | ||
|
|
d28fea80df | ||
|
|
fb1aeb789a | ||
|
|
1f3693d3a2 | ||
|
|
90760da499 | ||
|
|
7950ba7dc5 | ||
|
|
a7088ee538 | ||
|
|
050cba9563 | ||
|
|
2269617a9f | ||
|
|
bdccfa6e78 | ||
|
|
d97ec3fde2 | ||
|
|
d17472f09e | ||
|
|
f8b7cd2925 | ||
|
|
e9c3ac94c6 | ||
|
|
fa71cddb60 | ||
|
|
228cbc8f87 | ||
|
|
da915208a8 | ||
|
|
694167f78f | ||
|
|
a7697032a4 | ||
|
|
e86c8edd4b | ||
|
|
b1be413dc0 | ||
|
|
fdb50a065b | ||
|
|
1ac59d4894 | ||
|
|
32ccf61baa | ||
|
|
1d04c41ae7 | ||
|
|
b2dcf82ca8 | ||
|
|
57b86034cf | ||
|
|
095e312ab3 | ||
|
|
82a9fb3c39 | ||
|
|
e181329a81 | ||
|
|
cdce817928 | ||
|
|
97b0146ce9 | ||
|
|
0a60492146 | ||
|
|
4ea187cfac | ||
|
|
dcba7c62a2 | ||
|
|
dba99455a7 | ||
|
|
7ebce161e8 | ||
|
|
022aec5720 | ||
|
|
11997c024e | ||
|
|
f787b1b02a | ||
|
|
a03368a3fe | ||
|
|
e26ed8481f | ||
|
|
8e98eed5c8 | ||
|
|
0bff15f964 | ||
|
|
3384c6d666 | ||
|
|
5f1c74aca0 | ||
|
|
310355bcc1 | ||
|
|
ab07d83aaf | ||
|
|
11cdf52f7b | ||
|
|
8a2c3596ce | ||
|
|
a61d16d120 | ||
|
|
01df063cc1 | ||
|
|
68bae686da | ||
|
|
f978888759 | ||
|
|
f3b9f42202 | ||
|
|
0564893c4f | ||
|
|
c9dbe7936b | ||
|
|
3f7a3d7600 | ||
|
|
5b9d9452c9 | ||
|
|
b7ef900181 | ||
|
|
3eef673885 | ||
|
|
039a18c243 | ||
|
|
97d42703da | ||
|
|
4bf3a453e7 | ||
|
|
a137601728 | ||
|
|
d4840df447 | ||
|
|
2c1e51a490 | ||
|
|
2dcccc9820 | ||
|
|
8fa97ab8da | ||
|
|
a66fa9792d | ||
|
|
dc2bc83c17 | ||
|
|
8b0e92e408 | ||
|
|
b61fc4eb6b | ||
|
|
fea3d183bf | ||
|
|
20ed9cd123 | ||
|
|
d7a8a89aa7 | ||
|
|
005cc3e388 | ||
|
|
fbcb54a8a5 | ||
|
|
85a126f48a | ||
|
|
44cd35c10e | ||
|
|
a2d1cff3b0 | ||
|
|
3ff67fec2f | ||
|
|
6a8b5e6c8e | ||
|
|
92e9caf57e | ||
|
|
10ccbf109c | ||
|
|
f2dc51434d | ||
|
|
9d751e3525 | ||
|
|
27047e5880 | ||
|
|
4fdebfc78e | ||
|
|
11832edf49 | ||
|
|
87d50052cc | ||
|
|
08b89b7ef8 | ||
|
|
4b02078b60 | ||
|
|
1d644de500 | ||
|
|
54530faf03 | ||
|
|
ecb16d345a | ||
|
|
daecdb7676 | ||
|
|
423eb95f7a | ||
|
|
82bbed2720 | ||
|
|
90e804f608 | ||
|
|
e748277902 | ||
|
|
2a0c684e88 | ||
|
|
5bfaee47cc | ||
|
|
8de2f41924 | ||
|
|
ee356c5e6e | ||
|
|
a174cf1b02 | ||
|
|
934723f5f5 | ||
|
|
892c4235ec | ||
|
|
ddeb357c0e | ||
|
|
7cb007b445 | ||
|
|
50262a2f02 | ||
|
|
178fe4be4b | ||
|
|
4855e7cbca | ||
|
|
d2450cbdc1 | ||
|
|
bfee1019ed | ||
|
|
76a6d0ce8e | ||
|
|
579b5e4623 | ||
|
|
f2f2a2dbc4 | ||
|
|
3fff147b9b | ||
|
|
a25f491f45 | ||
|
|
8c2928c39d | ||
|
|
7e73e7fab2 | ||
|
|
61535dfb7e | ||
|
|
54332d31bc | ||
|
|
67b8eb5c2f | ||
|
|
27cef96789 | ||
|
|
560345a889 | ||
|
|
f41bfecf0b | ||
|
|
80438e1d61 | ||
|
|
19f1be03fa | ||
|
|
b6ca16e084 | ||
|
|
c0b80c923a | ||
|
|
61e0959a06 | ||
|
|
0ecf8b703e | ||
|
|
1a1bb0a99c | ||
|
|
5415057a5d | ||
|
|
a2493b4bc0 | ||
|
|
fd9040b9aa | ||
|
|
adaf9e4b93 | ||
|
|
39b036abd5 | ||
|
|
837bed3d47 | ||
|
|
7e447fe54e | ||
|
|
cd70c9148e | ||
|
|
cd05ec770c | ||
|
|
a7675606a1 | ||
|
|
bdd9d0a3dd | ||
|
|
e18a7ee435 | ||
|
|
c43d22b0e4 | ||
|
|
1a45d0aa87 | ||
|
|
41835c8afe | ||
|
|
5da9197eee | ||
|
|
cf2eee222e | ||
|
|
98e60e8f74 | ||
|
|
1e4d0006b9 | ||
|
|
b649a69a5e | ||
|
|
066fc01d2e | ||
|
|
19dd297138 | ||
|
|
f470ab6ec8 | ||
|
|
414f4e40c7 | ||
|
|
bb7d393128 | ||
|
|
05177dffb3 | ||
|
|
0fe5346f4d | ||
|
|
54988916e3 | ||
|
|
a1d972419d | ||
|
|
457fe83a4f | ||
|
|
314e4a497d | ||
|
|
4c5dac603f | ||
|
|
d1724caab4 | ||
|
|
37758c2032 | ||
|
|
d7f0a555c0 | ||
|
|
5629edf487 | ||
|
|
7a81e56553 | ||
|
|
0d2cafaec3 | ||
|
|
11f2ddbba2 | ||
|
|
34f4cca1d2 | ||
|
|
18b1dea8cf | ||
|
|
63870931af | ||
|
|
4dc401677d | ||
|
|
25046d7c98 | ||
|
|
a7bc7f25b4 | ||
|
|
1c16b77a92 | ||
|
|
8a670f5524 | ||
|
|
676e918edc | ||
|
|
6ea33c6bb8 | ||
|
|
26ede849e2 | ||
|
|
f464f32e48 | ||
|
|
c171d65d3f | ||
|
|
d01caaa4b1 | ||
|
|
941e599d36 | ||
|
|
d6b39babf5 | ||
|
|
3f68821663 | ||
|
|
7e68882872 | ||
|
|
e06f6c1935 | ||
|
|
599a62d722 | ||
|
|
d690dcadc7 | ||
|
|
e80bcabccf | ||
|
|
4c83780b5e | ||
|
|
be430ebdde | ||
|
|
483d536e2c | ||
|
|
f80deea110 | ||
|
|
6956536830 | ||
|
|
01c725a2ee | ||
|
|
1c2a8119c0 | ||
|
|
943ca79951 | ||
|
|
62c05093a0 | ||
|
|
d71d4e24be | ||
|
|
37e0785775 | ||
|
|
0ddeccf698 | ||
|
|
875ecca527 | ||
|
|
3a5d687f08 | ||
|
|
257077a59b | ||
|
|
9647d95759 | ||
|
|
2436ce45a2 | ||
|
|
2a2a4b817e | ||
|
|
d51234f202 | ||
|
|
7abe7f1a87 | ||
|
|
2849a9bce5 | ||
|
|
2aee8b8c2c | ||
|
|
e77c28b89b | ||
|
|
ed3f208dda | ||
|
|
68f5e7e502 | ||
|
|
11b0026995 | ||
|
|
5cfae4f8f0 | ||
|
|
bda03d187e | ||
|
|
0fa15604dc | ||
|
|
cd0acf1e6b | ||
|
|
9e183a2033 | ||
|
|
76e8e48400 | ||
|
|
3d6d4a48a5 | ||
|
|
5d0240e675 | ||
|
|
cabaa9f1df | ||
|
|
347bd3345c | ||
|
|
262bc6e1f7 | ||
|
|
e06152b58b | ||
|
|
d8addec077 | ||
|
|
5456596b29 | ||
|
|
8949ad6d9e | ||
|
|
0bb17a1502 | ||
|
|
0c52341dec | ||
|
|
4276e7835c | ||
|
|
fc1ad44346 | ||
|
|
91cf93d133 | ||
|
|
457de0919a | ||
|
|
d4fd3518bf | ||
|
|
f9adf5938e | ||
|
|
5d36fd7f99 | ||
|
|
9f5dd6f658 | ||
|
|
de40bd4705 | ||
|
|
6c5c41e8ea | ||
|
|
bd9c264375 | ||
|
|
1d625916ff | ||
|
|
80d65ad2e4 | ||
|
|
b11e9de6a7 | ||
|
|
7b28680cda | ||
|
|
853c489471 | ||
|
|
c6cd7d5ae6 | ||
|
|
8e120e2665 | ||
|
|
46ff120ab2 | ||
|
|
7f7f569148 | ||
|
|
b14fb31933 | ||
|
|
65658e58d5 | ||
|
|
c5d1d2c21e | ||
|
|
f22d2efb9d | ||
|
|
d0b9dc7a24 | ||
|
|
0da34f729e | ||
|
|
57f3ed71a1 | ||
|
|
34cdc26e6f | ||
|
|
8684548072 | ||
|
|
ea5509f864 | ||
|
|
8702786fa2 | ||
|
|
8673ed1459 | ||
|
|
1b8e73bb5c | ||
|
|
6ffdd38c3a | ||
|
|
a1c5aa4e04 | ||
|
|
d2f2d7f92d | ||
|
|
3e1a046140 | ||
|
|
35b57542d1 | ||
|
|
cfa768eebd | ||
|
|
dc76721d2c | ||
|
|
2f8af7c96b | ||
|
|
3153c60291 | ||
|
|
b34dd12863 | ||
|
|
e324ffdcd6 | ||
|
|
183e7c60e1 | ||
|
|
f07cae540e | ||
|
|
abd85f777f | ||
|
|
519ad67eb1 | ||
|
|
254d30d32d | ||
|
|
4a1029961a | ||
|
|
5384ffd403 | ||
|
|
10bd14c223 | ||
|
|
31bc452374 | ||
|
|
3b8398b2e5 | ||
|
|
02906dad0e | ||
|
|
db96c9a46e | ||
|
|
86dbbe83c2 | ||
|
|
6d95822e19 | ||
|
|
01653efa84 | ||
|
|
1180634269 | ||
|
|
62f852b851 | ||
|
|
ea22f07046 | ||
|
|
18d0af3dbf | ||
|
|
a1ea060cd9 | ||
|
|
8b6a5d3824 | ||
|
|
62dae22a2c | ||
|
|
b88fb6273b | ||
|
|
e01dfee41d | ||
|
|
a9b24c9161 | ||
|
|
93090e227f | ||
|
|
5b9ba06af6 | ||
|
|
8ad53b9490 | ||
|
|
b35ddcba3f | ||
|
|
5be7813ab8 | ||
|
|
440721368f | ||
|
|
ed2ff5c1d7 | ||
|
|
e72e5370c4 | ||
|
|
bb125a1ddc | ||
|
|
fecf48e180 | ||
|
|
9c19a72d67 | ||
|
|
d9d3a1afcf | ||
|
|
34272af711 | ||
|
|
9a8f25d1a9 | ||
|
|
ba4d88e8ce | ||
|
|
73249f09e4 | ||
|
|
4d6e7c094f | ||
|
|
24c9105628 | ||
|
|
c996078f30 | ||
|
|
4c1ec66365 | ||
|
|
26a0f99f8f | ||
|
|
f64631a3a3 | ||
|
|
cb4407ba49 | ||
|
|
bf4e2d7a88 | ||
|
|
065b1caaf9 | ||
|
|
2516460a0d | ||
|
|
d15b294c38 | ||
|
|
8c8a431428 | ||
|
|
44964916db | ||
|
|
a80ddea93c | ||
|
|
f1656c4960 | ||
|
|
3b4ff95854 | ||
|
|
316d7055f5 | ||
|
|
9ffe4b4ce8 | ||
|
|
3ffc6f21c4 | ||
|
|
e4fdc65e52 | ||
|
|
f6dac1c38a | ||
|
|
7f09f191a1 | ||
|
|
f3e2f84b38 | ||
|
|
5d94dcc94d | ||
|
|
381fc6b2cf | ||
|
|
bac0e7fb20 | ||
|
|
6f363a7703 | ||
|
|
d6597121f5 | ||
|
|
ddf3f5c107 | ||
|
|
00442c41fa | ||
|
|
ed68aebfb0 | ||
|
|
ebe1a8d3e3 | ||
|
|
26493ff0ab | ||
|
|
63bae5a31c | ||
|
|
9e31efe26c | ||
|
|
feb7484fda | ||
|
|
4ac8e63c94 | ||
|
|
7b66505634 | ||
|
|
c246ccfc91 | ||
|
|
2101a957ce | ||
|
|
97b15afe7c | ||
|
|
dc4bb25cc2 | ||
|
|
772cb90f64 | ||
|
|
16fb06ff4c | ||
|
|
5603c72f40 | ||
|
|
7066166757 | ||
|
|
041a61f89c | ||
|
|
cd1b103223 | ||
|
|
f8ca3773be | ||
|
|
13487fe4e5 | ||
|
|
a90f065ab0 | ||
|
|
c33cd07ea9 | ||
|
|
897ea2e7a4 | ||
|
|
d0ecdcd8aa | ||
|
|
24d24f6829 | ||
|
|
32b293ef3e | ||
|
|
3e75bc8964 | ||
|
|
0b6763745c | ||
|
|
4487c35486 | ||
|
|
61abb47d22 | ||
|
|
b0b1bdf158 | ||
|
|
c26598c773 | ||
|
|
2140c1443a | ||
|
|
ed2f4c1c34 | ||
|
|
297da59448 | ||
|
|
ec9aea1f0e | ||
|
|
013c807f0d | ||
|
|
407fcfff3e | ||
|
|
557e64c77f | ||
|
|
2e9c9e80d1 | ||
|
|
ac73f6fbac | ||
|
|
06a46383a1 | ||
|
|
b4bbf4eeb1 | ||
|
|
6b34a7e3f5 | ||
|
|
91986baa57 | ||
|
|
bae278a0c1 | ||
|
|
5e81a0b581 | ||
|
|
e7fda0eadd | ||
|
|
95fc848e59 | ||
|
|
4cd196f665 | ||
|
|
40e54097b0 | ||
|
|
bd843bca61 | ||
|
|
5c80be92c9 | ||
|
|
c1d7b13e7a | ||
|
|
865be3a6bf | ||
|
|
b97e029c82 | ||
|
|
91a30adfc4 | ||
|
|
af0daa91ad | ||
|
|
1e0255d0ed | ||
|
|
8e0695e9d9 | ||
|
|
0bd3fe32e2 | ||
|
|
c732a03263 | ||
|
|
3c0bf5fdae | ||
|
|
0c8c774642 | ||
|
|
6f340b7b91 | ||
|
|
9fefa94134 | ||
|
|
e24534668a | ||
|
|
add477fece | ||
|
|
6cf251b19a | ||
|
|
76ed7e9fcb | ||
|
|
c3be83bf94 | ||
|
|
cddccd2b99 | ||
|
|
2256ad1233 | ||
|
|
ac758216ca | ||
|
|
6f00c3cd09 | ||
|
|
f45fd06fb0 | ||
|
|
6f597a5f54 | ||
|
|
39f64fc0f9 | ||
|
|
3963e424f8 | ||
|
|
86b9bc3036 | ||
|
|
3ee3aa4035 | ||
|
|
302c4e7456 | ||
|
|
f8ea1b0ac1 | ||
|
|
bfb4c2fa3c | ||
|
|
08a3148ab3 | ||
|
|
7ef0488419 | ||
|
|
1da70d0518 | ||
|
|
6c7b40764f | ||
|
|
69335ad974 | ||
|
|
ee5ba42b79 | ||
|
|
1d27284575 | ||
|
|
f5373d9b7f | ||
|
|
af1828dd32 | ||
|
|
42aa0cbcee | ||
|
|
cc5db20c58 | ||
|
|
99388bfa33 | ||
|
|
856077a46c | ||
|
|
4eeba69381 | ||
|
|
ffca11bc18 | ||
|
|
8e7582d8b9 | ||
|
|
c8b256ded1 | ||
|
|
dd7cd68b51 | ||
|
|
24ef24558c | ||
|
|
edf8c663f1 | ||
|
|
c24f3d47f7 | ||
|
|
1874214927 | ||
|
|
c29d57622f | ||
|
|
6ae862980d | ||
|
|
c66ecd7517 | ||
|
|
6b8c1fcf47 | ||
|
|
427f173b38 | ||
|
|
b53e3f14d6 | ||
|
|
ffdc9a9585 | ||
|
|
b6a9969aee | ||
|
|
0265849ed3 | ||
|
|
6bc9cdc69d | ||
|
|
3d88dfd98a | ||
|
|
e2e14fd09c | ||
|
|
aa57f78561 | ||
|
|
0f7ba38655 | ||
|
|
1e71d355fc | ||
|
|
0c5c6b3288 | ||
|
|
b98fc8497f | ||
|
|
f002f8de03 | ||
|
|
4fa0201b7e | ||
|
|
8053cbd854 | ||
|
|
c71027c466 | ||
|
|
3cdf471473 | ||
|
|
b807c16395 | ||
|
|
aa59862a70 | ||
|
|
45f071abd9 | ||
|
|
019e148ddb | ||
|
|
b2d08f964d | ||
|
|
d5d74339dd | ||
|
|
9285dc896c | ||
|
|
6c98816f9f | ||
|
|
ed1d828f85 | ||
|
|
50343d5459 | ||
|
|
691a71ac05 | ||
|
|
e6a31a61db | ||
|
|
adc82532eb | ||
|
|
02770d56df | ||
|
|
1366c0b973 | ||
|
|
8bd3a894f1 | ||
|
|
6cc66e404b | ||
|
|
b1cf348abb | ||
|
|
8e11f3864a | ||
|
|
ea45918561 | ||
|
|
80887dd2cc | ||
|
|
b2731eddf3 | ||
|
|
18ac4fe2f3 | ||
|
|
81b7b74fd4 | ||
|
|
38607332b9 | ||
|
|
dd387735d4 | ||
|
|
d9645fa062 | ||
|
|
026f4b9c0e | ||
|
|
a223819dd7 | ||
|
|
bc16c0eec8 | ||
|
|
0d26fd8c5a | ||
|
|
bc2b430ec3 | ||
|
|
6de113b852 | ||
|
|
abe65ad01b | ||
|
|
76f648b4eb | ||
|
|
514bf7e3ed | ||
|
|
24b65d76b5 | ||
|
|
5a81be0756 | ||
|
|
dc2996988b | ||
|
|
c51f94ee87 | ||
|
|
5d34426428 | ||
|
|
4835ca13f5 | ||
|
|
84f46d5ea7 | ||
|
|
19209181dc | ||
|
|
154ec94c58 | ||
|
|
8909e6741a | ||
|
|
45a58a4a03 | ||
|
|
0348b7d915 | ||
|
|
c61c3147ba | ||
|
|
bec9c3a989 | ||
|
|
24ece00e93 | ||
|
|
1becf39269 | ||
|
|
b272109055 | ||
|
|
133f3108f0 | ||
|
|
0ed053b02f | ||
|
|
97d66a8d0c | ||
|
|
ec3c5994d9 | ||
|
|
1f6e384d2b | ||
|
|
5e4099a588 | ||
|
|
68737b9ed6 | ||
|
|
a8813cfb89 | ||
|
|
3b194d3c23 | ||
|
|
5374add5d8 | ||
|
|
8c4c5fa394 | ||
|
|
6fdb944b1e | ||
|
|
d488f809d5 | ||
|
|
5ed3695c48 | ||
|
|
18f39af769 | ||
|
|
1a0a3de2fc | ||
|
|
e536b7bc35 | ||
|
|
85decd7487 | ||
|
|
9fea71a70c | ||
|
|
bde5bf4c76 | ||
|
|
94abab3260 | ||
|
|
76ed136228 | ||
|
|
8d8b20aa47 | ||
|
|
aec0326d40 | ||
|
|
7faca5512a | ||
|
|
ad84272084 | ||
|
|
09e0f594ff | ||
|
|
dd2fbf4424 | ||
|
|
99b12a49c6 | ||
|
|
ea35efe440 | ||
|
|
bf09e740e9 | ||
|
|
60c77cec56 | ||
|
|
0e4a1dddb5 | ||
|
|
1cf18b6e12 | ||
|
|
f9a8be898a | ||
|
|
1521ce5a96 | ||
|
|
f2e62dd197 | ||
|
|
d378630b38 | ||
|
|
d9e6346911 | ||
|
|
238788e0e9 | ||
|
|
68ff828505 | ||
|
|
59447fc12b | ||
|
|
c8033fb6ab | ||
|
|
e33d5b952c | ||
|
|
4345ac2ba2 | ||
|
|
a12b43ce5c | ||
|
|
6885cf1f6d | ||
|
|
00f6fafcfc | ||
|
|
42dc64246c | ||
|
|
fbe303a3cd | ||
|
|
373845450b | ||
|
|
084bbc0bef | ||
|
|
0061fc04b7 | ||
|
|
f6a6410626 | ||
|
|
835be3d329 | ||
|
|
2395093394 | ||
|
|
28209e1c2a | ||
|
|
00562dd1d4 | ||
|
|
0f78d5cbf3 | ||
|
|
431c6de8d2 | ||
|
|
142e15bbcc | ||
|
|
31acc5c607 | ||
|
|
bfa0a26d41 | ||
|
|
93ab9b6a5e | ||
|
|
35e29d46bd | ||
|
|
465da6f818 | ||
|
|
e5f12fddd9 | ||
|
|
4fa9a1303a | ||
|
|
43f349d415 | ||
|
|
02069954de | ||
|
|
2e15875fed | ||
|
|
b34cfb676d | ||
|
|
3064497636 | ||
|
|
dec681fea0 | ||
|
|
523e27ba9a | ||
|
|
e7db76e581 | ||
|
|
689339117a | ||
|
|
b202765be4 | ||
|
|
3bbf3073df | ||
|
|
f46aaa2182 | ||
|
|
a2f33a6c35 | ||
|
|
b6bd6357ed | ||
|
|
c3a5878b1b | ||
|
|
3e4309eba3 | ||
|
|
414f45aa71 | ||
|
|
ebdc76346f | ||
|
|
64bfa955f4 | ||
|
|
612992fa1f | ||
|
|
c02ac56da8 | ||
|
|
9bfb295238 | ||
|
|
cddc22d2b3 | ||
|
|
11ded575d5 | ||
|
|
394cc536a9 | ||
|
|
6bd8cdb9cf | ||
|
|
e20a09f15a | ||
|
|
b89a4af0cf | ||
|
|
a56854af43 | ||
|
|
4a35d78c8d | ||
|
|
26b281271e | ||
|
|
96094cfde2 | ||
|
|
7e26af5476 | ||
|
|
c8dfb784bc | ||
|
|
fd3a5a5afe | ||
|
|
599b3d4c95 | ||
|
|
41719a00e7 | ||
|
|
b5c0f85dca | ||
|
|
7d6d262ed3 | ||
|
|
e21acd73eb | ||
|
|
702f9bc5f1 | ||
|
|
d0ce798881 | ||
|
|
2b1d197047 | ||
|
|
71bc2e6aab | ||
|
|
afb329934a | ||
|
|
1313af45a3 | ||
|
|
dddb327885 | ||
|
|
26b4a37323 | ||
|
|
9dad194130 | ||
|
|
03ad16ea8a | ||
|
|
2fa64b98e3 | ||
|
|
75d7e89cbb | ||
|
|
d73a443484 | ||
|
|
15a9b88fc8 | ||
|
|
03eb7203ec | ||
|
|
e38cd6819b | ||
|
|
d44cfaddf6 | ||
|
|
65225710a8 | ||
|
|
d7f5b16359 | ||
|
|
7185818724 | ||
|
|
868f3349e5 | ||
|
|
d7384e69d9 | ||
|
|
1d5c378343 | ||
|
|
4e1aed9976 | ||
|
|
e2e7996a54 | ||
|
|
df9f9a9f4f | ||
|
|
7553b0da80 | ||
|
|
8f30bf0bef | ||
|
|
8c12174521 |
@@ -11,6 +11,7 @@ ENV/
|
||||
.uv/
|
||||
*.egg-info/
|
||||
dist/
|
||||
!aether-hub/dist/aether-hub
|
||||
build/
|
||||
*.egg
|
||||
|
||||
|
||||
+81
-1
@@ -1,8 +1,16 @@
|
||||
# ==================== 必须配置(启动前) ====================
|
||||
# 以下配置项必须在项目启动前设置
|
||||
|
||||
# 数据库密码
|
||||
# 数据库配置
|
||||
DB_HOST=localhost
|
||||
DB_PORT=5432
|
||||
DB_USER=postgres
|
||||
DB_NAME=aether
|
||||
DB_PASSWORD=your_secure_password_here
|
||||
|
||||
# Redis 配置
|
||||
REDIS_HOST=localhost
|
||||
REDIS_PORT=6379
|
||||
REDIS_PASSWORD=your_redis_password_here
|
||||
|
||||
# JWT密钥(使用 python generate_keys.py 生成)
|
||||
@@ -13,6 +21,10 @@ JWT_SECRET_KEY=change-this-to-a-secure-random-string
|
||||
# 注意:更换此密钥后需要在管理面板重新配置所有 Provider API Key
|
||||
ENCRYPTION_KEY=change-this-to-another-secure-random-string
|
||||
|
||||
# 支付回调共享密钥(公开 /api/payment/callback/* 入口必须携带 x-payment-callback-token)
|
||||
# 建议使用 32+ 位随机字符串
|
||||
PAYMENT_CALLBACK_SECRET=change-this-to-a-secure-callback-secret
|
||||
|
||||
# 管理员账号(仅首次初始化时使用, 创建完成后可在系统内修改密码)
|
||||
ADMIN_EMAIL=[email protected]
|
||||
ADMIN_USERNAME=admin
|
||||
@@ -24,8 +36,76 @@ ADMIN_PASSWORD=admin123456
|
||||
# 应用端口(默认 8084)
|
||||
# APP_PORT=8084
|
||||
|
||||
# 生产部署镜像(deploy.sh 会读取)
|
||||
# APP_IMAGE=ghcr.io/fawney19/aether:latest
|
||||
|
||||
# Gunicorn Worker 数量(默认 2)
|
||||
# Tunnel 请求统一经 Hub 转发,可安全使用多 worker。
|
||||
# 非 Docker 运行时若使用 ProxyNode tunnel,请确保 aether-hub 可达(默认 ws://127.0.0.1:8085)。
|
||||
# GUNICORN_WORKERS=2
|
||||
|
||||
# Gunicorn Max Requests(默认 4000)
|
||||
# Worker 处理指定数量请求后自动重启,防止内存泄漏
|
||||
# max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
|
||||
# MAX_REQUESTS=4000
|
||||
|
||||
# glibc malloc arena 上限(默认 2)
|
||||
# 降低 malloc 内存碎片,减少 gunicorn worker RSS
|
||||
# MALLOC_ARENA_MAX=2
|
||||
|
||||
# HTTP 连接池上限(默认总预算约 200,按 worker 平分)
|
||||
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
|
||||
# HTTP_MAX_CONNECTIONS=100
|
||||
|
||||
# HTTP 保活连接数(默认约为 max_connections 的 30%)
|
||||
# HTTP_KEEPALIVE_CONNECTIONS=30
|
||||
|
||||
# HTTP 代理/Tunnel 客户端空闲清理(默认每 5 分钟扫描,空闲 600 秒即关闭)
|
||||
# HTTP_CLIENT_IDLE_CLEANUP_INTERVAL_MINUTES=5
|
||||
# HTTP_CLIENT_IDLE_CLEANUP_MAX_SECONDS=600
|
||||
|
||||
# curl_cffi session 池上限(默认 20,按 impersonate + proxy 组合缓存)
|
||||
# CURL_CFFI_MAX_SESSIONS=20
|
||||
|
||||
# 流式响应块缓存上限(单位 MB,默认 2)
|
||||
# 说明:
|
||||
# - 这是单个流式请求可保留的“解析后响应块”内存上限,不是全局上限
|
||||
# - 粗略峰值内存 ≈ 并发流数量 × RESPONSE_CHUNKS_MAX_SIZE_MB
|
||||
# 例如:100 并发、2MB 上限,理论峰值约 200MB
|
||||
# - 建议:
|
||||
# - 内存敏感环境:1
|
||||
# - 通用生产环境:2(默认)
|
||||
# - 需要更多调试上下文:4
|
||||
# RESPONSE_CHUNKS_MAX_SIZE_MB=2
|
||||
|
||||
# 流式空闲超时(单位秒,默认 30)
|
||||
# 当流已经开始但连续一段时间没有任何新 chunk 时,提前中断并返回 504,
|
||||
# 避免一直等到 worker 超时(如 300s)
|
||||
# STREAM_IDLE_TIMEOUT_SECONDS=30
|
||||
|
||||
# API Key 前缀(默认 sk)
|
||||
# API_KEY_PREFIX=sk
|
||||
|
||||
# 日志级别(默认 INFO,可选:DEBUG, INFO, WARNING, ERROR)
|
||||
# LOG_LEVEL=INFO
|
||||
|
||||
# CORS 配置(允许跨域的源,多个源用逗号分隔)
|
||||
# 示例: http://localhost:3000,https://example.com
|
||||
# 默认: * (允许所有源)
|
||||
# CORS_ORIGINS=*
|
||||
|
||||
# 启动预热配置(默认启用,降低首请求冷启动延迟)
|
||||
# 是否启用启动期预热任务(默认 true)
|
||||
# STARTUP_WARMUP_ENABLED=true
|
||||
# /readyz 是否等待预热完成(默认 true)
|
||||
# STARTUP_WARMUP_GATE_READINESS=true
|
||||
# 预热时优先 bootstrap 的 provider_type 列表(逗号分隔;留空表示自动探测)
|
||||
# STARTUP_WARMUP_PROVIDER_TYPES=codex,kiro
|
||||
|
||||
# ==================== 计费系统(可选) ====================
|
||||
# Video/Image/Audio 缺失 billing_rule 时是否拒绝请求(默认 false:允许请求但 cost=0 并告警)
|
||||
# BILLING_REQUIRE_RULE=false
|
||||
#
|
||||
# required 维度缺失时是否拒绝请求/标记任务失败(默认 false:cost=0 + 标记 incomplete)
|
||||
# BILLING_STRICT_MODE=false
|
||||
#
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
name: Build aether-hub
|
||||
|
||||
on:
|
||||
push:
|
||||
tags: ['hub-v*']
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
jobs:
|
||||
build:
|
||||
name: ${{ matrix.name }}
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- name: linux-amd64
|
||||
target: x86_64-unknown-linux-gnu
|
||||
use_cross: true
|
||||
- name: linux-arm64
|
||||
target: aarch64-unknown-linux-gnu
|
||||
use_cross: true
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Install Rust toolchain
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
targets: ${{ matrix.target }}
|
||||
|
||||
- name: Rust cache
|
||||
uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
workspaces: aether-hub -> target
|
||||
key: ${{ matrix.target }}
|
||||
|
||||
- name: Install cross
|
||||
if: matrix.use_cross
|
||||
uses: taiki-e/install-action@cross
|
||||
|
||||
- name: Build
|
||||
working-directory: aether-hub
|
||||
shell: bash
|
||||
run: |
|
||||
if [ "${{ matrix.use_cross }}" = "true" ]; then
|
||||
cross build --release --target ${{ matrix.target }}
|
||||
else
|
||||
cargo build --release --target ${{ matrix.target }}
|
||||
fi
|
||||
|
||||
- name: Package
|
||||
shell: bash
|
||||
run: |
|
||||
cd aether-hub/target/${{ matrix.target }}/release
|
||||
chmod +x aether-hub
|
||||
tar czf ../../../../aether-hub-${{ matrix.name }}.tar.gz aether-hub
|
||||
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-artifact@v5
|
||||
with:
|
||||
name: aether-hub-${{ matrix.name }}
|
||||
path: aether-hub-*.tar.gz
|
||||
if-no-files-found: error
|
||||
|
||||
release:
|
||||
needs: build
|
||||
runs-on: ubuntu-latest
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
steps:
|
||||
- name: Download all artifacts
|
||||
uses: actions/download-artifact@v5
|
||||
with:
|
||||
merge-multiple: true
|
||||
path: artifacts
|
||||
|
||||
- name: Generate checksums
|
||||
working-directory: artifacts
|
||||
run: sha256sum aether-hub-* > SHA256SUMS.txt
|
||||
|
||||
- name: Create GitHub Release
|
||||
uses: softprops/action-gh-release@v2
|
||||
with:
|
||||
name: "${{ github.ref_name }}"
|
||||
generate_release_notes: true
|
||||
files: |
|
||||
artifacts/aether-hub-*
|
||||
artifacts/SHA256SUMS.txt
|
||||
fail_on_unmatched_files: true
|
||||
@@ -0,0 +1,227 @@
|
||||
name: Build aether-proxy
|
||||
|
||||
on:
|
||||
push:
|
||||
tags: ['proxy-v*']
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
packages: write
|
||||
|
||||
env:
|
||||
REGISTRY: ghcr.io
|
||||
GHCR_IMAGE: fawney19/aether-proxy
|
||||
DOCKERHUB_IMAGE: fawney19/aether-proxy
|
||||
|
||||
jobs:
|
||||
build:
|
||||
name: ${{ matrix.name }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- name: linux-amd64
|
||||
target: x86_64-unknown-linux-gnu
|
||||
os: ubuntu-latest
|
||||
use_cross: true
|
||||
- name: linux-arm64
|
||||
target: aarch64-unknown-linux-gnu
|
||||
os: ubuntu-latest
|
||||
use_cross: true
|
||||
- name: macos-amd64
|
||||
target: x86_64-apple-darwin
|
||||
os: macos-latest
|
||||
use_cross: false
|
||||
- name: macos-arm64
|
||||
target: aarch64-apple-darwin
|
||||
os: macos-latest
|
||||
use_cross: false
|
||||
- name: windows-amd64
|
||||
target: x86_64-pc-windows-msvc
|
||||
os: windows-latest
|
||||
use_cross: false
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Install Rust toolchain
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
targets: ${{ matrix.target }}
|
||||
|
||||
- name: Rust cache
|
||||
uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
workspaces: aether-proxy -> target
|
||||
key: ${{ matrix.target }}
|
||||
|
||||
- name: Install cross
|
||||
if: matrix.use_cross
|
||||
uses: taiki-e/install-action@cross
|
||||
|
||||
- name: Build
|
||||
working-directory: aether-proxy
|
||||
shell: bash
|
||||
run: |
|
||||
if [ "${{ matrix.use_cross }}" = "true" ]; then
|
||||
cross build --release --target ${{ matrix.target }}
|
||||
else
|
||||
cargo build --release --target ${{ matrix.target }}
|
||||
fi
|
||||
|
||||
- name: Package (Unix)
|
||||
if: runner.os != 'Windows'
|
||||
shell: bash
|
||||
run: |
|
||||
cd aether-proxy/target/${{ matrix.target }}/release
|
||||
chmod +x aether-proxy
|
||||
tar czf ../../../../aether-proxy-${{ matrix.name }}.tar.gz aether-proxy
|
||||
|
||||
- name: Package (Windows)
|
||||
if: runner.os == 'Windows'
|
||||
shell: bash
|
||||
run: |
|
||||
cd aether-proxy/target/${{ matrix.target }}/release
|
||||
7z a ../../../../aether-proxy-${{ matrix.name }}.zip aether-proxy.exe
|
||||
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-artifact@v5
|
||||
with:
|
||||
name: aether-proxy-${{ matrix.name }}
|
||||
path: |
|
||||
aether-proxy-*.tar.gz
|
||||
aether-proxy-*.zip
|
||||
if-no-files-found: error
|
||||
|
||||
release:
|
||||
needs: build
|
||||
runs-on: ubuntu-latest
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
steps:
|
||||
- name: Download all artifacts
|
||||
uses: actions/download-artifact@v5
|
||||
with:
|
||||
merge-multiple: true
|
||||
path: artifacts
|
||||
|
||||
- name: Generate checksums
|
||||
working-directory: artifacts
|
||||
run: sha256sum aether-proxy-* > SHA256SUMS.txt
|
||||
|
||||
- name: Create GitHub Release
|
||||
uses: softprops/action-gh-release@v2
|
||||
with:
|
||||
name: "${{ github.ref_name }}"
|
||||
generate_release_notes: true
|
||||
files: |
|
||||
artifacts/aether-proxy-*
|
||||
artifacts/SHA256SUMS.txt
|
||||
fail_on_unmatched_files: true
|
||||
|
||||
docker:
|
||||
needs: build
|
||||
runs-on: ubuntu-latest
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
steps:
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Download Linux artifacts
|
||||
uses: actions/download-artifact@v5
|
||||
with:
|
||||
pattern: aether-proxy-linux-*
|
||||
merge-multiple: true
|
||||
path: artifacts
|
||||
|
||||
- name: Prepare binaries
|
||||
run: |
|
||||
mkdir -p aether-proxy/build/linux-amd64 aether-proxy/build/linux-arm64
|
||||
tar xzf artifacts/aether-proxy-linux-amd64.tar.gz -C aether-proxy/build/linux-amd64
|
||||
tar xzf artifacts/aether-proxy-linux-arm64.tar.gz -C aether-proxy/build/linux-arm64
|
||||
|
||||
- name: Set up QEMU
|
||||
uses: docker/setup-qemu-action@v3
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Log in to GHCR
|
||||
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: Extract metadata
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: |
|
||||
${{ env.REGISTRY }}/${{ env.GHCR_IMAGE }}
|
||||
docker.io/${{ env.DOCKERHUB_IMAGE }}
|
||||
tags: |
|
||||
type=match,pattern=proxy-v(.*),group=1
|
||||
type=match,pattern=proxy-v(\d+\.\d+),group=1
|
||||
type=sha,prefix=
|
||||
flavor: |
|
||||
latest=auto
|
||||
|
||||
- name: Build and push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: ./aether-proxy
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
platforms: linux/amd64,linux/arm64
|
||||
|
||||
update-readme:
|
||||
needs: release
|
||||
runs-on: ubuntu-latest
|
||||
if: startsWith(github.ref, 'refs/tags/')
|
||||
steps:
|
||||
- uses: actions/checkout@v5
|
||||
with:
|
||||
ref: master
|
||||
|
||||
- name: Update README download links
|
||||
env:
|
||||
TAG: ${{ github.ref_name }}
|
||||
run: |
|
||||
VERSION="${TAG#proxy-v}"
|
||||
BASE="https://github.com/fawney19/Aether/releases/download/${TAG}"
|
||||
cd aether-proxy
|
||||
|
||||
TABLE="| Platform | Download |\n|----------|----------|\n"
|
||||
TABLE+="| Linux x86_64 | [aether-proxy-linux-amd64.tar.gz](${BASE}/aether-proxy-linux-amd64.tar.gz) |\n"
|
||||
TABLE+="| Linux ARM64 | [aether-proxy-linux-arm64.tar.gz](${BASE}/aether-proxy-linux-arm64.tar.gz) |\n"
|
||||
TABLE+="| macOS x86_64 | [aether-proxy-macos-amd64.tar.gz](${BASE}/aether-proxy-macos-amd64.tar.gz) |\n"
|
||||
TABLE+="| macOS ARM64 | [aether-proxy-macos-arm64.tar.gz](${BASE}/aether-proxy-macos-arm64.tar.gz) |\n"
|
||||
TABLE+="| Windows x86_64 | [aether-proxy-windows-amd64.zip](${BASE}/aether-proxy-windows-amd64.zip) |"
|
||||
|
||||
# Replace content between markers
|
||||
if grep -q '<!-- DOWNLOAD_TABLE_START -->' README.md; then
|
||||
awk -v table="$TABLE" '
|
||||
/<!-- DOWNLOAD_TABLE_START -->/ { print; printf "%s\n", table; skip=1; next }
|
||||
/<!-- DOWNLOAD_TABLE_END -->/ { skip=0 }
|
||||
!skip { print }
|
||||
' README.md > README.tmp && mv README.tmp README.md
|
||||
fi
|
||||
|
||||
- name: Commit and push
|
||||
run: |
|
||||
cd aether-proxy
|
||||
git config user.name "github-actions[bot]"
|
||||
git config user.email "github-actions[bot]@users.noreply.github.com"
|
||||
git add README.md
|
||||
git diff --cached --quiet && exit 0
|
||||
TAG="${GITHUB_REF#refs/tags/}"
|
||||
git commit -m "chore(proxy): update download links for ${TAG}"
|
||||
git push
|
||||
@@ -18,12 +18,12 @@ jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Setup Node.js
|
||||
uses: actions/setup-node@v4
|
||||
uses: actions/setup-node@v5
|
||||
with:
|
||||
node-version: '20'
|
||||
node-version: '22'
|
||||
cache: 'npm'
|
||||
cache-dependency-path: frontend/package-lock.json
|
||||
|
||||
@@ -41,7 +41,7 @@ jobs:
|
||||
run: cp frontend/dist/index.html frontend/dist/404.html
|
||||
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
uses: actions/configure-pages@v5
|
||||
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
|
||||
@@ -15,16 +15,22 @@ env:
|
||||
REGISTRY: ghcr.io
|
||||
BASE_IMAGE_NAME: fawney19/aether-base
|
||||
APP_IMAGE_NAME: fawney19/aether
|
||||
# Files that affect base image - used for hash calculation
|
||||
BASE_FILES: "Dockerfile.base pyproject.toml frontend/package.json frontend/package-lock.json"
|
||||
GITHUB_REPO: fawney19/Aether
|
||||
# Base image hash inputs:
|
||||
# - Dockerfile.base
|
||||
# - pyproject.toml (dependency fingerprint only; ignores tool/optional deps)
|
||||
# - frontend/package-lock.json
|
||||
|
||||
jobs:
|
||||
check-base-changes:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
packages: read
|
||||
outputs:
|
||||
base_changed: ${{ steps.check.outputs.base_changed }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Log in to Container Registry
|
||||
uses: docker/login-action@v3
|
||||
@@ -41,15 +47,45 @@ jobs:
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Calculate current hash of base-related files
|
||||
CURRENT_HASH=$(cat ${{ env.BASE_FILES }} 2>/dev/null | sha256sum | cut -d' ' -f1)
|
||||
echo "Current base files hash: $CURRENT_HASH"
|
||||
# Calculate current hash of base-related inputs (dependency-only fingerprint)
|
||||
PY_FINGERPRINT=$(python3 - <<'PY'
|
||||
import json
|
||||
import pathlib
|
||||
import tomllib
|
||||
|
||||
data = tomllib.loads(pathlib.Path("pyproject.toml").read_text("utf-8"))
|
||||
project = data.get("project") or {}
|
||||
build = data.get("build-system") or {}
|
||||
|
||||
fingerprint = {
|
||||
"requires-python": project.get("requires-python"),
|
||||
"dependencies": sorted(project.get("dependencies") or []),
|
||||
"build-backend": build.get("build-backend"),
|
||||
"build-requires": sorted(build.get("requires") or []),
|
||||
}
|
||||
|
||||
print(json.dumps(fingerprint, sort_keys=True, separators=(",", ":")))
|
||||
PY
|
||||
)
|
||||
|
||||
CURRENT_HASH=$(
|
||||
(
|
||||
cat Dockerfile.base
|
||||
printf '%s\n' "$PY_FINGERPRINT"
|
||||
cat frontend/package-lock.json
|
||||
) | sha256sum | cut -d' ' -f1
|
||||
)
|
||||
echo "Current base hash: $CURRENT_HASH"
|
||||
|
||||
# Try to get hash label from remote image config
|
||||
# Pull the image config and extract labels
|
||||
REMOTE_HASH=""
|
||||
if docker pull ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest 2>/dev/null; then
|
||||
if docker pull ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest; then
|
||||
REMOTE_HASH=$(docker inspect ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest --format '{{ index .Config.Labels "org.opencontainers.image.base.hash" }}' 2>/dev/null) || true
|
||||
else
|
||||
echo "WARN: failed to pull remote base image; forcing base rebuild."
|
||||
echo "base_changed=true" >> $GITHUB_OUTPUT
|
||||
exit 0
|
||||
fi
|
||||
|
||||
if [ -z "$REMOTE_HASH" ] || [ "$REMOTE_HASH" == "<no value>" ]; then
|
||||
@@ -72,7 +108,7 @@ jobs:
|
||||
contents: read
|
||||
packages: write
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
@@ -87,7 +123,33 @@ jobs:
|
||||
- name: Calculate base files hash
|
||||
id: hash
|
||||
run: |
|
||||
HASH=$(cat ${{ env.BASE_FILES }} 2>/dev/null | sha256sum | cut -d' ' -f1)
|
||||
PY_FINGERPRINT=$(python3 - <<'PY'
|
||||
import json
|
||||
import pathlib
|
||||
import tomllib
|
||||
|
||||
data = tomllib.loads(pathlib.Path("pyproject.toml").read_text("utf-8"))
|
||||
project = data.get("project") or {}
|
||||
build = data.get("build-system") or {}
|
||||
|
||||
fingerprint = {
|
||||
"requires-python": project.get("requires-python"),
|
||||
"dependencies": sorted(project.get("dependencies") or []),
|
||||
"build-backend": build.get("build-backend"),
|
||||
"build-requires": sorted(build.get("requires") or []),
|
||||
}
|
||||
|
||||
print(json.dumps(fingerprint, sort_keys=True, separators=(",", ":")))
|
||||
PY
|
||||
)
|
||||
|
||||
HASH=$(
|
||||
(
|
||||
cat Dockerfile.base
|
||||
printf '%s\n' "$PY_FINGERPRINT"
|
||||
cat frontend/package-lock.json
|
||||
) | sha256sum | cut -d' ' -f1
|
||||
)
|
||||
echo "hash=$HASH" >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Extract metadata for base image
|
||||
@@ -102,26 +164,47 @@ jobs:
|
||||
org.opencontainers.image.base.hash=${{ steps.hash.outputs.hash }}
|
||||
|
||||
- name: Build and push base image
|
||||
uses: docker/build-push-action@v5
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: ./Dockerfile.base
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
cache-from: type=gha,scope=base
|
||||
cache-to: type=gha,mode=max,scope=base
|
||||
platforms: linux/amd64,linux/arm64
|
||||
|
||||
download-hub:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
outputs:
|
||||
hub_tag: ${{ steps.hub-tag.outputs.tag }}
|
||||
steps:
|
||||
- name: Get latest hub release tag
|
||||
id: hub-tag
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
TAG=$(gh release list --repo "${{ env.GITHUB_REPO }}" --limit 50 --json tagName,isDraft,isPrerelease \
|
||||
--jq '[.[] | select(.tagName | startswith("hub-v")) | select(.isDraft == false and .isPrerelease == false)] | .[0].tagName')
|
||||
if [ -z "$TAG" ] || [ "$TAG" = "null" ]; then
|
||||
echo "No hub release found"
|
||||
exit 1
|
||||
fi
|
||||
echo "tag=$TAG" >> $GITHUB_OUTPUT
|
||||
echo "Hub release tag: $TAG"
|
||||
|
||||
build-app:
|
||||
needs: [check-base-changes, build-base]
|
||||
if: always() && (needs.build-base.result == 'success' || needs.build-base.result == 'skipped')
|
||||
needs: [check-base-changes, build-base, download-hub]
|
||||
if: always() && (needs.build-base.result == 'success' || needs.build-base.result == 'skipped') && needs.download-hub.result == 'success'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/checkout@v5
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
@@ -133,31 +216,112 @@ jobs:
|
||||
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: Extract metadata for app image
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }}
|
||||
images: |
|
||||
${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }}
|
||||
docker.io/fawney19/aether
|
||||
tags: |
|
||||
type=raw,value=latest,enable={{is_default_branch}}
|
||||
type=ref,event=branch
|
||||
type=ref,event=pr
|
||||
type=semver,pattern={{version}}
|
||||
type=semver,pattern={{major}}.{{minor}}
|
||||
type=raw,value=pre,enable=${{ contains(github.ref, '-') }}
|
||||
type=raw,value=fix,enable=${{ contains(github.ref, '-fix') }}
|
||||
type=sha,prefix=
|
||||
flavor: |
|
||||
latest=auto
|
||||
|
||||
- name: Extract version from tag
|
||||
id: version
|
||||
run: |
|
||||
# 从 tag 提取版本号,如 v0.2.5 -> 0.2.5
|
||||
VERSION="${GITHUB_REF#refs/tags/v}"
|
||||
if [ "$VERSION" = "$GITHUB_REF" ]; then
|
||||
# 不是 tag 触发,使用 git describe
|
||||
VERSION=$(git describe --tags --always | sed 's/^v//')
|
||||
fi
|
||||
echo "version=$VERSION" >> $GITHUB_OUTPUT
|
||||
echo "Extracted version: $VERSION"
|
||||
|
||||
- name: Update Dockerfile.app to use registry base image
|
||||
run: |
|
||||
sed -i "s|FROM aether-base:latest AS builder|FROM ${{ env.REGISTRY }}/${{ env.BASE_IMAGE_NAME }}:latest AS builder|g" Dockerfile.app
|
||||
|
||||
- name: Build and push app image
|
||||
uses: docker/build-push-action@v5
|
||||
- name: Generate version file
|
||||
run: |
|
||||
# 生成 _version.py 文件
|
||||
cat > src/_version.py << EOF
|
||||
# Auto-generated by CI
|
||||
__version__ = '${{ steps.version.outputs.version }}'
|
||||
__version_tuple__ = tuple(int(x) for x in '${{ steps.version.outputs.version }}'.split('.') if x.isdigit())
|
||||
version = __version__
|
||||
version_tuple = __version_tuple__
|
||||
EOF
|
||||
|
||||
- name: Resolve hub release for build args
|
||||
run: |
|
||||
echo "Hub release tag: ${{ needs.download-hub.outputs.hub_tag }}"
|
||||
|
||||
- name: Build and push app image (amd64)
|
||||
id: build-amd64
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: ./Dockerfile.app
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
platforms: linux/amd64,linux/arm64
|
||||
no-cache-filters: builder
|
||||
cache-from: type=gha,scope=app-amd64
|
||||
cache-to: type=gha,mode=min,scope=app-amd64
|
||||
build-args: |
|
||||
HUB_RELEASE_REPO=${{ env.GITHUB_REPO }}
|
||||
HUB_TAG=${{ needs.download-hub.outputs.hub_tag }}
|
||||
platforms: linux/amd64
|
||||
outputs: type=image,"name=${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }},docker.io/fawney19/aether",push-by-digest=true,name-canonical=true,push=true
|
||||
|
||||
- name: Build and push app image (arm64)
|
||||
id: build-arm64
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: ./Dockerfile.app
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
no-cache-filters: builder
|
||||
cache-from: type=gha,scope=app-arm64
|
||||
cache-to: type=gha,mode=min,scope=app-arm64
|
||||
build-args: |
|
||||
HUB_RELEASE_REPO=${{ env.GITHUB_REPO }}
|
||||
HUB_TAG=${{ needs.download-hub.outputs.hub_tag }}
|
||||
platforms: linux/arm64
|
||||
outputs: type=image,"name=${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }},docker.io/fawney19/aether",push-by-digest=true,name-canonical=true,push=true
|
||||
|
||||
- name: Create multi-arch manifest and push
|
||||
run: |
|
||||
# Extract digests
|
||||
AMD64_DIGEST="${{ steps.build-amd64.outputs.digest }}"
|
||||
ARM64_DIGEST="${{ steps.build-arm64.outputs.digest }}"
|
||||
echo "amd64 digest: $AMD64_DIGEST"
|
||||
echo "arm64 digest: $ARM64_DIGEST"
|
||||
|
||||
# For each tag, create multi-arch manifest on each registry
|
||||
TAGS=$(echo "${{ steps.meta.outputs.tags }}" | tr '\n' ' ')
|
||||
for FULL_TAG in $TAGS; do
|
||||
# Determine which registry this tag belongs to
|
||||
if [[ "$FULL_TAG" == ghcr.io/* ]]; then
|
||||
REPO="${{ env.REGISTRY }}/${{ env.APP_IMAGE_NAME }}"
|
||||
elif [[ "$FULL_TAG" == docker.io/* ]]; then
|
||||
REPO="docker.io/fawney19/aether"
|
||||
else
|
||||
continue
|
||||
fi
|
||||
echo "Creating manifest for $FULL_TAG"
|
||||
docker buildx imagetools create -t "$FULL_TAG" \
|
||||
"$REPO@$AMD64_DIGEST" \
|
||||
"$REPO@$ARM64_DIGEST"
|
||||
done
|
||||
|
||||
+16
@@ -2,9 +2,11 @@
|
||||
# Edit at https://www.toptal.com/developers/gitignore?templates=python
|
||||
|
||||
# AI Assistant Configuration
|
||||
.codex/
|
||||
.claude/
|
||||
.serena/
|
||||
.gemini*/
|
||||
.plans
|
||||
|
||||
### Python ###
|
||||
*.db
|
||||
@@ -202,6 +204,7 @@ logs/
|
||||
|
||||
# Git backup
|
||||
.git.backup/
|
||||
.worktrees/
|
||||
|
||||
# Database backups
|
||||
backups/
|
||||
@@ -219,8 +222,21 @@ frontend/public/*-firework.svg
|
||||
# Debug and experimental files
|
||||
debug_*.html
|
||||
extracted_*.ts
|
||||
test.py
|
||||
|
||||
# Deploy script cache
|
||||
.deps-hash
|
||||
.code-hash
|
||||
.migration-hash
|
||||
.hub-hash
|
||||
|
||||
# Hub prebuilt binaries
|
||||
aether-hub/dist/
|
||||
|
||||
# Version file (auto-generated by hatch-vcs)
|
||||
src/_version.py
|
||||
|
||||
# Analysis folder (third-party code for reference)
|
||||
analysis/
|
||||
new-api/
|
||||
/aether-proxy/target/
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
3.13
|
||||
+249
-92
@@ -1,134 +1,291 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# 运行镜像:从 base 提取产物到精简运行时
|
||||
# 构建命令: docker build -f Dockerfile.app -t aether-app:latest .
|
||||
# 用于 GitHub Actions CI(官方源)
|
||||
|
||||
FROM aether-base:latest AS builder
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 复制前端源码并构建
|
||||
# 复制前端源码并构建(CI 通过 no-cache-filters=builder 确保每次重建)
|
||||
COPY frontend/ ./frontend/
|
||||
RUN cd frontend && npm run build
|
||||
|
||||
# ==================== 运行时镜像 ====================
|
||||
FROM python:3.12-slim
|
||||
|
||||
FROM python:3.13-slim
|
||||
WORKDIR /app
|
||||
|
||||
# 运行时依赖(无 gcc/nodejs/npm)
|
||||
RUN apt-get update && apt-get install -y \
|
||||
ARG HUB_RELEASE_REPO=fawney19/Aether
|
||||
ARG HUB_TAG
|
||||
ARG TARGETARCH
|
||||
ARG GITHUB_TOKEN
|
||||
|
||||
# 运行时依赖(无 gcc/nodejs/npm,使用 BuildKit 缓存加速)
|
||||
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
--mount=type=cache,target=/var/lib/apt,sharing=locked \
|
||||
apt-get update && apt-get install -y --no-install-recommends \
|
||||
nginx \
|
||||
supervisor \
|
||||
libpq5 \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
libjemalloc2
|
||||
RUN set -eux; \
|
||||
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
|
||||
[ -n "$jemalloc_path" ]; \
|
||||
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
|
||||
# 从 base 镜像复制 Python 包
|
||||
COPY --from=builder /usr/local/lib/python3.12/site-packages /usr/local/lib/python3.12/site-packages
|
||||
|
||||
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
||||
# 只复制需要的 Python 可执行文件
|
||||
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
||||
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
|
||||
COPY --from=builder /usr/local/bin/alembic /usr/local/bin/
|
||||
|
||||
# Hub 预编译二进制(构建时从 GitHub Release 下载)
|
||||
# GITHUB_TOKEN 可选:未认证 API 限流 60 次/小时,认证后 5000 次/小时
|
||||
RUN set -eux; \
|
||||
auth_header=""; \
|
||||
if [ -n "${GITHUB_TOKEN:-}" ]; then \
|
||||
auth_header="Authorization: token ${GITHUB_TOKEN}"; \
|
||||
fi; \
|
||||
tag="${HUB_TAG:-}"; \
|
||||
if [ -z "$tag" ]; then \
|
||||
tag="$(curl -sL ${auth_header:+-H "$auth_header"} "https://api.github.com/repos/${HUB_RELEASE_REPO}/releases" | python3 -c "import json,sys;print(next((r['tag_name'] for r in json.load(sys.stdin) if r.get('tag_name','').startswith('hub-v') and not r.get('draft') and not r.get('prerelease')),''))")"; \
|
||||
fi; \
|
||||
if [ -z "$tag" ]; then \
|
||||
echo "Failed to resolve hub release tag"; \
|
||||
exit 1; \
|
||||
fi; \
|
||||
arch="${TARGETARCH:-}"; \
|
||||
if [ -z "$arch" ]; then \
|
||||
arch="$(dpkg --print-architecture)"; \
|
||||
fi; \
|
||||
case "$arch" in \
|
||||
amd64|arm64) ;; \
|
||||
x86_64) arch="amd64" ;; \
|
||||
aarch64) arch="arm64" ;; \
|
||||
*) echo "Unsupported architecture: $arch"; exit 1 ;; \
|
||||
esac; \
|
||||
echo "Using Hub release tag: $tag"; \
|
||||
url="https://github.com/${HUB_RELEASE_REPO}/releases/download/${tag}/aether-hub-linux-${arch}.tar.gz"; \
|
||||
curl -L --fail -o /tmp/aether-hub.tar.gz "$url"; \
|
||||
tar xzf /tmp/aether-hub.tar.gz -C /usr/local/bin; \
|
||||
chmod +x /usr/local/bin/aether-hub; \
|
||||
rm -f /tmp/aether-hub.tar.gz
|
||||
# 从 builder 阶段复制前端构建产物
|
||||
COPY --from=builder /app/frontend/dist /usr/share/nginx/html
|
||||
|
||||
RUN chmod -R 755 /usr/share/nginx/html
|
||||
# 复制后端代码
|
||||
COPY src/ ./src/
|
||||
COPY alembic.ini ./
|
||||
COPY alembic/ ./alembic/
|
||||
|
||||
COPY gunicorn_conf.py ./
|
||||
# Nginx 配置模板
|
||||
# 策略:白名单后端路由 → 后端代理,其余全部 → 前端 SPA(index.html)
|
||||
# 智能处理 IP:有外层代理头就透传,没有就用直连 IP
|
||||
RUN printf '%s\n' \
|
||||
'server {' \
|
||||
' listen 80;' \
|
||||
' server_name _;' \
|
||||
' root /usr/share/nginx/html;' \
|
||||
' index index.html;' \
|
||||
' client_max_body_size 100M;' \
|
||||
'' \
|
||||
' location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {' \
|
||||
' expires 1y;' \
|
||||
' add_header Cache-Control "public, no-transform";' \
|
||||
' try_files $uri =404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location ~ ^/(src|node_modules)/ {' \
|
||||
' deny all;' \
|
||||
' return 404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location ~ ^/(dashboard|admin|login)(/|$) {' \
|
||||
' try_files $uri $uri/ /index.html;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location / {' \
|
||||
' try_files $uri $uri/ @backend;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location @backend {' \
|
||||
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $remote_addr;' \
|
||||
' proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Connection "";' \
|
||||
' proxy_set_header Accept $http_accept;' \
|
||||
' proxy_set_header Content-Type $content_type;' \
|
||||
' proxy_set_header Authorization $http_authorization;' \
|
||||
' proxy_set_header X-Api-Key $http_x_api_key;' \
|
||||
' proxy_buffering off;' \
|
||||
' proxy_cache off;' \
|
||||
' proxy_request_buffering off;' \
|
||||
' chunked_transfer_encoding on;' \
|
||||
' gzip off;' \
|
||||
' add_header X-Accel-Buffering no;' \
|
||||
' proxy_connect_timeout 600s;' \
|
||||
' proxy_send_timeout 600s;' \
|
||||
' proxy_read_timeout 600s;' \
|
||||
' }' \
|
||||
'}' > /etc/nginx/sites-available/default.template
|
||||
|
||||
'map $http_x_real_ip $real_ip {' \
|
||||
' default $http_x_real_ip;' \
|
||||
' "" $remote_addr;' \
|
||||
'}' \
|
||||
'' \
|
||||
'map $http_x_forwarded_for $forwarded_for {' \
|
||||
' default $http_x_forwarded_for;' \
|
||||
' "" $remote_addr;' \
|
||||
'}' \
|
||||
'' \
|
||||
'map $http_upgrade $connection_upgrade {' \
|
||||
' default upgrade;' \
|
||||
' "" "";' \
|
||||
'}' \
|
||||
'' \
|
||||
'server {' \
|
||||
' listen 80;' \
|
||||
' server_name _;' \
|
||||
' root /usr/share/nginx/html;' \
|
||||
' index index.html;' \
|
||||
' client_max_body_size 100M;' \
|
||||
'' \
|
||||
' # gzip 压缩配置(对 base64 图片等非流式响应有效)' \
|
||||
' gzip on;' \
|
||||
' gzip_min_length 256;' \
|
||||
' gzip_comp_level 5;' \
|
||||
' gzip_vary on;' \
|
||||
' gzip_proxied any;' \
|
||||
' gzip_types application/json text/plain text/css text/javascript application/javascript application/octet-stream;' \
|
||||
' gzip_disable "msie6";' \
|
||||
'' \
|
||||
' # 静态资源:长期缓存' \
|
||||
' location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {' \
|
||||
' expires 1y;' \
|
||||
' add_header Cache-Control "public, no-transform";' \
|
||||
' try_files $uri =404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # 安全:阻止访问源码目录' \
|
||||
' location ~ ^/(src|node_modules)/ {' \
|
||||
' deny all;' \
|
||||
' return 404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # WebSocket 隧道端点(aether-proxy tunnel 模式)' \
|
||||
' location = /api/internal/proxy-tunnel {' \
|
||||
' proxy_pass http://127.0.0.1:8085/proxy;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Upgrade $http_upgrade;' \
|
||||
' proxy_set_header Connection "upgrade";' \
|
||||
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
|
||||
' proxy_hide_header CF-Connecting-IP;' \
|
||||
' proxy_hide_header CF-IPCountry;' \
|
||||
' proxy_hide_header CF-Ray;' \
|
||||
' proxy_hide_header CF-Visitor;' \
|
||||
' proxy_hide_header CDN-Loop;' \
|
||||
' proxy_hide_header True-Client-IP;' \
|
||||
' proxy_hide_header CF-Worker;' \
|
||||
' proxy_hide_header CF-EW-Via;' \
|
||||
' proxy_hide_header CF-Warp-Tag-ID;' \
|
||||
' proxy_set_header CF-Connecting-IP "";' \
|
||||
' proxy_set_header CF-IPCountry "";' \
|
||||
' proxy_set_header CF-Ray "";' \
|
||||
' proxy_set_header CF-Visitor "";' \
|
||||
' proxy_set_header CDN-Loop "";' \
|
||||
' proxy_set_header True-Client-IP "";' \
|
||||
' proxy_set_header CF-Worker "";' \
|
||||
' proxy_set_header CF-EW-Via "";' \
|
||||
' proxy_set_header CF-Warp-Tag-ID "";' \
|
||||
' proxy_read_timeout 86400s;' \
|
||||
' proxy_send_timeout 86400s;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # 后端 API 路由(白名单)→ 代理到后端' \
|
||||
' location ~ ^/(api|v1|v1beta|upload|health)(/|$) {' \
|
||||
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Upgrade $http_upgrade;' \
|
||||
' proxy_set_header Connection $connection_upgrade;' \
|
||||
' proxy_set_header Accept $http_accept;' \
|
||||
' proxy_set_header Content-Type $content_type;' \
|
||||
' proxy_set_header Authorization $http_authorization;' \
|
||||
' proxy_set_header X-Api-Key $http_x_api_key;' \
|
||||
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
|
||||
' proxy_hide_header CF-Connecting-IP;' \
|
||||
' proxy_hide_header CF-IPCountry;' \
|
||||
' proxy_hide_header CF-Ray;' \
|
||||
' proxy_hide_header CF-Visitor;' \
|
||||
' proxy_hide_header CDN-Loop;' \
|
||||
' proxy_hide_header True-Client-IP;' \
|
||||
' proxy_hide_header CF-Worker;' \
|
||||
' proxy_hide_header CF-EW-Via;' \
|
||||
' proxy_hide_header CF-Warp-Tag-ID;' \
|
||||
' proxy_set_header CF-Connecting-IP "";' \
|
||||
' proxy_set_header CF-IPCountry "";' \
|
||||
' proxy_set_header CF-Ray "";' \
|
||||
' proxy_set_header CF-Visitor "";' \
|
||||
' proxy_set_header CDN-Loop "";' \
|
||||
' proxy_set_header True-Client-IP "";' \
|
||||
' proxy_set_header CF-Worker "";' \
|
||||
' proxy_set_header CF-EW-Via "";' \
|
||||
' proxy_set_header CF-Warp-Tag-ID "";' \
|
||||
' proxy_buffering off;' \
|
||||
' proxy_cache off;' \
|
||||
' proxy_request_buffering off;' \
|
||||
' chunked_transfer_encoding on;' \
|
||||
' gzip off;' \
|
||||
' add_header X-Accel-Buffering no;' \
|
||||
' proxy_connect_timeout 60s;' \
|
||||
' proxy_send_timeout 3600s;' \
|
||||
' proxy_read_timeout 3600s;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # API 文档路由 → 代理到后端' \
|
||||
' location ~ ^/(docs|redoc|openapi\\.json)$ {' \
|
||||
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
|
||||
' proxy_hide_header CF-Connecting-IP;' \
|
||||
' proxy_hide_header CF-IPCountry;' \
|
||||
' proxy_hide_header CF-Ray;' \
|
||||
' proxy_hide_header CF-Visitor;' \
|
||||
' proxy_hide_header CDN-Loop;' \
|
||||
' proxy_hide_header True-Client-IP;' \
|
||||
' proxy_hide_header CF-Worker;' \
|
||||
' proxy_hide_header CF-EW-Via;' \
|
||||
' proxy_hide_header CF-Warp-Tag-ID;' \
|
||||
' proxy_set_header CF-Connecting-IP "";' \
|
||||
' proxy_set_header CF-IPCountry "";' \
|
||||
' proxy_set_header CF-Ray "";' \
|
||||
' proxy_set_header CF-Visitor "";' \
|
||||
' proxy_set_header CDN-Loop "";' \
|
||||
' proxy_set_header True-Client-IP "";' \
|
||||
' proxy_set_header CF-Worker "";' \
|
||||
' proxy_set_header CF-EW-Via "";' \
|
||||
' proxy_set_header CF-Warp-Tag-ID "";' \
|
||||
' }' \
|
||||
'' \
|
||||
' # 所有其他路由 → 前端 SPA(先尝试静态文件,再回退到 index.html)' \
|
||||
' location / {' \
|
||||
' try_files $uri $uri/ /index.html;' \
|
||||
' }' \
|
||||
'}' > /etc/nginx/sites-available/default.template
|
||||
# Supervisor 配置
|
||||
RUN printf '%s\n' \
|
||||
'[supervisord]' \
|
||||
'nodaemon=true' \
|
||||
'logfile=/var/log/supervisor/supervisord.log' \
|
||||
'pidfile=/var/run/supervisord.pid' \
|
||||
'' \
|
||||
'[program:nginx]' \
|
||||
'command=/bin/bash -c "sed \"s/PORT_PLACEHOLDER/${PORT:-8084}/g\" /etc/nginx/sites-available/default.template > /etc/nginx/sites-available/default && /usr/sbin/nginx -g \"daemon off;\""' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/var/log/nginx/access.log' \
|
||||
'stderr_logfile=/var/log/nginx/error.log' \
|
||||
'' \
|
||||
'[program:app]' \
|
||||
'command=gunicorn src.main:app -w %(ENV_GUNICORN_WORKERS)s -k uvicorn.workers.UvicornWorker --bind 0.0.0.0:%(ENV_PORT)s --timeout 120 --access-logfile - --error-logfile - --log-level info' \
|
||||
'directory=/app' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/dev/stdout' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'stderr_logfile_maxbytes=0' \
|
||||
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' > /etc/supervisor/conf.d/supervisord.conf
|
||||
|
||||
'[supervisord]' \
|
||||
'nodaemon=true' \
|
||||
'logfile=/var/log/supervisor/supervisord.log' \
|
||||
'pidfile=/var/run/supervisord.pid' \
|
||||
'' \
|
||||
'[program:nginx]' \
|
||||
'command=/bin/bash -c "sed \"s/PORT_PLACEHOLDER/8084/g\" /etc/nginx/sites-available/default.template > /etc/nginx/sites-available/default && /usr/sbin/nginx -g \"daemon off;\""' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/var/log/nginx/access.log' \
|
||||
'stderr_logfile=/var/log/nginx/error.log' \
|
||||
'' \
|
||||
'[program:app]' \
|
||||
'command=/bin/bash -c "MAX_REQUESTS_JITTER=$((${MAX_REQUESTS:-50000}/20)); exec gunicorn src.main:app -c gunicorn_conf.py --preload -w %(ENV_GUNICORN_WORKERS)s -k uvicorn.workers.UvicornWorker --bind 127.0.0.1:8084 --max-requests ${MAX_REQUESTS:-50000} --max-requests-jitter $MAX_REQUESTS_JITTER --access-logfile - --error-logfile - --log-level info"' \
|
||||
'directory=/app' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/dev/stdout' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'stderr_logfile_maxbytes=0' \
|
||||
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
|
||||
'' \
|
||||
'[program:tunnel-hub]' \
|
||||
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/dev/stdout' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'stderr_logfile_maxbytes=0' > /etc/supervisor/conf.d/supervisord.conf
|
||||
# 创建目录
|
||||
RUN mkdir -p /var/log/supervisor /app/logs /app/data
|
||||
|
||||
# 入口脚本(启动前执行迁移)
|
||||
COPY entrypoint.sh /entrypoint.sh
|
||||
RUN chmod +x /entrypoint.sh
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONIOENCODING=utf-8 \
|
||||
LANG=C.UTF-8 \
|
||||
LC_ALL=C.UTF-8 \
|
||||
PORT=8084
|
||||
|
||||
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
|
||||
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
|
||||
PORT=8084 \
|
||||
GUNICORN_WORKERS=2 \
|
||||
MAX_REQUESTS=4000
|
||||
EXPOSE 80
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost/health || exit 1
|
||||
|
||||
ENTRYPOINT ["/entrypoint.sh"]
|
||||
CMD ["/usr/bin/supervisord", "-c", "/etc/supervisor/conf.d/supervisord.conf"]
|
||||
|
||||
+263
-78
@@ -1,6 +1,8 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# 运行镜像:从 base 提取产物到精简运行时(国内镜像源版本)
|
||||
# 构建命令: docker build -f Dockerfile.app.local -t aether-app:latest .
|
||||
# 用于本地/国内服务器部署
|
||||
|
||||
FROM aether-base:latest AS builder
|
||||
|
||||
WORKDIR /app
|
||||
@@ -10,126 +12,309 @@ COPY frontend/ ./frontend/
|
||||
RUN cd frontend && npm run build
|
||||
|
||||
# ==================== 运行时镜像 ====================
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.13-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 运行时依赖(使用清华镜像源)
|
||||
RUN sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
|
||||
apt-get update && apt-get install -y \
|
||||
ARG HUB_RELEASE_REPO=fawney19/Aether
|
||||
ARG HUB_TAG
|
||||
ARG TARGETARCH
|
||||
ARG GITHUB_TOKEN
|
||||
# GitHub 下载镜像前缀,国内构建时传入可用的镜像加速地址
|
||||
# 用法: --build-arg GITHUB_MIRROR=https://ghfast.top
|
||||
# 或: --build-arg GITHUB_MIRROR=https://gh-proxy.com
|
||||
# 或: --build-arg GITHUB_MIRROR=https://mirror.ghproxy.com
|
||||
ARG GITHUB_MIRROR
|
||||
|
||||
# 运行时依赖(使用清华镜像源 + BuildKit 缓存加速)
|
||||
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
--mount=type=cache,target=/var/lib/apt,sharing=locked \
|
||||
sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
|
||||
apt-get update && apt-get install -y --no-install-recommends \
|
||||
nginx \
|
||||
supervisor \
|
||||
libpq5 \
|
||||
curl \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
libjemalloc2
|
||||
RUN set -eux; \
|
||||
jemalloc_path="$(find /usr/lib -type f -name 'libjemalloc.so.2' | head -n1)"; \
|
||||
[ -n "$jemalloc_path" ]; \
|
||||
ln -sf "$jemalloc_path" /usr/local/lib/libjemalloc.so.2
|
||||
|
||||
# 从 base 镜像复制 Python 包
|
||||
COPY --from=builder /usr/local/lib/python3.12/site-packages /usr/local/lib/python3.12/site-packages
|
||||
COPY --from=builder /usr/local/lib/python3.13/site-packages /usr/local/lib/python3.13/site-packages
|
||||
|
||||
# 只复制需要的 Python 可执行文件
|
||||
COPY --from=builder /usr/local/bin/gunicorn /usr/local/bin/
|
||||
COPY --from=builder /usr/local/bin/uvicorn /usr/local/bin/
|
||||
COPY --from=builder /usr/local/bin/alembic /usr/local/bin/
|
||||
|
||||
# Hub 预编译二进制
|
||||
# 国内构建: --build-arg GITHUB_MIRROR=https://ghfast.top 即可走镜像下载
|
||||
# GITHUB_TOKEN 可选:未认证 API 限流 60 次/小时,认证后 5000 次/小时
|
||||
RUN set -eux; \
|
||||
arch="${TARGETARCH:-}"; \
|
||||
if [ -z "$arch" ]; then \
|
||||
arch="$(dpkg --print-architecture)"; \
|
||||
fi; \
|
||||
case "$arch" in \
|
||||
amd64|arm64) ;; \
|
||||
x86_64) arch="amd64" ;; \
|
||||
aarch64) arch="arm64" ;; \
|
||||
*) echo "Unsupported architecture: $arch"; exit 1 ;; \
|
||||
esac; \
|
||||
auth_header=""; \
|
||||
if [ -n "${GITHUB_TOKEN:-}" ]; then \
|
||||
auth_header="Authorization: token ${GITHUB_TOKEN}"; \
|
||||
fi; \
|
||||
tag="${HUB_TAG:-}"; \
|
||||
if [ -z "$tag" ]; then \
|
||||
tag="$(curl -sL ${auth_header:+-H "$auth_header"} "https://api.github.com/repos/${HUB_RELEASE_REPO}/releases" | python3 -c "import json,sys;print(next((r['tag_name'] for r in json.load(sys.stdin) if r.get('tag_name','').startswith('hub-v') and not r.get('draft') and not r.get('prerelease')),''))")"; \
|
||||
fi; \
|
||||
if [ -z "$tag" ]; then \
|
||||
echo "Failed to resolve hub release tag"; \
|
||||
exit 1; \
|
||||
fi; \
|
||||
echo "Using Hub release tag: $tag"; \
|
||||
origin_url="https://github.com/${HUB_RELEASE_REPO}/releases/download/${tag}/aether-hub-linux-${arch}.tar.gz"; \
|
||||
if [ -n "${GITHUB_MIRROR:-}" ]; then \
|
||||
url="${GITHUB_MIRROR}/https://github.com/${HUB_RELEASE_REPO}/releases/download/${tag}/aether-hub-linux-${arch}.tar.gz"; \
|
||||
echo "Using mirror: ${GITHUB_MIRROR}"; \
|
||||
else \
|
||||
url="$origin_url"; \
|
||||
fi; \
|
||||
curl -L --fail -o /tmp/aether-hub.tar.gz "$url"; \
|
||||
tar xzf /tmp/aether-hub.tar.gz -C /usr/local/bin; \
|
||||
chmod +x /usr/local/bin/aether-hub; \
|
||||
rm -f /tmp/aether-hub.tar.gz
|
||||
|
||||
# 从 builder 阶段复制前端构建产物
|
||||
COPY --from=builder /app/frontend/dist /usr/share/nginx/html
|
||||
RUN chmod -R 755 /usr/share/nginx/html
|
||||
|
||||
# 复制后端代码
|
||||
COPY src/ ./src/
|
||||
COPY alembic.ini ./
|
||||
COPY alembic/ ./alembic/
|
||||
COPY gunicorn_conf.py ./
|
||||
|
||||
# Nginx 配置模板
|
||||
# 策略:白名单后端路由 → 后端代理,其余全部 → 前端 SPA(index.html)
|
||||
# 智能处理 IP:有外层代理头就透传,没有就用直连 IP
|
||||
RUN printf '%s\n' \
|
||||
'server {' \
|
||||
' listen 80;' \
|
||||
' server_name _;' \
|
||||
' root /usr/share/nginx/html;' \
|
||||
' index index.html;' \
|
||||
' client_max_body_size 100M;' \
|
||||
'' \
|
||||
' location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {' \
|
||||
' expires 1y;' \
|
||||
' add_header Cache-Control "public, no-transform";' \
|
||||
' try_files $uri =404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location ~ ^/(src|node_modules)/ {' \
|
||||
' deny all;' \
|
||||
' return 404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location ~ ^/(dashboard|admin|login)(/|$) {' \
|
||||
' try_files $uri $uri/ /index.html;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location / {' \
|
||||
' try_files $uri $uri/ @backend;' \
|
||||
' }' \
|
||||
'' \
|
||||
' location @backend {' \
|
||||
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $remote_addr;' \
|
||||
' proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Connection "";' \
|
||||
' proxy_set_header Accept $http_accept;' \
|
||||
' proxy_set_header Content-Type $content_type;' \
|
||||
' proxy_set_header Authorization $http_authorization;' \
|
||||
' proxy_set_header X-Api-Key $http_x_api_key;' \
|
||||
' proxy_buffering off;' \
|
||||
' proxy_cache off;' \
|
||||
' proxy_request_buffering off;' \
|
||||
' chunked_transfer_encoding on;' \
|
||||
' gzip off;' \
|
||||
' add_header X-Accel-Buffering no;' \
|
||||
' proxy_connect_timeout 600s;' \
|
||||
' proxy_send_timeout 600s;' \
|
||||
' proxy_read_timeout 600s;' \
|
||||
' }' \
|
||||
'}' > /etc/nginx/sites-available/default.template
|
||||
'map $http_x_real_ip $real_ip {' \
|
||||
' default $http_x_real_ip;' \
|
||||
' "" $remote_addr;' \
|
||||
'}' \
|
||||
'' \
|
||||
'map $http_x_forwarded_for $forwarded_for {' \
|
||||
' default $http_x_forwarded_for;' \
|
||||
' "" $remote_addr;' \
|
||||
'}' \
|
||||
'' \
|
||||
'map $http_upgrade $connection_upgrade {' \
|
||||
' default upgrade;' \
|
||||
' "" "";' \
|
||||
'}' \
|
||||
'' \
|
||||
'server {' \
|
||||
' listen 80;' \
|
||||
' server_name _;' \
|
||||
' root /usr/share/nginx/html;' \
|
||||
' index index.html;' \
|
||||
' client_max_body_size 100M;' \
|
||||
'' \
|
||||
' # gzip 压缩配置(对 base64 图片等非流式响应有效)' \
|
||||
' gzip on;' \
|
||||
' gzip_min_length 256;' \
|
||||
' gzip_comp_level 5;' \
|
||||
' gzip_vary on;' \
|
||||
' gzip_proxied any;' \
|
||||
' gzip_types application/json text/plain text/css text/javascript application/javascript application/octet-stream;' \
|
||||
' gzip_disable "msie6";' \
|
||||
'' \
|
||||
' # 静态资源:长期缓存' \
|
||||
' location ~* \.(js|css|png|jpg|jpeg|gif|ico|svg|woff|woff2|ttf|eot)$ {' \
|
||||
' expires 1y;' \
|
||||
' add_header Cache-Control "public, no-transform";' \
|
||||
' try_files $uri =404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # 安全:阻止访问源码目录' \
|
||||
' location ~ ^/(src|node_modules)/ {' \
|
||||
' deny all;' \
|
||||
' return 404;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # WebSocket 隧道端点(aether-proxy tunnel 模式)' \
|
||||
' location = /api/internal/proxy-tunnel {' \
|
||||
' proxy_pass http://127.0.0.1:8085/proxy;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Upgrade $http_upgrade;' \
|
||||
' proxy_set_header Connection "upgrade";' \
|
||||
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
|
||||
' proxy_hide_header CF-Connecting-IP;' \
|
||||
' proxy_hide_header CF-IPCountry;' \
|
||||
' proxy_hide_header CF-Ray;' \
|
||||
' proxy_hide_header CF-Visitor;' \
|
||||
' proxy_hide_header CDN-Loop;' \
|
||||
' proxy_hide_header True-Client-IP;' \
|
||||
' proxy_hide_header CF-Worker;' \
|
||||
' proxy_hide_header CF-EW-Via;' \
|
||||
' proxy_hide_header CF-Warp-Tag-ID;' \
|
||||
' proxy_set_header CF-Connecting-IP "";' \
|
||||
' proxy_set_header CF-IPCountry "";' \
|
||||
' proxy_set_header CF-Ray "";' \
|
||||
' proxy_set_header CF-Visitor "";' \
|
||||
' proxy_set_header CDN-Loop "";' \
|
||||
' proxy_set_header True-Client-IP "";' \
|
||||
' proxy_set_header CF-Worker "";' \
|
||||
' proxy_set_header CF-EW-Via "";' \
|
||||
' proxy_set_header CF-Warp-Tag-ID "";' \
|
||||
' proxy_read_timeout 86400s;' \
|
||||
' proxy_send_timeout 86400s;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # 后端 API 路由(白名单)→ 代理到后端' \
|
||||
' location ~ ^/(api|v1|v1beta|upload|health)(/|$) {' \
|
||||
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' proxy_set_header Upgrade $http_upgrade;' \
|
||||
' proxy_set_header Connection $connection_upgrade;' \
|
||||
' proxy_set_header Accept $http_accept;' \
|
||||
' proxy_set_header Content-Type $content_type;' \
|
||||
' proxy_set_header Authorization $http_authorization;' \
|
||||
' proxy_set_header X-Api-Key $http_x_api_key;' \
|
||||
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
|
||||
' proxy_hide_header CF-Connecting-IP;' \
|
||||
' proxy_hide_header CF-IPCountry;' \
|
||||
' proxy_hide_header CF-Ray;' \
|
||||
' proxy_hide_header CF-Visitor;' \
|
||||
' proxy_hide_header CDN-Loop;' \
|
||||
' proxy_hide_header True-Client-IP;' \
|
||||
' proxy_hide_header CF-Worker;' \
|
||||
' proxy_hide_header CF-EW-Via;' \
|
||||
' proxy_hide_header CF-Warp-Tag-ID;' \
|
||||
' proxy_set_header CF-Connecting-IP "";' \
|
||||
' proxy_set_header CF-IPCountry "";' \
|
||||
' proxy_set_header CF-Ray "";' \
|
||||
' proxy_set_header CF-Visitor "";' \
|
||||
' proxy_set_header CDN-Loop "";' \
|
||||
' proxy_set_header True-Client-IP "";' \
|
||||
' proxy_set_header CF-Worker "";' \
|
||||
' proxy_set_header CF-EW-Via "";' \
|
||||
' proxy_set_header CF-Warp-Tag-ID "";' \
|
||||
' proxy_buffering off;' \
|
||||
' proxy_cache off;' \
|
||||
' proxy_request_buffering off;' \
|
||||
' chunked_transfer_encoding on;' \
|
||||
' gzip off;' \
|
||||
' add_header X-Accel-Buffering no;' \
|
||||
' proxy_connect_timeout 60s;' \
|
||||
' proxy_send_timeout 3600s;' \
|
||||
' proxy_read_timeout 3600s;' \
|
||||
' }' \
|
||||
'' \
|
||||
' # API 文档路由 → 代理到后端' \
|
||||
' location ~ ^/(docs|redoc|openapi\\.json)$ {' \
|
||||
' proxy_pass http://127.0.0.1:PORT_PLACEHOLDER;' \
|
||||
' proxy_http_version 1.1;' \
|
||||
' proxy_set_header Host $host;' \
|
||||
' proxy_set_header X-Real-IP $real_ip;' \
|
||||
' proxy_set_header X-Forwarded-For $forwarded_for;' \
|
||||
' proxy_set_header X-Forwarded-Proto $scheme;' \
|
||||
' # 剥离 CF 头,防止泄露给上游或返回给客户端' \
|
||||
' proxy_hide_header CF-Connecting-IP;' \
|
||||
' proxy_hide_header CF-IPCountry;' \
|
||||
' proxy_hide_header CF-Ray;' \
|
||||
' proxy_hide_header CF-Visitor;' \
|
||||
' proxy_hide_header CDN-Loop;' \
|
||||
' proxy_hide_header True-Client-IP;' \
|
||||
' proxy_hide_header CF-Worker;' \
|
||||
' proxy_hide_header CF-EW-Via;' \
|
||||
' proxy_hide_header CF-Warp-Tag-ID;' \
|
||||
' proxy_set_header CF-Connecting-IP "";' \
|
||||
' proxy_set_header CF-IPCountry "";' \
|
||||
' proxy_set_header CF-Ray "";' \
|
||||
' proxy_set_header CF-Visitor "";' \
|
||||
' proxy_set_header CDN-Loop "";' \
|
||||
' proxy_set_header True-Client-IP "";' \
|
||||
' proxy_set_header CF-Worker "";' \
|
||||
' proxy_set_header CF-EW-Via "";' \
|
||||
' proxy_set_header CF-Warp-Tag-ID "";' \
|
||||
' }' \
|
||||
'' \
|
||||
' # 所有其他路由 → 前端 SPA(先尝试静态文件,再回退到 index.html)' \
|
||||
' location / {' \
|
||||
' try_files $uri $uri/ /index.html;' \
|
||||
' }' \
|
||||
'}' > /etc/nginx/sites-available/default.template
|
||||
|
||||
# Supervisor 配置
|
||||
RUN printf '%s\n' \
|
||||
'[supervisord]' \
|
||||
'nodaemon=true' \
|
||||
'logfile=/var/log/supervisor/supervisord.log' \
|
||||
'pidfile=/var/run/supervisord.pid' \
|
||||
'' \
|
||||
'[program:nginx]' \
|
||||
'command=/bin/bash -c "sed \"s/PORT_PLACEHOLDER/${PORT:-8084}/g\" /etc/nginx/sites-available/default.template > /etc/nginx/sites-available/default && /usr/sbin/nginx -g \"daemon off;\""' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/var/log/nginx/access.log' \
|
||||
'stderr_logfile=/var/log/nginx/error.log' \
|
||||
'' \
|
||||
'[program:app]' \
|
||||
'command=gunicorn src.main:app -w %(ENV_GUNICORN_WORKERS)s -k uvicorn.workers.UvicornWorker --bind 0.0.0.0:%(ENV_PORT)s --timeout 120 --access-logfile - --error-logfile - --log-level info' \
|
||||
'directory=/app' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/dev/stdout' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'stderr_logfile_maxbytes=0' \
|
||||
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true' > /etc/supervisor/conf.d/supervisord.conf
|
||||
'[supervisord]' \
|
||||
'nodaemon=true' \
|
||||
'logfile=/var/log/supervisor/supervisord.log' \
|
||||
'pidfile=/var/run/supervisord.pid' \
|
||||
'' \
|
||||
'[program:nginx]' \
|
||||
'command=/bin/bash -c "sed \"s/PORT_PLACEHOLDER/${PORT:-8084}/g\" /etc/nginx/sites-available/default.template > /etc/nginx/sites-available/default && /usr/sbin/nginx -g \"daemon off;\""' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/var/log/nginx/access.log' \
|
||||
'stderr_logfile=/var/log/nginx/error.log' \
|
||||
'' \
|
||||
'[program:app]' \
|
||||
'command=/bin/bash -c "MAX_REQUESTS_JITTER=$((${MAX_REQUESTS:-50000}/20)); exec gunicorn src.main:app -c gunicorn_conf.py --preload -w %(ENV_GUNICORN_WORKERS)s -k uvicorn.workers.UvicornWorker --bind 0.0.0.0:%(ENV_PORT)s --max-requests ${MAX_REQUESTS:-50000} --max-requests-jitter $MAX_REQUESTS_JITTER --access-logfile - --error-logfile - --log-level info"' \
|
||||
'directory=/app' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/dev/stdout' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'stderr_logfile_maxbytes=0' \
|
||||
'environment=PYTHONUNBUFFERED=1,PYTHONIOENCODING=utf-8,LANG=C.UTF-8,LC_ALL=C.UTF-8,DOCKER_CONTAINER=true,LD_PRELOAD=/usr/local/lib/libjemalloc.so.2,MALLOC_CONF="background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000"' \
|
||||
'' \
|
||||
'[program:tunnel-hub]' \
|
||||
'command=/usr/local/bin/aether-hub --bind 0.0.0.0:8085' \
|
||||
'autostart=true' \
|
||||
'autorestart=true' \
|
||||
'stdout_logfile=/dev/stdout' \
|
||||
'stdout_logfile_maxbytes=0' \
|
||||
'stderr_logfile=/dev/stderr' \
|
||||
'stderr_logfile_maxbytes=0' > /etc/supervisor/conf.d/supervisord.conf
|
||||
|
||||
# 创建目录
|
||||
RUN mkdir -p /var/log/supervisor /app/logs /app/data
|
||||
|
||||
# 入口脚本(启动前执行迁移)
|
||||
COPY entrypoint.sh /entrypoint.sh
|
||||
RUN sed -i 's/\r$//' /entrypoint.sh && chmod +x /entrypoint.sh
|
||||
|
||||
# 环境变量
|
||||
ENV PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONIOENCODING=utf-8 \
|
||||
LANG=C.UTF-8 \
|
||||
LC_ALL=C.UTF-8 \
|
||||
PORT=8084
|
||||
LD_PRELOAD=/usr/local/lib/libjemalloc.so.2 \
|
||||
MALLOC_CONF=background_thread:true,dirty_decay_ms:5000,muzzy_decay_ms:5000 \
|
||||
PORT=8084 \
|
||||
GUNICORN_WORKERS=2 \
|
||||
MAX_REQUESTS=4000
|
||||
|
||||
EXPOSE 80
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
|
||||
CMD curl -f http://localhost/health || exit 1
|
||||
|
||||
ENTRYPOINT ["/entrypoint.sh"]
|
||||
CMD ["/usr/bin/supervisord", "-c", "/etc/supervisor/conf.d/supervisord.conf"]
|
||||
|
||||
+14
-11
@@ -1,25 +1,28 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# 构建镜像:编译环境 + 预编译的依赖
|
||||
# 用于 GitHub Actions CI 构建(不使用国内镜像源)
|
||||
# 构建命令: docker build -f Dockerfile.base -t aether-base:latest .
|
||||
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.13-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 构建工具
|
||||
RUN apt-get update && apt-get install -y \
|
||||
# 构建工具(使用 BuildKit 缓存加速)
|
||||
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
--mount=type=cache,target=/var/lib/apt,sharing=locked \
|
||||
apt-get update && apt-get install -y --no-install-recommends \
|
||||
libpq-dev \
|
||||
gcc \
|
||||
nodejs \
|
||||
npm \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
npm
|
||||
|
||||
# Python 依赖
|
||||
# Python 依赖(使用 BuildKit 缓存加速)
|
||||
COPY pyproject.toml README.md ./
|
||||
RUN mkdir -p src && touch src/__init__.py && \
|
||||
SETUPTOOLS_SCM_PRETEND_VERSION=0.1.0 pip install --no-cache-dir . && \
|
||||
pip cache purge
|
||||
RUN --mount=type=cache,target=/root/.cache/pip \
|
||||
mkdir -p src && touch src/__init__.py && \
|
||||
SETUPTOOLS_SCM_PRETEND_VERSION=0.1.0 pip install .
|
||||
|
||||
# 前端依赖(只安装,不构建)
|
||||
# 前端依赖(只安装,不构建,使用 BuildKit 缓存加速)
|
||||
COPY frontend/package*.json ./frontend/
|
||||
RUN cd frontend && npm ci
|
||||
RUN --mount=type=cache,target=/root/.npm \
|
||||
cd frontend && npm ci
|
||||
|
||||
+15
-12
@@ -1,28 +1,31 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# 构建镜像:编译环境 + 预编译的依赖(国内镜像源版本)
|
||||
# 构建命令: docker build -f Dockerfile.base.local -t aether-base:latest .
|
||||
# 只在 pyproject.toml 或 frontend/package*.json 变化时需要重建
|
||||
FROM python:3.12-slim
|
||||
FROM python:3.13-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 构建工具(使用清华镜像源)
|
||||
RUN sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
|
||||
apt-get update && apt-get install -y \
|
||||
# 构建工具(使用清华镜像源 + BuildKit 缓存加速)
|
||||
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked \
|
||||
--mount=type=cache,target=/var/lib/apt,sharing=locked \
|
||||
sed -i 's/deb.debian.org/mirrors.tuna.tsinghua.edu.cn/g' /etc/apt/sources.list.d/debian.sources && \
|
||||
apt-get update && apt-get install -y --no-install-recommends \
|
||||
libpq-dev \
|
||||
gcc \
|
||||
nodejs \
|
||||
npm \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
npm
|
||||
|
||||
# pip 镜像源
|
||||
RUN pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
|
||||
# Python 依赖
|
||||
# Python 依赖(使用 BuildKit 缓存加速)
|
||||
COPY pyproject.toml README.md ./
|
||||
RUN mkdir -p src && touch src/__init__.py && \
|
||||
SETUPTOOLS_SCM_PRETEND_VERSION=0.1.0 pip install --no-cache-dir . && \
|
||||
pip cache purge
|
||||
RUN --mount=type=cache,target=/root/.cache/pip \
|
||||
mkdir -p src && touch src/__init__.py && \
|
||||
SETUPTOOLS_SCM_PRETEND_VERSION=0.1.0 pip install .
|
||||
|
||||
# 前端依赖(只安装,不构建,使用淘宝镜像源)
|
||||
# 前端依赖(只安装,不构建,使用淘宝镜像源 + BuildKit 缓存加速)
|
||||
COPY frontend/package*.json ./frontend/
|
||||
RUN cd frontend && npm config set registry https://registry.npmmirror.com && npm ci
|
||||
RUN --mount=type=cache,target=/root/.npm \
|
||||
cd frontend && npm config set registry https://registry.npmmirror.com && npm ci
|
||||
|
||||
@@ -5,12 +5,17 @@ Aether 非商业开源许可证
|
||||
特此授予任何获得本软件及其相关文档文件(以下简称"软件")副本的人免费使用、
|
||||
复制、修改、合并、发布和分发本软件的权限,但须遵守以下条件:
|
||||
|
||||
1. 仅限非商业用途
|
||||
本软件不得用于商业目的。商业目的包括但不限于:
|
||||
1. 仅限非盈利用途
|
||||
本软件不得用于盈利目的。盈利目的包括但不限于:
|
||||
- 出售本软件或任何衍生作品
|
||||
- 使用本软件提供付费服务
|
||||
- 将本软件用于商业产品或服务
|
||||
- 将本软件用于任何旨在获取商业利益或金钱报酬的活动
|
||||
- 将本软件用于以盈利为目的的商业产品或服务
|
||||
|
||||
以下用途被明确允许:
|
||||
- 个人学习和研究
|
||||
- 教育机构的教学和研究
|
||||
- 非盈利组织的内部使用
|
||||
- 企业内部非盈利性质的使用(如内部工具、测试环境等)
|
||||
|
||||
2. 署名要求
|
||||
上述版权声明和本许可声明应包含在本软件的所有副本或主要部分中。
|
||||
@@ -22,7 +27,7 @@ Aether 非商业开源许可证
|
||||
您不得以不同的条款将本软件再许可给他人。
|
||||
|
||||
5. 商业许可
|
||||
如需商业使用,请联系版权持有人以获取单独的商业许可。
|
||||
如需将本软件用于盈利目的,请联系版权持有人以获取单独的商业许可。
|
||||
|
||||
本软件按"原样"提供,不提供任何明示或暗示的保证,包括但不限于对适销性、
|
||||
特定用途适用性和非侵权性的保证。在任何情况下,作者或版权持有人均不对任何
|
||||
|
||||
@@ -5,8 +5,8 @@
|
||||
<h1 align="center">Aether</h1>
|
||||
|
||||
<p align="center">
|
||||
<strong>开源 AI API 网关</strong><br>
|
||||
支持 Claude / OpenAI / Gemini 及其 CLI 客户端的统一接入层
|
||||
<strong>一站式 AI 基础设施平台</strong><br>
|
||||
支持 Claude / OpenAI / Gemini 及其 CLI 客户端的统一接入、格式转换、正/反向代理, 致力于成为用户驱动AI服务的底座
|
||||
</p>
|
||||
<p align="center">
|
||||
<a href="#简介">简介</a> •
|
||||
@@ -22,27 +22,15 @@
|
||||
|
||||
Aether 是一个自托管的 AI API 网关,为团队和个人提供多租户管理、智能负载均衡、成本配额控制和健康监控能力。通过统一的 API 入口,可以无缝对接 Claude、OpenAI、Gemini 等主流 AI 服务及其 CLI 工具。
|
||||
|
||||
### 页面预览
|
||||
<p align="center">
|
||||
<picture>
|
||||
<source media="(prefers-color-scheme: dark)" srcset="docs/architecture/architecture-dark.svg">
|
||||
<source media="(prefers-color-scheme: light)" srcset="docs/architecture/architecture-light.svg">
|
||||
<img src="docs/architecture/architecture-light.svg" width="680" alt="Aether Architecture">
|
||||
</picture>
|
||||
</p>
|
||||
|
||||
| 首页 | 仪表盘 |
|
||||
|:---:|:---:|
|
||||
|  |  |
|
||||
|
||||
| 健康监控 | 用户管理 |
|
||||
|:---:|:---:|
|
||||
|  |  |
|
||||
|
||||
| 提供商管理 | 使用记录 |
|
||||
|:---:|:---:|
|
||||
|  |  |
|
||||
|
||||
| 模型详情 | 关联提供商 |
|
||||
|:---:|:---:|
|
||||
|  |  |
|
||||
|
||||
| 链路追踪 | 系统设置 |
|
||||
|:---:|:---:|
|
||||
|  |  |
|
||||
页面预览: https://fawney19.github.io/Aether/
|
||||
|
||||
## 部署
|
||||
|
||||
@@ -51,20 +39,17 @@ Aether 是一个自托管的 AI API 网关,为团队和个人提供多租户
|
||||
```bash
|
||||
# 1. 克隆代码
|
||||
git clone https://github.com/fawney19/Aether.git
|
||||
cd aether
|
||||
cd Aether
|
||||
|
||||
# 2. 配置环境变量
|
||||
cp .env.example .env
|
||||
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
|
||||
|
||||
# 3. 部署
|
||||
docker-compose up -d
|
||||
# 3. 部署 / 更新(自动执行数据库迁移)
|
||||
docker compose pull && docker compose up -d
|
||||
|
||||
# 4. 首次部署时, 初始化数据库
|
||||
./migrate.sh
|
||||
|
||||
# 5. 更新
|
||||
docker-compose pull && docker-compose up -d && ./migrate.sh
|
||||
# 4. 升级前备份 (可选)
|
||||
docker compose exec postgres pg_dump -U postgres aether | gzip > backup_$(date +%Y%m%d_%H%M%S).sql.gz
|
||||
```
|
||||
|
||||
### Docker Compose(本地构建镜像)
|
||||
@@ -72,13 +57,14 @@ docker-compose pull && docker-compose up -d && ./migrate.sh
|
||||
```bash
|
||||
# 1. 克隆代码
|
||||
git clone https://github.com/fawney19/Aether.git
|
||||
cd aether
|
||||
cd Aether
|
||||
|
||||
# 2. 配置环境变量
|
||||
cp .env.example .env
|
||||
python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
|
||||
|
||||
# 3. 部署 / 更新(自动构建、启动、迁移)
|
||||
git pull
|
||||
./deploy.sh
|
||||
```
|
||||
|
||||
@@ -86,7 +72,7 @@ python generate_keys.py # 生成密钥, 并将生成的密钥填入 .env
|
||||
|
||||
```bash
|
||||
# 启动依赖
|
||||
docker-compose -f docker-compose.build.yml up -d postgres redis
|
||||
docker compose -f docker-compose.build.yml up -d postgres redis
|
||||
|
||||
# 后端
|
||||
uv sync
|
||||
@@ -96,6 +82,14 @@ uv sync
|
||||
cd frontend && npm install && npm run dev
|
||||
```
|
||||
|
||||
## Aether Proxy (可选)
|
||||
|
||||
Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙内的 Aether 实例中转 API 流量。或者部署在其他服务器为指定的提供商、账号、Key使用不同的节点访问。支持 TUI 向导一键配置、systemd 服务管理、TLS 加密、DNS 缓存及连接池调优。
|
||||
|
||||
- Docker Compose 部署或下载预编译二进制直接运行
|
||||
- 通过 `aether-proxy setup` 完成交互式配置,自动注册为系统服务
|
||||
- 详细文档见 [aether-proxy/README.md](aether-proxy/README.md)
|
||||
|
||||
## 环境变量
|
||||
|
||||
### 必需配置
|
||||
@@ -117,7 +111,7 @@ cd frontend && npm install && npm run dev
|
||||
| `APP_PORT` | 8084 | 应用端口 |
|
||||
| `API_KEY_PREFIX` | sk | API Key 前缀 |
|
||||
| `LOG_LEVEL` | INFO | 日志级别 (DEBUG/INFO/WARNING/ERROR) |
|
||||
| `GUNICORN_WORKERS` | 4 | Gunicorn 工作进程数 |
|
||||
| `GUNICORN_WORKERS` | 2 | Gunicorn 工作进程数 |
|
||||
| `DB_PORT` | 5432 | PostgreSQL 端口 |
|
||||
| `REDIS_PORT` | 6379 | Redis 端口 |
|
||||
|
||||
@@ -133,33 +127,57 @@ cd frontend && npm install && npm run dev
|
||||
| Headers | Base + 请求头 |
|
||||
| Full | Headers + 请求体 |
|
||||
|
||||
### Q: 管理员如何给模型配置 1M上下文 / 1H缓存 能力支持?
|
||||
### Q: 更新出问题如何回滚?
|
||||
|
||||
1. **模型管理**: 给模型设置 1M上下文 / 1H缓存 的能力支持, 并配置好价格
|
||||
2. **提供商管理**: 给端点添加支持该能力的密钥, 并勾选对应的能力标签
|
||||
**有备份的情况(推荐):**
|
||||
|
||||
### Q: 用户如何使用 1H缓存?
|
||||
```bash
|
||||
# 1. 停止应用
|
||||
docker compose stop app
|
||||
|
||||
- **模型级别**: 在模型管理中针对指定模型开启 1H缓存策略
|
||||
- **密钥级别**: 在密钥管理中针对指定密钥使用 1H缓存策略
|
||||
# 2. 恢复数据库(先清空再导入)
|
||||
docker compose exec -T postgres psql -U postgres -c "DROP DATABASE aether; CREATE DATABASE aether;"
|
||||
gunzip < backup_xxx.sql.gz | docker compose exec -T postgres psql -U postgres -d aether
|
||||
|
||||
> **注意**: 若对密钥设置强制 1H缓存, 则该密钥只能调用支持 1H缓存的模型
|
||||
# 3. 拉取旧版本镜像并重启
|
||||
# 方式一:使用具体版本 tag(如果有发布版本号)
|
||||
# 将 docker-compose.yml 中 image 从 ghcr.io/fawney19/aether:latest 改为指定版本
|
||||
# 方式二:使用之前记录的镜像 digest
|
||||
# 将 image 改为 ghcr.io/fawney19/aether@sha256:xxxxx
|
||||
docker compose up -d app
|
||||
```
|
||||
|
||||
### Q: 如何配置负载均衡?
|
||||
> 可以在升级前通过 `docker inspect ghcr.io/fawney19/aether:latest --format '{{index .RepoDigests 0}}'` 记录当前镜像 digest,方便回滚时使用。
|
||||
|
||||
在管理后台 **提供商管理** 中切换调度模式:
|
||||
**没有备份的情况:**
|
||||
|
||||
| 模式 | 说明 | 适用场景 |
|
||||
|------|------|----------|
|
||||
| **提供商优先** | 按 Provider 优先级排序, 同优先级内按 Key 优先级排序, 相同优先级哈希分散 | 优先使用特定供应商 |
|
||||
| **全局 Key 优先** | 忽略 Provider 层级, 所有 Key 按全局优先级统一排序, 相同优先级哈希分散 | 跨 Provider 统一调度, 最大化利用所有 Key |
|
||||
```bash
|
||||
# 1. 用当前容器回退数据库迁移(回退 1 步,按需调整数字)
|
||||
docker compose exec app alembic downgrade -1
|
||||
|
||||
### Q: 提供商免费套餐的计费模式会计入成本吗?
|
||||
# 2. 查看回退后的版本确认正确
|
||||
docker compose exec app alembic current
|
||||
|
||||
> **不会**。免费套餐的计费模式倍率为 0, 产生的记录不计入成本费用。
|
||||
# 3. 切回旧镜像并重启(同上方式修改 docker-compose.yml 中的 image)
|
||||
docker compose up -d app
|
||||
```
|
||||
|
||||
> 注意:没有备份的回滚依赖 alembic downgrade,如果迁移涉及不可逆的数据变更(如删除列),可能无法完全恢复数据。因此强烈建议升级前备份。
|
||||
|
||||
---
|
||||
|
||||
## 许可证
|
||||
|
||||
本项目采用 [Aether 非商业开源许可证](LICENSE)。
|
||||
本项目采用 [Aether 非商业开源许可证](LICENSE)。允许个人学习、教育研究、非盈利组织及企业内部非盈利性质的使用;禁止用于盈利目的。商业使用请联系获取商业许可。
|
||||
|
||||
## 联系作者
|
||||
|
||||
<p align="center">
|
||||
<img src="docs/author/qq_qrcode.jpg" width="200" alt="QQ二维码">
|
||||
|
||||
<img src="docs/author/qrcode_1770574997172.jpg" width="200" alt="QQ群二维码">
|
||||
</p>
|
||||
|
||||
## Star History
|
||||
|
||||
[](https://star-history.com/#fawney19/Aether&Date)
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
target/
|
||||
.git/
|
||||
.DS_Store
|
||||
Generated
+2012
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,27 @@
|
||||
[package]
|
||||
name = "aether-hub"
|
||||
version = "0.2.0"
|
||||
edition = "2021"
|
||||
description = "Tunnel Hub for Aether - frame router between workers and proxies"
|
||||
|
||||
[dependencies]
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
axum = { version = "0.8", features = ["ws"] }
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
tracing = "0.1"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
||||
clap = { version = "4", features = ["derive", "env"] }
|
||||
dashmap = "6"
|
||||
parking_lot = "0.12"
|
||||
flate2 = "1"
|
||||
futures-util = "0.3"
|
||||
bytes = "1"
|
||||
async-stream = "0.3"
|
||||
http-body-util = "0.1"
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
|
||||
|
||||
[profile.release]
|
||||
lto = true
|
||||
strip = true
|
||||
codegen-units = 1
|
||||
@@ -0,0 +1,34 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
|
||||
FROM rust:1.85-slim AS builder
|
||||
WORKDIR /build/aether-hub
|
||||
|
||||
# 可选:配置国内 Cargo 镜像源(本地构建时传 --build-arg CARGO_MIRROR=1)
|
||||
ARG CARGO_MIRROR
|
||||
RUN if [ -n "$CARGO_MIRROR" ]; then \
|
||||
printf '[source.crates-io]\nreplace-with = "tuna"\n\n[source.tuna]\nregistry = "sparse+https://mirrors.tuna.tsinghua.edu.cn/crates.io-index/"\n' \
|
||||
> /usr/local/cargo/config.toml; \
|
||||
fi
|
||||
|
||||
# 先构建依赖层,最大化后续代码变更时的缓存命中
|
||||
COPY Cargo.toml Cargo.lock ./
|
||||
RUN mkdir src && printf 'fn main() {}\n' > src/main.rs
|
||||
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
|
||||
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
|
||||
cargo build --release --locked
|
||||
RUN rm -rf src
|
||||
|
||||
COPY src ./src
|
||||
RUN --mount=type=cache,target=/usr/local/cargo/registry,sharing=locked \
|
||||
--mount=type=cache,target=/build/aether-hub/target,sharing=locked \
|
||||
cargo build --release --locked && \
|
||||
cp target/release/aether-hub /tmp/aether-hub
|
||||
|
||||
FROM debian:bookworm-slim
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
COPY --from=builder /tmp/aether-hub /usr/local/bin/aether-hub
|
||||
|
||||
EXPOSE 8085
|
||||
ENTRYPOINT ["/usr/local/bin/aether-hub"]
|
||||
CMD ["--bind", "0.0.0.0:8085"]
|
||||
@@ -0,0 +1,38 @@
|
||||
# aether-hub
|
||||
|
||||
`aether-hub` 是 Tunnel Hub 服务,负责在 proxy 与 worker 之间路由帧。
|
||||
|
||||
已集成在Docker镜像中, 无需单独部署。
|
||||
|
||||
## 部署端指定 Hub 版本并构建
|
||||
|
||||
```bash
|
||||
cd /path/to/Aether
|
||||
./deploy.sh --hub-tag hub-v0.1.0
|
||||
```
|
||||
|
||||
不指定 `--hub-tag` 时,`./deploy.sh` 会自动解析最新 `hub-v*` release,并在构建 app 镜像时从 GitHub Release 下载对应架构的 Hub 二进制。
|
||||
|
||||
## build.sh 模式说明
|
||||
|
||||
- 默认是 `binary` 模式(`cross` 构建二进制)。
|
||||
- `--upload <hub-vX.Y.Z>` 会把构建产物上传到 GitHub Release。
|
||||
- 加 `--image` 后进入镜像模式(`docker buildx`,可选)。
|
||||
|
||||
常用参数:
|
||||
|
||||
- `--tag <tag>`: 镜像 tag
|
||||
- `--image-name <name>`: 镜像名(默认 `ghcr.io/fawney19/aether-hub`)
|
||||
- `--platforms <list>`: 例如 `linux/amd64,linux/arm64`
|
||||
- `--push`: 推送镜像
|
||||
- `--load`: 加载到本地 Docker(单平台)
|
||||
- `--latest`: 额外打 `latest` tag
|
||||
|
||||
## 运行时参数
|
||||
|
||||
- `TUNNEL_HUB_WORKER_IDLE_TIMEOUT`:worker 心跳空闲超时,默认 `60` 秒
|
||||
- `TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY`:单连接出站队列容量,默认 `128`;队列打满时会把连接视为拥塞并主动关闭,避免 Hub 内存无限增长
|
||||
|
||||
## 与部署脚本关系
|
||||
|
||||
- `./deploy.sh`: 本地构建部署(会本地构建 app/base,并在构建 app 时从 GitHub Release 下载 Hub,可用 `--hub-tag` 固定版本)。
|
||||
Executable
+272
@@ -0,0 +1,272 @@
|
||||
#!/bin/bash
|
||||
# aether-hub 构建脚本
|
||||
#
|
||||
# 支持两种模式:
|
||||
# 1) binary 模式(默认): 构建多架构二进制并可上传 GitHub Release
|
||||
# 2) image 模式: 构建并推送/加载 Docker 镜像(推荐生产发布用)
|
||||
#
|
||||
# 示例:
|
||||
# # binary 模式(兼容旧行为)
|
||||
# ./build.sh
|
||||
# ./build.sh amd64
|
||||
# ./build.sh --upload hub-v0.1.0
|
||||
#
|
||||
# # image 模式(多架构推送)
|
||||
# ./build.sh --image --tag v0.2.5 --push --latest
|
||||
# ./build.sh --image --tag sha-abc123 --image-name ghcr.io/fawney19/aether-hub --push
|
||||
# ./build.sh --image --tag local-test --platforms linux/amd64 --load
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)"
|
||||
PROJECT_DIR="$(cd "$SCRIPT_DIR/.." && pwd)"
|
||||
DIST_DIR="$SCRIPT_DIR/dist"
|
||||
|
||||
# -------------------------------
|
||||
# Defaults
|
||||
# -------------------------------
|
||||
MODE="binary" # binary | image
|
||||
|
||||
# binary mode options
|
||||
UPLOAD=false
|
||||
UPLOAD_TAG=""
|
||||
BINARY_TARGETS=""
|
||||
|
||||
# image mode options
|
||||
IMAGE_NAME="${IMAGE_NAME:-ghcr.io/fawney19/aether-hub}"
|
||||
IMAGE_TAG=""
|
||||
IMAGE_PLATFORMS="linux/amd64,linux/arm64"
|
||||
IMAGE_PUSH=false
|
||||
IMAGE_LOAD=false
|
||||
IMAGE_LATEST=false
|
||||
|
||||
usage() {
|
||||
cat <<'EOF'
|
||||
用法:
|
||||
./build.sh [binary-args]
|
||||
./build.sh --image [image-args]
|
||||
|
||||
binary 模式(默认):
|
||||
amd64|arm64 仅构建指定架构(可重复)
|
||||
--upload <hub-vX.Y.Z> 上传到 GitHub Release(需要 gh CLI)
|
||||
|
||||
image 模式:
|
||||
--image 启用镜像模式
|
||||
--tag <tag> 镜像 tag(默认自动从 git describe 推导)
|
||||
--image-name <name> 镜像名(默认 ghcr.io/fawney19/aether-hub)
|
||||
--platforms <list> 平台列表,逗号分隔(默认 linux/amd64,linux/arm64)
|
||||
--push 推送镜像到仓库
|
||||
--load 加载到本地 Docker(仅单平台)
|
||||
--latest 额外打 latest tag
|
||||
|
||||
通用:
|
||||
-h, --help 显示帮助
|
||||
EOF
|
||||
}
|
||||
|
||||
while [ $# -gt 0 ]; do
|
||||
case "$1" in
|
||||
--image)
|
||||
MODE="image"
|
||||
shift
|
||||
;;
|
||||
--tag)
|
||||
IMAGE_TAG="${2:-}"
|
||||
shift 2
|
||||
;;
|
||||
--image-name)
|
||||
IMAGE_NAME="${2:-}"
|
||||
shift 2
|
||||
;;
|
||||
--platforms)
|
||||
IMAGE_PLATFORMS="${2:-}"
|
||||
shift 2
|
||||
;;
|
||||
--push)
|
||||
IMAGE_PUSH=true
|
||||
shift
|
||||
;;
|
||||
--load)
|
||||
IMAGE_LOAD=true
|
||||
shift
|
||||
;;
|
||||
--latest)
|
||||
IMAGE_LATEST=true
|
||||
shift
|
||||
;;
|
||||
--upload)
|
||||
UPLOAD=true
|
||||
UPLOAD_TAG="${2:-}"
|
||||
shift 2
|
||||
;;
|
||||
amd64|arm64)
|
||||
BINARY_TARGETS="$BINARY_TARGETS $1"
|
||||
shift
|
||||
;;
|
||||
-h|--help)
|
||||
usage
|
||||
exit 0
|
||||
;;
|
||||
*)
|
||||
echo "❌ 未知参数: $1"
|
||||
usage
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
build_binary() {
|
||||
if [ -z "$BINARY_TARGETS" ]; then
|
||||
BINARY_TARGETS="amd64 arm64"
|
||||
fi
|
||||
|
||||
if ! command -v cross >/dev/null 2>&1; then
|
||||
echo "❌ 需要安装 cross: cargo install cross --git https://github.com/cross-rs/cross"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
mkdir -p "$DIST_DIR"
|
||||
|
||||
echo "🔨 开始构建 aether-hub 二进制..."
|
||||
echo " 目标平台: $BINARY_TARGETS"
|
||||
echo ""
|
||||
|
||||
ARTIFACTS=""
|
||||
for arch in $BINARY_TARGETS; do
|
||||
case "$arch" in
|
||||
amd64) target="x86_64-unknown-linux-gnu" ;;
|
||||
arm64) target="aarch64-unknown-linux-gnu" ;;
|
||||
*) echo "❌ 未知架构: $arch"; exit 1 ;;
|
||||
esac
|
||||
|
||||
echo ">>> 构建 $arch ($target)..."
|
||||
cd "$SCRIPT_DIR"
|
||||
cross build --release --target "$target" --locked
|
||||
|
||||
BIN="target/$target/release/aether-hub"
|
||||
if [ ! -f "$BIN" ]; then
|
||||
echo "❌ 未找到二进制文件: $BIN"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
ARCHIVE="$DIST_DIR/aether-hub-linux-$arch.tar.gz"
|
||||
tar czf "$ARCHIVE" -C "target/$target/release" aether-hub
|
||||
ARTIFACTS="$ARTIFACTS $ARCHIVE"
|
||||
|
||||
SIZE=$(du -h "$ARCHIVE" | cut -f1)
|
||||
echo "✅ $arch 构建完成: $ARCHIVE ($SIZE)"
|
||||
echo ""
|
||||
done
|
||||
|
||||
cd "$DIST_DIR"
|
||||
shasum -a 256 aether-hub-*.tar.gz > SHA256SUMS.txt
|
||||
echo "📋 SHA256 校验和:"
|
||||
cat SHA256SUMS.txt
|
||||
echo ""
|
||||
|
||||
if [ "$UPLOAD" = true ]; then
|
||||
if [ -z "$UPLOAD_TAG" ]; then
|
||||
echo "❌ --upload 需要指定 tag,例如: ./build.sh --upload hub-v0.1.0"
|
||||
exit 1
|
||||
fi
|
||||
if ! command -v gh >/dev/null 2>&1; then
|
||||
echo "❌ 需要安装 GitHub CLI: brew install gh"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "📦 上传到 GitHub Release: $UPLOAD_TAG"
|
||||
cd "$PROJECT_DIR"
|
||||
|
||||
if ! git rev-parse "$UPLOAD_TAG" >/dev/null 2>&1; then
|
||||
git tag "$UPLOAD_TAG"
|
||||
git push origin "$UPLOAD_TAG"
|
||||
fi
|
||||
|
||||
gh release create "$UPLOAD_TAG" \
|
||||
--title "aether-hub ${UPLOAD_TAG#hub-}" \
|
||||
--generate-notes \
|
||||
$ARTIFACTS \
|
||||
"$DIST_DIR/SHA256SUMS.txt"
|
||||
|
||||
echo "✅ 上传完成!"
|
||||
fi
|
||||
|
||||
echo "🎉 binary 模式完成!"
|
||||
}
|
||||
|
||||
build_image() {
|
||||
if ! command -v docker >/dev/null 2>&1; then
|
||||
echo "❌ 未找到 docker,请先安装 Docker"
|
||||
exit 1
|
||||
fi
|
||||
if ! docker buildx version >/dev/null 2>&1; then
|
||||
echo "❌ 未找到 docker buildx,请先启用 buildx"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if [ "$IMAGE_PUSH" = true ] && [ "$IMAGE_LOAD" = true ]; then
|
||||
echo "❌ --push 与 --load 不能同时使用"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if [ "$IMAGE_PUSH" = false ] && [ "$IMAGE_LOAD" = false ]; then
|
||||
# image 模式默认走 push,符合发布场景
|
||||
IMAGE_PUSH=true
|
||||
fi
|
||||
|
||||
if [ -z "$IMAGE_TAG" ]; then
|
||||
IMAGE_TAG=$(git -C "$PROJECT_DIR" describe --tags --always 2>/dev/null | sed 's/^v//')
|
||||
if [ -z "$IMAGE_TAG" ]; then
|
||||
IMAGE_TAG=$(date +%Y%m%d%H%M%S)
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$IMAGE_LOAD" = true ] && [[ "$IMAGE_PLATFORMS" == *,* ]]; then
|
||||
echo "❌ --load 仅支持单平台,请用 --platforms linux/amd64(或 arm64)"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
local ref="${IMAGE_NAME}:${IMAGE_TAG}"
|
||||
local cmd=(docker buildx build
|
||||
--platform "$IMAGE_PLATFORMS"
|
||||
-f "$SCRIPT_DIR/Dockerfile"
|
||||
-t "$ref"
|
||||
)
|
||||
|
||||
if [ "$IMAGE_LATEST" = true ]; then
|
||||
cmd+=(-t "${IMAGE_NAME}:latest")
|
||||
fi
|
||||
|
||||
if [ "$IMAGE_PUSH" = true ]; then
|
||||
cmd+=(--push)
|
||||
else
|
||||
cmd+=(--load)
|
||||
fi
|
||||
|
||||
cmd+=("$SCRIPT_DIR")
|
||||
|
||||
echo "🔨 开始构建 aether-hub 镜像..."
|
||||
echo " image: $ref"
|
||||
echo " platforms: $IMAGE_PLATFORMS"
|
||||
echo " mode: $([ "$IMAGE_PUSH" = true ] && echo push || echo load)"
|
||||
echo ""
|
||||
|
||||
"${cmd[@]}"
|
||||
|
||||
if [ "$IMAGE_PUSH" = true ]; then
|
||||
echo "✅ 镜像已推送: $ref"
|
||||
if [ "$IMAGE_LATEST" = true ]; then
|
||||
echo "✅ 镜像已推送: ${IMAGE_NAME}:latest"
|
||||
fi
|
||||
else
|
||||
echo "✅ 镜像已加载到本地: $ref"
|
||||
fi
|
||||
|
||||
echo "🎉 image 模式完成!"
|
||||
}
|
||||
|
||||
if [ "$MODE" = "image" ]; then
|
||||
build_image
|
||||
else
|
||||
build_binary
|
||||
fi
|
||||
@@ -0,0 +1,85 @@
|
||||
use reqwest::Client;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ControlPlaneClient {
|
||||
client: Option<Client>,
|
||||
base_url: String,
|
||||
}
|
||||
|
||||
impl ControlPlaneClient {
|
||||
pub fn new(base_url: String) -> Self {
|
||||
let client = Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(10))
|
||||
.build()
|
||||
.ok();
|
||||
Self { client, base_url }
|
||||
}
|
||||
|
||||
pub fn disabled() -> Self {
|
||||
Self {
|
||||
client: None,
|
||||
base_url: String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn heartbeat_ack(&self, payload: &[u8]) -> Result<Vec<u8>, String> {
|
||||
let Some(client) = &self.client else {
|
||||
return Ok(b"{}".to_vec());
|
||||
};
|
||||
let url = format!(
|
||||
"{}/api/internal/hub/heartbeat",
|
||||
self.base_url.trim_end_matches('/')
|
||||
);
|
||||
let response = client
|
||||
.post(&url)
|
||||
.header("content-type", "application/json")
|
||||
.body(payload.to_vec())
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("heartbeat callback request failed: {e}"))?;
|
||||
if !response.status().is_success() {
|
||||
return Err(format!(
|
||||
"heartbeat callback failed with status {}",
|
||||
response.status()
|
||||
));
|
||||
}
|
||||
response
|
||||
.bytes()
|
||||
.await
|
||||
.map(|bytes| bytes.to_vec())
|
||||
.map_err(|e| format!("heartbeat callback body read failed: {e}"))
|
||||
}
|
||||
|
||||
pub async fn push_node_status(
|
||||
&self,
|
||||
node_id: &str,
|
||||
connected: bool,
|
||||
conn_count: usize,
|
||||
) -> Result<(), String> {
|
||||
let Some(client) = &self.client else {
|
||||
return Ok(());
|
||||
};
|
||||
let url = format!(
|
||||
"{}/api/internal/hub/node-status",
|
||||
self.base_url.trim_end_matches('/')
|
||||
);
|
||||
let response = client
|
||||
.post(&url)
|
||||
.json(&serde_json::json!({
|
||||
"node_id": node_id,
|
||||
"connected": connected,
|
||||
"conn_count": conn_count,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("node-status callback request failed: {e}"))?;
|
||||
if response.status().is_success() {
|
||||
Ok(())
|
||||
} else {
|
||||
Err(format!(
|
||||
"node-status callback failed with status {}",
|
||||
response.status()
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,878 @@
|
||||
use std::collections::HashMap;
|
||||
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, AtomicUsize, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::Message;
|
||||
use bytes::Bytes;
|
||||
use dashmap::DashMap;
|
||||
use parking_lot::{Mutex, RwLock};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::sync::mpsc::error::TrySendError;
|
||||
use tokio::sync::{watch, Notify};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use crate::control_plane::ControlPlaneClient;
|
||||
use crate::protocol;
|
||||
|
||||
const MAX_REQUEST_BODY_FRAME_SIZE: usize = 32 * 1024;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum SendStatus {
|
||||
Queued,
|
||||
Closed,
|
||||
Congested,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct ConnConfig {
|
||||
pub ping_interval: Duration,
|
||||
pub idle_timeout: Duration,
|
||||
pub outbound_queue_capacity: usize,
|
||||
}
|
||||
|
||||
pub struct BoundedOutbound {
|
||||
tx: mpsc::Sender<Message>,
|
||||
close_tx: watch::Sender<bool>,
|
||||
closing: AtomicBool,
|
||||
}
|
||||
|
||||
impl BoundedOutbound {
|
||||
pub fn new(tx: mpsc::Sender<Message>, close_tx: watch::Sender<bool>) -> Self {
|
||||
Self {
|
||||
tx,
|
||||
close_tx,
|
||||
closing: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn send(&self, msg: Message) -> SendStatus {
|
||||
if self.is_closing() {
|
||||
return SendStatus::Closed;
|
||||
}
|
||||
|
||||
match self.tx.try_send(msg) {
|
||||
Ok(()) => SendStatus::Queued,
|
||||
Err(TrySendError::Closed(_)) => {
|
||||
self.mark_closing();
|
||||
SendStatus::Closed
|
||||
}
|
||||
Err(TrySendError::Full(_)) => {
|
||||
self.mark_closing();
|
||||
SendStatus::Congested
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_closing(&self) -> bool {
|
||||
self.closing.load(Ordering::Acquire)
|
||||
}
|
||||
|
||||
pub fn mark_closing(&self) -> bool {
|
||||
if self.closing.swap(true, Ordering::AcqRel) {
|
||||
return false;
|
||||
}
|
||||
let _ = self.close_tx.send(true);
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ProxyConn {
|
||||
pub id: u64,
|
||||
pub node_id: String,
|
||||
pub node_name: String,
|
||||
pub outbound: BoundedOutbound,
|
||||
next_stream_id: AtomicU32,
|
||||
pub stream_count: AtomicUsize,
|
||||
pub max_streams: usize,
|
||||
}
|
||||
|
||||
impl ProxyConn {
|
||||
pub fn new(
|
||||
id: u64,
|
||||
node_id: String,
|
||||
node_name: String,
|
||||
tx: mpsc::Sender<Message>,
|
||||
close_tx: watch::Sender<bool>,
|
||||
max_streams: usize,
|
||||
) -> Self {
|
||||
Self {
|
||||
id,
|
||||
node_id,
|
||||
node_name,
|
||||
outbound: BoundedOutbound::new(tx, close_tx),
|
||||
next_stream_id: AtomicU32::new(2),
|
||||
stream_count: AtomicUsize::new(0),
|
||||
max_streams,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn alloc_stream_id(&self) -> Option<u32> {
|
||||
let mut current = self.stream_count.load(Ordering::Relaxed);
|
||||
loop {
|
||||
if current >= self.max_streams || !self.is_available() {
|
||||
return None;
|
||||
}
|
||||
match self.stream_count.compare_exchange_weak(
|
||||
current,
|
||||
current + 1,
|
||||
Ordering::AcqRel,
|
||||
Ordering::Relaxed,
|
||||
) {
|
||||
Ok(_) => break,
|
||||
Err(observed) => current = observed,
|
||||
}
|
||||
}
|
||||
|
||||
let sid = loop {
|
||||
let current_sid = self.next_stream_id.load(Ordering::Relaxed);
|
||||
let next_sid = if current_sid >= 0xFFFF_FFFE {
|
||||
2
|
||||
} else {
|
||||
current_sid + 2
|
||||
};
|
||||
if self
|
||||
.next_stream_id
|
||||
.compare_exchange_weak(current_sid, next_sid, Ordering::AcqRel, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
break current_sid;
|
||||
}
|
||||
};
|
||||
|
||||
Some(sid)
|
||||
}
|
||||
|
||||
pub fn release_stream(&self) {
|
||||
let mut current = self.stream_count.load(Ordering::Relaxed);
|
||||
while current > 0 {
|
||||
match self.stream_count.compare_exchange_weak(
|
||||
current,
|
||||
current - 1,
|
||||
Ordering::AcqRel,
|
||||
Ordering::Relaxed,
|
||||
) {
|
||||
Ok(_) => return,
|
||||
Err(observed) => current = observed,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn is_available(&self) -> bool {
|
||||
!self.outbound.is_closing()
|
||||
}
|
||||
|
||||
pub fn request_close(&self) {
|
||||
self.outbound.mark_closing();
|
||||
}
|
||||
|
||||
pub fn send(&self, msg: Message) -> SendStatus {
|
||||
let was_closing = self.outbound.is_closing();
|
||||
let status = self.outbound.send(msg);
|
||||
if status == SendStatus::Congested && !was_closing {
|
||||
warn!(
|
||||
conn_id = self.id,
|
||||
node_id = %self.node_id,
|
||||
node_name = %self.node_name,
|
||||
queued_streams = self.stream_count.load(Ordering::Relaxed),
|
||||
"proxy outbound queue full, closing congested connection"
|
||||
);
|
||||
}
|
||||
status
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct LocalResponseHead {
|
||||
pub status: u16,
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum LocalBodyEvent {
|
||||
Chunk(Bytes),
|
||||
End,
|
||||
Error(String),
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct LocalWaitState {
|
||||
response: Option<LocalResponseHead>,
|
||||
error: Option<String>,
|
||||
}
|
||||
|
||||
pub struct LocalStream {
|
||||
pub id: u64,
|
||||
proxy_conn_id: u64,
|
||||
proxy_stream_id: u32,
|
||||
wait_state: Mutex<LocalWaitState>,
|
||||
headers_notify: Notify,
|
||||
body_tx: mpsc::Sender<LocalBodyEvent>,
|
||||
body_rx: Mutex<Option<mpsc::Receiver<LocalBodyEvent>>>,
|
||||
terminal: AtomicBool,
|
||||
}
|
||||
|
||||
impl LocalStream {
|
||||
fn new(id: u64, proxy_conn_id: u64, proxy_stream_id: u32) -> Self {
|
||||
let (body_tx, body_rx) = mpsc::channel(128);
|
||||
Self {
|
||||
id,
|
||||
proxy_conn_id,
|
||||
proxy_stream_id,
|
||||
wait_state: Mutex::new(LocalWaitState::default()),
|
||||
headers_notify: Notify::new(),
|
||||
body_tx,
|
||||
body_rx: Mutex::new(Some(body_rx)),
|
||||
terminal: AtomicBool::new(false),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn wait_headers(&self, timeout: Duration) -> Result<LocalResponseHead, String> {
|
||||
tokio::time::timeout(timeout, async {
|
||||
loop {
|
||||
let outcome = {
|
||||
let state = self.wait_state.lock();
|
||||
if let Some(response) = &state.response {
|
||||
return Ok(response.clone());
|
||||
}
|
||||
state.error.clone()
|
||||
};
|
||||
if let Some(error) = outcome {
|
||||
return Err(error);
|
||||
}
|
||||
self.headers_notify.notified().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.map_err(|_| "timed out waiting for response headers".to_string())?
|
||||
}
|
||||
|
||||
pub fn take_body_receiver(&self) -> Option<mpsc::Receiver<LocalBodyEvent>> {
|
||||
self.body_rx.lock().take()
|
||||
}
|
||||
|
||||
fn set_response_headers(&self, meta: protocol::ResponseMeta) {
|
||||
let mut notify = false;
|
||||
{
|
||||
let mut state = self.wait_state.lock();
|
||||
if state.response.is_none() && state.error.is_none() {
|
||||
state.response = Some(LocalResponseHead {
|
||||
status: meta.status,
|
||||
headers: meta.headers,
|
||||
});
|
||||
notify = true;
|
||||
}
|
||||
}
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
}
|
||||
|
||||
fn push_body_chunk(&self, payload: Bytes) -> bool {
|
||||
if self.terminal.load(Ordering::Acquire) {
|
||||
return false;
|
||||
}
|
||||
self.body_tx
|
||||
.try_send(LocalBodyEvent::Chunk(payload))
|
||||
.is_ok()
|
||||
}
|
||||
|
||||
fn finish(&self) {
|
||||
if self.terminal.swap(true, Ordering::AcqRel) {
|
||||
return;
|
||||
}
|
||||
let mut notify = false;
|
||||
{
|
||||
let mut state = self.wait_state.lock();
|
||||
if state.response.is_none() && state.error.is_none() {
|
||||
state.error = Some("stream ended before response headers".to_string());
|
||||
notify = true;
|
||||
}
|
||||
}
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
let _ = self.body_tx.try_send(LocalBodyEvent::End);
|
||||
}
|
||||
|
||||
fn fail(&self, error: impl Into<String>) {
|
||||
if self.terminal.swap(true, Ordering::AcqRel) {
|
||||
return;
|
||||
}
|
||||
|
||||
let error = error.into();
|
||||
let mut notify = false;
|
||||
{
|
||||
let mut state = self.wait_state.lock();
|
||||
if state.response.is_none() && state.error.is_none() {
|
||||
state.error = Some(error.clone());
|
||||
notify = true;
|
||||
}
|
||||
}
|
||||
if notify {
|
||||
self.headers_notify.notify_waiters();
|
||||
}
|
||||
let _ = self.body_tx.try_send(LocalBodyEvent::Error(error));
|
||||
}
|
||||
}
|
||||
|
||||
pub struct HubRouter {
|
||||
proxy_conns: RwLock<HashMap<String, Vec<Arc<ProxyConn>>>>,
|
||||
proxy_conns_by_id: DashMap<u64, Arc<ProxyConn>>,
|
||||
local_streams: DashMap<u64, Arc<LocalStream>>,
|
||||
proxy_to_local: DashMap<(u64, u32), u64>,
|
||||
next_conn_id: AtomicU64,
|
||||
next_local_stream_id: AtomicU64,
|
||||
control_plane: ControlPlaneClient,
|
||||
}
|
||||
|
||||
impl HubRouter {
|
||||
pub fn new(control_plane: ControlPlaneClient) -> Arc<Self> {
|
||||
Arc::new(Self {
|
||||
proxy_conns: RwLock::new(HashMap::new()),
|
||||
proxy_conns_by_id: DashMap::new(),
|
||||
local_streams: DashMap::new(),
|
||||
proxy_to_local: DashMap::new(),
|
||||
next_conn_id: AtomicU64::new(1),
|
||||
next_local_stream_id: AtomicU64::new(1),
|
||||
control_plane,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn alloc_conn_id(&self) -> u64 {
|
||||
self.next_conn_id.fetch_add(1, Ordering::Relaxed)
|
||||
}
|
||||
|
||||
pub fn register_proxy(&self, conn: Arc<ProxyConn>) {
|
||||
let node_id = conn.node_id.clone();
|
||||
let node_name = conn.node_name.clone();
|
||||
let conn_id = conn.id;
|
||||
self.proxy_conns_by_id.insert(conn_id, conn.clone());
|
||||
|
||||
let pool_size = {
|
||||
let mut map = self.proxy_conns.write();
|
||||
map.entry(node_id.clone()).or_default().push(conn);
|
||||
map.get(&node_id).map(|v| v.len()).unwrap_or(0)
|
||||
};
|
||||
|
||||
info!(
|
||||
node_id = %node_id,
|
||||
node_name = %node_name,
|
||||
conn_id = conn_id,
|
||||
pool_size = pool_size,
|
||||
"proxy connected"
|
||||
);
|
||||
|
||||
self.notify_node_status(node_id, true, pool_size);
|
||||
}
|
||||
|
||||
pub fn unregister_proxy(&self, conn_id: u64, node_id: &str) {
|
||||
self.proxy_conns_by_id.remove(&conn_id);
|
||||
|
||||
let pool_size = {
|
||||
let mut map = self.proxy_conns.write();
|
||||
if let Some(conns) = map.get_mut(node_id) {
|
||||
conns.retain(|c| c.id != conn_id);
|
||||
if conns.is_empty() {
|
||||
map.remove(node_id);
|
||||
}
|
||||
}
|
||||
map.get(node_id).map(|v| v.len()).unwrap_or(0)
|
||||
};
|
||||
|
||||
info!(
|
||||
node_id = %node_id,
|
||||
conn_id = conn_id,
|
||||
remaining = pool_size,
|
||||
"proxy disconnected"
|
||||
);
|
||||
|
||||
self.cancel_streams_for_proxy(conn_id);
|
||||
self.notify_node_status(node_id.to_string(), pool_size > 0, pool_size);
|
||||
}
|
||||
|
||||
fn notify_node_status(&self, node_id: String, connected: bool, conn_count: usize) {
|
||||
let control_plane = self.control_plane.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(error) = control_plane
|
||||
.push_node_status(&node_id, connected, conn_count)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
node_id = %node_id,
|
||||
connected = connected,
|
||||
conn_count = conn_count,
|
||||
error = %error,
|
||||
"failed to push node status to app control plane"
|
||||
);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
fn get_proxy_conn(&self, node_id: &str) -> Option<Arc<ProxyConn>> {
|
||||
let map = self.proxy_conns.read();
|
||||
let conns = map.get(node_id)?;
|
||||
conns
|
||||
.iter()
|
||||
.filter(|c| c.is_available())
|
||||
.min_by_key(|c| c.stream_count.load(Ordering::Relaxed))
|
||||
.cloned()
|
||||
}
|
||||
|
||||
pub fn open_local_stream(
|
||||
&self,
|
||||
node_id: &str,
|
||||
meta: &protocol::RequestMeta,
|
||||
) -> Result<Arc<LocalStream>, String> {
|
||||
let proxy_conn = self
|
||||
.get_proxy_conn(node_id)
|
||||
.ok_or_else(|| format!("no proxy connection for node {node_id}"))?;
|
||||
let proxy_stream_id = proxy_conn
|
||||
.alloc_stream_id()
|
||||
.ok_or_else(|| format!("stream limit reached for node {node_id}"))?;
|
||||
|
||||
// Encode frames before registering the stream so that encoding failures
|
||||
// (practically impossible but theoretically possible) don't leak a stream
|
||||
// slot or orphan map entries.
|
||||
let meta_json = match serde_json::to_vec(meta) {
|
||||
Ok(json) => json,
|
||||
Err(e) => {
|
||||
proxy_conn.release_stream();
|
||||
return Err(format!("failed to encode request metadata: {e}"));
|
||||
}
|
||||
};
|
||||
let (meta_payload, meta_flags) = match protocol::compress_payload(&meta_json) {
|
||||
Ok(result) => result,
|
||||
Err(e) => {
|
||||
proxy_conn.release_stream();
|
||||
return Err(format!("failed to compress request metadata: {e}"));
|
||||
}
|
||||
};
|
||||
let header_frame = protocol::encode_frame(
|
||||
proxy_stream_id,
|
||||
protocol::REQUEST_HEADERS,
|
||||
meta_flags,
|
||||
&meta_payload,
|
||||
);
|
||||
|
||||
// Frames encoded successfully -- now register the stream.
|
||||
let local_stream_id = self.next_local_stream_id.fetch_add(1, Ordering::Relaxed);
|
||||
let local_stream = Arc::new(LocalStream::new(
|
||||
local_stream_id,
|
||||
proxy_conn.id,
|
||||
proxy_stream_id,
|
||||
));
|
||||
self.local_streams
|
||||
.insert(local_stream_id, local_stream.clone());
|
||||
self.proxy_to_local
|
||||
.insert((proxy_conn.id, proxy_stream_id), local_stream_id);
|
||||
|
||||
match proxy_conn.send(Message::Binary(header_frame.into())) {
|
||||
SendStatus::Queued => Ok(local_stream),
|
||||
SendStatus::Closed | SendStatus::Congested => {
|
||||
self.cleanup_local_stream(local_stream_id);
|
||||
proxy_conn.release_stream();
|
||||
Err("proxy connection congested".to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn push_local_request_body(
|
||||
&self,
|
||||
local_stream_id: u64,
|
||||
payload: Bytes,
|
||||
end_stream: bool,
|
||||
) -> Result<(), String> {
|
||||
let stream = self
|
||||
.local_streams
|
||||
.get(&local_stream_id)
|
||||
.map(|entry| entry.value().clone())
|
||||
.ok_or_else(|| "local stream not found".to_string())?;
|
||||
let proxy_conn = self
|
||||
.proxy_conns_by_id
|
||||
.get(&stream.proxy_conn_id)
|
||||
.map(|entry| entry.value().clone())
|
||||
.ok_or_else(|| "proxy connection unavailable".to_string())?;
|
||||
|
||||
let total_chunks = payload.len().div_ceil(MAX_REQUEST_BODY_FRAME_SIZE);
|
||||
if total_chunks == 0 {
|
||||
if end_stream {
|
||||
self.send_request_body_frame(&proxy_conn, stream.proxy_stream_id, &[], true)?;
|
||||
}
|
||||
} else {
|
||||
for (index, chunk) in payload.chunks(MAX_REQUEST_BODY_FRAME_SIZE).enumerate() {
|
||||
let is_last_chunk = index + 1 == total_chunks;
|
||||
self.send_request_body_frame(
|
||||
&proxy_conn,
|
||||
stream.proxy_stream_id,
|
||||
chunk,
|
||||
end_stream && is_last_chunk,
|
||||
)?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn send_request_body_frame(
|
||||
&self,
|
||||
proxy_conn: &Arc<ProxyConn>,
|
||||
proxy_stream_id: u32,
|
||||
payload: &[u8],
|
||||
end_stream: bool,
|
||||
) -> Result<(), String> {
|
||||
let (body_payload, body_flags) = protocol::compress_payload(payload)
|
||||
.map_err(|e| format!("failed to compress request body: {e}"))?;
|
||||
let body_frame = protocol::encode_frame(
|
||||
proxy_stream_id,
|
||||
protocol::REQUEST_BODY,
|
||||
body_flags
|
||||
| if end_stream {
|
||||
protocol::FLAG_END_STREAM
|
||||
} else {
|
||||
0
|
||||
},
|
||||
&body_payload,
|
||||
);
|
||||
match proxy_conn.send(Message::Binary(body_frame.into())) {
|
||||
SendStatus::Queued => Ok(()),
|
||||
SendStatus::Closed | SendStatus::Congested => {
|
||||
Err("proxy connection congested".to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn cancel_local_stream(&self, local_stream_id: u64, reason: &str) {
|
||||
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
||||
return;
|
||||
};
|
||||
|
||||
self.proxy_to_local
|
||||
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&stream.proxy_conn_id) {
|
||||
pc.release_stream();
|
||||
let frame = protocol::encode_stream_error(stream.proxy_stream_id, reason);
|
||||
let _ = pc.send(Message::Binary(frame.into()));
|
||||
}
|
||||
stream.fail(reason.to_string());
|
||||
}
|
||||
|
||||
fn cleanup_local_stream(&self, local_stream_id: u64) {
|
||||
let Some((_, stream)) = self.local_streams.remove(&local_stream_id) else {
|
||||
return;
|
||||
};
|
||||
self.proxy_to_local
|
||||
.remove(&(stream.proxy_conn_id, stream.proxy_stream_id));
|
||||
}
|
||||
|
||||
pub async fn handle_proxy_frame(&self, proxy_conn_id: u64, data: &mut [u8]) {
|
||||
let header = match protocol::FrameHeader::parse(data) {
|
||||
Some(h) => h,
|
||||
None => return,
|
||||
};
|
||||
let expected_len = protocol::HEADER_SIZE + header.payload_len as usize;
|
||||
if data.len() < expected_len {
|
||||
return;
|
||||
}
|
||||
|
||||
match header.msg_type {
|
||||
protocol::RESPONSE_HEADERS => {
|
||||
self.route_response_headers(proxy_conn_id, header, data);
|
||||
}
|
||||
protocol::RESPONSE_BODY => {
|
||||
self.route_response_body(proxy_conn_id, header, data);
|
||||
}
|
||||
protocol::STREAM_END => {
|
||||
self.finish_proxy_stream(proxy_conn_id, header.stream_id);
|
||||
}
|
||||
protocol::STREAM_ERROR => {
|
||||
let message = protocol::decode_payload(data, &header)
|
||||
.ok()
|
||||
.and_then(|payload| String::from_utf8(payload).ok())
|
||||
.unwrap_or_else(|| "stream error".to_string());
|
||||
self.fail_proxy_stream(proxy_conn_id, header.stream_id, message);
|
||||
}
|
||||
protocol::HEARTBEAT_DATA => {
|
||||
self.handle_heartbeat(proxy_conn_id, header.stream_id, data, &header)
|
||||
.await;
|
||||
}
|
||||
protocol::PING => {
|
||||
let payload = protocol::frame_payload_by_header(data, &header).unwrap_or(&[]);
|
||||
let pong = protocol::encode_pong(payload);
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
let _ = pc.send(Message::Binary(pong.into()));
|
||||
}
|
||||
}
|
||||
protocol::PONG => {}
|
||||
protocol::GOAWAY => {
|
||||
warn!(
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
"received GOAWAY from proxy connection"
|
||||
);
|
||||
}
|
||||
_ => {
|
||||
debug!(
|
||||
msg_type = header.msg_type,
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
"unexpected frame type from proxy"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn route_response_headers(
|
||||
&self,
|
||||
proxy_conn_id: u64,
|
||||
header: protocol::FrameHeader,
|
||||
data: &[u8],
|
||||
) {
|
||||
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
|
||||
return;
|
||||
};
|
||||
let Ok(payload) = protocol::decode_payload(data, &header) else {
|
||||
self.fail_proxy_stream(
|
||||
proxy_conn_id,
|
||||
header.stream_id,
|
||||
"failed to decode response headers",
|
||||
);
|
||||
return;
|
||||
};
|
||||
let Ok(meta) = serde_json::from_slice::<protocol::ResponseMeta>(&payload) else {
|
||||
self.fail_proxy_stream(
|
||||
proxy_conn_id,
|
||||
header.stream_id,
|
||||
"invalid response headers payload",
|
||||
);
|
||||
return;
|
||||
};
|
||||
if let Some(entry) = self.local_streams.get(&local_id) {
|
||||
entry.value().set_response_headers(meta);
|
||||
}
|
||||
}
|
||||
|
||||
fn route_response_body(&self, proxy_conn_id: u64, header: protocol::FrameHeader, data: &[u8]) {
|
||||
let Some(local_id) = self.lookup_local_stream(proxy_conn_id, header.stream_id) else {
|
||||
return;
|
||||
};
|
||||
let Ok(payload) = protocol::decode_payload(data, &header) else {
|
||||
self.fail_proxy_stream(
|
||||
proxy_conn_id,
|
||||
header.stream_id,
|
||||
"failed to decode response body",
|
||||
);
|
||||
return;
|
||||
};
|
||||
|
||||
let stream = match self.local_streams.get(&local_id) {
|
||||
Some(entry) => entry.value().clone(),
|
||||
None => return,
|
||||
};
|
||||
|
||||
if !stream.push_body_chunk(Bytes::from(payload)) {
|
||||
self.cancel_local_stream(local_id, "local relay response congested");
|
||||
}
|
||||
}
|
||||
|
||||
fn handle_stream_cleanup(
|
||||
&self,
|
||||
proxy_conn_id: u64,
|
||||
proxy_stream_id: u32,
|
||||
) -> Option<Arc<LocalStream>> {
|
||||
let local_id = self
|
||||
.proxy_to_local
|
||||
.remove(&(proxy_conn_id, proxy_stream_id))
|
||||
.map(|(_, local_id)| local_id)?;
|
||||
|
||||
let stream = self
|
||||
.local_streams
|
||||
.remove(&local_id)
|
||||
.map(|(_, stream)| stream)?;
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
pc.release_stream();
|
||||
}
|
||||
Some(stream)
|
||||
}
|
||||
|
||||
fn finish_proxy_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) {
|
||||
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
|
||||
stream.finish();
|
||||
}
|
||||
}
|
||||
|
||||
fn fail_proxy_stream(
|
||||
&self,
|
||||
proxy_conn_id: u64,
|
||||
proxy_stream_id: u32,
|
||||
error: impl Into<String>,
|
||||
) {
|
||||
if let Some(stream) = self.handle_stream_cleanup(proxy_conn_id, proxy_stream_id) {
|
||||
stream.fail(error.into());
|
||||
}
|
||||
}
|
||||
|
||||
fn lookup_local_stream(&self, proxy_conn_id: u64, proxy_stream_id: u32) -> Option<u64> {
|
||||
self.proxy_to_local
|
||||
.get(&(proxy_conn_id, proxy_stream_id))
|
||||
.map(|entry| *entry.value())
|
||||
}
|
||||
|
||||
async fn handle_heartbeat(
|
||||
&self,
|
||||
proxy_conn_id: u64,
|
||||
stream_id: u32,
|
||||
data: &[u8],
|
||||
header: &protocol::FrameHeader,
|
||||
) {
|
||||
let payload = match protocol::decode_payload(data, header) {
|
||||
Ok(payload) => payload,
|
||||
Err(error) => {
|
||||
warn!(proxy_conn_id = proxy_conn_id, error = %error, "failed to decode heartbeat payload");
|
||||
return;
|
||||
}
|
||||
};
|
||||
let ack_payload = match self.control_plane.heartbeat_ack(&payload).await {
|
||||
Ok(payload) => payload,
|
||||
Err(error) => {
|
||||
warn!(proxy_conn_id = proxy_conn_id, error = %error, "control-plane heartbeat callback failed");
|
||||
b"{}".to_vec()
|
||||
}
|
||||
};
|
||||
if let Some(pc) = self.proxy_conns_by_id.get(&proxy_conn_id) {
|
||||
let frame = protocol::encode_frame(stream_id, protocol::HEARTBEAT_ACK, 0, &ack_payload);
|
||||
let _ = pc.send(Message::Binary(frame.into()));
|
||||
}
|
||||
}
|
||||
|
||||
fn cancel_streams_for_proxy(&self, proxy_conn_id: u64) {
|
||||
let mut cancelled = 0usize;
|
||||
self.proxy_to_local.retain(|key, local_id| {
|
||||
if key.0 != proxy_conn_id {
|
||||
return true;
|
||||
}
|
||||
if let Some((_, stream)) = self.local_streams.remove(local_id) {
|
||||
stream.fail("proxy disconnected".to_string());
|
||||
}
|
||||
cancelled += 1;
|
||||
false
|
||||
});
|
||||
|
||||
if cancelled > 0 {
|
||||
warn!(
|
||||
proxy_conn_id = proxy_conn_id,
|
||||
streams_cancelled = cancelled,
|
||||
"cancelled in-flight streams due to proxy disconnect"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn stats(&self) -> HubStats {
|
||||
let proxy_conns = self.proxy_conns.read();
|
||||
let total_proxy = proxy_conns.values().map(|v| v.len()).sum();
|
||||
let nodes = proxy_conns.len();
|
||||
drop(proxy_conns);
|
||||
|
||||
HubStats {
|
||||
proxy_connections: total_proxy,
|
||||
nodes,
|
||||
active_streams: self.local_streams.len(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(serde::Serialize)]
|
||||
pub struct HubStats {
|
||||
pub proxy_connections: usize,
|
||||
pub nodes: usize,
|
||||
pub active_streams: usize,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn build_meta() -> protocol::RequestMeta {
|
||||
protocol::RequestMeta {
|
||||
method: "GET".to_string(),
|
||||
url: "https://example.com".to_string(),
|
||||
headers: HashMap::new(),
|
||||
timeout: 30,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn cancel_local_stream_notifies_proxy() {
|
||||
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||
|
||||
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let proxy = Arc::new(ProxyConn::new(
|
||||
100,
|
||||
"node-1".to_string(),
|
||||
"Node 1".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
));
|
||||
hub.register_proxy(proxy);
|
||||
|
||||
let stream = hub
|
||||
.open_local_stream("node-1", &build_meta())
|
||||
.expect("open local stream");
|
||||
let _ = proxy_rx.try_recv().expect("headers frame");
|
||||
hub.push_local_request_body(stream.id, Bytes::new(), true)
|
||||
.expect("finish empty body");
|
||||
let _ = proxy_rx.try_recv().expect("body frame");
|
||||
|
||||
hub.cancel_local_stream(stream.id, "client dropped");
|
||||
|
||||
let cancelled = proxy_rx.try_recv().expect("cancel frame");
|
||||
let cancelled_data = match cancelled {
|
||||
Message::Binary(data) => data.to_vec(),
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
let header = protocol::FrameHeader::parse(&cancelled_data).expect("cancel frame header");
|
||||
assert_eq!(header.msg_type, protocol::STREAM_ERROR);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn push_local_request_body_splits_large_payload_and_marks_end() {
|
||||
let hub = HubRouter::new(ControlPlaneClient::disabled());
|
||||
|
||||
let (proxy_tx, mut proxy_rx) = mpsc::channel(8);
|
||||
let (proxy_close_tx, _) = watch::channel(false);
|
||||
let proxy = Arc::new(ProxyConn::new(
|
||||
200,
|
||||
"node-2".to_string(),
|
||||
"Node 2".to_string(),
|
||||
proxy_tx,
|
||||
proxy_close_tx,
|
||||
16,
|
||||
));
|
||||
hub.register_proxy(proxy);
|
||||
|
||||
let stream = hub
|
||||
.open_local_stream("node-2", &build_meta())
|
||||
.expect("open local stream");
|
||||
let _ = proxy_rx.try_recv().expect("headers frame");
|
||||
|
||||
let payload = Bytes::from(vec![b'x'; MAX_REQUEST_BODY_FRAME_SIZE + 17]);
|
||||
hub.push_local_request_body(stream.id, payload, true)
|
||||
.expect("push request body");
|
||||
|
||||
let first = match proxy_rx.try_recv().expect("first body frame") {
|
||||
Message::Binary(data) => data.to_vec(),
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
let first_header = protocol::FrameHeader::parse(&first).expect("first body header");
|
||||
assert_eq!(first_header.msg_type, protocol::REQUEST_BODY);
|
||||
assert_eq!(first_header.flags & protocol::FLAG_END_STREAM, 0);
|
||||
|
||||
let second = match proxy_rx.try_recv().expect("second body frame") {
|
||||
Message::Binary(data) => data.to_vec(),
|
||||
other => panic!("unexpected message: {other:?}"),
|
||||
};
|
||||
let second_header = protocol::FrameHeader::parse(&second).expect("second body header");
|
||||
assert_eq!(second_header.msg_type, protocol::REQUEST_BODY);
|
||||
assert_ne!(second_header.flags & protocol::FLAG_END_STREAM, 0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,262 @@
|
||||
use std::io;
|
||||
use std::net::SocketAddr;
|
||||
use std::time::Duration;
|
||||
|
||||
use async_stream::stream;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::extract::{ConnectInfo, Path, Request, State};
|
||||
use axum::http::{HeaderMap, HeaderName, HeaderValue, Response, StatusCode};
|
||||
use axum::response::IntoResponse;
|
||||
use bytes::BytesMut;
|
||||
use futures_util::StreamExt;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::hub::{LocalBodyEvent, LocalStream};
|
||||
use crate::protocol;
|
||||
use crate::AppState;
|
||||
|
||||
pub const TUNNEL_ERROR_HEADER: &str = "x-aether-tunnel-error";
|
||||
const MAX_RELAY_META_LEN: usize = 256 * 1024;
|
||||
|
||||
struct StreamGuard {
|
||||
hub: std::sync::Arc<crate::hub::HubRouter>,
|
||||
stream_id: u64,
|
||||
finished: bool,
|
||||
}
|
||||
|
||||
impl Drop for StreamGuard {
|
||||
fn drop(&mut self) {
|
||||
if !self.finished {
|
||||
self.hub
|
||||
.cancel_local_stream(self.stream_id, "local relay client dropped");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn relay_request(
|
||||
Path(node_id): Path<String>,
|
||||
State(state): State<AppState>,
|
||||
ConnectInfo(addr): ConnectInfo<SocketAddr>,
|
||||
request: Request,
|
||||
) -> impl IntoResponse {
|
||||
if !addr.ip().is_loopback() {
|
||||
return tunnel_error_response(
|
||||
StatusCode::FORBIDDEN,
|
||||
"forbidden",
|
||||
"local relay only accepts loopback requests",
|
||||
);
|
||||
}
|
||||
|
||||
let mut body_stream = request.into_body().into_data_stream();
|
||||
let mut envelope_buf = BytesMut::new();
|
||||
let mut meta: Option<protocol::RequestMeta> = None;
|
||||
let mut stream: Option<std::sync::Arc<LocalStream>> = None;
|
||||
|
||||
while let Some(chunk_result) = body_stream.next().await {
|
||||
let chunk = match chunk_result {
|
||||
Ok(chunk) => chunk,
|
||||
Err(error) => {
|
||||
if let Some(active_stream) = &stream {
|
||||
state
|
||||
.hub
|
||||
.cancel_local_stream(active_stream.id, "failed to read relay request body");
|
||||
}
|
||||
warn!(error = %error, "failed to read local relay request body");
|
||||
return tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"failed to read relay request body",
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if stream.is_none() {
|
||||
envelope_buf.extend_from_slice(&chunk);
|
||||
let Some((parsed_meta, body_offset)) = (match try_decode_envelope_meta(&envelope_buf) {
|
||||
Ok(result) => result,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(StatusCode::BAD_REQUEST, "bad_request", &error);
|
||||
}
|
||||
}) else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let opened_stream = match state.hub.open_local_stream(&node_id, &parsed_meta) {
|
||||
Ok(stream) => stream,
|
||||
Err(error) => {
|
||||
return tunnel_error_response(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"connect",
|
||||
&error,
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if envelope_buf.len() > body_offset {
|
||||
let first_body_chunk = Bytes::copy_from_slice(&envelope_buf[body_offset..]);
|
||||
if let Err(error) =
|
||||
state
|
||||
.hub
|
||||
.push_local_request_body(opened_stream.id, first_body_chunk, false)
|
||||
{
|
||||
state.hub.cancel_local_stream(opened_stream.id, &error);
|
||||
return tunnel_error_response(
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"connect",
|
||||
&error,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
envelope_buf.clear();
|
||||
meta = Some(parsed_meta);
|
||||
stream = Some(opened_stream);
|
||||
continue;
|
||||
}
|
||||
|
||||
let Some(active_stream) = &stream else {
|
||||
continue;
|
||||
};
|
||||
if let Err(error) = state
|
||||
.hub
|
||||
.push_local_request_body(active_stream.id, chunk, false)
|
||||
{
|
||||
state.hub.cancel_local_stream(active_stream.id, &error);
|
||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||
}
|
||||
}
|
||||
|
||||
let (meta, stream) = match (meta, stream) {
|
||||
(Some(meta), Some(stream)) => (meta, stream),
|
||||
_ => {
|
||||
return tunnel_error_response(
|
||||
StatusCode::BAD_REQUEST,
|
||||
"bad_request",
|
||||
"relay envelope metadata truncated",
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
if let Err(error) = state
|
||||
.hub
|
||||
.push_local_request_body(stream.id, Bytes::new(), true)
|
||||
{
|
||||
state.hub.cancel_local_stream(stream.id, &error);
|
||||
return tunnel_error_response(StatusCode::SERVICE_UNAVAILABLE, "connect", &error);
|
||||
}
|
||||
|
||||
let request_guard = StreamGuard {
|
||||
hub: state.hub.clone(),
|
||||
stream_id: stream.id,
|
||||
finished: false,
|
||||
};
|
||||
|
||||
let wait_timeout = Duration::from_secs(meta.timeout.clamp(5, 300));
|
||||
let response_head = match stream.wait_headers(wait_timeout).await {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
state.hub.cancel_local_stream(stream.id, &error);
|
||||
return tunnel_error_response(StatusCode::GATEWAY_TIMEOUT, "timeout", &error);
|
||||
}
|
||||
};
|
||||
|
||||
let Some(mut body_rx) = stream.take_body_receiver() else {
|
||||
state
|
||||
.hub
|
||||
.cancel_local_stream(stream.id, "missing relay response body receiver");
|
||||
return tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"missing relay response body receiver",
|
||||
);
|
||||
};
|
||||
|
||||
let hub = state.hub.clone();
|
||||
let stream_id = stream.id;
|
||||
let body_stream = stream! {
|
||||
let mut guard = request_guard;
|
||||
guard.hub = hub;
|
||||
guard.stream_id = stream_id;
|
||||
while let Some(event) = body_rx.recv().await {
|
||||
match event {
|
||||
LocalBodyEvent::Chunk(chunk) => yield Ok::<Bytes, io::Error>(chunk),
|
||||
LocalBodyEvent::End => {
|
||||
guard.finished = true;
|
||||
break;
|
||||
}
|
||||
LocalBodyEvent::Error(error) => {
|
||||
guard.finished = true;
|
||||
yield Err(io::Error::other(error));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
guard.finished = true;
|
||||
};
|
||||
|
||||
let mut builder = Response::builder().status(response_head.status);
|
||||
if let Some(headers) = builder.headers_mut() {
|
||||
append_headers(headers, &response_head.headers);
|
||||
}
|
||||
match builder.body(Body::from_stream(body_stream)) {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
warn!(error = %error, "failed to build relay response");
|
||||
tunnel_error_response(
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"relay",
|
||||
"failed to build relay response",
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn try_decode_envelope_meta(
|
||||
buffer: &BytesMut,
|
||||
) -> Result<Option<(protocol::RequestMeta, usize)>, String> {
|
||||
if buffer.len() < 4 {
|
||||
return Ok(None);
|
||||
}
|
||||
let meta_len = u32::from_be_bytes([buffer[0], buffer[1], buffer[2], buffer[3]]) as usize;
|
||||
if meta_len > MAX_RELAY_META_LEN {
|
||||
return Err("relay metadata too large".to_string());
|
||||
}
|
||||
let meta_end = 4usize
|
||||
.checked_add(meta_len)
|
||||
.ok_or_else(|| "relay envelope length overflow".to_string())?;
|
||||
if buffer.len() < meta_end {
|
||||
return Ok(None);
|
||||
}
|
||||
let meta = serde_json::from_slice::<protocol::RequestMeta>(&buffer[4..meta_end])
|
||||
.map_err(|e| format!("invalid relay metadata: {e}"))?;
|
||||
Ok(Some((meta, meta_end)))
|
||||
}
|
||||
|
||||
fn append_headers(target: &mut HeaderMap, headers: &[(String, String)]) {
|
||||
for (name, value) in headers {
|
||||
let Ok(name) = HeaderName::from_bytes(name.as_bytes()) else {
|
||||
continue;
|
||||
};
|
||||
let Ok(value) = HeaderValue::from_str(value) else {
|
||||
continue;
|
||||
};
|
||||
target.append(name, value);
|
||||
}
|
||||
}
|
||||
|
||||
fn tunnel_error_response(status: StatusCode, kind: &str, message: &str) -> Response<Body> {
|
||||
let mut builder = Response::builder().status(status);
|
||||
if let Some(headers) = builder.headers_mut() {
|
||||
headers.insert(
|
||||
HeaderName::from_static(TUNNEL_ERROR_HEADER),
|
||||
HeaderValue::from_str(kind).unwrap_or_else(|_| HeaderValue::from_static("relay")),
|
||||
);
|
||||
headers.insert(
|
||||
axum::http::header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("text/plain; charset=utf-8"),
|
||||
);
|
||||
}
|
||||
builder
|
||||
.body(Body::from(message.to_string()))
|
||||
.unwrap_or_else(|_| Response::new(Body::from("relay error")))
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
mod control_plane;
|
||||
mod hub;
|
||||
mod local_relay;
|
||||
mod protocol;
|
||||
mod proxy_conn;
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::WebSocketUpgrade;
|
||||
use axum::extract::State;
|
||||
use axum::response::{IntoResponse, Json};
|
||||
use axum::routing::{get, post};
|
||||
use axum::Router;
|
||||
use clap::Parser;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use crate::control_plane::ControlPlaneClient;
|
||||
use crate::hub::{ConnConfig, HubRouter};
|
||||
use crate::local_relay::relay_request;
|
||||
|
||||
#[derive(Parser, Debug)]
|
||||
#[command(name = "aether-hub", about = "Tunnel Hub for Aether")]
|
||||
struct Args {
|
||||
/// Bind address
|
||||
#[arg(long, default_value = "0.0.0.0:8085", env = "TUNNEL_HUB_BIND")]
|
||||
bind: String,
|
||||
|
||||
/// Proxy-side idle timeout in seconds (0 to disable)
|
||||
#[arg(long, default_value_t = 0, env = "TUNNEL_HUB_PROXY_IDLE_TIMEOUT")]
|
||||
proxy_idle_timeout: u64,
|
||||
|
||||
/// Ping interval in seconds (for both sides)
|
||||
#[arg(long, default_value_t = 15, env = "TUNNEL_HUB_PING_INTERVAL")]
|
||||
ping_interval: u64,
|
||||
|
||||
/// Max concurrent streams per proxy connection
|
||||
#[arg(long, default_value_t = 2048, env = "TUNNEL_HUB_MAX_STREAMS")]
|
||||
max_streams: usize,
|
||||
|
||||
/// Per-connection outbound queue capacity before treating the socket as congested
|
||||
#[arg(
|
||||
long,
|
||||
default_value_t = 128,
|
||||
env = "TUNNEL_HUB_OUTBOUND_QUEUE_CAPACITY"
|
||||
)]
|
||||
outbound_queue_capacity: usize,
|
||||
|
||||
/// Local Aether app base URL for control-plane callbacks
|
||||
#[arg(
|
||||
long,
|
||||
default_value = "http://127.0.0.1:8084",
|
||||
env = "TUNNEL_HUB_APP_BASE_URL"
|
||||
)]
|
||||
app_base_url: String,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct AppState {
|
||||
pub hub: std::sync::Arc<HubRouter>,
|
||||
pub proxy_conn_cfg: ConnConfig,
|
||||
pub max_streams: usize,
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
// Initialize tracing
|
||||
tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| "aether_hub=info".into()),
|
||||
)
|
||||
.init();
|
||||
|
||||
let args = Args::parse();
|
||||
|
||||
let hub = HubRouter::new(ControlPlaneClient::new(args.app_base_url));
|
||||
let outbound_queue_capacity = args.outbound_queue_capacity.clamp(8, 4096);
|
||||
let ping_interval = Duration::from_secs(args.ping_interval);
|
||||
let state = AppState {
|
||||
hub,
|
||||
proxy_conn_cfg: ConnConfig {
|
||||
ping_interval,
|
||||
idle_timeout: Duration::from_secs(args.proxy_idle_timeout),
|
||||
outbound_queue_capacity,
|
||||
},
|
||||
max_streams: args.max_streams,
|
||||
};
|
||||
|
||||
let app = Router::new()
|
||||
.route("/health", get(health))
|
||||
.route("/stats", get(stats))
|
||||
.route("/proxy", get(ws_proxy))
|
||||
.route("/local/relay/{node_id}", post(relay_request))
|
||||
.with_state(state);
|
||||
|
||||
let listener = tokio::net::TcpListener::bind(&args.bind).await?;
|
||||
info!(bind = %args.bind, "aether-hub started");
|
||||
|
||||
axum::serve(
|
||||
listener,
|
||||
app.into_make_service_with_connect_info::<SocketAddr>(),
|
||||
)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// HTTP endpoints
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
async fn health() -> impl IntoResponse {
|
||||
Json(serde_json::json!({"status": "ok"}))
|
||||
}
|
||||
|
||||
async fn stats(State(state): State<AppState>) -> impl IntoResponse {
|
||||
Json(state.hub.stats())
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// WebSocket endpoints
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
async fn ws_proxy(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<AppState>,
|
||||
headers: axum::http::HeaderMap,
|
||||
) -> impl IntoResponse {
|
||||
let node_id = headers
|
||||
.get("x-node-id")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.unwrap_or("")
|
||||
.trim()
|
||||
.to_string();
|
||||
|
||||
let node_name = headers
|
||||
.get("x-node-name")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.unwrap_or(&node_id)
|
||||
.trim()
|
||||
.to_string();
|
||||
|
||||
let max_streams: usize = headers
|
||||
.get("x-tunnel-max-streams")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| v.parse().ok())
|
||||
.unwrap_or(state.max_streams)
|
||||
.clamp(64, 2048);
|
||||
|
||||
if node_id.is_empty() {
|
||||
warn!("proxy connection rejected: missing X-Node-ID header");
|
||||
return axum::http::StatusCode::BAD_REQUEST.into_response();
|
||||
}
|
||||
|
||||
ws.max_frame_size(64 * 1024 * 1024)
|
||||
.on_upgrade(move |socket| {
|
||||
proxy_conn::handle_proxy_connection(
|
||||
socket,
|
||||
state.hub,
|
||||
node_id,
|
||||
node_name,
|
||||
max_streams,
|
||||
state.proxy_conn_cfg,
|
||||
)
|
||||
})
|
||||
.into_response()
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
/// Tunnel binary frame protocol
|
||||
///
|
||||
/// Frame format (10-byte header + payload):
|
||||
/// | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
|
||||
use std::io::Read;
|
||||
|
||||
use flate2::read::GzDecoder;
|
||||
use flate2::write::GzEncoder;
|
||||
use flate2::Compression;
|
||||
|
||||
pub const HEADER_SIZE: usize = 10;
|
||||
|
||||
// Message types
|
||||
pub const REQUEST_HEADERS: u8 = 0x01;
|
||||
pub const REQUEST_BODY: u8 = 0x02;
|
||||
pub const RESPONSE_HEADERS: u8 = 0x03;
|
||||
pub const RESPONSE_BODY: u8 = 0x04;
|
||||
pub const STREAM_END: u8 = 0x05;
|
||||
pub const STREAM_ERROR: u8 = 0x06;
|
||||
pub const PING: u8 = 0x10;
|
||||
pub const PONG: u8 = 0x11;
|
||||
pub const GOAWAY: u8 = 0x12;
|
||||
pub const HEARTBEAT_DATA: u8 = 0x13;
|
||||
pub const HEARTBEAT_ACK: u8 = 0x14;
|
||||
// Flags
|
||||
pub const FLAG_END_STREAM: u8 = 0x01;
|
||||
pub const FLAG_GZIP_COMPRESSED: u8 = 0x02;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub struct FrameHeader {
|
||||
pub stream_id: u32,
|
||||
pub msg_type: u8,
|
||||
pub flags: u8,
|
||||
pub payload_len: u32,
|
||||
}
|
||||
|
||||
impl FrameHeader {
|
||||
/// Parse frame header from raw bytes (must be >= HEADER_SIZE)
|
||||
#[inline]
|
||||
pub fn parse(data: &[u8]) -> Option<Self> {
|
||||
if data.len() < HEADER_SIZE {
|
||||
return None;
|
||||
}
|
||||
Some(Self {
|
||||
stream_id: u32::from_be_bytes([data[0], data[1], data[2], data[3]]),
|
||||
msg_type: data[4],
|
||||
flags: data[5],
|
||||
payload_len: u32::from_be_bytes([data[6], data[7], data[8], data[9]]),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct RequestMeta {
|
||||
pub method: String,
|
||||
pub url: String,
|
||||
pub headers: std::collections::HashMap<String, String>,
|
||||
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
|
||||
pub timeout: u64,
|
||||
}
|
||||
|
||||
fn default_timeout() -> u64 {
|
||||
60
|
||||
}
|
||||
|
||||
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
#[derive(serde::Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum TimeoutValue {
|
||||
Int(u64),
|
||||
Float(f64),
|
||||
}
|
||||
|
||||
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
|
||||
TimeoutValue::Int(v) => Ok(v),
|
||||
TimeoutValue::Float(v) => {
|
||||
if !v.is_finite() || v < 0.0 {
|
||||
return Err(serde::de::Error::custom(
|
||||
"timeout must be a non-negative finite number",
|
||||
));
|
||||
}
|
||||
if v.fract() != 0.0 {
|
||||
return Err(serde::de::Error::custom("timeout must be integer seconds"));
|
||||
}
|
||||
if v > (u64::MAX as f64) {
|
||||
return Err(serde::de::Error::custom("timeout is too large"));
|
||||
}
|
||||
Ok(v as u64)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ResponseMeta {
|
||||
pub status: u16,
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
pub fn encode_frame(stream_id: u32, msg_type: u8, flags: u8, payload: &[u8]) -> Vec<u8> {
|
||||
let mut buf = Vec::with_capacity(HEADER_SIZE + payload.len());
|
||||
buf.extend_from_slice(&stream_id.to_be_bytes());
|
||||
buf.push(msg_type);
|
||||
buf.push(flags);
|
||||
buf.extend_from_slice(&(payload.len() as u32).to_be_bytes());
|
||||
buf.extend_from_slice(payload);
|
||||
buf
|
||||
}
|
||||
|
||||
/// Encode a STREAM_ERROR frame for a given stream_id with an error message
|
||||
pub fn encode_stream_error(stream_id: u32, msg: &str) -> Vec<u8> {
|
||||
encode_frame(stream_id, STREAM_ERROR, 0, msg.as_bytes())
|
||||
}
|
||||
|
||||
/// Encode a PING frame (stream_id=0)
|
||||
pub fn encode_ping() -> Vec<u8> {
|
||||
encode_frame(0, PING, 0, &[])
|
||||
}
|
||||
|
||||
/// Encode a PONG frame (stream_id=0, echo payload)
|
||||
pub fn encode_pong(payload: &[u8]) -> Vec<u8> {
|
||||
encode_frame(0, PONG, 0, payload)
|
||||
}
|
||||
|
||||
/// Encode a GOAWAY frame (stream_id=0)
|
||||
pub fn encode_goaway() -> Vec<u8> {
|
||||
encode_frame(0, GOAWAY, 0, &[])
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn frame_payload_by_header<'a>(data: &'a [u8], header: &FrameHeader) -> Option<&'a [u8]> {
|
||||
let payload_len = header.payload_len as usize;
|
||||
let end = HEADER_SIZE.checked_add(payload_len)?;
|
||||
if data.len() < end {
|
||||
return None;
|
||||
}
|
||||
Some(&data[HEADER_SIZE..end])
|
||||
}
|
||||
|
||||
pub fn decode_payload(data: &[u8], header: &FrameHeader) -> Result<Vec<u8>, String> {
|
||||
let payload = frame_payload_by_header(data, header)
|
||||
.ok_or_else(|| "incomplete frame payload".to_string())?;
|
||||
if header.flags & FLAG_GZIP_COMPRESSED != 0 {
|
||||
let mut decoder = GzDecoder::new(payload);
|
||||
let mut decoded = Vec::new();
|
||||
decoder
|
||||
.read_to_end(&mut decoded)
|
||||
.map_err(|e| format!("failed to decompress payload: {e}"))?;
|
||||
Ok(decoded)
|
||||
} else {
|
||||
Ok(payload.to_vec())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn compress_payload(payload: &[u8]) -> Result<(Vec<u8>, u8), std::io::Error> {
|
||||
maybe_recompress_payload(payload, true)
|
||||
}
|
||||
|
||||
fn maybe_recompress_payload(
|
||||
payload: &[u8],
|
||||
prefer_gzip: bool,
|
||||
) -> Result<(Vec<u8>, u8), std::io::Error> {
|
||||
if !prefer_gzip {
|
||||
return Ok((payload.to_vec(), 0));
|
||||
}
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
std::io::Write::write_all(&mut encoder, payload)?;
|
||||
let compressed = encoder.finish()?;
|
||||
if compressed.len() < payload.len() {
|
||||
Ok((compressed, FLAG_GZIP_COMPRESSED))
|
||||
} else {
|
||||
Ok((payload.to_vec(), 0))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
/// Proxy-side WebSocket connection handler
|
||||
///
|
||||
/// Handles the lifecycle of a single aether-proxy connection:
|
||||
/// accept -> authenticate (headers) -> read loop -> cleanup
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::{Message, WebSocket};
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use tokio::sync::{mpsc, watch};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use crate::hub::{ConnConfig, HubRouter, ProxyConn, SendStatus};
|
||||
use crate::protocol;
|
||||
|
||||
/// Maximum single frame size: 64 MB
|
||||
const MAX_FRAME_SIZE: usize = 64 * 1024 * 1024;
|
||||
|
||||
pub async fn handle_proxy_connection(
|
||||
ws: WebSocket,
|
||||
hub: Arc<HubRouter>,
|
||||
node_id: String,
|
||||
node_name: String,
|
||||
max_streams: usize,
|
||||
cfg: ConnConfig,
|
||||
) {
|
||||
let conn_id = hub.alloc_conn_id();
|
||||
let (mut ws_tx, ws_rx) = ws.split();
|
||||
|
||||
let (tx, mut rx) = mpsc::channel::<Message>(cfg.outbound_queue_capacity);
|
||||
let (close_tx, mut close_rx) = watch::channel(false);
|
||||
|
||||
let conn = Arc::new(ProxyConn::new(
|
||||
conn_id,
|
||||
node_id.clone(),
|
||||
node_name.clone(),
|
||||
tx,
|
||||
close_tx,
|
||||
max_streams,
|
||||
));
|
||||
|
||||
hub.register_proxy(conn.clone());
|
||||
|
||||
let writer = tokio::spawn(async move {
|
||||
loop {
|
||||
tokio::select! {
|
||||
msg = rx.recv() => match msg {
|
||||
Some(msg) => {
|
||||
if ws_tx.send(msg).await.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
None => break,
|
||||
},
|
||||
changed = close_rx.changed() => {
|
||||
if changed.is_err() || *close_rx.borrow() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
let _ = ws_tx.close().await;
|
||||
});
|
||||
|
||||
let ping_conn = conn.clone();
|
||||
let ping_interval = cfg.ping_interval;
|
||||
let ping_task = tokio::spawn(async move {
|
||||
loop {
|
||||
tokio::time::sleep(ping_interval).await;
|
||||
let ping = protocol::encode_ping();
|
||||
if !matches!(
|
||||
ping_conn.send(Message::Binary(ping.into())),
|
||||
SendStatus::Queued
|
||||
) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
let reader_hub = hub.clone();
|
||||
let reader_conn = conn.clone();
|
||||
let reader = tokio::spawn(async move {
|
||||
run_proxy_reader(ws_rx, reader_hub, reader_conn, cfg.idle_timeout).await;
|
||||
});
|
||||
|
||||
let _ = reader.await;
|
||||
ping_task.abort();
|
||||
conn.request_close();
|
||||
hub.unregister_proxy(conn_id, &node_id);
|
||||
drop(conn);
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
writer.abort();
|
||||
let _ = writer.await;
|
||||
}
|
||||
|
||||
async fn run_proxy_reader(
|
||||
mut ws_rx: futures_util::stream::SplitStream<WebSocket>,
|
||||
hub: Arc<HubRouter>,
|
||||
conn: Arc<ProxyConn>,
|
||||
idle_timeout: Duration,
|
||||
) {
|
||||
let idle_enabled = !idle_timeout.is_zero();
|
||||
let mut oversized_count = 0u32;
|
||||
loop {
|
||||
let msg = if idle_enabled {
|
||||
tokio::select! {
|
||||
msg = ws_rx.next() => msg,
|
||||
_ = tokio::time::sleep(idle_timeout) => {
|
||||
warn!(conn_id = conn.id, node_id = %conn.node_id, "proxy idle timeout");
|
||||
let _ = conn.send(Message::Binary(protocol::encode_goaway().into()));
|
||||
conn.request_close();
|
||||
break;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
ws_rx.next().await
|
||||
};
|
||||
|
||||
match msg {
|
||||
Some(Ok(Message::Binary(data))) => {
|
||||
let mut data = data.to_vec();
|
||||
if data.len() > MAX_FRAME_SIZE {
|
||||
oversized_count += 1;
|
||||
warn!(
|
||||
conn_id = conn.id,
|
||||
size = data.len(),
|
||||
"oversized frame from proxy"
|
||||
);
|
||||
if oversized_count >= 5 {
|
||||
warn!(conn_id = conn.id, "too many oversized frames, closing");
|
||||
conn.request_close();
|
||||
break;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
oversized_count = 0;
|
||||
|
||||
if data.len() < protocol::HEADER_SIZE {
|
||||
debug!(conn_id = conn.id, "frame too small, skipping");
|
||||
continue;
|
||||
}
|
||||
|
||||
hub.handle_proxy_frame(conn.id, &mut data).await;
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => {
|
||||
info!(conn_id = conn.id, node_id = %conn.node_id, "proxy WebSocket closed");
|
||||
break;
|
||||
}
|
||||
Some(Err(e)) => {
|
||||
warn!(conn_id = conn.id, error = %e, "proxy WebSocket error");
|
||||
break;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,8 @@
|
||||
# Aether server URL
|
||||
AETHER_PROXY_AETHER_URL=https://aether.example.com
|
||||
|
||||
# Management Token (ae_xxx, must belong to an ADMIN user)
|
||||
AETHER_PROXY_MANAGEMENT_TOKEN=ae_xxxxx
|
||||
|
||||
# Node identification
|
||||
AETHER_PROXY_NODE_NAME=proxy-01
|
||||
Generated
+3403
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,44 @@
|
||||
[package]
|
||||
name = "aether-proxy"
|
||||
version = "0.2.5"
|
||||
edition = "2021"
|
||||
description = "Tunnel proxy for Aether"
|
||||
|
||||
[dependencies]
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "stream", "http2"] }
|
||||
hyper = { version = "1", features = ["client", "http1", "http2"] }
|
||||
hyper-util = { version = "0.1", features = ["client", "client-legacy", "http1", "http2", "tokio"] }
|
||||
http-body-util = "0.1"
|
||||
tokio-tungstenite = { version = "0.24", features = ["rustls-tls-webpki-roots"] }
|
||||
tokio-rustls = "0.26"
|
||||
futures-util = "0.3"
|
||||
base64 = "0.22"
|
||||
clap = { version = "4", features = ["derive", "env"] }
|
||||
tracing = "0.1"
|
||||
tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] }
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
thiserror = "2"
|
||||
bytes = "1"
|
||||
sha2 = "0.10"
|
||||
hex = "0.4"
|
||||
anyhow = "1"
|
||||
arc-swap = "1"
|
||||
toml = "0.8"
|
||||
rustls = { version = "0.23", features = ["ring"] }
|
||||
ratatui = "0.30"
|
||||
crossterm = "0.28"
|
||||
url = "2"
|
||||
sysinfo = "0.32"
|
||||
libc = "0.2"
|
||||
flate2 = "1"
|
||||
tar = "0.4"
|
||||
socket2 = { version = "0.5", features = ["all"] }
|
||||
tower-service = "0.3"
|
||||
webpki-roots = "0.26"
|
||||
|
||||
[profile.release]
|
||||
lto = true
|
||||
strip = true
|
||||
codegen-units = 1
|
||||
@@ -0,0 +1,10 @@
|
||||
FROM debian:bookworm-slim
|
||||
|
||||
ARG TARGETARCH
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
COPY build/linux-${TARGETARCH}/aether-proxy /usr/local/bin/aether-proxy
|
||||
|
||||
ENTRYPOINT ["aether-proxy"]
|
||||
@@ -0,0 +1,154 @@
|
||||
# aether-proxy
|
||||
|
||||
Aether Tunnel 代理节点,部署在海外 VPS 上,通过 WebSocket 隧道为 Aether 实例中转 API 流量。
|
||||
|
||||
Tunnel 模式下代理节点**无需对外监听端口**,仅需出站连接到 Aether 服务器。
|
||||
|
||||
## 安装
|
||||
|
||||
### Docker Compose 部署
|
||||
|
||||
```bash
|
||||
cp .env.example .env
|
||||
# 编辑 .env 填入 AETHER_PROXY_AETHER_URL 和 AETHER_PROXY_MANAGEMENT_TOKEN
|
||||
docker compose up -d
|
||||
```
|
||||
|
||||
### 下载预编译二进制
|
||||
|
||||
<!-- DOWNLOAD_TABLE_START -->
|
||||
| Platform | Download |
|
||||
|----------|----------|
|
||||
| Linux x86_64 | [aether-proxy-linux-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-linux-amd64.tar.gz) |
|
||||
| Linux ARM64 | [aether-proxy-linux-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-linux-arm64.tar.gz) |
|
||||
| macOS x86_64 | [aether-proxy-macos-amd64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-macos-amd64.tar.gz) |
|
||||
| macOS ARM64 | [aether-proxy-macos-arm64.tar.gz](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-macos-arm64.tar.gz) |
|
||||
| Windows x86_64 | [aether-proxy-windows-amd64.zip](https://github.com/fawney19/Aether/releases/download/proxy-v0.2.5/aether-proxy-windows-amd64.zip) |
|
||||
<!-- DOWNLOAD_TABLE_END -->
|
||||
|
||||
## 快速开始
|
||||
|
||||
```bash
|
||||
# 1. 首次安装配置(TUI 向导,勾选 Install Service 随系统启动服务)
|
||||
sudo ./aether-proxy setup
|
||||
|
||||
# 2. 日常管理 (勾选 Install Service 作为系统服务的情况下)
|
||||
aether-proxy status # 看状态
|
||||
aether-proxy logs # 看日志
|
||||
|
||||
sudo aether-proxy start # 启动服务
|
||||
sudo aether-proxy stop # 停止服务
|
||||
sudo aether-proxy restart # 重启服务
|
||||
|
||||
# 3. 重新配置(改完自动重启服务)
|
||||
sudo aether-proxy setup
|
||||
|
||||
# 4. 彻底卸载
|
||||
sudo aether-proxy uninstall
|
||||
```
|
||||
|
||||
完成向导后, 配置自动保存到 `aether-proxy.toml`,如果启用了 Install Service,将自动注册并启动 systemd 服务。
|
||||
|
||||
### 直接运行
|
||||
|
||||
如果不需要安装为系统服务,可以直接运行。缺少必填参数时会自动进入 setup 向导:
|
||||
|
||||
```bash
|
||||
./aether-proxy
|
||||
```
|
||||
|
||||
## 配置
|
||||
|
||||
配置按以下优先级加载(高优先级覆盖低优先级):
|
||||
|
||||
1. CLI 参数
|
||||
2. 环境变量(`AETHER_PROXY_*`)
|
||||
3. 配置文件(`aether-proxy.toml`,或通过 `AETHER_PROXY_CONFIG` 指定路径)
|
||||
|
||||
### 参数一览
|
||||
|
||||
#### 基础配置
|
||||
|
||||
| 参数 | 环境变量 | 默认值 | 说明 |
|
||||
|------|----------|--------|------|
|
||||
| `--aether-url` | `AETHER_PROXY_AETHER_URL` | **必填** | Aether 服务器地址 |
|
||||
| `--management-token` | `AETHER_PROXY_MANAGEMENT_TOKEN` | **必填** | 管理员 Token(`ae_xxx` 格式) |
|
||||
| `--public-ip` | `AETHER_PROXY_PUBLIC_IP` | 自动检测 | 公网 IP |
|
||||
| `--node-name` | `AETHER_PROXY_NODE_NAME` | `proxy-01` | 节点名称标识 |
|
||||
| `--node-region` | `AETHER_PROXY_NODE_REGION` | 自动检测 | 地区标识 |
|
||||
| `--heartbeat-interval` | `AETHER_PROXY_HEARTBEAT_INTERVAL` | `30` | 心跳间隔(秒) |
|
||||
| `--allowed-ports` | `AETHER_PROXY_ALLOWED_PORTS` | `80,443,8080,8443` | 允许代理的目标端口 |
|
||||
|
||||
#### Tunnel 连接
|
||||
|
||||
| 参数 | 环境变量 | 默认值 | 说明 |
|
||||
|------|----------|--------|------|
|
||||
| `--tunnel-connections` | `AETHER_PROXY_TUNNEL_CONNECTIONS` | `3` | 到 Aether 的连接池大小 |
|
||||
| `--tunnel-max-streams` | `AETHER_PROXY_TUNNEL_MAX_STREAMS` | 自动(硬件估算) | 单连接最大并发 stream 数 |
|
||||
| `--tunnel-connect-timeout-secs` | `AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT_SECS` | `15` | TCP + TLS 握手超时(秒) |
|
||||
| `--tunnel-tcp-keepalive-secs` | `AETHER_PROXY_TUNNEL_TCP_KEEPALIVE_SECS` | `30` | TCP keepalive 初始延迟(秒) |
|
||||
| `--tunnel-tcp-nodelay` | `AETHER_PROXY_TUNNEL_TCP_NODELAY` | `true` | 禁用 Nagle 算法 |
|
||||
| `--tunnel-ping-interval-secs` | `AETHER_PROXY_TUNNEL_PING_INTERVAL_SECS` | `15` | WebSocket Ping 频率(秒) |
|
||||
| `--tunnel-stale-timeout-secs` | `AETHER_PROXY_TUNNEL_STALE_TIMEOUT_SECS` | `45` | 无数据断连阈值(秒) |
|
||||
| `--tunnel-reconnect-base-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS` | `500` | 指数退避基础延迟(毫秒) |
|
||||
| `--tunnel-reconnect-max-ms` | `AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS` | `30000` | 指数退避上限(毫秒) |
|
||||
|
||||
#### 上游 HTTP 请求
|
||||
|
||||
| 参数 | 环境变量 | 默认值 | 说明 |
|
||||
|------|----------|--------|------|
|
||||
| `--upstream-connect-timeout-secs` | `AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT_SECS` | `30` | 上游建连超时(秒) |
|
||||
| `--upstream-pool-max-idle-per-host` | `AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST` | `64` | 每 Host 最大空闲连接数 |
|
||||
| `--upstream-pool-idle-timeout-secs` | `AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT_SECS` | `300` | 连接池空闲超时(秒) |
|
||||
| `--upstream-tcp-keepalive-secs` | `AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE_SECS` | `60` | TCP keepalive(秒,0 关闭) |
|
||||
| `--upstream-tcp-nodelay` | `AETHER_PROXY_UPSTREAM_TCP_NODELAY` | `true` | 启用 TCP_NODELAY |
|
||||
|
||||
#### Aether API 客户端
|
||||
|
||||
| 参数 | 环境变量 | 默认值 | 说明 |
|
||||
|------|----------|--------|------|
|
||||
| `--aether-request-timeout-secs` | `AETHER_PROXY_AETHER_REQUEST_TIMEOUT_SECS` | `10` | 请求总超时(秒) |
|
||||
| `--aether-connect-timeout-secs` | `AETHER_PROXY_AETHER_CONNECT_TIMEOUT_SECS` | `10` | 建连超时(秒) |
|
||||
| `--aether-retry-max-attempts` | `AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS` | `3` | 最大重试次数 |
|
||||
|
||||
#### DNS 与安全
|
||||
|
||||
| 参数 | 环境变量 | 默认值 | 说明 |
|
||||
|------|----------|--------|------|
|
||||
| `--dns-cache-ttl-secs` | `AETHER_PROXY_DNS_CACHE_TTL_SECS` | `60` | DNS 缓存 TTL(秒) |
|
||||
| `--dns-cache-capacity` | `AETHER_PROXY_DNS_CACHE_CAPACITY` | `1024` | DNS 缓存容量(条目数) |
|
||||
|
||||
#### 日志
|
||||
|
||||
| 参数 | 环境变量 | 默认值 | 说明 |
|
||||
|------|----------|--------|------|
|
||||
| `--log-level` | `AETHER_PROXY_LOG_LEVEL` | `info` | 日志级别 |
|
||||
| `--log-json` | `AETHER_PROXY_LOG_JSON` | `false` | JSON 格式日志 |
|
||||
|
||||
### 多服务器配置
|
||||
|
||||
在 `aether-proxy.toml` 中使用 `[[servers]]` 配置多个 Aether 服务器:
|
||||
|
||||
```toml
|
||||
[[servers]]
|
||||
aether_url = "https://aether-1.example.com"
|
||||
management_token = "ae_xxx"
|
||||
node_name = "jp-proxy-01"
|
||||
|
||||
[[servers]]
|
||||
aether_url = "https://aether-2.example.com"
|
||||
management_token = "ae_yyy"
|
||||
node_name = "jp-proxy-02"
|
||||
```
|
||||
|
||||
## 发布新版本
|
||||
|
||||
推送 `proxy-v*` 格式的 tag,GitHub Actions 会自动:
|
||||
- 编译所有平台二进制并发布到 Releases
|
||||
- 构建 Docker 镜像并推送到 GHCR 和 Docker Hub
|
||||
- 更新 README 中的下载链接表格
|
||||
|
||||
```bash
|
||||
git tag proxy-v0.2.0
|
||||
git push origin proxy-v0.2.0
|
||||
```
|
||||
@@ -0,0 +1,14 @@
|
||||
services:
|
||||
aether-proxy:
|
||||
image: ghcr.io/fawney19/aether-proxy:latest
|
||||
container_name: aether-proxy
|
||||
restart: unless-stopped
|
||||
env_file:
|
||||
- .env
|
||||
environment:
|
||||
AETHER_PROXY_LOG_JSON: "true"
|
||||
logging:
|
||||
driver: json-file
|
||||
options:
|
||||
max-size: "50m"
|
||||
max-file: "3"
|
||||
@@ -0,0 +1,360 @@
|
||||
//! Application lifecycle: initialization, task orchestration, and shutdown.
|
||||
|
||||
use std::sync::atomic::AtomicU64;
|
||||
use std::sync::{Arc, RwLock};
|
||||
use std::time::Duration;
|
||||
|
||||
use arc_swap::ArcSwap;
|
||||
use tokio::signal;
|
||||
use tokio::sync::{watch, Mutex};
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use crate::config::{Config, ServerEntry};
|
||||
use crate::net;
|
||||
use crate::registration::client::AetherClient;
|
||||
use crate::runtime::{self, DynamicConfig};
|
||||
use crate::state::{AppState, ProxyMetrics, ServerContext};
|
||||
use crate::upstream_client;
|
||||
use crate::{hardware, target_filter, tunnel};
|
||||
|
||||
/// Run the full application lifecycle after config has been parsed.
|
||||
pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Result<()> {
|
||||
config.validate()?;
|
||||
init_tracing(&config);
|
||||
|
||||
info!(
|
||||
version = env!("CARGO_PKG_VERSION"),
|
||||
node_name = %config.node_name,
|
||||
server_count = servers.len(),
|
||||
"aether-proxy starting (tunnel mode)"
|
||||
);
|
||||
|
||||
// Resolve public IP (best-effort for region info)
|
||||
let public_ip = match &config.public_ip {
|
||||
Some(ip) => ip.clone(),
|
||||
None => net::detect_public_ip()
|
||||
.await
|
||||
.unwrap_or_else(|_| "0.0.0.0".to_string()),
|
||||
};
|
||||
|
||||
// Auto-detect region if not configured
|
||||
if config.node_region.is_none() {
|
||||
if let Some(region) = net::detect_region(&public_ip).await {
|
||||
config.node_region = Some(region);
|
||||
}
|
||||
}
|
||||
|
||||
// Collect hardware info (once at startup, sent during registration)
|
||||
let hw_info = hardware::collect();
|
||||
|
||||
// Auto-detect tunnel_max_streams from hardware if not explicitly set
|
||||
if config.tunnel_max_streams.is_none() {
|
||||
let auto = (hw_info.estimated_max_concurrency / 10).clamp(64, 1024) as u32;
|
||||
config.tunnel_max_streams = Some(auto);
|
||||
info!(
|
||||
tunnel_max_streams = auto,
|
||||
"auto-detected tunnel_max_streams from hardware"
|
||||
);
|
||||
}
|
||||
|
||||
info!(
|
||||
max_concurrency = hw_info.estimated_max_concurrency,
|
||||
"hardware info collected"
|
||||
);
|
||||
|
||||
let dns_cache = Arc::new(target_filter::DnsCache::new(
|
||||
Duration::from_secs(config.dns_cache_ttl_secs),
|
||||
config.dns_cache_capacity,
|
||||
));
|
||||
|
||||
// Build Hyper client for tunnel upstream requests (shared).
|
||||
// DNS still flows through validated addresses from DnsCache, while the
|
||||
// custom connector exposes per-request connect/TLS timing when available.
|
||||
let upstream_client = upstream_client::build_upstream_client(&config, Arc::clone(&dns_cache));
|
||||
|
||||
// Register with each Aether server and build per-server contexts.
|
||||
// Wrapped in Arc<Mutex> so retry_failed_registrations can append later.
|
||||
let server_contexts: Arc<Mutex<Vec<Arc<ServerContext>>>> = Arc::new(Mutex::new(Vec::new()));
|
||||
let mut failed_entries: Vec<(String, ServerEntry)> = Vec::new();
|
||||
for (i, entry) in servers.iter().enumerate() {
|
||||
let label = if servers.len() == 1 {
|
||||
"server".to_string()
|
||||
} else {
|
||||
format!("server-{}", i)
|
||||
};
|
||||
let node_name = entry
|
||||
.node_name
|
||||
.clone()
|
||||
.unwrap_or_else(|| config.node_name.clone());
|
||||
let client = Arc::new(AetherClient::new(
|
||||
&config,
|
||||
&entry.aether_url,
|
||||
&entry.management_token,
|
||||
));
|
||||
match client
|
||||
.register(&config, &node_name, &public_ip, Some(&hw_info))
|
||||
.await
|
||||
{
|
||||
Ok(node_id) => {
|
||||
info!(server = %label, node_id = %node_id, url = %entry.aether_url, node_name = %node_name, "registered");
|
||||
// Initialize dynamic config with per-server node_name (not global),
|
||||
// so that the heartbeat and reconnect use the correct name.
|
||||
let mut dynamic = DynamicConfig::from_config(&config);
|
||||
dynamic.node_name = node_name.clone();
|
||||
server_contexts.lock().await.push(Arc::new(ServerContext {
|
||||
server_label: label,
|
||||
aether_url: entry.aether_url.clone(),
|
||||
management_token: entry.management_token.clone(),
|
||||
node_name,
|
||||
node_id: Arc::new(RwLock::new(node_id)),
|
||||
aether_client: client,
|
||||
dynamic: Arc::new(ArcSwap::from_pointee(dynamic)),
|
||||
active_connections: Arc::new(AtomicU64::new(0)),
|
||||
metrics: Arc::new(ProxyMetrics::new()),
|
||||
}));
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(
|
||||
server = %label,
|
||||
url = %entry.aether_url,
|
||||
error = %e,
|
||||
"registration failed, will retry in background"
|
||||
);
|
||||
failed_entries.push((label, entry.clone()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
{
|
||||
let ctx_count = server_contexts.lock().await.len();
|
||||
if ctx_count == 0 && failed_entries.is_empty() {
|
||||
anyhow::bail!("no servers configured");
|
||||
}
|
||||
if ctx_count == 0 {
|
||||
anyhow::bail!(
|
||||
"no servers registered successfully (all {} failed)",
|
||||
failed_entries.len()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Build shared application state
|
||||
let tunnel_tls_config = Arc::new(crate::tunnel::client::build_tls_config());
|
||||
let state = Arc::new(AppState {
|
||||
config: Arc::new(config),
|
||||
dns_cache,
|
||||
upstream_client,
|
||||
tunnel_tls_config,
|
||||
});
|
||||
|
||||
// Shutdown signal channel
|
||||
let (shutdown_tx, shutdown_rx) = watch::channel(false);
|
||||
|
||||
info!(
|
||||
active_servers = server_contexts.lock().await.len(),
|
||||
"running in tunnel mode"
|
||||
);
|
||||
|
||||
// Spawn tunnel connections per server (pool_size connections each)
|
||||
let pool_size = state.config.tunnel_connections.max(1) as usize;
|
||||
let mut tunnel_handles = Vec::new();
|
||||
for server in server_contexts.lock().await.iter() {
|
||||
for conn_idx in 0..pool_size {
|
||||
let s = Arc::clone(&state);
|
||||
let srv = Arc::clone(server);
|
||||
let rx = shutdown_rx.clone();
|
||||
tunnel_handles.push(tokio::spawn(async move {
|
||||
tunnel::run(&s, &srv, conn_idx, rx).await;
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
// Spawn background retry for failed server registrations
|
||||
if !failed_entries.is_empty() {
|
||||
let retry_state = Arc::clone(&state);
|
||||
let retry_contexts = Arc::clone(&server_contexts);
|
||||
let retry_public_ip = public_ip.clone();
|
||||
let retry_hw_info = hw_info.clone();
|
||||
let retry_shutdown = shutdown_rx.clone();
|
||||
let retry_pool_size = pool_size;
|
||||
tokio::spawn(async move {
|
||||
retry_failed_registrations(
|
||||
retry_state,
|
||||
retry_contexts,
|
||||
failed_entries,
|
||||
retry_public_ip,
|
||||
retry_hw_info,
|
||||
retry_pool_size,
|
||||
retry_shutdown,
|
||||
)
|
||||
.await;
|
||||
});
|
||||
}
|
||||
|
||||
// Wait for shutdown signal
|
||||
wait_for_shutdown().await;
|
||||
info!("shutdown signal received, cleaning up...");
|
||||
let _ = shutdown_tx.send(true);
|
||||
|
||||
// Graceful unregister from all servers (including retry-registered ones)
|
||||
for server in server_contexts.lock().await.iter() {
|
||||
let node_id = server.node_id.read().unwrap().clone();
|
||||
if let Err(e) = server.aether_client.unregister(&node_id).await {
|
||||
error!(
|
||||
server = %server.server_label,
|
||||
error = %e,
|
||||
"unregister failed during shutdown"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Wait for all tunnel tasks
|
||||
for h in tunnel_handles {
|
||||
let _ = h.await;
|
||||
}
|
||||
|
||||
info!("aether-proxy stopped");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Retry interval for failed server registrations (5 minutes).
|
||||
const REGISTRATION_RETRY_INTERVAL: Duration = Duration::from_secs(300);
|
||||
/// Max registration retry attempts before giving up.
|
||||
const REGISTRATION_RETRY_MAX: u32 = 12;
|
||||
|
||||
/// Background task that retries registration for servers that failed at startup.
|
||||
async fn retry_failed_registrations(
|
||||
state: Arc<AppState>,
|
||||
server_contexts: Arc<Mutex<Vec<Arc<ServerContext>>>>,
|
||||
failed: Vec<(String, ServerEntry)>,
|
||||
public_ip: String,
|
||||
hw_info: crate::hardware::HardwareInfo,
|
||||
pool_size: usize,
|
||||
mut shutdown: watch::Receiver<bool>,
|
||||
) {
|
||||
for (label, entry) in &failed {
|
||||
let node_name = entry
|
||||
.node_name
|
||||
.clone()
|
||||
.unwrap_or_else(|| state.config.node_name.clone());
|
||||
let client = Arc::new(AetherClient::new(
|
||||
&state.config,
|
||||
&entry.aether_url,
|
||||
&entry.management_token,
|
||||
));
|
||||
|
||||
let mut attempt = 0u32;
|
||||
let node_id = loop {
|
||||
attempt += 1;
|
||||
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(REGISTRATION_RETRY_INTERVAL) => {}
|
||||
_ = shutdown.changed() => {
|
||||
info!(server = %label, "shutdown during registration retry");
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
match client
|
||||
.register(&state.config, &node_name, &public_ip, Some(&hw_info))
|
||||
.await
|
||||
{
|
||||
Ok(id) => {
|
||||
info!(server = %label, node_id = %id, attempt, "registration retry succeeded");
|
||||
break id;
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(
|
||||
server = %label,
|
||||
attempt,
|
||||
max = REGISTRATION_RETRY_MAX,
|
||||
error = %e,
|
||||
"registration retry failed"
|
||||
);
|
||||
if attempt >= REGISTRATION_RETRY_MAX {
|
||||
error!(server = %label, "giving up registration after {} attempts", attempt);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// Build server context and spawn tunnels
|
||||
let mut dynamic = DynamicConfig::from_config(&state.config);
|
||||
dynamic.node_name = node_name.clone();
|
||||
let server = Arc::new(ServerContext {
|
||||
server_label: label.clone(),
|
||||
aether_url: entry.aether_url.clone(),
|
||||
management_token: entry.management_token.clone(),
|
||||
node_name,
|
||||
node_id: Arc::new(RwLock::new(node_id)),
|
||||
aether_client: client,
|
||||
dynamic: Arc::new(ArcSwap::from_pointee(dynamic)),
|
||||
active_connections: Arc::new(AtomicU64::new(0)),
|
||||
metrics: Arc::new(ProxyMetrics::new()),
|
||||
});
|
||||
|
||||
// Add to shared list so shutdown can unregister this server
|
||||
server_contexts.lock().await.push(Arc::clone(&server));
|
||||
|
||||
for conn_idx in 0..pool_size {
|
||||
let s = Arc::clone(&state);
|
||||
let srv = Arc::clone(&server);
|
||||
let rx = shutdown.clone();
|
||||
tokio::spawn(async move {
|
||||
tunnel::run(&s, &srv, conn_idx, rx).await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn init_tracing(config: &Config) {
|
||||
use tracing_subscriber::prelude::*;
|
||||
use tracing_subscriber::{reload, EnvFilter};
|
||||
|
||||
let filter = EnvFilter::try_new(&config.log_level).unwrap_or_else(|_| EnvFilter::new("info"));
|
||||
|
||||
let (filter_layer, reload_handle) = reload::Layer::new(filter);
|
||||
|
||||
runtime::set_log_reloader(Box::new(move |level: &str| {
|
||||
if let Ok(new_filter) = EnvFilter::try_new(level) {
|
||||
let _ = reload_handle.modify(|f| *f = new_filter);
|
||||
}
|
||||
}));
|
||||
|
||||
if config.log_json {
|
||||
tracing_subscriber::registry()
|
||||
.with(filter_layer)
|
||||
.with(tracing_subscriber::fmt::layer().json())
|
||||
.init();
|
||||
} else {
|
||||
tracing_subscriber::registry()
|
||||
.with(filter_layer)
|
||||
.with(tracing_subscriber::fmt::layer())
|
||||
.init();
|
||||
}
|
||||
}
|
||||
|
||||
async fn wait_for_shutdown() {
|
||||
let ctrl_c = async {
|
||||
signal::ctrl_c()
|
||||
.await
|
||||
.expect("failed to install Ctrl+C handler");
|
||||
};
|
||||
|
||||
#[cfg(unix)]
|
||||
let terminate = async {
|
||||
signal::unix::signal(signal::unix::SignalKind::terminate())
|
||||
.expect("failed to install SIGTERM handler")
|
||||
.recv()
|
||||
.await;
|
||||
};
|
||||
|
||||
#[cfg(not(unix))]
|
||||
let terminate = std::future::pending::<()>();
|
||||
|
||||
tokio::select! {
|
||||
_ = ctrl_c => {},
|
||||
_ = terminate => {},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,656 @@
|
||||
use std::path::Path;
|
||||
|
||||
use clap::Parser;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// Fields that existed in 0.1.x but were removed in 0.2.0.
|
||||
const LEGACY_ONLY_KEYS: &[&str] = &[
|
||||
"hmac_key",
|
||||
"listen_port",
|
||||
"timestamp_tolerance",
|
||||
"connect_timeout_secs",
|
||||
"tls_handshake_timeout_secs",
|
||||
"enable_tls",
|
||||
"tls_cert",
|
||||
"tls_key",
|
||||
];
|
||||
|
||||
/// Fields renamed from 0.1.x `delegate_*` to 0.2.0 `upstream_*`.
|
||||
const DELEGATE_TO_UPSTREAM: &[(&str, &str)] = &[
|
||||
(
|
||||
"delegate_connect_timeout_secs",
|
||||
"upstream_connect_timeout_secs",
|
||||
),
|
||||
(
|
||||
"delegate_pool_max_idle_per_host",
|
||||
"upstream_pool_max_idle_per_host",
|
||||
),
|
||||
(
|
||||
"delegate_pool_idle_timeout_secs",
|
||||
"upstream_pool_idle_timeout_secs",
|
||||
),
|
||||
("delegate_tcp_keepalive_secs", "upstream_tcp_keepalive_secs"),
|
||||
("delegate_tcp_nodelay", "upstream_tcp_nodelay"),
|
||||
];
|
||||
|
||||
/// Aether tunnel proxy.
|
||||
///
|
||||
/// Deployed on overseas VPS to relay API traffic for Aether instances
|
||||
/// behind the GFW. Connects to Aether via WebSocket tunnel, registers
|
||||
/// with Aether, and relays upstream requests.
|
||||
#[derive(Parser, Debug, Clone)]
|
||||
#[command(version, about)]
|
||||
pub struct Config {
|
||||
/// Aether server URL (e.g. https://aether.example.com)
|
||||
#[arg(long, env = "AETHER_PROXY_AETHER_URL")]
|
||||
pub aether_url: String,
|
||||
|
||||
/// Management Token for Aether admin API (ae_xxx)
|
||||
#[arg(long, env = "AETHER_PROXY_MANAGEMENT_TOKEN")]
|
||||
pub management_token: String,
|
||||
|
||||
/// Public IP address of this node (auto-detected if omitted)
|
||||
#[arg(long, env = "AETHER_PROXY_PUBLIC_IP")]
|
||||
pub public_ip: Option<String>,
|
||||
|
||||
/// Human-readable node name
|
||||
#[arg(long, env = "AETHER_PROXY_NODE_NAME", default_value = "proxy-01")]
|
||||
pub node_name: String,
|
||||
|
||||
/// Region label (e.g. ap-northeast-1)
|
||||
#[arg(long, env = "AETHER_PROXY_NODE_REGION")]
|
||||
pub node_region: Option<String>,
|
||||
|
||||
/// Heartbeat interval in seconds
|
||||
#[arg(long, env = "AETHER_PROXY_HEARTBEAT_INTERVAL", default_value_t = 30)]
|
||||
pub heartbeat_interval: u64,
|
||||
|
||||
/// Allowed destination ports (default: 80,443,8080,8443)
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_ALLOWED_PORTS",
|
||||
value_delimiter = ',',
|
||||
default_values_t = vec![80, 443, 8080, 8443]
|
||||
)]
|
||||
pub allowed_ports: Vec<u16>,
|
||||
|
||||
/// Aether API request timeout in seconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_REQUEST_TIMEOUT",
|
||||
default_value_t = 10
|
||||
)]
|
||||
pub aether_request_timeout_secs: u64,
|
||||
|
||||
/// Aether API connect timeout in seconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_CONNECT_TIMEOUT",
|
||||
default_value_t = 10
|
||||
)]
|
||||
pub aether_connect_timeout_secs: u64,
|
||||
|
||||
/// Aether API max idle connections per host
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_POOL_MAX_IDLE_PER_HOST",
|
||||
default_value_t = 8
|
||||
)]
|
||||
pub aether_pool_max_idle_per_host: usize,
|
||||
|
||||
/// Aether API idle timeout in seconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_POOL_IDLE_TIMEOUT",
|
||||
default_value_t = 90
|
||||
)]
|
||||
pub aether_pool_idle_timeout_secs: u64,
|
||||
|
||||
/// Aether API TCP keepalive in seconds (0 disables)
|
||||
#[arg(long, env = "AETHER_PROXY_AETHER_TCP_KEEPALIVE", default_value_t = 60)]
|
||||
pub aether_tcp_keepalive_secs: u64,
|
||||
|
||||
/// Aether API TCP_NODELAY
|
||||
#[arg(long, env = "AETHER_PROXY_AETHER_TCP_NODELAY", default_value_t = true)]
|
||||
pub aether_tcp_nodelay: bool,
|
||||
|
||||
/// Enable HTTP/2 when talking to Aether API
|
||||
#[arg(long, env = "AETHER_PROXY_AETHER_HTTP2", default_value_t = true)]
|
||||
pub aether_http2: bool,
|
||||
|
||||
/// Aether API retry attempts (including initial)
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS",
|
||||
default_value_t = 3
|
||||
)]
|
||||
pub aether_retry_max_attempts: u32,
|
||||
|
||||
/// Aether API retry base delay in milliseconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_RETRY_BASE_DELAY_MS",
|
||||
default_value_t = 200
|
||||
)]
|
||||
pub aether_retry_base_delay_ms: u64,
|
||||
|
||||
/// Aether API retry max delay in milliseconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_AETHER_RETRY_MAX_DELAY_MS",
|
||||
default_value_t = 2000
|
||||
)]
|
||||
pub aether_retry_max_delay_ms: u64,
|
||||
|
||||
/// Maximum concurrent TCP connections (defaults to hardware estimate)
|
||||
#[arg(long, env = "AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS")]
|
||||
pub max_concurrent_connections: Option<u64>,
|
||||
|
||||
/// DNS cache TTL in seconds
|
||||
#[arg(long, env = "AETHER_PROXY_DNS_CACHE_TTL", default_value_t = 60)]
|
||||
pub dns_cache_ttl_secs: u64,
|
||||
|
||||
/// DNS cache capacity (entries)
|
||||
#[arg(long, env = "AETHER_PROXY_DNS_CACHE_CAPACITY", default_value_t = 1024)]
|
||||
pub dns_cache_capacity: usize,
|
||||
|
||||
/// Upstream HTTP client connect timeout in seconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT",
|
||||
default_value_t = 30
|
||||
)]
|
||||
pub upstream_connect_timeout_secs: u64,
|
||||
|
||||
/// Upstream HTTP client max idle connections per host
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST",
|
||||
default_value_t = 64
|
||||
)]
|
||||
pub upstream_pool_max_idle_per_host: usize,
|
||||
|
||||
/// Upstream HTTP client idle timeout in seconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT",
|
||||
default_value_t = 300
|
||||
)]
|
||||
pub upstream_pool_idle_timeout_secs: u64,
|
||||
|
||||
/// Upstream TCP keepalive in seconds (0 disables)
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE",
|
||||
default_value_t = 60
|
||||
)]
|
||||
pub upstream_tcp_keepalive_secs: u64,
|
||||
|
||||
/// Upstream TCP_NODELAY
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_UPSTREAM_TCP_NODELAY",
|
||||
default_value_t = true
|
||||
)]
|
||||
pub upstream_tcp_nodelay: bool,
|
||||
|
||||
/// Log level (trace, debug, info, warn, error)
|
||||
#[arg(long, env = "AETHER_PROXY_LOG_LEVEL", default_value = "info")]
|
||||
pub log_level: String,
|
||||
|
||||
/// Output logs as JSON
|
||||
#[arg(long, env = "AETHER_PROXY_LOG_JSON", default_value_t = false)]
|
||||
pub log_json: bool,
|
||||
|
||||
/// Tunnel reconnect base delay in milliseconds (used by exponential backoff)
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS",
|
||||
default_value_t = 500
|
||||
)]
|
||||
pub tunnel_reconnect_base_ms: u64,
|
||||
|
||||
/// Tunnel reconnect max delay in milliseconds (cap for exponential backoff)
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS",
|
||||
default_value_t = 30000
|
||||
)]
|
||||
pub tunnel_reconnect_max_ms: u64,
|
||||
|
||||
/// WebSocket tunnel ping interval in seconds
|
||||
#[arg(long, env = "AETHER_PROXY_TUNNEL_PING_INTERVAL", default_value_t = 15)]
|
||||
pub tunnel_ping_interval_secs: u64,
|
||||
|
||||
/// Maximum concurrent streams over tunnel (auto-detected from hardware if omitted)
|
||||
#[arg(long, env = "AETHER_PROXY_TUNNEL_MAX_STREAMS")]
|
||||
pub tunnel_max_streams: Option<u32>,
|
||||
|
||||
/// WebSocket tunnel TCP connect timeout in seconds
|
||||
#[arg(
|
||||
long,
|
||||
env = "AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT",
|
||||
default_value_t = 15
|
||||
)]
|
||||
pub tunnel_connect_timeout_secs: u64,
|
||||
|
||||
/// WebSocket tunnel TCP keepalive in seconds (0 disables)
|
||||
#[arg(long, env = "AETHER_PROXY_TUNNEL_TCP_KEEPALIVE", default_value_t = 30)]
|
||||
pub tunnel_tcp_keepalive_secs: u64,
|
||||
|
||||
/// WebSocket tunnel TCP_NODELAY
|
||||
#[arg(long, env = "AETHER_PROXY_TUNNEL_TCP_NODELAY", default_value_t = true)]
|
||||
pub tunnel_tcp_nodelay: bool,
|
||||
|
||||
/// Tunnel connection staleness timeout in seconds (triggers reconnect if no data received)
|
||||
#[arg(long, env = "AETHER_PROXY_TUNNEL_STALE_TIMEOUT", default_value_t = 45)]
|
||||
pub tunnel_stale_timeout_secs: u64,
|
||||
|
||||
/// Number of parallel WebSocket tunnel connections per server (connection pool)
|
||||
#[arg(long, env = "AETHER_PROXY_TUNNEL_CONNECTIONS", default_value_t = 3)]
|
||||
pub tunnel_connections: u32,
|
||||
}
|
||||
|
||||
impl Config {
|
||||
/// Validate configuration values are within sane ranges.
|
||||
/// Called after parsing to catch misconfigurations early.
|
||||
pub fn validate(&self) -> anyhow::Result<()> {
|
||||
if self.heartbeat_interval == 0 {
|
||||
anyhow::bail!("heartbeat_interval must be > 0");
|
||||
}
|
||||
if self.heartbeat_interval > 3600 {
|
||||
anyhow::bail!("heartbeat_interval must be <= 3600");
|
||||
}
|
||||
if self.allowed_ports.is_empty() {
|
||||
anyhow::bail!("allowed_ports must not be empty");
|
||||
}
|
||||
for &port in &self.allowed_ports {
|
||||
if port == 0 {
|
||||
anyhow::bail!("allowed_ports: port 0 is not valid");
|
||||
}
|
||||
}
|
||||
if self.tunnel_connect_timeout_secs == 0 {
|
||||
anyhow::bail!("tunnel_connect_timeout_secs must be > 0");
|
||||
}
|
||||
if self.tunnel_ping_interval_secs == 0 {
|
||||
anyhow::bail!("tunnel_ping_interval_secs must be > 0");
|
||||
}
|
||||
if self.tunnel_stale_timeout_secs <= self.tunnel_ping_interval_secs {
|
||||
anyhow::bail!(
|
||||
"tunnel_stale_timeout_secs ({}) must be > tunnel_ping_interval_secs ({})",
|
||||
self.tunnel_stale_timeout_secs,
|
||||
self.tunnel_ping_interval_secs
|
||||
);
|
||||
}
|
||||
if self.tunnel_connections == 0 {
|
||||
anyhow::bail!("tunnel_connections must be > 0");
|
||||
}
|
||||
if self.aether_retry_max_attempts == 0 {
|
||||
anyhow::bail!("aether_retry_max_attempts must be >= 1");
|
||||
}
|
||||
if self.upstream_connect_timeout_secs == 0 {
|
||||
anyhow::bail!("upstream_connect_timeout_secs must be > 0");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Per-server connection config (used in multi-server TOML `[[servers]]`).
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct ServerEntry {
|
||||
pub aether_url: String,
|
||||
pub management_token: String,
|
||||
/// Per-server node name override. Falls back to the global `node_name`.
|
||||
pub node_name: Option<String>,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// TOML config file support
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Serializable config for TOML file persistence.
|
||||
/// All fields are optional -- only populated values are written.
|
||||
#[derive(Debug, Default, Serialize, Deserialize)]
|
||||
pub struct ConfigFile {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_url: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub management_token: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub public_ip: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub node_name: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub node_region: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub heartbeat_interval: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub allowed_ports: Option<Vec<u16>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_request_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_connect_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_pool_max_idle_per_host: Option<usize>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_pool_idle_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_tcp_keepalive_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_tcp_nodelay: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_http2: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_retry_max_attempts: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_retry_base_delay_ms: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub aether_retry_max_delay_ms: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_concurrent_connections: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub dns_cache_ttl_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub dns_cache_capacity: Option<usize>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub upstream_connect_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub upstream_pool_max_idle_per_host: Option<usize>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub upstream_pool_idle_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub upstream_tcp_keepalive_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub upstream_tcp_nodelay: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub log_level: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub log_json: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_reconnect_base_ms: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_reconnect_max_ms: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_ping_interval_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_max_streams: Option<u32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_connect_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_tcp_keepalive_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_tcp_nodelay: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_stale_timeout_secs: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub tunnel_connections: Option<u32>,
|
||||
|
||||
/// Multi-server config: each entry connects to a separate Aether instance.
|
||||
/// When present, top-level aether_url/management_token are ignored for
|
||||
/// tunnel connections (but still injected as env for clap compatibility).
|
||||
#[serde(default, skip_serializing_if = "Vec::is_empty")]
|
||||
pub servers: Vec<ServerEntry>,
|
||||
}
|
||||
|
||||
impl ConfigFile {
|
||||
/// Load from a TOML file.
|
||||
pub fn load(path: &Path) -> anyhow::Result<Self> {
|
||||
let content = std::fs::read_to_string(path)?;
|
||||
Ok(toml::from_str(&content)?)
|
||||
}
|
||||
|
||||
/// Save to a TOML file.
|
||||
pub fn save(&self, path: &Path) -> anyhow::Result<()> {
|
||||
let content = toml::to_string_pretty(self)?;
|
||||
std::fs::write(path, content)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Detect and migrate a 0.1.x config file to 0.2.0 format in-place.
|
||||
///
|
||||
/// Returns `true` if migration was performed, `false` if already current.
|
||||
/// The original file is backed up as `<name>.v1.bak` before rewriting.
|
||||
pub fn migrate_legacy(path: &Path) -> anyhow::Result<bool> {
|
||||
let content = match std::fs::read_to_string(path) {
|
||||
Ok(c) => c,
|
||||
Err(_) => return Ok(false),
|
||||
};
|
||||
let mut table: toml::map::Map<String, toml::Value> = toml::from_str(&content)?;
|
||||
|
||||
// Detect legacy format: presence of any 0.1.x-only key.
|
||||
let is_legacy = LEGACY_ONLY_KEYS.iter().any(|k| table.contains_key(*k))
|
||||
|| DELEGATE_TO_UPSTREAM
|
||||
.iter()
|
||||
.any(|(old, _)| table.contains_key(*old));
|
||||
|
||||
if !is_legacy {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
// 1. Rename delegate_* -> upstream_* (carry over user-customized values)
|
||||
for &(old, new) in DELEGATE_TO_UPSTREAM {
|
||||
if let Some(val) = table.remove(old) {
|
||||
table.entry(new.to_string()).or_insert(val);
|
||||
}
|
||||
}
|
||||
|
||||
// 2. Build [[servers]] from top-level aether_url + management_token + node_name
|
||||
if !table.contains_key("servers") {
|
||||
let aether_url = table.get("aether_url").and_then(|v| v.as_str());
|
||||
let management_token = table.get("management_token").and_then(|v| v.as_str());
|
||||
if let (Some(url), Some(token)) = (aether_url, management_token) {
|
||||
let mut entry = toml::map::Map::new();
|
||||
entry.insert("aether_url".into(), toml::Value::String(url.to_string()));
|
||||
entry.insert(
|
||||
"management_token".into(),
|
||||
toml::Value::String(token.to_string()),
|
||||
);
|
||||
if let Some(name) = table.get("node_name").and_then(|v| v.as_str()) {
|
||||
entry.insert("node_name".into(), toml::Value::String(name.to_string()));
|
||||
}
|
||||
table.insert(
|
||||
"servers".into(),
|
||||
toml::Value::Array(vec![toml::Value::Table(entry)]),
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Remove top-level fields that are now in [[servers]] or obsolete
|
||||
table.remove("aether_url");
|
||||
table.remove("management_token");
|
||||
table.remove("node_name");
|
||||
for &key in LEGACY_ONLY_KEYS {
|
||||
table.remove(key);
|
||||
}
|
||||
|
||||
// 4. Backup original file (abort migration if backup fails)
|
||||
let backup_path = path.with_extension("v1.bak");
|
||||
std::fs::copy(path, &backup_path).map_err(|e| {
|
||||
anyhow::anyhow!(
|
||||
"failed to backup config before migration: {} -> {}: {}",
|
||||
path.display(),
|
||||
backup_path.display(),
|
||||
e
|
||||
)
|
||||
})?;
|
||||
|
||||
// 5. Write migrated config
|
||||
let new_content = toml::to_string_pretty(&table)?;
|
||||
std::fs::write(path, &new_content)?;
|
||||
|
||||
eprintln!(" Config migrated from 0.1.x to 0.2.0 format.");
|
||||
eprintln!(" Backup saved: {}", backup_path.display());
|
||||
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
/// Resolve the effective server list.
|
||||
///
|
||||
/// If `[[servers]]` is present, use it. Otherwise fall back to the
|
||||
/// top-level `aether_url` + `management_token` as a single server.
|
||||
pub fn effective_servers(&self) -> Vec<ServerEntry> {
|
||||
if !self.servers.is_empty() {
|
||||
return self.servers.clone();
|
||||
}
|
||||
match (&self.aether_url, &self.management_token) {
|
||||
(Some(url), Some(token)) => vec![ServerEntry {
|
||||
aether_url: url.clone(),
|
||||
management_token: token.clone(),
|
||||
node_name: None,
|
||||
}],
|
||||
_ => vec![],
|
||||
}
|
||||
}
|
||||
|
||||
/// Inject values as environment variables so clap picks them up.
|
||||
///
|
||||
/// Only sets variables that are **not** already present in the
|
||||
/// environment, preserving the precedence: CLI > env > config file.
|
||||
pub fn inject_env(&self) {
|
||||
self.inject_env_inner(false);
|
||||
}
|
||||
|
||||
/// Inject values as environment variables, **overriding** any existing
|
||||
/// values. Used after setup to ensure the freshly-saved config takes
|
||||
/// effect before re-parsing.
|
||||
pub fn inject_env_override(&self) {
|
||||
self.inject_env_inner(true);
|
||||
}
|
||||
|
||||
fn inject_env_inner(&self, force: bool) {
|
||||
macro_rules! set {
|
||||
($env:expr, $val:expr) => {
|
||||
if let Some(ref v) = $val {
|
||||
if force || std::env::var($env).is_err() {
|
||||
std::env::set_var($env, v.to_string());
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
// When top-level fields are absent, fall back to the first [[servers]]
|
||||
// entry so that clap's required `aether_url` / `management_token` are
|
||||
// satisfied even with the new config format.
|
||||
let first_server = self.servers.first();
|
||||
let aether_url = self
|
||||
.aether_url
|
||||
.as_deref()
|
||||
.or(first_server.map(|s| s.aether_url.as_str()));
|
||||
let management_token = self
|
||||
.management_token
|
||||
.as_deref()
|
||||
.or(first_server.map(|s| s.management_token.as_str()));
|
||||
let node_name = self
|
||||
.node_name
|
||||
.as_deref()
|
||||
.or(first_server.and_then(|s| s.node_name.as_deref()));
|
||||
|
||||
set!("AETHER_PROXY_AETHER_URL", aether_url);
|
||||
set!("AETHER_PROXY_MANAGEMENT_TOKEN", management_token);
|
||||
set!("AETHER_PROXY_PUBLIC_IP", self.public_ip);
|
||||
set!("AETHER_PROXY_NODE_NAME", node_name);
|
||||
set!("AETHER_PROXY_NODE_REGION", self.node_region);
|
||||
set!("AETHER_PROXY_HEARTBEAT_INTERVAL", self.heartbeat_interval);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_REQUEST_TIMEOUT",
|
||||
self.aether_request_timeout_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_CONNECT_TIMEOUT",
|
||||
self.aether_connect_timeout_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_POOL_MAX_IDLE_PER_HOST",
|
||||
self.aether_pool_max_idle_per_host
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_POOL_IDLE_TIMEOUT",
|
||||
self.aether_pool_idle_timeout_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_TCP_KEEPALIVE",
|
||||
self.aether_tcp_keepalive_secs
|
||||
);
|
||||
set!("AETHER_PROXY_AETHER_TCP_NODELAY", self.aether_tcp_nodelay);
|
||||
set!("AETHER_PROXY_AETHER_HTTP2", self.aether_http2);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_RETRY_MAX_ATTEMPTS",
|
||||
self.aether_retry_max_attempts
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_RETRY_BASE_DELAY_MS",
|
||||
self.aether_retry_base_delay_ms
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_AETHER_RETRY_MAX_DELAY_MS",
|
||||
self.aether_retry_max_delay_ms
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_MAX_CONCURRENT_CONNECTIONS",
|
||||
self.max_concurrent_connections
|
||||
);
|
||||
set!("AETHER_PROXY_DNS_CACHE_TTL", self.dns_cache_ttl_secs);
|
||||
set!("AETHER_PROXY_DNS_CACHE_CAPACITY", self.dns_cache_capacity);
|
||||
set!(
|
||||
"AETHER_PROXY_UPSTREAM_CONNECT_TIMEOUT",
|
||||
self.upstream_connect_timeout_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_UPSTREAM_POOL_MAX_IDLE_PER_HOST",
|
||||
self.upstream_pool_max_idle_per_host
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_UPSTREAM_POOL_IDLE_TIMEOUT",
|
||||
self.upstream_pool_idle_timeout_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_UPSTREAM_TCP_KEEPALIVE",
|
||||
self.upstream_tcp_keepalive_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_UPSTREAM_TCP_NODELAY",
|
||||
self.upstream_tcp_nodelay
|
||||
);
|
||||
set!("AETHER_PROXY_LOG_LEVEL", self.log_level);
|
||||
set!("AETHER_PROXY_LOG_JSON", self.log_json);
|
||||
set!(
|
||||
"AETHER_PROXY_TUNNEL_RECONNECT_BASE_MS",
|
||||
self.tunnel_reconnect_base_ms
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_TUNNEL_RECONNECT_MAX_MS",
|
||||
self.tunnel_reconnect_max_ms
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_TUNNEL_PING_INTERVAL",
|
||||
self.tunnel_ping_interval_secs
|
||||
);
|
||||
set!("AETHER_PROXY_TUNNEL_MAX_STREAMS", self.tunnel_max_streams);
|
||||
set!(
|
||||
"AETHER_PROXY_TUNNEL_CONNECT_TIMEOUT",
|
||||
self.tunnel_connect_timeout_secs
|
||||
);
|
||||
set!(
|
||||
"AETHER_PROXY_TUNNEL_TCP_KEEPALIVE",
|
||||
self.tunnel_tcp_keepalive_secs
|
||||
);
|
||||
set!("AETHER_PROXY_TUNNEL_TCP_NODELAY", self.tunnel_tcp_nodelay);
|
||||
set!(
|
||||
"AETHER_PROXY_TUNNEL_STALE_TIMEOUT",
|
||||
self.tunnel_stale_timeout_secs
|
||||
);
|
||||
set!("AETHER_PROXY_TUNNEL_CONNECTIONS", self.tunnel_connections);
|
||||
|
||||
// allowed_ports needs special handling (comma-separated)
|
||||
if let Some(ref ports) = self.allowed_ports {
|
||||
if force || std::env::var("AETHER_PROXY_ALLOWED_PORTS").is_err() {
|
||||
let s: String = ports
|
||||
.iter()
|
||||
.map(|p| p.to_string())
|
||||
.collect::<Vec<_>>()
|
||||
.join(",");
|
||||
std::env::set_var("AETHER_PROXY_ALLOWED_PORTS", s);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
use serde::Serialize;
|
||||
use sysinfo::System;
|
||||
use tracing::info;
|
||||
|
||||
/// Hardware information collected at startup.
|
||||
///
|
||||
/// The struct is `Serialize`-able so it can be sent directly as the
|
||||
/// `hardware_info` JSON bag in the registration request. New fields
|
||||
/// can be added without database schema migrations.
|
||||
#[derive(Debug, Clone, Serialize)]
|
||||
pub struct HardwareInfo {
|
||||
pub cpu_cores: u32,
|
||||
pub total_memory_mb: u64,
|
||||
pub os_info: String,
|
||||
pub fd_limit: u64,
|
||||
#[serde(skip)]
|
||||
pub estimated_max_concurrency: u64,
|
||||
}
|
||||
|
||||
/// Collect hardware information and estimate max concurrency.
|
||||
///
|
||||
/// Should be called once at startup -- hardware does not change at runtime.
|
||||
pub fn collect() -> HardwareInfo {
|
||||
let sys = System::new_all();
|
||||
|
||||
let cpu_cores = sys.cpus().len() as u32;
|
||||
let total_memory_mb = sys.total_memory() / (1024 * 1024);
|
||||
let os_info = format!(
|
||||
"{} {}",
|
||||
System::name().unwrap_or_else(|| "Unknown".into()),
|
||||
System::os_version().unwrap_or_default(),
|
||||
)
|
||||
.trim()
|
||||
.to_string();
|
||||
|
||||
// Estimate max concurrent connections:
|
||||
// - Each tokio async task uses ~8-16 KB stack + heap buffers
|
||||
// - OS file descriptor limit is often the real bottleneck
|
||||
// - Conservative formula: min(fd_limit - 100, ram_mb * 40, cpu_cores * 2000)
|
||||
let fd_limit = get_fd_limit();
|
||||
let by_fd = fd_limit.saturating_sub(100);
|
||||
let by_ram = total_memory_mb.saturating_mul(40);
|
||||
let by_cpu = (cpu_cores as u64).saturating_mul(2000);
|
||||
let estimated_max_concurrency = by_fd.min(by_ram).min(by_cpu);
|
||||
|
||||
info!(
|
||||
cpu_cores,
|
||||
total_memory_mb,
|
||||
os_info = %os_info,
|
||||
fd_limit,
|
||||
estimated_max_concurrency,
|
||||
"hardware info collected"
|
||||
);
|
||||
|
||||
HardwareInfo {
|
||||
cpu_cores,
|
||||
total_memory_mb,
|
||||
os_info,
|
||||
fd_limit,
|
||||
estimated_max_concurrency,
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the soft file-descriptor limit (RLIMIT_NOFILE).
|
||||
fn get_fd_limit() -> u64 {
|
||||
#[cfg(unix)]
|
||||
{
|
||||
let mut rlim = libc::rlimit {
|
||||
rlim_cur: 0,
|
||||
rlim_max: 0,
|
||||
};
|
||||
let ret = unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut rlim) };
|
||||
if ret == 0 {
|
||||
return rlim.rlim_cur;
|
||||
}
|
||||
}
|
||||
// Fallback for non-unix or error
|
||||
1024
|
||||
}
|
||||
@@ -0,0 +1,168 @@
|
||||
mod app;
|
||||
mod config;
|
||||
mod hardware;
|
||||
mod net;
|
||||
mod registration;
|
||||
mod runtime;
|
||||
mod setup;
|
||||
mod state;
|
||||
mod target_filter;
|
||||
mod tunnel;
|
||||
mod upstream_client;
|
||||
|
||||
use std::path::PathBuf;
|
||||
|
||||
use clap::{CommandFactory, FromArgMatches, Parser};
|
||||
|
||||
use config::Config;
|
||||
|
||||
/// Default config file name.
|
||||
const DEFAULT_CONFIG: &str = "aether-proxy.toml";
|
||||
|
||||
/// Build the full clap command: Config args + discoverable subcommands.
|
||||
///
|
||||
/// `subcommand_negates_reqs` lets subcommands bypass the required Config
|
||||
/// flags so that e.g. `aether-proxy setup` doesn't demand `--aether-url`.
|
||||
fn build_command() -> clap::Command {
|
||||
Config::command()
|
||||
.subcommand(
|
||||
clap::Command::new("setup")
|
||||
.about("Interactive setup wizard (TUI)")
|
||||
.arg(
|
||||
clap::Arg::new("config_path")
|
||||
.help("Path to config file")
|
||||
.default_value(DEFAULT_CONFIG),
|
||||
),
|
||||
)
|
||||
.subcommand(clap::Command::new("start").about("Start the systemd service"))
|
||||
.subcommand(clap::Command::new("status").about("Show service status"))
|
||||
.subcommand(clap::Command::new("logs").about("Tail service logs"))
|
||||
.subcommand(clap::Command::new("restart").about("Restart the systemd service"))
|
||||
.subcommand(clap::Command::new("stop").about("Stop the systemd service"))
|
||||
.subcommand(clap::Command::new("uninstall").about("Uninstall the systemd service"))
|
||||
.subcommand(
|
||||
clap::Command::new("upgrade")
|
||||
.about("Self-upgrade from GitHub releases")
|
||||
.arg(clap::Arg::new("version").help("Target version (e.g. 0.2.0)")),
|
||||
)
|
||||
.subcommand_negates_reqs(true)
|
||||
}
|
||||
|
||||
#[tokio::main]
|
||||
async fn main() -> anyhow::Result<()> {
|
||||
rustls::crypto::ring::default_provider()
|
||||
.install_default()
|
||||
.map_err(|_| anyhow::anyhow!("Failed to install rustls CryptoProvider"))?;
|
||||
|
||||
// Load config file as env-var defaults (before clap parsing)
|
||||
let config_file_path =
|
||||
std::env::var("AETHER_PROXY_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
|
||||
let config_path = std::path::Path::new(&config_file_path);
|
||||
if config_path.exists() {
|
||||
// Migrate legacy 0.1.x config to 0.2.0 format if needed
|
||||
if let Err(e) = config::ConfigFile::migrate_legacy(config_path) {
|
||||
eprintln!(" WARNING: config migration failed: {}", e);
|
||||
}
|
||||
if let Ok(file_cfg) = config::ConfigFile::load(config_path) {
|
||||
file_cfg.inject_env();
|
||||
}
|
||||
}
|
||||
|
||||
// Parse CLI (subcommands + config args in one pass)
|
||||
match build_command().try_get_matches() {
|
||||
Ok(matches) => match matches.subcommand() {
|
||||
Some(("setup", sub_m)) => {
|
||||
let path = sub_m
|
||||
.get_one::<String>("config_path")
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| PathBuf::from(DEFAULT_CONFIG));
|
||||
handle_setup_result(setup::run(path)?).await
|
||||
}
|
||||
Some(("start", _)) => setup::service::cmd_start(),
|
||||
Some(("status", _)) => setup::service::cmd_status(),
|
||||
Some(("logs", _)) => setup::service::cmd_logs(),
|
||||
Some(("restart", _)) => setup::service::cmd_restart(),
|
||||
Some(("stop", _)) => setup::service::cmd_stop(),
|
||||
Some(("uninstall", _)) => setup::service::cmd_uninstall(),
|
||||
Some(("upgrade", sub_m)) => {
|
||||
let version = sub_m.get_one::<String>("version").cloned();
|
||||
setup::upgrade::cmd_upgrade(version).await
|
||||
}
|
||||
Some(_) => unreachable!(),
|
||||
None => {
|
||||
// No subcommand — run the proxy with parsed config.
|
||||
let config = Config::from_arg_matches(&matches)?;
|
||||
run_proxy(config).await
|
||||
}
|
||||
},
|
||||
Err(e) => {
|
||||
if e.kind() == clap::error::ErrorKind::MissingRequiredArgument {
|
||||
eprintln!("Missing required config, launching setup wizard...\n");
|
||||
handle_setup_result(setup::run(PathBuf::from(&config_file_path))?).await
|
||||
} else {
|
||||
e.exit();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Decide what to do after the setup wizard completes.
|
||||
async fn handle_setup_result(outcome: setup::SetupOutcome) -> anyhow::Result<()> {
|
||||
match outcome {
|
||||
setup::SetupOutcome::ServiceInstalled => Ok(()),
|
||||
setup::SetupOutcome::ReadyToRun(config_path) => {
|
||||
// Reload config from the file that setup just wrote, overriding
|
||||
// any stale env vars from a previous config.
|
||||
match config::ConfigFile::load(&config_path) {
|
||||
Ok(file_cfg) => file_cfg.inject_env_override(),
|
||||
Err(e) => anyhow::bail!("failed to reload config after setup: {}", e),
|
||||
}
|
||||
// Parse from env-only (argv may still contain "setup" etc.)
|
||||
let config = Config::try_parse_from(["aether-proxy"])
|
||||
.map_err(|e| anyhow::anyhow!("config invalid after setup: {}", e))?;
|
||||
eprintln!(" Starting proxy...\n");
|
||||
run_proxy(config).await
|
||||
}
|
||||
setup::SetupOutcome::Cancelled => {
|
||||
eprintln!(" Setup cancelled.");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Start the proxy server, checking for systemd conflicts first.
|
||||
async fn run_proxy(config: Config) -> anyhow::Result<()> {
|
||||
// Warn if systemd service is already running (would cause port conflict).
|
||||
// Skip this check when we ARE the systemd service (INVOCATION_ID is set by systemd).
|
||||
if std::env::var_os("INVOCATION_ID").is_none() && setup::service::is_service_active() {
|
||||
eprintln!("Warning: systemd service is already running.");
|
||||
eprintln!("Use `./aether-proxy stop` to stop it first, or manage via subcommands:");
|
||||
eprintln!(" ./aether-proxy status / logs / restart / stop");
|
||||
std::process::exit(1);
|
||||
}
|
||||
|
||||
// Resolve server list: prefer [[servers]] from TOML, fall back to CLI/env single server.
|
||||
let config_path =
|
||||
std::env::var("AETHER_PROXY_CONFIG").unwrap_or_else(|_| DEFAULT_CONFIG.to_string());
|
||||
let servers = if std::path::Path::new(&config_path).exists() {
|
||||
config::ConfigFile::load(std::path::Path::new(&config_path))
|
||||
.ok()
|
||||
.map(|f| f.effective_servers())
|
||||
.filter(|s| !s.is_empty())
|
||||
.unwrap_or_else(|| {
|
||||
vec![config::ServerEntry {
|
||||
aether_url: config.aether_url.clone(),
|
||||
management_token: config.management_token.clone(),
|
||||
node_name: None,
|
||||
}]
|
||||
})
|
||||
} else {
|
||||
vec![config::ServerEntry {
|
||||
aether_url: config.aether_url.clone(),
|
||||
management_token: config.management_token.clone(),
|
||||
node_name: None,
|
||||
}]
|
||||
};
|
||||
|
||||
app::run(config, servers).await
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
//! Network utility functions (public IP detection, region detection).
|
||||
//!
|
||||
//! These are standalone helpers not tied to any specific client or service.
|
||||
|
||||
use reqwest::Client;
|
||||
use tracing::{debug, info};
|
||||
|
||||
/// Auto-detect public IP by querying external services.
|
||||
pub async fn detect_public_ip() -> anyhow::Result<String> {
|
||||
let endpoints = [
|
||||
"https://api.ipify.org",
|
||||
"https://ifconfig.me/ip",
|
||||
"https://icanhazip.com",
|
||||
];
|
||||
|
||||
let client = Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(5))
|
||||
.build()?;
|
||||
|
||||
for endpoint in &endpoints {
|
||||
match client.get(*endpoint).send().await {
|
||||
Ok(resp) if resp.status().is_success() => {
|
||||
let ip = resp.text().await?.trim().to_string();
|
||||
if !ip.is_empty() {
|
||||
info!(ip = %ip, source = %endpoint, "detected public IP");
|
||||
return Ok(ip);
|
||||
}
|
||||
}
|
||||
Ok(resp) => {
|
||||
debug!(endpoint = %endpoint, status = %resp.status(), "IP detection failed");
|
||||
}
|
||||
Err(e) => {
|
||||
debug!(endpoint = %endpoint, error = %e, "IP detection failed");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
anyhow::bail!("failed to detect public IP from any source; use --public-ip")
|
||||
}
|
||||
|
||||
/// Auto-detect geographic region from a public IP address.
|
||||
///
|
||||
/// Uses multiple providers with HTTPS preferred. Falls back to ip-api.com
|
||||
/// over plain HTTP (their free tier doesn't support HTTPS).
|
||||
/// This is best-effort and non-sensitive -- region detection should never
|
||||
/// block startup.
|
||||
pub async fn detect_region(ip: &str) -> Option<String> {
|
||||
// Try HTTPS provider first
|
||||
let https_url = format!("https://ipinfo.io/{}/country", ip);
|
||||
|
||||
let client = Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(5))
|
||||
.build()
|
||||
.ok()?;
|
||||
|
||||
// Try ipinfo.io (HTTPS, returns plain text country code)
|
||||
if let Ok(resp) = client.get(&https_url).send().await {
|
||||
if resp.status().is_success() {
|
||||
if let Ok(text) = resp.text().await {
|
||||
let code = text.trim();
|
||||
if !code.is_empty() && code.len() <= 3 {
|
||||
info!(region = %code, ip = %ip, source = "ipinfo.io", "detected region");
|
||||
return Some(code.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback: ip-api.com (HTTP only on free tier, non-sensitive data)
|
||||
let http_url = format!("http://ip-api.com/json/{}?fields=countryCode", ip);
|
||||
match client.get(&http_url).send().await {
|
||||
Ok(resp) if resp.status().is_success() => {
|
||||
let body: serde_json::Value = resp.json().await.ok()?;
|
||||
let code = body.get("countryCode")?.as_str()?;
|
||||
if code.is_empty() {
|
||||
return None;
|
||||
}
|
||||
info!(region = %code, ip = %ip, source = "ip-api.com", "detected region");
|
||||
Some(code.to_string())
|
||||
}
|
||||
_ => {
|
||||
debug!(ip = %ip, "region detection failed");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,257 @@
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use reqwest::{Client, StatusCode};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tokio::time::sleep;
|
||||
use tracing::{debug, error, info};
|
||||
|
||||
use crate::config::Config;
|
||||
use crate::hardware::HardwareInfo;
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct RegisterRequest {
|
||||
name: String,
|
||||
ip: String,
|
||||
port: u16,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
region: Option<String>,
|
||||
heartbeat_interval: u64,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
hardware_info: Option<serde_json::Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
estimated_max_concurrency: Option<u64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
proxy_metadata: Option<serde_json::Value>,
|
||||
tunnel_mode: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct RegisterResponse {
|
||||
pub node_id: String,
|
||||
}
|
||||
|
||||
/// Remote configuration pushed by the Aether management backend.
|
||||
#[derive(Debug, Clone, Deserialize)]
|
||||
pub struct RemoteConfig {
|
||||
pub node_name: Option<String>,
|
||||
pub allowed_ports: Option<Vec<u16>>,
|
||||
pub log_level: Option<String>,
|
||||
pub heartbeat_interval: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct UnregisterRequest {
|
||||
node_id: String,
|
||||
}
|
||||
|
||||
/// Aether API client for proxy node lifecycle management.
|
||||
pub struct AetherClient {
|
||||
http: Client,
|
||||
base_url: String,
|
||||
token: String,
|
||||
retry_max_attempts: u32,
|
||||
retry_base_delay: Duration,
|
||||
retry_max_delay: Duration,
|
||||
}
|
||||
|
||||
impl AetherClient {
|
||||
pub fn new(config: &Config, aether_url: &str, management_token: &str) -> Self {
|
||||
let mut builder = Client::builder()
|
||||
.timeout(Duration::from_secs(config.aether_request_timeout_secs))
|
||||
.connect_timeout(Duration::from_secs(config.aether_connect_timeout_secs))
|
||||
.pool_max_idle_per_host(config.aether_pool_max_idle_per_host)
|
||||
.pool_idle_timeout(Duration::from_secs(config.aether_pool_idle_timeout_secs))
|
||||
.tcp_nodelay(config.aether_tcp_nodelay);
|
||||
|
||||
if config.aether_tcp_keepalive_secs > 0 {
|
||||
builder =
|
||||
builder.tcp_keepalive(Some(Duration::from_secs(config.aether_tcp_keepalive_secs)));
|
||||
} else {
|
||||
builder = builder.tcp_keepalive(None);
|
||||
}
|
||||
|
||||
if config.aether_http2 {
|
||||
builder = builder.http2_adaptive_window(true);
|
||||
}
|
||||
|
||||
let http = builder.build().expect("failed to create HTTP client");
|
||||
|
||||
let retry_base_delay = Duration::from_millis(config.aether_retry_base_delay_ms);
|
||||
let retry_max_delay =
|
||||
Duration::from_millis(config.aether_retry_max_delay_ms).max(retry_base_delay);
|
||||
|
||||
Self {
|
||||
http,
|
||||
base_url: aether_url.trim_end_matches('/').to_string(),
|
||||
token: management_token.to_string(),
|
||||
retry_max_attempts: config.aether_retry_max_attempts.max(1),
|
||||
retry_base_delay,
|
||||
retry_max_delay,
|
||||
}
|
||||
}
|
||||
|
||||
/// Register this node with Aether (idempotent upsert by ip:port).
|
||||
///
|
||||
/// Returns the stable node_id assigned by Aether.
|
||||
pub async fn register(
|
||||
&self,
|
||||
config: &Config,
|
||||
node_name: &str,
|
||||
public_ip: &str,
|
||||
hw: Option<&HardwareInfo>,
|
||||
) -> anyhow::Result<String> {
|
||||
let url = format!("{}/api/admin/proxy-nodes/register", self.base_url);
|
||||
let body = RegisterRequest {
|
||||
name: node_name.to_string(),
|
||||
ip: public_ip.to_string(),
|
||||
port: 0,
|
||||
region: config.node_region.clone(),
|
||||
heartbeat_interval: config.heartbeat_interval,
|
||||
hardware_info: hw.and_then(|h| serde_json::to_value(h).ok()),
|
||||
estimated_max_concurrency: hw.map(|h| h.estimated_max_concurrency),
|
||||
proxy_metadata: Some(serde_json::json!({
|
||||
"version": env!("CARGO_PKG_VERSION"),
|
||||
})),
|
||||
tunnel_mode: true,
|
||||
};
|
||||
|
||||
info!(
|
||||
url = %url,
|
||||
name = %body.name,
|
||||
ip = %body.ip,
|
||||
"registering with Aether"
|
||||
);
|
||||
|
||||
let resp = self
|
||||
.send_with_retry(
|
||||
|| {
|
||||
self.http
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", self.token))
|
||||
.json(&body)
|
||||
},
|
||||
"register",
|
||||
)
|
||||
.await?;
|
||||
|
||||
let status = resp.status();
|
||||
if !status.is_success() {
|
||||
let text = resp.text().await.unwrap_or_default();
|
||||
anyhow::bail!("register failed (HTTP {}): {}", status, text);
|
||||
}
|
||||
|
||||
let data: RegisterResponse = resp.json().await?;
|
||||
info!(node_id = %data.node_id, "registered successfully");
|
||||
Ok(data.node_id)
|
||||
}
|
||||
|
||||
/// Unregister this node from Aether (graceful shutdown).
|
||||
pub async fn unregister(&self, node_id: &str) -> anyhow::Result<()> {
|
||||
let url = format!("{}/api/admin/proxy-nodes/unregister", self.base_url);
|
||||
let body = UnregisterRequest {
|
||||
node_id: node_id.to_string(),
|
||||
};
|
||||
|
||||
info!(node_id = %node_id, "unregistering from Aether");
|
||||
|
||||
let resp = self
|
||||
.send_with_retry(
|
||||
|| {
|
||||
self.http
|
||||
.post(&url)
|
||||
.header("Authorization", format!("Bearer {}", self.token))
|
||||
.json(&body)
|
||||
},
|
||||
"unregister",
|
||||
)
|
||||
.await;
|
||||
|
||||
match resp {
|
||||
Ok(r) if r.status().is_success() => {
|
||||
info!(node_id = %node_id, "unregistered successfully");
|
||||
Ok(())
|
||||
}
|
||||
Ok(r) => {
|
||||
let text = r.text().await.unwrap_or_default();
|
||||
error!(body = %text, "unregister failed");
|
||||
anyhow::bail!("unregister failed: {}", text);
|
||||
}
|
||||
Err(e) => {
|
||||
// Best-effort during shutdown
|
||||
error!(error = %e, "unregister request failed");
|
||||
anyhow::bail!("unregister request failed: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_with_retry<F>(
|
||||
&self,
|
||||
mut make_req: F,
|
||||
label: &str,
|
||||
) -> Result<reqwest::Response, reqwest::Error>
|
||||
where
|
||||
F: FnMut() -> reqwest::RequestBuilder,
|
||||
{
|
||||
let mut attempt: u32 = 0;
|
||||
let mut delay = self.retry_base_delay;
|
||||
|
||||
loop {
|
||||
attempt = attempt.saturating_add(1);
|
||||
let resp = make_req().send().await;
|
||||
match resp {
|
||||
Ok(resp) => {
|
||||
if should_retry_status(resp.status()) && attempt < self.retry_max_attempts {
|
||||
let sleep_for = jitter_delay(delay);
|
||||
debug!(
|
||||
attempt,
|
||||
status = %resp.status(),
|
||||
sleep_ms = sleep_for.as_millis(),
|
||||
label,
|
||||
"Aether request retrying"
|
||||
);
|
||||
sleep(sleep_for).await;
|
||||
let next_delay = delay.checked_mul(2).unwrap_or(self.retry_max_delay);
|
||||
delay = std::cmp::min(next_delay, self.retry_max_delay);
|
||||
continue;
|
||||
}
|
||||
return Ok(resp);
|
||||
}
|
||||
Err(e) => {
|
||||
if attempt < self.retry_max_attempts {
|
||||
let sleep_for = jitter_delay(delay);
|
||||
debug!(
|
||||
attempt,
|
||||
error = %e,
|
||||
sleep_ms = sleep_for.as_millis(),
|
||||
label,
|
||||
"Aether request retrying"
|
||||
);
|
||||
sleep(sleep_for).await;
|
||||
let next_delay = delay.checked_mul(2).unwrap_or(self.retry_max_delay);
|
||||
delay = std::cmp::min(next_delay, self.retry_max_delay);
|
||||
continue;
|
||||
}
|
||||
return Err(e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn should_retry_status(status: StatusCode) -> bool {
|
||||
status.is_server_error()
|
||||
|| status == StatusCode::TOO_MANY_REQUESTS
|
||||
|| status == StatusCode::REQUEST_TIMEOUT
|
||||
}
|
||||
|
||||
fn jitter_delay(base: Duration) -> Duration {
|
||||
if base.is_zero() {
|
||||
return base;
|
||||
}
|
||||
let nanos = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.subsec_nanos() as u64)
|
||||
.unwrap_or(0);
|
||||
let jitter_ms = nanos % 100;
|
||||
base + Duration::from_millis(jitter_ms)
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
pub mod client;
|
||||
@@ -0,0 +1,121 @@
|
||||
//! Runtime-mutable configuration that can be updated remotely via heartbeat.
|
||||
//!
|
||||
//! Fields in [`DynamicConfig`] are initially populated from the static
|
||||
//! [`Config`](crate::config::Config) and may be overridden by the Aether
|
||||
//! management backend through the heartbeat response.
|
||||
|
||||
use std::collections::HashSet;
|
||||
use std::sync::{Arc, OnceLock};
|
||||
|
||||
use arc_swap::ArcSwap;
|
||||
use tracing::info;
|
||||
|
||||
use crate::config::Config;
|
||||
|
||||
/// Configuration that can be changed at runtime without restart.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct DynamicConfig {
|
||||
pub node_name: String,
|
||||
pub allowed_ports: Arc<HashSet<u16>>,
|
||||
pub log_level: String,
|
||||
pub heartbeat_interval: u64,
|
||||
/// Monotonically increasing version from the backend.
|
||||
/// `0` means no remote config has ever been applied.
|
||||
pub config_version: u64,
|
||||
}
|
||||
|
||||
impl DynamicConfig {
|
||||
/// Initialize from static config (startup defaults).
|
||||
pub fn from_config(config: &Config) -> Self {
|
||||
Self {
|
||||
node_name: config.node_name.clone(),
|
||||
allowed_ports: Arc::new(config.allowed_ports.iter().copied().collect()),
|
||||
log_level: config.log_level.clone(),
|
||||
heartbeat_interval: config.heartbeat_interval,
|
||||
config_version: 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared dynamic config handle (lock-free reads via ArcSwap).
|
||||
pub type SharedDynamicConfig = Arc<ArcSwap<DynamicConfig>>;
|
||||
|
||||
// -- Log-level hot-reload -----
|
||||
|
||||
/// Global log-level reloader function, set during tracing init.
|
||||
type LogReloader = Box<dyn Fn(&str) + Send + Sync>;
|
||||
|
||||
static LOG_RELOADER: OnceLock<LogReloader> = OnceLock::new();
|
||||
|
||||
/// Register the log-level reload function (called once from `init_tracing`).
|
||||
pub fn set_log_reloader(f: LogReloader) {
|
||||
let _ = LOG_RELOADER.set(f);
|
||||
}
|
||||
|
||||
/// Apply a remote config update to the dynamic config.
|
||||
///
|
||||
/// Uses copy-on-write: loads the current snapshot, clones it, applies changes,
|
||||
/// and stores the new Arc. Reads are always lock-free.
|
||||
///
|
||||
/// Returns `true` if the config was actually changed.
|
||||
pub fn apply_remote_config(
|
||||
dynamic: &SharedDynamicConfig,
|
||||
remote: &crate::registration::client::RemoteConfig,
|
||||
version: u64,
|
||||
) -> bool {
|
||||
let current = dynamic.load();
|
||||
|
||||
if version <= current.config_version {
|
||||
return false;
|
||||
}
|
||||
|
||||
let mut new_cfg = (**current).clone();
|
||||
let mut changed = Vec::new();
|
||||
|
||||
if let Some(ref name) = remote.node_name {
|
||||
if *name != new_cfg.node_name {
|
||||
changed.push(format!("node_name -> {}", name));
|
||||
new_cfg.node_name = name.clone();
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(ref ports) = remote.allowed_ports {
|
||||
let new_set: HashSet<u16> = ports.iter().copied().collect();
|
||||
if new_set != *new_cfg.allowed_ports {
|
||||
changed.push(format!("allowed_ports -> {:?}", ports));
|
||||
new_cfg.allowed_ports = Arc::new(new_set);
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(interval) = remote.heartbeat_interval {
|
||||
if interval != new_cfg.heartbeat_interval {
|
||||
changed.push(format!("heartbeat_interval -> {}s", interval));
|
||||
new_cfg.heartbeat_interval = interval;
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(ref level) = remote.log_level {
|
||||
if *level != new_cfg.log_level {
|
||||
changed.push(format!("log_level -> {}", level));
|
||||
new_cfg.log_level = level.clone();
|
||||
// Hot-reload tracing filter
|
||||
if let Some(reloader) = LOG_RELOADER.get() {
|
||||
reloader(level);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let has_changes = !changed.is_empty();
|
||||
|
||||
if has_changes {
|
||||
new_cfg.config_version = version;
|
||||
info!(
|
||||
version,
|
||||
changes = %changed.join(", "),
|
||||
"remote config applied"
|
||||
);
|
||||
dynamic.store(Arc::new(new_cfg));
|
||||
}
|
||||
|
||||
has_changes
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
//! Safe DNS resolver for reqwest that reuses validated addresses from DnsCache.
|
||||
//!
|
||||
//! This resolver ensures reqwest connects only to addresses that have been
|
||||
//! previously validated by `target_filter::validate_target()`, eliminating
|
||||
//! the TOCTTOU gap where DNS rebinding could redirect traffic to private IPs.
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::Arc;
|
||||
|
||||
use reqwest::dns::{Addrs, Name, Resolve, Resolving};
|
||||
|
||||
use crate::target_filter::{self, DnsCache};
|
||||
|
||||
/// A DNS resolver that serves validated public addresses from the shared DnsCache.
|
||||
///
|
||||
/// When reqwest needs to resolve a hostname, this resolver returns addresses
|
||||
/// from the cache (populated by `validate_target()` during request validation).
|
||||
/// If the hostname is not in cache (shouldn't happen in normal flow), it
|
||||
/// performs a fresh resolution with private-IP filtering.
|
||||
pub struct SafeDnsResolver {
|
||||
dns_cache: Arc<DnsCache>,
|
||||
}
|
||||
|
||||
impl SafeDnsResolver {
|
||||
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
|
||||
Self { dns_cache }
|
||||
}
|
||||
}
|
||||
|
||||
impl Resolve for SafeDnsResolver {
|
||||
fn resolve(&self, name: Name) -> Resolving {
|
||||
let dns_cache = Arc::clone(&self.dns_cache);
|
||||
Box::pin(async move {
|
||||
let host = name.as_str();
|
||||
|
||||
// Try cache first (should be populated by validate_target).
|
||||
// reqwest resolves by hostname only (no port), so use host-only lookup.
|
||||
if let Some(addrs) = dns_cache.get_by_host(host).await {
|
||||
let socket_addrs: Vec<SocketAddr> = (*addrs).clone();
|
||||
return Ok(Box::new(socket_addrs.into_iter()) as Addrs);
|
||||
}
|
||||
|
||||
// Fallback: resolve with private-IP filtering (defensive).
|
||||
// This path should rarely be hit since validate_target() runs first.
|
||||
// We don't know the real port here (reqwest Resolve only gives hostname),
|
||||
// so resolve directly without caching to avoid polluting the cache with
|
||||
// an incorrect port-based key.
|
||||
let addr_str = format!("{}:0", host);
|
||||
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
|
||||
.await
|
||||
.map_err(|e| -> Box<dyn std::error::Error + Send + Sync> { Box::new(e) })?
|
||||
.filter(|addr| !target_filter::is_private_ip(&addr.ip()))
|
||||
.collect();
|
||||
|
||||
if resolved.is_empty() {
|
||||
return Err(Box::new(std::io::Error::other(format!(
|
||||
"all resolved addresses for {} are private/reserved",
|
||||
host
|
||||
)))
|
||||
as Box<dyn std::error::Error + Send + Sync>);
|
||||
}
|
||||
|
||||
Ok(Box::new(resolved.into_iter()) as Addrs)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
pub(crate) mod service;
|
||||
mod tui;
|
||||
pub(crate) mod upgrade;
|
||||
|
||||
pub use self::tui::{run, SetupOutcome};
|
||||
@@ -0,0 +1,256 @@
|
||||
//! Systemd service installation for aether-proxy.
|
||||
//!
|
||||
//! Called from the setup TUI when the user enables "Install Service".
|
||||
//! The unit file points to the binary and config at their current
|
||||
//! absolute paths -- no files are copied.
|
||||
|
||||
use std::path::Path;
|
||||
use std::process::Command;
|
||||
|
||||
const UNIT_PATH: &str = "/etc/systemd/system/aether-proxy.service";
|
||||
const SERVICE_NAME: &str = "aether-proxy";
|
||||
|
||||
/// Whether systemd service installation is possible (systemd present + root).
|
||||
pub fn is_available() -> bool {
|
||||
is_systemd_available() && is_root()
|
||||
}
|
||||
|
||||
/// Install aether-proxy as a systemd service. Must be run as root.
|
||||
pub fn install_service(config_path: &Path) -> anyhow::Result<()> {
|
||||
if !is_systemd_available() {
|
||||
anyhow::bail!("systemd not available");
|
||||
}
|
||||
if !is_root() {
|
||||
anyhow::bail!("root required, use: sudo ./aether-proxy setup");
|
||||
}
|
||||
|
||||
let exe_path = std::env::current_exe()?.canonicalize()?;
|
||||
let exe_str = exe_path
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("binary path contains invalid UTF-8"))?;
|
||||
|
||||
let config_abs = std::fs::canonicalize(config_path)?;
|
||||
let config_str = config_abs
|
||||
.to_str()
|
||||
.ok_or_else(|| anyhow::anyhow!("config path contains invalid UTF-8"))?;
|
||||
|
||||
let working_dir = config_abs
|
||||
.parent()
|
||||
.unwrap_or_else(|| Path::new("/"))
|
||||
.to_str()
|
||||
.unwrap_or("/");
|
||||
|
||||
// Stop existing service if running (ignore errors)
|
||||
if Path::new(UNIT_PATH).exists() {
|
||||
eprintln!(" Stopping existing service...");
|
||||
let _ = Command::new("systemctl")
|
||||
.args(["stop", SERVICE_NAME])
|
||||
.status();
|
||||
}
|
||||
|
||||
// Write unit file
|
||||
eprintln!(" Generating systemd unit file...");
|
||||
eprintln!(" Binary: {}", exe_str);
|
||||
eprintln!(" Config: {}", config_str);
|
||||
eprintln!(" WorkDir: {}", working_dir);
|
||||
|
||||
let unit_content = format!(
|
||||
"[Unit]\n\
|
||||
Description=Aether Proxy\n\
|
||||
After=network.target\n\
|
||||
\n\
|
||||
[Service]\n\
|
||||
Type=simple\n\
|
||||
WorkingDirectory={working_dir}\n\
|
||||
Environment=AETHER_PROXY_CONFIG={config_str}\n\
|
||||
ExecStart={exe_str}\n\
|
||||
Restart=on-failure\n\
|
||||
RestartSec=5\n\
|
||||
LimitNOFILE=65535\n\
|
||||
UMask=0077\n\
|
||||
\n\
|
||||
[Install]\n\
|
||||
WantedBy=multi-user.target\n",
|
||||
);
|
||||
std::fs::write(UNIT_PATH, &unit_content)?;
|
||||
|
||||
// Reload and enable
|
||||
eprintln!(" Enabling and starting service...");
|
||||
run_cmd("systemctl", &["daemon-reload"])?;
|
||||
run_cmd("systemctl", &["enable", "--now", SERVICE_NAME])?;
|
||||
|
||||
// Verify
|
||||
eprintln!();
|
||||
let output = Command::new("systemctl")
|
||||
.args(["is-active", SERVICE_NAME])
|
||||
.output()?;
|
||||
let state = String::from_utf8_lossy(&output.stdout).trim().to_string();
|
||||
|
||||
if state == "active" {
|
||||
eprintln!(" Service started successfully!");
|
||||
} else {
|
||||
eprintln!(" Service state: {} (check logs)", state);
|
||||
}
|
||||
|
||||
eprintln!();
|
||||
eprintln!(" Commands:");
|
||||
eprintln!(" ./aether-proxy status # service status");
|
||||
eprintln!(" ./aether-proxy logs # tail logs");
|
||||
eprintln!(" sudo ./aether-proxy restart # restart");
|
||||
eprintln!(" sudo ./aether-proxy stop # stop");
|
||||
eprintln!(" sudo ./aether-proxy uninstall # remove service");
|
||||
eprintln!();
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn is_systemd_available() -> bool {
|
||||
Command::new("systemctl")
|
||||
.arg("--version")
|
||||
.stdout(std::process::Stdio::null())
|
||||
.stderr(std::process::Stdio::null())
|
||||
.status()
|
||||
.map(|s| s.success())
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
pub(crate) fn is_root() -> bool {
|
||||
#[cfg(unix)]
|
||||
{
|
||||
unsafe { libc::geteuid() == 0 }
|
||||
}
|
||||
#[cfg(not(unix))]
|
||||
{
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// Whether a systemd unit file is currently installed.
|
||||
pub fn is_installed() -> bool {
|
||||
Path::new(UNIT_PATH).exists()
|
||||
}
|
||||
|
||||
/// Remove the systemd service (called from setup TUI when Install Service is toggled off).
|
||||
pub fn uninstall_service() -> anyhow::Result<()> {
|
||||
if !Path::new(UNIT_PATH).exists() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
eprintln!(" Stopping and removing existing service...");
|
||||
let _ = Command::new("systemctl")
|
||||
.args(["disable", "--now", SERVICE_NAME])
|
||||
.status();
|
||||
|
||||
std::fs::remove_file(UNIT_PATH)?;
|
||||
eprintln!(" Removed {}", UNIT_PATH);
|
||||
run_cmd("systemctl", &["daemon-reload"])?;
|
||||
eprintln!(" Service uninstalled.");
|
||||
eprintln!();
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Check if the systemd service is currently active.
|
||||
pub fn is_service_active() -> bool {
|
||||
std::path::Path::new(UNIT_PATH).exists()
|
||||
&& Command::new("systemctl")
|
||||
.args(["is-active", "--quiet", SERVICE_NAME])
|
||||
.stdout(std::process::Stdio::null())
|
||||
.stderr(std::process::Stdio::null())
|
||||
.status()
|
||||
.map(|s| s.success())
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
// ── CLI subcommands (systemd wrappers) ──────────────────────────────────────
|
||||
|
||||
fn ensure_service_installed() -> anyhow::Result<()> {
|
||||
if !std::path::Path::new(UNIT_PATH).exists() {
|
||||
anyhow::bail!("service not installed, run `sudo ./aether-proxy setup` first");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn ensure_root_and_service() -> anyhow::Result<()> {
|
||||
ensure_service_installed()?;
|
||||
if !is_root() {
|
||||
anyhow::bail!("root required, use: sudo ./aether-proxy <command>");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// `aether-proxy status` -- show service status.
|
||||
pub fn cmd_status() -> anyhow::Result<()> {
|
||||
ensure_service_installed()?;
|
||||
let status = Command::new("systemctl")
|
||||
.args(["status", SERVICE_NAME])
|
||||
.status()?;
|
||||
// systemctl status returns non-zero when inactive; that's fine
|
||||
std::process::exit(status.code().unwrap_or(1));
|
||||
}
|
||||
|
||||
/// `aether-proxy logs` -- tail service logs.
|
||||
pub fn cmd_logs() -> anyhow::Result<()> {
|
||||
ensure_service_installed()?;
|
||||
let status = Command::new("journalctl")
|
||||
.args(["-u", SERVICE_NAME, "-f", "--no-pager", "-n", "100"])
|
||||
.status()?;
|
||||
std::process::exit(status.code().unwrap_or(1));
|
||||
}
|
||||
|
||||
/// `aether-proxy start` -- start the service.
|
||||
pub fn cmd_start() -> anyhow::Result<()> {
|
||||
ensure_root_and_service()?;
|
||||
run_cmd("systemctl", &["start", SERVICE_NAME])?;
|
||||
eprintln!(" Service started.");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// `aether-proxy restart` -- restart the service.
|
||||
pub fn cmd_restart() -> anyhow::Result<()> {
|
||||
ensure_root_and_service()?;
|
||||
run_cmd("systemctl", &["restart", SERVICE_NAME])?;
|
||||
eprintln!(" Service restarted.");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// `aether-proxy stop` -- stop the service.
|
||||
pub fn cmd_stop() -> anyhow::Result<()> {
|
||||
ensure_root_and_service()?;
|
||||
run_cmd("systemctl", &["stop", SERVICE_NAME])?;
|
||||
eprintln!(" Service stopped.");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// `aether-proxy uninstall` -- disable and remove the systemd service.
|
||||
pub fn cmd_uninstall() -> anyhow::Result<()> {
|
||||
ensure_root_and_service()?;
|
||||
|
||||
eprintln!(" Stopping and disabling service...");
|
||||
let _ = Command::new("systemctl")
|
||||
.args(["disable", "--now", SERVICE_NAME])
|
||||
.status();
|
||||
|
||||
if std::path::Path::new(UNIT_PATH).exists() {
|
||||
std::fs::remove_file(UNIT_PATH)?;
|
||||
eprintln!(" Removed {}", UNIT_PATH);
|
||||
}
|
||||
|
||||
run_cmd("systemctl", &["daemon-reload"])?;
|
||||
eprintln!(" Service uninstalled.");
|
||||
eprintln!();
|
||||
eprintln!(" Config file and TLS certs are preserved. Remove manually if needed.");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn run_cmd(program: &str, args: &[&str]) -> anyhow::Result<()> {
|
||||
let display = format!("{} {}", program, args.join(" "));
|
||||
eprintln!(" > {}", display);
|
||||
|
||||
let status = Command::new(program).args(args).status()?;
|
||||
if !status.success() {
|
||||
anyhow::bail!("command failed: {}", display);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,868 @@
|
||||
//! Interactive TUI for configuring aether-proxy.
|
||||
//!
|
||||
//! Launched via `aether-proxy setup [path]`. Presents a full-screen form
|
||||
//! backed by ratatui where the user can navigate fields, edit values, and
|
||||
//! save to a TOML config file. Supports multi-server configuration via
|
||||
//! a tabbed interface.
|
||||
|
||||
use std::io;
|
||||
use std::path::PathBuf;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use crossterm::event::{self, Event, KeyCode, KeyEvent, KeyEventKind, KeyModifiers};
|
||||
use crossterm::execute;
|
||||
use crossterm::terminal::{self, EnterAlternateScreen, LeaveAlternateScreen};
|
||||
use ratatui::backend::CrosstermBackend;
|
||||
use ratatui::layout::{Constraint, Layout, Rect};
|
||||
use ratatui::style::{Color, Modifier, Style};
|
||||
use ratatui::text::{Line, Span};
|
||||
use ratatui::widgets::{Block, Borders, Paragraph};
|
||||
use ratatui::Frame;
|
||||
use ratatui::Terminal;
|
||||
|
||||
use crate::config::{ConfigFile, ServerEntry};
|
||||
|
||||
/// Outcome of the setup wizard, returned to the caller.
|
||||
pub enum SetupOutcome {
|
||||
/// Config saved; systemd service installed and started.
|
||||
ServiceInstalled,
|
||||
/// Config saved; no service -- caller should start the proxy directly.
|
||||
ReadyToRun(PathBuf),
|
||||
/// User quit without saving.
|
||||
Cancelled,
|
||||
}
|
||||
|
||||
/// Column width reserved for the field label (chars).
|
||||
const LABEL_WIDTH: usize = 22;
|
||||
|
||||
// -- Field types --------------------------------------------------------------
|
||||
|
||||
#[derive(Clone, Copy, PartialEq)]
|
||||
enum FieldKind {
|
||||
Text,
|
||||
Secret,
|
||||
Bool,
|
||||
LogLevel,
|
||||
}
|
||||
|
||||
struct Field {
|
||||
label: &'static str,
|
||||
key: &'static str,
|
||||
value: String,
|
||||
kind: FieldKind,
|
||||
required: bool,
|
||||
help: &'static str,
|
||||
}
|
||||
// -- Server tab ---------------------------------------------------------------
|
||||
|
||||
/// A single server tab's editable fields.
|
||||
struct ServerTab {
|
||||
fields: Vec<Field>,
|
||||
}
|
||||
|
||||
impl ServerTab {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
fields: vec![
|
||||
Field {
|
||||
label: "Aether URL",
|
||||
key: "aether_url",
|
||||
value: String::new(),
|
||||
kind: FieldKind::Text,
|
||||
required: true,
|
||||
help: "Aether URL (e.g. https://aether.example.com)",
|
||||
},
|
||||
Field {
|
||||
label: "Management Token",
|
||||
key: "management_token",
|
||||
value: String::new(),
|
||||
kind: FieldKind::Secret,
|
||||
required: true,
|
||||
help: "Aether Management Token (ae_xxx)",
|
||||
},
|
||||
Field {
|
||||
label: "Node Name",
|
||||
key: "node_name",
|
||||
value: "proxy-01".into(),
|
||||
kind: FieldKind::Text,
|
||||
required: true,
|
||||
help: "Node name for identification in Aether dashboard",
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
|
||||
fn from_entry(entry: &ServerEntry) -> Self {
|
||||
let mut tab = Self::new();
|
||||
tab.fields[0].value = entry.aether_url.clone();
|
||||
tab.fields[1].value = entry.management_token.clone();
|
||||
if let Some(ref name) = entry.node_name {
|
||||
tab.fields[2].value = name.clone();
|
||||
}
|
||||
tab
|
||||
}
|
||||
}
|
||||
|
||||
// -- App state ----------------------------------------------------------------
|
||||
|
||||
#[derive(PartialEq)]
|
||||
enum Mode {
|
||||
Normal,
|
||||
Editing,
|
||||
}
|
||||
|
||||
struct App {
|
||||
server_tabs: Vec<ServerTab>,
|
||||
active_tab: usize,
|
||||
global_fields: Vec<Field>,
|
||||
selected: usize,
|
||||
mode: Mode,
|
||||
edit_buffer: String,
|
||||
edit_cursor: usize,
|
||||
config_path: PathBuf,
|
||||
modified: bool,
|
||||
message: Option<(String, Instant, bool)>,
|
||||
scroll_offset: usize,
|
||||
saved_once: bool,
|
||||
pending_quit: bool,
|
||||
confirm_delete: bool,
|
||||
}
|
||||
impl App {
|
||||
fn new(config_path: PathBuf) -> Self {
|
||||
Self {
|
||||
server_tabs: vec![ServerTab::new()],
|
||||
active_tab: 0,
|
||||
global_fields: vec![
|
||||
Field {
|
||||
label: "Log Level",
|
||||
key: "log_level",
|
||||
value: "info".into(),
|
||||
kind: FieldKind::LogLevel,
|
||||
required: true,
|
||||
help: "Log level -- Enter to cycle: trace / debug / info / warn / error",
|
||||
},
|
||||
Field {
|
||||
label: "Log JSON",
|
||||
key: "log_json",
|
||||
value: "false".into(),
|
||||
kind: FieldKind::Bool,
|
||||
required: true,
|
||||
help: "Output logs as JSON -- Enter to toggle",
|
||||
},
|
||||
Field {
|
||||
label: "Install Service",
|
||||
key: "install_service",
|
||||
value: if super::service::is_available() {
|
||||
"true"
|
||||
} else {
|
||||
"false"
|
||||
}
|
||||
.into(),
|
||||
kind: FieldKind::Bool,
|
||||
required: true,
|
||||
help: "Install as systemd service (requires root) -- Enter to toggle",
|
||||
},
|
||||
],
|
||||
selected: 0,
|
||||
mode: Mode::Normal,
|
||||
edit_buffer: String::new(),
|
||||
edit_cursor: 0,
|
||||
config_path,
|
||||
modified: false,
|
||||
message: None,
|
||||
scroll_offset: 0,
|
||||
saved_once: false,
|
||||
pending_quit: false,
|
||||
confirm_delete: false,
|
||||
}
|
||||
}
|
||||
|
||||
// -- Field accessors (unified index across server + global) ---------------
|
||||
|
||||
fn server_field_count(&self) -> usize {
|
||||
self.server_tabs[self.active_tab].fields.len()
|
||||
}
|
||||
|
||||
fn total_field_count(&self) -> usize {
|
||||
self.server_field_count() + self.global_fields.len()
|
||||
}
|
||||
|
||||
fn selected_field(&self) -> &Field {
|
||||
let sc = self.server_field_count();
|
||||
if self.selected < sc {
|
||||
&self.server_tabs[self.active_tab].fields[self.selected]
|
||||
} else {
|
||||
&self.global_fields[self.selected - sc]
|
||||
}
|
||||
}
|
||||
|
||||
fn selected_field_mut(&mut self) -> &mut Field {
|
||||
let sc = self.server_field_count();
|
||||
if self.selected < sc {
|
||||
&mut self.server_tabs[self.active_tab].fields[self.selected]
|
||||
} else {
|
||||
&mut self.global_fields[self.selected - sc]
|
||||
}
|
||||
}
|
||||
|
||||
fn clamp_selection(&mut self) {
|
||||
let max = self.total_field_count();
|
||||
if self.selected >= max {
|
||||
self.selected = max.saturating_sub(1);
|
||||
}
|
||||
self.scroll_offset = 0;
|
||||
self.confirm_delete = false;
|
||||
}
|
||||
// -- Config <-> fields -----------------------------------------------------
|
||||
|
||||
fn load_from_file(&mut self) {
|
||||
if let Ok(cfg) = ConfigFile::load(&self.config_path) {
|
||||
self.apply_config(&cfg);
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_config(&mut self, cfg: &ConfigFile) {
|
||||
// Global fields
|
||||
for field in &mut self.global_fields {
|
||||
let val: Option<String> = match field.key {
|
||||
"log_level" => cfg.log_level.clone(),
|
||||
"log_json" => cfg.log_json.map(|v| v.to_string()),
|
||||
_ => None,
|
||||
};
|
||||
if let Some(v) = val {
|
||||
field.value = v;
|
||||
}
|
||||
}
|
||||
|
||||
// Server tabs
|
||||
let servers = cfg.effective_servers();
|
||||
if servers.is_empty() {
|
||||
let mut tab = ServerTab::new();
|
||||
// Single-server fallback: use top-level node_name
|
||||
if let Some(ref name) = cfg.node_name {
|
||||
tab.fields[2].value = name.clone();
|
||||
}
|
||||
self.server_tabs = vec![tab];
|
||||
} else {
|
||||
self.server_tabs = servers.iter().map(ServerTab::from_entry).collect();
|
||||
// For single-server mode, node_name might be in top-level only
|
||||
if self.server_tabs.len() == 1 && self.server_tabs[0].fields[2].value.is_empty() {
|
||||
if let Some(ref name) = cfg.node_name {
|
||||
self.server_tabs[0].fields[2].value = name.clone();
|
||||
}
|
||||
}
|
||||
}
|
||||
self.active_tab = 0;
|
||||
self.selected = 0;
|
||||
self.scroll_offset = 0;
|
||||
}
|
||||
|
||||
fn to_config(&self) -> ConfigFile {
|
||||
let get_global = |key: &str| -> Option<String> {
|
||||
self.global_fields
|
||||
.iter()
|
||||
.find(|f| f.key == key)
|
||||
.map(|f| f.value.clone())
|
||||
.filter(|v| !v.is_empty())
|
||||
};
|
||||
|
||||
let get_tab = |tab: &ServerTab, key: &str| -> Option<String> {
|
||||
tab.fields
|
||||
.iter()
|
||||
.find(|f| f.key == key)
|
||||
.map(|f| f.value.clone())
|
||||
.filter(|v| !v.is_empty())
|
||||
};
|
||||
|
||||
let mut cfg = ConfigFile {
|
||||
log_level: get_global("log_level"),
|
||||
log_json: get_global("log_json").and_then(|v| v.parse().ok()),
|
||||
..ConfigFile::default()
|
||||
};
|
||||
|
||||
// Always write [[servers]] format; old top-level fields are read-only compat
|
||||
cfg.servers = self
|
||||
.server_tabs
|
||||
.iter()
|
||||
.map(|tab| ServerEntry {
|
||||
aether_url: get_tab(tab, "aether_url").unwrap_or_default(),
|
||||
management_token: get_tab(tab, "management_token").unwrap_or_default(),
|
||||
node_name: get_tab(tab, "node_name"),
|
||||
})
|
||||
.collect();
|
||||
cfg
|
||||
}
|
||||
|
||||
fn save(&mut self) -> anyhow::Result<()> {
|
||||
let cfg = self.to_config();
|
||||
cfg.save(&self.config_path)?;
|
||||
// Restrict config file permissions to owner-only (contains management token).
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
let _ =
|
||||
std::fs::set_permissions(&self.config_path, std::fs::Permissions::from_mode(0o600));
|
||||
}
|
||||
self.modified = false;
|
||||
self.saved_once = true;
|
||||
self.message = Some((
|
||||
format!("saved to {}", self.config_path.display()),
|
||||
Instant::now(),
|
||||
false,
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
// -- Scrolling ---------------------------------------------------------------
|
||||
|
||||
fn ensure_visible(&mut self, visible_rows: usize) {
|
||||
if visible_rows == 0 {
|
||||
return;
|
||||
}
|
||||
// Account for separator line between server and global fields
|
||||
let display_row = if self.selected >= self.server_field_count() {
|
||||
self.selected + 1
|
||||
} else {
|
||||
self.selected
|
||||
};
|
||||
if display_row < self.scroll_offset {
|
||||
self.scroll_offset = display_row;
|
||||
} else if display_row >= self.scroll_offset + visible_rows {
|
||||
self.scroll_offset = display_row - visible_rows + 1;
|
||||
}
|
||||
}
|
||||
|
||||
// -- Key handling -------------------------------------------------------------
|
||||
|
||||
/// Returns `true` when the app should exit.
|
||||
fn handle_key(&mut self, key: KeyEvent) -> bool {
|
||||
// Expire old messages (but keep quit-confirmation messages alive)
|
||||
if let Some((_, when, _)) = &self.message {
|
||||
if !self.pending_quit && !self.confirm_delete && when.elapsed() > Duration::from_secs(4)
|
||||
{
|
||||
self.message = None;
|
||||
}
|
||||
}
|
||||
|
||||
match self.mode {
|
||||
Mode::Normal => self.handle_normal(key),
|
||||
Mode::Editing => {
|
||||
self.handle_edit(key);
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn handle_normal(&mut self, key: KeyEvent) -> bool {
|
||||
// -- Quit handling (with unsaved-changes confirmation) -----------------
|
||||
let is_quit_key = matches!(key.code, KeyCode::Char('q') | KeyCode::Esc);
|
||||
|
||||
if is_quit_key {
|
||||
if !self.modified || self.pending_quit {
|
||||
return true;
|
||||
}
|
||||
self.pending_quit = true;
|
||||
self.confirm_delete = false;
|
||||
self.message = Some((
|
||||
"unsaved changes! q again to discard, ^S to save".into(),
|
||||
Instant::now(),
|
||||
true,
|
||||
));
|
||||
return false;
|
||||
}
|
||||
|
||||
// Any other key cancels pending quit / pending delete
|
||||
if self.pending_quit {
|
||||
self.pending_quit = false;
|
||||
self.message = None;
|
||||
}
|
||||
if self.confirm_delete && !matches!(key.code, KeyCode::Delete | KeyCode::Char('x')) {
|
||||
self.confirm_delete = false;
|
||||
self.message = None;
|
||||
}
|
||||
|
||||
match key.code {
|
||||
KeyCode::Char('s')
|
||||
if key.modifiers.contains(KeyModifiers::CONTROL)
|
||||
|| key.modifiers.contains(KeyModifiers::SUPER) =>
|
||||
{
|
||||
if let Err(e) = self.save() {
|
||||
self.message = Some((format!("error: {}", e), Instant::now(), true));
|
||||
}
|
||||
}
|
||||
KeyCode::Up | KeyCode::Char('k') => {
|
||||
self.selected = self.selected.saturating_sub(1);
|
||||
}
|
||||
KeyCode::Down | KeyCode::Char('j') => {
|
||||
if self.selected + 1 < self.total_field_count() {
|
||||
self.selected += 1;
|
||||
}
|
||||
}
|
||||
KeyCode::Home => self.selected = 0,
|
||||
KeyCode::End => self.selected = self.total_field_count() - 1,
|
||||
KeyCode::Enter | KeyCode::Char(' ') => {
|
||||
let kind = self.selected_field().kind;
|
||||
let key_str = self.selected_field().key;
|
||||
let value = self.selected_field().value.clone();
|
||||
match kind {
|
||||
FieldKind::Bool => {
|
||||
let toggled = if value == "true" { "false" } else { "true" };
|
||||
if key_str == "install_service"
|
||||
&& toggled == "true"
|
||||
&& !super::service::is_available()
|
||||
{
|
||||
self.message = Some((
|
||||
"requires root with systemd, use: sudo aether-proxy setup".into(),
|
||||
Instant::now(),
|
||||
true,
|
||||
));
|
||||
} else {
|
||||
self.selected_field_mut().value = toggled.into();
|
||||
self.modified = true;
|
||||
}
|
||||
}
|
||||
FieldKind::LogLevel => {
|
||||
const LEVELS: &[&str] = &["trace", "debug", "info", "warn", "error"];
|
||||
let idx = LEVELS.iter().position(|l| *l == value).unwrap_or(2);
|
||||
self.selected_field_mut().value = LEVELS[(idx + 1) % LEVELS.len()].into();
|
||||
self.modified = true;
|
||||
}
|
||||
_ => {
|
||||
self.edit_buffer = value;
|
||||
self.edit_cursor = self.edit_buffer.chars().count();
|
||||
self.mode = Mode::Editing;
|
||||
}
|
||||
}
|
||||
}
|
||||
// -- Tab navigation --
|
||||
KeyCode::Tab => {
|
||||
if self.server_tabs.len() > 1 {
|
||||
self.active_tab = (self.active_tab + 1) % self.server_tabs.len();
|
||||
self.clamp_selection();
|
||||
}
|
||||
}
|
||||
KeyCode::BackTab => {
|
||||
if self.server_tabs.len() > 1 {
|
||||
self.active_tab = if self.active_tab == 0 {
|
||||
self.server_tabs.len() - 1
|
||||
} else {
|
||||
self.active_tab - 1
|
||||
};
|
||||
self.clamp_selection();
|
||||
}
|
||||
}
|
||||
KeyCode::Char(c @ '1'..='9') if !key.modifiers.contains(KeyModifiers::CONTROL) => {
|
||||
let idx = (c as usize) - ('1' as usize);
|
||||
if idx < self.server_tabs.len() && idx != self.active_tab {
|
||||
self.active_tab = idx;
|
||||
self.clamp_selection();
|
||||
}
|
||||
}
|
||||
// -- Add / remove server --
|
||||
KeyCode::Char('+') | KeyCode::Char('a') => {
|
||||
self.server_tabs.push(ServerTab::new());
|
||||
self.active_tab = self.server_tabs.len() - 1;
|
||||
self.selected = 0;
|
||||
self.scroll_offset = 0;
|
||||
self.modified = true;
|
||||
self.message = Some((
|
||||
format!("added server {}", self.server_tabs.len()),
|
||||
Instant::now(),
|
||||
false,
|
||||
));
|
||||
}
|
||||
KeyCode::Delete | KeyCode::Char('x') => {
|
||||
if self.server_tabs.len() <= 1 {
|
||||
self.message =
|
||||
Some(("cannot remove the last server".into(), Instant::now(), true));
|
||||
} else if self.confirm_delete {
|
||||
let removed = self.active_tab + 1;
|
||||
self.server_tabs.remove(self.active_tab);
|
||||
self.active_tab = self.active_tab.min(self.server_tabs.len() - 1);
|
||||
self.clamp_selection();
|
||||
self.modified = true;
|
||||
self.message =
|
||||
Some((format!("server {} removed", removed), Instant::now(), false));
|
||||
} else {
|
||||
self.confirm_delete = true;
|
||||
self.message = Some((
|
||||
"press Delete/x again to remove this server".into(),
|
||||
Instant::now(),
|
||||
true,
|
||||
));
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
fn handle_edit(&mut self, key: KeyEvent) {
|
||||
match key.code {
|
||||
KeyCode::Esc => {
|
||||
self.mode = Mode::Normal;
|
||||
}
|
||||
KeyCode::Enter => {
|
||||
if self.validate_edit() {
|
||||
self.selected_field_mut().value = self.edit_buffer.clone();
|
||||
self.modified = true;
|
||||
self.mode = Mode::Normal;
|
||||
} else {
|
||||
self.message = Some(("invalid format".into(), Instant::now(), true));
|
||||
}
|
||||
}
|
||||
KeyCode::Backspace => {
|
||||
if self.edit_cursor > 0 {
|
||||
self.edit_cursor -= 1;
|
||||
let byte = self.char_byte_pos(self.edit_cursor);
|
||||
self.edit_buffer.remove(byte);
|
||||
}
|
||||
}
|
||||
KeyCode::Delete => {
|
||||
if self.edit_cursor < self.edit_buffer.chars().count() {
|
||||
let byte = self.char_byte_pos(self.edit_cursor);
|
||||
self.edit_buffer.remove(byte);
|
||||
}
|
||||
}
|
||||
KeyCode::Left => {
|
||||
self.edit_cursor = self.edit_cursor.saturating_sub(1);
|
||||
}
|
||||
KeyCode::Right => {
|
||||
let len = self.edit_buffer.chars().count();
|
||||
if self.edit_cursor < len {
|
||||
self.edit_cursor += 1;
|
||||
}
|
||||
}
|
||||
KeyCode::Home => self.edit_cursor = 0,
|
||||
KeyCode::End => self.edit_cursor = self.edit_buffer.chars().count(),
|
||||
KeyCode::Char(c) => {
|
||||
let byte = self.char_byte_pos(self.edit_cursor);
|
||||
self.edit_buffer.insert(byte, c);
|
||||
self.edit_cursor += 1;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_edit(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
/// Byte offset of the char at `char_idx`.
|
||||
fn char_byte_pos(&self, char_idx: usize) -> usize {
|
||||
self.edit_buffer
|
||||
.char_indices()
|
||||
.nth(char_idx)
|
||||
.map(|(i, _)| i)
|
||||
.unwrap_or(self.edit_buffer.len())
|
||||
}
|
||||
}
|
||||
// -- Rendering ----------------------------------------------------------------
|
||||
|
||||
fn ui(f: &mut Frame, app: &mut App) {
|
||||
let area = f.area();
|
||||
|
||||
let title = if app.modified {
|
||||
" Aether Proxy Setup [*] "
|
||||
} else {
|
||||
" Aether Proxy Setup "
|
||||
};
|
||||
|
||||
let outer = Block::default()
|
||||
.borders(Borders::ALL)
|
||||
.title(title)
|
||||
.title_alignment(ratatui::layout::Alignment::Center)
|
||||
.border_style(Style::default().fg(Color::Cyan));
|
||||
|
||||
let inner = outer.inner(area);
|
||||
f.render_widget(outer, area);
|
||||
|
||||
// Split: fields | tab bar | footer
|
||||
let chunks = Layout::vertical([
|
||||
Constraint::Min(1),
|
||||
Constraint::Length(1),
|
||||
Constraint::Length(4),
|
||||
])
|
||||
.split(inner);
|
||||
|
||||
render_fields(f, app, chunks[0]);
|
||||
render_tab_bar(f, app, chunks[1]);
|
||||
render_footer(f, app, chunks[2]);
|
||||
}
|
||||
|
||||
fn render_fields(f: &mut Frame, app: &mut App, area: Rect) {
|
||||
let visible = area.height as usize;
|
||||
app.ensure_visible(visible);
|
||||
|
||||
let server_count = app.server_field_count();
|
||||
let mut lines: Vec<Line> = Vec::new();
|
||||
// display_row tracks the actual row index (including separator)
|
||||
let mut display_row: usize = 0;
|
||||
|
||||
// Server fields
|
||||
for i in 0..server_count {
|
||||
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
|
||||
lines.push(build_field_line(app, i, display_row));
|
||||
}
|
||||
display_row += 1;
|
||||
}
|
||||
|
||||
// Separator line
|
||||
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
|
||||
lines.push(Line::from(Span::styled(
|
||||
" ----------------------------------------",
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)));
|
||||
}
|
||||
display_row += 1;
|
||||
|
||||
// Global fields
|
||||
for i in 0..app.global_fields.len() {
|
||||
let field_idx = server_count + i;
|
||||
if display_row >= app.scroll_offset && display_row < app.scroll_offset + visible {
|
||||
lines.push(build_field_line(app, field_idx, display_row));
|
||||
}
|
||||
display_row += 1;
|
||||
}
|
||||
|
||||
let paragraph = Paragraph::new(lines);
|
||||
f.render_widget(paragraph, area);
|
||||
|
||||
// Cursor position while editing
|
||||
if app.mode == Mode::Editing {
|
||||
let sel_display_row = if app.selected >= server_count {
|
||||
app.selected + 1
|
||||
} else {
|
||||
app.selected
|
||||
};
|
||||
let row_in_view = sel_display_row.saturating_sub(app.scroll_offset);
|
||||
let prefix: u16 = 3 + LABEL_WIDTH as u16 + 2;
|
||||
let cx = area.x + prefix + app.edit_cursor as u16;
|
||||
let cy = area.y + row_in_view as u16;
|
||||
if cx < area.x + area.width && cy < area.y + area.height {
|
||||
f.set_cursor_position((cx, cy));
|
||||
}
|
||||
}
|
||||
}
|
||||
fn build_field_line(app: &App, field_idx: usize, _display_row: usize) -> Line<'static> {
|
||||
let sc = app.server_field_count();
|
||||
let field = if field_idx < sc {
|
||||
&app.server_tabs[app.active_tab].fields[field_idx]
|
||||
} else {
|
||||
&app.global_fields[field_idx - sc]
|
||||
};
|
||||
|
||||
let selected = field_idx == app.selected;
|
||||
let indicator = if selected { " > " } else { " " };
|
||||
|
||||
let label_style = if selected {
|
||||
Style::default()
|
||||
.fg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD)
|
||||
} else {
|
||||
Style::default().fg(Color::DarkGray)
|
||||
};
|
||||
|
||||
let padded_label = format!("{:<width$}", field.label, width = LABEL_WIDTH);
|
||||
|
||||
let (value_text, value_style) = if app.mode == Mode::Editing && selected {
|
||||
(app.edit_buffer.clone(), Style::default().fg(Color::Yellow))
|
||||
} else {
|
||||
field_display(field)
|
||||
};
|
||||
|
||||
Line::from(vec![
|
||||
Span::styled(indicator.to_string(), label_style),
|
||||
Span::styled(padded_label, label_style),
|
||||
Span::raw(" "),
|
||||
Span::styled(value_text, value_style),
|
||||
])
|
||||
}
|
||||
|
||||
/// Returns (display_text, style) for a field in normal mode.
|
||||
fn field_display(field: &Field) -> (String, Style) {
|
||||
if field.value.is_empty() {
|
||||
let text = if field.required {
|
||||
"(required)".into()
|
||||
} else {
|
||||
"-".into()
|
||||
};
|
||||
let color = if field.required {
|
||||
Color::Red
|
||||
} else {
|
||||
Color::DarkGray
|
||||
};
|
||||
return (text, Style::default().fg(color));
|
||||
}
|
||||
|
||||
match field.kind {
|
||||
FieldKind::Secret => (
|
||||
"*".repeat(field.value.len().min(20)),
|
||||
Style::default().fg(Color::White),
|
||||
),
|
||||
FieldKind::Bool => {
|
||||
if field.value == "true" {
|
||||
("[x] on".into(), Style::default().fg(Color::Green))
|
||||
} else {
|
||||
("[ ] off".into(), Style::default().fg(Color::DarkGray))
|
||||
}
|
||||
}
|
||||
FieldKind::LogLevel => {
|
||||
let color = match field.value.as_str() {
|
||||
"trace" => Color::Magenta,
|
||||
"debug" => Color::Blue,
|
||||
"info" => Color::Green,
|
||||
"warn" => Color::Yellow,
|
||||
"error" => Color::Red,
|
||||
_ => Color::White,
|
||||
};
|
||||
(field.value.clone(), Style::default().fg(color))
|
||||
}
|
||||
_ => (field.value.clone(), Style::default().fg(Color::White)),
|
||||
}
|
||||
}
|
||||
fn render_tab_bar(f: &mut Frame, app: &App, area: Rect) {
|
||||
let mut spans: Vec<Span> = Vec::new();
|
||||
spans.push(Span::raw(" "));
|
||||
|
||||
for (i, tab) in app.server_tabs.iter().enumerate() {
|
||||
let num = i + 1;
|
||||
let name = tab
|
||||
.fields
|
||||
.iter()
|
||||
.find(|f| f.key == "node_name")
|
||||
.filter(|f| !f.value.is_empty())
|
||||
.map(|f| f.value.clone())
|
||||
.unwrap_or_else(|| format!("Server {}", num));
|
||||
|
||||
let label = format!(" {} {} ", num, name);
|
||||
|
||||
if i == app.active_tab {
|
||||
spans.push(Span::styled(
|
||||
label,
|
||||
Style::default()
|
||||
.fg(Color::Black)
|
||||
.bg(Color::Cyan)
|
||||
.add_modifier(Modifier::BOLD),
|
||||
));
|
||||
} else {
|
||||
spans.push(Span::styled(label, Style::default().fg(Color::DarkGray)));
|
||||
}
|
||||
spans.push(Span::raw(" "));
|
||||
}
|
||||
|
||||
spans.push(Span::styled(" + Add ", Style::default().fg(Color::Green)));
|
||||
|
||||
f.render_widget(Paragraph::new(Line::from(spans)), area);
|
||||
}
|
||||
|
||||
fn render_footer(f: &mut Frame, app: &App, area: Rect) {
|
||||
let help = app.selected_field().help;
|
||||
|
||||
let keybindings = if app.mode == Mode::Editing {
|
||||
"Enter confirm Esc cancel"
|
||||
} else if app.server_tabs.len() > 1 {
|
||||
"j/k select Enter edit Tab switch + add x remove ^S save q quit"
|
||||
} else {
|
||||
"j/k select Enter edit + add server ^S save q quit"
|
||||
};
|
||||
|
||||
let mut status_spans: Vec<Span> = vec![Span::styled(
|
||||
format!(" {}", keybindings),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)];
|
||||
|
||||
if let Some((msg, _, is_err)) = &app.message {
|
||||
let color = if *is_err { Color::Red } else { Color::Green };
|
||||
status_spans.push(Span::raw(" "));
|
||||
status_spans.push(Span::styled(msg.clone(), Style::default().fg(color)));
|
||||
}
|
||||
|
||||
let footer_text = vec![
|
||||
Line::raw(""),
|
||||
Line::from(Span::styled(
|
||||
format!(" {}", help),
|
||||
Style::default().fg(Color::DarkGray),
|
||||
)),
|
||||
Line::from(status_spans),
|
||||
];
|
||||
|
||||
let footer = Paragraph::new(footer_text).block(
|
||||
Block::default()
|
||||
.borders(Borders::TOP)
|
||||
.border_style(Style::default().fg(Color::DarkGray)),
|
||||
);
|
||||
|
||||
f.render_widget(footer, area);
|
||||
}
|
||||
// -- Entry point --------------------------------------------------------------
|
||||
|
||||
pub fn run(config_path: PathBuf) -> anyhow::Result<SetupOutcome> {
|
||||
terminal::enable_raw_mode()?;
|
||||
let mut stdout = io::stdout();
|
||||
execute!(stdout, EnterAlternateScreen)?;
|
||||
let backend = CrosstermBackend::new(stdout);
|
||||
let mut terminal = Terminal::new(backend)?;
|
||||
|
||||
let mut app = App::new(config_path.clone());
|
||||
app.load_from_file();
|
||||
|
||||
let result = event_loop(&mut terminal, &mut app);
|
||||
|
||||
terminal::disable_raw_mode()?;
|
||||
execute!(terminal.backend_mut(), LeaveAlternateScreen)?;
|
||||
terminal.show_cursor()?;
|
||||
|
||||
result?;
|
||||
|
||||
// -- Post-TUI: decide outcome ---------------------------------------------
|
||||
|
||||
if !app.saved_once {
|
||||
return Ok(SetupOutcome::Cancelled);
|
||||
}
|
||||
|
||||
eprintln!();
|
||||
eprintln!(" Config saved to {}", config_path.display());
|
||||
eprintln!();
|
||||
|
||||
let wants_service = app
|
||||
.global_fields
|
||||
.iter()
|
||||
.find(|f| f.key == "install_service")
|
||||
.map(|f| f.value == "true")
|
||||
.unwrap_or(false);
|
||||
|
||||
if wants_service {
|
||||
match super::service::install_service(&config_path) {
|
||||
Ok(()) => return Ok(SetupOutcome::ServiceInstalled),
|
||||
Err(e) => {
|
||||
eprintln!(" Service install failed: {}", e);
|
||||
eprintln!(" Starting proxy directly instead.\n");
|
||||
}
|
||||
}
|
||||
} else if super::service::is_installed() {
|
||||
if let Err(e) = super::service::uninstall_service() {
|
||||
eprintln!(" Service uninstall failed: {}", e);
|
||||
eprintln!();
|
||||
}
|
||||
}
|
||||
|
||||
Ok(SetupOutcome::ReadyToRun(config_path))
|
||||
}
|
||||
|
||||
fn event_loop(
|
||||
terminal: &mut Terminal<CrosstermBackend<io::Stdout>>,
|
||||
app: &mut App,
|
||||
) -> anyhow::Result<()> {
|
||||
loop {
|
||||
terminal.draw(|f| ui(f, app))?;
|
||||
|
||||
if event::poll(Duration::from_millis(200))? {
|
||||
if let Event::Key(key) = event::read()? {
|
||||
if key.kind == KeyEventKind::Press && app.handle_key(key) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,411 @@
|
||||
//! Self-upgrade for aether-proxy.
|
||||
//!
|
||||
//! Downloads a release from GitHub, verifies SHA256 checksum, and atomically
|
||||
//! replaces the running binary. Restarts the systemd service if active.
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
const GITHUB_API_BASE: &str = "https://api.github.com";
|
||||
const GITHUB_REPO: &str = "fawney19/Aether";
|
||||
const CURRENT_VERSION: &str = env!("CARGO_PKG_VERSION");
|
||||
|
||||
// ── GitHub API types ─────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct GithubRelease {
|
||||
tag_name: String,
|
||||
name: String,
|
||||
}
|
||||
|
||||
// ── Platform detection ───────────────────────────────────────────────────────
|
||||
|
||||
fn detect_platform() -> &'static str {
|
||||
if cfg!(target_os = "linux") && cfg!(target_arch = "x86_64") {
|
||||
"linux-amd64"
|
||||
} else if cfg!(target_os = "linux") && cfg!(target_arch = "aarch64") {
|
||||
"linux-arm64"
|
||||
} else if cfg!(target_os = "macos") && cfg!(target_arch = "x86_64") {
|
||||
"macos-amd64"
|
||||
} else if cfg!(target_os = "macos") && cfg!(target_arch = "aarch64") {
|
||||
"macos-arm64"
|
||||
} else if cfg!(target_os = "windows") && cfg!(target_arch = "x86_64") {
|
||||
"windows-amd64"
|
||||
} else {
|
||||
// All supported targets are covered above; this is unreachable for
|
||||
// any platform we actually build for.
|
||||
panic!("unsupported platform: compile-time target not in the supported matrix")
|
||||
}
|
||||
}
|
||||
|
||||
// ── GitHub HTTP client ───────────────────────────────────────────────────────
|
||||
|
||||
fn build_github_client() -> anyhow::Result<reqwest::Client> {
|
||||
let mut headers = reqwest::header::HeaderMap::new();
|
||||
|
||||
if let Ok(token) = std::env::var("GITHUB_TOKEN") {
|
||||
headers.insert(
|
||||
reqwest::header::AUTHORIZATION,
|
||||
reqwest::header::HeaderValue::from_str(&format!("Bearer {}", token))?,
|
||||
);
|
||||
}
|
||||
|
||||
headers.insert(
|
||||
reqwest::header::ACCEPT,
|
||||
reqwest::header::HeaderValue::from_static("application/vnd.github+json"),
|
||||
);
|
||||
|
||||
Ok(reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(300))
|
||||
.user_agent(format!("aether-proxy/{}", CURRENT_VERSION))
|
||||
.default_headers(headers)
|
||||
.build()?)
|
||||
}
|
||||
|
||||
// ── Release fetching ─────────────────────────────────────────────────────────
|
||||
|
||||
async fn fetch_release(
|
||||
client: &reqwest::Client,
|
||||
version: Option<&str>,
|
||||
) -> anyhow::Result<GithubRelease> {
|
||||
match version {
|
||||
Some(ver) => {
|
||||
// Accept both "proxy-v0.2.0" and bare "0.2.0"
|
||||
let tag = if ver.starts_with("proxy-v") {
|
||||
ver.to_string()
|
||||
} else {
|
||||
format!("proxy-v{}", ver)
|
||||
};
|
||||
let url = format!(
|
||||
"{}/repos/{}/releases/tags/{}",
|
||||
GITHUB_API_BASE, GITHUB_REPO, tag
|
||||
);
|
||||
let resp = client.get(&url).send().await?;
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
anyhow::bail!("release '{}' not found (HTTP {}): {}", tag, status, body);
|
||||
}
|
||||
Ok(resp.json().await?)
|
||||
}
|
||||
None => {
|
||||
// List releases and find the latest proxy-v* tag
|
||||
let url = format!(
|
||||
"{}/repos/{}/releases?per_page=20",
|
||||
GITHUB_API_BASE, GITHUB_REPO
|
||||
);
|
||||
let resp = client.get(&url).send().await?;
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
anyhow::bail!("failed to list releases (HTTP {}): {}", status, body);
|
||||
}
|
||||
let releases: Vec<GithubRelease> = resp.json().await?;
|
||||
releases
|
||||
.into_iter()
|
||||
.find(|r| r.tag_name.starts_with("proxy-v"))
|
||||
.ok_or_else(|| anyhow::anyhow!("no proxy-v* release found"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Download via GitHub release direct links ─────────────────────────────────
|
||||
|
||||
/// Download a release asset via the public direct download URL:
|
||||
/// `https://github.com/{repo}/releases/download/{tag}/{filename}`
|
||||
async fn download_release_file(
|
||||
client: &reqwest::Client,
|
||||
tag: &str,
|
||||
filename: &str,
|
||||
) -> anyhow::Result<Vec<u8>> {
|
||||
let url = format!(
|
||||
"https://github.com/{}/releases/download/{}/{}",
|
||||
GITHUB_REPO, tag, filename
|
||||
);
|
||||
let resp = client
|
||||
.get(&url)
|
||||
.header(reqwest::header::ACCEPT, "application/octet-stream")
|
||||
.send()
|
||||
.await?;
|
||||
if !resp.status().is_success() {
|
||||
anyhow::bail!(
|
||||
"download failed for '{}' (HTTP {})",
|
||||
filename,
|
||||
resp.status(),
|
||||
);
|
||||
}
|
||||
Ok(resp.bytes().await?.to_vec())
|
||||
}
|
||||
|
||||
fn parse_checksum(sums_text: &str, filename: &str) -> anyhow::Result<String> {
|
||||
for line in sums_text.lines() {
|
||||
// Format: "<hash> <filename>" (GNU coreutils convention)
|
||||
let mut parts = line.split_ascii_whitespace();
|
||||
let (Some(hash), Some(name)) = (parts.next(), parts.next()) else {
|
||||
continue;
|
||||
};
|
||||
if name == filename || name.ends_with(filename) {
|
||||
return Ok(hash.to_lowercase());
|
||||
}
|
||||
}
|
||||
anyhow::bail!("checksum for '{}' not found in SHA256SUMS.txt", filename);
|
||||
}
|
||||
|
||||
async fn download_and_verify(
|
||||
client: &reqwest::Client,
|
||||
tag: &str,
|
||||
platform: &str,
|
||||
dest: &Path,
|
||||
) -> anyhow::Result<()> {
|
||||
let archive_name = format!("aether-proxy-{}.tar.gz", platform);
|
||||
|
||||
eprintln!(" Downloading {}...", archive_name);
|
||||
let (archive_bytes, checksum_bytes) = tokio::try_join!(
|
||||
download_release_file(client, tag, &archive_name),
|
||||
download_release_file(client, tag, "SHA256SUMS.txt"),
|
||||
)?;
|
||||
let checksum_text = String::from_utf8(checksum_bytes)?;
|
||||
|
||||
eprintln!(
|
||||
" Downloaded {} ({} bytes)",
|
||||
archive_name,
|
||||
archive_bytes.len()
|
||||
);
|
||||
|
||||
// Verify SHA256
|
||||
let expected_hash = parse_checksum(&checksum_text, &archive_name)?;
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(&archive_bytes);
|
||||
let actual_hash = hex::encode(hasher.finalize());
|
||||
|
||||
if actual_hash != expected_hash {
|
||||
anyhow::bail!(
|
||||
"SHA256 mismatch for {}:\n expected: {}\n actual: {}",
|
||||
archive_name,
|
||||
expected_hash,
|
||||
actual_hash
|
||||
);
|
||||
}
|
||||
eprintln!(" SHA256 verified: {}", &actual_hash[..16]);
|
||||
|
||||
extract_binary(&archive_bytes, dest)?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── Archive extraction ───────────────────────────────────────────────────────
|
||||
|
||||
fn extract_binary(archive_bytes: &[u8], dest: &Path) -> anyhow::Result<()> {
|
||||
use flate2::read::GzDecoder;
|
||||
use tar::Archive;
|
||||
|
||||
// Guard against decompression bombs
|
||||
const MAX_BINARY_SIZE: u64 = 100 * 1024 * 1024; // 100 MB
|
||||
|
||||
let decoder = GzDecoder::new(archive_bytes);
|
||||
let mut archive = Archive::new(decoder);
|
||||
|
||||
let binary_name = if cfg!(target_os = "windows") {
|
||||
"aether-proxy.exe"
|
||||
} else {
|
||||
"aether-proxy"
|
||||
};
|
||||
|
||||
for entry in archive.entries()? {
|
||||
let mut entry = entry?;
|
||||
// Only accept regular files -- reject symlinks to prevent write-through attacks
|
||||
if entry.header().entry_type() != tar::EntryType::Regular {
|
||||
continue;
|
||||
}
|
||||
let path = entry.path()?;
|
||||
if path.file_name().and_then(|n| n.to_str()) == Some(binary_name) {
|
||||
let size = entry.header().size()?;
|
||||
if size > MAX_BINARY_SIZE {
|
||||
anyhow::bail!(
|
||||
"binary too large ({} bytes, max {} bytes)",
|
||||
size,
|
||||
MAX_BINARY_SIZE
|
||||
);
|
||||
}
|
||||
let mut file = std::fs::File::create(dest)?;
|
||||
std::io::copy(&mut entry, &mut file)?;
|
||||
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
std::fs::set_permissions(dest, std::fs::Permissions::from_mode(0o755))?;
|
||||
}
|
||||
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
anyhow::bail!("'{}' not found in archive", binary_name);
|
||||
}
|
||||
|
||||
// ── Atomic binary replacement ────────────────────────────────────────────────
|
||||
|
||||
fn atomic_replace(new_binary: &Path) -> anyhow::Result<PathBuf> {
|
||||
let current_exe = std::env::current_exe()?.canonicalize()?;
|
||||
let backup_path = current_exe.with_extension("bak");
|
||||
|
||||
// Remove stale backup
|
||||
let _ = std::fs::remove_file(&backup_path);
|
||||
|
||||
// current -> .bak
|
||||
std::fs::rename(¤t_exe, &backup_path).map_err(|e| {
|
||||
anyhow::anyhow!(
|
||||
"failed to backup current binary '{}' -> '{}': {}",
|
||||
current_exe.display(),
|
||||
backup_path.display(),
|
||||
e
|
||||
)
|
||||
})?;
|
||||
|
||||
// new -> current
|
||||
if let Err(e) = std::fs::rename(new_binary, ¤t_exe) {
|
||||
eprintln!(" ERROR: failed to place new binary, rolling back...");
|
||||
let _ = std::fs::rename(&backup_path, ¤t_exe);
|
||||
anyhow::bail!(
|
||||
"failed to install new binary '{}' -> '{}': {}",
|
||||
new_binary.display(),
|
||||
current_exe.display(),
|
||||
e
|
||||
);
|
||||
}
|
||||
|
||||
eprintln!(" Binary replaced: {}", current_exe.display());
|
||||
Ok(backup_path)
|
||||
}
|
||||
|
||||
// ── Public entry point ───────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum RestartMode {
|
||||
BestEffort,
|
||||
Required,
|
||||
}
|
||||
|
||||
async fn execute_upgrade(
|
||||
version: Option<&str>,
|
||||
require_root: bool,
|
||||
restart_mode: RestartMode,
|
||||
) -> anyhow::Result<()> {
|
||||
// Resolve exe path once; reuse throughout the function
|
||||
let current_exe = std::env::current_exe()?.canonicalize()?;
|
||||
let exe_dir = current_exe
|
||||
.parent()
|
||||
.ok_or_else(|| anyhow::anyhow!("cannot determine binary directory"))?;
|
||||
let temp_path = exe_dir.join(".aether-proxy.upgrade.tmp");
|
||||
|
||||
if require_root {
|
||||
if !super::service::is_root() {
|
||||
anyhow::bail!("automatic upgrade requires root privileges");
|
||||
}
|
||||
} else if !super::service::is_root() {
|
||||
// Check write permission to binary directory for manual upgrade mode.
|
||||
let test_path = exe_dir.join(".aether-proxy.write-test");
|
||||
match std::fs::File::create(&test_path) {
|
||||
Ok(_) => {
|
||||
let _ = std::fs::remove_file(&test_path);
|
||||
}
|
||||
Err(_) => {
|
||||
anyhow::bail!(
|
||||
"no write access to {}. Use: sudo aether-proxy upgrade",
|
||||
exe_dir.display()
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let platform = detect_platform();
|
||||
eprintln!(" Platform: {}", platform);
|
||||
eprintln!(" Current version: {}", CURRENT_VERSION);
|
||||
|
||||
let client = build_github_client()?;
|
||||
let release = fetch_release(&client, version).await?;
|
||||
let target_tag = &release.tag_name;
|
||||
let target_semver = target_tag.strip_prefix("proxy-v").unwrap_or(target_tag);
|
||||
|
||||
eprintln!(" Target version: {} ({})", target_tag, release.name);
|
||||
|
||||
if target_semver == CURRENT_VERSION {
|
||||
eprintln!(
|
||||
" Already running version {}, nothing to do.",
|
||||
CURRENT_VERSION
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
eprintln!();
|
||||
eprintln!(" Upgrading: {} -> {}", CURRENT_VERSION, target_semver);
|
||||
eprintln!();
|
||||
|
||||
if let Err(e) = download_and_verify(&client, target_tag, platform, &temp_path).await {
|
||||
let _ = std::fs::remove_file(&temp_path);
|
||||
return Err(e);
|
||||
}
|
||||
let backup_path = match atomic_replace(&temp_path) {
|
||||
Ok(backup) => backup,
|
||||
Err(e) => {
|
||||
let _ = std::fs::remove_file(&temp_path);
|
||||
return Err(e);
|
||||
}
|
||||
};
|
||||
|
||||
match restart_mode {
|
||||
RestartMode::BestEffort => {
|
||||
// Restart systemd service if running.
|
||||
// Use best-effort: binary is already replaced, so a restart failure should
|
||||
// not abort the whole upgrade -- the user can restart manually.
|
||||
if super::service::is_service_active() {
|
||||
if super::service::is_root() {
|
||||
eprintln!(" Restarting systemd service...");
|
||||
match super::service::run_cmd("systemctl", &["restart", "aether-proxy"]) {
|
||||
Ok(()) => eprintln!(" Service restarted."),
|
||||
Err(e) => {
|
||||
eprintln!(" WARNING: failed to restart service: {}", e);
|
||||
eprintln!(" Run manually: sudo systemctl restart aether-proxy");
|
||||
}
|
||||
}
|
||||
} else {
|
||||
eprintln!(" Systemd service is active, but restart requires root.");
|
||||
eprintln!(" Run: sudo systemctl restart aether-proxy");
|
||||
eprintln!(" Skipping restart.");
|
||||
}
|
||||
} else {
|
||||
eprintln!(" No active systemd service detected, skipping restart.");
|
||||
}
|
||||
}
|
||||
RestartMode::Required => {
|
||||
if !super::service::is_root() {
|
||||
anyhow::bail!("automatic upgrade requires root privileges");
|
||||
}
|
||||
eprintln!(" Restarting systemd service...");
|
||||
super::service::run_cmd("systemctl", &["restart", "aether-proxy"])?;
|
||||
eprintln!(" Service restarted.");
|
||||
}
|
||||
}
|
||||
|
||||
eprintln!();
|
||||
eprintln!(" Upgrade complete!");
|
||||
eprintln!(
|
||||
" Backup kept at: {} (will be cleaned up on next upgrade)",
|
||||
backup_path.display()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// `aether-proxy upgrade [version]` -- self-upgrade from GitHub releases.
|
||||
pub async fn cmd_upgrade(version: Option<String>) -> anyhow::Result<()> {
|
||||
execute_upgrade(version.as_deref(), false, RestartMode::BestEffort).await
|
||||
}
|
||||
|
||||
/// Perform automatic upgrade to a specific version.
|
||||
///
|
||||
/// This path is designed for server-pushed upgrades in systemd/root scenarios:
|
||||
/// it requires root and requires a successful `systemctl restart aether-proxy`.
|
||||
pub async fn perform_upgrade(version: &str) -> anyhow::Result<()> {
|
||||
execute_upgrade(Some(version), true, RestartMode::Required).await
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
//! Shared application state passed to all subsystems.
|
||||
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, RwLock};
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::config::Config;
|
||||
use crate::registration::client::AetherClient;
|
||||
use crate::runtime::SharedDynamicConfig;
|
||||
use crate::target_filter::DnsCache;
|
||||
use crate::upstream_client::UpstreamClient;
|
||||
|
||||
/// Central application state shared across all servers/tunnels.
|
||||
pub struct AppState {
|
||||
pub config: Arc<Config>,
|
||||
/// DNS cache for upstream target resolution (shared).
|
||||
pub dns_cache: Arc<DnsCache>,
|
||||
/// Hyper client for tunnel upstream requests with validated DNS and connection timing.
|
||||
pub upstream_client: UpstreamClient,
|
||||
/// Shared TLS config for tunnel WebSocket connections (avoids re-parsing root CAs on each reconnect).
|
||||
pub tunnel_tls_config: Arc<rustls::ClientConfig>,
|
||||
}
|
||||
|
||||
/// Per-server state: one instance per Aether server connection.
|
||||
pub struct ServerContext {
|
||||
/// Human-readable label for logging (e.g. "server-0").
|
||||
pub server_label: String,
|
||||
/// Aether server URL for this connection.
|
||||
pub aether_url: String,
|
||||
/// Management token for this server.
|
||||
pub management_token: String,
|
||||
/// Resolved node name at registration time (per-server override or global fallback).
|
||||
/// After startup, the active node_name is read from `dynamic` (may be updated remotely).
|
||||
#[allow(dead_code)]
|
||||
pub node_name: String,
|
||||
/// Node ID assigned by this Aether server.
|
||||
pub node_id: Arc<RwLock<String>>,
|
||||
/// API client for this server.
|
||||
pub aether_client: Arc<AetherClient>,
|
||||
/// Dynamic config from this server's heartbeat ACKs.
|
||||
pub dynamic: SharedDynamicConfig,
|
||||
/// Per-server active connection count.
|
||||
pub active_connections: Arc<AtomicU64>,
|
||||
/// Per-server request/latency metrics.
|
||||
pub metrics: Arc<ProxyMetrics>,
|
||||
}
|
||||
|
||||
/// Aggregate metrics for reporting to Aether.
|
||||
pub struct ProxyMetrics {
|
||||
pub total_requests: AtomicU64,
|
||||
/// Cumulative connection-establishment latency in nanoseconds
|
||||
/// (DNS + TCP/TLS + TTFB, excludes response body streaming).
|
||||
pub total_latency_ns: AtomicU64,
|
||||
pub failed_requests: AtomicU64,
|
||||
pub dns_failures: AtomicU64,
|
||||
pub stream_errors: AtomicU64,
|
||||
}
|
||||
|
||||
impl ProxyMetrics {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
total_requests: AtomicU64::new(0),
|
||||
total_latency_ns: AtomicU64::new(0),
|
||||
failed_requests: AtomicU64::new(0),
|
||||
dns_failures: AtomicU64::new(0),
|
||||
stream_errors: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Record a completed request with its connection-establishment latency
|
||||
/// (DNS + TCP/TLS + TTFB, excludes response body streaming).
|
||||
pub fn record_request(&self, connect_elapsed: Duration) {
|
||||
let nanos = u64::try_from(connect_elapsed.as_nanos()).unwrap_or(u64::MAX);
|
||||
self.total_requests.fetch_add(1, Ordering::Release);
|
||||
self.total_latency_ns.fetch_add(nanos, Ordering::Release);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,381 @@
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
/// Check if an IP address belongs to a private/reserved network.
|
||||
pub fn is_private_ip(ip: &IpAddr) -> bool {
|
||||
match ip {
|
||||
IpAddr::V4(v4) => is_private_ipv4(v4),
|
||||
IpAddr::V6(v6) => is_private_ipv6(v6),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_private_ipv4(ip: &Ipv4Addr) -> bool {
|
||||
let octets = ip.octets();
|
||||
// 10.0.0.0/8
|
||||
if octets[0] == 10 {
|
||||
return true;
|
||||
}
|
||||
// 172.16.0.0/12
|
||||
if octets[0] == 172 && (16..=31).contains(&octets[1]) {
|
||||
return true;
|
||||
}
|
||||
// 192.168.0.0/16
|
||||
if octets[0] == 192 && octets[1] == 168 {
|
||||
return true;
|
||||
}
|
||||
// 127.0.0.0/8
|
||||
if octets[0] == 127 {
|
||||
return true;
|
||||
}
|
||||
// 169.254.0.0/16 (link-local)
|
||||
if octets[0] == 169 && octets[1] == 254 {
|
||||
return true;
|
||||
}
|
||||
// 0.0.0.0/8
|
||||
if octets[0] == 0 {
|
||||
return true;
|
||||
}
|
||||
// 100.64.0.0/10 (CGNAT / shared address space)
|
||||
if octets[0] == 100 && (64..=127).contains(&octets[1]) {
|
||||
return true;
|
||||
}
|
||||
// 192.0.0.0/24 (IETF protocol assignments)
|
||||
if octets[0] == 192 && octets[1] == 0 && octets[2] == 0 {
|
||||
return true;
|
||||
}
|
||||
// 198.18.0.0/15 (benchmark testing)
|
||||
if octets[0] == 198 && (18..=19).contains(&octets[1]) {
|
||||
return true;
|
||||
}
|
||||
// 240.0.0.0/4 (reserved for future use)
|
||||
if octets[0] >= 240 {
|
||||
return true;
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
fn is_private_ipv6(ip: &Ipv6Addr) -> bool {
|
||||
// ::1 loopback
|
||||
if ip.is_loopback() {
|
||||
return true;
|
||||
}
|
||||
// :: unspecified
|
||||
if ip.is_unspecified() {
|
||||
return true;
|
||||
}
|
||||
let segments = ip.segments();
|
||||
// fc00::/7 (ULA) - first byte is 0xfc or 0xfd
|
||||
if segments[0] & 0xfe00 == 0xfc00 {
|
||||
return true;
|
||||
}
|
||||
// fe80::/10 (link-local)
|
||||
if segments[0] & 0xffc0 == 0xfe80 {
|
||||
return true;
|
||||
}
|
||||
// IPv4-mapped IPv6 (::ffff:x.x.x.x) - check the embedded IPv4
|
||||
if let Some(v4) = ip.to_ipv4_mapped() {
|
||||
return is_private_ipv4(&v4);
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum FilterError {
|
||||
PrivateIp(IpAddr),
|
||||
PortNotAllowed(u16),
|
||||
DnsResolutionFailed(String),
|
||||
NoPublicAddrs(String),
|
||||
}
|
||||
|
||||
impl std::fmt::Display for FilterError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::PrivateIp(ip) => write!(f, "target IP {} is in private/reserved range", ip),
|
||||
Self::PortNotAllowed(port) => write!(f, "port {} not in allowed list", port),
|
||||
Self::DnsResolutionFailed(host) => write!(f, "DNS resolution failed for {}", host),
|
||||
Self::NoPublicAddrs(host) => {
|
||||
write!(
|
||||
f,
|
||||
"all resolved addresses for {} are private/reserved",
|
||||
host
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct DnsCacheEntry {
|
||||
addrs: Arc<Vec<SocketAddr>>,
|
||||
expires_at: Instant,
|
||||
inserted_at: Instant,
|
||||
}
|
||||
|
||||
/// Lightweight DNS cache with TTL + capacity bounds.
|
||||
/// Stores all public resolved addresses per host (used by SafeDnsResolver
|
||||
/// to ensure reqwest connects to the same validated addresses).
|
||||
pub struct DnsCache {
|
||||
ttl: Duration,
|
||||
capacity: usize,
|
||||
entries: RwLock<HashMap<String, DnsCacheEntry>>,
|
||||
}
|
||||
|
||||
impl DnsCache {
|
||||
pub fn new(ttl: Duration, capacity: usize) -> Self {
|
||||
Self {
|
||||
ttl,
|
||||
capacity,
|
||||
entries: RwLock::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Look up cached public addresses for a host (any port).
|
||||
///
|
||||
/// Used by `SafeDnsResolver` which only knows the hostname — returns the
|
||||
/// first unexpired entry whose key starts with `host:`.
|
||||
pub async fn get_by_host(&self, host: &str) -> Option<Arc<Vec<SocketAddr>>> {
|
||||
if self.capacity == 0 || self.ttl.is_zero() {
|
||||
return None;
|
||||
}
|
||||
let prefix = format!("{}:", host.to_ascii_lowercase());
|
||||
let now = Instant::now();
|
||||
let entries = self.entries.read().await;
|
||||
for (key, entry) in entries.iter() {
|
||||
if key.starts_with(&prefix) && entry.expires_at > now {
|
||||
return Some(Arc::clone(&entry.addrs));
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
/// Look up cached public addresses for a host + port.
|
||||
pub async fn get(&self, host: &str, port: u16) -> Option<Arc<Vec<SocketAddr>>> {
|
||||
if self.capacity == 0 || self.ttl.is_zero() {
|
||||
return None;
|
||||
}
|
||||
let key = Self::key(host, port);
|
||||
let now = Instant::now();
|
||||
|
||||
// Fast path: read lock for cache hit
|
||||
{
|
||||
let entries = self.entries.read().await;
|
||||
match entries.get(&key) {
|
||||
Some(entry) if entry.expires_at > now => return Some(Arc::clone(&entry.addrs)),
|
||||
None => return None,
|
||||
Some(_) => {} // expired, fall through to evict
|
||||
}
|
||||
}
|
||||
|
||||
// Slow path: write lock to remove expired entry
|
||||
let mut entries = self.entries.write().await;
|
||||
entries.remove(&key);
|
||||
None
|
||||
}
|
||||
|
||||
/// Insert resolved public addresses into cache.
|
||||
pub async fn insert(&self, host: &str, port: u16, addrs: Arc<Vec<SocketAddr>>) {
|
||||
if self.capacity == 0 || self.ttl.is_zero() || addrs.is_empty() {
|
||||
return;
|
||||
}
|
||||
let key = Self::key(host, port);
|
||||
let now = Instant::now();
|
||||
let mut entries = self.entries.write().await;
|
||||
entries.retain(|_, entry| entry.expires_at > now);
|
||||
while entries.len() >= self.capacity {
|
||||
let oldest_key = entries
|
||||
.iter()
|
||||
.min_by_key(|(_, entry)| entry.inserted_at)
|
||||
.map(|(key, _)| key.clone());
|
||||
if let Some(key) = oldest_key {
|
||||
entries.remove(&key);
|
||||
} else {
|
||||
break;
|
||||
}
|
||||
}
|
||||
entries.insert(
|
||||
key,
|
||||
DnsCacheEntry {
|
||||
addrs,
|
||||
expires_at: now + self.ttl,
|
||||
inserted_at: now,
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
fn key(host: &str, port: u16) -> String {
|
||||
format!("{}:{}", host.to_ascii_lowercase(), port)
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve a hostname to public (non-private) socket addresses.
|
||||
///
|
||||
/// Results are cached in `dns_cache`. Private/reserved IPs are filtered out.
|
||||
/// Returns an error if no public addresses remain after filtering.
|
||||
pub async fn resolve_public_addrs(
|
||||
host: &str,
|
||||
port: u16,
|
||||
dns_cache: &DnsCache,
|
||||
) -> Result<Vec<SocketAddr>, FilterError> {
|
||||
// Cache hit
|
||||
if let Some(addrs) = dns_cache.get(host, port).await {
|
||||
return Ok((*addrs).clone());
|
||||
}
|
||||
|
||||
// Async DNS resolution
|
||||
let addr_str = format!("{}:{}", host, port);
|
||||
let resolved: Vec<SocketAddr> = tokio::net::lookup_host(&addr_str)
|
||||
.await
|
||||
.map_err(|_| FilterError::DnsResolutionFailed(host.to_string()))?
|
||||
.collect();
|
||||
|
||||
if resolved.is_empty() {
|
||||
return Err(FilterError::DnsResolutionFailed(host.to_string()));
|
||||
}
|
||||
|
||||
// Filter out private/reserved addresses
|
||||
let public: Vec<SocketAddr> = resolved
|
||||
.into_iter()
|
||||
.filter(|addr| !is_private_ip(&addr.ip()))
|
||||
.collect();
|
||||
|
||||
if public.is_empty() {
|
||||
return Err(FilterError::NoPublicAddrs(host.to_string()));
|
||||
}
|
||||
|
||||
// Cache the validated public addresses
|
||||
let arc_addrs = Arc::new(public);
|
||||
dns_cache.insert(host, port, Arc::clone(&arc_addrs)).await;
|
||||
Ok((*arc_addrs).clone())
|
||||
}
|
||||
|
||||
/// Validate that the target host:port is allowed.
|
||||
///
|
||||
/// Performs port whitelist check, private IP filtering, and DNS resolution
|
||||
/// with caching. The resolved addresses are stored in the shared DnsCache
|
||||
/// so that the SafeDnsResolver can reuse them, eliminating the TOCTTOU gap.
|
||||
pub async fn validate_target(
|
||||
host: &str,
|
||||
port: u16,
|
||||
allowed_ports: &HashSet<u16>,
|
||||
dns_cache: &DnsCache,
|
||||
) -> Result<Vec<SocketAddr>, FilterError> {
|
||||
// Port whitelist check
|
||||
if !allowed_ports.contains(&port) {
|
||||
return Err(FilterError::PortNotAllowed(port));
|
||||
}
|
||||
|
||||
// Try parsing as IP directly (no DNS needed)
|
||||
if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
if is_private_ip(&ip) {
|
||||
return Err(FilterError::PrivateIp(ip));
|
||||
}
|
||||
return Ok(vec![SocketAddr::new(ip, port)]);
|
||||
}
|
||||
|
||||
// Resolve and validate DNS (populates cache for SafeDnsResolver)
|
||||
resolve_public_addrs(host, port, dns_cache).await
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn ports() -> HashSet<u16> {
|
||||
[80, 443, 8080, 8443].into_iter().collect()
|
||||
}
|
||||
|
||||
fn cache() -> DnsCache {
|
||||
DnsCache::new(Duration::from_secs(60), 128)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_private_ipv4() {
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1))));
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(172, 16, 0, 1))));
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1))));
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1))));
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(169, 254, 1, 1))));
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(0, 0, 0, 0))));
|
||||
// CGNAT
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1))));
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(
|
||||
100, 127, 255, 254
|
||||
))));
|
||||
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(
|
||||
100, 63, 255, 254
|
||||
))));
|
||||
// Benchmark testing
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(198, 18, 0, 1))));
|
||||
// Reserved
|
||||
assert!(is_private_ip(&IpAddr::V4(Ipv4Addr::new(240, 0, 0, 1))));
|
||||
// Public
|
||||
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
|
||||
assert!(!is_private_ip(&IpAddr::V4(Ipv4Addr::new(203, 0, 113, 1))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_private_ipv6() {
|
||||
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::LOCALHOST)));
|
||||
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::UNSPECIFIED)));
|
||||
// fc00::1 (ULA)
|
||||
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new(
|
||||
0xfc00, 0, 0, 0, 0, 0, 0, 1
|
||||
))));
|
||||
// fe80::1 (link-local)
|
||||
assert!(is_private_ip(&IpAddr::V6(Ipv6Addr::new(
|
||||
0xfe80, 0, 0, 0, 0, 0, 0, 1
|
||||
))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_port_not_allowed() {
|
||||
let cache = cache();
|
||||
let result = validate_target("8.8.8.8", 22, &ports(), &cache).await;
|
||||
assert!(matches!(result, Err(FilterError::PortNotAllowed(22))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_private_ip_blocked() {
|
||||
let cache = cache();
|
||||
let result = validate_target("127.0.0.1", 80, &ports(), &cache).await;
|
||||
assert!(matches!(result, Err(FilterError::PrivateIp(_))));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_public_ip_allowed() {
|
||||
let cache = cache();
|
||||
let result = validate_target("8.8.8.8", 443, &ports(), &cache).await;
|
||||
assert!(result.is_ok());
|
||||
let addrs = result.unwrap();
|
||||
assert_eq!(addrs.len(), 1);
|
||||
assert_eq!(addrs[0].ip(), IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_stores_multiple_addrs() {
|
||||
let cache = cache();
|
||||
let addrs = vec![
|
||||
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443),
|
||||
SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 0, 0, 1)), 443),
|
||||
];
|
||||
cache
|
||||
.insert("example.com", 443, Arc::new(addrs.clone()))
|
||||
.await;
|
||||
let cached = cache.get("example.com", 443).await.unwrap();
|
||||
assert_eq!(*cached, addrs);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_cache_key_case_insensitive() {
|
||||
let cache = cache();
|
||||
let addrs = vec![SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), 443)];
|
||||
cache
|
||||
.insert("Example.COM", 443, Arc::new(addrs.clone()))
|
||||
.await;
|
||||
let cached = cache.get("example.com", 443).await.unwrap();
|
||||
assert_eq!(*cached, addrs);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,238 @@
|
||||
//! WebSocket tunnel client: connect, authenticate, and run the tunnel.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use tokio::net::TcpStream;
|
||||
use tokio::sync::watch;
|
||||
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
|
||||
use tokio_tungstenite::tungstenite::http;
|
||||
use tokio_tungstenite::tungstenite::protocol::WebSocketConfig;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use crate::state::{AppState, ServerContext};
|
||||
|
||||
use super::{dispatcher, heartbeat, writer};
|
||||
|
||||
/// Outcome of a tunnel session.
|
||||
pub enum TunnelOutcome {
|
||||
/// Graceful shutdown requested by the local process.
|
||||
Shutdown,
|
||||
/// Remote side disconnected or connection lost — should reconnect.
|
||||
Disconnected,
|
||||
}
|
||||
|
||||
/// Connect to Aether's WebSocket tunnel endpoint and run until disconnected.
|
||||
///
|
||||
/// `conn_idx` identifies which connection in the pool this is (0-based).
|
||||
/// Only connection 0 sends heartbeats to avoid resetting shared metrics.
|
||||
pub async fn connect_and_run(
|
||||
state: &Arc<AppState>,
|
||||
server: &Arc<ServerContext>,
|
||||
conn_idx: usize,
|
||||
shutdown: &mut watch::Receiver<bool>,
|
||||
) -> Result<TunnelOutcome, anyhow::Error> {
|
||||
let ws_url = build_tunnel_url(server);
|
||||
info!(url = %ws_url, conn = conn_idx, "connecting tunnel");
|
||||
|
||||
// Build WebSocket request with auth headers
|
||||
let mut request = ws_url.clone().into_client_request()?;
|
||||
let headers = request.headers_mut();
|
||||
headers.insert(
|
||||
"Authorization",
|
||||
http::HeaderValue::from_str(&format!("Bearer {}", server.management_token))?,
|
||||
);
|
||||
let node_id = server.node_id.read().unwrap().clone();
|
||||
headers.insert("X-Node-Id", http::HeaderValue::from_str(&node_id)?);
|
||||
// Use dynamic node_name (may be updated by remote config) instead of
|
||||
// the static server.node_name, so that remote name changes take effect
|
||||
// on the next reconnect.
|
||||
let dynamic_node_name = server.dynamic.load().node_name.clone();
|
||||
headers.insert(
|
||||
"X-Node-Name",
|
||||
http::HeaderValue::from_str(&dynamic_node_name)?,
|
||||
);
|
||||
// Advertise per-connection max concurrent streams so the backend can
|
||||
// respect the proxy's capacity limit (backward-compatible: old backends
|
||||
// ignore this header).
|
||||
let max_streams = state.config.tunnel_max_streams.unwrap_or(128);
|
||||
headers.insert("X-Tunnel-Max-Streams", http::HeaderValue::from(max_streams));
|
||||
|
||||
// Parse host:port from URL
|
||||
let uri: http::Uri = ws_url.parse()?;
|
||||
let host = uri
|
||||
.host()
|
||||
.ok_or_else(|| anyhow::anyhow!("missing host in tunnel URL"))?;
|
||||
let is_tls = uri.scheme_str() == Some("wss");
|
||||
let port = uri.port_u16().unwrap_or(if is_tls { 443 } else { 80 });
|
||||
|
||||
// TCP connect with timeout
|
||||
let connect_timeout = Duration::from_secs(state.config.tunnel_connect_timeout_secs);
|
||||
let tcp_stream = tokio::time::timeout(connect_timeout, TcpStream::connect((host, port)))
|
||||
.await
|
||||
.map_err(|_| {
|
||||
anyhow::anyhow!(
|
||||
"tunnel TCP connect timeout ({}s)",
|
||||
connect_timeout.as_secs()
|
||||
)
|
||||
})??;
|
||||
|
||||
// Configure TCP parameters via socket2
|
||||
configure_tcp_socket(&tcp_stream, state);
|
||||
|
||||
// WebSocket upgrade (with TLS if wss://)
|
||||
let connector = if is_tls {
|
||||
Some(tokio_tungstenite::Connector::Rustls(Arc::clone(
|
||||
&state.tunnel_tls_config,
|
||||
)))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
// Match Python-side _MAX_FRAME_SIZE (64 MiB) to prevent tungstenite's
|
||||
// default 16 MiB limit from rejecting large AI API payloads (multi-image
|
||||
// base64 requests can exceed 16 MiB).
|
||||
let ws_config = WebSocketConfig {
|
||||
max_frame_size: Some(64 << 20),
|
||||
max_message_size: Some(64 << 20),
|
||||
..Default::default()
|
||||
};
|
||||
let handshake_timeout = Duration::from_secs(state.config.tunnel_connect_timeout_secs);
|
||||
let (ws_stream, _response) = tokio::time::timeout(
|
||||
handshake_timeout,
|
||||
tokio_tungstenite::client_async_tls_with_config(
|
||||
request,
|
||||
tcp_stream,
|
||||
Some(ws_config),
|
||||
connector,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| {
|
||||
anyhow::anyhow!(
|
||||
"tunnel WebSocket handshake timeout ({}s)",
|
||||
handshake_timeout.as_secs()
|
||||
)
|
||||
})??;
|
||||
info!(
|
||||
conn = conn_idx,
|
||||
tcp_keepalive_secs = state.config.tunnel_tcp_keepalive_secs,
|
||||
tcp_nodelay = state.config.tunnel_tcp_nodelay,
|
||||
connect_timeout_secs = state.config.tunnel_connect_timeout_secs,
|
||||
stale_timeout_secs = state.config.tunnel_stale_timeout_secs,
|
||||
"tunnel connected"
|
||||
);
|
||||
|
||||
// NOTE: reconnect_attempts reset is handled by the caller (mod.rs)
|
||||
// based on how long the connection stayed alive.
|
||||
|
||||
// Split into read/write halves
|
||||
let (ws_sink, ws_read) = futures_util::StreamExt::split(ws_stream);
|
||||
|
||||
// Spawn writer task (with WebSocket ping keepalive)
|
||||
let ping_interval = Duration::from_secs(state.config.tunnel_ping_interval_secs);
|
||||
let (frame_tx, mut writer_handle) = writer::spawn_writer(ws_sink, ping_interval);
|
||||
|
||||
// Spawn heartbeat task (only for primary connection to avoid
|
||||
// resetting shared atomic metrics via swap(0))
|
||||
let hb_handle = if conn_idx == 0 {
|
||||
heartbeat::spawn(
|
||||
Arc::clone(&state.config),
|
||||
Arc::clone(server),
|
||||
frame_tx.clone(),
|
||||
shutdown.clone(),
|
||||
)
|
||||
} else {
|
||||
heartbeat::spawn_noop()
|
||||
};
|
||||
|
||||
// Run dispatcher (blocks until disconnect or shutdown).
|
||||
// Also watch for writer exit — if the write half dies (e.g. the peer
|
||||
// closed the connection) but the read half stays open, dispatcher would
|
||||
// block forever on `ws_stream.next()`. Monitoring `writer_handle`
|
||||
// ensures we detect this and trigger a reconnect promptly.
|
||||
let state_clone = Arc::clone(state);
|
||||
let server_clone = Arc::clone(server);
|
||||
let outcome = tokio::select! {
|
||||
result = dispatcher::run(state_clone, server_clone, ws_read, frame_tx.clone(), hb_handle) => {
|
||||
match result {
|
||||
Ok(()) => TunnelOutcome::Disconnected,
|
||||
Err(e) => return Err(e),
|
||||
}
|
||||
}
|
||||
writer_result = &mut writer_handle => {
|
||||
match writer_result {
|
||||
Ok(()) => warn!("writer task exited normally, triggering reconnect"),
|
||||
Err(e) => {
|
||||
if e.is_panic() {
|
||||
tracing::error!(error = %e, "writer task panicked, triggering reconnect");
|
||||
} else {
|
||||
warn!(error = %e, "writer task cancelled, triggering reconnect");
|
||||
}
|
||||
}
|
||||
}
|
||||
TunnelOutcome::Disconnected
|
||||
}
|
||||
_ = shutdown.changed() => {
|
||||
debug!("shutdown during tunnel dispatch");
|
||||
TunnelOutcome::Shutdown
|
||||
}
|
||||
};
|
||||
|
||||
// Drop our sender; the writer will exit once all stream handler clones
|
||||
// are also dropped (i.e. after they finish their in-flight work).
|
||||
drop(frame_tx);
|
||||
|
||||
// Wait for the writer task to finish with a generous timeout — the
|
||||
// dispatcher already waits up to 30s for stream handlers, so 35s here
|
||||
// covers that plus a small margin.
|
||||
// Skip if the writer already exited (the select branch that fired).
|
||||
if !writer_handle.is_finished() {
|
||||
let _ = tokio::time::timeout(Duration::from_secs(35), writer_handle).await;
|
||||
}
|
||||
|
||||
info!("tunnel disconnected");
|
||||
Ok(outcome)
|
||||
}
|
||||
|
||||
/// Configure TCP keepalive and NODELAY on an established socket.
|
||||
fn configure_tcp_socket(stream: &TcpStream, state: &Arc<AppState>) {
|
||||
let sock_ref = socket2::SockRef::from(stream);
|
||||
|
||||
if state.config.tunnel_tcp_keepalive_secs > 0 {
|
||||
let keepalive = socket2::TcpKeepalive::new()
|
||||
.with_time(Duration::from_secs(state.config.tunnel_tcp_keepalive_secs))
|
||||
.with_interval(Duration::from_secs(5));
|
||||
#[cfg(not(target_os = "windows"))]
|
||||
let keepalive = keepalive.with_retries(3);
|
||||
if let Err(e) = sock_ref.set_tcp_keepalive(&keepalive) {
|
||||
warn!(error = %e, "failed to set TCP keepalive on tunnel socket");
|
||||
}
|
||||
}
|
||||
|
||||
if state.config.tunnel_tcp_nodelay {
|
||||
if let Err(e) = sock_ref.set_nodelay(true) {
|
||||
warn!(error = %e, "failed to set TCP_NODELAY on tunnel socket");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Build rustls ClientConfig with system root certificates.
|
||||
pub fn build_tls_config() -> rustls::ClientConfig {
|
||||
let root_store =
|
||||
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
|
||||
rustls::ClientConfig::builder()
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth()
|
||||
}
|
||||
|
||||
fn build_tunnel_url(server: &ServerContext) -> String {
|
||||
let base = server.aether_url.trim_end_matches('/');
|
||||
let ws_base = if base.starts_with("https://") {
|
||||
base.replacen("https://", "wss://", 1)
|
||||
} else if base.starts_with("http://") {
|
||||
base.replacen("http://", "ws://", 1)
|
||||
} else {
|
||||
format!("wss://{}", base)
|
||||
};
|
||||
format!("{}/api/internal/proxy-tunnel", ws_base)
|
||||
}
|
||||
@@ -0,0 +1,249 @@
|
||||
//! Frame dispatcher: reads incoming WebSocket frames and routes them.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::StreamExt;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
use crate::state::{AppState, ServerContext};
|
||||
|
||||
use super::heartbeat::HeartbeatHandle;
|
||||
use super::protocol::{decompress_if_gzip, Frame, MsgType, RequestMeta};
|
||||
use super::stream_handler;
|
||||
use super::writer::FrameSender;
|
||||
|
||||
/// Run the dispatcher loop, reading from the WebSocket stream.
|
||||
pub async fn run<S>(
|
||||
state: Arc<AppState>,
|
||||
server: Arc<ServerContext>,
|
||||
mut ws_stream: S,
|
||||
frame_tx: FrameSender,
|
||||
heartbeat: HeartbeatHandle,
|
||||
) -> Result<(), anyhow::Error>
|
||||
where
|
||||
S: StreamExt<Item = Result<Message, tokio_tungstenite::tungstenite::Error>>
|
||||
+ Unpin
|
||||
+ Send
|
||||
+ 'static,
|
||||
{
|
||||
// Active streams: stream_id -> body sender
|
||||
let mut streams: HashMap<u32, mpsc::Sender<Frame>> = HashMap::new();
|
||||
// Track spawned stream handlers so we can wait for them on shutdown
|
||||
let mut handler_handles: Vec<JoinHandle<()>> = Vec::new();
|
||||
let max_streams = state.config.tunnel_max_streams.unwrap_or(128) as usize;
|
||||
let mut frames_since_cleanup: u32 = 0;
|
||||
let stale_timeout = Duration::from_secs(state.config.tunnel_stale_timeout_secs);
|
||||
|
||||
// Track last time we received any data to detect stale connections
|
||||
let mut last_data_at = tokio::time::Instant::now();
|
||||
|
||||
let read_err = loop {
|
||||
let msg_result = tokio::select! {
|
||||
msg = ws_stream.next() => {
|
||||
match msg {
|
||||
Some(r) => r,
|
||||
None => break None,
|
||||
}
|
||||
}
|
||||
_ = tokio::time::sleep_until(last_data_at + stale_timeout) => {
|
||||
warn!(
|
||||
stale_secs = stale_timeout.as_secs(),
|
||||
"tunnel connection stale, no data received"
|
||||
);
|
||||
break None;
|
||||
}
|
||||
};
|
||||
|
||||
let msg = match msg_result {
|
||||
Ok(m) => m,
|
||||
Err(e) => {
|
||||
error!(error = %e, "WebSocket read error");
|
||||
break Some(e);
|
||||
}
|
||||
};
|
||||
|
||||
// Any successfully received message proves the connection is alive
|
||||
last_data_at = tokio::time::Instant::now();
|
||||
|
||||
let data = match msg {
|
||||
Message::Binary(data) => Bytes::from(data),
|
||||
Message::Ping(_) => continue,
|
||||
Message::Pong(_) => continue,
|
||||
Message::Close(_) => {
|
||||
info!("received WebSocket close");
|
||||
break None;
|
||||
}
|
||||
_ => continue,
|
||||
};
|
||||
|
||||
let frame = match Frame::decode(data) {
|
||||
Ok(f) => f,
|
||||
Err(e) => {
|
||||
warn!(error = %e, "failed to decode frame");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
match frame.msg_type {
|
||||
MsgType::RequestHeaders => {
|
||||
// Decompress if the frame is gzip-compressed, then parse metadata
|
||||
let payload = match decompress_if_gzip(&frame) {
|
||||
Ok(p) => p,
|
||||
Err(e) => {
|
||||
warn!(stream_id = frame.stream_id, error = %e, "frame decompress failed");
|
||||
continue;
|
||||
}
|
||||
};
|
||||
let meta: RequestMeta = match serde_json::from_slice(&payload) {
|
||||
Ok(m) => m,
|
||||
Err(e) => {
|
||||
warn!(stream_id = frame.stream_id, error = %e, "invalid request metadata");
|
||||
// Use try_send to avoid blocking the read loop
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from(format!("invalid request metadata: {e}")),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped"
|
||||
);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
if streams.len() >= max_streams {
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"max concurrent streams reached"
|
||||
);
|
||||
if frame_tx
|
||||
.try_send(Frame::new(
|
||||
frame.stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from("max concurrent streams reached"),
|
||||
))
|
||||
.is_err()
|
||||
{
|
||||
warn!(
|
||||
stream_id = frame.stream_id,
|
||||
"writer channel full, StreamError dropped"
|
||||
);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
// Create body channel and spawn handler
|
||||
let (body_tx, body_rx) = mpsc::channel::<Frame>(64);
|
||||
streams.insert(frame.stream_id, body_tx);
|
||||
|
||||
let state_clone = Arc::clone(&state);
|
||||
let server_clone = Arc::clone(&server);
|
||||
let tx_clone = frame_tx.clone();
|
||||
let sid = frame.stream_id;
|
||||
let handle = tokio::spawn(async move {
|
||||
stream_handler::handle_stream(
|
||||
state_clone,
|
||||
server_clone,
|
||||
sid,
|
||||
meta,
|
||||
body_rx,
|
||||
tx_clone,
|
||||
)
|
||||
.await;
|
||||
});
|
||||
handler_handles.push(handle);
|
||||
|
||||
debug!(stream_id = frame.stream_id, "new stream started");
|
||||
}
|
||||
|
||||
MsgType::RequestBody => {
|
||||
if let Some(tx) = streams.get(&frame.stream_id) {
|
||||
let is_end = frame.is_end_stream();
|
||||
let sid = frame.stream_id;
|
||||
let _ = tx.send(frame).await;
|
||||
if is_end {
|
||||
streams.remove(&sid);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::StreamEnd | MsgType::StreamError => {
|
||||
// Client-side cancellation or end
|
||||
if let Some(tx) = streams.remove(&frame.stream_id) {
|
||||
let _ = tx.send(frame).await;
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::Ping => {
|
||||
// Use try_send to avoid blocking the read loop when writer is congested
|
||||
if frame_tx
|
||||
.try_send(Frame::control(MsgType::Pong, frame.payload))
|
||||
.is_err()
|
||||
{
|
||||
warn!("writer channel full, Pong dropped");
|
||||
}
|
||||
}
|
||||
|
||||
MsgType::HeartbeatAck => {
|
||||
heartbeat.on_ack(frame.payload).await;
|
||||
}
|
||||
|
||||
MsgType::GoAway => {
|
||||
info!("received GOAWAY");
|
||||
break None;
|
||||
}
|
||||
|
||||
_ => {
|
||||
debug!(msg_type = ?frame.msg_type, "ignoring unexpected frame type");
|
||||
}
|
||||
}
|
||||
|
||||
// Periodically clean up finished handles to avoid unbounded growth.
|
||||
// Trigger every 64 frames OR when the count exceeds max_streams.
|
||||
frames_since_cleanup += 1;
|
||||
if frames_since_cleanup >= 64 || handler_handles.len() > max_streams {
|
||||
handler_handles.retain(|h| !h.is_finished());
|
||||
frames_since_cleanup = 0;
|
||||
}
|
||||
};
|
||||
|
||||
// Drop body senders so stream handlers waiting on body_rx will unblock
|
||||
streams.clear();
|
||||
|
||||
// Wait for active stream handlers to finish so their frame_tx clones
|
||||
// are dropped before the writer closes the sink.
|
||||
drain_handlers(handler_handles).await;
|
||||
|
||||
match read_err {
|
||||
Some(e) => Err(e.into()),
|
||||
None => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Wait for all active stream handlers to finish (with a timeout).
|
||||
async fn drain_handlers(handles: Vec<JoinHandle<()>>) {
|
||||
if handles.is_empty() {
|
||||
return;
|
||||
}
|
||||
let count = handles.len();
|
||||
debug!(count, "waiting for active stream handlers to finish");
|
||||
let _ = tokio::time::timeout(Duration::from_secs(30), async {
|
||||
for h in handles {
|
||||
let _ = h.await;
|
||||
}
|
||||
})
|
||||
.await;
|
||||
}
|
||||
@@ -0,0 +1,340 @@
|
||||
//! Tunnel heartbeat: sends metrics over the tunnel, processes ACKs.
|
||||
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
use std::time::SystemTime;
|
||||
use std::time::UNIX_EPOCH;
|
||||
|
||||
use bytes::Bytes;
|
||||
use tokio::sync::watch;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use crate::config::Config;
|
||||
use crate::registration::client::RemoteConfig;
|
||||
use crate::runtime;
|
||||
use crate::state::ServerContext;
|
||||
|
||||
use super::protocol::{Frame, MsgType};
|
||||
use super::writer::FrameSender;
|
||||
|
||||
const CURRENT_VERSION: &str = env!("CARGO_PKG_VERSION");
|
||||
static UPGRADE_IN_PROGRESS: AtomicBool = AtomicBool::new(false);
|
||||
static NON_ROOT_UPGRADE_WARNED: AtomicBool = AtomicBool::new(false);
|
||||
|
||||
enum AckDecision {
|
||||
Accept {
|
||||
heartbeat_id: Option<u64>,
|
||||
upgrade_to: Option<String>,
|
||||
},
|
||||
Ignore,
|
||||
}
|
||||
|
||||
/// Handle for the dispatcher to forward HeartbeatAck frames.
|
||||
#[derive(Clone)]
|
||||
pub struct HeartbeatHandle {
|
||||
ack_tx: tokio::sync::mpsc::Sender<Bytes>,
|
||||
}
|
||||
|
||||
impl HeartbeatHandle {
|
||||
pub async fn on_ack(&self, payload: Bytes) {
|
||||
let _ = self.ack_tx.send(payload).await;
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a no-op heartbeat handle that silently discards ACKs.
|
||||
/// Used for non-primary tunnel connections (conn_idx > 0) to avoid
|
||||
/// resetting shared atomic metrics via `swap(0)`.
|
||||
pub fn spawn_noop() -> HeartbeatHandle {
|
||||
let (ack_tx, _) = tokio::sync::mpsc::channel::<Bytes>(1);
|
||||
// receiver is immediately dropped; on_ack() calls will silently fail
|
||||
HeartbeatHandle { ack_tx }
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct HeartbeatSnapshot {
|
||||
requests: u64,
|
||||
latency_ns: u64,
|
||||
failed: u64,
|
||||
dns_failures: u64,
|
||||
stream_errors: u64,
|
||||
}
|
||||
|
||||
/// Spawn the heartbeat task. Returns a handle for forwarding ACKs.
|
||||
pub fn spawn(
|
||||
_config: Arc<Config>,
|
||||
server: Arc<ServerContext>,
|
||||
frame_tx: FrameSender,
|
||||
mut shutdown: watch::Receiver<bool>,
|
||||
) -> HeartbeatHandle {
|
||||
let (ack_tx, mut ack_rx) = tokio::sync::mpsc::channel::<Bytes>(4);
|
||||
|
||||
tokio::spawn(async move {
|
||||
// Read initial interval from dynamic config (may be updated by remote config).
|
||||
let initial_interval = Duration::from_secs(server.dynamic.load().heartbeat_interval);
|
||||
let mut current_interval = initial_interval;
|
||||
// At most one in-flight heartbeat snapshot is tracked at a time.
|
||||
// Snapshot is only cleared after receiving an ACK, which avoids losing
|
||||
// interval counters when ACK/frame delivery is temporarily unstable.
|
||||
let mut pending: Option<(u64, HeartbeatSnapshot)> = None;
|
||||
let mut next_heartbeat_id: u64 = 1;
|
||||
let heartbeat_session_id = format!(
|
||||
"{}-{}",
|
||||
std::process::id(),
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_nanos()
|
||||
);
|
||||
|
||||
// Skip first immediate tick by sleeping first.
|
||||
tokio::time::sleep(current_interval).await;
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(current_interval) => {
|
||||
let (heartbeat_id, snapshot) = if let Some((id, snap)) = pending {
|
||||
(id, snap)
|
||||
} else {
|
||||
let snap = collect_snapshot(&server);
|
||||
let id = next_heartbeat_id;
|
||||
next_heartbeat_id = next_heartbeat_id.wrapping_add(1);
|
||||
if next_heartbeat_id == 0 {
|
||||
next_heartbeat_id = 1;
|
||||
}
|
||||
pending = Some((id, snap));
|
||||
(id, snap)
|
||||
};
|
||||
|
||||
let payload = build_heartbeat_payload(
|
||||
&server,
|
||||
&heartbeat_session_id,
|
||||
heartbeat_id,
|
||||
snapshot
|
||||
);
|
||||
let frame = Frame::control(MsgType::HeartbeatData, payload);
|
||||
if frame_tx.send(frame).await.is_err() {
|
||||
if let Some((_, snap)) = pending.take() {
|
||||
restore_snapshot(&server, snap);
|
||||
}
|
||||
break; // Writer closed
|
||||
}
|
||||
debug!("sent heartbeat data");
|
||||
|
||||
// Re-read interval from dynamic config (remote config may have
|
||||
// updated it since the last heartbeat).
|
||||
let new_interval = Duration::from_secs(
|
||||
server.dynamic.load().heartbeat_interval
|
||||
);
|
||||
if new_interval != current_interval {
|
||||
debug!(
|
||||
old_secs = current_interval.as_secs(),
|
||||
new_secs = new_interval.as_secs(),
|
||||
"heartbeat interval updated from dynamic config"
|
||||
);
|
||||
current_interval = new_interval;
|
||||
}
|
||||
}
|
||||
Some(ack_payload) = ack_rx.recv() => {
|
||||
match handle_ack(&server, &ack_payload) {
|
||||
AckDecision::Accept {
|
||||
heartbeat_id: ack_id,
|
||||
upgrade_to,
|
||||
} => {
|
||||
if let Some((pending_id, _)) = pending {
|
||||
match ack_id {
|
||||
Some(id) if id == pending_id => {
|
||||
pending = None;
|
||||
}
|
||||
None => {
|
||||
// Backward-compatible with servers that don't echo
|
||||
// heartbeat_id in ACK payload yet.
|
||||
pending = None;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
maybe_trigger_upgrade(upgrade_to);
|
||||
}
|
||||
AckDecision::Ignore => {}
|
||||
}
|
||||
}
|
||||
_ = shutdown.changed() => {
|
||||
debug!("heartbeat task shutting down");
|
||||
if let Some((_, snap)) = pending.take() {
|
||||
restore_snapshot(&server, snap);
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
HeartbeatHandle { ack_tx }
|
||||
}
|
||||
|
||||
fn collect_snapshot(server: &ServerContext) -> HeartbeatSnapshot {
|
||||
HeartbeatSnapshot {
|
||||
requests: server.metrics.total_requests.swap(0, Ordering::AcqRel),
|
||||
latency_ns: server.metrics.total_latency_ns.swap(0, Ordering::AcqRel),
|
||||
failed: server.metrics.failed_requests.swap(0, Ordering::AcqRel),
|
||||
dns_failures: server.metrics.dns_failures.swap(0, Ordering::AcqRel),
|
||||
stream_errors: server.metrics.stream_errors.swap(0, Ordering::AcqRel),
|
||||
}
|
||||
}
|
||||
|
||||
fn restore_snapshot(server: &ServerContext, snap: HeartbeatSnapshot) {
|
||||
if snap.requests > 0 {
|
||||
server
|
||||
.metrics
|
||||
.total_requests
|
||||
.fetch_add(snap.requests, Ordering::Release);
|
||||
}
|
||||
if snap.latency_ns > 0 {
|
||||
server
|
||||
.metrics
|
||||
.total_latency_ns
|
||||
.fetch_add(snap.latency_ns, Ordering::Release);
|
||||
}
|
||||
if snap.failed > 0 {
|
||||
server
|
||||
.metrics
|
||||
.failed_requests
|
||||
.fetch_add(snap.failed, Ordering::Release);
|
||||
}
|
||||
if snap.dns_failures > 0 {
|
||||
server
|
||||
.metrics
|
||||
.dns_failures
|
||||
.fetch_add(snap.dns_failures, Ordering::Release);
|
||||
}
|
||||
if snap.stream_errors > 0 {
|
||||
server
|
||||
.metrics
|
||||
.stream_errors
|
||||
.fetch_add(snap.stream_errors, Ordering::Release);
|
||||
}
|
||||
}
|
||||
|
||||
fn build_heartbeat_payload(
|
||||
server: &ServerContext,
|
||||
heartbeat_session_id: &str,
|
||||
heartbeat_id: u64,
|
||||
snapshot: HeartbeatSnapshot,
|
||||
) -> Bytes {
|
||||
let node_id = server.node_id.read().unwrap().clone();
|
||||
|
||||
let avg_latency_ms = if snapshot.requests > 0 {
|
||||
Some(snapshot.latency_ns as f64 / snapshot.requests as f64 / 1_000_000.0)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let payload = serde_json::json!({
|
||||
"node_id": node_id,
|
||||
"heartbeat_session_id": heartbeat_session_id,
|
||||
"heartbeat_id": heartbeat_id,
|
||||
"active_connections": server.active_connections.load(Ordering::Acquire),
|
||||
"total_requests": snapshot.requests,
|
||||
"avg_latency_ms": avg_latency_ms,
|
||||
"failed_requests": snapshot.failed,
|
||||
"dns_failures": snapshot.dns_failures,
|
||||
"stream_errors": snapshot.stream_errors,
|
||||
"proxy_metadata": {
|
||||
"version": CURRENT_VERSION,
|
||||
},
|
||||
});
|
||||
|
||||
Bytes::from(serde_json::to_vec(&payload).unwrap_or_default())
|
||||
}
|
||||
|
||||
fn handle_ack(server: &ServerContext, payload: &[u8]) -> AckDecision {
|
||||
if payload.is_empty() {
|
||||
return AckDecision::Accept {
|
||||
heartbeat_id: None,
|
||||
upgrade_to: None,
|
||||
};
|
||||
}
|
||||
|
||||
#[derive(serde::Deserialize)]
|
||||
struct AckPayload {
|
||||
#[serde(default)]
|
||||
remote_config: Option<RemoteConfig>,
|
||||
#[serde(default)]
|
||||
config_version: u64,
|
||||
#[serde(default)]
|
||||
heartbeat_id: Option<u64>,
|
||||
#[serde(default)]
|
||||
upgrade_to: Option<String>,
|
||||
}
|
||||
|
||||
match serde_json::from_slice::<AckPayload>(payload) {
|
||||
Ok(ack) => {
|
||||
if let Some(ref rc) = ack.remote_config {
|
||||
runtime::apply_remote_config(&server.dynamic, rc, ack.config_version);
|
||||
}
|
||||
AckDecision::Accept {
|
||||
heartbeat_id: ack.heartbeat_id,
|
||||
upgrade_to: ack.upgrade_to.and_then(normalize_upgrade_target),
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(error = %e, "failed to parse heartbeat ACK");
|
||||
AckDecision::Ignore
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_upgrade_target(raw: String) -> Option<String> {
|
||||
let trimmed = raw.trim();
|
||||
if trimmed.is_empty() {
|
||||
return None;
|
||||
}
|
||||
let normalized = trimmed.strip_prefix("proxy-v").unwrap_or(trimmed);
|
||||
if normalized == CURRENT_VERSION {
|
||||
return None;
|
||||
}
|
||||
Some(normalized.to_string())
|
||||
}
|
||||
|
||||
fn maybe_trigger_upgrade(version: Option<String>) {
|
||||
let Some(target_version) = version else {
|
||||
return;
|
||||
};
|
||||
if !crate::setup::service::is_root() {
|
||||
if NON_ROOT_UPGRADE_WARNED
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.is_ok()
|
||||
{
|
||||
warn!(
|
||||
target_version = %target_version,
|
||||
"remote upgrade skipped: root privileges are required"
|
||||
);
|
||||
}
|
||||
return;
|
||||
}
|
||||
if UPGRADE_IN_PROGRESS
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.is_err()
|
||||
{
|
||||
debug!(target_version = %target_version, "upgrade already in progress, ignoring");
|
||||
return;
|
||||
}
|
||||
|
||||
tokio::spawn(async move {
|
||||
info!(target_version = %target_version, "received remote upgrade instruction");
|
||||
match crate::setup::upgrade::perform_upgrade(&target_version).await {
|
||||
Ok(()) => {
|
||||
info!(target_version = %target_version, "remote upgrade finished");
|
||||
}
|
||||
Err(e) => {
|
||||
warn!(
|
||||
target_version = %target_version,
|
||||
error = %e,
|
||||
"remote upgrade failed"
|
||||
);
|
||||
UPGRADE_IN_PROGRESS.store(false, Ordering::Release);
|
||||
}
|
||||
}
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
pub mod client;
|
||||
pub mod dispatcher;
|
||||
pub mod heartbeat;
|
||||
pub mod protocol;
|
||||
pub mod stream_handler;
|
||||
pub mod writer;
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use tokio::sync::watch;
|
||||
use tracing::{error, info};
|
||||
|
||||
use crate::state::{AppState, ServerContext};
|
||||
|
||||
/// If a tunnel stays connected at least this long, treat the next disconnect
|
||||
/// as a non-failure and reset reconnect backoff.
|
||||
const STABLE_SESSION_RESET_AFTER: Duration = Duration::from_secs(30);
|
||||
/// Startup staggering step per secondary connection, used to avoid
|
||||
/// simultaneous bursts when a pool of tunnels starts together.
|
||||
const STARTUP_STAGGER_STEP_MS: u64 = 150;
|
||||
/// Upper bound for startup staggering.
|
||||
const MAX_STARTUP_STAGGER_MS: u64 = 1_500;
|
||||
/// Keep a tiny floor for repeated reconnects; first retry is still immediate.
|
||||
const MIN_RECONNECT_DELAY_MS: u64 = 50;
|
||||
/// Even under sustained failures, keep probing frequently so recovery is fast
|
||||
/// once cross-border network quality improves.
|
||||
const RECONNECT_PROBE_MAX_DELAY_MS: u64 = 3_000;
|
||||
|
||||
/// Run the tunnel mode main loop (connect, dispatch, reconnect).
|
||||
///
|
||||
/// `conn_idx` identifies which connection in the pool this is (0-based).
|
||||
/// Only connection 0 sends heartbeats to avoid resetting shared metrics.
|
||||
pub async fn run(
|
||||
state: &Arc<AppState>,
|
||||
server: &Arc<ServerContext>,
|
||||
conn_idx: usize,
|
||||
mut shutdown: watch::Receiver<bool>,
|
||||
) {
|
||||
info!(server = %server.server_label, conn = conn_idx, "starting tunnel");
|
||||
let reconnect_salt = compute_connection_salt(server, conn_idx);
|
||||
|
||||
let startup_delay = compute_startup_stagger(conn_idx, reconnect_salt);
|
||||
if !startup_delay.is_zero() {
|
||||
info!(
|
||||
server = %server.server_label,
|
||||
conn = conn_idx,
|
||||
delay_ms = startup_delay.as_millis(),
|
||||
"startup stagger before first connect"
|
||||
);
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(startup_delay) => {}
|
||||
_ = shutdown.changed() => {
|
||||
info!(server = %server.server_label, conn = conn_idx, "shutdown requested during startup stagger");
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut consecutive_failures: u32 = 0;
|
||||
|
||||
loop {
|
||||
let started_at = Instant::now();
|
||||
match client::connect_and_run(state, server, conn_idx, &mut shutdown).await {
|
||||
Ok(client::TunnelOutcome::Shutdown) => {
|
||||
info!(server = %server.server_label, conn = conn_idx, "tunnel shut down gracefully");
|
||||
return;
|
||||
}
|
||||
Ok(client::TunnelOutcome::Disconnected) => {
|
||||
info!(server = %server.server_label, conn = conn_idx, "tunnel disconnected, reconnecting");
|
||||
}
|
||||
Err(e) => {
|
||||
error!(server = %server.server_label, conn = conn_idx, error = %e, "tunnel connection error, reconnecting");
|
||||
}
|
||||
}
|
||||
|
||||
if *shutdown.borrow() {
|
||||
info!(server = %server.server_label, conn = conn_idx, "shutdown requested, not reconnecting");
|
||||
return;
|
||||
}
|
||||
|
||||
// Reset backoff after a stable session to keep recovery snappy when
|
||||
// failures are only occasional.
|
||||
let connected_for = started_at.elapsed();
|
||||
if connected_for >= STABLE_SESSION_RESET_AFTER {
|
||||
consecutive_failures = 0;
|
||||
} else {
|
||||
consecutive_failures = consecutive_failures.saturating_add(1);
|
||||
}
|
||||
|
||||
let reconnect_delay = compute_reconnect_delay(
|
||||
state.config.tunnel_reconnect_base_ms,
|
||||
state.config.tunnel_reconnect_max_ms,
|
||||
consecutive_failures,
|
||||
reconnect_salt,
|
||||
);
|
||||
info!(
|
||||
server = %server.server_label,
|
||||
conn = conn_idx,
|
||||
failures = consecutive_failures,
|
||||
delay_ms = reconnect_delay.as_millis(),
|
||||
"waiting before reconnect"
|
||||
);
|
||||
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(reconnect_delay) => {}
|
||||
_ = shutdown.changed() => {
|
||||
info!(server = %server.server_label, conn = conn_idx, "shutdown requested during reconnect wait");
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn compute_connection_salt(server: &ServerContext, conn_idx: usize) -> u64 {
|
||||
// FNV-1a style hash over server label + connection index.
|
||||
let mut h: u64 = 0xcbf29ce484222325;
|
||||
for &b in server.server_label.as_bytes() {
|
||||
h ^= b as u64;
|
||||
h = h.wrapping_mul(0x100000001b3);
|
||||
}
|
||||
h ^= conn_idx as u64;
|
||||
mix_u64(h)
|
||||
}
|
||||
|
||||
fn compute_startup_stagger(conn_idx: usize, salt: u64) -> Duration {
|
||||
if conn_idx == 0 {
|
||||
return Duration::ZERO;
|
||||
}
|
||||
let base = (conn_idx as u64).saturating_mul(STARTUP_STAGGER_STEP_MS);
|
||||
let jitter = mix_u64(salt) % 301; // 0..=300ms
|
||||
Duration::from_millis((base + jitter).min(MAX_STARTUP_STAGGER_MS))
|
||||
}
|
||||
|
||||
fn compute_reconnect_delay(
|
||||
base_ms: u64,
|
||||
max_ms: u64,
|
||||
consecutive_failures: u32,
|
||||
salt: u64,
|
||||
) -> Duration {
|
||||
// First retry should be immediate to maximize recovery speed on transient
|
||||
// blips (the user's primary expectation in poor networks).
|
||||
if consecutive_failures <= 1 {
|
||||
return Duration::ZERO;
|
||||
}
|
||||
|
||||
// Keep a sane minimum for repeated failures.
|
||||
let base_ms = base_ms.max(MIN_RECONNECT_DELAY_MS);
|
||||
let max_ms = max_ms.max(base_ms);
|
||||
let cap_ms = compute_reconnect_cap_ms(base_ms, max_ms, consecutive_failures)
|
||||
.min(RECONNECT_PROBE_MAX_DELAY_MS.max(base_ms));
|
||||
|
||||
// Equal-jitter: randomize in [cap/2, cap], preventing synchronized reconnect
|
||||
// storms while keeping reconnect latency bounded.
|
||||
if cap_ms <= 1 {
|
||||
return Duration::from_millis(cap_ms);
|
||||
}
|
||||
|
||||
let half = cap_ms / 2;
|
||||
let span = cap_ms - half;
|
||||
let now_nanos = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|d| d.subsec_nanos() as u64)
|
||||
.unwrap_or(0);
|
||||
let mixed = mix_u64(now_nanos ^ salt);
|
||||
let jitter = if span == 0 { 0 } else { mixed % (span + 1) };
|
||||
Duration::from_millis(half + jitter)
|
||||
}
|
||||
|
||||
fn compute_reconnect_cap_ms(base_ms: u64, max_ms: u64, consecutive_failures: u32) -> u64 {
|
||||
if consecutive_failures <= 1 {
|
||||
return base_ms.min(max_ms);
|
||||
}
|
||||
|
||||
let shift = (consecutive_failures - 1).min(31);
|
||||
let factor = 1u64 << shift;
|
||||
base_ms.saturating_mul(factor).min(max_ms)
|
||||
}
|
||||
|
||||
fn mix_u64(mut x: u64) -> u64 {
|
||||
// SplitMix64 finalizer - cheap bit mixing for pseudo-random jitter.
|
||||
x ^= x >> 30;
|
||||
x = x.wrapping_mul(0xbf58476d1ce4e5b9);
|
||||
x ^= x >> 27;
|
||||
x = x.wrapping_mul(0x94d049bb133111eb);
|
||||
x ^ (x >> 31)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
|
||||
use super::{
|
||||
compute_reconnect_cap_ms, compute_reconnect_delay, compute_startup_stagger,
|
||||
MAX_STARTUP_STAGGER_MS, RECONNECT_PROBE_MAX_DELAY_MS, STARTUP_STAGGER_STEP_MS,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn reconnect_cap_grows_exponentially_and_caps() {
|
||||
let base = 500;
|
||||
let max = 30_000;
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 0), 500);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 1), 500);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 2), 1_000);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 3), 2_000);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 4), 4_000);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 5), 8_000);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 6), 16_000);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 7), 30_000);
|
||||
assert_eq!(compute_reconnect_cap_ms(base, max, 20), 30_000);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn startup_stagger_is_zero_for_primary_and_bounded_for_secondary() {
|
||||
assert_eq!(compute_startup_stagger(0, 42), Duration::ZERO);
|
||||
|
||||
let d1 = compute_startup_stagger(1, 42);
|
||||
let d2 = compute_startup_stagger(2, 42);
|
||||
|
||||
assert!(d1 >= Duration::from_millis(STARTUP_STAGGER_STEP_MS));
|
||||
assert!(d1 <= Duration::from_millis(MAX_STARTUP_STAGGER_MS));
|
||||
assert!(d2 >= Duration::from_millis(STARTUP_STAGGER_STEP_MS * 2));
|
||||
assert!(d2 <= Duration::from_millis(MAX_STARTUP_STAGGER_MS));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reconnect_delay_is_immediate_on_first_failure() {
|
||||
assert_eq!(compute_reconnect_delay(700, 45_000, 1, 123), Duration::ZERO);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reconnect_delay_stays_within_probe_ceiling_after_many_failures() {
|
||||
let d = compute_reconnect_delay(500, 45_000, 100, 12345);
|
||||
assert!(d <= Duration::from_millis(RECONNECT_PROBE_MAX_DELAY_MS));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,260 @@
|
||||
//! Binary frame protocol for WebSocket tunnel multiplexing.
|
||||
//!
|
||||
//! Frame layout (10-byte header + variable payload):
|
||||
//! ```text
|
||||
//! | stream_id (4B) | msg_type (1B) | flags (1B) | payload_len (4B) | payload (NB) |
|
||||
//! ```
|
||||
|
||||
use bytes::{Buf, BufMut, Bytes, BytesMut};
|
||||
|
||||
pub const HEADER_SIZE: usize = 10;
|
||||
|
||||
/// Frame flags.
|
||||
pub mod flags {
|
||||
pub const END_STREAM: u8 = 0x01;
|
||||
pub const GZIP_COMPRESSED: u8 = 0x02;
|
||||
}
|
||||
|
||||
/// Message types for the tunnel protocol.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
#[repr(u8)]
|
||||
pub enum MsgType {
|
||||
RequestHeaders = 0x01,
|
||||
RequestBody = 0x02,
|
||||
ResponseHeaders = 0x03,
|
||||
ResponseBody = 0x04,
|
||||
StreamEnd = 0x05,
|
||||
StreamError = 0x06,
|
||||
Ping = 0x10,
|
||||
Pong = 0x11,
|
||||
GoAway = 0x12,
|
||||
HeartbeatData = 0x13,
|
||||
HeartbeatAck = 0x14,
|
||||
}
|
||||
|
||||
impl MsgType {
|
||||
pub fn from_u8(v: u8) -> Option<Self> {
|
||||
match v {
|
||||
0x01 => Some(Self::RequestHeaders),
|
||||
0x02 => Some(Self::RequestBody),
|
||||
0x03 => Some(Self::ResponseHeaders),
|
||||
0x04 => Some(Self::ResponseBody),
|
||||
0x05 => Some(Self::StreamEnd),
|
||||
0x06 => Some(Self::StreamError),
|
||||
0x10 => Some(Self::Ping),
|
||||
0x11 => Some(Self::Pong),
|
||||
0x12 => Some(Self::GoAway),
|
||||
0x13 => Some(Self::HeartbeatData),
|
||||
0x14 => Some(Self::HeartbeatAck),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A single multiplexed frame.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Frame {
|
||||
pub stream_id: u32,
|
||||
pub msg_type: MsgType,
|
||||
pub flags: u8,
|
||||
pub payload: Bytes,
|
||||
}
|
||||
|
||||
impl Frame {
|
||||
pub fn new(stream_id: u32, msg_type: MsgType, flags: u8, payload: impl Into<Bytes>) -> Self {
|
||||
Self {
|
||||
stream_id,
|
||||
msg_type,
|
||||
flags,
|
||||
payload: payload.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Control frame (stream_id = 0).
|
||||
pub fn control(msg_type: MsgType, payload: impl Into<Bytes>) -> Self {
|
||||
Self::new(0, msg_type, 0, payload)
|
||||
}
|
||||
|
||||
pub fn is_end_stream(&self) -> bool {
|
||||
self.flags & flags::END_STREAM != 0
|
||||
}
|
||||
|
||||
pub fn is_gzip(&self) -> bool {
|
||||
self.flags & flags::GZIP_COMPRESSED != 0
|
||||
}
|
||||
|
||||
/// Encode into a binary buffer.
|
||||
pub fn encode(&self) -> Bytes {
|
||||
let mut buf = BytesMut::with_capacity(HEADER_SIZE + self.payload.len());
|
||||
buf.put_u32(self.stream_id);
|
||||
buf.put_u8(self.msg_type as u8);
|
||||
buf.put_u8(self.flags);
|
||||
buf.put_u32(self.payload.len() as u32);
|
||||
buf.put(self.payload.clone());
|
||||
buf.freeze()
|
||||
}
|
||||
|
||||
/// Decode from a binary buffer.
|
||||
pub fn decode(mut data: Bytes) -> Result<Self, ProtocolError> {
|
||||
if data.len() < HEADER_SIZE {
|
||||
return Err(ProtocolError::TooShort {
|
||||
expected: HEADER_SIZE,
|
||||
actual: data.len(),
|
||||
});
|
||||
}
|
||||
let stream_id = data.get_u32();
|
||||
let msg_type_raw = data.get_u8();
|
||||
let frame_flags = data.get_u8();
|
||||
let payload_len = data.get_u32() as usize;
|
||||
|
||||
if data.remaining() < payload_len {
|
||||
return Err(ProtocolError::Incomplete {
|
||||
expected: HEADER_SIZE + payload_len,
|
||||
actual: HEADER_SIZE + data.remaining(),
|
||||
});
|
||||
}
|
||||
|
||||
let msg_type =
|
||||
MsgType::from_u8(msg_type_raw).ok_or(ProtocolError::UnknownMsgType(msg_type_raw))?;
|
||||
let payload = data.split_to(payload_len);
|
||||
|
||||
Ok(Self {
|
||||
stream_id,
|
||||
msg_type,
|
||||
flags: frame_flags,
|
||||
payload,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Protocol errors.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum ProtocolError {
|
||||
#[error("frame too short: expected {expected} bytes, got {actual}")]
|
||||
TooShort { expected: usize, actual: usize },
|
||||
#[error("frame incomplete: expected {expected} bytes, got {actual}")]
|
||||
Incomplete { expected: usize, actual: usize },
|
||||
#[error("unknown message type: 0x{0:02x}")]
|
||||
UnknownMsgType(u8),
|
||||
}
|
||||
|
||||
/// JSON payload for REQUEST_HEADERS frames.
|
||||
#[derive(Debug, serde::Deserialize)]
|
||||
pub struct RequestMeta {
|
||||
pub method: String,
|
||||
pub url: String,
|
||||
pub headers: std::collections::HashMap<String, String>,
|
||||
#[serde(default = "default_timeout", deserialize_with = "deserialize_timeout")]
|
||||
pub timeout: u64,
|
||||
}
|
||||
|
||||
fn default_timeout() -> u64 {
|
||||
60
|
||||
}
|
||||
|
||||
fn deserialize_timeout<'de, D>(deserializer: D) -> Result<u64, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
#[derive(serde::Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum TimeoutValue {
|
||||
Int(u64),
|
||||
Float(f64),
|
||||
}
|
||||
|
||||
match <TimeoutValue as serde::Deserialize>::deserialize(deserializer)? {
|
||||
TimeoutValue::Int(v) => Ok(v),
|
||||
TimeoutValue::Float(v) => {
|
||||
if !v.is_finite() || v < 0.0 {
|
||||
return Err(serde::de::Error::custom(
|
||||
"timeout must be a non-negative finite number",
|
||||
));
|
||||
}
|
||||
if v.fract() != 0.0 {
|
||||
return Err(serde::de::Error::custom("timeout must be integer seconds"));
|
||||
}
|
||||
if v > (u64::MAX as f64) {
|
||||
return Err(serde::de::Error::custom("timeout is too large"));
|
||||
}
|
||||
Ok(v as u64)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// JSON payload for RESPONSE_HEADERS frames.
|
||||
#[derive(Debug, serde::Serialize)]
|
||||
pub struct ResponseMeta {
|
||||
pub status: u16,
|
||||
/// Header list preserving duplicates (e.g. multiple Set-Cookie).
|
||||
pub headers: Vec<(String, String)>,
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Tunnel frame compression helpers
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
/// Minimum payload size to attempt gzip compression (bytes).
|
||||
const COMPRESS_MIN_SIZE: usize = 512;
|
||||
|
||||
/// If the frame has the GZIP_COMPRESSED flag, decompress the payload; otherwise
|
||||
/// return a clone of the raw payload bytes.
|
||||
pub fn decompress_if_gzip(frame: &Frame) -> Result<Bytes, std::io::Error> {
|
||||
if frame.is_gzip() {
|
||||
decompress_gzip(&frame.payload)
|
||||
} else {
|
||||
Ok(frame.payload.clone())
|
||||
}
|
||||
}
|
||||
|
||||
/// Gzip-compress `data` if it is large enough and compression actually shrinks
|
||||
/// the payload. Returns `(payload, extra_flags)` where `extra_flags` contains
|
||||
/// `GZIP_COMPRESSED` when compression was applied.
|
||||
pub fn compress_payload(data: Bytes) -> (Bytes, u8) {
|
||||
if data.len() >= COMPRESS_MIN_SIZE {
|
||||
if let Ok(compressed) = compress_gzip(&data) {
|
||||
if compressed.len() < data.len() {
|
||||
return (compressed, flags::GZIP_COMPRESSED);
|
||||
}
|
||||
}
|
||||
}
|
||||
(data, 0)
|
||||
}
|
||||
|
||||
fn decompress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
|
||||
use flate2::read::GzDecoder;
|
||||
use std::io::Read;
|
||||
let mut decoder = GzDecoder::new(data);
|
||||
let mut buf = Vec::new();
|
||||
decoder.read_to_end(&mut buf)?;
|
||||
Ok(Bytes::from(buf))
|
||||
}
|
||||
|
||||
fn compress_gzip(data: &[u8]) -> Result<Bytes, std::io::Error> {
|
||||
use flate2::write::GzEncoder;
|
||||
use flate2::Compression;
|
||||
use std::io::Write;
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::fast());
|
||||
encoder.write_all(data)?;
|
||||
let compressed = encoder.finish()?;
|
||||
Ok(Bytes::from(compressed))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::RequestMeta;
|
||||
|
||||
#[test]
|
||||
fn request_meta_accepts_integer_timeout() {
|
||||
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15}"#;
|
||||
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
|
||||
assert_eq!(meta.timeout, 15);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_meta_accepts_integer_like_float_timeout() {
|
||||
let raw = br#"{"method":"GET","url":"https://example.com","headers":{},"timeout":15.0}"#;
|
||||
let meta: RequestMeta = serde_json::from_slice(raw).expect("parse request meta");
|
||||
assert_eq!(meta.timeout, 15);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,506 @@
|
||||
//! Per-stream request handler.
|
||||
//!
|
||||
//! Receives request frames, executes the upstream HTTP request,
|
||||
//! and sends response frames back through the writer channel.
|
||||
|
||||
use std::io;
|
||||
use std::sync::atomic::AtomicUsize;
|
||||
use std::sync::atomic::Ordering;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::stream;
|
||||
use futures_util::StreamExt;
|
||||
use http_body_util::BodyExt;
|
||||
use hyper::body::Frame as BodyFrame;
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::state::{AppState, ServerContext};
|
||||
use crate::target_filter;
|
||||
use crate::upstream_client;
|
||||
|
||||
use super::protocol::{
|
||||
compress_payload, decompress_if_gzip, flags, Frame as TunnelFrame, MsgType, RequestMeta,
|
||||
ResponseMeta,
|
||||
};
|
||||
use super::writer::FrameSender;
|
||||
|
||||
/// Maximum response body chunk size per frame (32 KB).
|
||||
const MAX_CHUNK_SIZE: usize = 32 * 1024;
|
||||
|
||||
/// Timeout for sending a single frame to the writer channel.
|
||||
/// If the writer is congested (TCP backpressure), we abandon the stream
|
||||
/// rather than blocking indefinitely and exhausting the stream pool.
|
||||
const FRAME_SEND_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
/// Minimum allowed upstream request timeout (seconds).
|
||||
const MIN_TIMEOUT_SECS: u64 = 5;
|
||||
/// Maximum allowed upstream request timeout (seconds).
|
||||
const MAX_TIMEOUT_SECS: u64 = 300;
|
||||
|
||||
/// Headers that must not be forwarded to upstream (hop-by-hop or security-sensitive).
|
||||
///
|
||||
/// `host` and `content-length` are managed by the HTTP client (reqwest/hyper):
|
||||
/// - `host` → translated to `:authority` pseudo-header in HTTP/2; forwarding
|
||||
/// the original `host` alongside `:authority` triggers PROTOCOL_ERROR on
|
||||
/// strict H2 implementations (e.g. Google APIs).
|
||||
/// - `content-length` → recalculated by hyper from the actual body; a stale
|
||||
/// value from the tunnel (body may have been re-compressed) causes H2
|
||||
/// PROTOCOL_ERROR when it mismatches the real frame length.
|
||||
const BLOCKED_HEADERS: &[&str] = &[
|
||||
"connection",
|
||||
"content-length",
|
||||
"host",
|
||||
"keep-alive",
|
||||
"proxy-authenticate",
|
||||
"proxy-authorization",
|
||||
"proxy-connection",
|
||||
"te",
|
||||
"trailer",
|
||||
"transfer-encoding",
|
||||
"upgrade",
|
||||
];
|
||||
|
||||
/// Handle a single stream: receive body, execute upstream, send response.
|
||||
pub async fn handle_stream(
|
||||
state: Arc<AppState>,
|
||||
server: Arc<ServerContext>,
|
||||
stream_id: u32,
|
||||
meta: RequestMeta,
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
frame_tx: FrameSender,
|
||||
) {
|
||||
server.active_connections.fetch_add(1, Ordering::Release);
|
||||
|
||||
let connect_elapsed =
|
||||
handle_stream_inner(&state, &server, stream_id, meta, body_rx, &frame_tx).await;
|
||||
|
||||
server.active_connections.fetch_sub(1, Ordering::Release);
|
||||
if let Some(d) = connect_elapsed {
|
||||
server.metrics.record_request(d);
|
||||
}
|
||||
}
|
||||
|
||||
/// Send a frame to the writer with a timeout. Returns false if send failed.
|
||||
async fn send_frame(tx: &FrameSender, frame: TunnelFrame) -> bool {
|
||||
match tokio::time::timeout(FRAME_SEND_TIMEOUT, tx.send(frame)).await {
|
||||
Ok(Ok(())) => true,
|
||||
Ok(Err(_)) => {
|
||||
// Channel closed (writer exited)
|
||||
false
|
||||
}
|
||||
Err(_) => {
|
||||
// Timeout — writer is congested
|
||||
warn!("frame send timeout (writer congested), abandoning stream");
|
||||
false
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns the connection-establishment duration (DNS + TCP/TLS + TTFB) if the
|
||||
/// upstream request succeeded, or `None` if the request never reached the
|
||||
/// response-headers stage.
|
||||
async fn handle_stream_inner(
|
||||
state: &AppState,
|
||||
server: &ServerContext,
|
||||
stream_id: u32,
|
||||
meta: RequestMeta,
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
frame_tx: &FrameSender,
|
||||
) -> Option<Duration> {
|
||||
// Validate target
|
||||
let target_url = match url::Url::parse(&meta.url) {
|
||||
Ok(u) => u,
|
||||
Err(e) => {
|
||||
send_error(frame_tx, stream_id, &format!("invalid URL: {e}")).await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
// Only allow http/https schemes (block file://, data://, etc.)
|
||||
match target_url.scheme() {
|
||||
"http" | "https" => {}
|
||||
other => {
|
||||
send_error(
|
||||
frame_tx,
|
||||
stream_id,
|
||||
&format!("unsupported URL scheme: {other}"),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
}
|
||||
}
|
||||
|
||||
let host = match target_url.host_str() {
|
||||
Some(h) => h.to_string(),
|
||||
None => {
|
||||
send_error(frame_tx, stream_id, "missing host in URL").await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
let port = target_url.port_or_known_default().unwrap_or(443);
|
||||
|
||||
// DNS + target validation (populates dns_cache for SafeDnsResolver)
|
||||
let connect_start = Instant::now();
|
||||
{
|
||||
let allowed_ports = Arc::clone(&server.dynamic.load().allowed_ports);
|
||||
if let Err(e) =
|
||||
target_filter::validate_target(&host, port, &allowed_ports, &state.dns_cache).await
|
||||
{
|
||||
server.metrics.dns_failures.fetch_add(1, Ordering::Release);
|
||||
send_error(frame_tx, stream_id, &format!("target blocked: {e}")).await;
|
||||
return None;
|
||||
}
|
||||
}
|
||||
let dns_ms = connect_start.elapsed().as_millis() as u64;
|
||||
|
||||
// Execute upstream request
|
||||
let client = &state.upstream_client;
|
||||
let timeout = Duration::from_secs(meta.timeout.clamp(MIN_TIMEOUT_SECS, MAX_TIMEOUT_SECS));
|
||||
let request_body_size = Arc::new(AtomicUsize::new(0));
|
||||
let request_body = build_streaming_request_body(body_rx, Arc::clone(&request_body_size));
|
||||
|
||||
let method: hyper::Method = meta.method.parse().unwrap_or(hyper::Method::GET);
|
||||
let mut request = match hyper::Request::builder()
|
||||
.method(method)
|
||||
.uri(meta.url.as_str())
|
||||
.body(request_body)
|
||||
{
|
||||
Ok(request) => request,
|
||||
Err(e) => {
|
||||
send_error(
|
||||
frame_tx,
|
||||
stream_id,
|
||||
&format!("invalid upstream request: {e}"),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
let headers = request.headers_mut();
|
||||
for (k, v) in &meta.headers {
|
||||
let k_lower = k.to_ascii_lowercase();
|
||||
if BLOCKED_HEADERS.contains(&k_lower.as_str()) {
|
||||
continue;
|
||||
}
|
||||
if let (Ok(name), Ok(value)) = (
|
||||
hyper::header::HeaderName::from_bytes(k.as_bytes()),
|
||||
hyper::header::HeaderValue::from_str(v),
|
||||
) {
|
||||
headers.insert(name, value);
|
||||
}
|
||||
}
|
||||
|
||||
let mut captured_connection = upstream_client::capture_connection(&mut request);
|
||||
let connection_start = Instant::now();
|
||||
let connection_capture = tokio::spawn(async move {
|
||||
let connected = captured_connection.wait_for_connection_metadata().await;
|
||||
connected
|
||||
.as_ref()
|
||||
.map(|_| connection_start.elapsed().as_millis() as u64)
|
||||
});
|
||||
|
||||
let upstream_start = Instant::now();
|
||||
let response = match tokio::time::timeout(timeout, client.request(request)).await {
|
||||
Ok(Ok(response)) => response,
|
||||
Ok(Err(e)) => {
|
||||
connection_capture.abort();
|
||||
server
|
||||
.metrics
|
||||
.failed_requests
|
||||
.fetch_add(1, Ordering::Release);
|
||||
let msg = if e.is_connect() {
|
||||
format!("upstream connect error: {e}")
|
||||
} else {
|
||||
format!("upstream error: {e}")
|
||||
};
|
||||
send_error(frame_tx, stream_id, &msg).await;
|
||||
return None;
|
||||
}
|
||||
Err(_) => {
|
||||
connection_capture.abort();
|
||||
server
|
||||
.metrics
|
||||
.failed_requests
|
||||
.fetch_add(1, Ordering::Release);
|
||||
send_error(frame_tx, stream_id, "upstream timeout").await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
// Capture connection-establishment duration (DNS + TCP/TLS + TTFB)
|
||||
// before proceeding to stream the response body.
|
||||
let connect_elapsed = connect_start.elapsed();
|
||||
|
||||
// Send RESPONSE_HEADERS
|
||||
let status = response.status().as_u16();
|
||||
let ttfb_ms = upstream_start.elapsed().as_millis() as u64;
|
||||
// Short timeout: on connection reuse hyper may never fire the connect
|
||||
// callback, so avoid blocking indefinitely.
|
||||
let connection_acquire_ms =
|
||||
match tokio::time::timeout(Duration::from_millis(100), connection_capture).await {
|
||||
Ok(Ok(ms)) => ms,
|
||||
Ok(Err(_)) => None, // JoinError (task panicked / cancelled)
|
||||
Err(_) => None, // timeout -- task is detached but lightweight
|
||||
};
|
||||
let request_timing =
|
||||
upstream_client::resolve_request_timing(&response, connection_acquire_ms, ttfb_ms);
|
||||
let mut resp_headers: Vec<(String, String)> = Vec::with_capacity(response.headers().len() + 1);
|
||||
for (k, v) in response.headers() {
|
||||
if let Ok(vs) = v.to_str() {
|
||||
resp_headers.push((k.as_str().to_string(), vs.to_string()));
|
||||
}
|
||||
}
|
||||
let timing = serde_json::json!({
|
||||
"dns_ms": dns_ms,
|
||||
"connection_acquire_ms": request_timing.connection_acquire_ms,
|
||||
"connection_reused": request_timing.connection_reused,
|
||||
"connect_ms": request_timing.connect_ms,
|
||||
"tls_ms": request_timing.tls_ms,
|
||||
"ttfb_ms": ttfb_ms,
|
||||
"upstream_ms": ttfb_ms,
|
||||
"response_wait_ms": request_timing.response_wait_ms,
|
||||
"upstream_processing_ms": request_timing.response_wait_ms,
|
||||
"timing_source": "instrumented_connector",
|
||||
"total_ms": connect_elapsed.as_millis() as u64,
|
||||
"body_size": request_body_size.load(Ordering::Relaxed),
|
||||
"mode": "tunnel",
|
||||
});
|
||||
resp_headers.push(("x-proxy-timing".to_string(), timing.to_string()));
|
||||
let resp_meta = ResponseMeta {
|
||||
status,
|
||||
headers: resp_headers,
|
||||
};
|
||||
let meta_json: Bytes = serde_json::to_vec(&resp_meta).unwrap_or_default().into();
|
||||
let (meta_payload, meta_flags) = compress_payload(meta_json);
|
||||
if !send_frame(
|
||||
frame_tx,
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::ResponseHeaders,
|
||||
meta_flags,
|
||||
meta_payload,
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Some(connect_elapsed);
|
||||
}
|
||||
|
||||
// Stream response body — relay upstream bytes through the tunnel.
|
||||
// Apply tunnel-level frame compression for chunks that benefit from it
|
||||
// (e.g. uncompressed SSE text). Already-compressed data (gzip/br from
|
||||
// upstream Content-Encoding) won't shrink further and will be sent as-is
|
||||
// thanks to the size check in compress_payload().
|
||||
let mut stream = response.into_body().into_data_stream();
|
||||
while let Some(chunk_result) = stream.next().await {
|
||||
match chunk_result {
|
||||
Ok(chunk) => {
|
||||
if chunk.len() <= MAX_CHUNK_SIZE {
|
||||
let (payload, extra_flags) = compress_payload(chunk);
|
||||
if !send_frame(
|
||||
frame_tx,
|
||||
TunnelFrame::new(stream_id, MsgType::ResponseBody, extra_flags, payload),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Some(connect_elapsed);
|
||||
}
|
||||
} else {
|
||||
// Split oversized chunks, compress each slice
|
||||
let mut offset = 0;
|
||||
while offset < chunk.len() {
|
||||
let end = (offset + MAX_CHUNK_SIZE).min(chunk.len());
|
||||
let slice = chunk.slice(offset..end);
|
||||
let (payload, extra_flags) = compress_payload(slice);
|
||||
if !send_frame(
|
||||
frame_tx,
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::ResponseBody,
|
||||
extra_flags,
|
||||
payload,
|
||||
),
|
||||
)
|
||||
.await
|
||||
{
|
||||
return Some(connect_elapsed);
|
||||
}
|
||||
offset = end;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
server.metrics.stream_errors.fetch_add(1, Ordering::Release);
|
||||
warn!(stream_id, error = %e, "upstream body read error");
|
||||
send_error(frame_tx, stream_id, &format!("body read error: {e}")).await;
|
||||
return Some(connect_elapsed);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Send STREAM_END
|
||||
let _ = send_frame(
|
||||
frame_tx,
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::StreamEnd,
|
||||
flags::END_STREAM,
|
||||
Bytes::new(),
|
||||
),
|
||||
)
|
||||
.await;
|
||||
|
||||
debug!(stream_id, status, "stream completed");
|
||||
Some(connect_elapsed)
|
||||
}
|
||||
|
||||
async fn send_error(tx: &FrameSender, stream_id: u32, msg: &str) {
|
||||
// Error frames use best-effort delivery — don't block if writer is congested
|
||||
let _ = send_frame(
|
||||
tx,
|
||||
TunnelFrame::new(
|
||||
stream_id,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from(msg.to_string()),
|
||||
),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
fn build_streaming_request_body(
|
||||
body_rx: mpsc::Receiver<TunnelFrame>,
|
||||
body_size: Arc<AtomicUsize>,
|
||||
) -> upstream_client::UpstreamRequestBody {
|
||||
let body_stream = stream::unfold(
|
||||
(body_rx, body_size, false),
|
||||
|(mut body_rx, body_size, finished)| async move {
|
||||
if finished {
|
||||
return None;
|
||||
}
|
||||
|
||||
loop {
|
||||
let frame = match body_rx.recv().await {
|
||||
Some(frame) => frame,
|
||||
None => return None,
|
||||
};
|
||||
|
||||
match frame.msg_type {
|
||||
MsgType::RequestBody => {
|
||||
let end_stream = frame.is_end_stream();
|
||||
let payload = match decompress_if_gzip(&frame) {
|
||||
Ok(payload) => payload,
|
||||
Err(error) => {
|
||||
let err =
|
||||
io::Error::other(format!("gzip decompress failed: {error}"));
|
||||
return Some((Err(err), (body_rx, body_size, true)));
|
||||
}
|
||||
};
|
||||
|
||||
if payload.is_empty() {
|
||||
if end_stream {
|
||||
return None;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
body_size.fetch_add(payload.len(), Ordering::Relaxed);
|
||||
return Some((
|
||||
Ok(BodyFrame::data(payload)),
|
||||
(body_rx, body_size, end_stream),
|
||||
));
|
||||
}
|
||||
MsgType::StreamError => {
|
||||
let message = String::from_utf8(frame.payload.to_vec())
|
||||
.unwrap_or_else(|_| "client cancelled request body".to_string());
|
||||
return Some((Err(io::Error::other(message)), (body_rx, body_size, true)));
|
||||
}
|
||||
MsgType::StreamEnd => return None,
|
||||
_ => continue,
|
||||
}
|
||||
}
|
||||
},
|
||||
);
|
||||
|
||||
upstream_client::stream_request_body(body_stream)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_request_body_yields_chunks_and_tracks_size() {
|
||||
let (tx, rx) = mpsc::channel(4);
|
||||
let body_size = Arc::new(AtomicUsize::new(0));
|
||||
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
|
||||
|
||||
tx.send(TunnelFrame::new(
|
||||
1,
|
||||
MsgType::RequestBody,
|
||||
0,
|
||||
Bytes::from_static(b"abc"),
|
||||
))
|
||||
.await
|
||||
.expect("send first chunk");
|
||||
tx.send(TunnelFrame::new(
|
||||
1,
|
||||
MsgType::RequestBody,
|
||||
flags::END_STREAM,
|
||||
Bytes::from_static(b"def"),
|
||||
))
|
||||
.await
|
||||
.expect("send final chunk");
|
||||
drop(tx);
|
||||
|
||||
let first = body
|
||||
.frame()
|
||||
.await
|
||||
.expect("first frame")
|
||||
.expect("first frame ok")
|
||||
.into_data()
|
||||
.expect("first data frame");
|
||||
let second = body
|
||||
.frame()
|
||||
.await
|
||||
.expect("second frame")
|
||||
.expect("second frame ok")
|
||||
.into_data()
|
||||
.expect("second data frame");
|
||||
|
||||
assert_eq!(first, Bytes::from_static(b"abc"));
|
||||
assert_eq!(second, Bytes::from_static(b"def"));
|
||||
assert!(body.frame().await.is_none());
|
||||
assert_eq!(body_size.load(Ordering::Relaxed), 6);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn streaming_request_body_surfaces_client_cancel_as_error() {
|
||||
let (tx, rx) = mpsc::channel(4);
|
||||
let body_size = Arc::new(AtomicUsize::new(0));
|
||||
let mut body = build_streaming_request_body(rx, Arc::clone(&body_size));
|
||||
|
||||
tx.send(TunnelFrame::new(
|
||||
1,
|
||||
MsgType::StreamError,
|
||||
0,
|
||||
Bytes::from_static(b"client cancelled"),
|
||||
))
|
||||
.await
|
||||
.expect("send cancel frame");
|
||||
drop(tx);
|
||||
|
||||
let err = body
|
||||
.frame()
|
||||
.await
|
||||
.expect("error frame present")
|
||||
.expect_err("body should surface cancellation error");
|
||||
assert!(err.to_string().contains("client cancelled"));
|
||||
assert!(body.frame().await.is_none());
|
||||
assert_eq!(body_size.load(Ordering::Relaxed), 0);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
//! Dedicated WebSocket writer task.
|
||||
//!
|
||||
//! All frame writes go through an mpsc channel to a single writer task,
|
||||
//! avoiding contention on the WebSocket sink. The writer also sends
|
||||
//! periodic WebSocket Ping frames to keep the connection alive through
|
||||
//! intermediary proxies (Nginx, Cloudflare, etc.).
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::SinkExt;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio_tungstenite::tungstenite::Message;
|
||||
use tracing::{debug, error, trace};
|
||||
|
||||
use super::protocol::Frame;
|
||||
|
||||
/// Sender half — cloned by stream handlers and heartbeat.
|
||||
pub type FrameSender = mpsc::Sender<Frame>;
|
||||
|
||||
/// Spawn the writer task. Returns the sender and a JoinHandle for cleanup.
|
||||
///
|
||||
/// `ping_interval` controls WebSocket-level Ping frequency (typically 15s).
|
||||
/// This keeps the connection alive through intermediary proxies/load-balancers.
|
||||
pub fn spawn_writer<S>(mut sink: S, ping_interval: Duration) -> (FrameSender, JoinHandle<()>)
|
||||
where
|
||||
S: SinkExt<Message, Error = tokio_tungstenite::tungstenite::Error> + Unpin + Send + 'static,
|
||||
{
|
||||
let (tx, mut rx) = mpsc::channel::<Frame>(256);
|
||||
|
||||
let handle = tokio::spawn(async move {
|
||||
let mut ping_ticker = tokio::time::interval(ping_interval);
|
||||
ping_ticker.tick().await; // skip first immediate tick
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
frame = rx.recv() => {
|
||||
match frame {
|
||||
Some(frame) => {
|
||||
let data = frame.encode();
|
||||
if let Err(e) = sink.send(Message::Binary(data.into())).await {
|
||||
error!(error = %e, "failed to write frame to WebSocket");
|
||||
break;
|
||||
}
|
||||
}
|
||||
None => break, // all senders dropped
|
||||
}
|
||||
}
|
||||
_ = ping_ticker.tick() => {
|
||||
if let Err(e) = sink.send(Message::Ping(vec![])).await {
|
||||
error!(error = %e, "failed to send WebSocket ping");
|
||||
break;
|
||||
}
|
||||
trace!("sent WebSocket ping");
|
||||
}
|
||||
}
|
||||
}
|
||||
debug!("writer task exiting");
|
||||
let _ = sink.close().await;
|
||||
});
|
||||
|
||||
(tx, handle)
|
||||
}
|
||||
@@ -0,0 +1,444 @@
|
||||
use std::future::Future;
|
||||
use std::io;
|
||||
use std::net::IpAddr;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
use std::time::Duration;
|
||||
|
||||
use bytes::Bytes;
|
||||
use futures_util::Stream;
|
||||
use http_body_util::combinators::UnsyncBoxBody;
|
||||
use http_body_util::{BodyExt, StreamBody};
|
||||
use hyper::body::Frame;
|
||||
use hyper::rt;
|
||||
use hyper::Response;
|
||||
use hyper::Uri;
|
||||
pub use hyper_util::client::legacy::connect::capture_connection;
|
||||
use hyper_util::client::legacy::connect::dns::Name;
|
||||
use hyper_util::client::legacy::connect::{Connected, Connection, HttpConnector};
|
||||
use hyper_util::client::legacy::Client;
|
||||
use hyper_util::rt::{TokioExecutor, TokioIo, TokioTimer};
|
||||
use rustls::pki_types::ServerName;
|
||||
use rustls::ClientConfig;
|
||||
use tokio::net::TcpStream;
|
||||
use tokio_rustls::TlsConnector;
|
||||
use tower_service::Service;
|
||||
|
||||
use crate::config::Config;
|
||||
use crate::target_filter::{self, DnsCache};
|
||||
|
||||
type BoxError = Box<dyn std::error::Error + Send + Sync>;
|
||||
|
||||
type PlainStream = TokioIo<TcpStream>;
|
||||
type TlsStream = TokioIo<tokio_rustls::client::TlsStream<TcpStream>>;
|
||||
|
||||
pub type UpstreamRequestBody = UnsyncBoxBody<Bytes, io::Error>;
|
||||
pub type UpstreamClient = Client<InstrumentedConnector, UpstreamRequestBody>;
|
||||
|
||||
pub fn stream_request_body<S>(stream: S) -> UpstreamRequestBody
|
||||
where
|
||||
S: Stream<Item = Result<Frame<Bytes>, io::Error>> + Send + 'static,
|
||||
{
|
||||
StreamBody::new(stream).boxed_unsync()
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct ConnectTiming {
|
||||
pub connect_ms: u64,
|
||||
pub tls_ms: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct RequestTiming {
|
||||
pub connection_acquire_ms: u64,
|
||||
pub connect_ms: u64,
|
||||
pub tls_ms: u64,
|
||||
pub response_wait_ms: u64,
|
||||
pub connection_reused: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct ValidatedResolver {
|
||||
dns_cache: Arc<DnsCache>,
|
||||
}
|
||||
|
||||
impl ValidatedResolver {
|
||||
pub fn new(dns_cache: Arc<DnsCache>) -> Self {
|
||||
Self { dns_cache }
|
||||
}
|
||||
}
|
||||
|
||||
pub struct ValidatedAddrs {
|
||||
inner: std::vec::IntoIter<std::net::SocketAddr>,
|
||||
}
|
||||
|
||||
impl Iterator for ValidatedAddrs {
|
||||
type Item = std::net::SocketAddr;
|
||||
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.inner.next()
|
||||
}
|
||||
}
|
||||
|
||||
impl Service<Name> for ValidatedResolver {
|
||||
type Response = ValidatedAddrs;
|
||||
type Error = io::Error;
|
||||
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
|
||||
|
||||
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||
Poll::Ready(Ok(()))
|
||||
}
|
||||
|
||||
fn call(&mut self, name: Name) -> Self::Future {
|
||||
let dns_cache = Arc::clone(&self.dns_cache);
|
||||
let host = name.as_str().to_string();
|
||||
Box::pin(async move {
|
||||
if let Some(addrs) = dns_cache.get_by_host(&host).await {
|
||||
return Ok(ValidatedAddrs {
|
||||
inner: (*addrs).clone().into_iter(),
|
||||
});
|
||||
}
|
||||
|
||||
let resolved = target_filter::resolve_public_addrs(&host, 0, dns_cache.as_ref())
|
||||
.await
|
||||
.map_err(|err| io::Error::other(err.to_string()))?;
|
||||
Ok(ValidatedAddrs {
|
||||
inner: resolved.into_iter(),
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct InstrumentedConnector {
|
||||
http: HttpConnector<ValidatedResolver>,
|
||||
tls_config: Arc<ClientConfig>,
|
||||
}
|
||||
|
||||
impl Service<Uri> for InstrumentedConnector {
|
||||
type Response = TimedConn;
|
||||
type Error = BoxError;
|
||||
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
|
||||
|
||||
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||
self.http.poll_ready(cx).map_err(Into::into)
|
||||
}
|
||||
|
||||
fn call(&mut self, dst: Uri) -> Self::Future {
|
||||
let scheme = dst.scheme_str().map(|value| value.to_ascii_lowercase());
|
||||
let tls_config = Arc::clone(&self.tls_config);
|
||||
let connecting = self.http.call(dst.clone());
|
||||
let connect_start = std::time::Instant::now();
|
||||
|
||||
Box::pin(async move {
|
||||
match scheme.as_deref() {
|
||||
Some("http") => {
|
||||
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
|
||||
let connect_ms = connect_start.elapsed().as_millis() as u64;
|
||||
Ok(TimedConn::new(
|
||||
MaybeHttpsStream::Http(tcp),
|
||||
ConnectTiming {
|
||||
connect_ms,
|
||||
tls_ms: 0,
|
||||
},
|
||||
))
|
||||
}
|
||||
Some("https") => {
|
||||
let server_name = resolve_server_name(&dst)?;
|
||||
let tcp = connecting.await.map_err(|err| Box::new(err) as BoxError)?;
|
||||
let connect_ms = connect_start.elapsed().as_millis() as u64;
|
||||
|
||||
let tls_start = std::time::Instant::now();
|
||||
let tls_stream = TlsConnector::from(tls_config)
|
||||
.connect(server_name, tcp.into_inner())
|
||||
.await
|
||||
.map_err(io::Error::other)?;
|
||||
let tls_ms = tls_start.elapsed().as_millis() as u64;
|
||||
|
||||
Ok(TimedConn::new(
|
||||
MaybeHttpsStream::Https(TokioIo::new(tls_stream)),
|
||||
ConnectTiming { connect_ms, tls_ms },
|
||||
))
|
||||
}
|
||||
Some(other) => Err(io::Error::other(format!("unsupported scheme {other}")).into()),
|
||||
None => Err(io::Error::other("missing scheme").into()),
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn build_upstream_client(config: &Config, dns_cache: Arc<DnsCache>) -> UpstreamClient {
|
||||
let mut http = HttpConnector::new_with_resolver(ValidatedResolver::new(dns_cache));
|
||||
http.enforce_http(false);
|
||||
http.set_connect_timeout(Some(Duration::from_secs(
|
||||
config.upstream_connect_timeout_secs,
|
||||
)));
|
||||
http.set_nodelay(config.upstream_tcp_nodelay);
|
||||
if config.upstream_tcp_keepalive_secs > 0 {
|
||||
http.set_keepalive(Some(Duration::from_secs(
|
||||
config.upstream_tcp_keepalive_secs,
|
||||
)));
|
||||
} else {
|
||||
http.set_keepalive(None);
|
||||
}
|
||||
|
||||
let connector = InstrumentedConnector {
|
||||
http,
|
||||
tls_config: build_tls_config(),
|
||||
};
|
||||
|
||||
let mut builder = Client::builder(TokioExecutor::new());
|
||||
builder.pool_max_idle_per_host(config.upstream_pool_max_idle_per_host);
|
||||
builder.pool_idle_timeout(Duration::from_secs(config.upstream_pool_idle_timeout_secs));
|
||||
builder.pool_timer(TokioTimer::new());
|
||||
builder.build(connector)
|
||||
}
|
||||
|
||||
pub fn resolve_request_timing<B>(
|
||||
response: &Response<B>,
|
||||
connection_acquire_ms: Option<u64>,
|
||||
ttfb_ms: u64,
|
||||
) -> RequestTiming {
|
||||
let raw = response
|
||||
.extensions()
|
||||
.get::<ConnectTiming>()
|
||||
.copied()
|
||||
.unwrap_or_default();
|
||||
|
||||
let raw_connection_ms = raw.connect_ms.saturating_add(raw.tls_ms);
|
||||
let measured_acquire_ms = connection_acquire_ms.unwrap_or(raw_connection_ms.min(ttfb_ms));
|
||||
let likely_reused = measured_acquire_ms <= 5 && raw_connection_ms > 0;
|
||||
let connector_matches_request = raw_connection_ms <= measured_acquire_ms.saturating_add(25);
|
||||
|
||||
let (connect_ms, tls_ms) = if likely_reused || !connector_matches_request {
|
||||
(0, 0)
|
||||
} else {
|
||||
(raw.connect_ms, raw.tls_ms)
|
||||
};
|
||||
|
||||
RequestTiming {
|
||||
connection_acquire_ms: measured_acquire_ms,
|
||||
connect_ms,
|
||||
tls_ms,
|
||||
response_wait_ms: ttfb_ms.saturating_sub(measured_acquire_ms),
|
||||
connection_reused: likely_reused,
|
||||
}
|
||||
}
|
||||
|
||||
fn build_tls_config() -> Arc<ClientConfig> {
|
||||
let root_store =
|
||||
rustls::RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
|
||||
let mut config = ClientConfig::builder()
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth();
|
||||
config.alpn_protocols = vec![b"h2".to_vec(), b"http/1.1".to_vec()];
|
||||
Arc::new(config)
|
||||
}
|
||||
|
||||
fn resolve_server_name(uri: &Uri) -> Result<ServerName<'static>, BoxError> {
|
||||
let host = uri.host().ok_or_else(|| io::Error::other("missing host"))?;
|
||||
let host = host.trim_start_matches('[').trim_end_matches(']');
|
||||
|
||||
if let Ok(ip) = host.parse::<IpAddr>() {
|
||||
return Ok(ServerName::from(ip));
|
||||
}
|
||||
|
||||
Ok(ServerName::try_from(host.to_string())?)
|
||||
}
|
||||
|
||||
pub struct TimedConn {
|
||||
inner: MaybeHttpsStream,
|
||||
timing: ConnectTiming,
|
||||
}
|
||||
|
||||
impl TimedConn {
|
||||
fn new(inner: MaybeHttpsStream, timing: ConnectTiming) -> Self {
|
||||
Self { inner, timing }
|
||||
}
|
||||
}
|
||||
|
||||
impl Connection for TimedConn {
|
||||
fn connected(&self) -> Connected {
|
||||
self.inner.connected().extra(self.timing)
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Read for TimedConn {
|
||||
fn poll_read(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: rt::ReadBufCursor<'_>,
|
||||
) -> Poll<Result<(), io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_read(cx, buf)
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Write for TimedConn {
|
||||
fn poll_write(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_write(cx, buf)
|
||||
}
|
||||
|
||||
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_flush(cx)
|
||||
}
|
||||
|
||||
fn poll_shutdown(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
) -> Poll<Result<(), io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_shutdown(cx)
|
||||
}
|
||||
|
||||
fn is_write_vectored(&self) -> bool {
|
||||
self.inner.is_write_vectored()
|
||||
}
|
||||
|
||||
fn poll_write_vectored(
|
||||
mut self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
bufs: &[std::io::IoSlice<'_>],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
|
||||
}
|
||||
}
|
||||
|
||||
pub enum MaybeHttpsStream {
|
||||
Http(PlainStream),
|
||||
Https(TlsStream),
|
||||
}
|
||||
|
||||
impl Connection for MaybeHttpsStream {
|
||||
fn connected(&self) -> Connected {
|
||||
match self {
|
||||
Self::Http(stream) => stream.connected(),
|
||||
Self::Https(stream) => {
|
||||
let (tcp, tls) = stream.inner().get_ref();
|
||||
if tls.alpn_protocol() == Some(b"h2") {
|
||||
tcp.connected().negotiated_h2()
|
||||
} else {
|
||||
tcp.connected()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Read for MaybeHttpsStream {
|
||||
fn poll_read(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: rt::ReadBufCursor<'_>,
|
||||
) -> Poll<Result<(), io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_read(cx, buf),
|
||||
Self::Https(stream) => Pin::new(stream).poll_read(cx, buf),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl rt::Write for MaybeHttpsStream {
|
||||
fn poll_write(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
buf: &[u8],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_write(cx, buf),
|
||||
Self::Https(stream) => Pin::new(stream).poll_write(cx, buf),
|
||||
}
|
||||
}
|
||||
|
||||
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_flush(cx),
|
||||
Self::Https(stream) => Pin::new(stream).poll_flush(cx),
|
||||
}
|
||||
}
|
||||
|
||||
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_shutdown(cx),
|
||||
Self::Https(stream) => Pin::new(stream).poll_shutdown(cx),
|
||||
}
|
||||
}
|
||||
|
||||
fn is_write_vectored(&self) -> bool {
|
||||
match self {
|
||||
Self::Http(stream) => stream.is_write_vectored(),
|
||||
Self::Https(stream) => stream.is_write_vectored(),
|
||||
}
|
||||
}
|
||||
|
||||
fn poll_write_vectored(
|
||||
self: Pin<&mut Self>,
|
||||
cx: &mut Context<'_>,
|
||||
bufs: &[std::io::IoSlice<'_>],
|
||||
) -> Poll<Result<usize, io::Error>> {
|
||||
match Pin::get_mut(self) {
|
||||
Self::Http(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
|
||||
Self::Https(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use hyper::Response;
|
||||
|
||||
#[test]
|
||||
fn fresh_connection_uses_connector_breakdown() {
|
||||
let mut response = Response::new(());
|
||||
response.extensions_mut().insert(ConnectTiming {
|
||||
connect_ms: 80,
|
||||
tls_ms: 40,
|
||||
});
|
||||
|
||||
let timing = resolve_request_timing(&response, Some(125), 600);
|
||||
|
||||
assert_eq!(timing.connection_acquire_ms, 125);
|
||||
assert_eq!(timing.connect_ms, 80);
|
||||
assert_eq!(timing.tls_ms, 40);
|
||||
assert_eq!(timing.response_wait_ms, 475);
|
||||
assert!(!timing.connection_reused);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reused_connection_zeroes_stale_connect_timings() {
|
||||
let mut response = Response::new(());
|
||||
response.extensions_mut().insert(ConnectTiming {
|
||||
connect_ms: 70,
|
||||
tls_ms: 30,
|
||||
});
|
||||
|
||||
let timing = resolve_request_timing(&response, Some(0), 310);
|
||||
|
||||
assert_eq!(timing.connection_acquire_ms, 0);
|
||||
assert_eq!(timing.connect_ms, 0);
|
||||
assert_eq!(timing.tls_ms, 0);
|
||||
assert_eq!(timing.response_wait_ms, 310);
|
||||
assert!(timing.connection_reused);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn falls_back_to_connector_timings_when_capture_missing() {
|
||||
let mut response = Response::new(());
|
||||
response.extensions_mut().insert(ConnectTiming {
|
||||
connect_ms: 55,
|
||||
tls_ms: 25,
|
||||
});
|
||||
|
||||
let timing = resolve_request_timing(&response, None, 400);
|
||||
|
||||
assert_eq!(timing.connection_acquire_ms, 80);
|
||||
assert_eq!(timing.connect_ms, 55);
|
||||
assert_eq!(timing.tls_ms, 25);
|
||||
assert_eq!(timing.response_wait_ms, 320);
|
||||
assert!(!timing.connection_reused);
|
||||
}
|
||||
}
|
||||
+32
-12
@@ -3,13 +3,15 @@ Alembic 环境配置
|
||||
用于数据库迁移的运行时环境设置
|
||||
"""
|
||||
|
||||
from logging.config import fileConfig
|
||||
from sqlalchemy import engine_from_config, pool
|
||||
from alembic import context
|
||||
import os
|
||||
import sys
|
||||
from logging.config import fileConfig
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy import engine_from_config, pool, text
|
||||
|
||||
from alembic import context
|
||||
|
||||
# 添加项目根目录到 Python 路径
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(__file__)))
|
||||
|
||||
@@ -30,7 +32,7 @@ from src.models.database import Base
|
||||
config = context.config
|
||||
|
||||
# 从环境变量获取数据库 URL
|
||||
# 优先使用 DATABASE_URL,否则从 DB_PASSWORD 自动构建(与 docker-compose 保持一致)
|
||||
# 优先使用 DATABASE_URL,否则从 DB_PASSWORD 自动构建(与 docker compose 保持一致)
|
||||
database_url = os.getenv("DATABASE_URL")
|
||||
if not database_url:
|
||||
db_password = os.getenv("DB_PASSWORD", "")
|
||||
@@ -48,6 +50,11 @@ if config.config_file_name is not None:
|
||||
# 目标元数据(包含所有表定义)
|
||||
target_metadata = Base.metadata
|
||||
|
||||
# PostgreSQL 全局迁移锁,避免多进程并发执行 Alembic 导致竞态(重复加列/索引等)
|
||||
# 使用会话级 advisory lock(pg_advisory_lock),在迁移完成后手动释放。
|
||||
# ID 由 crc32("aether-alembic-migration") 拼接生成,仅需全局唯一即可。
|
||||
MIGRATION_ADVISORY_LOCK_ID = 582694137405821
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
"""
|
||||
@@ -83,15 +90,28 @@ def run_migrations_online() -> None:
|
||||
)
|
||||
|
||||
with connectable.connect() as connection:
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata,
|
||||
compare_type=True, # 比较列类型变更
|
||||
compare_server_default=True, # 比较默认值变更
|
||||
)
|
||||
try:
|
||||
# 使用会话级 advisory lock(非事务级),避免干扰 Alembic 的事务管理。
|
||||
# pg_advisory_lock 在会话结束时自动释放,不受 COMMIT/ROLLBACK 影响。
|
||||
if connection.dialect.name == "postgresql":
|
||||
connection.execute(
|
||||
text("SELECT pg_advisory_lock(:lock_id)"),
|
||||
{"lock_id": MIGRATION_ADVISORY_LOCK_ID},
|
||||
)
|
||||
connection.commit()
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata,
|
||||
compare_type=True,
|
||||
compare_server_default=True,
|
||||
transaction_per_migration=True, # 每个迁移文件独立事务,完成即提交
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
|
||||
# 根据模式选择运行方式
|
||||
|
||||
@@ -394,6 +394,10 @@ def upgrade() -> None:
|
||||
index=True,
|
||||
),
|
||||
)
|
||||
# usage 表复合索引(优化常见查询)
|
||||
op.create_index("idx_usage_user_created", "usage", ["user_id", "created_at"])
|
||||
op.create_index("idx_usage_apikey_created", "usage", ["api_key_id", "created_at"])
|
||||
op.create_index("idx_usage_provider_model_created", "usage", ["provider", "model", "created_at"])
|
||||
|
||||
# ==================== user_quotas ====================
|
||||
op.create_table(
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
"""add stats_daily_model table and rename provider_model_aliases
|
||||
|
||||
Revision ID: a1b2c3d4e5f6
|
||||
Revises: f30f9936f6a2
|
||||
Create Date: 2025-12-20 12:00:00.000000+00:00
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'a1b2c3d4e5f6'
|
||||
down_revision = 'f30f9936f6a2'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col['name'] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""创建 stats_daily_model 表,重命名 provider_model_aliases 为 provider_model_mappings"""
|
||||
# 1. 创建 stats_daily_model 表
|
||||
if not table_exists('stats_daily_model'):
|
||||
op.create_table(
|
||||
'stats_daily_model',
|
||||
sa.Column('id', sa.String(36), primary_key=True),
|
||||
sa.Column('date', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('model', sa.String(100), nullable=False),
|
||||
sa.Column('total_requests', sa.Integer(), nullable=False, default=0),
|
||||
sa.Column('input_tokens', sa.BigInteger(), nullable=False, default=0),
|
||||
sa.Column('output_tokens', sa.BigInteger(), nullable=False, default=0),
|
||||
sa.Column('cache_creation_tokens', sa.BigInteger(), nullable=False, default=0),
|
||||
sa.Column('cache_read_tokens', sa.BigInteger(), nullable=False, default=0),
|
||||
sa.Column('total_cost', sa.Float(), nullable=False, default=0.0),
|
||||
sa.Column('avg_response_time_ms', sa.Float(), nullable=False, default=0.0),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False,
|
||||
server_default=sa.func.now()),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False,
|
||||
server_default=sa.func.now(), onupdate=sa.func.now()),
|
||||
sa.UniqueConstraint('date', 'model', name='uq_stats_daily_model'),
|
||||
)
|
||||
|
||||
# 创建索引
|
||||
op.create_index('idx_stats_daily_model_date', 'stats_daily_model', ['date'])
|
||||
op.create_index('idx_stats_daily_model_date_model', 'stats_daily_model', ['date', 'model'])
|
||||
|
||||
# 2. 重命名 models 表的 provider_model_aliases 为 provider_model_mappings
|
||||
if column_exists('models', 'provider_model_aliases') and not column_exists('models', 'provider_model_mappings'):
|
||||
op.alter_column('models', 'provider_model_aliases', new_column_name='provider_model_mappings')
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
"""检查索引是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
indexes = [idx['name'] for idx in inspector.get_indexes(table_name)]
|
||||
return index_name in indexes
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""删除 stats_daily_model 表,恢复 provider_model_aliases 列名"""
|
||||
# 恢复列名
|
||||
if column_exists('models', 'provider_model_mappings') and not column_exists('models', 'provider_model_aliases'):
|
||||
op.alter_column('models', 'provider_model_mappings', new_column_name='provider_model_aliases')
|
||||
|
||||
# 删除表
|
||||
if table_exists('stats_daily_model'):
|
||||
if index_exists('stats_daily_model', 'idx_stats_daily_model_date_model'):
|
||||
op.drop_index('idx_stats_daily_model_date_model', table_name='stats_daily_model')
|
||||
if index_exists('stats_daily_model', 'idx_stats_daily_model_date'):
|
||||
op.drop_index('idx_stats_daily_model_date', table_name='stats_daily_model')
|
||||
op.drop_table('stats_daily_model')
|
||||
@@ -0,0 +1,65 @@
|
||||
"""add usage table composite indexes for query optimization
|
||||
|
||||
Revision ID: b2c3d4e5f6g7
|
||||
Revises: a1b2c3d4e5f6
|
||||
Create Date: 2025-12-20 15:00:00.000000+00:00
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
from sqlalchemy import text
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'b2c3d4e5f6g7'
|
||||
down_revision = 'a1b2c3d4e5f6'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""为 usage 表添加复合索引以优化常见查询
|
||||
|
||||
注意:这些索引已经在 baseline 迁移中创建。
|
||||
此迁移仅用于从旧版本升级的场景,新安装会跳过。
|
||||
"""
|
||||
conn = op.get_bind()
|
||||
|
||||
# 检查 usage 表是否存在
|
||||
result = conn.execute(text(
|
||||
"SELECT EXISTS (SELECT FROM information_schema.tables WHERE table_name = 'usage')"
|
||||
))
|
||||
if not result.scalar():
|
||||
# 表不存在,跳过
|
||||
return
|
||||
|
||||
# 定义需要创建的索引
|
||||
indexes = [
|
||||
("idx_usage_user_created", "ON usage (user_id, created_at)"),
|
||||
("idx_usage_apikey_created", "ON usage (api_key_id, created_at)"),
|
||||
("idx_usage_provider_model_created", "ON usage (provider, model, created_at)"),
|
||||
]
|
||||
|
||||
# 分别检查并创建每个索引
|
||||
for index_name, index_def in indexes:
|
||||
result = conn.execute(text(
|
||||
f"SELECT EXISTS (SELECT 1 FROM pg_indexes WHERE indexname = '{index_name}')"
|
||||
))
|
||||
if result.scalar():
|
||||
continue # 索引已存在,跳过
|
||||
|
||||
conn.execute(text(f"CREATE INDEX {index_name} {index_def}"))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""删除复合索引"""
|
||||
conn = op.get_bind()
|
||||
|
||||
# 使用 IF EXISTS 避免索引不存在时报错
|
||||
conn.execute(text(
|
||||
"DROP INDEX IF EXISTS idx_usage_provider_model_created"
|
||||
))
|
||||
conn.execute(text(
|
||||
"DROP INDEX IF EXISTS idx_usage_apikey_created"
|
||||
))
|
||||
conn.execute(text(
|
||||
"DROP INDEX IF EXISTS idx_usage_user_created"
|
||||
))
|
||||
@@ -0,0 +1,161 @@
|
||||
"""add ldap authentication support
|
||||
|
||||
Revision ID: c3d4e5f6g7h8
|
||||
Revises: b2c3d4e5f6g7
|
||||
Create Date: 2026-01-01 14:00:00.000000+00:00
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import text
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'c3d4e5f6g7h8'
|
||||
down_revision = 'b2c3d4e5f6g7'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _type_exists(conn, type_name: str) -> bool:
|
||||
"""检查 PostgreSQL 类型是否存在"""
|
||||
result = conn.execute(
|
||||
text("SELECT 1 FROM pg_type WHERE typname = :name"),
|
||||
{"name": type_name}
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def _column_exists(conn, table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否存在"""
|
||||
result = conn.execute(
|
||||
text("""
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_name = :table AND column_name = :column
|
||||
"""),
|
||||
{"table": table_name, "column": column_name}
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def _index_exists(conn, index_name: str) -> bool:
|
||||
"""检查索引是否存在"""
|
||||
result = conn.execute(
|
||||
text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
|
||||
{"name": index_name}
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def _table_exists(conn, table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
result = conn.execute(
|
||||
text("""
|
||||
SELECT 1 FROM information_schema.tables
|
||||
WHERE table_name = :name AND table_schema = 'public'
|
||||
"""),
|
||||
{"name": table_name}
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""添加 LDAP 认证支持
|
||||
|
||||
1. 创建 authsource 枚举类型
|
||||
2. 在 users 表添加 auth_source 字段和 LDAP 标识字段
|
||||
3. 创建 ldap_configs 表
|
||||
"""
|
||||
conn = op.get_bind()
|
||||
|
||||
# 1. 创建 authsource 枚举类型(幂等)
|
||||
if not _type_exists(conn, 'authsource'):
|
||||
conn.execute(text("CREATE TYPE authsource AS ENUM ('local', 'ldap')"))
|
||||
|
||||
# 2. 在 users 表添加字段(幂等)
|
||||
if not _column_exists(conn, 'users', 'auth_source'):
|
||||
op.add_column('users', sa.Column(
|
||||
'auth_source',
|
||||
sa.Enum('local', 'ldap', name='authsource', create_type=False),
|
||||
nullable=False,
|
||||
server_default='local'
|
||||
))
|
||||
|
||||
if not _column_exists(conn, 'users', 'ldap_dn'):
|
||||
op.add_column('users', sa.Column('ldap_dn', sa.String(length=512), nullable=True))
|
||||
|
||||
if not _column_exists(conn, 'users', 'ldap_username'):
|
||||
op.add_column('users', sa.Column('ldap_username', sa.String(length=255), nullable=True))
|
||||
|
||||
# 创建索引(幂等)
|
||||
if not _index_exists(conn, 'ix_users_ldap_dn'):
|
||||
op.create_index('ix_users_ldap_dn', 'users', ['ldap_dn'])
|
||||
|
||||
if not _index_exists(conn, 'ix_users_ldap_username'):
|
||||
op.create_index('ix_users_ldap_username', 'users', ['ldap_username'])
|
||||
|
||||
# 3. 创建 ldap_configs 表(幂等)
|
||||
if not _table_exists(conn, 'ldap_configs'):
|
||||
op.create_table(
|
||||
'ldap_configs',
|
||||
sa.Column('id', sa.Integer(), autoincrement=True, nullable=False),
|
||||
sa.Column('server_url', sa.String(length=255), nullable=False),
|
||||
sa.Column('bind_dn', sa.String(length=255), nullable=False),
|
||||
sa.Column('bind_password_encrypted', sa.Text(), nullable=True),
|
||||
sa.Column('base_dn', sa.String(length=255), nullable=False),
|
||||
sa.Column('user_search_filter', sa.String(length=500), nullable=False, server_default='(uid={username})'),
|
||||
sa.Column('username_attr', sa.String(length=50), nullable=False, server_default='uid'),
|
||||
sa.Column('email_attr', sa.String(length=50), nullable=False, server_default='mail'),
|
||||
sa.Column('display_name_attr', sa.String(length=50), nullable=False, server_default='cn'),
|
||||
sa.Column('is_enabled', sa.Boolean(), nullable=False, server_default='false'),
|
||||
sa.Column('is_exclusive', sa.Boolean(), nullable=False, server_default='false'),
|
||||
sa.Column('use_starttls', sa.Boolean(), nullable=False, server_default='false'),
|
||||
sa.Column('connect_timeout', sa.Integer(), nullable=False, server_default='10'),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False, server_default=sa.text('now()')),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False, server_default=sa.text('now()')),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""回滚 LDAP 认证支持
|
||||
|
||||
警告:回滚前请确保:
|
||||
1. 已备份数据库
|
||||
2. 没有 LDAP 用户需要保留
|
||||
"""
|
||||
conn = op.get_bind()
|
||||
|
||||
# 检查是否存在 LDAP 用户,防止数据丢失
|
||||
if _column_exists(conn, 'users', 'auth_source'):
|
||||
result = conn.execute(text("SELECT COUNT(*) FROM users WHERE auth_source = 'ldap'"))
|
||||
ldap_user_count = result.scalar()
|
||||
if ldap_user_count and ldap_user_count > 0:
|
||||
raise RuntimeError(
|
||||
f"无法回滚:存在 {ldap_user_count} 个 LDAP 用户。"
|
||||
f"请先删除或转换这些用户,或使用 --force 参数强制回滚(将丢失数据)。"
|
||||
)
|
||||
|
||||
# 1. 删除 ldap_configs 表(幂等)
|
||||
if _table_exists(conn, 'ldap_configs'):
|
||||
op.drop_table('ldap_configs')
|
||||
|
||||
# 2. 删除 users 表的 LDAP 相关字段(幂等)
|
||||
if _index_exists(conn, 'ix_users_ldap_username'):
|
||||
op.drop_index('ix_users_ldap_username', table_name='users')
|
||||
|
||||
if _index_exists(conn, 'ix_users_ldap_dn'):
|
||||
op.drop_index('ix_users_ldap_dn', table_name='users')
|
||||
|
||||
if _column_exists(conn, 'users', 'ldap_username'):
|
||||
op.drop_column('users', 'ldap_username')
|
||||
|
||||
if _column_exists(conn, 'users', 'ldap_dn'):
|
||||
op.drop_column('users', 'ldap_dn')
|
||||
|
||||
if _column_exists(conn, 'users', 'auth_source'):
|
||||
op.drop_column('users', 'auth_source')
|
||||
|
||||
# 3. 删除 authsource 枚举类型(幂等)
|
||||
# 注意:不使用 CASCADE,因为此时所有依赖应该已被删除
|
||||
if _type_exists(conn, 'authsource'):
|
||||
conn.execute(text("DROP TYPE authsource"))
|
||||
@@ -0,0 +1,131 @@
|
||||
"""add_management_tokens_table
|
||||
|
||||
Revision ID: ad55f1d008b7
|
||||
Revises: c3d4e5f6g7h8
|
||||
Create Date: 2026-01-06 15:24:10.660394+00:00
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'ad55f1d008b7'
|
||||
down_revision = 'c3d4e5f6g7h8'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
conn = op.get_bind()
|
||||
inspector = inspect(conn)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
"""检查索引是否存在"""
|
||||
conn = op.get_bind()
|
||||
inspector = inspect(conn)
|
||||
try:
|
||||
indexes = inspector.get_indexes(table_name)
|
||||
return any(idx["name"] == index_name for idx in indexes)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def constraint_exists(table_name: str, constraint_name: str) -> bool:
|
||||
"""检查约束是否存在"""
|
||||
conn = op.get_bind()
|
||||
inspector = inspect(conn)
|
||||
try:
|
||||
constraints = inspector.get_unique_constraints(table_name)
|
||||
if any(c["name"] == constraint_name for c in constraints):
|
||||
return True
|
||||
# 也检查 check 约束
|
||||
check_constraints = inspector.get_check_constraints(table_name)
|
||||
if any(c["name"] == constraint_name for c in check_constraints):
|
||||
return True
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""应用迁移:创建 management_tokens 表"""
|
||||
# 幂等性检查
|
||||
if table_exists("management_tokens"):
|
||||
# 表已存在,检查是否需要添加约束
|
||||
if not constraint_exists("management_tokens", "uq_management_tokens_user_name"):
|
||||
op.create_unique_constraint(
|
||||
"uq_management_tokens_user_name",
|
||||
"management_tokens",
|
||||
["user_id", "name"],
|
||||
)
|
||||
# 添加 IP 白名单非空检查约束
|
||||
if not constraint_exists("management_tokens", "check_allowed_ips_not_empty"):
|
||||
op.create_check_constraint(
|
||||
"check_allowed_ips_not_empty",
|
||||
"management_tokens",
|
||||
"allowed_ips IS NULL OR allowed_ips::text = 'null' OR json_array_length(allowed_ips) > 0",
|
||||
)
|
||||
return
|
||||
|
||||
op.create_table('management_tokens',
|
||||
sa.Column('id', sa.String(length=36), nullable=False),
|
||||
sa.Column('user_id', sa.String(length=36), nullable=False),
|
||||
sa.Column('token_hash', sa.String(length=64), nullable=False),
|
||||
sa.Column('token_prefix', sa.String(length=12), nullable=True),
|
||||
sa.Column('name', sa.String(length=100), nullable=False),
|
||||
sa.Column('description', sa.Text(), nullable=True),
|
||||
sa.Column('allowed_ips', sa.JSON(), nullable=True),
|
||||
sa.Column('expires_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('last_used_at', sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column('last_used_ip', sa.String(length=45), nullable=True),
|
||||
sa.Column('usage_count', sa.Integer(), server_default='0', nullable=False),
|
||||
sa.Column('is_active', sa.Boolean(), server_default='true', nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(['user_id'], ['users.id'], ondelete='CASCADE'),
|
||||
sa.PrimaryKeyConstraint('id')
|
||||
)
|
||||
op.create_index('idx_management_tokens_is_active', 'management_tokens', ['is_active'], unique=False)
|
||||
op.create_index('idx_management_tokens_user_id', 'management_tokens', ['user_id'], unique=False)
|
||||
op.create_index(op.f('ix_management_tokens_token_hash'), 'management_tokens', ['token_hash'], unique=True)
|
||||
# 添加用户名称唯一约束
|
||||
op.create_unique_constraint(
|
||||
"uq_management_tokens_user_name",
|
||||
"management_tokens",
|
||||
["user_id", "name"],
|
||||
)
|
||||
# 添加 IP 白名单非空检查约束
|
||||
# 注意:JSON 类型的 NULL 可能被序列化为 JSON 'null',需要同时处理
|
||||
op.create_check_constraint(
|
||||
"check_allowed_ips_not_empty",
|
||||
"management_tokens",
|
||||
"allowed_ips IS NULL OR allowed_ips::text = 'null' OR json_array_length(allowed_ips) > 0",
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""回滚迁移:删除 management_tokens 表"""
|
||||
# 幂等性检查
|
||||
if not table_exists("management_tokens"):
|
||||
return
|
||||
|
||||
# 删除约束
|
||||
if constraint_exists("management_tokens", "check_allowed_ips_not_empty"):
|
||||
op.drop_constraint("check_allowed_ips_not_empty", "management_tokens", type_="check")
|
||||
if constraint_exists("management_tokens", "uq_management_tokens_user_name"):
|
||||
op.drop_constraint("uq_management_tokens_user_name", "management_tokens", type_="unique")
|
||||
|
||||
# 删除索引
|
||||
if index_exists("management_tokens", "ix_management_tokens_token_hash"):
|
||||
op.drop_index(op.f('ix_management_tokens_token_hash'), table_name='management_tokens')
|
||||
if index_exists("management_tokens", "idx_management_tokens_user_id"):
|
||||
op.drop_index('idx_management_tokens_user_id', table_name='management_tokens')
|
||||
if index_exists("management_tokens", "idx_management_tokens_is_active"):
|
||||
op.drop_index('idx_management_tokens_is_active', table_name='management_tokens')
|
||||
|
||||
# 删除表
|
||||
op.drop_table('management_tokens')
|
||||
@@ -0,0 +1,73 @@
|
||||
"""cleanup ambiguous database fields
|
||||
|
||||
Revision ID: 02a45b66b7c4
|
||||
Revises: ad55f1d008b7
|
||||
Create Date: 2026-01-07 11:20:12.684426+00:00
|
||||
|
||||
变更内容:
|
||||
1. users 表:重命名 allowed_endpoints 为 allowed_api_formats(修正历史命名错误)
|
||||
2. api_keys 表:删除 allowed_endpoints 字段(未使用的功能)
|
||||
3. providers 表:删除 rate_limit 字段(与 rpm_limit 功能重复,且未使用)
|
||||
4. usage 表:重命名 provider 为 provider_name(避免与 provider_id 外键混淆)
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '02a45b66b7c4'
|
||||
down_revision = 'ad55f1d008b7'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col['name'] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""
|
||||
1. users.allowed_endpoints -> allowed_api_formats(重命名)
|
||||
2. api_keys.allowed_endpoints 删除
|
||||
3. providers.rate_limit 删除(与 rpm_limit 重复)
|
||||
4. usage.provider -> provider_name(重命名)
|
||||
"""
|
||||
# 1. users 表:重命名 allowed_endpoints 为 allowed_api_formats
|
||||
if _column_exists('users', 'allowed_endpoints'):
|
||||
op.alter_column('users', 'allowed_endpoints', new_column_name='allowed_api_formats')
|
||||
|
||||
# 2. api_keys 表:删除 allowed_endpoints 字段
|
||||
if _column_exists('api_keys', 'allowed_endpoints'):
|
||||
op.drop_column('api_keys', 'allowed_endpoints')
|
||||
|
||||
# 3. providers 表:删除 rate_limit 字段(与 rpm_limit 功能重复)
|
||||
if _column_exists('providers', 'rate_limit'):
|
||||
op.drop_column('providers', 'rate_limit')
|
||||
|
||||
# 4. usage 表:重命名 provider 为 provider_name
|
||||
if _column_exists('usage', 'provider'):
|
||||
op.alter_column('usage', 'provider', new_column_name='provider_name')
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""回滚:恢复原字段"""
|
||||
# 4. usage 表:将 provider_name 改回 provider
|
||||
if _column_exists('usage', 'provider_name'):
|
||||
op.alter_column('usage', 'provider_name', new_column_name='provider')
|
||||
|
||||
# 3. providers 表:恢复 rate_limit 字段
|
||||
if not _column_exists('providers', 'rate_limit'):
|
||||
op.add_column('providers', sa.Column('rate_limit', sa.Integer(), nullable=True))
|
||||
|
||||
# 2. api_keys 表:恢复 allowed_endpoints 字段
|
||||
if not _column_exists('api_keys', 'allowed_endpoints'):
|
||||
op.add_column('api_keys', sa.Column('allowed_endpoints', sa.JSON(), nullable=True))
|
||||
|
||||
# 1. users 表:将 allowed_api_formats 改回 allowed_endpoints
|
||||
if _column_exists('users', 'allowed_api_formats'):
|
||||
op.alter_column('users', 'allowed_api_formats', new_column_name='allowed_endpoints')
|
||||
@@ -0,0 +1,604 @@
|
||||
"""consolidated schema updates
|
||||
|
||||
Revision ID: m4n5o6p7q8r9
|
||||
Revises: 02a45b66b7c4
|
||||
Create Date: 2026-01-10 20:00:00.000000
|
||||
|
||||
This migration consolidates all schema changes from 2026-01-08 to 2026-01-10:
|
||||
|
||||
1. provider_api_keys: Key 直接关联 Provider (provider_id, api_formats)
|
||||
2. provider_api_keys: 添加 rate_multipliers JSON 字段(按格式费率)
|
||||
3. models: global_model_id 改为可空(支持独立 ProviderModel)
|
||||
4. providers: 添加 timeout, max_retries, proxy(从 endpoint 迁移)
|
||||
5. providers: display_name 重命名为 name,删除原 name
|
||||
6. provider_api_keys: max_concurrent -> rpm_limit(并发改 RPM)
|
||||
7. provider_api_keys: 健康度改为按格式存储(health_by_format, circuit_breaker_by_format)
|
||||
8. provider_endpoints: 删除废弃的 rate_limit 列
|
||||
9. usage: 添加 client_response_headers 字段
|
||||
10. provider_api_keys: 删除 endpoint_id(Key 不再与 Endpoint 绑定)
|
||||
11. provider_endpoints: 删除废弃的 max_concurrent 列
|
||||
12. providers: 删除废弃的 rpm_limit, rpm_used, rpm_reset_at 列
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects import postgresql
|
||||
from sqlalchemy.exc import ProgrammingError
|
||||
|
||||
from alembic import op
|
||||
|
||||
# 配置日志
|
||||
alembic_logger = logging.getLogger("alembic.runtime.migration")
|
||||
|
||||
revision = "m4n5o6p7q8r9"
|
||||
down_revision = "02a45b66b7c4"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""Check if a column exists in the table (bypasses inspector cache)"""
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text(
|
||||
"SELECT 1 FROM information_schema.columns "
|
||||
"WHERE table_name = :table AND column_name = :col"
|
||||
),
|
||||
{"table": table_name, "col": column_name},
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def _constraint_exists(table_name: str, constraint_name: str) -> bool:
|
||||
"""Check if a constraint exists (bypasses inspector cache)"""
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text(
|
||||
"SELECT 1 FROM information_schema.table_constraints "
|
||||
"WHERE table_name = :table AND constraint_name = :name"
|
||||
),
|
||||
{"table": table_name, "name": constraint_name},
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def _index_exists(table_name: str, index_name: str) -> bool:
|
||||
"""Check if an index exists (bypasses inspector cache)"""
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
|
||||
{"name": index_name},
|
||||
)
|
||||
return result.scalar() is not None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Apply all consolidated schema changes"""
|
||||
bind = op.get_bind()
|
||||
|
||||
# ========== 1. provider_api_keys: 添加 provider_id 和 api_formats ==========
|
||||
if not _column_exists("provider_api_keys", "provider_id"):
|
||||
conn = op.get_bind()
|
||||
conn.execute(sa.text("SAVEPOINT sp_add_provider_id"))
|
||||
try:
|
||||
op.add_column(
|
||||
"provider_api_keys", sa.Column("provider_id", sa.String(36), nullable=True)
|
||||
)
|
||||
conn.execute(sa.text("RELEASE SAVEPOINT sp_add_provider_id"))
|
||||
except ProgrammingError as exc:
|
||||
if getattr(getattr(exc, "orig", None), "pgcode", None) == "42701":
|
||||
conn.execute(sa.text("ROLLBACK TO SAVEPOINT sp_add_provider_id"))
|
||||
alembic_logger.warning("provider_api_keys.provider_id already exists; skipping add")
|
||||
else:
|
||||
conn.execute(sa.text("ROLLBACK TO SAVEPOINT sp_add_provider_id"))
|
||||
raise
|
||||
|
||||
# 数据迁移:从 endpoint 获取 provider_id(如果 endpoint_id 仍存在)
|
||||
if _column_exists("provider_api_keys", "endpoint_id"):
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys k
|
||||
SET provider_id = e.provider_id
|
||||
FROM provider_endpoints e
|
||||
WHERE k.endpoint_id = e.id AND k.provider_id IS NULL
|
||||
""")
|
||||
|
||||
# 检查无法关联的孤儿 Key
|
||||
result = bind.execute(
|
||||
sa.text("SELECT COUNT(*) FROM provider_api_keys WHERE provider_id IS NULL")
|
||||
)
|
||||
orphan_count = result.scalar() or 0
|
||||
if orphan_count > 0:
|
||||
# 使用 logger 记录更明显的告警
|
||||
alembic_logger.warning("=" * 60)
|
||||
alembic_logger.warning(
|
||||
f"[MIGRATION WARNING] 发现 {orphan_count} 个无法关联 Provider 的孤儿 Key"
|
||||
)
|
||||
alembic_logger.warning("=" * 60)
|
||||
alembic_logger.info("正在备份孤儿 Key 到 _orphan_api_keys_backup 表...")
|
||||
|
||||
# 先备份孤儿数据到临时表,避免数据丢失
|
||||
op.execute("""
|
||||
CREATE TABLE IF NOT EXISTS _orphan_api_keys_backup AS
|
||||
SELECT *, NOW() as backup_at
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id IS NULL
|
||||
""")
|
||||
|
||||
# 记录备份的 Key ID
|
||||
orphan_ids = bind.execute(
|
||||
sa.text("SELECT id, name FROM provider_api_keys WHERE provider_id IS NULL")
|
||||
).fetchall()
|
||||
alembic_logger.info("备份的孤儿 Key 列表:")
|
||||
for key_id, key_name in orphan_ids:
|
||||
alembic_logger.info(f" - Key: {key_name} (ID: {key_id})")
|
||||
|
||||
# 删除孤儿数据
|
||||
op.execute("DELETE FROM provider_api_keys WHERE provider_id IS NULL")
|
||||
alembic_logger.info(f"已备份并删除 {orphan_count} 个孤儿 Key")
|
||||
|
||||
# 提供恢复指南
|
||||
alembic_logger.warning("-" * 60)
|
||||
alembic_logger.warning("[恢复指南] 如需恢复孤儿 Key:")
|
||||
alembic_logger.warning(" 1. 查询备份表: SELECT * FROM _orphan_api_keys_backup;")
|
||||
alembic_logger.warning(" 2. 确定正确的 provider_id")
|
||||
alembic_logger.warning(" 3. 执行恢复:")
|
||||
alembic_logger.warning(" INSERT INTO provider_api_keys (...)")
|
||||
alembic_logger.warning(" SELECT ... FROM _orphan_api_keys_backup WHERE ...;")
|
||||
alembic_logger.warning("-" * 60)
|
||||
|
||||
# 设置 NOT NULL 并创建外键
|
||||
op.alter_column("provider_api_keys", "provider_id", nullable=False)
|
||||
|
||||
if not _constraint_exists("provider_api_keys", "fk_provider_api_keys_provider"):
|
||||
op.create_foreign_key(
|
||||
"fk_provider_api_keys_provider",
|
||||
"provider_api_keys",
|
||||
"providers",
|
||||
["provider_id"],
|
||||
["id"],
|
||||
ondelete="CASCADE",
|
||||
)
|
||||
|
||||
if not _index_exists("provider_api_keys", "idx_provider_api_keys_provider_id"):
|
||||
op.create_index("idx_provider_api_keys_provider_id", "provider_api_keys", ["provider_id"])
|
||||
|
||||
if not _column_exists("provider_api_keys", "api_formats"):
|
||||
op.add_column("provider_api_keys", sa.Column("api_formats", sa.JSON(), nullable=True))
|
||||
|
||||
# 数据迁移:从 endpoint 获取 api_format
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys k
|
||||
SET api_formats = json_build_array(e.api_format)
|
||||
FROM provider_endpoints e
|
||||
WHERE k.endpoint_id = e.id AND k.api_formats IS NULL
|
||||
""")
|
||||
|
||||
op.alter_column("provider_api_keys", "api_formats", nullable=False, server_default="[]")
|
||||
|
||||
# 修改 endpoint_id 为可空,外键改为 SET NULL
|
||||
if _constraint_exists("provider_api_keys", "provider_api_keys_endpoint_id_fkey"):
|
||||
op.drop_constraint(
|
||||
"provider_api_keys_endpoint_id_fkey", "provider_api_keys", type_="foreignkey"
|
||||
)
|
||||
op.alter_column("provider_api_keys", "endpoint_id", nullable=True)
|
||||
# 不再重建外键,因为后面会删除这个字段
|
||||
|
||||
# ========== 2. provider_api_keys: 添加 rate_multipliers ==========
|
||||
if not _column_exists("provider_api_keys", "rate_multipliers"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("rate_multipliers", postgresql.JSON(astext_type=sa.Text()), nullable=True),
|
||||
)
|
||||
|
||||
# 数据迁移:将 rate_multiplier 按 api_formats 转换
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys
|
||||
SET rate_multipliers = (
|
||||
SELECT jsonb_object_agg(elem, rate_multiplier)
|
||||
FROM jsonb_array_elements_text(api_formats::jsonb) AS elem
|
||||
)
|
||||
WHERE api_formats IS NOT NULL
|
||||
AND api_formats::text != '[]'
|
||||
AND api_formats::text != 'null'
|
||||
AND rate_multipliers IS NULL
|
||||
""")
|
||||
|
||||
# ========== 3. models: global_model_id 改为可空 ==========
|
||||
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=True)
|
||||
|
||||
# ========== 4. providers: 添加 timeout, max_retries, proxy ==========
|
||||
if not _column_exists("providers", "timeout"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("timeout", sa.Integer(), nullable=True, comment="请求超时(秒)"),
|
||||
)
|
||||
|
||||
if not _column_exists("providers", "max_retries"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("max_retries", sa.Integer(), nullable=True, comment="最大重试次数"),
|
||||
)
|
||||
|
||||
if not _column_exists("providers", "proxy"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("proxy", postgresql.JSONB(), nullable=True, comment="代理配置"),
|
||||
)
|
||||
|
||||
# 从端点迁移数据到 provider(动态构建 SQL,仅引用存在的列)
|
||||
ep_has_timeout = _column_exists("provider_endpoints", "timeout")
|
||||
ep_has_max_retries = _column_exists("provider_endpoints", "max_retries")
|
||||
ep_has_proxy = _column_exists("provider_endpoints", "proxy")
|
||||
|
||||
set_clauses = []
|
||||
if _column_exists("providers", "timeout"):
|
||||
if ep_has_timeout:
|
||||
set_clauses.append("""
|
||||
timeout = COALESCE(
|
||||
p.timeout,
|
||||
(SELECT MAX(e.timeout) FROM provider_endpoints e WHERE e.provider_id = p.id AND e.timeout IS NOT NULL),
|
||||
300
|
||||
)""")
|
||||
else:
|
||||
set_clauses.append("timeout = COALESCE(p.timeout, 300)")
|
||||
|
||||
if _column_exists("providers", "max_retries"):
|
||||
if ep_has_max_retries:
|
||||
set_clauses.append("""
|
||||
max_retries = COALESCE(
|
||||
p.max_retries,
|
||||
(SELECT MAX(e.max_retries) FROM provider_endpoints e WHERE e.provider_id = p.id AND e.max_retries IS NOT NULL),
|
||||
2
|
||||
)""")
|
||||
else:
|
||||
set_clauses.append("max_retries = COALESCE(p.max_retries, 2)")
|
||||
|
||||
if _column_exists("providers", "proxy") and ep_has_proxy:
|
||||
set_clauses.append("""
|
||||
proxy = COALESCE(
|
||||
p.proxy,
|
||||
(SELECT e.proxy FROM provider_endpoints e WHERE e.provider_id = p.id AND e.proxy IS NOT NULL ORDER BY e.created_at LIMIT 1)
|
||||
)""")
|
||||
|
||||
if set_clauses:
|
||||
where_parts = []
|
||||
if _column_exists("providers", "timeout"):
|
||||
where_parts.append("p.timeout IS NULL")
|
||||
if _column_exists("providers", "max_retries"):
|
||||
where_parts.append("p.max_retries IS NULL")
|
||||
where_clause = " OR ".join(where_parts) if where_parts else "TRUE"
|
||||
sql = "UPDATE providers p SET " + ", ".join(set_clauses) + " WHERE " + where_clause
|
||||
op.execute(sql)
|
||||
|
||||
# ========== 5. providers: display_name -> name ==========
|
||||
# 注意:这里假设 display_name 已经被重命名为 name
|
||||
# 如果 display_name 仍然存在,则需要执行重命名
|
||||
if _column_exists("providers", "display_name"):
|
||||
# 删除旧的 name 索引
|
||||
if _index_exists("providers", "ix_providers_name"):
|
||||
op.drop_index("ix_providers_name", table_name="providers")
|
||||
|
||||
# 如果存在旧的 name 列,先删除
|
||||
if _column_exists("providers", "name"):
|
||||
op.drop_column("providers", "name")
|
||||
|
||||
# 重命名 display_name 为 name
|
||||
op.alter_column("providers", "display_name", new_column_name="name")
|
||||
|
||||
# 创建新索引
|
||||
op.create_index("ix_providers_name", "providers", ["name"], unique=True)
|
||||
|
||||
# ========== 6. provider_api_keys: max_concurrent -> rpm_limit ==========
|
||||
if _column_exists("provider_api_keys", "max_concurrent"):
|
||||
op.alter_column("provider_api_keys", "max_concurrent", new_column_name="rpm_limit")
|
||||
|
||||
if _column_exists("provider_api_keys", "learned_max_concurrent"):
|
||||
op.alter_column(
|
||||
"provider_api_keys", "learned_max_concurrent", new_column_name="learned_rpm_limit"
|
||||
)
|
||||
|
||||
if _column_exists("provider_api_keys", "last_concurrent_peak"):
|
||||
op.alter_column(
|
||||
"provider_api_keys", "last_concurrent_peak", new_column_name="last_rpm_peak"
|
||||
)
|
||||
|
||||
# 删除废弃字段
|
||||
for col in ["rate_limit", "daily_limit", "monthly_limit"]:
|
||||
if _column_exists("provider_api_keys", col):
|
||||
op.drop_column("provider_api_keys", col)
|
||||
|
||||
# ========== 7. provider_api_keys: 健康度改为按格式存储 ==========
|
||||
if not _column_exists("provider_api_keys", "health_by_format"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column(
|
||||
"health_by_format",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
nullable=True,
|
||||
comment="按API格式存储的健康度数据",
|
||||
),
|
||||
)
|
||||
|
||||
if not _column_exists("provider_api_keys", "circuit_breaker_by_format"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column(
|
||||
"circuit_breaker_by_format",
|
||||
postgresql.JSONB(astext_type=sa.Text()),
|
||||
nullable=True,
|
||||
comment="按API格式存储的熔断器状态",
|
||||
),
|
||||
)
|
||||
|
||||
# 数据迁移:如果存在旧字段,迁移数据到新结构
|
||||
if _column_exists("provider_api_keys", "health_score"):
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys
|
||||
SET health_by_format = (
|
||||
SELECT jsonb_object_agg(
|
||||
elem,
|
||||
jsonb_build_object(
|
||||
'health_score', COALESCE(health_score, 1.0),
|
||||
'consecutive_failures', COALESCE(consecutive_failures, 0),
|
||||
'last_failure_at', last_failure_at,
|
||||
'request_results_window', COALESCE(request_results_window::jsonb, '[]'::jsonb)
|
||||
)
|
||||
)
|
||||
FROM jsonb_array_elements_text(api_formats::jsonb) AS elem
|
||||
)
|
||||
WHERE api_formats IS NOT NULL
|
||||
AND api_formats::text != '[]'
|
||||
AND health_by_format IS NULL
|
||||
""")
|
||||
|
||||
# Circuit Breaker 迁移策略:
|
||||
# 不复制旧的 circuit_breaker_open 状态到所有 format,而是全部重置为 closed
|
||||
# 原因:旧的单一 circuit breaker 状态可能因某一个 format 失败而打开,
|
||||
# 如果复制到所有 format,会导致其他正常工作的 format 被错误标记为不可用
|
||||
if _column_exists("provider_api_keys", "circuit_breaker_open"):
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys
|
||||
SET circuit_breaker_by_format = (
|
||||
SELECT jsonb_object_agg(
|
||||
elem,
|
||||
jsonb_build_object(
|
||||
'open', false,
|
||||
'open_at', NULL,
|
||||
'next_probe_at', NULL,
|
||||
'half_open_until', NULL,
|
||||
'half_open_successes', 0,
|
||||
'half_open_failures', 0
|
||||
)
|
||||
)
|
||||
FROM jsonb_array_elements_text(api_formats::jsonb) AS elem
|
||||
)
|
||||
WHERE api_formats IS NOT NULL
|
||||
AND api_formats::text != '[]'
|
||||
AND circuit_breaker_by_format IS NULL
|
||||
""")
|
||||
|
||||
# 设置默认空对象
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys
|
||||
SET health_by_format = '{}'::jsonb
|
||||
WHERE health_by_format IS NULL
|
||||
""")
|
||||
op.execute("""
|
||||
UPDATE provider_api_keys
|
||||
SET circuit_breaker_by_format = '{}'::jsonb
|
||||
WHERE circuit_breaker_by_format IS NULL
|
||||
""")
|
||||
|
||||
# 创建 GIN 索引
|
||||
if not _index_exists("provider_api_keys", "ix_provider_api_keys_health_by_format"):
|
||||
op.create_index(
|
||||
"ix_provider_api_keys_health_by_format",
|
||||
"provider_api_keys",
|
||||
["health_by_format"],
|
||||
postgresql_using="gin",
|
||||
)
|
||||
if not _index_exists("provider_api_keys", "ix_provider_api_keys_circuit_breaker_by_format"):
|
||||
op.create_index(
|
||||
"ix_provider_api_keys_circuit_breaker_by_format",
|
||||
"provider_api_keys",
|
||||
["circuit_breaker_by_format"],
|
||||
postgresql_using="gin",
|
||||
)
|
||||
|
||||
# 删除旧字段
|
||||
old_health_columns = [
|
||||
"health_score",
|
||||
"consecutive_failures",
|
||||
"last_failure_at",
|
||||
"request_results_window",
|
||||
"circuit_breaker_open",
|
||||
"circuit_breaker_open_at",
|
||||
"next_probe_at",
|
||||
"half_open_until",
|
||||
"half_open_successes",
|
||||
"half_open_failures",
|
||||
]
|
||||
for col in old_health_columns:
|
||||
if _column_exists("provider_api_keys", col):
|
||||
op.drop_column("provider_api_keys", col)
|
||||
|
||||
# ========== 8. provider_endpoints: 删除废弃的 rate_limit 列 ==========
|
||||
if _column_exists("provider_endpoints", "rate_limit"):
|
||||
op.drop_column("provider_endpoints", "rate_limit")
|
||||
|
||||
# ========== 9. usage: 添加 client_response_headers ==========
|
||||
if not _column_exists("usage", "client_response_headers"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column("client_response_headers", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# ========== 10. provider_api_keys: 删除 endpoint_id ==========
|
||||
# Key 不再与 Endpoint 绑定,通过 provider_id + api_formats 关联
|
||||
if _column_exists("provider_api_keys", "endpoint_id"):
|
||||
# 查找 endpoint_id 上的外键并删除(用 savepoint 保护,避免事务中止)
|
||||
conn = op.get_bind()
|
||||
fk_rows = conn.execute(
|
||||
sa.text(
|
||||
"SELECT con.conname FROM pg_constraint con "
|
||||
"JOIN pg_attribute att ON att.attnum = ANY(con.conkey) "
|
||||
" AND att.attrelid = con.conrelid "
|
||||
"WHERE con.conrelid = 'provider_api_keys'::regclass "
|
||||
" AND con.contype = 'f' AND att.attname = 'endpoint_id'"
|
||||
)
|
||||
).fetchall()
|
||||
for (fk_name,) in fk_rows:
|
||||
conn.execute(sa.text(f"SAVEPOINT sp_drop_fk_{fk_name}"))
|
||||
try:
|
||||
op.drop_constraint(fk_name, "provider_api_keys", type_="foreignkey")
|
||||
conn.execute(sa.text(f"RELEASE SAVEPOINT sp_drop_fk_{fk_name}"))
|
||||
except Exception:
|
||||
conn.execute(sa.text(f"ROLLBACK TO SAVEPOINT sp_drop_fk_{fk_name}"))
|
||||
op.drop_column("provider_api_keys", "endpoint_id")
|
||||
|
||||
# ========== 11. provider_endpoints: 删除废弃的 max_concurrent 列 ==========
|
||||
if _column_exists("provider_endpoints", "max_concurrent"):
|
||||
op.drop_column("provider_endpoints", "max_concurrent")
|
||||
|
||||
# ========== 12. providers: 删除废弃的 RPM 相关字段 ==========
|
||||
if _column_exists("providers", "rpm_limit"):
|
||||
op.drop_column("providers", "rpm_limit")
|
||||
if _column_exists("providers", "rpm_used"):
|
||||
op.drop_column("providers", "rpm_used")
|
||||
if _column_exists("providers", "rpm_reset_at"):
|
||||
op.drop_column("providers", "rpm_reset_at")
|
||||
|
||||
alembic_logger.info("[OK] Consolidated migration completed successfully")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""
|
||||
Downgrade is complex due to data migrations.
|
||||
For safety, this only removes new columns without restoring old structure.
|
||||
Manual intervention may be required for full rollback.
|
||||
"""
|
||||
bind = op.get_bind()
|
||||
|
||||
# 12. 恢复 providers RPM 相关字段
|
||||
if not _column_exists("providers", "rpm_limit"):
|
||||
op.add_column("providers", sa.Column("rpm_limit", sa.Integer(), nullable=True))
|
||||
if not _column_exists("providers", "rpm_used"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("rpm_used", sa.Integer(), server_default="0", nullable=True),
|
||||
)
|
||||
if not _column_exists("providers", "rpm_reset_at"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("rpm_reset_at", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
|
||||
# 11. 恢复 provider_endpoints.max_concurrent
|
||||
if not _column_exists("provider_endpoints", "max_concurrent"):
|
||||
op.add_column(
|
||||
"provider_endpoints", sa.Column("max_concurrent", sa.Integer(), nullable=True)
|
||||
)
|
||||
|
||||
# 10. 恢复 endpoint_id
|
||||
if not _column_exists("provider_api_keys", "endpoint_id"):
|
||||
op.add_column("provider_api_keys", sa.Column("endpoint_id", sa.String(36), nullable=True))
|
||||
|
||||
# 9. 删除 client_response_headers
|
||||
if _column_exists("usage", "client_response_headers"):
|
||||
op.drop_column("usage", "client_response_headers")
|
||||
|
||||
# 8. 恢复 provider_endpoints.rate_limit(如果需要)
|
||||
if not _column_exists("provider_endpoints", "rate_limit"):
|
||||
op.add_column("provider_endpoints", sa.Column("rate_limit", sa.Integer(), nullable=True))
|
||||
|
||||
# 7. 删除健康度 JSON 字段
|
||||
bind.execute(sa.text("DROP INDEX IF EXISTS ix_provider_api_keys_health_by_format"))
|
||||
bind.execute(sa.text("DROP INDEX IF EXISTS ix_provider_api_keys_circuit_breaker_by_format"))
|
||||
if _column_exists("provider_api_keys", "health_by_format"):
|
||||
op.drop_column("provider_api_keys", "health_by_format")
|
||||
if _column_exists("provider_api_keys", "circuit_breaker_by_format"):
|
||||
op.drop_column("provider_api_keys", "circuit_breaker_by_format")
|
||||
|
||||
# 6. rpm_limit -> max_concurrent(简化版:仅重命名)
|
||||
if _column_exists("provider_api_keys", "rpm_limit"):
|
||||
op.alter_column("provider_api_keys", "rpm_limit", new_column_name="max_concurrent")
|
||||
if _column_exists("provider_api_keys", "learned_rpm_limit"):
|
||||
op.alter_column(
|
||||
"provider_api_keys", "learned_rpm_limit", new_column_name="learned_max_concurrent"
|
||||
)
|
||||
if _column_exists("provider_api_keys", "last_rpm_peak"):
|
||||
op.alter_column(
|
||||
"provider_api_keys", "last_rpm_peak", new_column_name="last_concurrent_peak"
|
||||
)
|
||||
|
||||
# 恢复已删除的字段
|
||||
if not _column_exists("provider_api_keys", "rate_limit"):
|
||||
op.add_column("provider_api_keys", sa.Column("rate_limit", sa.Integer(), nullable=True))
|
||||
if not _column_exists("provider_api_keys", "daily_limit"):
|
||||
op.add_column("provider_api_keys", sa.Column("daily_limit", sa.Integer(), nullable=True))
|
||||
if not _column_exists("provider_api_keys", "monthly_limit"):
|
||||
op.add_column("provider_api_keys", sa.Column("monthly_limit", sa.Integer(), nullable=True))
|
||||
|
||||
# 5. name -> display_name (需要先删除索引)
|
||||
if _column_exists("providers", "name") and not _column_exists("providers", "display_name"):
|
||||
if _index_exists("providers", "ix_providers_name"):
|
||||
op.drop_index("ix_providers_name", table_name="providers")
|
||||
op.alter_column("providers", "name", new_column_name="display_name")
|
||||
|
||||
if not _column_exists("providers", "name"):
|
||||
op.add_column("providers", sa.Column("name", sa.String(100), nullable=True))
|
||||
op.execute("""
|
||||
UPDATE providers
|
||||
SET name = LOWER(REPLACE(REPLACE(display_name, ' ', '_'), '-', '_'))
|
||||
""")
|
||||
op.alter_column("providers", "name", nullable=False)
|
||||
if not _index_exists("providers", "ix_providers_name"):
|
||||
op.create_index("ix_providers_name", "providers", ["name"], unique=True)
|
||||
|
||||
# 4. 删除 providers 的 timeout, max_retries, proxy
|
||||
if _column_exists("providers", "proxy"):
|
||||
op.drop_column("providers", "proxy")
|
||||
if _column_exists("providers", "max_retries"):
|
||||
op.drop_column("providers", "max_retries")
|
||||
if _column_exists("providers", "timeout"):
|
||||
op.drop_column("providers", "timeout")
|
||||
|
||||
# 3. models: global_model_id 改回 NOT NULL
|
||||
result = bind.execute(sa.text("SELECT COUNT(*) FROM models WHERE global_model_id IS NULL"))
|
||||
orphan_model_count = result.scalar() or 0
|
||||
if orphan_model_count > 0:
|
||||
alembic_logger.warning(
|
||||
f"[WARN] 发现 {orphan_model_count} 个无 global_model_id 的独立模型,将被删除"
|
||||
)
|
||||
op.execute("DELETE FROM models WHERE global_model_id IS NULL")
|
||||
alembic_logger.info(f"已删除 {orphan_model_count} 个独立模型")
|
||||
op.alter_column("models", "global_model_id", nullable=False)
|
||||
|
||||
# 2. 删除 rate_multipliers
|
||||
if _column_exists("provider_api_keys", "rate_multipliers"):
|
||||
op.drop_column("provider_api_keys", "rate_multipliers")
|
||||
|
||||
# 1. 删除 provider_id 和 api_formats
|
||||
if _index_exists("provider_api_keys", "idx_provider_api_keys_provider_id"):
|
||||
op.drop_index("idx_provider_api_keys_provider_id", table_name="provider_api_keys")
|
||||
if _constraint_exists("provider_api_keys", "fk_provider_api_keys_provider"):
|
||||
op.drop_constraint("fk_provider_api_keys_provider", "provider_api_keys", type_="foreignkey")
|
||||
if _column_exists("provider_api_keys", "api_formats"):
|
||||
op.drop_column("provider_api_keys", "api_formats")
|
||||
if _column_exists("provider_api_keys", "provider_id"):
|
||||
op.drop_column("provider_api_keys", "provider_id")
|
||||
|
||||
# 恢复 endpoint_id 外键(简化版:仅创建外键,不强制 NOT NULL)
|
||||
if _column_exists("provider_api_keys", "endpoint_id"):
|
||||
if not _constraint_exists("provider_api_keys", "provider_api_keys_endpoint_id_fkey"):
|
||||
op.create_foreign_key(
|
||||
"provider_api_keys_endpoint_id_fkey",
|
||||
"provider_api_keys",
|
||||
"provider_endpoints",
|
||||
["endpoint_id"],
|
||||
["id"],
|
||||
ondelete="SET NULL",
|
||||
)
|
||||
|
||||
alembic_logger.info("[OK] Downgrade completed (simplified version)")
|
||||
+95
@@ -0,0 +1,95 @@
|
||||
"""add auto_fetch_models and locked_models to provider_api_keys
|
||||
|
||||
Revision ID: e4ebe3233b40
|
||||
Revises: m4n5o6p7q8r9
|
||||
Create Date: 2026-01-13 17:59:53.119479+00:00
|
||||
|
||||
为 provider_api_keys 表添加自动获取模型相关字段:
|
||||
1. auto_fetch_models: 是否启用自动获取模型
|
||||
2. last_models_fetch_at: 最后获取时间
|
||||
3. last_models_fetch_error: 最后获取错误信息
|
||||
4. locked_models: 被锁定的模型列表(刷新时不会被删除)
|
||||
|
||||
注意: downgrade 操作会永久删除 auto_fetch_models 配置和 locked_models 数据
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
def _index_exists(index_name: str) -> bool:
|
||||
"""Check if an index exists"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
indexes = inspector.get_indexes("provider_api_keys")
|
||||
return any(idx["name"] == index_name for idx in indexes)
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'e4ebe3233b40'
|
||||
down_revision = 'm4n5o6p7q8r9'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""Check if a column exists in the table"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""添加自动获取模型相关字段"""
|
||||
if not _column_exists("provider_api_keys", "auto_fetch_models"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("auto_fetch_models", sa.Boolean(), nullable=False, server_default="false"),
|
||||
)
|
||||
|
||||
if not _column_exists("provider_api_keys", "last_models_fetch_at"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("last_models_fetch_at", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
|
||||
if not _column_exists("provider_api_keys", "last_models_fetch_error"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("last_models_fetch_error", sa.Text(), nullable=True),
|
||||
)
|
||||
|
||||
if not _column_exists("provider_api_keys", "locked_models"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("locked_models", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# 添加复合索引以优化调度器查询
|
||||
if not _index_exists("ix_provider_api_keys_auto_fetch_active"):
|
||||
op.create_index(
|
||||
"ix_provider_api_keys_auto_fetch_active",
|
||||
"provider_api_keys",
|
||||
["auto_fetch_models", "is_active"],
|
||||
postgresql_where=sa.text("auto_fetch_models = true AND is_active = true"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""移除自动获取模型相关字段"""
|
||||
# 先删除索引
|
||||
if _index_exists("ix_provider_api_keys_auto_fetch_active"):
|
||||
op.drop_index("ix_provider_api_keys_auto_fetch_active", table_name="provider_api_keys")
|
||||
|
||||
if _column_exists("provider_api_keys", "locked_models"):
|
||||
op.drop_column("provider_api_keys", "locked_models")
|
||||
|
||||
if _column_exists("provider_api_keys", "last_models_fetch_error"):
|
||||
op.drop_column("provider_api_keys", "last_models_fetch_error")
|
||||
|
||||
if _column_exists("provider_api_keys", "last_models_fetch_at"):
|
||||
op.drop_column("provider_api_keys", "last_models_fetch_at")
|
||||
|
||||
if _column_exists("provider_api_keys", "auto_fetch_models"):
|
||||
op.drop_column("provider_api_keys", "auto_fetch_models")
|
||||
@@ -0,0 +1,104 @@
|
||||
"""add header_rules to provider_endpoints and is_locked to api_keys
|
||||
|
||||
Revision ID: 6d579000e511
|
||||
Revises: e4ebe3233b40
|
||||
Create Date: 2026-01-15 23:00:00.000000+00:00
|
||||
|
||||
变更:
|
||||
1. provider_endpoints 表: 添加 header_rules 字段,迁移 headers 数据
|
||||
2. api_keys 表: 添加 is_locked 字段(管理员锁定标志)
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects.postgresql import JSON
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = '6d579000e511'
|
||||
down_revision = 'e4ebe3233b40'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(connection, table: str, column: str) -> bool:
|
||||
"""检查列是否存在"""
|
||||
result = connection.execute(
|
||||
sa.text("""
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_name = :table AND column_name = :column
|
||||
"""),
|
||||
{"table": table, "column": column}
|
||||
)
|
||||
return result.fetchone() is not None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""添加 header_rules 字段并迁移现有 headers 数据;添加 is_locked 字段"""
|
||||
connection = op.get_bind()
|
||||
|
||||
# ========== provider_endpoints.header_rules ==========
|
||||
# 1. 添加 header_rules 列(幂等)
|
||||
if not _column_exists(connection, 'provider_endpoints', 'header_rules'):
|
||||
op.add_column('provider_endpoints', sa.Column('header_rules', JSON, nullable=True))
|
||||
|
||||
# 2. 批量迁移:headers -> header_rules
|
||||
# 使用纯 SQL 将 {"k1":"v1", "k2":"v2"} 转换为 [{"action":"set","key":"k1","value":"v1"}, ...]
|
||||
if _column_exists(connection, 'provider_endpoints', 'headers'):
|
||||
connection.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints
|
||||
SET header_rules = (
|
||||
SELECT jsonb_agg(
|
||||
jsonb_build_object('action', 'set', 'key', key, 'value', value)
|
||||
)
|
||||
FROM jsonb_each_text(headers::jsonb)
|
||||
)
|
||||
WHERE headers IS NOT NULL
|
||||
AND headers::text != '{}'
|
||||
AND jsonb_typeof(headers::jsonb) = 'object'
|
||||
AND header_rules IS NULL
|
||||
""")
|
||||
)
|
||||
|
||||
# 3. 删除旧列
|
||||
op.drop_column('provider_endpoints', 'headers')
|
||||
|
||||
# ========== api_keys.is_locked ==========
|
||||
if not _column_exists(connection, 'api_keys', 'is_locked'):
|
||||
op.add_column(
|
||||
'api_keys',
|
||||
sa.Column('is_locked', sa.Boolean(), nullable=False, server_default='false')
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""移除 header_rules 字段,恢复 headers 字段;移除 is_locked 字段"""
|
||||
connection = op.get_bind()
|
||||
|
||||
# ========== api_keys.is_locked ==========
|
||||
if _column_exists(connection, 'api_keys', 'is_locked'):
|
||||
op.drop_column('api_keys', 'is_locked')
|
||||
|
||||
# ========== provider_endpoints.header_rules ==========
|
||||
# 1. 添加 headers 列(幂等)
|
||||
if not _column_exists(connection, 'provider_endpoints', 'headers'):
|
||||
op.add_column('provider_endpoints', sa.Column('headers', JSON, nullable=True))
|
||||
|
||||
# 2. 批量迁移:header_rules -> headers(仅提取 set 操作)
|
||||
if _column_exists(connection, 'provider_endpoints', 'header_rules'):
|
||||
connection.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints
|
||||
SET headers = (
|
||||
SELECT jsonb_object_agg(rule->>'key', rule->>'value')
|
||||
FROM jsonb_array_elements(header_rules::jsonb) AS rule
|
||||
WHERE rule->>'action' = 'set'
|
||||
AND rule->>'key' IS NOT NULL
|
||||
)
|
||||
WHERE header_rules IS NOT NULL
|
||||
AND jsonb_typeof(header_rules::jsonb) = 'array'
|
||||
AND jsonb_array_length(header_rules::jsonb) > 0
|
||||
""")
|
||||
)
|
||||
|
||||
# 3. 删除 header_rules 列
|
||||
op.drop_column('provider_endpoints', 'header_rules')
|
||||
@@ -0,0 +1,127 @@
|
||||
"""add global_priority_by_format and remove deprecated fields
|
||||
|
||||
Revision ID: ddd59cdf0349
|
||||
Revises: 6d579000e511
|
||||
Create Date: 2026-01-16 12:00:00.000000+00:00
|
||||
|
||||
变更:
|
||||
1. provider_api_keys 表: 添加 global_priority_by_format 字段(按 API 格式的全局优先级)
|
||||
2. 迁移现有 global_priority 数据到新字段
|
||||
3. 删除已废弃的 global_priority 字段
|
||||
4. 删除已废弃的 rate_multiplier 字段(已被 rate_multipliers 替代)
|
||||
5. 删除已废弃的 providers.timeout 字段(由环境变量控制)
|
||||
6. 删除已废弃的 provider_endpoints.timeout 字段(由环境变量控制)
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.dialects.postgresql import JSON
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'ddd59cdf0349'
|
||||
down_revision = '6d579000e511'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(connection, table: str, column: str) -> bool:
|
||||
"""检查列是否存在"""
|
||||
result = connection.execute(
|
||||
sa.text("""
|
||||
SELECT 1 FROM information_schema.columns
|
||||
WHERE table_name = :table AND column_name = :column
|
||||
"""),
|
||||
{"table": table, "column": column}
|
||||
)
|
||||
return result.fetchone() is not None
|
||||
|
||||
|
||||
def upgrade():
|
||||
connection = op.get_bind()
|
||||
|
||||
# 1. 添加 global_priority_by_format 字段
|
||||
if not _column_exists(connection, 'provider_api_keys', 'global_priority_by_format'):
|
||||
op.add_column(
|
||||
'provider_api_keys',
|
||||
sa.Column('global_priority_by_format', JSON, nullable=True)
|
||||
)
|
||||
|
||||
# 2. 迁移现有 global_priority 数据到新字段
|
||||
# 对于有 global_priority 的 Key,将其值应用到所有支持的 api_formats
|
||||
if _column_exists(connection, 'provider_api_keys', 'global_priority'):
|
||||
# 将 JSON 数组转换为 text[] 后使用 unnest
|
||||
connection.execute(sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET global_priority_by_format = (
|
||||
SELECT jsonb_object_agg(format, global_priority)
|
||||
FROM jsonb_array_elements_text(api_formats::jsonb) AS format
|
||||
)
|
||||
WHERE global_priority IS NOT NULL
|
||||
AND api_formats IS NOT NULL
|
||||
AND jsonb_array_length(api_formats::jsonb) > 0
|
||||
AND global_priority_by_format IS NULL
|
||||
"""))
|
||||
|
||||
# 3. 删除 global_priority 字段
|
||||
op.drop_column('provider_api_keys', 'global_priority')
|
||||
|
||||
# 4. 删除 rate_multiplier 字段(已被 rate_multipliers 替代)
|
||||
if _column_exists(connection, 'provider_api_keys', 'rate_multiplier'):
|
||||
op.drop_column('provider_api_keys', 'rate_multiplier')
|
||||
|
||||
# 5. 删除 providers.timeout 字段(由环境变量控制)
|
||||
if _column_exists(connection, 'providers', 'timeout'):
|
||||
op.drop_column('providers', 'timeout')
|
||||
|
||||
# 6. 删除 provider_endpoints.timeout 字段(由环境变量控制)
|
||||
if _column_exists(connection, 'provider_endpoints', 'timeout'):
|
||||
op.drop_column('provider_endpoints', 'timeout')
|
||||
|
||||
|
||||
def downgrade():
|
||||
connection = op.get_bind()
|
||||
|
||||
# 1. 恢复 rate_multiplier 字段
|
||||
if not _column_exists(connection, 'provider_api_keys', 'rate_multiplier'):
|
||||
op.add_column(
|
||||
'provider_api_keys',
|
||||
sa.Column('rate_multiplier', sa.Float, nullable=False, server_default='1.0')
|
||||
)
|
||||
|
||||
# 2. 恢复 global_priority 字段并迁移数据
|
||||
if not _column_exists(connection, 'provider_api_keys', 'global_priority'):
|
||||
op.add_column(
|
||||
'provider_api_keys',
|
||||
sa.Column('global_priority', sa.Integer, nullable=True)
|
||||
)
|
||||
|
||||
# 从 global_priority_by_format 迁移数据(取第一个格式的优先级值)
|
||||
if _column_exists(connection, 'provider_api_keys', 'global_priority_by_format'):
|
||||
connection.execute(sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET global_priority = (
|
||||
SELECT (value::text)::integer
|
||||
FROM jsonb_each(global_priority_by_format::jsonb)
|
||||
LIMIT 1
|
||||
)
|
||||
WHERE global_priority_by_format IS NOT NULL
|
||||
AND jsonb_typeof(global_priority_by_format::jsonb) = 'object'
|
||||
AND global_priority IS NULL
|
||||
"""))
|
||||
|
||||
# 3. 删除 global_priority_by_format 字段
|
||||
if _column_exists(connection, 'provider_api_keys', 'global_priority_by_format'):
|
||||
op.drop_column('provider_api_keys', 'global_priority_by_format')
|
||||
|
||||
# 4. 恢复 providers.timeout 字段
|
||||
if not _column_exists(connection, 'providers', 'timeout'):
|
||||
op.add_column(
|
||||
'providers',
|
||||
sa.Column('timeout', sa.Integer, nullable=True, server_default='300')
|
||||
)
|
||||
|
||||
# 5. 恢复 provider_endpoints.timeout 字段
|
||||
if not _column_exists(connection, 'provider_endpoints', 'timeout'):
|
||||
op.add_column(
|
||||
'provider_endpoints',
|
||||
sa.Column('timeout', sa.Integer, nullable=True, server_default='300')
|
||||
)
|
||||
+223
@@ -0,0 +1,223 @@
|
||||
"""make users email/password nullable add email_verified and oauth tables
|
||||
|
||||
Revision ID: 33e347f97c0c
|
||||
Revises: ddd59cdf0349
|
||||
Create Date: 2026-01-18 11:18:15.940559+00:00
|
||||
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "33e347f97c0c"
|
||||
down_revision = "ddd59cdf0349"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_is_nullable(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否允许 NULL"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
for col in inspector.get_columns(table_name):
|
||||
if col["name"] == column_name:
|
||||
return col["nullable"]
|
||||
return False
|
||||
|
||||
|
||||
def enum_value_exists(enum_name: str, value: str) -> bool:
|
||||
"""检查 PostgreSQL ENUM 是否包含指定值"""
|
||||
bind = op.get_bind()
|
||||
if bind.dialect.name != "postgresql":
|
||||
return True # 非 PostgreSQL 跳过检查
|
||||
result = bind.execute(
|
||||
sa.text(
|
||||
"SELECT 1 FROM pg_enum WHERE enumlabel = :value "
|
||||
"AND enumtypid = (SELECT oid FROM pg_type WHERE typname = :enum_name)"
|
||||
),
|
||||
{"value": value, "enum_name": enum_name},
|
||||
).first()
|
||||
return result is not None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""应用迁移:升级到新版本"""
|
||||
bind = op.get_bind()
|
||||
|
||||
# ========== Part 1: users 表修改 ==========
|
||||
|
||||
# 1) 新增 email_verified
|
||||
if not column_exists("users", "email_verified"):
|
||||
op.add_column("users", sa.Column("email_verified", sa.Boolean(), nullable=True))
|
||||
# 历史数据回填:已有邮箱的用户默认视为已验证
|
||||
op.execute(sa.text("UPDATE users SET email_verified = true WHERE email IS NOT NULL"))
|
||||
op.execute(sa.text("UPDATE users SET email_verified = false WHERE email IS NULL"))
|
||||
# 收紧约束
|
||||
op.alter_column("users", "email_verified", existing_type=sa.Boolean(), nullable=False)
|
||||
|
||||
# 2) email 放宽为可空
|
||||
if not column_is_nullable("users", "email"):
|
||||
op.alter_column(
|
||||
"users",
|
||||
"email",
|
||||
existing_type=sa.String(length=255),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
# 3) password_hash 放宽为可空
|
||||
if not column_is_nullable("users", "password_hash"):
|
||||
op.alter_column(
|
||||
"users",
|
||||
"password_hash",
|
||||
existing_type=sa.String(length=255),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
# ========== Part 2: OAuth 相关 ==========
|
||||
|
||||
# 4) 扩展 authsource enum
|
||||
if bind.dialect.name == "postgresql" and not enum_value_exists("authsource", "oauth"):
|
||||
ctx = op.get_context()
|
||||
with ctx.autocommit_block():
|
||||
op.execute("ALTER TYPE authsource ADD VALUE IF NOT EXISTS 'oauth'")
|
||||
|
||||
# 5) OAuth provider 配置表
|
||||
if not table_exists("oauth_providers"):
|
||||
op.create_table(
|
||||
"oauth_providers",
|
||||
sa.Column("provider_type", sa.String(length=50), primary_key=True),
|
||||
sa.Column("display_name", sa.String(length=100), nullable=False),
|
||||
sa.Column("client_id", sa.String(length=255), nullable=False),
|
||||
sa.Column("client_secret_encrypted", sa.Text(), nullable=True),
|
||||
sa.Column("authorization_url_override", sa.String(length=500), nullable=True),
|
||||
sa.Column("token_url_override", sa.String(length=500), nullable=True),
|
||||
sa.Column("userinfo_url_override", sa.String(length=500), nullable=True),
|
||||
sa.Column("scopes", sa.JSON(), nullable=True),
|
||||
sa.Column("redirect_uri", sa.String(length=500), nullable=False),
|
||||
sa.Column("frontend_callback_url", sa.String(length=500), nullable=False),
|
||||
sa.Column("attribute_mapping", sa.JSON(), nullable=True),
|
||||
sa.Column("extra_config", sa.JSON(), nullable=True),
|
||||
sa.Column(
|
||||
"is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("false")
|
||||
),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
)
|
||||
|
||||
# 6) 用户 OAuth 绑定关系表
|
||||
if not table_exists("user_oauth_links"):
|
||||
op.create_table(
|
||||
"user_oauth_links",
|
||||
sa.Column("id", sa.String(length=36), primary_key=True),
|
||||
sa.Column(
|
||||
"user_id",
|
||||
sa.String(length=36),
|
||||
sa.ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"provider_type",
|
||||
sa.String(length=50),
|
||||
sa.ForeignKey("oauth_providers.provider_type", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("provider_user_id", sa.String(length=255), nullable=False),
|
||||
sa.Column("provider_username", sa.String(length=255), nullable=True),
|
||||
sa.Column("provider_email", sa.String(length=255), nullable=True),
|
||||
sa.Column("extra_data", sa.JSON(), nullable=True),
|
||||
sa.Column(
|
||||
"linked_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.Column("last_login_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.UniqueConstraint(
|
||||
"provider_type", "provider_user_id", name="uq_oauth_provider_user"
|
||||
),
|
||||
sa.UniqueConstraint("user_id", "provider_type", name="uq_user_oauth_provider"),
|
||||
)
|
||||
op.create_index("ix_user_oauth_links_user_id", "user_oauth_links", ["user_id"])
|
||||
op.create_index(
|
||||
"ix_user_oauth_links_provider_type", "user_oauth_links", ["provider_type"]
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""回滚迁移:降级到旧版本"""
|
||||
bind = op.get_bind()
|
||||
|
||||
# ========== Part 2: OAuth 相关(先删除,因为有外键依赖) ==========
|
||||
|
||||
if table_exists("user_oauth_links"):
|
||||
op.drop_index("ix_user_oauth_links_provider_type", table_name="user_oauth_links")
|
||||
op.drop_index("ix_user_oauth_links_user_id", table_name="user_oauth_links")
|
||||
op.drop_table("user_oauth_links")
|
||||
|
||||
if table_exists("oauth_providers"):
|
||||
op.drop_table("oauth_providers")
|
||||
|
||||
# 注意:Postgres 不支持从 ENUM 删除值,authsource 不回退
|
||||
|
||||
# ========== Part 1: users 表修改 ==========
|
||||
|
||||
# 降级前检查:避免把包含 NULL 的列强制改回 NOT NULL
|
||||
has_null_email = bind.execute(
|
||||
sa.text("SELECT 1 FROM users WHERE email IS NULL LIMIT 1")
|
||||
).first()
|
||||
if has_null_email:
|
||||
raise RuntimeError("Cannot downgrade: users.email contains NULL values")
|
||||
|
||||
has_null_password = bind.execute(
|
||||
sa.text("SELECT 1 FROM users WHERE password_hash IS NULL LIMIT 1")
|
||||
).first()
|
||||
if has_null_password:
|
||||
raise RuntimeError("Cannot downgrade: users.password_hash contains NULL values")
|
||||
|
||||
# 恢复 NOT NULL 约束
|
||||
if column_is_nullable("users", "email"):
|
||||
op.alter_column(
|
||||
"users",
|
||||
"email",
|
||||
existing_type=sa.String(length=255),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
if column_is_nullable("users", "password_hash"):
|
||||
op.alter_column(
|
||||
"users",
|
||||
"password_hash",
|
||||
existing_type=sa.String(length=255),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
if column_exists("users", "email_verified"):
|
||||
op.drop_column("users", "email_verified")
|
||||
@@ -0,0 +1,65 @@
|
||||
"""add_stats_daily_provider_table
|
||||
|
||||
Revision ID: c868729753ad
|
||||
Revises: 33e347f97c0c
|
||||
Create Date: 2026-01-19 05:19:49.634662+00:00
|
||||
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = 'c868729753ad'
|
||||
down_revision = '33e347f97c0c'
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
"""检查索引是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
indexes = [idx['name'] for idx in inspector.get_indexes(table_name)]
|
||||
return index_name in indexes
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""应用迁移:升级到新版本"""
|
||||
if not table_exists('stats_daily_provider'):
|
||||
op.create_table(
|
||||
'stats_daily_provider',
|
||||
sa.Column('id', sa.String(length=36), nullable=False),
|
||||
sa.Column('date', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('provider_name', sa.String(length=100), nullable=False),
|
||||
sa.Column('total_requests', sa.Integer(), nullable=False),
|
||||
sa.Column('input_tokens', sa.BigInteger(), nullable=False),
|
||||
sa.Column('output_tokens', sa.BigInteger(), nullable=False),
|
||||
sa.Column('cache_creation_tokens', sa.BigInteger(), nullable=False),
|
||||
sa.Column('cache_read_tokens', sa.BigInteger(), nullable=False),
|
||||
sa.Column('total_cost', sa.Float(), nullable=False),
|
||||
sa.Column('created_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column('updated_at', sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint('id'),
|
||||
sa.UniqueConstraint('date', 'provider_name', name='uq_stats_daily_provider')
|
||||
)
|
||||
op.create_index('idx_stats_daily_provider_date', 'stats_daily_provider', ['date'], unique=False)
|
||||
op.create_index('idx_stats_daily_provider_date_provider', 'stats_daily_provider', ['date', 'provider_name'], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""回滚迁移:降级到旧版本"""
|
||||
if table_exists('stats_daily_provider'):
|
||||
if index_exists('stats_daily_provider', 'idx_stats_daily_provider_date_provider'):
|
||||
op.drop_index('idx_stats_daily_provider_date_provider', table_name='stats_daily_provider')
|
||||
if index_exists('stats_daily_provider', 'idx_stats_daily_provider_date'):
|
||||
op.drop_index('idx_stats_daily_provider_date', table_name='stats_daily_provider')
|
||||
op.drop_table('stats_daily_provider')
|
||||
+51
@@ -0,0 +1,51 @@
|
||||
"""add_format_acceptance_config_to_provider_endpoints
|
||||
|
||||
Revision ID: 4b4c7b0df1a2
|
||||
Revises: c868729753ad
|
||||
Create Date: 2026-01-21 18:45:00+00:00
|
||||
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "4b4c7b0df1a2"
|
||||
down_revision = "c868729753ad"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not table_exists("provider_endpoints"):
|
||||
return
|
||||
if column_exists("provider_endpoints", "format_acceptance_config"):
|
||||
return
|
||||
op.add_column(
|
||||
"provider_endpoints",
|
||||
sa.Column("format_acceptance_config", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if not table_exists("provider_endpoints"):
|
||||
return
|
||||
if not column_exists("provider_endpoints", "format_acceptance_config"):
|
||||
return
|
||||
op.drop_column("provider_endpoints", "format_acceptance_config")
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
"""add_format_conversion_tracking_and_model_filter_patterns_and_provider_timeout
|
||||
|
||||
Revision ID: f7c8d9e0a1b2
|
||||
Revises: 4b4c7b0df1a2
|
||||
Create Date: 2026-01-27 10:00:00+00:00
|
||||
|
||||
Changes:
|
||||
1. usage 表: 添加 endpoint_api_format 和 has_format_conversion 字段
|
||||
2. provider_api_keys 表: 添加 model_include_patterns 和 model_exclude_patterns 字段
|
||||
- 支持通配符规则自动过滤从上游获取的模型列表
|
||||
- 包含规则和排除规则(支持 * 和 ? 通配符)
|
||||
3. providers 表: 添加 stream_first_byte_timeout 和 request_timeout 字段
|
||||
- 允许每个提供商单独配置超时时间
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "f7c8d9e0a1b2"
|
||||
down_revision = "4b4c7b0df1a2"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# === usage 表: 格式转换追踪 ===
|
||||
if table_exists("usage"):
|
||||
# 添加 endpoint_api_format 字段(端点原生 API 格式)
|
||||
if not column_exists("usage", "endpoint_api_format"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column("endpoint_api_format", sa.String(50), nullable=True),
|
||||
)
|
||||
|
||||
# 添加 has_format_conversion 字段(是否发生了格式转换)
|
||||
if not column_exists("usage", "has_format_conversion"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column("has_format_conversion", sa.Boolean(), nullable=True, server_default="false"),
|
||||
)
|
||||
|
||||
# === provider_api_keys 表: 模型过滤规则 ===
|
||||
if table_exists("provider_api_keys"):
|
||||
# 添加 model_include_patterns 字段(包含规则,支持 * 和 ? 通配符)
|
||||
if not column_exists("provider_api_keys", "model_include_patterns"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("model_include_patterns", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# 添加 model_exclude_patterns 字段(排除规则,支持 * 和 ? 通配符)
|
||||
if not column_exists("provider_api_keys", "model_exclude_patterns"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("model_exclude_patterns", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# === providers 表: 超时配置 ===
|
||||
if table_exists("providers"):
|
||||
# 添加 stream_first_byte_timeout 字段(流式请求首字节超时)
|
||||
if not column_exists("providers", "stream_first_byte_timeout"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("stream_first_byte_timeout", sa.Float(), nullable=True),
|
||||
)
|
||||
|
||||
# 添加 request_timeout 字段(非流式请求整体超时)
|
||||
if not column_exists("providers", "request_timeout"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("request_timeout", sa.Float(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# === providers 表: 移除超时配置 ===
|
||||
if table_exists("providers"):
|
||||
if column_exists("providers", "request_timeout"):
|
||||
op.drop_column("providers", "request_timeout")
|
||||
|
||||
if column_exists("providers", "stream_first_byte_timeout"):
|
||||
op.drop_column("providers", "stream_first_byte_timeout")
|
||||
|
||||
# === provider_api_keys 表: 移除模型过滤规则 ===
|
||||
if table_exists("provider_api_keys"):
|
||||
if column_exists("provider_api_keys", "model_exclude_patterns"):
|
||||
op.drop_column("provider_api_keys", "model_exclude_patterns")
|
||||
|
||||
if column_exists("provider_api_keys", "model_include_patterns"):
|
||||
op.drop_column("provider_api_keys", "model_include_patterns")
|
||||
|
||||
# === usage 表: 移除格式转换追踪 ===
|
||||
if table_exists("usage"):
|
||||
if column_exists("usage", "has_format_conversion"):
|
||||
op.drop_column("usage", "has_format_conversion")
|
||||
|
||||
if column_exists("usage", "endpoint_api_format"):
|
||||
op.drop_column("usage", "endpoint_api_format")
|
||||
@@ -0,0 +1,58 @@
|
||||
"""add_keep_priority_on_conversion_to_providers
|
||||
|
||||
Revision ID: 364680d1bc99
|
||||
Revises: f7c8d9e0a1b2
|
||||
Create Date: 2026-01-28 12:00:00+00:00
|
||||
|
||||
Changes:
|
||||
1. providers 表: 添加 keep_priority_on_conversion 字段
|
||||
- 格式转换时是否保持提供商原优先级
|
||||
- 默认 False:需要格式转换时,候选会被降级到不需要转换的候选之后
|
||||
- 设为 True:即使需要格式转换,也保持原优先级排名
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "364680d1bc99"
|
||||
down_revision = "f7c8d9e0a1b2"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# === providers 表: 添加格式转换优先级保持配置 ===
|
||||
if table_exists("providers"):
|
||||
if not column_exists("providers", "keep_priority_on_conversion"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column(
|
||||
"keep_priority_on_conversion",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default="false",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# === providers 表: 移除格式转换优先级保持配置 ===
|
||||
if table_exists("providers"):
|
||||
if column_exists("providers", "keep_priority_on_conversion"):
|
||||
op.drop_column("providers", "keep_priority_on_conversion")
|
||||
@@ -0,0 +1,51 @@
|
||||
"""Add auth_type and auth_config fields to provider_api_keys table
|
||||
|
||||
Revision ID: 7f6f8065f517
|
||||
Revises: 364680d1bc99
|
||||
Create Date: 2026-01-30 10:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy import inspect
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "7f6f8065f517"
|
||||
down_revision: Union[str, None] = "364680d1bc99"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否已存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 添加 auth_type 字段,默认值为 "api_key"
|
||||
if not column_exists("provider_api_keys", "auth_type"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("auth_type", sa.String(20), nullable=False, server_default="api_key"),
|
||||
)
|
||||
|
||||
# 添加 auth_config 字段(Text,存储加密后的认证配置)
|
||||
if not column_exists("provider_api_keys", "auth_config"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("auth_config", sa.Text, nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if column_exists("provider_api_keys", "auth_config"):
|
||||
op.drop_column("provider_api_keys", "auth_config")
|
||||
|
||||
if column_exists("provider_api_keys", "auth_type"):
|
||||
op.drop_column("provider_api_keys", "auth_type")
|
||||
@@ -0,0 +1,112 @@
|
||||
"""Add video_tasks table
|
||||
|
||||
Revision ID: b6f1a2c5d8e9
|
||||
Revises: 7f6f8065f517
|
||||
Create Date: 2026-01-30 18:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b6f1a2c5d8e9"
|
||||
down_revision: Union[str, None] = "7f6f8065f517"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if table_exists("video_tasks"):
|
||||
return
|
||||
|
||||
op.create_table(
|
||||
"video_tasks",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("external_task_id", sa.String(200), nullable=True, index=False),
|
||||
sa.Column("user_id", sa.String(36), sa.ForeignKey("users.id"), nullable=False),
|
||||
sa.Column("api_key_id", sa.String(36), sa.ForeignKey("api_keys.id"), nullable=True),
|
||||
sa.Column("provider_id", sa.String(36), sa.ForeignKey("providers.id"), nullable=True),
|
||||
sa.Column(
|
||||
"endpoint_id", sa.String(36), sa.ForeignKey("provider_endpoints.id"), nullable=True
|
||||
),
|
||||
sa.Column("key_id", sa.String(36), sa.ForeignKey("provider_api_keys.id"), nullable=True),
|
||||
sa.Column("client_api_format", sa.String(50), nullable=False),
|
||||
sa.Column("provider_api_format", sa.String(50), nullable=False),
|
||||
sa.Column("format_converted", sa.Boolean(), server_default=sa.false()),
|
||||
sa.Column("model", sa.String(100), nullable=False),
|
||||
sa.Column("prompt", sa.Text(), nullable=False),
|
||||
sa.Column("original_request_body", sa.JSON(), nullable=True),
|
||||
sa.Column("converted_request_body", sa.JSON(), nullable=True),
|
||||
sa.Column("duration_seconds", sa.Integer(), server_default=sa.text("4")),
|
||||
sa.Column("resolution", sa.String(20), server_default=sa.text("'720p'")),
|
||||
sa.Column("aspect_ratio", sa.String(10), server_default=sa.text("'16:9'")),
|
||||
sa.Column("size", sa.String(20), nullable=True),
|
||||
sa.Column("status", sa.String(20), server_default=sa.text("'pending'")),
|
||||
sa.Column("progress_percent", sa.Integer(), server_default=sa.text("0")),
|
||||
sa.Column("progress_message", sa.String(500), nullable=True),
|
||||
sa.Column("video_url", sa.String(2000), nullable=True),
|
||||
sa.Column("video_urls", sa.JSON(), nullable=True),
|
||||
sa.Column("thumbnail_url", sa.String(2000), nullable=True),
|
||||
sa.Column("video_size_bytes", sa.BigInteger(), nullable=True),
|
||||
sa.Column("video_expires_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("stored_video_path", sa.String(500), nullable=True),
|
||||
sa.Column("storage_provider", sa.String(50), nullable=True),
|
||||
sa.Column("error_code", sa.String(50), nullable=True),
|
||||
sa.Column("error_message", sa.Text(), nullable=True),
|
||||
sa.Column("retry_count", sa.Integer(), server_default=sa.text("0")),
|
||||
sa.Column("max_retries", sa.Integer(), server_default=sa.text("3")),
|
||||
sa.Column("poll_interval_seconds", sa.Integer(), server_default=sa.text("10")),
|
||||
sa.Column("next_poll_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("poll_count", sa.Integer(), server_default=sa.text("0")),
|
||||
sa.Column("max_poll_count", sa.Integer(), server_default=sa.text("360")),
|
||||
sa.Column(
|
||||
"remixed_from_task_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("video_tasks.id", ondelete="SET NULL"),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.Column("submitted_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
)
|
||||
|
||||
op.create_index("idx_video_tasks_user_status", "video_tasks", ["user_id", "status"])
|
||||
op.create_index("idx_video_tasks_next_poll", "video_tasks", ["next_poll_at"])
|
||||
op.create_index("idx_video_tasks_external_id", "video_tasks", ["external_task_id"])
|
||||
# 唯一约束:同一用户不能有重复的 external_task_id
|
||||
op.create_unique_constraint(
|
||||
"uq_video_tasks_user_external_id",
|
||||
"video_tasks",
|
||||
["user_id", "external_task_id"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if not table_exists("video_tasks"):
|
||||
return
|
||||
|
||||
op.drop_constraint("uq_video_tasks_user_external_id", "video_tasks", type_="unique")
|
||||
op.drop_index("idx_video_tasks_external_id", table_name="video_tasks")
|
||||
op.drop_index("idx_video_tasks_next_poll", table_name="video_tasks")
|
||||
op.drop_index("idx_video_tasks_user_status", table_name="video_tasks")
|
||||
op.drop_table("video_tasks")
|
||||
+180
@@ -0,0 +1,180 @@
|
||||
"""Add billing system tables and video_tasks.request_metadata
|
||||
|
||||
Revision ID: c8d2e4f6a1b3
|
||||
Revises: b6f1a2c5d8e9
|
||||
Create Date: 2026-01-31 12:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c8d2e4f6a1b3"
|
||||
down_revision: Union[str, None] = "b6f1a2c5d8e9"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
try:
|
||||
indexes = inspector.get_indexes(table_name)
|
||||
except Exception:
|
||||
return False
|
||||
return any(idx.get("name") == index_name for idx in indexes)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ==================== video_tasks.request_metadata ====================
|
||||
if not column_exists("video_tasks", "request_metadata"):
|
||||
op.add_column(
|
||||
"video_tasks",
|
||||
sa.Column("request_metadata", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# ==================== billing_rules ====================
|
||||
if not table_exists("billing_rules"):
|
||||
op.create_table(
|
||||
"billing_rules",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column(
|
||||
"global_model_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("global_models.id", ondelete="CASCADE"),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column(
|
||||
"model_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("models.id", ondelete="CASCADE"),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column("name", sa.String(100), nullable=False),
|
||||
sa.Column("task_type", sa.String(20), nullable=False, server_default="chat"),
|
||||
sa.Column("expression", sa.Text(), nullable=False),
|
||||
sa.Column("variables", JSONB, nullable=False, server_default=sa.text("'{}'::jsonb")),
|
||||
sa.Column(
|
||||
"dimension_mappings", JSONB, nullable=False, server_default=sa.text("'{}'::jsonb")
|
||||
),
|
||||
sa.Column("is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("true")),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(global_model_id IS NOT NULL AND model_id IS NULL) OR "
|
||||
"(global_model_id IS NULL AND model_id IS NOT NULL)",
|
||||
name="chk_billing_rules_model_ref",
|
||||
),
|
||||
)
|
||||
|
||||
# Partial unique indexes for enabled rules
|
||||
if table_exists("billing_rules"):
|
||||
if not index_exists("billing_rules", "uq_billing_rules_global_model_task"):
|
||||
op.create_index(
|
||||
"uq_billing_rules_global_model_task",
|
||||
"billing_rules",
|
||||
["global_model_id", "task_type"],
|
||||
unique=True,
|
||||
postgresql_where=sa.text("is_enabled = TRUE AND global_model_id IS NOT NULL"),
|
||||
)
|
||||
if not index_exists("billing_rules", "uq_billing_rules_model_task"):
|
||||
op.create_index(
|
||||
"uq_billing_rules_model_task",
|
||||
"billing_rules",
|
||||
["model_id", "task_type"],
|
||||
unique=True,
|
||||
postgresql_where=sa.text("is_enabled = TRUE AND model_id IS NOT NULL"),
|
||||
)
|
||||
|
||||
# ==================== dimension_collectors ====================
|
||||
if not table_exists("dimension_collectors"):
|
||||
op.create_table(
|
||||
"dimension_collectors",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("api_format", sa.String(50), nullable=False),
|
||||
sa.Column("task_type", sa.String(20), nullable=False),
|
||||
sa.Column("dimension_name", sa.String(100), nullable=False),
|
||||
sa.Column("source_type", sa.String(20), nullable=False),
|
||||
sa.Column("source_path", sa.String(200), nullable=True),
|
||||
sa.Column("value_type", sa.String(20), nullable=False, server_default="float"),
|
||||
sa.Column("transform_expression", sa.Text(), nullable=True),
|
||||
sa.Column("default_value", sa.String(100), nullable=True),
|
||||
sa.Column("priority", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("is_enabled", sa.Boolean(), nullable=False, server_default=sa.text("true")),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"(source_type = 'computed' AND source_path IS NULL AND transform_expression IS NOT NULL) OR "
|
||||
"(source_type != 'computed' AND source_path IS NOT NULL)",
|
||||
name="chk_dimension_collectors_source_config",
|
||||
),
|
||||
)
|
||||
|
||||
if table_exists("dimension_collectors"):
|
||||
if not index_exists("dimension_collectors", "uq_dimension_collectors_enabled"):
|
||||
op.create_index(
|
||||
"uq_dimension_collectors_enabled",
|
||||
"dimension_collectors",
|
||||
["api_format", "task_type", "dimension_name", "priority"],
|
||||
unique=True,
|
||||
postgresql_where=sa.text("is_enabled = TRUE"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop in reverse order
|
||||
if table_exists("dimension_collectors"):
|
||||
if index_exists("dimension_collectors", "uq_dimension_collectors_enabled"):
|
||||
op.drop_index("uq_dimension_collectors_enabled", table_name="dimension_collectors")
|
||||
op.drop_table("dimension_collectors")
|
||||
|
||||
if table_exists("billing_rules"):
|
||||
if index_exists("billing_rules", "uq_billing_rules_model_task"):
|
||||
op.drop_index("uq_billing_rules_model_task", table_name="billing_rules")
|
||||
if index_exists("billing_rules", "uq_billing_rules_global_model_task"):
|
||||
op.drop_index("uq_billing_rules_global_model_task", table_name="billing_rules")
|
||||
op.drop_table("billing_rules")
|
||||
|
||||
if column_exists("video_tasks", "request_metadata"):
|
||||
op.drop_column("video_tasks", "request_metadata")
|
||||
+462
@@ -0,0 +1,462 @@
|
||||
"""Add api_family/endpoint_kind and migrate api_format to endpoint signature keys
|
||||
|
||||
Revision ID: cf40e6a5c5b1
|
||||
Revises: c8d2e4f6a1b3
|
||||
Create Date: 2026-01-31 15:30:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import Sequence, Union
|
||||
from uuid import uuid4
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect, text
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "cf40e6a5c5b1"
|
||||
down_revision: Union[str, None] = "c8d2e4f6a1b3"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _json_loads(val):
|
||||
if val is None:
|
||||
return None
|
||||
if isinstance(val, (dict, list)):
|
||||
return val
|
||||
if isinstance(val, str):
|
||||
try:
|
||||
return json.loads(val)
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _json_dumps(val):
|
||||
"""将 dict/list 转为 JSON 字符串,None 保持 None"""
|
||||
if val is None:
|
||||
return None
|
||||
if isinstance(val, str):
|
||||
return val
|
||||
return json.dumps(val)
|
||||
|
||||
|
||||
def _normalize_signature(value: str | None) -> str | None:
|
||||
"""
|
||||
Normalize legacy api_format / signature-ish strings to canonical signature key.
|
||||
|
||||
- canonical: `<family>:<kind>` (lowercase)
|
||||
- legacy examples: "OPENAI", "OPENAI_CLI", "GEMINI_VIDEO"
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
raw = str(value).strip()
|
||||
if not raw:
|
||||
return None
|
||||
|
||||
if ":" in raw:
|
||||
fam, kind = raw.split(":", 1)
|
||||
fam = fam.strip().lower()
|
||||
kind = kind.strip().lower()
|
||||
if not fam or not kind:
|
||||
return None
|
||||
return f"{fam}:{kind}"
|
||||
|
||||
upper = raw.upper()
|
||||
if upper.startswith("CLAUDE"):
|
||||
fam = "claude"
|
||||
elif upper.startswith("OPENAI"):
|
||||
fam = "openai"
|
||||
elif upper.startswith("GEMINI"):
|
||||
fam = "gemini"
|
||||
else:
|
||||
return None
|
||||
|
||||
kind = "chat"
|
||||
if upper.endswith("_CLI"):
|
||||
kind = "cli"
|
||||
elif upper.endswith("_VIDEO"):
|
||||
kind = "video"
|
||||
|
||||
return f"{fam}:{kind}"
|
||||
|
||||
|
||||
def _normalize_signature_list(values) -> list[str] | None:
|
||||
if values is None:
|
||||
return None
|
||||
if isinstance(values, str):
|
||||
values = _json_loads(values)
|
||||
if not isinstance(values, list):
|
||||
return None
|
||||
|
||||
out: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for v in values:
|
||||
sig = _normalize_signature(str(v) if v is not None else None)
|
||||
if not sig:
|
||||
continue
|
||||
if sig in seen:
|
||||
continue
|
||||
seen.add(sig)
|
||||
out.append(sig)
|
||||
return out
|
||||
|
||||
|
||||
def _normalize_signature_dict(values) -> dict | None:
|
||||
if values is None:
|
||||
return None
|
||||
if isinstance(values, str):
|
||||
values = _json_loads(values)
|
||||
if not isinstance(values, dict):
|
||||
return None
|
||||
|
||||
out: dict = {}
|
||||
for k, v in values.items():
|
||||
sig = _normalize_signature(str(k) if k is not None else None)
|
||||
if not sig:
|
||||
continue
|
||||
out[sig] = v
|
||||
return out
|
||||
|
||||
|
||||
def _add_video_variants(formats: list[str]) -> list[str]:
|
||||
"""
|
||||
保持原有格式,不自动补齐 video 变体。
|
||||
"""
|
||||
return formats
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
try:
|
||||
indexes = inspector.get_indexes(table_name)
|
||||
except Exception:
|
||||
return False
|
||||
return any(idx.get("name") == index_name for idx in indexes)
|
||||
|
||||
|
||||
def _migrate_format_acceptance_config(cfg) -> dict | None:
|
||||
cfg_obj = _json_loads(cfg)
|
||||
if not isinstance(cfg_obj, dict):
|
||||
return cfg_obj if cfg_obj is None else None
|
||||
|
||||
for key in ("accept_formats", "reject_formats"):
|
||||
raw = cfg_obj.get(key)
|
||||
if not isinstance(raw, list):
|
||||
continue
|
||||
normalized = _normalize_signature_list(raw) or []
|
||||
cfg_obj[key] = normalized
|
||||
|
||||
return cfg_obj
|
||||
|
||||
|
||||
def migrate_provider_endpoints(connection) -> None:
|
||||
"""
|
||||
- 将 provider_endpoints.api_format 统一迁移为 signature key(小写)
|
||||
- 填充/校准 api_family / endpoint_kind
|
||||
- 迁移 format_acceptance_config 中的 accept/reject formats
|
||||
"""
|
||||
rows = connection.execute(text("""
|
||||
SELECT
|
||||
id,
|
||||
api_format,
|
||||
api_family,
|
||||
endpoint_kind,
|
||||
format_acceptance_config
|
||||
FROM provider_endpoints
|
||||
""")).fetchall()
|
||||
|
||||
for row in rows:
|
||||
sig = _normalize_signature(row.api_format)
|
||||
if not sig:
|
||||
continue
|
||||
fam, kind = sig.split(":", 1)
|
||||
|
||||
cfg = _migrate_format_acceptance_config(row.format_acceptance_config)
|
||||
|
||||
connection.execute(
|
||||
text("""
|
||||
UPDATE provider_endpoints
|
||||
SET
|
||||
api_format = :api_format,
|
||||
api_family = :api_family,
|
||||
endpoint_kind = :endpoint_kind,
|
||||
format_acceptance_config = CAST(:format_acceptance_config AS json)
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{
|
||||
"id": row.id,
|
||||
"api_format": sig,
|
||||
"api_family": fam,
|
||||
"endpoint_kind": kind,
|
||||
"format_acceptance_config": _json_dumps(cfg),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def create_video_endpoints(connection) -> None:
|
||||
"""
|
||||
不再自动创建 video endpoint,保持原有配置。
|
||||
"""
|
||||
pass
|
||||
|
||||
|
||||
def migrate_provider_api_keys(connection) -> None:
|
||||
"""
|
||||
迁移 provider_api_keys:
|
||||
- api_formats -> signature keys(并补齐 video 变体)
|
||||
- dict 字段 key -> signature keys(rate_multipliers/global_priority/health/circuit_breaker)
|
||||
- rate_multipliers/global_priority_by_format 复制 chat -> video(如 openai:chat -> openai:video)
|
||||
"""
|
||||
rows = connection.execute(text("""
|
||||
SELECT
|
||||
id,
|
||||
api_formats,
|
||||
rate_multipliers,
|
||||
global_priority_by_format,
|
||||
health_by_format,
|
||||
circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
""")).fetchall()
|
||||
|
||||
for row in rows:
|
||||
api_formats = _normalize_signature_list(row.api_formats)
|
||||
if api_formats is not None:
|
||||
api_formats = _add_video_variants(api_formats)
|
||||
|
||||
rate_multipliers = _normalize_signature_dict(row.rate_multipliers)
|
||||
global_priority_by_format = _normalize_signature_dict(row.global_priority_by_format)
|
||||
|
||||
health_by_format = _normalize_signature_dict(row.health_by_format)
|
||||
circuit_breaker_by_format = _normalize_signature_dict(row.circuit_breaker_by_format)
|
||||
|
||||
connection.execute(
|
||||
text("""
|
||||
UPDATE provider_api_keys
|
||||
SET
|
||||
api_formats = CAST(:api_formats AS json),
|
||||
rate_multipliers = CAST(:rate_multipliers AS json),
|
||||
global_priority_by_format = CAST(:global_priority_by_format AS json),
|
||||
health_by_format = CAST(:health_by_format AS json),
|
||||
circuit_breaker_by_format = CAST(:circuit_breaker_by_format AS json)
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{
|
||||
"id": row.id,
|
||||
"api_formats": _json_dumps(api_formats),
|
||||
"rate_multipliers": _json_dumps(rate_multipliers),
|
||||
"global_priority_by_format": _json_dumps(global_priority_by_format),
|
||||
"health_by_format": _json_dumps(health_by_format),
|
||||
"circuit_breaker_by_format": _json_dumps(circuit_breaker_by_format),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def migrate_allowed_api_formats(connection, *, table_name: str) -> None:
|
||||
"""迁移 users/api_keys.allowed_api_formats 为 signature keys(并补齐 video 变体)。"""
|
||||
if not table_exists(table_name):
|
||||
return
|
||||
rows = connection.execute(text(f"""
|
||||
SELECT id, allowed_api_formats
|
||||
FROM {table_name}
|
||||
""")).fetchall()
|
||||
|
||||
for row in rows:
|
||||
allowed = _normalize_signature_list(row.allowed_api_formats)
|
||||
if allowed is None:
|
||||
continue
|
||||
allowed = _add_video_variants(allowed)
|
||||
connection.execute(
|
||||
text(f"""
|
||||
UPDATE {table_name}
|
||||
SET allowed_api_formats = CAST(:allowed_api_formats AS json)
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{"id": row.id, "allowed_api_formats": _json_dumps(allowed)},
|
||||
)
|
||||
|
||||
|
||||
def migrate_video_tasks(connection) -> None:
|
||||
"""
|
||||
video_tasks.*_api_format 迁移为 signature keys。
|
||||
|
||||
注意:video_tasks 表天然是 video 任务,因此将 openai/gemini 的 kind 强制归一为 video,
|
||||
以兼容历史上复用 chat 格式存储的旧记录。
|
||||
"""
|
||||
if not table_exists("video_tasks"):
|
||||
return
|
||||
|
||||
rows = connection.execute(text("""
|
||||
SELECT id, client_api_format, provider_api_format
|
||||
FROM video_tasks
|
||||
""")).fetchall()
|
||||
|
||||
for row in rows:
|
||||
client_sig = _normalize_signature(row.client_api_format) or ""
|
||||
provider_sig = _normalize_signature(row.provider_api_format) or ""
|
||||
|
||||
def _force_video(sig: str) -> str:
|
||||
if not sig or ":" not in sig:
|
||||
return sig
|
||||
fam, _kind = sig.split(":", 1)
|
||||
fam = fam.strip().lower()
|
||||
if fam in ("openai", "gemini"):
|
||||
return f"{fam}:video"
|
||||
return sig
|
||||
|
||||
client_sig = _force_video(client_sig)
|
||||
provider_sig = _force_video(provider_sig)
|
||||
|
||||
if not client_sig or not provider_sig:
|
||||
continue
|
||||
|
||||
connection.execute(
|
||||
text("""
|
||||
UPDATE video_tasks
|
||||
SET client_api_format = :client_api_format,
|
||||
provider_api_format = :provider_api_format
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{
|
||||
"id": row.id,
|
||||
"client_api_format": client_sig,
|
||||
"provider_api_format": provider_sig,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def migrate_model_provider_mappings(connection) -> None:
|
||||
"""迁移 models.provider_model_mappings[*].api_formats 为 signature keys。"""
|
||||
if not table_exists("models"):
|
||||
return
|
||||
|
||||
rows = connection.execute(text("""
|
||||
SELECT id, provider_model_mappings
|
||||
FROM models
|
||||
WHERE provider_model_mappings IS NOT NULL
|
||||
""")).fetchall()
|
||||
|
||||
for row in rows:
|
||||
mappings = _json_loads(row.provider_model_mappings)
|
||||
if not isinstance(mappings, list):
|
||||
continue
|
||||
|
||||
changed = False
|
||||
new_mappings: list = []
|
||||
for item in mappings:
|
||||
if not isinstance(item, dict):
|
||||
new_mappings.append(item)
|
||||
continue
|
||||
raw_formats = item.get("api_formats")
|
||||
if isinstance(raw_formats, list):
|
||||
normalized = _normalize_signature_list(raw_formats) or []
|
||||
# 内容比较(而非引用比较),避免已迁移数据被无意义地重复 UPDATE
|
||||
if set(normalized) != set(raw_formats):
|
||||
changed = True
|
||||
item = dict(item)
|
||||
item["api_formats"] = normalized
|
||||
new_mappings.append(item)
|
||||
|
||||
if not changed:
|
||||
continue
|
||||
|
||||
connection.execute(
|
||||
text("""
|
||||
UPDATE models
|
||||
SET provider_model_mappings = CAST(:provider_model_mappings AS json)
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{"id": row.id, "provider_model_mappings": _json_dumps(new_mappings)},
|
||||
)
|
||||
|
||||
|
||||
def migrate_dimension_collectors(connection) -> None:
|
||||
"""迁移 dimension_collectors.api_format 为 signature keys(如果存在历史数据)。"""
|
||||
if not table_exists("dimension_collectors"):
|
||||
return
|
||||
|
||||
rows = connection.execute(text("""
|
||||
SELECT id, api_format
|
||||
FROM dimension_collectors
|
||||
WHERE api_format IS NOT NULL
|
||||
""")).fetchall()
|
||||
|
||||
for row in rows:
|
||||
sig = _normalize_signature(row.api_format)
|
||||
if not sig:
|
||||
continue
|
||||
connection.execute(
|
||||
text("""
|
||||
UPDATE dimension_collectors
|
||||
SET api_format = :api_format
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{"id": row.id, "api_format": sig},
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not table_exists("provider_endpoints"):
|
||||
return
|
||||
|
||||
# ==================== provider_endpoints.api_family / endpoint_kind ====================
|
||||
if not column_exists("provider_endpoints", "api_family"):
|
||||
op.add_column("provider_endpoints", sa.Column("api_family", sa.String(50), nullable=True))
|
||||
if not column_exists("provider_endpoints", "endpoint_kind"):
|
||||
op.add_column(
|
||||
"provider_endpoints", sa.Column("endpoint_kind", sa.String(50), nullable=True)
|
||||
)
|
||||
|
||||
# ==================== idx_provider_family_kind ====================
|
||||
if not index_exists("provider_endpoints", "idx_provider_family_kind"):
|
||||
op.create_index(
|
||||
"idx_provider_family_kind",
|
||||
"provider_endpoints",
|
||||
["provider_id", "api_family", "endpoint_kind"],
|
||||
)
|
||||
|
||||
# ==================== data migrations (idempotent) ====================
|
||||
conn = op.get_bind()
|
||||
|
||||
migrate_provider_endpoints(conn)
|
||||
create_video_endpoints(conn)
|
||||
|
||||
if table_exists("provider_api_keys"):
|
||||
migrate_provider_api_keys(conn)
|
||||
|
||||
migrate_allowed_api_formats(conn, table_name="users")
|
||||
migrate_allowed_api_formats(conn, table_name="api_keys")
|
||||
migrate_video_tasks(conn)
|
||||
migrate_model_provider_mappings(conn)
|
||||
migrate_dimension_collectors(conn)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Drop index/columns only; data changes are intentionally kept (safe rollback strategy).
|
||||
if table_exists("provider_endpoints"):
|
||||
if index_exists("provider_endpoints", "idx_provider_family_kind"):
|
||||
op.drop_index("idx_provider_family_kind", table_name="provider_endpoints")
|
||||
if column_exists("provider_endpoints", "endpoint_kind"):
|
||||
op.drop_column("provider_endpoints", "endpoint_kind")
|
||||
if column_exists("provider_endpoints", "api_family"):
|
||||
op.drop_column("provider_endpoints", "api_family")
|
||||
+329
@@ -0,0 +1,329 @@
|
||||
"""Add usage billing, video_tasks fields, gemini_file_mappings, provider format conversion, and indexes
|
||||
|
||||
Revision ID: a2f1b3c4d5e6
|
||||
Revises: cf40e6a5c5b1
|
||||
Create Date: 2026-02-01 12:00:00+00:00
|
||||
|
||||
Changes:
|
||||
1. usage 表:
|
||||
- 添加 billing_status (pending/settled/void),用于表示结算状态
|
||||
- 添加 finalized_at,用于记录结算完成时间
|
||||
- 添加 (provider_name, created_at) 和 (model, created_at) 索引
|
||||
|
||||
2. video_tasks 表:
|
||||
- 添加 request_id(全局唯一),用于与 Usage/RequestCandidate 建立稳定关联
|
||||
- 添加 short_id (Gemini-style short ID)
|
||||
|
||||
3. gemini_file_mappings 表:
|
||||
- 创建新表用于文件映射
|
||||
- 添加 source_hash 字段用于关联相同源文件
|
||||
|
||||
4. providers 表:
|
||||
- 添加 enable_format_conversion 开关字段
|
||||
|
||||
5. request_candidates 表:
|
||||
- 添加 created_at 索引
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import string
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect, text
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "a2f1b3c4d5e6"
|
||||
down_revision: Union[str, None] = "cf40e6a5c5b1"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
inspector.clear_cache()
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
# Clear cached schema info to get fresh data
|
||||
inspector.clear_cache()
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
inspector.clear_cache()
|
||||
indexes = inspector.get_indexes(table_name)
|
||||
return any(idx.get("name") == index_name for idx in indexes)
|
||||
|
||||
|
||||
def unique_constraint_exists(table_name: str, constraint_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
inspector.clear_cache()
|
||||
constraints = inspector.get_unique_constraints(table_name)
|
||||
return any(c.get("name") == constraint_name for c in constraints)
|
||||
|
||||
|
||||
def generate_short_id(length: int = 12) -> str:
|
||||
"""Generate a Gemini-style short ID (lowercase letters + digits)"""
|
||||
alphabet = string.ascii_lowercase + string.digits
|
||||
return "".join(secrets.choice(alphabet) for _ in range(length))
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
dialect = bind.dialect.name
|
||||
|
||||
# =========================================================================
|
||||
# 1. usage 表: billing_status + finalized_at + 索引
|
||||
# =========================================================================
|
||||
if table_exists("usage"):
|
||||
if not column_exists("usage", "billing_status"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column(
|
||||
"billing_status",
|
||||
sa.String(20),
|
||||
nullable=False,
|
||||
server_default="settled",
|
||||
),
|
||||
)
|
||||
|
||||
if not column_exists("usage", "finalized_at"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column("finalized_at", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
|
||||
if not index_exists("usage", "idx_usage_billing_status"):
|
||||
op.create_index("idx_usage_billing_status", "usage", ["billing_status"])
|
||||
|
||||
# (provider_name, created_at) — provider list / dashboard queries
|
||||
if (
|
||||
column_exists("usage", "provider_name")
|
||||
and column_exists("usage", "created_at")
|
||||
and not index_exists("usage", "idx_usage_provider_created")
|
||||
):
|
||||
op.create_index("idx_usage_provider_created", "usage", ["provider_name", "created_at"])
|
||||
|
||||
# (model, created_at) — model analytics / recent requests queries
|
||||
if (
|
||||
column_exists("usage", "model")
|
||||
and column_exists("usage", "created_at")
|
||||
and not index_exists("usage", "idx_usage_model_created")
|
||||
):
|
||||
op.create_index("idx_usage_model_created", "usage", ["model", "created_at"])
|
||||
|
||||
# =========================================================================
|
||||
# 2. video_tasks 表: request_id + short_id
|
||||
# =========================================================================
|
||||
if table_exists("video_tasks"):
|
||||
# --- request_id ---
|
||||
if not column_exists("video_tasks", "request_id"):
|
||||
op.add_column(
|
||||
"video_tasks",
|
||||
sa.Column("request_id", sa.String(100), nullable=True),
|
||||
)
|
||||
|
||||
# 回填 request_id
|
||||
if dialect == "postgresql":
|
||||
op.execute("""
|
||||
UPDATE video_tasks
|
||||
SET request_id = COALESCE(request_metadata->>'request_id', id)
|
||||
WHERE request_id IS NULL
|
||||
""")
|
||||
elif dialect == "sqlite":
|
||||
op.execute("""
|
||||
UPDATE video_tasks
|
||||
SET request_id = COALESCE(json_extract(request_metadata, '$.request_id'), id)
|
||||
WHERE request_id IS NULL
|
||||
""")
|
||||
else:
|
||||
op.execute("""
|
||||
UPDATE video_tasks
|
||||
SET request_id = id
|
||||
WHERE request_id IS NULL
|
||||
""")
|
||||
|
||||
if dialect == "postgresql":
|
||||
op.alter_column("video_tasks", "request_id", nullable=False)
|
||||
|
||||
if not index_exists("video_tasks", "idx_video_tasks_request_id"):
|
||||
op.create_index("idx_video_tasks_request_id", "video_tasks", ["request_id"])
|
||||
|
||||
if not unique_constraint_exists("video_tasks", "uq_video_tasks_request_id"):
|
||||
op.create_unique_constraint(
|
||||
"uq_video_tasks_request_id",
|
||||
"video_tasks",
|
||||
["request_id"],
|
||||
)
|
||||
|
||||
# --- short_id ---
|
||||
if not column_exists("video_tasks", "short_id"):
|
||||
op.add_column(
|
||||
"video_tasks",
|
||||
sa.Column("short_id", sa.String(16), nullable=True),
|
||||
)
|
||||
|
||||
# Populate existing rows with unique short_ids
|
||||
result = bind.execute(text("SELECT id FROM video_tasks WHERE short_id IS NULL"))
|
||||
for row in result:
|
||||
short_id = generate_short_id()
|
||||
bind.execute(
|
||||
text("UPDATE video_tasks SET short_id = :short_id WHERE id = :id"),
|
||||
{"short_id": short_id, "id": row[0]},
|
||||
)
|
||||
|
||||
op.alter_column("video_tasks", "short_id", nullable=False)
|
||||
op.create_index("ix_video_tasks_short_id", "video_tasks", ["short_id"], unique=True)
|
||||
|
||||
# =========================================================================
|
||||
# 3. gemini_file_mappings 表
|
||||
# =========================================================================
|
||||
if not table_exists("gemini_file_mappings"):
|
||||
op.create_table(
|
||||
"gemini_file_mappings",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("file_name", sa.String(255), nullable=False, unique=True),
|
||||
sa.Column(
|
||||
"key_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("provider_api_keys.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"user_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column("display_name", sa.String(255), nullable=True),
|
||||
sa.Column("mime_type", sa.String(100), nullable=True),
|
||||
sa.Column("source_hash", sa.String(64), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||
)
|
||||
|
||||
op.create_index("ix_gemini_file_mappings_id", "gemini_file_mappings", ["id"])
|
||||
op.create_index(
|
||||
"ix_gemini_file_mappings_file_name", "gemini_file_mappings", ["file_name"], unique=True
|
||||
)
|
||||
op.create_index("ix_gemini_file_mappings_key_id", "gemini_file_mappings", ["key_id"])
|
||||
op.create_index("ix_gemini_file_mappings_user_id", "gemini_file_mappings", ["user_id"])
|
||||
op.create_index("idx_gemini_file_mappings_expires", "gemini_file_mappings", ["expires_at"])
|
||||
op.create_index(
|
||||
"idx_gemini_file_mappings_source_hash", "gemini_file_mappings", ["source_hash"]
|
||||
)
|
||||
else:
|
||||
# 表已存在,只添加 source_hash
|
||||
if not column_exists("gemini_file_mappings", "source_hash"):
|
||||
op.add_column(
|
||||
"gemini_file_mappings",
|
||||
sa.Column("source_hash", sa.String(64), nullable=True),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_gemini_file_mappings_source_hash",
|
||||
"gemini_file_mappings",
|
||||
["source_hash"],
|
||||
)
|
||||
|
||||
# =========================================================================
|
||||
# 4. providers 表: enable_format_conversion
|
||||
# =========================================================================
|
||||
if table_exists("providers") and not column_exists("providers", "enable_format_conversion"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column(
|
||||
"enable_format_conversion",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
),
|
||||
)
|
||||
|
||||
# =========================================================================
|
||||
# 5. request_candidates 表: created_at 索引
|
||||
# =========================================================================
|
||||
if table_exists("request_candidates"):
|
||||
if not index_exists("request_candidates", "idx_request_candidates_created_at"):
|
||||
op.create_index(
|
||||
"idx_request_candidates_created_at",
|
||||
"request_candidates",
|
||||
["created_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
dialect = bind.dialect.name
|
||||
|
||||
# =========================================================================
|
||||
# 5. request_candidates 表回滚
|
||||
# =========================================================================
|
||||
if table_exists("request_candidates"):
|
||||
if index_exists("request_candidates", "idx_request_candidates_created_at"):
|
||||
op.drop_index("idx_request_candidates_created_at", table_name="request_candidates")
|
||||
|
||||
# =========================================================================
|
||||
# 4. providers 表回滚
|
||||
# =========================================================================
|
||||
if table_exists("providers") and column_exists("providers", "enable_format_conversion"):
|
||||
op.drop_column("providers", "enable_format_conversion")
|
||||
|
||||
# =========================================================================
|
||||
# 3. gemini_file_mappings 表回滚
|
||||
# =========================================================================
|
||||
if table_exists("gemini_file_mappings"):
|
||||
op.drop_index("idx_gemini_file_mappings_source_hash", table_name="gemini_file_mappings")
|
||||
op.drop_index("idx_gemini_file_mappings_expires", table_name="gemini_file_mappings")
|
||||
op.drop_index("ix_gemini_file_mappings_user_id", table_name="gemini_file_mappings")
|
||||
op.drop_index("ix_gemini_file_mappings_key_id", table_name="gemini_file_mappings")
|
||||
op.drop_index("ix_gemini_file_mappings_file_name", table_name="gemini_file_mappings")
|
||||
op.drop_index("ix_gemini_file_mappings_id", table_name="gemini_file_mappings")
|
||||
op.drop_table("gemini_file_mappings")
|
||||
|
||||
# =========================================================================
|
||||
# 2. video_tasks 表回滚
|
||||
# =========================================================================
|
||||
if table_exists("video_tasks"):
|
||||
# short_id
|
||||
if column_exists("video_tasks", "short_id"):
|
||||
if index_exists("video_tasks", "ix_video_tasks_short_id"):
|
||||
op.drop_index("ix_video_tasks_short_id", table_name="video_tasks")
|
||||
op.drop_column("video_tasks", "short_id")
|
||||
|
||||
# request_id
|
||||
if column_exists("video_tasks", "request_id"):
|
||||
if dialect == "postgresql":
|
||||
if unique_constraint_exists("video_tasks", "uq_video_tasks_request_id"):
|
||||
op.drop_constraint("uq_video_tasks_request_id", "video_tasks", type_="unique")
|
||||
if index_exists("video_tasks", "idx_video_tasks_request_id"):
|
||||
op.drop_index("idx_video_tasks_request_id", table_name="video_tasks")
|
||||
op.drop_column("video_tasks", "request_id")
|
||||
|
||||
# =========================================================================
|
||||
# 1. usage 表回滚
|
||||
# =========================================================================
|
||||
if table_exists("usage"):
|
||||
if index_exists("usage", "idx_usage_model_created"):
|
||||
op.drop_index("idx_usage_model_created", table_name="usage")
|
||||
if index_exists("usage", "idx_usage_provider_created"):
|
||||
op.drop_index("idx_usage_provider_created", table_name="usage")
|
||||
if index_exists("usage", "idx_usage_billing_status"):
|
||||
op.drop_index("idx_usage_billing_status", table_name="usage")
|
||||
if column_exists("usage", "finalized_at"):
|
||||
op.drop_column("usage", "finalized_at")
|
||||
if column_exists("usage", "billing_status"):
|
||||
op.drop_column("usage", "billing_status")
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Add video_duration_seconds to video_tasks and body_rules to provider_endpoints
|
||||
|
||||
Revision ID: b3c4d5e6f7a8
|
||||
Revises: a2f1b3c4d5e6
|
||||
Create Date: 2026-02-03 15:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b3c4d5e6f7a8"
|
||||
down_revision: Union[str, None] = "a2f1b3c4d5e6"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""Check if a column exists in a table."""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 1. Add video_duration_seconds to video_tasks
|
||||
if not _column_exists("video_tasks", "video_duration_seconds"):
|
||||
op.add_column(
|
||||
"video_tasks",
|
||||
sa.Column("video_duration_seconds", sa.Float(), nullable=True),
|
||||
)
|
||||
|
||||
# 2. Add body_rules to provider_endpoints
|
||||
# 请求体规则支持三种操作:
|
||||
# - set: 设置/覆盖字段 {"action": "set", "path": "metadata", "value": {"custom": "val"}}
|
||||
# - drop: 删除字段 {"action": "drop", "path": "unwanted_field"}
|
||||
# - rename: 重命名字段 {"action": "rename", "from": "old_key", "to": "new_key"}
|
||||
if not _column_exists("provider_endpoints", "body_rules"):
|
||||
op.add_column(
|
||||
"provider_endpoints",
|
||||
sa.Column("body_rules", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Remove body_rules from provider_endpoints
|
||||
if _column_exists("provider_endpoints", "body_rules"):
|
||||
op.drop_column("provider_endpoints", "body_rules")
|
||||
|
||||
# Remove video_duration_seconds from video_tasks
|
||||
if _column_exists("video_tasks", "video_duration_seconds"):
|
||||
op.drop_column("video_tasks", "video_duration_seconds")
|
||||
@@ -0,0 +1,347 @@
|
||||
"""add_stats_hourly_and_daily_complete_flag
|
||||
|
||||
Revision ID: c4e8f9a1b2c3
|
||||
Revises: b3c4d5e6f7a8
|
||||
Create Date: 2026-02-04 12:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c4e8f9a1b2c3"
|
||||
down_revision: Union[str, None] = "b3c4d5e6f7a8"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def _index_exists(table_name: str, index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
indexes = [idx["name"] for idx in inspector.get_indexes(table_name)]
|
||||
return index_name in indexes
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
# Use information_schema for more reliable detection (inspector can have caching issues)
|
||||
result = bind.execute(
|
||||
sa.text(
|
||||
"SELECT EXISTS ("
|
||||
"SELECT 1 FROM information_schema.columns "
|
||||
"WHERE table_name = :table AND column_name = :column"
|
||||
")"
|
||||
),
|
||||
{"table": table_name, "column": column_name},
|
||||
)
|
||||
return bool(result.scalar())
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if _table_exists("stats_daily"):
|
||||
if not _column_exists("stats_daily", "is_complete"):
|
||||
op.add_column(
|
||||
"stats_daily",
|
||||
sa.Column("is_complete", sa.Boolean(), nullable=False, server_default=sa.false()),
|
||||
)
|
||||
op.execute("UPDATE stats_daily SET is_complete = true")
|
||||
if not _column_exists("stats_daily", "aggregated_at"):
|
||||
op.add_column(
|
||||
"stats_daily",
|
||||
sa.Column("aggregated_at", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
|
||||
if not _table_exists("stats_hourly"):
|
||||
op.create_table(
|
||||
"stats_hourly",
|
||||
sa.Column("id", sa.String(length=36), nullable=False),
|
||||
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("total_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("success_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("error_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("cache_creation_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("cache_read_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("total_cost", sa.Float(), nullable=False),
|
||||
sa.Column("actual_total_cost", sa.Float(), nullable=False),
|
||||
sa.Column("avg_response_time_ms", sa.Float(), nullable=False),
|
||||
sa.Column("is_complete", sa.Boolean(), nullable=False),
|
||||
sa.Column("aggregated_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("hour_utc", name="uq_stats_hourly_hour"),
|
||||
)
|
||||
op.create_index("idx_stats_hourly_hour", "stats_hourly", ["hour_utc"], unique=False)
|
||||
|
||||
if not _table_exists("stats_hourly_user"):
|
||||
op.create_table(
|
||||
"stats_hourly_user",
|
||||
sa.Column("id", sa.String(length=36), nullable=False),
|
||||
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("user_id", sa.String(length=36), nullable=False),
|
||||
sa.Column("total_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("success_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("error_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("total_cost", sa.Float(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("hour_utc", "user_id", name="uq_stats_hourly_user"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_stats_hourly_user_hour", "stats_hourly_user", ["hour_utc"], unique=False
|
||||
)
|
||||
op.create_index(
|
||||
"idx_stats_hourly_user_user_hour",
|
||||
"stats_hourly_user",
|
||||
["user_id", "hour_utc"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
if not _table_exists("stats_hourly_model"):
|
||||
op.create_table(
|
||||
"stats_hourly_model",
|
||||
sa.Column("id", sa.String(length=36), nullable=False),
|
||||
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("model", sa.String(length=100), nullable=False),
|
||||
sa.Column("total_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("total_cost", sa.Float(), nullable=False),
|
||||
sa.Column("avg_response_time_ms", sa.Float(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("hour_utc", "model", name="uq_stats_hourly_model"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_stats_hourly_model_hour", "stats_hourly_model", ["hour_utc"], unique=False
|
||||
)
|
||||
op.create_index(
|
||||
"idx_stats_hourly_model_model_hour",
|
||||
"stats_hourly_model",
|
||||
["model", "hour_utc"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
if not _table_exists("stats_hourly_provider"):
|
||||
op.create_table(
|
||||
"stats_hourly_provider",
|
||||
sa.Column("id", sa.String(length=36), nullable=False),
|
||||
sa.Column("hour_utc", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("provider_name", sa.String(length=100), nullable=False),
|
||||
sa.Column("total_requests", sa.Integer(), nullable=False),
|
||||
sa.Column("input_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("output_tokens", sa.BigInteger(), nullable=False),
|
||||
sa.Column("total_cost", sa.Float(), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("hour_utc", "provider_name", name="uq_stats_hourly_provider"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_stats_hourly_provider_hour",
|
||||
"stats_hourly_provider",
|
||||
["hour_utc"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
if not _table_exists("stats_daily_api_key"):
|
||||
op.create_table(
|
||||
"stats_daily_api_key",
|
||||
sa.Column("id", sa.String(length=36), primary_key=True),
|
||||
sa.Column(
|
||||
"api_key_id",
|
||||
sa.String(length=36),
|
||||
sa.ForeignKey("api_keys.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("date", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("total_requests", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("success_requests", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("error_requests", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("input_tokens", sa.BigInteger(), nullable=False, server_default="0"),
|
||||
sa.Column("output_tokens", sa.BigInteger(), nullable=False, server_default="0"),
|
||||
sa.Column("cache_creation_tokens", sa.BigInteger(), nullable=False, server_default="0"),
|
||||
sa.Column("cache_read_tokens", sa.BigInteger(), nullable=False, server_default="0"),
|
||||
sa.Column("total_cost", sa.Float(), nullable=False, server_default="0"),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.UniqueConstraint("api_key_id", "date", name="uq_stats_daily_api_key"),
|
||||
)
|
||||
|
||||
if _table_exists("stats_daily_api_key"):
|
||||
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date"):
|
||||
op.create_index("idx_stats_daily_api_key_date", "stats_daily_api_key", ["date"])
|
||||
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_key_date"):
|
||||
op.create_index(
|
||||
"idx_stats_daily_api_key_key_date",
|
||||
"stats_daily_api_key",
|
||||
["api_key_id", "date"],
|
||||
)
|
||||
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_requests"):
|
||||
op.create_index(
|
||||
"idx_stats_daily_api_key_date_requests",
|
||||
"stats_daily_api_key",
|
||||
["date", "total_requests"],
|
||||
)
|
||||
if not _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_cost"):
|
||||
op.create_index(
|
||||
"idx_stats_daily_api_key_date_cost",
|
||||
"stats_daily_api_key",
|
||||
["date", "total_cost"],
|
||||
)
|
||||
|
||||
if _table_exists("usage"):
|
||||
if not _column_exists("usage", "error_category"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column("error_category", sa.String(length=50), nullable=True),
|
||||
)
|
||||
op.create_index("idx_usage_error_category", "usage", ["error_category"], unique=False)
|
||||
|
||||
if _table_exists("stats_daily"):
|
||||
for name in (
|
||||
"p50_response_time_ms",
|
||||
"p90_response_time_ms",
|
||||
"p99_response_time_ms",
|
||||
"p50_first_byte_time_ms",
|
||||
"p90_first_byte_time_ms",
|
||||
"p99_first_byte_time_ms",
|
||||
):
|
||||
if not _column_exists("stats_daily", name):
|
||||
op.add_column("stats_daily", sa.Column(name, sa.Integer(), nullable=True))
|
||||
|
||||
if not _table_exists("stats_daily_error"):
|
||||
op.create_table(
|
||||
"stats_daily_error",
|
||||
sa.Column("id", sa.String(length=36), primary_key=True),
|
||||
sa.Column("date", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("error_category", sa.String(length=50), nullable=False),
|
||||
sa.Column("provider_name", sa.String(length=100), nullable=True),
|
||||
sa.Column("model", sa.String(length=100), nullable=True),
|
||||
sa.Column("count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
),
|
||||
sa.UniqueConstraint(
|
||||
"date",
|
||||
"error_category",
|
||||
"provider_name",
|
||||
"model",
|
||||
name="uq_stats_daily_error",
|
||||
),
|
||||
)
|
||||
|
||||
if _table_exists("stats_daily_error"):
|
||||
if not _index_exists("stats_daily_error", "idx_stats_daily_error_date"):
|
||||
op.create_index("idx_stats_daily_error_date", "stats_daily_error", ["date"])
|
||||
if not _index_exists("stats_daily_error", "idx_stats_daily_error_category"):
|
||||
op.create_index(
|
||||
"idx_stats_daily_error_category",
|
||||
"stats_daily_error",
|
||||
["date", "error_category"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if _table_exists("stats_daily_error"):
|
||||
if _index_exists("stats_daily_error", "idx_stats_daily_error_category"):
|
||||
op.drop_index("idx_stats_daily_error_category", table_name="stats_daily_error")
|
||||
if _index_exists("stats_daily_error", "idx_stats_daily_error_date"):
|
||||
op.drop_index("idx_stats_daily_error_date", table_name="stats_daily_error")
|
||||
op.drop_table("stats_daily_error")
|
||||
|
||||
if _table_exists("stats_daily"):
|
||||
for name in (
|
||||
"p50_response_time_ms",
|
||||
"p90_response_time_ms",
|
||||
"p99_response_time_ms",
|
||||
"p50_first_byte_time_ms",
|
||||
"p90_first_byte_time_ms",
|
||||
"p99_first_byte_time_ms",
|
||||
):
|
||||
if _column_exists("stats_daily", name):
|
||||
op.drop_column("stats_daily", name)
|
||||
|
||||
if _table_exists("usage") and _column_exists("usage", "error_category"):
|
||||
if _index_exists("usage", "idx_usage_error_category"):
|
||||
op.drop_index("idx_usage_error_category", table_name="usage")
|
||||
op.drop_column("usage", "error_category")
|
||||
|
||||
if _table_exists("stats_daily_api_key"):
|
||||
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_cost"):
|
||||
op.drop_index("idx_stats_daily_api_key_date_cost", table_name="stats_daily_api_key")
|
||||
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date_requests"):
|
||||
op.drop_index("idx_stats_daily_api_key_date_requests", table_name="stats_daily_api_key")
|
||||
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_key_date"):
|
||||
op.drop_index("idx_stats_daily_api_key_key_date", table_name="stats_daily_api_key")
|
||||
if _index_exists("stats_daily_api_key", "idx_stats_daily_api_key_date"):
|
||||
op.drop_index("idx_stats_daily_api_key_date", table_name="stats_daily_api_key")
|
||||
op.drop_table("stats_daily_api_key")
|
||||
|
||||
if _table_exists("stats_hourly_provider"):
|
||||
if _index_exists("stats_hourly_provider", "idx_stats_hourly_provider_hour"):
|
||||
op.drop_index("idx_stats_hourly_provider_hour", table_name="stats_hourly_provider")
|
||||
op.drop_table("stats_hourly_provider")
|
||||
|
||||
if _table_exists("stats_hourly_model"):
|
||||
if _index_exists("stats_hourly_model", "idx_stats_hourly_model_model_hour"):
|
||||
op.drop_index("idx_stats_hourly_model_model_hour", table_name="stats_hourly_model")
|
||||
if _index_exists("stats_hourly_model", "idx_stats_hourly_model_hour"):
|
||||
op.drop_index("idx_stats_hourly_model_hour", table_name="stats_hourly_model")
|
||||
op.drop_table("stats_hourly_model")
|
||||
|
||||
if _table_exists("stats_hourly_user"):
|
||||
if _index_exists("stats_hourly_user", "idx_stats_hourly_user_user_hour"):
|
||||
op.drop_index("idx_stats_hourly_user_user_hour", table_name="stats_hourly_user")
|
||||
if _index_exists("stats_hourly_user", "idx_stats_hourly_user_hour"):
|
||||
op.drop_index("idx_stats_hourly_user_hour", table_name="stats_hourly_user")
|
||||
op.drop_table("stats_hourly_user")
|
||||
|
||||
if _table_exists("stats_hourly"):
|
||||
if _index_exists("stats_hourly", "idx_stats_hourly_hour"):
|
||||
op.drop_index("idx_stats_hourly_hour", table_name="stats_hourly")
|
||||
op.drop_table("stats_hourly")
|
||||
|
||||
if _table_exists("stats_daily"):
|
||||
if _column_exists("stats_daily", "aggregated_at"):
|
||||
op.drop_column("stats_daily", "aggregated_at")
|
||||
if _column_exists("stats_daily", "is_complete"):
|
||||
op.drop_column("stats_daily", "is_complete")
|
||||
@@ -0,0 +1,205 @@
|
||||
"""Add provider_type, upstream_metadata, oauth_invalid fields and expand string columns to TEXT
|
||||
|
||||
- Add providers.provider_type (String(20), server_default="custom")
|
||||
- Add provider_api_keys.upstream_metadata (JSON, nullable)
|
||||
- Add provider_api_keys.oauth_invalid_at (DateTime, nullable) - OAuth Token 失效时间
|
||||
- Add provider_api_keys.oauth_invalid_reason (String(255), nullable) - OAuth Token 失效原因
|
||||
- Expand multiple VARCHAR columns to TEXT for long values (OAuth tokens, LDAP DN, URLs, etc.)
|
||||
|
||||
Revision ID: b5c6d7e8f9a0
|
||||
Revises: c4e8f9a1b2c3
|
||||
Create Date: 2026-02-04 15:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "b5c6d7e8f9a0"
|
||||
down_revision: Union[str, None] = "c4e8f9a1b2c3"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
# 需要扩展为 TEXT 的列(表名, 列名, 原始类型长度)
|
||||
COLUMNS_TO_EXPAND = [
|
||||
("provider_api_keys", "api_key", 500), # OAuth tokens can be very long
|
||||
(
|
||||
"provider_api_keys",
|
||||
"auth_config",
|
||||
None,
|
||||
), # 确保 auth_config 是 TEXT 类型(可能从 JSON 迁移过来)
|
||||
("ldap_configs", "bind_dn", 255), # LDAP DN can be deeply nested
|
||||
("ldap_configs", "base_dn", 255), # LDAP DN can be deeply nested
|
||||
("ldap_configs", "user_search_filter", 500), # Complex LDAP filters
|
||||
("oauth_providers", "client_id", 255), # Some OAuth providers use JWT client_id
|
||||
]
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
"""检查列是否已存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [col["name"] for col in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
"""检查表是否存在"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def is_sqlite() -> bool:
|
||||
"""检查是否为 SQLite 数据库"""
|
||||
bind = op.get_bind()
|
||||
return bind.dialect.name == "sqlite"
|
||||
|
||||
|
||||
def get_column_type(table_name: str, column_name: str) -> str | None:
|
||||
"""获取列的数据类型"""
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
for col in inspector.get_columns(table_name):
|
||||
if col["name"] == column_name:
|
||||
return str(col["type"]).upper()
|
||||
return None
|
||||
|
||||
|
||||
def expand_column_to_text(table_name: str, column_name: str, original_length: int | None) -> None:
|
||||
"""将 VARCHAR 列扩展为 TEXT(兼容 SQLite)"""
|
||||
if not table_exists(table_name):
|
||||
return
|
||||
if not column_exists(table_name, column_name):
|
||||
return
|
||||
|
||||
# 检查当前列类型,如果已经是 TEXT 则跳过
|
||||
col_type = get_column_type(table_name, column_name)
|
||||
if col_type and "TEXT" in col_type:
|
||||
return
|
||||
|
||||
# 如果是 JSON 类型(可能是历史遗留),先将 JSON 数据转为文本表示再变更类型
|
||||
is_json_col = col_type and "JSON" in col_type
|
||||
|
||||
if is_json_col and not is_sqlite():
|
||||
# PostgreSQL: 先用 CAST 把 JSON 值转为 TEXT,保留数据
|
||||
op.execute(
|
||||
sa.text(
|
||||
f"ALTER TABLE {table_name} ALTER COLUMN {column_name} "
|
||||
f"TYPE TEXT USING {column_name}::TEXT"
|
||||
)
|
||||
)
|
||||
return
|
||||
|
||||
if is_sqlite():
|
||||
# SQLite 不支持直接 ALTER COLUMN,需要用 batch 模式
|
||||
# batch 模式会自动处理 JSON->TEXT 的数据迁移
|
||||
with op.batch_alter_table(table_name) as batch_op:
|
||||
batch_op.alter_column(
|
||||
column_name,
|
||||
type_=sa.Text(),
|
||||
existing_type=sa.String(original_length) if original_length else sa.Text(),
|
||||
)
|
||||
else:
|
||||
op.alter_column(
|
||||
table_name,
|
||||
column_name,
|
||||
type_=sa.Text(),
|
||||
existing_type=sa.String(original_length) if original_length else sa.Text(),
|
||||
existing_nullable=True,
|
||||
)
|
||||
|
||||
|
||||
def shrink_column_to_varchar(
|
||||
table_name: str, column_name: str, target_length: int, nullable: bool = False
|
||||
) -> None:
|
||||
"""将 TEXT 列缩小为 VARCHAR(兼容 SQLite)
|
||||
WARNING: 如果数据超过 target_length 会失败
|
||||
"""
|
||||
if not table_exists(table_name):
|
||||
return
|
||||
if not column_exists(table_name, column_name):
|
||||
return
|
||||
|
||||
if is_sqlite():
|
||||
with op.batch_alter_table(table_name) as batch_op:
|
||||
batch_op.alter_column(
|
||||
column_name,
|
||||
type_=sa.String(target_length),
|
||||
existing_type=sa.Text(),
|
||||
existing_nullable=nullable,
|
||||
)
|
||||
else:
|
||||
op.alter_column(
|
||||
table_name,
|
||||
column_name,
|
||||
type_=sa.String(target_length),
|
||||
existing_type=sa.Text(),
|
||||
existing_nullable=nullable,
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# Add providers.provider_type
|
||||
if not column_exists("providers", "provider_type"):
|
||||
op.add_column(
|
||||
"providers",
|
||||
sa.Column("provider_type", sa.String(20), nullable=False, server_default="custom"),
|
||||
)
|
||||
|
||||
# Add provider_api_keys.upstream_metadata
|
||||
if not column_exists("provider_api_keys", "upstream_metadata"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("upstream_metadata", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
# Add provider_api_keys.oauth_invalid_at
|
||||
if not column_exists("provider_api_keys", "oauth_invalid_at"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("oauth_invalid_at", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
|
||||
# Add provider_api_keys.oauth_invalid_reason
|
||||
if not column_exists("provider_api_keys", "oauth_invalid_reason"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("oauth_invalid_reason", sa.String(255), nullable=True),
|
||||
)
|
||||
|
||||
# Expand VARCHAR columns to TEXT
|
||||
for table_name, column_name, original_length in COLUMNS_TO_EXPAND:
|
||||
expand_column_to_text(table_name, column_name, original_length)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Shrink TEXT columns back to VARCHAR
|
||||
# WARNING: Downgrade may fail if any values exceed original length
|
||||
for table_name, column_name, original_length in reversed(COLUMNS_TO_EXPAND):
|
||||
# 跳过没有原始长度的列(如 auth_config,由其他迁移创建)
|
||||
if original_length is None:
|
||||
continue
|
||||
shrink_column_to_varchar(table_name, column_name, original_length)
|
||||
|
||||
# Drop provider_api_keys.oauth_invalid_reason
|
||||
if column_exists("provider_api_keys", "oauth_invalid_reason"):
|
||||
op.drop_column("provider_api_keys", "oauth_invalid_reason")
|
||||
|
||||
# Drop provider_api_keys.oauth_invalid_at
|
||||
if column_exists("provider_api_keys", "oauth_invalid_at"):
|
||||
op.drop_column("provider_api_keys", "oauth_invalid_at")
|
||||
|
||||
# Drop provider_api_keys.upstream_metadata
|
||||
if column_exists("provider_api_keys", "upstream_metadata"):
|
||||
op.drop_column("provider_api_keys", "upstream_metadata")
|
||||
|
||||
# Drop providers.provider_type
|
||||
if column_exists("providers", "provider_type"):
|
||||
op.drop_column("providers", "provider_type")
|
||||
@@ -0,0 +1,254 @@
|
||||
"""Antigravity endpoint signature to gemini:chat & add proxy_nodes table (with manual fields)
|
||||
|
||||
Revision ID: e1b2c3d4f5a6
|
||||
Revises: b5c6d7e8f9a0
|
||||
Create Date: 2026-02-06 23:45:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect, text
|
||||
from sqlalchemy.dialects import postgresql
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "e1b2c3d4f5a6"
|
||||
down_revision: str | None = "b5c6d7e8f9a0"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
return table_name in inspector.get_table_names()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# =========================================================================
|
||||
# Part 1: Antigravity endpoint signature migration (gemini:cli -> gemini:chat)
|
||||
# =========================================================================
|
||||
|
||||
# --- provider_endpoints ---
|
||||
# Update only when there is no conflicting gemini:chat endpoint for the same provider
|
||||
# (provider_endpoints has a unique constraint on (provider_id, api_format)).
|
||||
conn.execute(text("""
|
||||
UPDATE provider_endpoints pe
|
||||
SET
|
||||
api_format = 'gemini:chat',
|
||||
api_family = 'gemini',
|
||||
endpoint_kind = 'chat'
|
||||
WHERE pe.api_format = 'gemini:cli'
|
||||
AND pe.provider_id IN (
|
||||
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
|
||||
)
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM provider_endpoints pe2
|
||||
WHERE pe2.provider_id = pe.provider_id
|
||||
AND pe2.api_format = 'gemini:chat'
|
||||
)
|
||||
"""))
|
||||
|
||||
# Best-effort normalization for already-existing Antigravity gemini:chat endpoints.
|
||||
conn.execute(text("""
|
||||
UPDATE provider_endpoints pe
|
||||
SET
|
||||
api_family = 'gemini',
|
||||
endpoint_kind = 'chat'
|
||||
WHERE pe.api_format = 'gemini:chat'
|
||||
AND pe.provider_id IN (
|
||||
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
|
||||
)
|
||||
"""))
|
||||
|
||||
# --- provider_api_keys.api_formats (JSON array) ---
|
||||
# Replace "gemini:cli" with "gemini:chat" in the JSON array for Antigravity keys.
|
||||
# Uses text-level replace on the serialized JSON -- safe because the value is a
|
||||
# simple string with no special characters that could cause ambiguous replacements.
|
||||
conn.execute(text("""
|
||||
UPDATE provider_api_keys pak
|
||||
SET api_formats = replace(pak.api_formats::text, '"gemini:cli"', '"gemini:chat"')::json
|
||||
WHERE pak.provider_id IN (
|
||||
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
|
||||
)
|
||||
AND pak.api_formats IS NOT NULL
|
||||
AND pak.api_formats::text LIKE '%"gemini:cli"%'
|
||||
"""))
|
||||
|
||||
# =========================================================================
|
||||
# Part 2: Create proxy_nodes table with manual proxy fields (idempotent)
|
||||
# =========================================================================
|
||||
|
||||
# Create ENUM type (idempotent)
|
||||
op.execute(
|
||||
"DO $$ BEGIN "
|
||||
"CREATE TYPE proxynodestatus AS ENUM ('online', 'unhealthy', 'offline'); "
|
||||
"EXCEPTION WHEN duplicate_object THEN NULL; "
|
||||
"END $$"
|
||||
)
|
||||
|
||||
if table_exists("proxy_nodes"):
|
||||
# Table already exists — ensure manual proxy columns are present
|
||||
inspector = inspect(conn)
|
||||
existing_columns = {c["name"] for c in inspector.get_columns("proxy_nodes")}
|
||||
|
||||
# ip 列扩容:手动节点的 ip 存储 "socks5://hostname" 形式,45 字符可能不够
|
||||
ip_col = next((c for c in inspector.get_columns("proxy_nodes") if c["name"] == "ip"), None)
|
||||
if ip_col and hasattr(ip_col["type"], "length") and (ip_col["type"].length or 0) < 512:
|
||||
op.alter_column("proxy_nodes", "ip", type_=sa.String(512), existing_nullable=False)
|
||||
|
||||
manual_columns = [
|
||||
("is_manual", sa.Boolean(), False, sa.text("false"), "是否为手动添加的代理节点"),
|
||||
("proxy_url", sa.String(500), True, None, "手动节点的完整代理 URL"),
|
||||
("proxy_username", sa.String(255), True, None, "手动节点的代理用户名"),
|
||||
("proxy_password", sa.String(500), True, None, "手动节点的代理密码"),
|
||||
]
|
||||
for col_name, col_type, nullable, default, comment in manual_columns:
|
||||
if col_name not in existing_columns:
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
col_name,
|
||||
col_type, # type: ignore[arg-type]
|
||||
nullable=nullable,
|
||||
server_default=default,
|
||||
comment=comment,
|
||||
),
|
||||
)
|
||||
return
|
||||
|
||||
op.create_table(
|
||||
"proxy_nodes",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("name", sa.String(100), nullable=False),
|
||||
sa.Column("ip", sa.String(512), nullable=False),
|
||||
sa.Column("port", sa.Integer(), nullable=False),
|
||||
sa.Column("region", sa.String(100), nullable=True),
|
||||
sa.Column(
|
||||
"status",
|
||||
postgresql.ENUM(
|
||||
"online",
|
||||
"unhealthy",
|
||||
"offline",
|
||||
name="proxynodestatus",
|
||||
create_type=False,
|
||||
),
|
||||
nullable=False,
|
||||
server_default=sa.text("'online'"),
|
||||
),
|
||||
sa.Column(
|
||||
"registered_by",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("users.id", ondelete="SET NULL"),
|
||||
nullable=True,
|
||||
),
|
||||
sa.Column("last_heartbeat_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("heartbeat_interval", sa.Integer(), nullable=False, server_default=sa.text("30")),
|
||||
sa.Column("active_connections", sa.Integer(), nullable=False, server_default=sa.text("0")),
|
||||
sa.Column("total_requests", sa.BigInteger(), nullable=False, server_default=sa.text("0")),
|
||||
sa.Column("avg_latency_ms", sa.Float(), nullable=True),
|
||||
# --- Manual proxy node fields ---
|
||||
sa.Column(
|
||||
"is_manual",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
comment="是否为手动添加的代理节点",
|
||||
),
|
||||
sa.Column(
|
||||
"proxy_url",
|
||||
sa.String(500),
|
||||
nullable=True,
|
||||
comment="手动节点的完整代理 URL",
|
||||
),
|
||||
sa.Column(
|
||||
"proxy_username",
|
||||
sa.String(255),
|
||||
nullable=True,
|
||||
comment="手动节点的代理用户名",
|
||||
),
|
||||
sa.Column(
|
||||
"proxy_password",
|
||||
sa.String(500),
|
||||
nullable=True,
|
||||
comment="手动节点的代理密码",
|
||||
),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.text("CURRENT_TIMESTAMP"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.UniqueConstraint("ip", "port", name="uq_proxy_node_ip_port"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# =========================================================================
|
||||
# Part 2 rollback: Drop proxy_nodes table (and manual columns if present)
|
||||
# =========================================================================
|
||||
if table_exists("proxy_nodes"):
|
||||
op.drop_table("proxy_nodes")
|
||||
|
||||
# Best-effort: drop type (only used by proxy_nodes)
|
||||
op.execute("DROP TYPE IF EXISTS proxynodestatus")
|
||||
|
||||
# =========================================================================
|
||||
# Part 1 rollback: Revert Antigravity endpoint signature (gemini:chat -> gemini:cli)
|
||||
# =========================================================================
|
||||
|
||||
# --- provider_endpoints ---
|
||||
conn.execute(text("""
|
||||
UPDATE provider_endpoints pe
|
||||
SET
|
||||
api_format = 'gemini:cli',
|
||||
api_family = 'gemini',
|
||||
endpoint_kind = 'cli'
|
||||
WHERE pe.api_format = 'gemini:chat'
|
||||
AND pe.provider_id IN (
|
||||
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
|
||||
)
|
||||
AND NOT EXISTS (
|
||||
SELECT 1 FROM provider_endpoints pe2
|
||||
WHERE pe2.provider_id = pe.provider_id
|
||||
AND pe2.api_format = 'gemini:cli'
|
||||
)
|
||||
"""))
|
||||
|
||||
# Best-effort normalization for already-existing Antigravity gemini:cli endpoints.
|
||||
conn.execute(text("""
|
||||
UPDATE provider_endpoints pe
|
||||
SET
|
||||
api_family = 'gemini',
|
||||
endpoint_kind = 'cli'
|
||||
WHERE pe.api_format = 'gemini:cli'
|
||||
AND pe.provider_id IN (
|
||||
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
|
||||
)
|
||||
"""))
|
||||
|
||||
# --- provider_api_keys.api_formats (JSON array) ---
|
||||
conn.execute(text("""
|
||||
UPDATE provider_api_keys pak
|
||||
SET api_formats = replace(pak.api_formats::text, '"gemini:chat"', '"gemini:cli"')::json
|
||||
WHERE pak.provider_id IN (
|
||||
SELECT id FROM providers WHERE lower(provider_type) = 'antigravity'
|
||||
)
|
||||
AND pak.api_formats IS NOT NULL
|
||||
AND pak.api_formats::text LIKE '%"gemini:chat"%'
|
||||
"""))
|
||||
@@ -0,0 +1,61 @@
|
||||
"""Add remote_config and config_version to proxy_nodes
|
||||
|
||||
Revision ID: 3aff3ffc4a0e
|
||||
Revises: e1b2c3d4f5a6
|
||||
Create Date: 2026-02-07 15:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "3aff3ffc4a0e"
|
||||
down_revision: str | None = "e1b2c3d4f5a6"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [c["name"] for c in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not column_exists("proxy_nodes", "remote_config"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"remote_config",
|
||||
sa.JSON(),
|
||||
nullable=True,
|
||||
comment="管理端下发的远程配置 (allowed_ports, log_level, heartbeat_interval, timestamp_tolerance)",
|
||||
),
|
||||
)
|
||||
|
||||
if not column_exists("proxy_nodes", "config_version"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"config_version",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
comment="远程配置版本号,每次更新 +1",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if column_exists("proxy_nodes", "config_version"):
|
||||
op.drop_column("proxy_nodes", "config_version")
|
||||
if column_exists("proxy_nodes", "remote_config"):
|
||||
op.drop_column("proxy_nodes", "remote_config")
|
||||
@@ -0,0 +1,61 @@
|
||||
"""Add tls_enabled and tls_cert_fingerprint to proxy_nodes
|
||||
|
||||
Revision ID: 4b5c6d7e8f9a
|
||||
Revises: 3aff3ffc4a0e
|
||||
Create Date: 2026-02-07 18:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "4b5c6d7e8f9a"
|
||||
down_revision: str | None = "3aff3ffc4a0e"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [c["name"] for c in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not column_exists("proxy_nodes", "tls_enabled"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tls_enabled",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default="false",
|
||||
comment="是否启用 TLS 加密",
|
||||
),
|
||||
)
|
||||
|
||||
if not column_exists("proxy_nodes", "tls_cert_fingerprint"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tls_cert_fingerprint",
|
||||
sa.String(128),
|
||||
nullable=True,
|
||||
comment="TLS 证书 SHA-256 指纹(hex)",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if column_exists("proxy_nodes", "tls_cert_fingerprint"):
|
||||
op.drop_column("proxy_nodes", "tls_cert_fingerprint")
|
||||
if column_exists("proxy_nodes", "tls_enabled"):
|
||||
op.drop_column("proxy_nodes", "tls_enabled")
|
||||
@@ -0,0 +1,60 @@
|
||||
"""Add hardware_info and estimated_max_concurrency to proxy_nodes
|
||||
|
||||
Revision ID: 5c6d7e8f9a0b
|
||||
Revises: 4b5c6d7e8f9a
|
||||
Create Date: 2026-02-08 12:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "5c6d7e8f9a0b"
|
||||
down_revision: str | None = "4b5c6d7e8f9a"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [c["name"] for c in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not column_exists("proxy_nodes", "hardware_info"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"hardware_info",
|
||||
sa.JSON(),
|
||||
nullable=True,
|
||||
comment="硬件信息 (cpu_cores, total_memory_mb, os_info, fd_limit, ...)",
|
||||
),
|
||||
)
|
||||
|
||||
if not column_exists("proxy_nodes", "estimated_max_concurrency"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"estimated_max_concurrency",
|
||||
sa.Integer(),
|
||||
nullable=True,
|
||||
comment="基于硬件估算的最大并发连接数",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if column_exists("proxy_nodes", "estimated_max_concurrency"):
|
||||
op.drop_column("proxy_nodes", "estimated_max_concurrency")
|
||||
if column_exists("proxy_nodes", "hardware_info"):
|
||||
op.drop_column("proxy_nodes", "hardware_info")
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Add proxy column to provider_api_keys for per-key proxy configuration
|
||||
|
||||
Revision ID: 6d7e8f9a0b1c
|
||||
Revises: 5c6d7e8f9a0b
|
||||
Create Date: 2026-02-08 15:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "6d7e8f9a0b1c"
|
||||
down_revision: str | None = "5c6d7e8f9a0b"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [c["name"] for c in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not column_exists("provider_api_keys", "proxy"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column(
|
||||
"proxy",
|
||||
sa.JSON(),
|
||||
nullable=True,
|
||||
comment="Key 级别代理配置(覆盖 Provider 级别代理),如 {node_id, enabled}",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if column_exists("provider_api_keys", "proxy"):
|
||||
op.drop_column("provider_api_keys", "proxy")
|
||||
+46
@@ -0,0 +1,46 @@
|
||||
"""Add provider_request_body and client_response_body columns to usage table
|
||||
|
||||
Revision ID: 7e8f9a0b1c2d
|
||||
Revises: 6d7e8f9a0b1c
|
||||
Create Date: 2026-02-20 18:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "7e8f9a0b1c2d"
|
||||
down_revision: str | None = "6d7e8f9a0b1c"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
# Use PostgreSQL native IF NOT EXISTS to avoid duplicate-column races
|
||||
# when migrations are triggered concurrently (e.g. startup + manual run).
|
||||
conn.execute(text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_request_body JSON"))
|
||||
conn.execute(
|
||||
text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_request_body_compressed BYTEA")
|
||||
)
|
||||
conn.execute(text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS client_response_body JSON"))
|
||||
conn.execute(
|
||||
text("ALTER TABLE usage ADD COLUMN IF NOT EXISTS client_response_body_compressed BYTEA")
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
for col in (
|
||||
"client_response_body_compressed",
|
||||
"client_response_body",
|
||||
"provider_request_body_compressed",
|
||||
"provider_request_body",
|
||||
):
|
||||
conn.execute(text(f"ALTER TABLE usage DROP COLUMN IF EXISTS {col}"))
|
||||
@@ -0,0 +1,91 @@
|
||||
"""Add api_family and endpoint_kind columns to usage table
|
||||
|
||||
Revision ID: 8f9a0b1c2d3e
|
||||
Revises: 7e8f9a0b1c2d
|
||||
Create Date: 2026-02-21 15:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect, text
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "8f9a0b1c2d3e"
|
||||
down_revision: str | None = "7e8f9a0b1c2d"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
# Usage 表新增列
|
||||
NEW_COLUMNS = [
|
||||
("api_family", sa.String(50)),
|
||||
("endpoint_kind", sa.String(50)),
|
||||
("provider_api_family", sa.String(50)),
|
||||
("provider_endpoint_kind", sa.String(50)),
|
||||
]
|
||||
|
||||
# 新增索引
|
||||
NEW_INDEXES = [
|
||||
("idx_usage_api_family", "usage", ["api_family"]),
|
||||
("idx_usage_endpoint_kind", "usage", ["endpoint_kind"]),
|
||||
("idx_usage_family_kind", "usage", ["api_family", "endpoint_kind"]),
|
||||
]
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# 使用 PostgreSQL 原生 IF NOT EXISTS,比 inspect 更可靠(避免同一事务内缓存问题)
|
||||
col_definitions = {
|
||||
"api_family": "VARCHAR(50)",
|
||||
"endpoint_kind": "VARCHAR(50)",
|
||||
"provider_api_family": "VARCHAR(50)",
|
||||
"provider_endpoint_kind": "VARCHAR(50)",
|
||||
}
|
||||
for col_name, col_type_sql in col_definitions.items():
|
||||
conn.execute(text(f"ALTER TABLE usage ADD COLUMN IF NOT EXISTS {col_name} {col_type_sql}"))
|
||||
|
||||
# 数据迁移:从 api_format 解析 api_family + endpoint_kind
|
||||
conn.execute(text("""
|
||||
UPDATE usage SET
|
||||
api_family = lower(split_part(api_format, ':', 1)),
|
||||
endpoint_kind = lower(split_part(api_format, ':', 2))
|
||||
WHERE api_format IS NOT NULL
|
||||
AND api_format LIKE '%%:%%'
|
||||
AND api_family IS NULL
|
||||
"""))
|
||||
conn.execute(text("""
|
||||
UPDATE usage SET
|
||||
provider_api_family = lower(split_part(endpoint_api_format, ':', 1)),
|
||||
provider_endpoint_kind = lower(split_part(endpoint_api_format, ':', 2))
|
||||
WHERE endpoint_api_format IS NOT NULL
|
||||
AND endpoint_api_format LIKE '%%:%%'
|
||||
AND provider_api_family IS NULL
|
||||
"""))
|
||||
|
||||
# 创建索引
|
||||
inspector = inspect(conn)
|
||||
existing_indexes = {idx["name"] for idx in inspector.get_indexes("usage")}
|
||||
for idx_name, table, columns in NEW_INDEXES:
|
||||
if idx_name not in existing_indexes:
|
||||
op.create_index(idx_name, table, columns)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = inspect(conn)
|
||||
|
||||
existing_indexes = {idx["name"] for idx in inspector.get_indexes("usage")}
|
||||
for idx_name, _, _ in reversed(NEW_INDEXES):
|
||||
if idx_name in existing_indexes:
|
||||
op.drop_index(idx_name, table_name="usage")
|
||||
|
||||
existing_columns = {col["name"] for col in inspector.get_columns("usage")}
|
||||
for col_name, _ in reversed(NEW_COLUMNS):
|
||||
if col_name in existing_columns:
|
||||
op.drop_column("usage", col_name)
|
||||
@@ -0,0 +1,106 @@
|
||||
"""Add tunnel mode fields and remove IP forwarding fields
|
||||
|
||||
Revision ID: 9a0b1c2d3e4f
|
||||
Revises: 8f9a0b1c2d3e
|
||||
Create Date: 2026-02-24 17:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision: str = "9a0b1c2d3e4f"
|
||||
down_revision: str | None = "8f9a0b1c2d3e"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [c["name"] for c in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# 添加 tunnel 模式字段
|
||||
if not column_exists("proxy_nodes", "tunnel_mode"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tunnel_mode",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
comment="是否使用 WebSocket 隧道模式",
|
||||
),
|
||||
)
|
||||
if not column_exists("proxy_nodes", "tunnel_connected"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tunnel_connected",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
comment="隧道是否已连接",
|
||||
),
|
||||
)
|
||||
if not column_exists("proxy_nodes", "tunnel_connected_at"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tunnel_connected_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=True,
|
||||
comment="隧道最近一次建立时间",
|
||||
),
|
||||
)
|
||||
|
||||
# tunnel 模式节点不需要 port,将其置零
|
||||
op.execute("UPDATE proxy_nodes SET port = 0 WHERE tunnel_mode = true")
|
||||
|
||||
# 移除旧的 IP 转发字段
|
||||
if column_exists("proxy_nodes", "tls_enabled"):
|
||||
op.drop_column("proxy_nodes", "tls_enabled")
|
||||
if column_exists("proxy_nodes", "tls_cert_fingerprint"):
|
||||
op.drop_column("proxy_nodes", "tls_cert_fingerprint")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# 恢复 IP 转发字段
|
||||
if not column_exists("proxy_nodes", "tls_cert_fingerprint"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tls_cert_fingerprint",
|
||||
sa.String(128),
|
||||
nullable=True,
|
||||
comment="TLS 证书 SHA-256 指纹(hex)",
|
||||
),
|
||||
)
|
||||
if not column_exists("proxy_nodes", "tls_enabled"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"tls_enabled",
|
||||
sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
comment="是否启用 TLS 加密",
|
||||
),
|
||||
)
|
||||
|
||||
# 移除 tunnel 模式字段
|
||||
if column_exists("proxy_nodes", "tunnel_connected_at"):
|
||||
op.drop_column("proxy_nodes", "tunnel_connected_at")
|
||||
if column_exists("proxy_nodes", "tunnel_connected"):
|
||||
op.drop_column("proxy_nodes", "tunnel_connected")
|
||||
if column_exists("proxy_nodes", "tunnel_mode"):
|
||||
op.drop_column("proxy_nodes", "tunnel_mode")
|
||||
@@ -0,0 +1,224 @@
|
||||
"""Add cache_creation columns, clean up capability settings, add user_model_usage_counts,
|
||||
enforce global_model_id NOT NULL
|
||||
|
||||
1. Add cache_creation_input_tokens_5m and cache_creation_input_tokens_1h to usage table.
|
||||
2. Clean up cache_1h/context_1m/gemini_files from user-configurable settings
|
||||
(now auto-detected via REQUEST_PARAM mode).
|
||||
3. Create user_model_usage_counts table for per-user per-model atomic usage counters.
|
||||
4. Enforce models.global_model_id NOT NULL (delete orphan models without global model).
|
||||
|
||||
Revision ID: b2c3d4e5f6a7
|
||||
Revises: 9a0b1c2d3e4f
|
||||
Create Date: 2026-02-28 14:00:00.000000
|
||||
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from datetime import datetime, timezone
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision: str = "b2c3d4e5f6a7"
|
||||
down_revision: str | None = "9a0b1c2d3e4f"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
columns = [c["name"] for c in insp.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
return table_name in insp.get_table_names()
|
||||
|
||||
|
||||
def index_exists(table_name: str, index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
return any(idx["name"] == index_name for idx in insp.get_indexes(table_name))
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# --- 1. Add cache_creation columns ---
|
||||
if not column_exists("usage", "cache_creation_input_tokens_5m"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column(
|
||||
"cache_creation_input_tokens_5m",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default=sa.text("0"),
|
||||
comment="5min TTL cache creation input tokens",
|
||||
),
|
||||
)
|
||||
if not column_exists("usage", "cache_creation_input_tokens_1h"):
|
||||
op.add_column(
|
||||
"usage",
|
||||
sa.Column(
|
||||
"cache_creation_input_tokens_1h",
|
||||
sa.Integer(),
|
||||
nullable=False,
|
||||
server_default=sa.text("0"),
|
||||
comment="1h TTL cache creation input tokens",
|
||||
),
|
||||
)
|
||||
|
||||
# --- 2. Clean up stale capability settings (pure Python, DB-agnostic) ---
|
||||
stale_keys = {"cache_1h", "context_1m", "gemini_files"}
|
||||
conn = op.get_bind()
|
||||
|
||||
# ApiKey.force_capabilities: dict-like JSON, remove stale keys
|
||||
rows = conn.execute(
|
||||
sa.text("SELECT id, force_capabilities FROM api_keys WHERE force_capabilities IS NOT NULL")
|
||||
).fetchall()
|
||||
for row in rows:
|
||||
raw = row[1]
|
||||
if raw is None:
|
||||
continue
|
||||
data = raw if isinstance(raw, dict) else json.loads(raw)
|
||||
cleaned = {k: v for k, v in data.items() if k not in stale_keys}
|
||||
new_val = json.dumps(cleaned) if cleaned else None
|
||||
conn.execute(
|
||||
sa.text("UPDATE api_keys SET force_capabilities = :val WHERE id = :id"),
|
||||
{"val": new_val, "id": row[0]},
|
||||
)
|
||||
|
||||
# User.model_capability_settings: nested dict {model_key: {cap: val}}, remove stale keys
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id, model_capability_settings FROM users"
|
||||
" WHERE model_capability_settings IS NOT NULL"
|
||||
)
|
||||
).fetchall()
|
||||
for row in rows:
|
||||
raw = row[1]
|
||||
if raw is None:
|
||||
continue
|
||||
data = raw if isinstance(raw, dict) else json.loads(raw)
|
||||
cleaned = {}
|
||||
for model_key, caps in data.items():
|
||||
cap_cleaned = {k: v for k, v in caps.items() if k not in stale_keys}
|
||||
if cap_cleaned:
|
||||
cleaned[model_key] = cap_cleaned
|
||||
new_val = json.dumps(cleaned) if cleaned else None
|
||||
conn.execute(
|
||||
sa.text("UPDATE users SET model_capability_settings = :val WHERE id = :id"),
|
||||
{"val": new_val, "id": row[0]},
|
||||
)
|
||||
|
||||
# GlobalModel.supported_capabilities: JSON array, remove stale entries
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id, supported_capabilities FROM global_models"
|
||||
" WHERE supported_capabilities IS NOT NULL"
|
||||
)
|
||||
).fetchall()
|
||||
for row in rows:
|
||||
raw = row[1]
|
||||
if raw is None:
|
||||
continue
|
||||
data = raw if isinstance(raw, list) else json.loads(raw)
|
||||
cleaned = [c for c in data if c not in stale_keys]
|
||||
new_val = json.dumps(cleaned) if cleaned else None
|
||||
conn.execute(
|
||||
sa.text("UPDATE global_models SET supported_capabilities = :val WHERE id = :id"),
|
||||
{"val": new_val, "id": row[0]},
|
||||
)
|
||||
|
||||
# --- 3. Create user_model_usage_counts table ---
|
||||
if not table_exists("user_model_usage_counts"):
|
||||
op.create_table(
|
||||
"user_model_usage_counts",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column(
|
||||
"user_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("model", sa.String(100), nullable=False),
|
||||
sa.Column("usage_count", sa.Integer, nullable=False, server_default="0"),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
sa.UniqueConstraint("user_id", "model", name="uq_user_model_usage_count"),
|
||||
)
|
||||
if not index_exists("user_model_usage_counts", "idx_user_model_usage_user"):
|
||||
op.create_index("idx_user_model_usage_user", "user_model_usage_counts", ["user_id"])
|
||||
if not index_exists("user_model_usage_counts", "idx_user_model_usage_model"):
|
||||
op.create_index("idx_user_model_usage_model", "user_model_usage_counts", ["model"])
|
||||
|
||||
# Backfill from existing usage records (truncate first for idempotency)
|
||||
conn.execute(sa.text("DELETE FROM user_model_usage_counts"))
|
||||
rows = conn.execute(
|
||||
sa.text(
|
||||
"SELECT user_id, model, COUNT(*) AS cnt FROM usage"
|
||||
" WHERE user_id IS NOT NULL GROUP BY user_id, model"
|
||||
)
|
||||
).fetchall()
|
||||
now = datetime.now(timezone.utc)
|
||||
for row in rows:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT INTO user_model_usage_counts"
|
||||
" (id, user_id, model, usage_count, created_at, updated_at)"
|
||||
" VALUES (:id, :user_id, :model, :cnt, :now, :now)"
|
||||
),
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"user_id": row[0],
|
||||
"model": row[1],
|
||||
"cnt": row[2],
|
||||
"now": now,
|
||||
},
|
||||
)
|
||||
|
||||
# --- 4. Enforce models.global_model_id NOT NULL ---
|
||||
conn = op.get_bind()
|
||||
insp = inspect(conn)
|
||||
model_cols = {c["name"]: c for c in insp.get_columns("models")}
|
||||
if model_cols.get("global_model_id", {}).get("nullable", True):
|
||||
op.execute("DELETE FROM models WHERE global_model_id IS NULL")
|
||||
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Revert models.global_model_id to nullable
|
||||
if column_exists("models", "global_model_id"):
|
||||
op.alter_column("models", "global_model_id", existing_type=sa.String(36), nullable=True)
|
||||
|
||||
# Drop user_model_usage_counts
|
||||
if table_exists("user_model_usage_counts"):
|
||||
if index_exists("user_model_usage_counts", "idx_user_model_usage_model"):
|
||||
op.drop_index("idx_user_model_usage_model", table_name="user_model_usage_counts")
|
||||
if index_exists("user_model_usage_counts", "idx_user_model_usage_user"):
|
||||
op.drop_index("idx_user_model_usage_user", table_name="user_model_usage_counts")
|
||||
op.drop_table("user_model_usage_counts")
|
||||
|
||||
# Drop cache_creation columns
|
||||
if column_exists("usage", "cache_creation_input_tokens_1h"):
|
||||
op.drop_column("usage", "cache_creation_input_tokens_1h")
|
||||
if column_exists("usage", "cache_creation_input_tokens_5m"):
|
||||
op.drop_column("usage", "cache_creation_input_tokens_5m")
|
||||
# capability settings cleanup is not reversible
|
||||
@@ -0,0 +1,153 @@
|
||||
"""proxy_node_metrics_and_events
|
||||
|
||||
Revision ID: 48afe197cc15
|
||||
Revises: b2c3d4e5f6a7
|
||||
Create Date: 2026-02-28 04:33:11.201185+00:00
|
||||
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "48afe197cc15"
|
||||
down_revision = "b2c3d4e5f6a7"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
columns = [c["name"] for c in insp.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def _table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
return table_name in insp.get_table_names()
|
||||
|
||||
|
||||
def _enum_has_value(enum_name: str, value: str) -> bool:
|
||||
"""检查 PostgreSQL 枚举类型是否包含指定值"""
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text(
|
||||
"SELECT 1 FROM pg_enum e JOIN pg_type t ON e.enumtypid = t.oid"
|
||||
" WHERE t.typname = :enum_name AND e.enumlabel = :value"
|
||||
),
|
||||
{"enum_name": enum_name, "value": value},
|
||||
)
|
||||
return result.fetchone() is not None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# proxy_nodes: 将已废弃的 unhealthy 状态迁移为 offline,然后从枚举中移除
|
||||
if _enum_has_value("proxynodestatus", "unhealthy"):
|
||||
op.execute("UPDATE proxy_nodes SET status = 'offline' WHERE status = 'unhealthy'")
|
||||
op.execute("ALTER TYPE proxynodestatus RENAME TO proxynodestatus_old")
|
||||
op.execute("CREATE TYPE proxynodestatus AS ENUM ('online', 'offline')")
|
||||
# 必须先移除旧枚举类型的 DEFAULT,否则 ALTER TYPE 会因无法转换默认值而报错
|
||||
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status DROP DEFAULT")
|
||||
op.execute(
|
||||
"ALTER TABLE proxy_nodes ALTER COLUMN status TYPE proxynodestatus"
|
||||
" USING status::text::proxynodestatus"
|
||||
)
|
||||
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status SET DEFAULT 'online'::proxynodestatus")
|
||||
op.execute("DROP TYPE proxynodestatus_old")
|
||||
|
||||
# proxy_nodes: 新增错误指标字段
|
||||
if not _column_exists("proxy_nodes", "failed_requests"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"failed_requests",
|
||||
sa.BigInteger(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
comment="累计失败请求数",
|
||||
),
|
||||
)
|
||||
if not _column_exists("proxy_nodes", "dns_failures"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"dns_failures",
|
||||
sa.BigInteger(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
comment="累计 DNS 失败数",
|
||||
),
|
||||
)
|
||||
if not _column_exists("proxy_nodes", "stream_errors"):
|
||||
op.add_column(
|
||||
"proxy_nodes",
|
||||
sa.Column(
|
||||
"stream_errors",
|
||||
sa.BigInteger(),
|
||||
nullable=False,
|
||||
server_default="0",
|
||||
comment="累计流错误数",
|
||||
),
|
||||
)
|
||||
|
||||
# proxy_node_events: 连接事件表
|
||||
if not _table_exists("proxy_node_events"):
|
||||
op.create_table(
|
||||
"proxy_node_events",
|
||||
sa.Column("id", sa.BigInteger(), autoincrement=True, nullable=False),
|
||||
sa.Column("node_id", sa.String(length=36), nullable=False),
|
||||
sa.Column(
|
||||
"event_type",
|
||||
sa.String(length=20),
|
||||
nullable=False,
|
||||
comment="事件类型: connected, disconnected, error",
|
||||
),
|
||||
sa.Column(
|
||||
"detail",
|
||||
sa.String(length=500),
|
||||
nullable=True,
|
||||
comment="事件详情(如断开原因)",
|
||||
),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.ForeignKeyConstraint(["node_id"], ["proxy_nodes.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index(
|
||||
"idx_proxy_node_events_node_created",
|
||||
"proxy_node_events",
|
||||
["node_id", "created_at"],
|
||||
)
|
||||
op.create_index(
|
||||
op.f("ix_proxy_node_events_node_id"),
|
||||
"proxy_node_events",
|
||||
["node_id"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# 恢复 proxynodestatus 枚举,加回 unhealthy
|
||||
if not _enum_has_value("proxynodestatus", "unhealthy"):
|
||||
op.execute("ALTER TYPE proxynodestatus RENAME TO proxynodestatus_old")
|
||||
op.execute("CREATE TYPE proxynodestatus AS ENUM ('online', 'unhealthy', 'offline')")
|
||||
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status DROP DEFAULT")
|
||||
op.execute(
|
||||
"ALTER TABLE proxy_nodes ALTER COLUMN status TYPE proxynodestatus"
|
||||
" USING status::text::proxynodestatus"
|
||||
)
|
||||
op.execute("ALTER TABLE proxy_nodes ALTER COLUMN status SET DEFAULT 'online'::proxynodestatus")
|
||||
op.execute("DROP TYPE proxynodestatus_old")
|
||||
|
||||
if _table_exists("proxy_node_events"):
|
||||
op.drop_index(op.f("ix_proxy_node_events_node_id"), table_name="proxy_node_events")
|
||||
op.drop_index("idx_proxy_node_events_node_created", table_name="proxy_node_events")
|
||||
op.drop_table("proxy_node_events")
|
||||
if _column_exists("proxy_nodes", "stream_errors"):
|
||||
op.drop_column("proxy_nodes", "stream_errors")
|
||||
if _column_exists("proxy_nodes", "dns_failures"):
|
||||
op.drop_column("proxy_nodes", "dns_failures")
|
||||
if _column_exists("proxy_nodes", "failed_requests"):
|
||||
op.drop_column("proxy_nodes", "failed_requests")
|
||||
+49
@@ -0,0 +1,49 @@
|
||||
"""add_request_candidates_composite_indexes
|
||||
|
||||
Revision ID: 00b9161b8729
|
||||
Revises: 48afe197cc15
|
||||
Create Date: 2026-02-28 14:48:00.000000+00:00
|
||||
|
||||
"""
|
||||
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "00b9161b8729"
|
||||
down_revision = "48afe197cc15"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _index_exists(index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
indexes = insp.get_indexes("request_candidates")
|
||||
return any(idx["name"] == index_name for idx in indexes)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# (request_id, status) - fallback/retry 查询优化
|
||||
if not _index_exists("idx_rc_request_id_status"):
|
||||
op.create_index(
|
||||
"idx_rc_request_id_status",
|
||||
"request_candidates",
|
||||
["request_id", "status"],
|
||||
)
|
||||
|
||||
# (provider_id, status, created_at) - provider 聚合统计优化
|
||||
if not _index_exists("idx_rc_provider_status_created"):
|
||||
op.create_index(
|
||||
"idx_rc_provider_status_created",
|
||||
"request_candidates",
|
||||
["provider_id", "status", "created_at"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if _index_exists("idx_rc_provider_status_created"):
|
||||
op.drop_index("idx_rc_provider_status_created", table_name="request_candidates")
|
||||
if _index_exists("idx_rc_request_id_status"):
|
||||
op.drop_index("idx_rc_request_id_status", table_name="request_candidates")
|
||||
@@ -0,0 +1,263 @@
|
||||
"""vertex_ai_provider_type
|
||||
|
||||
Migrate legacy Vertex auth_type/provider_type into the new model:
|
||||
- provider_type=vertex_ai
|
||||
- auth_type=service_account (legacy vertex_ai renamed)
|
||||
- fixed Vertex endpoints: gemini:chat + claude:chat
|
||||
|
||||
Revision ID: 2a624af8dd3a
|
||||
Revises: 00b9161b8729
|
||||
Create Date: 2026-02-28 15:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "2a624af8dd3a"
|
||||
down_revision = "00b9161b8729"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_VERTEX_BASE_URL = "https://aiplatform.googleapis.com"
|
||||
_VERTEX_ENDPOINTS: tuple[tuple[str, str, str], ...] = (
|
||||
("gemini:chat", "gemini", "chat"),
|
||||
("claude:chat", "claude", "chat"),
|
||||
)
|
||||
_VERTEX_KEY_FORMATS_SA = '["gemini:chat","claude:chat"]'
|
||||
_VERTEX_KEY_FORMATS_API_KEY = '["gemini:chat"]'
|
||||
|
||||
|
||||
def _select_vertex_provider_ids(conn: sa.Connection) -> list[str]:
|
||||
"""Collect providers that should be treated as Vertex after migration."""
|
||||
rows = conn.execute(sa.text("""
|
||||
SELECT DISTINCT p.id
|
||||
FROM providers p
|
||||
LEFT JOIN provider_api_keys pak ON pak.provider_id = p.id
|
||||
WHERE lower(COALESCE(p.provider_type, '')) = 'vertex_ai'
|
||||
OR pak.auth_type = 'vertex_ai'
|
||||
"""))
|
||||
return [str(row[0]) for row in rows if row[0]]
|
||||
|
||||
|
||||
def _ensure_fixed_vertex_endpoints(conn: sa.Connection, provider_ids: list[str]) -> None:
|
||||
"""Ensure every Vertex provider has fixed gemini:chat + claude:chat endpoints."""
|
||||
for provider_id in provider_ids:
|
||||
provider_max_retries = (
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
SELECT COALESCE(max_retries, 2)
|
||||
FROM providers
|
||||
WHERE id = :provider_id
|
||||
"""),
|
||||
{"provider_id": provider_id},
|
||||
).scalar()
|
||||
or 2
|
||||
)
|
||||
|
||||
for api_format, api_family, endpoint_kind in _VERTEX_ENDPOINTS:
|
||||
# Normalize existing fixed endpoint fields.
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints
|
||||
SET
|
||||
api_family = :api_family,
|
||||
endpoint_kind = :endpoint_kind,
|
||||
base_url = :base_url,
|
||||
custom_path = NULL,
|
||||
is_active = TRUE,
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE provider_id = :provider_id
|
||||
AND api_format = :api_format
|
||||
"""),
|
||||
{
|
||||
"provider_id": provider_id,
|
||||
"api_format": api_format,
|
||||
"api_family": api_family,
|
||||
"endpoint_kind": endpoint_kind,
|
||||
"base_url": _VERTEX_BASE_URL,
|
||||
},
|
||||
)
|
||||
|
||||
exists = conn.execute(
|
||||
sa.text("""
|
||||
SELECT 1
|
||||
FROM provider_endpoints
|
||||
WHERE provider_id = :provider_id
|
||||
AND api_format = :api_format
|
||||
LIMIT 1
|
||||
"""),
|
||||
{"provider_id": provider_id, "api_format": api_format},
|
||||
).first()
|
||||
|
||||
if not exists:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO provider_endpoints (
|
||||
id,
|
||||
provider_id,
|
||||
api_format,
|
||||
api_family,
|
||||
endpoint_kind,
|
||||
base_url,
|
||||
custom_path,
|
||||
header_rules,
|
||||
body_rules,
|
||||
max_retries,
|
||||
is_active,
|
||||
config,
|
||||
format_acceptance_config,
|
||||
proxy,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (
|
||||
:id,
|
||||
:provider_id,
|
||||
:api_format,
|
||||
:api_family,
|
||||
:endpoint_kind,
|
||||
:base_url,
|
||||
NULL,
|
||||
NULL,
|
||||
NULL,
|
||||
:max_retries,
|
||||
TRUE,
|
||||
NULL,
|
||||
NULL,
|
||||
NULL,
|
||||
CURRENT_TIMESTAMP,
|
||||
CURRENT_TIMESTAMP
|
||||
)
|
||||
"""),
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"provider_id": provider_id,
|
||||
"api_format": api_format,
|
||||
"api_family": api_family,
|
||||
"endpoint_kind": endpoint_kind,
|
||||
"base_url": _VERTEX_BASE_URL,
|
||||
"max_retries": int(provider_max_retries),
|
||||
},
|
||||
)
|
||||
|
||||
# Vertex fixed-provider model: disable non-fixed endpoints.
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints
|
||||
SET
|
||||
is_active = FALSE,
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE provider_id = :provider_id
|
||||
AND api_format NOT IN ('gemini:chat', 'claude:chat')
|
||||
"""),
|
||||
{"provider_id": provider_id},
|
||||
)
|
||||
|
||||
|
||||
def _normalize_vertex_key_formats(conn: sa.Connection, provider_ids: list[str]) -> None:
|
||||
"""Normalize key.api_formats for Vertex keys by auth type."""
|
||||
for provider_id in provider_ids:
|
||||
# Service Account (and legacy vertex_ai) keys: allow Gemini + Claude models.
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET
|
||||
api_formats = CAST(:api_formats AS json),
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE provider_id = :provider_id
|
||||
AND auth_type IN ('service_account', 'vertex_ai')
|
||||
"""),
|
||||
{
|
||||
"provider_id": provider_id,
|
||||
"api_formats": _VERTEX_KEY_FORMATS_SA,
|
||||
},
|
||||
)
|
||||
|
||||
# API Key mode on Vertex 仅支持 Gemini(Google publisher)。
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET
|
||||
api_formats = CAST(:api_formats AS json),
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE provider_id = :provider_id
|
||||
AND auth_type = 'api_key'
|
||||
"""),
|
||||
{
|
||||
"provider_id": provider_id,
|
||||
"api_formats": _VERTEX_KEY_FORMATS_API_KEY,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# 1) 收集目标 Provider(兼容重复执行,先识别 legacy/new 两种来源)。
|
||||
provider_ids = _select_vertex_provider_ids(conn)
|
||||
|
||||
# 2) 先重命名 auth_type(legacy vertex_ai -> service_account)。
|
||||
conn.execute(sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET auth_type = 'service_account'
|
||||
WHERE auth_type = 'vertex_ai'
|
||||
"""))
|
||||
|
||||
if not provider_ids:
|
||||
return
|
||||
|
||||
# 3) 归一 provider_type,并启用格式转换(Vertex 同时承载 Gemini/Claude)。
|
||||
for provider_id in provider_ids:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE providers
|
||||
SET
|
||||
provider_type = 'vertex_ai',
|
||||
enable_format_conversion = TRUE
|
||||
WHERE id = :provider_id
|
||||
"""),
|
||||
{"provider_id": provider_id},
|
||||
)
|
||||
|
||||
# 4) 固定端点落地:gemini:chat + claude:chat。
|
||||
_ensure_fixed_vertex_endpoints(conn, provider_ids)
|
||||
|
||||
# 5) 归一 key 的 api_formats,避免调度命中旧格式。
|
||||
_normalize_vertex_key_formats(conn, provider_ids)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
provider_rows = conn.execute(sa.text("""
|
||||
SELECT id
|
||||
FROM providers
|
||||
WHERE lower(COALESCE(provider_type, '')) = 'vertex_ai'
|
||||
"""))
|
||||
provider_ids = [str(row[0]) for row in provider_rows if row[0]]
|
||||
|
||||
if provider_ids:
|
||||
for provider_id in provider_ids:
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET auth_type = 'vertex_ai'
|
||||
WHERE provider_id = :provider_id
|
||||
AND auth_type = 'service_account'
|
||||
"""),
|
||||
{"provider_id": provider_id},
|
||||
)
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE providers
|
||||
SET provider_type = 'custom'
|
||||
WHERE id = :provider_id
|
||||
"""),
|
||||
{"provider_id": provider_id},
|
||||
)
|
||||
@@ -0,0 +1,199 @@
|
||||
"""backfill_codex_compact_endpoint
|
||||
|
||||
Backfill Codex reverse-proxy endpoints:
|
||||
- ensure `openai:cli` endpoint is pinned to force_stream
|
||||
- ensure `openai:compact` endpoint exists
|
||||
|
||||
Revision ID: f0c3a7b9d1e2
|
||||
Revises: 2a624af8dd3a
|
||||
Create Date: 2026-03-01 17:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "f0c3a7b9d1e2"
|
||||
down_revision = "2a624af8dd3a"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_CODEX_BASE_URL = "https://chatgpt.com/backend-api/codex"
|
||||
_COMPACT_FORMAT = "openai:compact"
|
||||
_CLI_FORMAT = "openai:cli"
|
||||
_FORCE_STREAM = "force_stream"
|
||||
|
||||
|
||||
def _find_codex_provider_ids(conn: sa.Connection) -> list[str]:
|
||||
"""Find Codex providers (by provider_type or legacy base_url pattern)."""
|
||||
rows = conn.execute(sa.text("""
|
||||
SELECT DISTINCT p.id
|
||||
FROM providers p
|
||||
LEFT JOIN provider_endpoints pe ON pe.provider_id = p.id
|
||||
WHERE lower(COALESCE(p.provider_type, '')) = 'codex'
|
||||
OR (
|
||||
lower(COALESCE(pe.api_format, '')) = 'openai:cli'
|
||||
AND lower(COALESCE(pe.base_url, '')) LIKE '%/backend-api/codex%'
|
||||
)
|
||||
"""))
|
||||
return [str(r[0]) for r in rows if r[0]]
|
||||
|
||||
|
||||
def _get_cli_endpoint(conn: sa.Connection, provider_id: str) -> dict[str, Any] | None:
|
||||
"""Load existing openai:cli endpoint for the provider."""
|
||||
row = (
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
SELECT base_url, header_rules, body_rules, max_retries, proxy, config
|
||||
FROM provider_endpoints
|
||||
WHERE provider_id = :pid AND api_format = :fmt
|
||||
LIMIT 1
|
||||
"""),
|
||||
{"pid": provider_id, "fmt": _CLI_FORMAT},
|
||||
)
|
||||
.mappings()
|
||||
.first()
|
||||
)
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def _pin_cli_force_stream(conn: sa.Connection, provider_id: str, cli: dict[str, Any]) -> None:
|
||||
"""Set upstream_stream_policy=force_stream on existing cli endpoint."""
|
||||
cfg = dict(cli.get("config") or {}) if isinstance(cli.get("config"), dict) else {}
|
||||
cfg.pop("upstreamStreamPolicy", None)
|
||||
cfg.pop("upstream_stream", None)
|
||||
cfg["upstream_stream_policy"] = _FORCE_STREAM
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints
|
||||
SET api_family = 'openai',
|
||||
endpoint_kind = 'cli',
|
||||
config = CAST(:config AS json),
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE provider_id = :pid AND api_format = :fmt
|
||||
"""),
|
||||
{
|
||||
"pid": provider_id,
|
||||
"fmt": _CLI_FORMAT,
|
||||
"config": json.dumps(cfg, ensure_ascii=False),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _ensure_compact_endpoint(conn: sa.Connection, provider_id: str, cli: dict[str, Any]) -> None:
|
||||
"""Create openai:compact endpoint if missing (clone from cli)."""
|
||||
exists = conn.execute(
|
||||
sa.text(
|
||||
"SELECT 1 FROM provider_endpoints WHERE provider_id = :pid AND api_format = :fmt LIMIT 1"
|
||||
),
|
||||
{"pid": provider_id, "fmt": _COMPACT_FORMAT},
|
||||
).first()
|
||||
if exists:
|
||||
# Already exists, just ensure api_family/endpoint_kind are set.
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints
|
||||
SET api_family = 'openai', endpoint_kind = 'compact',
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
WHERE provider_id = :pid AND api_format = :fmt
|
||||
"""),
|
||||
{"pid": provider_id, "fmt": _COMPACT_FORMAT},
|
||||
)
|
||||
return
|
||||
|
||||
# Clone from cli endpoint, strip stream policy.
|
||||
cfg = dict(cli.get("config") or {}) if isinstance(cli.get("config"), dict) else {}
|
||||
for k in ("upstream_stream_policy", "upstreamStreamPolicy", "upstream_stream"):
|
||||
cfg.pop(k, None)
|
||||
|
||||
def _json(val: Any) -> str | None:
|
||||
return json.dumps(val, ensure_ascii=False) if val is not None else None
|
||||
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO provider_endpoints (
|
||||
id, provider_id, api_format, api_family, endpoint_kind,
|
||||
base_url, custom_path, header_rules, body_rules,
|
||||
max_retries, is_active, config, format_acceptance_config,
|
||||
proxy, created_at, updated_at
|
||||
) VALUES (
|
||||
:id, :pid, :fmt, 'openai', 'compact',
|
||||
:base_url, NULL, CAST(:header_rules AS json), CAST(:body_rules AS json),
|
||||
:max_retries, TRUE, CAST(:config AS json), NULL,
|
||||
CAST(:proxy AS jsonb), CURRENT_TIMESTAMP, CURRENT_TIMESTAMP
|
||||
)
|
||||
"""),
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"pid": provider_id,
|
||||
"fmt": _COMPACT_FORMAT,
|
||||
"base_url": cli.get("base_url") or _CODEX_BASE_URL,
|
||||
"header_rules": _json(cli.get("header_rules")),
|
||||
"body_rules": _json(cli.get("body_rules")),
|
||||
"max_retries": cli.get("max_retries") or 2,
|
||||
"config": _json(cfg or None),
|
||||
"proxy": _json(cli.get("proxy")),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _add_compact_to_key_formats(conn: sa.Connection, provider_id: str) -> None:
|
||||
"""Ensure provider keys include openai:compact in api_formats."""
|
||||
rows = (
|
||||
conn.execute(
|
||||
sa.text("SELECT id, api_formats FROM provider_api_keys WHERE provider_id = :pid"),
|
||||
{"pid": provider_id},
|
||||
)
|
||||
.mappings()
|
||||
.all()
|
||||
)
|
||||
for row in rows:
|
||||
raw = row["api_formats"]
|
||||
formats: list[str] = []
|
||||
if isinstance(raw, list):
|
||||
for item in raw:
|
||||
v = str(item or "").strip().lower()
|
||||
if v and v not in formats:
|
||||
formats.append(v)
|
||||
|
||||
if _COMPACT_FORMAT in formats:
|
||||
continue
|
||||
|
||||
# Insert compact right after cli, or at end.
|
||||
if _CLI_FORMAT in formats:
|
||||
idx = formats.index(_CLI_FORMAT) + 1
|
||||
formats.insert(idx, _COMPACT_FORMAT)
|
||||
else:
|
||||
formats.append(_COMPACT_FORMAT)
|
||||
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_api_keys
|
||||
SET api_formats = CAST(:fmts AS json), updated_at = CURRENT_TIMESTAMP
|
||||
WHERE id = :id
|
||||
"""),
|
||||
{"id": row["id"], "fmts": json.dumps(formats, ensure_ascii=False)},
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
for provider_id in _find_codex_provider_ids(conn):
|
||||
cli = _get_cli_endpoint(conn, provider_id)
|
||||
if not cli:
|
||||
continue # No cli endpoint to clone from; skip.
|
||||
_pin_cli_force_stream(conn, provider_id, cli)
|
||||
_ensure_compact_endpoint(conn, provider_id, cli)
|
||||
_add_compact_to_key_formats(conn, provider_id)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Data backfill: no-op to avoid deleting user-managed data.
|
||||
return
|
||||
@@ -0,0 +1,29 @@
|
||||
"""add_proxy_metadata_to_proxy_nodes
|
||||
|
||||
Revision ID: 1d2e3f4a5b6c
|
||||
Revises: f0c3a7b9d1e2
|
||||
Create Date: 2026-03-02 13:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "1d2e3f4a5b6c"
|
||||
down_revision: str | None = "f0c3a7b9d1e2"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TABLE public.proxy_nodes ADD COLUMN IF NOT EXISTS proxy_metadata json")
|
||||
op.execute(
|
||||
"COMMENT ON COLUMN public.proxy_nodes.proxy_metadata IS 'aether-proxy 上报元数据(版本等)'"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("ALTER TABLE public.proxy_nodes DROP COLUMN IF EXISTS proxy_metadata")
|
||||
@@ -0,0 +1,72 @@
|
||||
"""backfill_codex_default_body_rules
|
||||
|
||||
Backfill default body_rules for codex providers with openai:cli endpoints
|
||||
that currently have body_rules IS NULL.
|
||||
|
||||
Rules:
|
||||
- drop max_output_tokens
|
||||
- drop temperature
|
||||
- drop top_p
|
||||
- set store = false
|
||||
- set instructions = "You are GPT-5." (when instructions not exists)
|
||||
|
||||
Revision ID: dd0278c0a28c
|
||||
Revises: 1d2e3f4a5b6c
|
||||
Create Date: 2026-03-02 15:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision = "dd0278c0a28c"
|
||||
down_revision = "1d2e3f4a5b6c"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_TARGET_FORMATS = ("openai:cli",)
|
||||
|
||||
_DEFAULT_BODY_RULES = [
|
||||
{"action": "drop", "path": "max_output_tokens"},
|
||||
{"action": "drop", "path": "temperature"},
|
||||
{"action": "drop", "path": "top_p"},
|
||||
{"action": "set", "path": "store", "value": False},
|
||||
{
|
||||
"action": "set",
|
||||
"path": "instructions",
|
||||
"value": "You are GPT-5.",
|
||||
"condition": {"path": "instructions", "op": "not_exists"},
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# 幂等性: 仅回填 codex 提供商中 body_rules 为空(SQL NULL 或 JSON null)的记录
|
||||
rules_json = json.dumps(_DEFAULT_BODY_RULES, ensure_ascii=False)
|
||||
result = conn.execute(
|
||||
sa.text("""
|
||||
UPDATE provider_endpoints pe
|
||||
SET body_rules = CAST(:rules AS json),
|
||||
updated_at = CURRENT_TIMESTAMP
|
||||
FROM providers p
|
||||
WHERE pe.provider_id = p.id
|
||||
AND p.provider_type = :ptype
|
||||
AND pe.api_format = :fmt
|
||||
AND (pe.body_rules IS NULL OR pe.body_rules::text = 'null')
|
||||
"""),
|
||||
{"rules": rules_json, "ptype": "codex", "fmt": _TARGET_FORMATS[0]},
|
||||
)
|
||||
if result.rowcount:
|
||||
print(f" backfilled body_rules for {result.rowcount} endpoint(s)")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# Data backfill: no-op to avoid removing user-customized rules.
|
||||
return
|
||||
@@ -0,0 +1,45 @@
|
||||
"""add_idx_usage_provider_key
|
||||
|
||||
Add composite index on usage(provider_id, provider_api_key_id) to support
|
||||
the pool management page's per-key usage stats aggregation query.
|
||||
|
||||
Revision ID: 0ba031f328de
|
||||
Revises: dd0278c0a28c
|
||||
Create Date: 2026-03-03 10:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "0ba031f328de"
|
||||
down_revision = "dd0278c0a28c"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
INDEX_NAME = "idx_usage_provider_key"
|
||||
TABLE = "usage"
|
||||
COLUMNS = ["provider_id", "provider_api_key_id"]
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
|
||||
{"name": INDEX_NAME},
|
||||
).fetchone()
|
||||
if result:
|
||||
return
|
||||
op.create_index(INDEX_NAME, TABLE, COLUMNS)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
|
||||
{"name": INDEX_NAME},
|
||||
).fetchone()
|
||||
if not result:
|
||||
return
|
||||
op.drop_index(INDEX_NAME, table_name=TABLE)
|
||||
@@ -0,0 +1,45 @@
|
||||
"""add_idx_usage_status_user_created
|
||||
|
||||
Add composite index on usage(status, user_id, created_at) to speed up
|
||||
interval timeline and active usage analytics queries.
|
||||
|
||||
Revision ID: 5f1d2e3c4b5a
|
||||
Revises: 0ba031f328de
|
||||
Create Date: 2026-03-03 17:30:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
|
||||
revision = "5f1d2e3c4b5a"
|
||||
down_revision = "0ba031f328de"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
INDEX_NAME = "idx_usage_status_user_created"
|
||||
TABLE = "usage"
|
||||
COLUMNS = ["status", "user_id", "created_at"]
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
|
||||
{"name": INDEX_NAME},
|
||||
).fetchone()
|
||||
if result:
|
||||
return
|
||||
op.create_index(INDEX_NAME, TABLE, COLUMNS)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
result = bind.execute(
|
||||
sa.text("SELECT 1 FROM pg_indexes WHERE indexname = :name"),
|
||||
{"name": INDEX_NAME},
|
||||
).fetchone()
|
||||
if not result:
|
||||
return
|
||||
op.drop_index(INDEX_NAME, table_name=TABLE)
|
||||
@@ -0,0 +1,41 @@
|
||||
"""add fingerprint column to provider_api_keys
|
||||
|
||||
Revision ID: 6a9b8c7d5e4f
|
||||
Revises: 5f1d2e3c4b5a
|
||||
Create Date: 2026-03-04 23:50:00.000000
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "6a9b8c7d5e4f"
|
||||
down_revision: str | None = "5f1d2e3c4b5a"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
inspector = inspect(bind)
|
||||
columns = [c["name"] for c in inspector.get_columns(table_name)]
|
||||
return column_name in columns
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not column_exists("provider_api_keys", "fingerprint"):
|
||||
op.add_column(
|
||||
"provider_api_keys",
|
||||
sa.Column("fingerprint", sa.JSON(), nullable=True),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if column_exists("provider_api_keys", "fingerprint"):
|
||||
op.drop_column("provider_api_keys", "fingerprint")
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,65 @@
|
||||
"""remove standalone api key locking
|
||||
|
||||
Revision ID: 7c91d2e4f8a1
|
||||
Revises: 6f7a8b9c0d1e
|
||||
Create Date: 2026-03-05 17:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "7c91d2e4f8a1"
|
||||
down_revision: str | None = "6f7a8b9c0d1e"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
_CONSTRAINT_NAME = "ck_api_keys_standalone_not_locked"
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
insp.clear_cache()
|
||||
return column_name in [c["name"] for c in insp.get_columns(table_name)]
|
||||
|
||||
|
||||
def _check_constraint_exists(table_name: str, constraint_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
insp.clear_cache()
|
||||
return any(c.get("name") == constraint_name for c in insp.get_check_constraints(table_name))
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
if not (
|
||||
_column_exists("api_keys", "is_standalone")
|
||||
and _column_exists("api_keys", "is_locked")
|
||||
and _column_exists("api_keys", "is_active")
|
||||
):
|
||||
return
|
||||
|
||||
op.execute(sa.text("""
|
||||
UPDATE api_keys
|
||||
SET is_active = FALSE,
|
||||
is_locked = FALSE
|
||||
WHERE is_standalone IS TRUE AND is_locked IS TRUE
|
||||
"""))
|
||||
|
||||
if not _check_constraint_exists("api_keys", _CONSTRAINT_NAME):
|
||||
op.create_check_constraint(
|
||||
_CONSTRAINT_NAME,
|
||||
"api_keys",
|
||||
"(NOT is_standalone) OR (NOT is_locked)",
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if _check_constraint_exists("api_keys", _CONSTRAINT_NAME):
|
||||
op.drop_constraint(_CONSTRAINT_NAME, "api_keys", type_="check")
|
||||
+220
@@ -0,0 +1,220 @@
|
||||
"""tighten wallet transaction snapshots and remove wallet version
|
||||
|
||||
Revision ID: 8e71f2a4c9b0
|
||||
Revises: 7c91d2e4f8a1
|
||||
Create Date: 2026-03-07 13:00:00.000000+00:00
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy import inspect
|
||||
|
||||
from alembic import op
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "8e71f2a4c9b0"
|
||||
down_revision: str | None = "7c91d2e4f8a1"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
_WALLET_TX_BEFORE_CHECK = "ck_wallet_tx_balance_before_consistent"
|
||||
_WALLET_TX_AFTER_CHECK = "ck_wallet_tx_balance_after_consistent"
|
||||
_WALLET_LIMIT_MODE_INDEX = "idx_wallets_limit_mode"
|
||||
|
||||
|
||||
def _table_exists(table_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
insp.clear_cache()
|
||||
return table_name in insp.get_table_names()
|
||||
|
||||
|
||||
def _column_exists(table_name: str, column_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
insp.clear_cache()
|
||||
return column_name in [c["name"] for c in insp.get_columns(table_name)]
|
||||
|
||||
|
||||
def _index_exists(table_name: str, index_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
insp.clear_cache()
|
||||
return any(index.get("name") == index_name for index in insp.get_indexes(table_name))
|
||||
|
||||
|
||||
def _check_constraint_exists(table_name: str, constraint_name: str) -> bool:
|
||||
bind = op.get_bind()
|
||||
insp = inspect(bind)
|
||||
insp.clear_cache()
|
||||
return any(c.get("name") == constraint_name for c in insp.get_check_constraints(table_name))
|
||||
|
||||
|
||||
def _tighten_wallet_transaction_snapshots() -> None:
|
||||
if not _table_exists("wallet_transactions"):
|
||||
return
|
||||
|
||||
required_columns = {
|
||||
"balance_before",
|
||||
"balance_after",
|
||||
"recharge_balance_before",
|
||||
"recharge_balance_after",
|
||||
"gift_balance_before",
|
||||
"gift_balance_after",
|
||||
}
|
||||
existing_columns = {
|
||||
column["name"] for column in inspect(op.get_bind()).get_columns("wallet_transactions")
|
||||
}
|
||||
if not required_columns.issubset(existing_columns):
|
||||
return
|
||||
|
||||
op.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE wallet_transactions
|
||||
SET recharge_balance_before = balance_before
|
||||
WHERE recharge_balance_before IS NULL
|
||||
"""
|
||||
)
|
||||
)
|
||||
op.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE wallet_transactions
|
||||
SET recharge_balance_after = balance_after
|
||||
WHERE recharge_balance_after IS NULL
|
||||
"""
|
||||
)
|
||||
)
|
||||
op.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE wallet_transactions
|
||||
SET gift_balance_before = 0
|
||||
WHERE gift_balance_before IS NULL
|
||||
"""
|
||||
)
|
||||
)
|
||||
op.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE wallet_transactions
|
||||
SET gift_balance_after = 0
|
||||
WHERE gift_balance_after IS NULL
|
||||
"""
|
||||
)
|
||||
)
|
||||
op.execute(
|
||||
sa.text(
|
||||
"""
|
||||
UPDATE wallet_transactions
|
||||
SET balance_before = recharge_balance_before + gift_balance_before,
|
||||
balance_after = recharge_balance_after + gift_balance_after
|
||||
"""
|
||||
)
|
||||
)
|
||||
|
||||
if not _check_constraint_exists("wallet_transactions", _WALLET_TX_BEFORE_CHECK):
|
||||
op.create_check_constraint(
|
||||
_WALLET_TX_BEFORE_CHECK,
|
||||
"wallet_transactions",
|
||||
"balance_before = recharge_balance_before + gift_balance_before",
|
||||
)
|
||||
if not _check_constraint_exists("wallet_transactions", _WALLET_TX_AFTER_CHECK):
|
||||
op.create_check_constraint(
|
||||
_WALLET_TX_AFTER_CHECK,
|
||||
"wallet_transactions",
|
||||
"balance_after = recharge_balance_after + gift_balance_after",
|
||||
)
|
||||
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"recharge_balance_before",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=False,
|
||||
)
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"recharge_balance_after",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=False,
|
||||
)
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"gift_balance_before",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=False,
|
||||
)
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"gift_balance_after",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
|
||||
def _drop_wallet_cleanup_artifacts() -> None:
|
||||
if not _table_exists("wallets"):
|
||||
return
|
||||
|
||||
if _index_exists("wallets", _WALLET_LIMIT_MODE_INDEX):
|
||||
op.drop_index(_WALLET_LIMIT_MODE_INDEX, table_name="wallets")
|
||||
|
||||
if _column_exists("wallets", "version"):
|
||||
op.drop_column("wallets", "version")
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
_tighten_wallet_transaction_snapshots()
|
||||
_drop_wallet_cleanup_artifacts()
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
if _table_exists("wallets"):
|
||||
if not _column_exists("wallets", "version"):
|
||||
op.add_column(
|
||||
"wallets",
|
||||
sa.Column("version", sa.Integer(), nullable=False, server_default="0"),
|
||||
)
|
||||
if not _index_exists("wallets", _WALLET_LIMIT_MODE_INDEX):
|
||||
op.create_index(_WALLET_LIMIT_MODE_INDEX, "wallets", ["limit_mode"])
|
||||
|
||||
if not _table_exists("wallet_transactions"):
|
||||
return
|
||||
|
||||
if _column_exists("wallet_transactions", "recharge_balance_before"):
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"recharge_balance_before",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=True,
|
||||
)
|
||||
if _column_exists("wallet_transactions", "recharge_balance_after"):
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"recharge_balance_after",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=True,
|
||||
)
|
||||
if _column_exists("wallet_transactions", "gift_balance_before"):
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"gift_balance_before",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=True,
|
||||
)
|
||||
if _column_exists("wallet_transactions", "gift_balance_after"):
|
||||
op.alter_column(
|
||||
"wallet_transactions",
|
||||
"gift_balance_after",
|
||||
existing_type=sa.Numeric(20, 8),
|
||||
nullable=True,
|
||||
)
|
||||
|
||||
if _check_constraint_exists("wallet_transactions", _WALLET_TX_AFTER_CHECK):
|
||||
op.drop_constraint(_WALLET_TX_AFTER_CHECK, "wallet_transactions", type_="check")
|
||||
if _check_constraint_exists("wallet_transactions", _WALLET_TX_BEFORE_CHECK):
|
||||
op.drop_constraint(_WALLET_TX_BEFORE_CHECK, "wallet_transactions", type_="check")
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user